diff --git a/.agents/skills/backend-code-review/SKILL.md b/.agents/skills/backend-code-review/SKILL.md index 35dc54173e4..66ab8dc6099 100644 --- a/.agents/skills/backend-code-review/SKILL.md +++ b/.agents/skills/backend-code-review/SKILL.md @@ -1,168 +1,40 @@ --- name: backend-code-review -description: Review backend code for quality, security, maintainability, and best practices based on established checklist rules. Use when the user requests a review, analysis, or improvement of backend files (e.g., `.py`) under the `api/` directory. Do NOT use for frontend files (e.g., `.tsx`, `.ts`, `.js`). Supports pending-change review, code snippets review, and file-focused review. +description: Use only when the user explicitly requests a review or audit of backend code under `api/`. Supports pending-change, file-focused, and pasted-diff reviews. Do not use for implementation-only requests, diagnosis without review intent, frontend code, or backend code outside `api/`. --- # Backend Code Review -## When to use this skill +Review the requested scope for concrete, reproducible defects. The nearest `AGENTS.md` owns package facts and commands; this skill owns the review workflow and routes to its bundled rule packs. -Use this skill whenever the user asks to **review, analyze, or improve** backend code (e.g., `.py`) under the `api/` directory. Supports the following review modes: +## Evidence First -- **Pending-change review**: when the user asks to review current changes (inspect staged/working-tree files slated for commit to get the changes). -- **Code snippets review**: when the user pastes code snippets (e.g., a function/class/module excerpt) into the chat and asks for a review. -- **File-focused review**: when the user points to specific files and asks for a review of those files (one file or a small, explicit set of files, e.g., `api/...`, `api/app.py`). +1. Establish the requested review scope and inspect the relevant diff or files. +2. Read the changed lines, their behavior owner, nearby tests, and local docstrings or comments that define contracts. +3. Trace callers, persistence boundaries, authorization, generated schemas, or external I/O only when they decide correctness. +4. Report only findings tied to an observable failure, violated contract, security boundary, data integrity risk, or demonstrated maintenance problem. -Do NOT use this skill when: +## Rule Routing -- The request is about frontend code or UI (e.g., `.tsx`, `.ts`, `.js`, `web/`). -- The user is not asking for a review/analysis/improvement of backend code. -- The scope is not under `api/` (unless the user explicitly asks to review backend-related changes outside `api/`). +Read only the packs matched by the diff: -## How to use this skill +- Models or migrations: [`references/db-schema-rule.md`][db-schema] +- Controller, service, core/domain, library, or model dependency direction: [`references/architecture-rule.md`][architecture] +- Table access outside an established repository boundary: [`references/repositories-rule.md`][repositories] +- SQLAlchemy sessions, queries, transactions, CRUD, concurrency, or raw SQL: [`references/sqlalchemy-rule.md`][sqlalchemy] -Follow these steps when using this skill: +When no pack applies, review correctness, security, behavior changes, and test evidence directly. Check current official documentation only when local code and contracts do not settle framework or library behavior. -1. **Identify the review mode** (pending-change vs snippet vs file-focused) based on the user’s input. Keep the scope tight: review only what the user provided or explicitly referenced. -2. Follow the rules defined in **Checklist** to perform the review. If no Checklist rule matches, apply **General Review Rules** as a fallback to perform the best-effort review. -3. Compose the final output strictly follow the **Required Output Format**. +## Severity And Output -Notes when using this skill: -- Always include actionable fixes or suggestions (including possible code snippets). -- Use best-effort `File:Line` references when a file path and line numbers are available; otherwise, use the most specific identifier you can. +- **P0**: security or privacy exposure, data loss, or a production-wide outage. +- **P1**: user-visible regression, broken authorization or tenant isolation, invalid public contract, or failed primary workflow. +- **P2**: concrete correctness, performance, maintainability, or test defect likely to cause incorrect behavior. +- **P3**: minor actionable cleanup; omit unless the user requested a thorough audit. -## Checklist +Lead with findings ordered by severity. Include a tight file and line reference, the failing contract or reproduction path, impact, and a concrete fix direction. If there are no findings, say `No issues found.` and state any material verification gap. Do not add praise sections, speculative risks, or an unsolicited offer to implement fixes. -- db schema design: if the review scope includes code/files under `api/models/` or `api/migrations/`, follow [references/db-schema-rule.md](references/db-schema-rule.md) to perform the review -- architecture: if the review scope involves controller/service/core-domain/libs/model layering, dependency direction, or moving responsibilities across modules, follow [references/architecture-rule.md](references/architecture-rule.md) to perform the review -- repositories abstraction: if the review scope contains table/model operations (e.g., `select(...)`, `session.execute(...)`, joins, CRUD) and is not under `api/repositories`, `api/core/repositories`, or `api/extensions/*/repositories/`, follow [references/repositories-rule.md](references/repositories-rule.md) to perform the review -- sqlalchemy patterns: if the review scope involves SQLAlchemy session/query usage, db transaction/crud usage, or raw SQL usage, follow [references/sqlalchemy-rule.md](references/sqlalchemy-rule.md) to perform the review - -## General Review Rules - -### 1. Security Review - -Check for: -- SQL injection vulnerabilities -- Server-Side Request Forgery (SSRF) -- Command injection -- Insecure deserialization -- Hardcoded secrets/credentials -- Improper authentication/authorization -- Insecure direct object references - -### 2. Performance Review - -Check for: -- N+1 queries -- Missing database indexes -- Memory leaks -- Blocking operations in async code -- Missing caching opportunities - -### 3. Code Quality Review - -Check for: -- Code forward compatibility -- Code duplication (DRY violations) -- Functions doing too much (SRP violations) -- Deep nesting / complex conditionals -- Magic numbers/strings -- Poor naming -- Missing error handling -- Incomplete type coverage - -### 4. Testing Review - -Check for: -- Missing test coverage for new code -- Tests that don't test behavior -- Flaky test patterns -- Missing edge cases - -## Required Output Format - -When this skill invoked, the response must exactly follow one of the two templates: - -### Template A (any findings) - -```markdown -# Code Review Summary - -Found critical issues need to be fixed: - -## 🔴 Critical (Must Fix) - -### 1. - -FilePath: line - - -#### Explanation - - - -#### Suggested Fix - -1. -2. (optional, omit if not applicable) - ---- -... (repeat for each critical issue) ... - -Found suggestions for improvement: - -## 🟡 Suggestions (Should Consider) - -### 1. - -FilePath: line - - -#### Explanation - - - -#### Suggested Fix - -1. -2. (optional, omit if not applicable) - ---- -... (repeat for each suggestion) ... - -Found optional nits: - -## 🟢 Nits (Optional) -### 1. - -FilePath: line - - -#### Explanation - - - -#### Suggested Fix - -- - ---- -... (repeat for each nits) ... - -## ✅ What's Good - -- -``` - -- If there are no critical issues or suggestions or option nits or good points, just omit that section. -- If the issue number is more than 10, summarize as "Found 10+ critical issues/suggestions/optional nits" and only output the first 10 items. -- Don't compress the blank lines between sections; keep them as-is for readability. -- If there is any issue requires code changes, append a brief follow-up question to ask whether the user wants to apply the fix(es) after the structured output. For example: "Would you like me to use the Suggested fix(es) to address these issues?" - -### Template B (no issues) - -```markdown -## Code Review Summary -✅ No issues found. -``` \ No newline at end of file +[architecture]: references/architecture-rule.md +[db-schema]: references/db-schema-rule.md +[repositories]: references/repositories-rule.md +[sqlalchemy]: references/sqlalchemy-rule.md diff --git a/.agents/skills/backend-code-review/references/architecture-rule.md b/.agents/skills/backend-code-review/references/architecture-rule.md index c3fd08bf033..02ee49d8fd9 100644 --- a/.agents/skills/backend-code-review/references/architecture-rule.md +++ b/.agents/skills/backend-code-review/references/architecture-rule.md @@ -7,7 +7,6 @@ ### Keep business logic out of controllers - Category: maintainability -- Severity: critical - Description: Controllers should parse input, call services, and return serialized responses. Business decisions inside controllers make behavior hard to reuse and test. - Suggested fix: Move domain/business logic into the service or core/domain layer. Keep controller handlers thin and orchestration-focused. - Example: @@ -34,7 +33,6 @@ ### Preserve layer dependency direction - Category: best practices -- Severity: critical - Description: Controllers may depend on services, and services may depend on core/domain abstractions. Reversing this direction (for example, core importing controller/web modules) creates cycles and leaks transport concerns into domain code. - Suggested fix: Extract shared contracts into core/domain or service-level modules and make upper layers depend on lower, not the reverse. - Example: @@ -58,7 +56,6 @@ ### Keep libs business-agnostic - Category: maintainability -- Severity: critical - Description: Modules under `api/libs/` should remain reusable, business-agnostic building blocks. They must not encode product/domain-specific rules, workflow orchestration, or business decisions. - Suggested fix: - If business logic appears in `api/libs/`, extract it into the appropriate `services/` or `core/` module and keep `libs` focused on generic, cross-cutting helpers. @@ -88,4 +85,4 @@ def should_archive_conversation(conversation, tenant_id: str) -> bool: threshold_days = 90 if has_paid_plan(tenant_id) else 30 return older_than_days(conversation.idle_days, threshold_days) - ``` \ No newline at end of file + ``` diff --git a/.agents/skills/backend-code-review/references/db-schema-rule.md b/.agents/skills/backend-code-review/references/db-schema-rule.md index 8feae2596a1..8f922bc3dc0 100644 --- a/.agents/skills/backend-code-review/references/db-schema-rule.md +++ b/.agents/skills/backend-code-review/references/db-schema-rule.md @@ -8,7 +8,6 @@ ### Do not query other tables inside `@property` - Category: [maintainability, performance] -- Severity: critical - Description: A model `@property` must not open sessions or query other tables. This hides dependencies across models, tightly couples schema objects to data access, and can cause N+1 query explosions when iterating collections. - Suggested fix: - Keep model properties pure and local to already-loaded fields. @@ -41,7 +40,6 @@ ### Prefer including `tenant_id` in model definitions - Category: maintainability -- Severity: suggestion - Description: In multi-tenant domains, include `tenant_id` in schema definitions whenever the entity belongs to tenant-owned data. This improves data isolation safety and keeps future partitioning/sharding strategies practical as data volume grows. - Suggested fix: - Add a `tenant_id` column and ensure related unique/index constraints include tenant dimension when applicable. @@ -70,7 +68,6 @@ ### Detect and avoid duplicate/redundant indexes - Category: performance -- Severity: suggestion - Description: Review index definitions for leftmost-prefix redundancy. For example, index `(a, b, c)` can safely cover most lookups for `(a, b)`. Keeping both may increase write overhead and can mislead the optimizer into suboptimal execution plans. - Suggested fix: - Before adding an index, compare against existing composite indexes by leftmost-prefix rules. @@ -94,7 +91,6 @@ ### Avoid PostgreSQL-only dialect usage in models; wrap in `models.types` - Category: maintainability -- Severity: critical - Description: Model/schema definitions should avoid PostgreSQL-only constructs directly in business models. When database-specific behavior is required, encapsulate it in `api/models/types.py` using both PostgreSQL and MySQL dialect implementations, then consume that abstraction from model code. - Suggested fix: - Do not directly place dialect-only types/operators in model columns when a portable wrapper can be used. @@ -122,7 +118,6 @@ ### Guard migration incompatibilities with dialect checks and shared types - Category: maintainability -- Severity: critical - Description: Migration scripts under `api/migrations/versions/` must account for PostgreSQL/MySQL incompatibilities explicitly. For dialect-sensitive DDL or defaults, branch on the active dialect (for example, `conn.dialect.name == "postgresql"`), and prefer reusable compatibility abstractions from `models.types` where applicable. - Suggested fix: - In migration upgrades/downgrades, bind connection and branch by dialect for incompatible SQL fragments. diff --git a/.agents/skills/backend-code-review/references/repositories-rule.md b/.agents/skills/backend-code-review/references/repositories-rule.md index 555de98eb04..c0f16a21282 100644 --- a/.agents/skills/backend-code-review/references/repositories-rule.md +++ b/.agents/skills/backend-code-review/references/repositories-rule.md @@ -8,7 +8,6 @@ ### Introduce repositories abstraction - Category: maintainability -- Severity: suggestion - Description: If a table/model already has a repository abstraction, all reads/writes/queries for that table should use the existing repository. If no repository exists, introduce one only when complexity justifies it, such as large/high-volume tables, repeated complex query logic, or likely storage-strategy variation. - Suggested fix: - First check `api/repositories`, `api/core/repositories`, and `api/extensions/*/repositories/` to verify whether the table/model already has a repository abstraction. If it exists, route all operations through it and add missing repository methods instead of bypassing it with ad-hoc SQLAlchemy access. diff --git a/.agents/skills/backend-code-review/references/sqlalchemy-rule.md b/.agents/skills/backend-code-review/references/sqlalchemy-rule.md index cda3a5dc98d..2ed3be4bbb5 100644 --- a/.agents/skills/backend-code-review/references/sqlalchemy-rule.md +++ b/.agents/skills/backend-code-review/references/sqlalchemy-rule.md @@ -8,7 +8,6 @@ ### Use Session context manager with explicit transaction control behavior - Category: best practices -- Severity: critical - Description: Session and transaction lifecycle must be explicit and bounded on write paths. Missing commits can silently drop intended updates, while ad-hoc or long-lived transactions increase contention, lock duration, and deadlock risk. - Suggested fix: - Use **explicit `session.commit()`** after completing a related write unit. @@ -47,7 +46,6 @@ ### Enforce tenant_id scoping on shared-resource queries - Category: security -- Severity: critical - Description: Reads and writes against shared tables must be scoped by `tenant_id` to prevent cross-tenant data leakage or corruption. - Suggested fix: Add `tenant_id` predicate to all tenant-owned entity queries and propagate tenant context through service/repository interfaces. - Example: @@ -67,7 +65,6 @@ ### Prefer SQLAlchemy expressions over raw SQL by default - Category: maintainability -- Severity: suggestion - Description: Raw SQL should be exceptional. ORM/Core expressions are easier to evolve, safer to compose, and more consistent with the codebase. - Suggested fix: Rewrite straightforward raw SQL into SQLAlchemy `select/update/delete` expressions; keep raw SQL only when required by clear technical constraints. - Example: @@ -89,7 +86,6 @@ ### Protect write paths with concurrency safeguards - Category: quality -- Severity: critical - Description: Multi-writer paths without explicit concurrency control can silently overwrite data. Choose the safeguard based on contention level, lock scope, and throughput cost instead of defaulting to one strategy. - Suggested fix: - **Optimistic locking**: Use when contention is usually low and retries are acceptable. Add a version (or updated_at) guard in `WHERE` and treat `rowcount == 0` as a conflict. @@ -136,4 +132,4 @@ ).scalar_one() run.status = "cancelled" session.commit() - ``` \ No newline at end of file + ``` diff --git a/.agents/skills/e2e-cucumber-playwright/SKILL.md b/.agents/skills/e2e-cucumber-playwright/SKILL.md index 5762bf2076d..62c2012dce8 100644 --- a/.agents/skills/e2e-cucumber-playwright/SKILL.md +++ b/.agents/skills/e2e-cucumber-playwright/SKILL.md @@ -1,87 +1,31 @@ --- name: e2e-cucumber-playwright -description: Write, update, or review Dify end-to-end tests under `e2e/` that use Cucumber, Gherkin, and Playwright. Use when the task involves `.feature` files, `features/step-definitions/`, `features/support/`, `DifyWorld`, scenario tags, locator/assertion choices, or E2E testing best practices for this repository. +description: Use when writing, changing, or reviewing Cucumber and Playwright tests under `e2e/`, including feature files, step definitions, support code, scenario tags, locators, and assertions. Do not use for Vitest, React Testing Library, backend tests, or generic browser automation outside the E2E suite. --- -# Dify E2E Cucumber + Playwright +# E2E Cucumber And Playwright -Use this skill for Dify's repository-level E2E suite in `e2e/`. Use [`e2e/AGENTS.md`](../../../e2e/AGENTS.md) as the canonical package guide for local architecture and conventions, then read any feature-scoped `AGENTS.md` that owns the target area. Apply Playwright/Cucumber best practices only where they fit the current suite. +`e2e/AGENTS.md` owns the suite architecture, lifecycle, commands, tags, generated-client boundaries, fixtures, and cleanup contracts. Read the nearest feature-scoped `AGENTS.md` when one exists. This skill adds no parallel package policy. -## Scope +## Topic Routing -- Use this skill for `.feature` files, Cucumber step definitions, `DifyWorld`, hooks, tags, and E2E review work under `e2e/`. -- Do not use this skill for Vitest or React Testing Library work under `web/`; use `frontend-testing` instead. -- Do not use this skill for backend test or API review tasks under `api/`. +Read only the bundled reference required by the change: -## Read Order +- Locator, assertion, isolation, or waiting decisions: [`references/playwright-best-practices.md`][playwright] +- Scenario wording, step granularity, expressions, or tag design: [`references/cucumber-best-practices.md`][cucumber] -1. Read [`e2e/AGENTS.md`](../../../e2e/AGENTS.md) first. -2. Read only the files directly involved in the task: - - target `.feature` files under `e2e/features/` - - related step files under `e2e/features/step-definitions/` - - `e2e/features/support/hooks.ts` and `e2e/features/support/world.ts` when session lifecycle or shared state matters - - `e2e/scripts/run-cucumber.ts` and `e2e/cucumber.config.ts` when tags or execution flow matter -3. Read [`references/playwright-best-practices.md`](references/playwright-best-practices.md) only when locator, assertion, isolation, or waiting choices are involved. -4. Read [`references/cucumber-best-practices.md`](references/cucumber-best-practices.md) only when scenario wording, step granularity, tags, or expression design are involved. -5. Re-check official Playwright or Cucumber docs with the available documentation tools before introducing a new framework pattern. - -Keep this skill focused on Cucumber, Playwright, and package-level E2E guidance. Put feature-specific conventions in the owning feature's `AGENTS.md` instead of adding them here. - -## Local Rules - -- `e2e/` uses Cucumber for scenarios and Playwright as the browser layer. -- `DifyWorld` is the per-scenario context object. Type `this` as `DifyWorld` and use `async function`, not arrow functions. -- Keep glue organized by capability under `e2e/features/step-definitions/`; use `common/` only for broadly reusable steps. -- Treat `e2e/AGENTS.md`, `features/support/hooks.ts`, and the Cucumber configuration as the owners of current session and tag semantics. Verify them when behavior depends on session state instead of copying a tag inventory into this skill. -- 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. +Check current official Playwright or Cucumber documentation before introducing a framework pattern that local code and references do not already establish. ## Workflow -1. Rebuild local context. - - Inspect the target feature area. - - Reuse an existing step when wording and behavior already match. - - Add a new step only for a genuinely new user action or assertion. - - Before adding several similar steps, scan the target capability for an existing domain noun that can be parameterized without hiding behavior. - - Keep edits close to the current capability folder unless the step is broadly reusable. -2. Write behavior-first scenarios. - - Describe user-observable behavior, not DOM mechanics. - - Keep each scenario focused on one workflow or outcome. - - Keep scenarios independent and re-runnable. -3. Write step definitions in the local style. - - Keep one step to one user-visible action or one assertion. - - Prefer Cucumber Expressions such as `{string}` and `{int}`. - - Use a bounded regex only when the accepted values are a small explicit domain set and Cucumber Expressions would make the Gherkin less natural. - - Do not create one-off steps for each case variant when the same domain action or outcome applies to named surfaces, modes, or resources. - - Scope locators to stable containers when the page has repeated elements. - - Avoid page-object layers or extra helper abstractions unless repeated complexity clearly justifies them. -4. Use Playwright in the local style. - - Prefer user-facing locators: `getByRole`, `getByLabel`, `getByPlaceholder`, `getByText`, then `getByTestId` for explicit contracts. - - Use web-first `expect(...)` assertions. - - Do not use `waitForTimeout`, manual polling, or raw visibility checks when a locator action or retrying assertion already expresses the behavior. - - Use `expect.poll` for API persistence, backend eventual consistency, captured browser events, or other non-DOM state; prefer locator assertions for DOM readiness and visible UI state. - - If a product element has real user-facing semantics but no accessible name, prefer fixing that accessible contract over adding a test id. -5. Validate narrowly. - - Run the narrowest tagged scenario or flow that exercises the change. - - Run the package-required static checks documented in `e2e/AGENTS.md`. - - Broaden verification only when the change affects hooks, tags, setup, or shared step semantics. +1. Add E2E coverage only for a critical user journey with a cross-boundary outcome that cheaper owner-level tests do not already prove. +2. Identify the user-visible behavior and its feature owner. Start from real product defaults and actor roles; setup may establish preconditions but must not manufacture the opposite state to make the scenario meaningful. +3. Read the target scenario, matching step definitions, and lifecycle files only when session or shared state matters. +4. Reuse an existing step when wording and behavior match; add one coherent scenario or step when they do not. +5. Keep browser actions and assertions at the public user boundary; keep setup, seed, polling, and cleanup at their package-defined owners. +6. Run the narrowest tagged scenario and package checks documented in `e2e/AGENTS.md`; broaden only for shared hooks, tags, or support changes. -## Review Checklist +For review requests, lead with reproducible correctness failures, flake sources, or demonstrated architecture drift. Report the behavior verified and any external-runtime, browser, or environment gap. -- Does the scenario describe behavior rather than implementation? -- Does it fit the current session model, tags, and `DifyWorld` usage? -- Should an existing step be reused instead of adding a new one? -- 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. - -## References - -- [`references/playwright-best-practices.md`](references/playwright-best-practices.md) -- [`references/cucumber-best-practices.md`](references/cucumber-best-practices.md) +[cucumber]: references/cucumber-best-practices.md +[playwright]: references/playwright-best-practices.md diff --git a/.agents/skills/e2e-cucumber-playwright/references/cucumber-best-practices.md b/.agents/skills/e2e-cucumber-playwright/references/cucumber-best-practices.md index 06177faa5c7..a02e2e3b0e2 100644 --- a/.agents/skills/e2e-cucumber-playwright/references/cucumber-best-practices.md +++ b/.agents/skills/e2e-cucumber-playwright/references/cucumber-best-practices.md @@ -1,12 +1,12 @@ -# Cucumber Best Practices For Dify E2E +# Cucumber Best Practices -Use this reference when writing or reviewing Gherkin scenarios, step definitions, parameter expressions, and step reuse in Dify's `e2e/` suite. +Use this reference when writing or reviewing Gherkin scenarios, step definitions, parameter expressions, and step reuse. Official sources: -- https://cucumber.io/docs/guides/10-minute-tutorial/ -- https://cucumber.io/docs/cucumber/step-definitions/ -- https://cucumber.io/docs/cucumber/cucumber-expressions/ +- https://cucumber.io/docs/guides/10-minute-tutorial +- https://cucumber.io/docs/cucumber/step-definitions +- https://cucumber.io/docs/cucumber/cucumber-expressions ## What Matters Most @@ -24,11 +24,7 @@ Apply it like this: A scenario should usually prove one workflow or business outcome. If a scenario wanders across several unrelated behaviors, split it. -In Dify's suite, this means: - -- one capability-focused scenario per feature path -- no long setup chains when existing bootstrap or reusable steps already cover them -- no hidden dependency on another scenario's side effects +Keep each scenario centered on one coherent outcome. Avoid hidden dependencies on another scenario's side effects, and keep unavoidable setup outside the behavior narrative unless the precondition matters to the specification. ### 3. Reuse steps, but only when behavior really matches @@ -68,26 +64,11 @@ Use regex for a bounded natural-language alternative only when it keeps Gherkin Step definitions are glue between Gherkin and automation, not a second abstraction language. -For Dify: - -- type `this` as `DifyWorld` -- use `async function` -- keep each step to one user-visible action or assertion -- rely on `DifyWorld` and existing support code for shared context -- avoid leaking cross-scenario state +Keep each step to one user-visible action or assertion. In JavaScript and TypeScript, use `async function` when the step reads Cucumber World state because Cucumber binds `this`; do not leak state across scenarios through module globals. ### 6. Use tags intentionally -Tags should communicate run scope or session semantics, not become ad hoc metadata. - -In Dify's current suite: - -- capability tags group related scenarios -- `@unauthenticated` changes session behavior -- `@authenticated` is descriptive/selective, not a behavior switch by itself -- `@fresh` belongs to reset/full-install flows only - -If a proposed tag implies behavior, verify that hooks or runner configuration actually implement it. +Tags should communicate selection or execution intent, not become ad hoc metadata. A tag does not change runtime behavior unless configuration or hooks implement it. ## Review Questions diff --git a/.agents/skills/e2e-cucumber-playwright/references/playwright-best-practices.md b/.agents/skills/e2e-cucumber-playwright/references/playwright-best-practices.md index deb29e72430..e6709e00d75 100644 --- a/.agents/skills/e2e-cucumber-playwright/references/playwright-best-practices.md +++ b/.agents/skills/e2e-cucumber-playwright/references/playwright-best-practices.md @@ -1,6 +1,6 @@ -# Playwright Best Practices For Dify E2E +# Playwright Best Practices -Use this reference when writing or reviewing locator, assertion, isolation, or synchronization logic for Dify's Cucumber-based E2E suite. +Use this reference when writing or reviewing locator, assertion, isolation, or synchronization logic. Official sources: @@ -13,20 +13,19 @@ Official sources: ### 1. Keep scenarios isolated -Playwright's model is built around clean browser contexts so one test does not leak into another. In Dify's suite, that principle maps to per-scenario session setup in `features/support/hooks.ts` and `DifyWorld`. +Playwright's model is built around clean browser contexts so one test does not leak into another. Apply it like this: - do not depend on another scenario having run first -- do not persist ad hoc scenario state outside `DifyWorld` -- do not couple ordinary scenarios to `@fresh` behavior -- when a flow needs special auth/session semantics, express that through the existing tag model or explicit hook changes +- keep scenario state in the runner's scenario-owned context rather than module globals +- model special authentication or session setup through explicit per-scenario fixtures rather than shared mutable state ### 2. Prefer user-facing locators Playwright recommends built-in locators that reflect what users perceive on the page. -Preferred order in this repository: +Preferred order: 1. `getByRole` 2. `getByLabel` @@ -79,16 +78,9 @@ Bad pattern: - stack arbitrary waits before every action - wait on unstable implementation details instead of the visible state the user cares about -### 5. Match debugging to the current suite +### 5. Match debugging to the active harness -Playwright's wider ecosystem supports traces and rich debugging tools. Dify's current suite already captures: - -- full-page screenshots -- page HTML -- console errors -- page errors - -Use the existing artifact flow by default. If a task is specifically about improving diagnostics, confirm the change fits the current Cucumber architecture before importing broader Playwright tooling. +Playwright supports traces, screenshots, page snapshots, and browser logs. Configure artifact capture at the runner boundary instead of adding parallel diagnostics to individual scenarios. ## Review Questions @@ -96,4 +88,4 @@ Use the existing artifact flow by default. If a task is specifically about impro - Is this assertion using Playwright's retrying semantics? - Is any explicit wait masking a real readiness problem? - Does this code preserve per-scenario isolation? -- Is a new abstraction really needed, or does it bypass the existing `DifyWorld` + step-definition model? +- Is a new abstraction really needed, or does it bypass the runner's scenario-owned context and lifecycle? diff --git a/.agents/skills/frontend-code-review/SKILL.md b/.agents/skills/frontend-code-review/SKILL.md index 85a8b1d9ef6..b5e262affc7 100644 --- a/.agents/skills/frontend-code-review/SKILL.md +++ b/.agents/skills/frontend-code-review/SKILL.md @@ -1,94 +1,48 @@ --- name: frontend-code-review -description: Review Dify frontend code for correctness, accessibility, component design, dify-ui usage, data/query boundaries, performance, and tests. Trigger for `.tsx`, `.ts`, `.js`, UI, React, Next.js, pending-change, or focused frontend review requests. +description: Use only when the user explicitly requests a review or audit of frontend code under `web/` or `packages/dify-ui/`. Supports pending-change, file-focused, and pasted-diff reviews. Do not use for implementation-only requests, diagnosis without review intent, or backend-only code. --- # Frontend Code Review -## When To Use +Review the requested scope for concrete, reproducible regressions. This skill owns the review phase and routes directly to its bundled rule packs. For a combined review-and-fix request, establish findings before applying implementation or testing guidance. -Use this skill when the user asks to review, audit, analyze, or sanity-check frontend code under `web/`, `packages/dify-ui/`, or frontend-adjacent TypeScript files. +## Evidence First -Supported modes: +1. Establish the review scope from the requested files or current diff. +2. Read the changed lines, their behavior owner, and the nearest scoped `AGENTS.md`. +3. Trace public consumers, generated contracts, primitive APIs, or runtime configuration only when they decide correctness. +4. Report only findings tied to an observable failure, violated contract, security boundary, or demonstrated maintenance risk. -- **Pending-change review**: inspect staged and working-tree changes. -- **File-focused review**: inspect explicitly named files or paths. -- **Diff/snippet review**: review pasted diffs or snippets using best-effort references. +## Rule Routing -Do not use this skill for backend-only code under `api/`; use `backend-code-review` instead. +Read only the packs matched by the diff: -## Required Context +- DOM semantics, focus, keyboard, forms, disabled state, or visible interaction: [`references/accessibility-ui.md`][accessibility] +- Dify UI imports, Base UI wrappers, overlays, tokens, or primitive contracts: [`references/dify-ui.md`][dify-ui] +- Component ownership, props, state, Effects, navigation, or module boundaries: [`references/component-architecture.md`][component-architecture] +- Generated clients, Query, mutations, auth, SSR, URL state, or persistence: [`references/data-query-contracts.md`][data-query] +- Test files or a concrete missing-regression-test finding: [`references/testing.md`][testing] +- Bundle, waterfall, rendering, or subscription cost supported by evidence: [`references/performance.md`][performance] +- Stable Dify runtime invariants in the named paths: [`references/dify-invariants.md`][dify-invariants] +- General TypeScript or styling quality not owned above: [`references/code-quality.md`][code-quality] -Before reviewing, read the relevant local contracts: +Read `packages/dify-ui/README.md`, `packages/dify-ui/AGENTS.md`, `web/docs/overlay.md`, or `web/docs/test.md` only when the reviewed code falls under that contract. Check current official documentation when local code and bundled references do not settle a framework, browser, or accessibility behavior. -- `web/AGENTS.md` for Dify frontend workflow, overlays, design tokens, state, and tests. -- `packages/dify-ui/README.md` and `packages/dify-ui/AGENTS.md` when code uses or changes `@langgenius/dify-ui/*`. -- `web/docs/overlay.md` when reviewing dialogs, drawers, popovers, tooltips, menus, selects, comboboxes, or other floating UI. -- `web/docs/test.md` and the `frontend-testing` skill when reviewing tests or testability. -- `karpathy-guidelines` for scope control and focused, verifiable changes. -- `how-to-write-component` when reviewing React component structure, ownership, effects, query/mutation contracts, or memoization. +## Severity And Output -For any UI, UX, or accessibility review, fetch the latest Web Interface Guidelines before finalizing findings. Treat them as a required baseline, not the complete source of accessibility truth: +- **P0**: security or privacy leak, data loss, production crash, or inaccessible critical workflow. +- **P1**: user-visible regression, invalid API or authorization contract, hydration failure, or broken primary interaction. +- **P2**: concrete maintainability, performance, test, or accessibility defect likely to cause incorrect behavior. +- **P3**: minor actionable cleanup; omit unless the user requested a thorough audit. -```text -https://raw.githubusercontent.com/vercel-labs/web-interface-guidelines/main/command.md -``` +Lead with findings ordered by severity. Include a tight file and line reference, the failing contract or reproduction path, impact, and a concrete fix direction. If there are no findings, say `No issues found.` and state any material verification gap. Do not add praise sections, speculative risks, or an unsolicited offer to implement fixes. -If the review depends on a current framework, SDK, browser API, or accessibility behavior and local code does not settle it, check the current official docs first. For browser compatibility, deprecation, or behavior-sensitive frontend APIs, verify MDN or the relevant standard. - -## Rule Packs - -Apply every relevant rule pack: - -- [references/accessibility-ui.md](references/accessibility-ui.md) — accessibility, semantic HTML, focus, forms, keyboard, disabled states, copy, and long-content behavior. Combines Web Interface Guidelines with Dify UI, Base UI, MDN, and local primitive contracts. -- [references/dify-ui.md](references/dify-ui.md) — Dify UI primitive usage, Base UI semantics, overlays, forms, tokens, radius mapping, and primitive boundaries. -- [references/component-architecture.md](references/component-architecture.md) — component ownership, props, state, effects, exports, wrappers, and feature organization. -- [references/data-query-contracts.md](references/data-query-contracts.md) — generated contracts, TanStack Query, mutations, workspace/auth/SSR boundaries, URL/local storage state. -- [references/performance.md](references/performance.md) — React/Next performance review rules from Vercel guidance, scoped to real risk. -- [references/testing.md](references/testing.md) — frontend test review rules. -- [references/dify-invariants.md](references/dify-invariants.md) — stable Dify-specific runtime invariants that generic React/a11y rules will not catch. -- [references/code-quality.md](references/code-quality.md) — general TypeScript, styling, naming, and maintainability rules. - -## Review Process - -1. Identify the review scope. For pending changes, inspect `git diff --stat`, `git diff`, and staged diff if relevant. For file-focused reviews, stay within the named files unless a referenced owner/contract must be read. -2. Read code around the changed lines and the owning module. Do not review by isolated snippets when nearby ownership, labels, query inputs, or overlay structure decide correctness. -3. Check user-visible regressions first: accessibility, broken interaction, auth/permission leaks, query/hydration errors, data loss, navigation mistakes, and impossible states. -4. Then check maintainability and performance: ownership, effects, wrappers, memoization, bundle/waterfall risks, tests, and design-system drift. -5. Report only actionable findings. Do not list speculative risks, style preferences, or broad refactors unless they are directly tied to a reproducible issue in scope. - -## Severity - -- **P0**: security/privacy/auth leak, data loss, production crash, inaccessible critical flow, or broken primary workflow. -- **P1**: user-visible regression, hydration/SSR failure, invalid API/query contract, broken keyboard/focus behavior, or serious design-system/a11y violation. -- **P2**: maintainability or performance issue likely to cause bugs, duplicated state, incorrect ownership, missing tests for risky behavior, or non-critical a11y issue. -- **P3**: minor cleanup with clear value. Omit unless the user asked for a thorough audit. - -## Output Format - -Lead with findings, ordered by severity. Use this structure: - -```markdown -## Findings - -- [P1] Short issue title - File: `path/to/file.tsx:123` - Why it matters and how to reproduce or reason about it. - Suggested fix: concrete fix direction. - -## Open Questions - -- Question or assumption, if any. - -## Summary - -Brief secondary context. Mention tests not run or residual risk. -``` - -Rules: - -- If there are no findings, say `No issues found.` and mention any test gaps or residual risk. -- Always include file and line when available. -- Keep findings concrete and reproducible. -- Do not include praise sections by default. -- Do not ask to apply fixes unless the user explicitly wants review plus implementation. +[accessibility]: references/accessibility-ui.md +[code-quality]: references/code-quality.md +[component-architecture]: references/component-architecture.md +[data-query]: references/data-query-contracts.md +[dify-invariants]: references/dify-invariants.md +[dify-ui]: references/dify-ui.md +[performance]: references/performance.md +[testing]: references/testing.md diff --git a/.agents/skills/frontend-code-review/references/component-architecture.md b/.agents/skills/frontend-code-review/references/component-architecture.md index 9f66533d215..ffabbf4cd69 100644 --- a/.agents/skills/frontend-code-review/references/component-architecture.md +++ b/.agents/skills/frontend-code-review/references/component-architecture.md @@ -46,14 +46,13 @@ When existing components already own interaction logic, prefer reusing or extend Flag: -- `React.FC` / `FC`. -- Default exports outside framework-required files. +- Declaration or export rewrites made only for stylistic uniformity, without changing an owned behavior or contract. - Named `Props` types for trivial one-off props where inline typing is clearer. - Props named by UI implementation instead of domain/API role. - API data converted too early or under a generic name that breaks traceability. - Callers duplicating fallback checks that the lowest rendering component already handles. -Prefer top-level `function` declarations for components and module helpers. Use arrow functions for callbacks and local lambdas. +Do not flag `FC`, `React.FC`, function declarations, arrow functions, named exports, or default exports by syntax alone. Report them only when the chosen form causes a concrete type, lifecycle, export, framework, or enforced package-contract defect. ## Effects diff --git a/.agents/skills/frontend-code-review/references/data-query-contracts.md b/.agents/skills/frontend-code-review/references/data-query-contracts.md index c1dadfafaf6..db2e2017d27 100644 --- a/.agents/skills/frontend-code-review/references/data-query-contracts.md +++ b/.agents/skills/frontend-code-review/references/data-query-contracts.md @@ -12,7 +12,7 @@ Flag: - Re-declaring API DTOs in components. - Adding compatibility layers instead of migrating the pointed line and deleting the old layer. -Use `web/contract/*` as the API shape source of truth. Follow existing `{ params, query?, body? }` input shape. +Backend Pydantic and OpenAPI schemas own API shape. Generated clients and schemas under `packages/contracts/generated/*` are authoritative at frontend boundaries and use the `{ params, query?, body? }` input shape. ## Queries diff --git a/.agents/skills/frontend-code-review/references/dify-ui.md b/.agents/skills/frontend-code-review/references/dify-ui.md index 27eced33b6a..93484a6a0fc 100644 --- a/.agents/skills/frontend-code-review/references/dify-ui.md +++ b/.agents/skills/frontend-code-review/references/dify-ui.md @@ -122,7 +122,7 @@ Flag: - Manual class strings that duplicate primitive variants. - `min-w-(--anchor-width)` on picker popups when it defeats viewport clamping. -Use the Figma radius mapping from `packages/dify-ui/AGENTS.md`; for example `--radius/sm` maps to `rounded-md`, and `--radius/md` maps to `rounded-lg`. +Use the Figma radius mapping from `packages/dify-ui/README.md`; for example `--radius/sm` maps to `rounded-md`, and `--radius/md` maps to `rounded-lg`. Use `!` only for a tightly scoped compatibility override after confirming the primitive API, data attributes, and selector structure cannot express the state. diff --git a/.agents/skills/frontend-testing/SKILL.md b/.agents/skills/frontend-testing/SKILL.md index cb71dc14772..168a3906c16 100644 --- a/.agents/skills/frontend-testing/SKILL.md +++ b/.agents/skills/frontend-testing/SKILL.md @@ -1,35 +1,16 @@ --- name: frontend-testing -description: Write, update, or review Dify frontend tests using Vitest and Testing Library. Trigger for frontend specs, test coverage requests, regressions, testability, or testing strategy under web/ or packages/dify-ui/. +description: Use when writing or changing Vitest or React Testing Library tests under `web/` or `packages/dify-ui/`, or when the user explicitly requests frontend test strategy, including evaluation of an existing strategy. Do not use for frontend code-review-only requests, general testability discussion, Python tests, or Cucumber/Playwright E2E. --- -# Dify Frontend Testing +# Frontend Testing -Use this skill for Vitest work under `web/` and `packages/dify-ui/`. Do not use it for Python tests or Cucumber/Playwright tests under `e2e/`. +`web/docs/test.md` is the single policy owner. Read it before changing frontend tests; this skill adds no separate requirements. -## Required Source +1. Identify the observable contract and regression risk. +2. Choose the smallest boundary that includes the behavior owner. +3. Establish the failing case first when practical, then implement one coherent scenario. +4. Run the focused spec before the affected suite and relevant static checks. +5. Report the behavior verified and any remaining browser, visual, or end-to-end risk. -Before writing, changing, or reviewing frontend tests, read `web/docs/test.md` completely. It is the single source of truth. This skill provides an execution checklist and must not redefine or extend that policy. - -## Workflow - -1. Read the source, its behavior owner, nearby specs, and relevant public dependencies. -1. Apply the canonical guide to decide whether a test is needed and choose its boundary. -1. For a behavior change or bug fix, write or identify the failing scenario first when practical. -1. Implement one coherent scenario at a time and run the focused spec before expanding scope. -1. Finish with the affected suite and relevant repository checks. -1. Report what behavior was verified and any risk that still requires browser, visual, or end-to-end validation. - -When reviewing existing tests, recommend deleting low-value tests as readily as adding missing behavior coverage. - -Run focused tests from the owning workspace: - -```bash -# web/ -vp test run path/to/spec-or-directory - -# packages/dify-ui/ -vp test run --project unit src/path/to/spec -``` - -Run Dify UI Storybook tests with `vp test --project storybook --run`. Run broader checks only after the focused behavior passes. +Recommend deleting low-value tests as readily as adding missing behavior coverage. Use `web/docs/test.md` for policy and Web commands; use the `packages/dify-ui/README.md` Development section for Dify UI commands. diff --git a/.agents/skills/how-to-write-component/SKILL.md b/.agents/skills/how-to-write-component/SKILL.md index 3572d03e757..f9c0a2ccc64 100644 --- a/.agents/skills/how-to-write-component/SKILL.md +++ b/.agents/skills/how-to-write-component/SKILL.md @@ -1,144 +1,41 @@ --- name: how-to-write-component -description: Use when writing, refactoring, or reviewing React/TypeScript components in Dify web, especially decisions about component ownership, props/types, URL/query state, Jotai state, async state, generated API contracts, queries/mutations, overlays, effects, navigation, performance, and empty states. +description: Use when implementing or refactoring React/TypeScript components and the task requires decisions about component ownership, feature boundaries, state, data flow, effects, or interaction ownership. Do not use for review-only requests, test-only work, copy-only edits, or styling-only changes. --- # How To Write A Component -Use this as the component decision guide for Dify web. Existing code is reference material, not automatic precedent; if touched code violates these rules, adapt it and fix equivalent patterns in the same feature branch. +Use this skill to route component architecture decisions to its bundled references. Read only the references required by the current change. ## First Decisions -| Question | Default | Promote or extract only when | +| Question | Default | Promote only when | | --- | --- | --- | -| Where should code live? | Keep it local to the feature workflow, route, or owner. | Multiple verticals need the same stable primitive. | -| How should route/tab folders be named? | Match the current route segment, tab name, or user-visible surface. | Keep a historical or broader parent only when it still owns multiple surfaces. | -| Who owns state, data, and handlers? | The lowest component that uses them. | A parent coordinates shared loading, errors, empty UI, selection, submission, navigation, or one consistent snapshot. | -| Should this become Jotai state? | Keep synchronous UI/form state in component or DOM state. | Siblings need one source of truth, the value drives atoms, or scoped workflow state must survive hidden/unmounted steps. | -| Should URL state enter Jotai? | Let Next.js route params and `nuqs` own URL state and updates. | Query atoms or shared derived atoms need a read-only bridge hydrated at the route/surface boundary. | -| Should this query/mutation become an atom? | Use TanStack Query hooks at the lowest owner. | It reads atom state, feeds derived atoms, or participates in shared Jotai workflow orchestration. | -| Should this be a helper/wrapper? | Prefer direct readable code at the use site. | The name captures a stable domain rule or the wrapper owns real behavior, validation, state, error handling, or semantics. | -| Where should a hotkey live? | Keep a single-owner hotkey constant in its component. | Multiple production files share one command, or the feature owns a real command registry with shared metadata and behavior. | -| Is an Effect needed? | No. Derive during render or handle the user action in the event handler. | It synchronizes with an external system such as browser APIs, subscriptions, timers, analytics, or imperative DOM/non-React widgets. | +| Where should code live? | In the product workflow, route, or feature owner. | Several verticals need the same stable contract. | +| Who owns state and handlers? | The lowest visual owner that consumes them. | A parent coordinates one workflow or consistent snapshot. | +| Should state enter Jotai? | Keep component and form state local. | Siblings need one source of truth or scoped workflow persistence. | +| Who owns URL state? | Next.js route APIs and `nuqs`. | Atoms require a read-only route-identity bridge. | +| Who owns remote state? | TanStack Query at the lowest consumer. | Atom state drives the query or shared derivations consume it. | +| Is a wrapper needed? | Use the primitive or direct code. | The wrapper owns behavior, validation, state, or semantics. | +| Is an Effect needed? | Derive during render or handle the user action. | A named external system must be synchronized. | -## Core Defaults +## Topic Routing -- Search before adding UI, hooks, helpers, query utilities, or styling patterns. Reuse existing base components, feature components, hooks, utilities, and design styles when they fit. -- Follow Dify's CSS-first Tailwind v4 contract from `packages/dify-ui/README.md` and `packages/dify-ui/AGENTS.md`. Prefer design-system tokens, utilities, and radius mappings over generic Tailwind choices. -- Preserve visible keyboard focus states on the final focusable element. Prefer styled `@langgenius/dify-ui/*` controls when available, because components such as `Button` and form/control primitives carry the standard Dify UI `focus-visible` styling. Do not assume every Dify UI export provides visual focus styles: headless anatomy parts and direct Base UI re-exports such as dialog/popover/tooltip/drawer triggers usually only provide behavior and semantics. When using native `button` / `a`, custom trigger `render` props, clickable rows, icon buttons, menu-like items, or direct trigger parts, verify the rendered focusable element has a visible focus state. If it does not, add the standard Dify UI focus style: `outline-hidden focus-visible:ring-2 focus-visible:ring-state-accent-solid`. Do not hide outlines without an equivalent visible `focus-visible` indicator. Component-specific focus styles should follow an existing styled primitive pattern or a concrete design constraint, not a new ad hoc style. -- Group feature code by workflow, route, or ownership area with route-aligned names: components, hooks, local types, query helpers, atoms, constants, tests, and small utilities should live near the code that changes with them. -- For each feature module, keep a module-local `README.md` as a boundary note. Start with the module name, a brief one-sentence description, then split dependencies into `Internal Modules` and `External Modules` sections; keep both sections and write `None.` when one category is empty. `Internal Modules` lists modules inside the same overall feature using paths from that feature root, such as `shared/domain/runtime-status`; `External Modules` lists project modules outside the feature using paths from the web root without a `web/` prefix, such as `app/components/base/skeleton`. Omit npm packages, workspace package dependencies, and whitelisted plumbing modules. Do not copy caller-relative import paths into the README. -- Module README whitelist: `@/service/client`, `@/next/*`. -- Keep source/default selection, validation, dirty checks, and payload shaping close to the workflow that owns submit behavior. Do not hide flow-specific priority order, fallback behavior, or submit semantics in generic utilities. -- Prefer direct conditionals for small branch-specific decisions, especially form source selection and request payload assembly. -- Loading states for page sections, cards, lists, tables, forms, and drawers should be skeletons scoped to the content being loaded. Use spinners only for small inline busy indicators. +- Component moves, module boundaries, props, types, or owner placement: read [`references/ownership.md`][ownership]. +- Jotai, form drafts, route identity, URL state, or persistence: read [`references/state.md`][state]. +- Generated contracts, nullable API data, Query, mutations, SSR, auth, or workspace state: read [`references/data.md`][data]. +- Hotkeys, focus, dialogs, menus, popovers, or other secondary surfaces: read [`references/interactions.md`][interactions] and the overlay guide it references when applicable. +- Effects, navigation, memoization, preloading, or render cost: read [`references/runtime.md`][runtime]. -## Layout And Ownership +## Workflow -- State-heavy wizards, drawers, modals, and secondary workflows can be a small feature surface: an entry file, one feature-local state file when Jotai is actually needed, and shallow `ui/` owners that match real visual regions. -- The entry file handles route integration, provider wiring, close behavior, and surface mounting. The composition owner handles high-level workflow branching. The closest visual owner handles section branching. -- When a page or tab maps to a route segment, name its feature folder after that route/tab surface instead of a stale parent grouping. Remove misleading intermediate folders when only one surface remains. -- When a tab folder grows into several independent sections or action areas, split the first level by product/visual owners. Keep the root for the entry component and cross-owner state, colocate tests with the owner folder, and put truly shared local UI under a specifically named `components/` file. -- Repeated TanStack query calls in sibling components are acceptable when each component independently consumes the data; TanStack Query deduplicates and shares cache. -- Pass stable domain identity across boundaries. Do not forward derived presentation state when the receiver can derive it from its own data source. -- A component that owns a visual surface should also own data access, loading, empty, and error states for content rendered inside it unless a parent truly coordinates that state. -- Avoid prop drilling. One pass-through layer is acceptable; repeated forwarding means ownership should move down or into feature-scoped Jotai UI state. Keep server/cache state in Query and API flow. -- Do not replace prop drilling with one large view-model hook threaded through section props. Move each hook, query, derived value, and handler to the concrete section that consumes it. -- Keep callbacks in a parent only for workflow coordination such as form submission, shared selection, batch behavior, or navigation. Otherwise let the child, menu, or row own the action. +1. Identify the behavior owner and the public contract being changed. +2. Read the nearby implementation, tests, and only the routed skill references. +3. Implement one coherent vertical slice. Do not expand into equivalent patterns elsewhere unless the current contract cannot be completed without them. +4. Verify observable behavior at the narrowest sufficient boundary, then run the checks documented by the owning package: `web/docs/test.md` or `web/docs/lint.md` for Web, and the `packages/dify-ui/README.md` Development section for Dify UI. -## Feature-Scoped Jotai - -- A Jotai-backed feature has one feature-local state file for shared primitive atoms, query atoms, derived atoms, write-only actions, mutation atoms, submission orchestration, provider exports, and optional scope configuration. -- Keep component-owned synchronous UI state local even inside Jotai features: dialog open flags, menus/popovers, confirmations, field drafts, and selected local options usually belong in component state. -- Use uncontrolled `@langgenius/dify-ui/form` and `@langgenius/dify-ui/field` controls for edit/create forms whose fields are read only at submit time. Initialize query-backed defaults with `defaultValue` and keyed remounts. -- Promote form state to atoms only when another component must react to in-progress values, a draft must survive unmount/remount in the scoped workflow, or multiple steps share the same editable draft before submit. -- Treat `useParams`, route args, and `nuqs` query state as framework-owned state. When atom logic needs those values, hydrate primitive atoms at the route or surface boundary, such as with `useHydrateAtoms(..., { dangerouslyForceHydrate: true })`; keep URL updates in the route/query-state APIs instead of write atoms. -- Within a route-owned feature, choose one source for route identity. If route params are bridged into feature atoms, use that bridge consistently for route-derived queries and actions instead of also threading the same route id through page, tab, and section props. -- For async work tied to atom state, use `atomWithQuery` or `atomWithMutation`; write atoms should update only the inputs that drive those atoms. This applies to pure frontend async work as well as network requests, so do not hand-roll loading/error/in-flight state with `useState` or `useRef` for atom-orchestrated async behavior. For component-owned remote work, use `useQuery` or `useMutation` directly. -- `jotai-tanstack-query` query atoms do not support TanStack Query tracked properties. A component that reads `useAtomValue(queryAtom)` subscribes to the whole query result, even if it only accesses `data`, `isLoading`, or `isError`. Export field-specific derived atoms and have components read the exact fields they render; use `selectAtom(queryAtom, result => result.field)` for query-result fields so unchanged selections do not notify subscribers. Keep direct `useAtomValue(queryAtom)` only when the component or hook genuinely needs the full observer result. -- Row-local async state belongs to the row owner unless it participates in a shared Jotai workflow or needs atom-scoped reset semantics. -- Leave query and mutation atoms unscoped so they keep shared QueryClient cache and invalidation behavior. Scope resettable primitives and explicit hydration tuples; scope a derived atom only when every dependency should be private to that surface. -- For scoped primitives that are always hydrated by `ScopeProvider`, prefer `atomWithLazy(() => { throw new Error(...) })` when consumers should see a non-null type. -- Order state files by dependency graph: types/constants, primitives, query atoms, query-data derived atoms, business/readiness derived atoms, write actions, mutation atoms, submission orchestration, provider exports. -- Name derived atoms as business facts and write atoms as user or workflow commands. Components should read or write the exact atom they need with `useAtomValue` or `useSetAtom`. -- Menu/dialog `open` state usually stays local, but a scoped atom is acceptable when a composed menu plus secondary surface would otherwise pass confusing `open`/`onClose` props through unrelated layers. Scope that primitive with the surface instance so reset behavior stays local. -- Keep independent dialog lifecycles separate. Avoid one discriminated "current action dialog" atom when dialogs have separate open state, loading guards, or reset behavior. - -## Components, Props, And Types - -- Type component signatures directly; do not use `FC` or `React.FC`. -- Prefer `function` for top-level components and module helpers. Use arrow functions for local callbacks, handlers, and lambda-style APIs. -- Prefer named exports. Use default exports only where the framework requires them, such as Next.js route files. -- Avoid barrel files that only re-export secondary owners. `index.tsx` is acceptable for a route/tab entry component; import header controls, switches, sections, and row owners from their concrete owner files. -- Type simple one-off props inline. Use a named `Props` type only when reused, exported, complex, or clearer. -- Use API-generated or API-returned types at component boundaries. Keep small UI conversion helpers and one-off UI extensions beside the component that needs them. -- Preserve domain value types for selection components. Do not widen enum, union, boolean, numeric, object, or nullable select/radio values to `string`; keep wrappers and option value carriers typed from their feature option collection. -- Avoid `common.tsx` buckets for shared UI. Use a feature-local `components/` folder with concrete filenames that describe the shared role. -- Do not create type aliases that only rename another type. Use aliases only for real UI concepts, refinements, or reusable local contracts. -- Name values by their domain role and backend API contract, especially persistent IDs and route params. Normalize framework or route params at the boundary. -- Put fallback and invariant checks in the lowest component that already handles that state. Do not extract helpers whose only behavior is hiding missing display data. - -## Keyboard Shortcuts - -- Distinguish application commands from local keyboard semantics before choosing an API. Use `@tanstack/react-hotkeys` for application commands. Keep menu navigation, dialog Escape handling owned by a primitive, editor commands, and other widget-scoped ARIA interactions in their local component or primitive. -- Use `useHotkey` or `useHotkeys` for registered commands. For a command intentionally owned by an existing `onKeyDown`, use `matchesKeyboardEvent` instead of hand-written `metaKey` / `ctrlKey` parsing or a second global listener. -- Define a reusable string command with `satisfies Hotkey` and an object-form command with `satisfies RawHotkey`. Reserve `RegisterableHotkey` for API boundaries that intentionally accept either form. A one-time inline literal passed directly to TanStack is already type-checked; extract it when registration, display, metadata, or another production consumer needs the same source. -- Keep registered hotkeys distinct from held keys and display-only accelerators. Use `IndividualKey` with `useKeyHold` for held-key interactions, and use an explicitly named `displayKey` for local widget accelerators that are not registered `Hotkey` values. -- Keep registration and keycap/menu display derived from one canonical command. Do not maintain a hotkey string beside a separate `['Mod', ...]` display array. -- Keep a single-owner command constant in its owning component. Create a feature-local `hotkeys.ts` only when multiple production files consume the same command. Keep a dedicated definitions/registry module when a feature owns a real command system with IDs, metadata, alternate bindings, and centralized registration. Tests do not count as another production owner, and file-name uniformity alone is not a reason to extract. -- Make scope and availability explicit. Use `enabled` for business or surface lifecycle, `ignoreInputs` for whether input-like elements may trigger the command, and `target` when the command belongs to a concrete DOM subtree. Global application commands may use the document target; inline editors and composed overlays should prefer the actual editor or Base UI Popup ref when that owner is exposed. -- Put a scoped ref on the real behavior owner. Do not add a wrapper DOM element solely to obtain a hotkey target. If a shared overlay convenience component hides the Popup ref, either rely on its modal lifecycle/focus boundary when that is sufficient or design the primitive API separately; do not create a fake owner at the call site. -- Set `preventDefault` and `stopPropagation` according to the existing product behavior and browser interaction. Do not silently accept TanStack defaults when migrating from another listener if that changes typing, submission, or propagation semantics. -- Test observable command behavior, disabled/input/target scope, and the shared registration/display contract at the owning feature boundary. Prefer partial mocks that retain TanStack formatting and matching behavior when a registration boundary must be isolated. - -## Generated API And Nullable Data - -- Treat generated contracts as authoritative at API, query, mutation, cache, and service boundaries. For enterprise APIs, use `packages/contracts/generated/enterprise/*`. -- Do not hand-write DTO mirrors, widen generated fields/enums, or add parallel frontend enum/status layers unless they model product state not represented by the API. -- Use generated enum objects and union types directly in props, comparisons, status logic, and i18n keys. Presentation-only tone maps should be keyed by generated enums. -- Normalize or coerce only at real boundaries: user-entered forms, search, URL/query params, file names, DOM IDs, or legacy adapters. -- Do not coerce nullable or optional API strings to `''` in query, derived model, or payload-building code. Keep `null` or `undefined` until the final boundary requiring a string. -- Do not use `value || undefined` for mutation fields where `''` means "clear this value". Trim or normalize at the form boundary, then preserve intentional empty strings. -- Prefer nullable-tolerant render props for API-returned rows. Narrow only where a real value is required, such as mutation params, route hrefs, select values, query input, or required React keys. -- Build required values in the same branch that proves them, using `flatMap`, a local loop, or an early return. Avoid truthiness guards, `filter(Boolean)`, `filter(item => item.id)`, and `!` after filters. -- Use conditional spreads or explicit pushes for conditional array items instead of `undefined` placeholders followed by narrowing filters. -- Empty collection fallbacks are for not-yet-loaded query data or genuinely nullable collections at the owning render boundary, not for hiding required API fields. - -## Queries And Mutations - -- Keep `web/contract/*` as the API shape source of truth and follow the `{ params, query?, body? }` input shape. -- Consume generated queries with `useQuery(consoleQuery.xxx.queryOptions(...))` or `useQuery(marketplaceQuery.xxx.queryOptions(...))`. -- If a generated query input comes from an atom, including a route-identity bridge atom, keep the query in `atomWithQuery`; do not unwrap the atom in a component just to call `useQuery`. -- Consume owner-local mutations with `useMutation(consoleQuery.xxx.mutationOptions(...))` or `useMutation(marketplaceQuery.xxx.mutationOptions(...))` when pending/error state is not consumed by feature atoms. -- In `atomWithQuery`, `atomWithInfiniteQuery`, and `atomWithMutation`, return generated `queryOptions()`, `infiniteOptions()`, or `mutationOptions()` directly. Pass `enabled`, `retry`, `placeholderData`, `select`, and pagination options into the generated call instead of spreading options into a hand-built object. -- For generated oRPC options with missing required input, branch the whole input with `input: condition ? validInput : skipToken` and `enabled: Boolean(condition)`. Never place `skipToken` inside a nested placeholder payload or coerce required IDs to `''`. -- When prefetch and render use the same request, extract local query options or a query-options atom so `prefetchQuery` and `useQuery`/`atomWithQuery` share the exact options. -- For custom query or mutation functions, wrap options with TanStack `queryOptions(...)` or `mutationOptions(...)`. -- Do not extract generated `queryOptions(...)` into a helper solely to share input construction; extract only when prefetch/render must share exact options or the helper owns real domain behavior. -- Avoid pass-through hooks and thin `web/service/use-*` wrappers that only rename generated options. Keep feature hooks for real orchestration, workflow state, or shared domain behavior. -- Put shared cache behavior in `createTanstackQueryUtils(...experimental_defaults...)`. Component or atom callbacks may handle local toasts, closing dialogs, and navigation, but should not replace shared invalidation or patch shared server state locally. -- For overlays that may open heavier secondary content, prefetch from the trigger/menu open event with `queryClient.prefetchQuery(queryOptions)` when `onOpenChange` is available. Do not mount hidden subscribers just to warm cache. -- Do not use deprecated `useInvalid` or `useReset`. -- Prefer `mutate(...)`; use `mutateAsync(...)` only when Promise semantics are required, and wrap awaited calls in `try/catch`. - -## Boundaries And Overlays - -- Use the first level below a page or tab to organize independent page sections when it adds structure or the root folder becomes noisy. This layer is layout/semantic first, not automatically the data owner. -- Treat component names, semantic roles, and user- or design-marked visual regions as boundary constraints. Keep adjacent UI as a sibling owner or introduce a correctly named broader owner. -- Keep cohesive forms, menu bodies, and one-off helpers local unless they need their own state, reuse, or semantic boundary. -- Separate hidden secondary surfaces from the trigger's main flow. For dialogs, dropdowns, popovers, and similar branches, extract a small local component when hidden content would obscure the parent. -- Preserve composability by separating behavior ownership from placement ownership: an action can own trigger/open/menu content while the caller owns slots, offsets, and alignment. -- When a dialog, dropdown, or popover accepts controlled `open`, mount it unconditionally unless unmounting is required for performance or reset semantics. Use keyed scope or local state reset instead of `{open && }` wrappers. -- When opening a dialog from a menu item, keep the menu and dialog as sibling surfaces. Let the menu command open the dialog, and mount the dialog outside menu popup content. -- For dialogs and alert dialogs, keep the root responsible for `open` wiring and put query/mutation hooks inside the content component when work should mount only after the overlay opens. -- Prefer uncontrolled overlay roots when the library can own open state. Use `onOpenChange` for side effects and CSS/data selectors for open-state styling. -- Avoid wrapper DOM unless it provides layout, semantics, accessibility, state ownership, or library integration. Avoid shallow wrappers, hook-to-props adapters, layout-only render props, children pass-through wrappers, and prop renaming unless they add real behavior or a real boundary. - -## Effects, Navigation, And Performance - -- Use Effects only to synchronize with external systems. Do not use Effects to transform props/state for rendering, handle user actions, copy state, reset state from props, or fetch data. -- For forms initialized from query data, prefer keyed remounts or surface-entry atom hydration over Effects that copy query data into form state. -- Prefer framework data APIs or TanStack Query for data fetching. -- Prefer `Link` for normal navigation. Use router APIs only for command-flow side effects such as mutation success, guarded redirects, or form submission. -- Before using `memo`, move changing state down to the smallest component that uses it. If state must wrap stable content, lift the stable content up and pass it as `children`. -- Avoid `memo`, `useMemo`, and `useCallback` unless there is a clear performance reason. +[data]: references/data.md +[interactions]: references/interactions.md +[ownership]: references/ownership.md +[runtime]: references/runtime.md +[state]: references/state.md diff --git a/.agents/skills/how-to-write-component/references/data.md b/.agents/skills/how-to-write-component/references/data.md new file mode 100644 index 00000000000..e151c2c0b5b --- /dev/null +++ b/.agents/skills/how-to-write-component/references/data.md @@ -0,0 +1,42 @@ +# Component Data And Queries + +Read this document when a component consumes generated contracts, nullable API values, TanStack Query, mutations, prefetching, authentication, or workspace state. + +## Generated Contracts + +- Treat generated contracts as authoritative at API, query, mutation, cache, and service boundaries. Enterprise APIs use `packages/contracts/generated/enterprise/*`. +- Backend Pydantic and OpenAPI schemas own API shape. Follow the generated `{ params, query?, body? }` input shape; when it is wrong, fix the backend schema and regenerate `packages/contracts/generated/*`. +- Do not hand-write DTO mirrors, widen generated fields or enums, edit generated output, or add a parallel frontend status layer unless it models product state absent from the API. +- Check deprecated markers, schema shape, and the actual consumer before assuming that a generated operation is ready to use. +- Normalize only at real boundaries such as user input, search, URL params, filenames, DOM IDs, or a required legacy adapter. +- Preserve `null`, `undefined`, and intentional empty strings until the final boundary. Do not use `value || undefined` when an empty string means clearing a field. +- Build required values in the branch that proves them. Avoid `filter(Boolean)`, truthiness filters, non-null assertions after filters, and placeholder values used only to satisfy types. + +## Queries + +- Use generated options directly with `useQuery(consoleQuery.xxx.queryOptions(...))`, `marketplaceQuery`, or the equivalent generated client. +- If query input comes from atom state, keep it in `atomWithQuery`; do not unwrap the atom in a component solely to call `useQuery`. +- For missing required input, branch the whole generated input with `skipToken`. Add `enabled` only for an independent execution condition; do not put `skipToken` inside a placeholder payload or coerce IDs to empty strings. +- Return generated `queryOptions()`, `infiniteOptions()`, or `mutationOptions()` directly from TanStack Query atoms. Pass supported options into the generated call instead of spreading into a parallel object. +- Share the exact options between prefetch and render when they represent the same request. Do not extract option helpers merely to reuse input construction. +- Avoid pass-through service hooks that only rename generated options. Keep feature hooks for actual orchestration or shared domain behavior. + +## Mutations And Cache + +- Use generated `mutationOptions()` directly for owner-local mutations. +- Put shared invalidation, retries, and cache behavior in `createTanstackQueryUtils(...experimental_defaults...)`. Local callbacks may own toast, close, and navigation effects but must not replace shared cache policy. +- Prefer `mutate(...)`. Use `mutateAsync(...)` only when Promise composition is required, and catch awaited failures. +- Preserve intentional empty values and current list/detail ownership when updating data. Do not add optimistic updates without a verified owner contract. + +## Prefetch And Hidden Surfaces + +- Prefetch expensive secondary content from the trigger or menu-open event when it benefits the visible path. Do not mount hidden subscribers solely to warm the cache. +- `prefetchQuery` is cache warmup, not an authorization or availability gate. Use a hard fetch boundary when the server must decide whether rendering may proceed. + +## SSR, Authentication, And Workspace + +- Static configuration owns path-invariant routing. Request-dependent authentication, setup, role, and tenant decisions belong to SSR or runtime decision boundaries. +- Distinguish soft SSR cache warming from authoritative decisions. Prefetched or placeholder data must not grant access or represent successful availability. +- Never reuse tenant-scoped state after switching workspaces. Discard it at the switch boundary or isolate it by workspace identity. +- Do not make product or authorization decisions from bootstrap defaults. Wait for authoritative data, or render an explicit loading or error state. +- Keep loading and Suspense behavior inside the feature that owns the request. Do not add fake global data merely to bypass that boundary. diff --git a/.agents/skills/how-to-write-component/references/interactions.md b/.agents/skills/how-to-write-component/references/interactions.md new file mode 100644 index 00000000000..a89f2b88633 --- /dev/null +++ b/.agents/skills/how-to-write-component/references/interactions.md @@ -0,0 +1,31 @@ +# Component Interactions And Overlays + +Read this document when a change involves application hotkeys, focus, dialogs, menus, popovers, or other secondary surfaces. Overlay primitive selection and layering are owned by the [overlay guide]. + +## Focus And Semantics + +- Preserve a visible focus indicator on the final focusable element. Styled Dify UI controls usually provide it; headless anatomy parts and direct trigger exports may not. +- Native buttons, links, custom trigger renderers, clickable rows, icon controls, and menu-like items must retain their correct native semantics and accessible name. +- Do not hide an outline without an equivalent visible `focus-visible` treatment. Follow an existing Dify UI pattern rather than inventing a call-site style. + +## Keyboard Commands + +- Distinguish application commands from widget-local keyboard semantics. Use `@tanstack/react-hotkeys` for application commands; keep menu navigation, dialog Escape handling, editor behavior, and ARIA widget keys in their local primitive or owner. +- Use `useHotkey` or `useHotkeys` for registered commands. When an existing `onKeyDown` intentionally owns the command, use `matchesKeyboardEvent` rather than duplicating modifier parsing or adding another global listener. +- Keep registration and keycap or menu display derived from one canonical command. Distinguish registered commands, held keys, and display-only accelerators. +- Keep a single-owner command beside its component. Create a feature-local hotkey module only when several production files share it; tests alone do not justify extraction. +- Make availability and scope explicit with `enabled`, `ignoreInputs`, and `target`. Put a target ref on the actual behavior owner rather than creating wrapper DOM solely for hotkey scope. +- Preserve existing `preventDefault` and propagation behavior when migrating command APIs. +- Test observable command behavior, disabled and input scope, target scope, and the registration/display contract at the owning feature boundary. + +## Secondary Surfaces + +- Follow `web/docs/overlay.md` for primitive choice. Dify UI primitives are the default, with package-approved Web wrappers such as `Infotip` where the overlay guide allows them. +- Separate behavior ownership from placement ownership: the action may own trigger, open state, and menu content while the caller owns slots, offsets, and alignment. +- Keep menu and dialog surfaces as siblings when a menu command opens a dialog. Mount the dialog outside popup content. +- Mount controlled overlays unconditionally unless unmounting is required for performance or reset semantics. Prefer keyed or owner-local reset over conditional wrappers. +- Put query and mutation work inside dialog or alert-dialog content when it should mount only after opening. +- Prefer uncontrolled roots when the primitive can own open state. Use controlled state only for business coordination, analytics, cleanup, or explicit reset behavior. +- Do not add manual portals or call-site z-index escalation. Fix ownership and stacking structure at the shared boundary. + +[overlay guide]: ../../../../web/docs/overlay.md diff --git a/.agents/skills/how-to-write-component/references/ownership.md b/.agents/skills/how-to-write-component/references/ownership.md new file mode 100644 index 00000000000..d8da521aafc --- /dev/null +++ b/.agents/skills/how-to-write-component/references/ownership.md @@ -0,0 +1,39 @@ +# Component Ownership And Modules + +Read this document when adding, moving, splitting, or refactoring React components or feature modules. + +## Vertical Modules + +- Organize code by product workflow, route, or behavior owner. Keep components, hooks, local types, atoms, query helpers, tests, and small utilities beside the code that changes with them. +- Name page and tab folders after the current route, tab, or user-visible surface. Do not preserve stale parent groupings that no longer own multiple surfaces. +- Split a growing page or tab by product or visual owners. Keep the feature root for its public entrypoint and genuinely cross-owner coordination. +- Import other features only through explicit public entrypoints. Avoid barrels that merely re-export secondary owners. +- Promote code outside a feature only when multiple verticals use the same stable contract. Possible future reuse is not sufficient. + +## Component Ownership + +- Put state, data access, loading, empty, error, and handlers in the lowest visual owner that uses them. +- Keep coordination in a parent only when it needs one consistent snapshot or coordinates submission, shared selection, batch behavior, navigation, or cross-section loading and errors. +- Repeated TanStack Query calls in siblings are acceptable when each sibling independently consumes the data; the cache already deduplicates requests. +- Pass stable domain identity across boundaries. Do not pass raw server data together with separately derived flags for the same concept. +- One pass-through prop layer is acceptable. Repeated forwarding means ownership should move closer to the consumer or into feature-scoped shared state. +- Do not replace prop drilling with one large view-model hook. Move each query, derived value, and handler to the concrete owner that consumes it. +- Keep source selection, defaults, validation, dirty checks, and payload shaping beside the workflow that owns submission. + +## Boundaries + +- State-heavy wizards, drawers, modals, and secondary workflows can form a small vertical surface with an entrypoint, optional feature-local state, and shallow owners matching real visual regions. +- The entrypoint owns route integration, provider wiring, close behavior, and mounting. Composition owners handle workflow branches; the closest visual owner handles section branches. +- Separate hidden dialogs, dropdowns, and popovers into small local owners when their content obscures the parent flow. +- Keep cohesive forms, menu bodies, and one-off helpers local unless they have their own state, reuse, or semantic boundary. +- Avoid wrapper components and wrapper DOM that only rename props, pass children through, or hide the real primitive. A wrapper must own behavior, validation, state, accessibility, layout, or library integration. +- Loading states for page sections, cards, lists, tables, forms, and drawers should use skeletons scoped to the loaded content. Reserve spinners for small inline busy indicators. + +## Components And Types + +- Choose component declaration and export forms from the actual component contract, framework requirements, and enforced package rules. Existing style is context, not authority; do not rewrite unaffected code solely to normalize `FC`, `function`, arrow-function, named-export, or default-export forms. +- Type simple one-off props inline. Name a `Props` type when it is reused, exported, complex, or materially clearer. +- Use API-generated or API-returned types at component boundaries. Keep one-off UI refinements and conversions beside their owner. +- Preserve domain value types for selections. Do not widen enums, unions, booleans, numbers, objects, or nullable values to `string` before a real boundary requires it. +- Avoid generic `common.tsx` buckets and aliases that only rename another type. Name files, values, and public types after their domain role. +- Put fallback and invariant checks in the lowest component that already renders that state. Do not extract helpers whose only purpose is hiding missing display data. diff --git a/.agents/skills/how-to-write-component/references/runtime.md b/.agents/skills/how-to-write-component/references/runtime.md new file mode 100644 index 00000000000..24554d895f5 --- /dev/null +++ b/.agents/skills/how-to-write-component/references/runtime.md @@ -0,0 +1,24 @@ +# Component Effects, Navigation, And Runtime Cost + +Read this document when a change introduces Effects, navigation side effects, memoization, preloading, or render-cost optimizations. + +## Effects + +- Keep render pure: do not read or write `ref.current` during render except for predictable null-guarded lazy initialization. Update interaction-owned refs in event handlers, synchronize external-system refs after commit, and use state or derivation for rendered values. +- Use Effects only to synchronize with a named external system such as a browser API, subscription, timer, analytics integration, non-React widget, or imperative DOM API. +- Do not use Effects to transform render state, handle user actions, copy query data, reset state from props, or fetch data owned by framework APIs or TanStack Query. +- Initialize query-backed forms with keyed remounts or surface-entry hydration instead of copying data through Effects. + +## Navigation + +- Use `Link` for ordinary navigation. +- Use router APIs for command-flow side effects such as mutation success, guarded redirects, or form submission. +- Keep shareable navigation state in the URL rather than hidden component state. + +## Runtime Cost + +- Move changing state to the smallest consumer before considering memoization. Stable parent content can be lifted and passed as children. +- Avoid `memo`, `useMemo`, and `useCallback` unless identity or computation has a demonstrated consumer or measurable cost. +- Start independent remote work together and await it near the branch that consumes it. Avoid introducing request waterfalls. +- Load heavy optional surfaces on demand when they sit behind a dialog, tab, command, or feature activation. +- Use narrow selectors or field-level atoms for broad stores and subscriptions. Do not optimize simple primitive expressions merely for stylistic consistency. diff --git a/.agents/skills/how-to-write-component/references/state.md b/.agents/skills/how-to-write-component/references/state.md new file mode 100644 index 00000000000..0ede0a85e92 --- /dev/null +++ b/.agents/skills/how-to-write-component/references/state.md @@ -0,0 +1,38 @@ +# Component State And URL Ownership + +Read this document when a change involves Jotai, form drafts, route identity, shared client state, or local persistence. + +## Choose The Owner + +- Keep synchronous state local when one component owns it: dialog and menu state, confirmations, field drafts, and local selections usually belong to the component or DOM. +- Use feature-scoped Jotai when siblings need one source of truth, values drive other atoms, or a scoped workflow must preserve state across hidden or unmounted steps. +- Keep server and cache state in TanStack Query. Use existing feature stores for complex, high-frequency interaction state such as workflow canvas drag, resize, and runtime panels. +- Use feature-owned storage only for low-frequency client preferences, dismissed notices, and UI defaults. Live application state does not belong in local storage. + +## Forms + +- Prefer uncontrolled Dify UI form and field controls when values are only read at submit time. Initialize query-backed defaults with `defaultValue` and keyed remounts. +- Promote form values to atoms only when another owner reacts to in-progress values, the draft must survive scoped unmounting, or several workflow steps edit the same draft. +- Keep validation, source priority, fallback behavior, dirty checks, and payload assembly in the workflow that owns submission. + +## Route And URL State + +- Treat `useParams`, route arguments, and `nuqs` as the owners of URL identity and updates. +- Hydrate a primitive atom at the route or surface boundary only when query atoms or shared derived atoms require route identity. Keep URL writes in route and query-state APIs. +- Within one route-owned feature, choose one route-identity source. Do not hydrate route identity into atoms while also threading the same ID through multiple component layers. +- Put shareable filters, tabs, pagination, and search state in the URL. Keep one-shot navigation signals and transient UI state out of persistent subscriptions. + +## Jotai And Query + +- A Jotai-backed feature may keep one feature-local state module ordered by dependency: types and constants, primitives, query atoms, query-data derivations, business facts, commands, mutations, submission orchestration, and provider exports. +- Use `atomWithQuery` or `atomWithMutation` for async work driven by atom state. Do not hand-roll loading, error, or in-flight state for atom-orchestrated work. +- Use field-specific derived atoms for query results. `jotai-tanstack-query` does not provide TanStack Query tracked properties, so reading a whole query atom subscribes to the entire observer result. +- Leave query and mutation atoms unscoped so they retain the shared QueryClient cache. Scope resettable primitives and hydration tuples; scope a derived atom only when all dependencies should be private to the surface. +- Use non-null lazy primitives for values always hydrated by a scope provider. Name derived atoms as business facts and write atoms as user or workflow commands. +- Keep independent dialog lifecycles separate. A scoped open-state atom is acceptable only when composed sibling surfaces would otherwise pass confusing lifecycle props through unrelated owners. + +## Persistence + +- Use feature-owned storage modules built on `createLocalStorageState`; callers should not scatter direct storage access or raw keys. +- Persist high-frequency interaction state only on commit or after updates settle. +- Do not add ad hoc global event listeners for shared state. Centralize subscriptions through the owning atom, store, or subscription hook. diff --git a/.agents/skills/karpathy-guidelines/SKILL.md b/.agents/skills/karpathy-guidelines/SKILL.md deleted file mode 100644 index 2b7330f5b83..00000000000 --- a/.agents/skills/karpathy-guidelines/SKILL.md +++ /dev/null @@ -1,33 +0,0 @@ ---- -name: karpathy-guidelines -description: Lightweight coding guardrails for making focused, simple, and verifiable changes in this repo. Use for all coding work. ---- - -# Karpathy Guidelines - -Use this skill whenever you touch code in this repository. - -## Principles - -- Keep the change small and directly tied to the user request. -- Prefer the simplest implementation that fits the existing codebase. -- Read the nearby code first, then match its patterns. -- Avoid unrelated refactors, broad rewrites, or style churn. -- Preserve existing behavior unless the user explicitly asked to change it. -- Treat regressions as a signal to narrow the change, not to add workaround layers. - -## Workflow - -1. Inspect the current implementation and tests around the change. -2. Make the smallest coherent edit. -3. Add or update focused tests when the behavior changes or the risk is non-trivial. -4. Run the narrowest relevant verification first. -5. Report exactly what was verified and anything left unverified. - -## Review Checklist - -- Does this change solve the stated problem without expanding scope? -- Did it preserve existing route/component/data-flow semantics? -- Are new abstractions justified by real complexity? -- Are tests focused on the behavior that could regress? -- Are unrelated files and generated artifacts left alone? diff --git a/.claude/skills/component-refactoring b/.claude/skills/component-refactoring deleted file mode 120000 index 53ae67e2f2e..00000000000 --- a/.claude/skills/component-refactoring +++ /dev/null @@ -1 +0,0 @@ -../../.agents/skills/component-refactoring \ No newline at end of file diff --git a/.claude/skills/karpathy-guidelines b/.claude/skills/karpathy-guidelines deleted file mode 120000 index 743bef5277d..00000000000 --- a/.claude/skills/karpathy-guidelines +++ /dev/null @@ -1 +0,0 @@ -../../.agents/skills/karpathy-guidelines \ No newline at end of file diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 8b6d343878d..b18d16c6191 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -245,7 +245,7 @@ # Frontend - Billing and Education /web/app/components/billing/ @iamjoel @zxhlyh -/web/app/education-apply/ @iamjoel @zxhlyh +/web/app/education/ @iamjoel @zxhlyh # Frontend - Workspace /web/app/components/header/account-dropdown/workplace-selector/ @iamjoel @zxhlyh diff --git a/.github/DISCUSSION_TEMPLATE/general.yml b/.github/DISCUSSION_TEMPLATE/general.yml index 176b207db40..7573367eb1f 100644 --- a/.github/DISCUSSION_TEMPLATE/general.yml +++ b/.github/DISCUSSION_TEMPLATE/general.yml @@ -1,17 +1,17 @@ -title: "General Discussion" +title: 'General Discussion' body: - type: checkboxes attributes: label: Self Checks - description: "To make sure we get to you in time, please check the following :)" + description: 'To make sure we get to you in time, please check the following :)' options: - label: I have searched for existing issues [search for existing issues](https://github.com/langgenius/dify/issues), including closed ones. required: true - label: I confirm that I am using English to submit this report (我已阅读并同意 [Language Policy](https://github.com/langgenius/dify/issues/1542)). required: true - - label: "[FOR CHINESE USERS] 请务必使用英文提交 Issue,否则会被关闭。谢谢!:)" + - label: '[FOR CHINESE USERS] 请务必使用英文提交 Issue,否则会被关闭。谢谢!:)' required: true - - label: "Please do not modify this template :) and fill in all the required fields." + - label: 'Please do not modify this template :) and fill in all the required fields.' required: true - type: textarea attributes: diff --git a/.github/DISCUSSION_TEMPLATE/help.yml b/.github/DISCUSSION_TEMPLATE/help.yml index 85c0debb828..7426ad197ca 100644 --- a/.github/DISCUSSION_TEMPLATE/help.yml +++ b/.github/DISCUSSION_TEMPLATE/help.yml @@ -1,17 +1,17 @@ -title: "Help" +title: 'Help' body: - type: checkboxes attributes: label: Self Checks - description: "To make sure we get to you in time, please check the following :)" + description: 'To make sure we get to you in time, please check the following :)' options: - label: I have searched for existing issues [search for existing issues](https://github.com/langgenius/dify/issues), including closed ones. required: true - label: I confirm that I am using English to submit this report (我已阅读并同意 [Language Policy](https://github.com/langgenius/dify/issues/1542)). required: true - - label: "[FOR CHINESE USERS] 请务必使用英文提交 Issue,否则会被关闭。谢谢!:)" + - label: '[FOR CHINESE USERS] 请务必使用英文提交 Issue,否则会被关闭。谢谢!:)' required: true - - label: "Please do not modify this template :) and fill in all the required fields." + - label: 'Please do not modify this template :) and fill in all the required fields.' required: true - type: textarea attributes: diff --git a/.github/DISCUSSION_TEMPLATE/suggestion.yml b/.github/DISCUSSION_TEMPLATE/suggestion.yml index 86b5b3e98b6..31a8736d8dc 100644 --- a/.github/DISCUSSION_TEMPLATE/suggestion.yml +++ b/.github/DISCUSSION_TEMPLATE/suggestion.yml @@ -3,15 +3,15 @@ body: - type: checkboxes attributes: label: Self Checks - description: "To make sure we get to you in time, please check the following :)" + description: 'To make sure we get to you in time, please check the following :)' options: - label: I have searched for existing issues [search for existing issues](https://github.com/langgenius/dify/issues), including closed ones. required: true - label: I confirm that I am using English to submit this report (我已阅读并同意 [Language Policy](https://github.com/langgenius/dify/issues/1542)). required: true - - label: "[FOR CHINESE USERS] 请务必使用英文提交 Issue,否则会被关闭。谢谢!:)" + - label: '[FOR CHINESE USERS] 请务必使用英文提交 Issue,否则会被关闭。谢谢!:)' required: true - - label: "Please do not modify this template :) and fill in all the required fields." + - label: 'Please do not modify this template :) and fill in all the required fields.' required: true - type: textarea attributes: diff --git a/.github/ISSUE_TEMPLATE/bug_report.yml b/.github/ISSUE_TEMPLATE/bug_report.yml index d684fe91443..467d75ad153 100644 --- a/.github/ISSUE_TEMPLATE/bug_report.yml +++ b/.github/ISSUE_TEMPLATE/bug_report.yml @@ -1,4 +1,4 @@ -name: "🕷️ Bug report" +name: '🕷️ Bug report' description: Report errors or unexpected behavior labels: - bug @@ -6,7 +6,7 @@ body: - type: checkboxes attributes: label: Self Checks - description: "To make sure we get to you in time, please check the following :)" + description: 'To make sure we get to you in time, please check the following :)' options: - label: I have read the [Contributing Guide](https://github.com/langgenius/dify/blob/main/CONTRIBUTING.md) and [Language Policy](https://github.com/langgenius/dify/issues/1542). required: true @@ -18,7 +18,7 @@ body: required: true - label: 【中文用户 & Non English User】请使用英语提交,否则会被关闭 :) required: true - - label: "Please do not modify this template :) and fill in all the required fields." + - label: 'Please do not modify this template :) and fill in all the required fields.' required: true - type: input diff --git a/.github/ISSUE_TEMPLATE/config.yml b/.github/ISSUE_TEMPLATE/config.yml index 859f499b8e8..c38441f09a3 100644 --- a/.github/ISSUE_TEMPLATE/config.yml +++ b/.github/ISSUE_TEMPLATE/config.yml @@ -1,13 +1,13 @@ blank_issues_enabled: false contact_links: - name: "\U0001F510 Security Vulnerabilities" - url: "https://github.com/langgenius/dify/security/advisories/new" + url: 'https://github.com/langgenius/dify/security/advisories/new' about: Report security vulnerabilities through GitHub Security Advisories to ensure responsible disclosure. 💡 Please do not report security vulnerabilities in public issues. - name: "\U0001F4A1 Model Providers & Plugins" - url: "https://github.com/langgenius/dify-official-plugins/issues/new/choose" + url: 'https://github.com/langgenius/dify-official-plugins/issues/new/choose' about: Report issues with official plugins or model providers, you will need to provide the plugin version and other relevant details. - name: "\U0001F4AC Documentation Issues" - url: "https://github.com/langgenius/dify-docs/issues/new" + url: 'https://github.com/langgenius/dify-docs/issues/new' about: Report issues with the documentation, such as typos, outdated information, or missing content. Please provide the specific section and details of the issue. - name: "\U0001F4E7 Discussions" url: https://github.com/langgenius/dify/discussions/categories/general diff --git a/.github/ISSUE_TEMPLATE/feature_request.yml b/.github/ISSUE_TEMPLATE/feature_request.yml index bd293e24429..05cad470e7c 100644 --- a/.github/ISSUE_TEMPLATE/feature_request.yml +++ b/.github/ISSUE_TEMPLATE/feature_request.yml @@ -1,4 +1,4 @@ -name: "⭐ Feature or enhancement request" +name: '⭐ Feature or enhancement request' description: Propose something new. labels: - enhancement @@ -6,7 +6,7 @@ body: - type: checkboxes attributes: label: Self Checks - description: "To make sure we get to you in time, please check the following :)" + description: 'To make sure we get to you in time, please check the following :)' options: - label: I have read the [Contributing Guide](https://github.com/langgenius/dify/blob/main/CONTRIBUTING.md) and [Language Policy](https://github.com/langgenius/dify/issues/1542). required: true @@ -14,7 +14,7 @@ body: required: true - label: I confirm that I am using English to submit this report, otherwise it will be closed. required: true - - label: "Please do not modify this template :) and fill in all the required fields." + - label: 'Please do not modify this template :) and fill in all the required fields.' required: true - type: textarea attributes: diff --git a/.github/ISSUE_TEMPLATE/refactor.yml b/.github/ISSUE_TEMPLATE/refactor.yml index dbe8cbb6027..7b70b1d6290 100644 --- a/.github/ISSUE_TEMPLATE/refactor.yml +++ b/.github/ISSUE_TEMPLATE/refactor.yml @@ -1,11 +1,11 @@ -name: "✨ Refactor or Chore" +name: '✨ Refactor or Chore' description: Refactor existing code or perform maintenance chores to improve readability and reliability. -title: "[Refactor/Chore] " +title: '[Refactor/Chore] ' body: - type: checkboxes attributes: label: Self Checks - description: "To make sure we get to you in time, please check the following :)" + description: 'To make sure we get to you in time, please check the following :)' options: - label: I have read the [Contributing Guide](https://github.com/langgenius/dify/blob/main/CONTRIBUTING.md) and [Language Policy](https://github.com/langgenius/dify/issues/1542). required: true @@ -17,26 +17,26 @@ body: required: true - label: 【中文用户 & Non English User】请使用英语提交,否则会被关闭 :) required: true - - label: "Please do not modify this template :) and fill in all the required fields." + - label: 'Please do not modify this template :) and fill in all the required fields.' required: true - type: textarea id: description attributes: label: Description - placeholder: "Describe the refactor or chore you are proposing." + placeholder: 'Describe the refactor or chore you are proposing.' validations: required: true - type: textarea id: motivation attributes: label: Motivation - placeholder: "Explain why this refactor or chore is necessary." + placeholder: 'Explain why this refactor or chore is necessary.' validations: required: false - type: textarea id: additional-context attributes: label: Additional Context - placeholder: "Add any other context or screenshots about the request here." + placeholder: 'Add any other context or screenshots about the request here.' validations: required: false diff --git a/.github/actions/setup-web/action.yml b/.github/actions/setup-web/action.yml index 3ff297dcd1e..bb2e92e719e 100644 --- a/.github/actions/setup-web/action.yml +++ b/.github/actions/setup-web/action.yml @@ -5,11 +5,11 @@ runs: using: composite steps: - name: Setup pnpm - uses: pnpm/action-setup@0e279bb959325dab635dd2c09392533439d90093 # v6.0.8 + uses: pnpm/action-setup@0977fd99725f1db4007ccb2928dbb4e90d06cc86 # v6.0.10 with: run_install: false - name: Setup Vite+ - uses: voidzero-dev/setup-vp@ca1c46663915d6c1042ae23bd39ab85718bfb0fa # v1.10.0 + uses: voidzero-dev/setup-vp@143f5f385f39b1b753ffed1a01ad443811855c8b # v1.16.1 with: node-version-file: .nvmrc cache: true diff --git a/.github/dependabot.yml b/.github/dependabot.yml index 3c22088ffe5..945cfab861e 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -1,223 +1,223 @@ version: 2 updates: - - package-ecosystem: "uv" - directory: "/api" + - package-ecosystem: 'uv' + directory: '/api' open-pull-requests-limit: 10 schedule: - interval: "weekly" + interval: 'weekly' groups: flask: patterns: - - "flask" - - "flask-*" - - "werkzeug" - - "gunicorn" + - 'flask' + - 'flask-*' + - 'werkzeug' + - 'gunicorn' google: patterns: - - "google-*" - - "googleapis-*" + - 'google-*' + - 'googleapis-*' opentelemetry: patterns: - - "opentelemetry-*" + - 'opentelemetry-*' pydantic: patterns: - - "pydantic" - - "pydantic-*" + - 'pydantic' + - 'pydantic-*' llm: patterns: - - "langfuse" - - "langsmith" - - "litellm" - - "mlflow*" - - "opik" - - "weave*" - - "arize*" - - "tiktoken" - - "transformers" + - 'langfuse' + - 'langsmith' + - 'litellm' + - 'mlflow*' + - 'opik' + - 'weave*' + - 'arize*' + - 'tiktoken' + - 'transformers' database: patterns: - - "sqlalchemy" - - "psycopg2*" - - "psycogreen" - - "redis*" - - "alembic*" + - 'sqlalchemy' + - 'psycopg2*' + - 'psycogreen' + - 'redis*' + - 'alembic*' storage: patterns: - - "boto3*" - - "botocore*" - - "azure-*" - - "bce-*" - - "cos-python-*" - - "esdk-obs-*" - - "google-cloud-storage" - - "opendal" - - "oss2" - - "supabase*" - - "tos*" + - 'boto3*' + - 'botocore*' + - 'azure-*' + - 'bce-*' + - 'cos-python-*' + - 'esdk-obs-*' + - 'google-cloud-storage' + - 'opendal' + - 'oss2' + - 'supabase*' + - 'tos*' vdb: patterns: - - "alibabacloud*" - - "chromadb" - - "clickhouse-*" - - "clickzetta-*" - - "couchbase" - - "elasticsearch" - - "opensearch-py" - - "oracledb" - - "pgvect*" - - "pymilvus" - - "pymochow" - - "pyobvector" - - "qdrant-client" - - "intersystems-*" - - "tablestore" - - "tcvectordb" - - "tidb-vector" - - "upstash-*" - - "volcengine-*" - - "weaviate-*" - - "xinference-*" - - "mo-vector" - - "mysql-connector-*" + - 'alibabacloud*' + - 'chromadb' + - 'clickhouse-*' + - 'clickzetta-*' + - 'couchbase' + - 'elasticsearch' + - 'opensearch-py' + - 'oracledb' + - 'pgvect*' + - 'pymilvus' + - 'pymochow' + - 'pyobvector' + - 'qdrant-client' + - 'intersystems-*' + - 'tablestore' + - 'tcvectordb' + - 'tidb-vector' + - 'upstash-*' + - 'volcengine-*' + - 'weaviate-*' + - 'xinference-*' + - 'mo-vector' + - 'mysql-connector-*' dev: patterns: - - "coverage" - - "dotenv-linter" - - "faker" - - "lxml-stubs" - - "basedpyright" - - "ruff" - - "pytest*" - - "types-*" - - "boto3-stubs" - - "hypothesis" - - "pandas-stubs" - - "scipy-stubs" - - "import-linter" - - "celery-types" - - "mypy*" - - "pyrefly" + - 'coverage' + - 'dotenv-linter' + - 'faker' + - 'lxml-stubs' + - 'basedpyright' + - 'ruff' + - 'pytest*' + - 'types-*' + - 'boto3-stubs' + - 'hypothesis' + - 'pandas-stubs' + - 'scipy-stubs' + - 'import-linter' + - 'celery-types' + - 'mypy*' + - 'pyrefly' python-packages: patterns: - - "*" - - package-ecosystem: "github-actions" - directory: "/" + - '*' + - package-ecosystem: 'github-actions' + directory: '/' open-pull-requests-limit: 5 schedule: - interval: "weekly" + interval: 'weekly' groups: github-actions-dependencies: patterns: - - "*" - - package-ecosystem: "uv" - directory: "/api" - target-branch: "lts/1.13.x" + - '*' + - package-ecosystem: 'uv' + directory: '/api' + target-branch: 'lts/1.13.x' open-pull-requests-limit: 10 schedule: - interval: "weekly" + interval: 'weekly' groups: flask: patterns: - - "flask" - - "flask-*" - - "werkzeug" - - "gunicorn" + - 'flask' + - 'flask-*' + - 'werkzeug' + - 'gunicorn' google: patterns: - - "google-*" - - "googleapis-*" + - 'google-*' + - 'googleapis-*' opentelemetry: patterns: - - "opentelemetry-*" + - 'opentelemetry-*' pydantic: patterns: - - "pydantic" - - "pydantic-*" + - 'pydantic' + - 'pydantic-*' llm: patterns: - - "langfuse" - - "langsmith" - - "litellm" - - "mlflow*" - - "opik" - - "weave*" - - "arize*" - - "tiktoken" - - "transformers" + - 'langfuse' + - 'langsmith' + - 'litellm' + - 'mlflow*' + - 'opik' + - 'weave*' + - 'arize*' + - 'tiktoken' + - 'transformers' database: patterns: - - "sqlalchemy" - - "psycopg2*" - - "psycogreen" - - "redis*" - - "alembic*" + - 'sqlalchemy' + - 'psycopg2*' + - 'psycogreen' + - 'redis*' + - 'alembic*' storage: patterns: - - "boto3*" - - "botocore*" - - "azure-*" - - "bce-*" - - "cos-python-*" - - "esdk-obs-*" - - "google-cloud-storage" - - "opendal" - - "oss2" - - "supabase*" - - "tos*" + - 'boto3*' + - 'botocore*' + - 'azure-*' + - 'bce-*' + - 'cos-python-*' + - 'esdk-obs-*' + - 'google-cloud-storage' + - 'opendal' + - 'oss2' + - 'supabase*' + - 'tos*' vdb: patterns: - - "alibabacloud*" - - "chromadb" - - "clickhouse-*" - - "clickzetta-*" - - "couchbase" - - "elasticsearch" - - "opensearch-py" - - "oracledb" - - "pgvect*" - - "pymilvus" - - "pymochow" - - "pyobvector" - - "qdrant-client" - - "intersystems-*" - - "tablestore" - - "tcvectordb" - - "tidb-vector" - - "upstash-*" - - "volcengine-*" - - "weaviate-*" - - "xinference-*" - - "mo-vector" - - "mysql-connector-*" + - 'alibabacloud*' + - 'chromadb' + - 'clickhouse-*' + - 'clickzetta-*' + - 'couchbase' + - 'elasticsearch' + - 'opensearch-py' + - 'oracledb' + - 'pgvect*' + - 'pymilvus' + - 'pymochow' + - 'pyobvector' + - 'qdrant-client' + - 'intersystems-*' + - 'tablestore' + - 'tcvectordb' + - 'tidb-vector' + - 'upstash-*' + - 'volcengine-*' + - 'weaviate-*' + - 'xinference-*' + - 'mo-vector' + - 'mysql-connector-*' dev: patterns: - - "coverage" - - "dotenv-linter" - - "faker" - - "lxml-stubs" - - "basedpyright" - - "ruff" - - "pytest*" - - "types-*" - - "boto3-stubs" - - "hypothesis" - - "pandas-stubs" - - "scipy-stubs" - - "import-linter" - - "celery-types" - - "mypy*" - - "pyrefly" + - 'coverage' + - 'dotenv-linter' + - 'faker' + - 'lxml-stubs' + - 'basedpyright' + - 'ruff' + - 'pytest*' + - 'types-*' + - 'boto3-stubs' + - 'hypothesis' + - 'pandas-stubs' + - 'scipy-stubs' + - 'import-linter' + - 'celery-types' + - 'mypy*' + - 'pyrefly' python-packages: patterns: - - "*" - - package-ecosystem: "github-actions" - directory: "/" - target-branch: "lts/1.13.x" + - '*' + - package-ecosystem: 'github-actions' + directory: '/' + target-branch: 'lts/1.13.x' open-pull-requests-limit: 5 schedule: - interval: "weekly" + interval: 'weekly' groups: github-actions-dependencies: patterns: - - "*" + - '*' diff --git a/.github/linters/.hadolint.yaml b/.github/linters/.hadolint.yaml index a607204b136..c17c41d5cea 100644 --- a/.github/linters/.hadolint.yaml +++ b/.github/linters/.hadolint.yaml @@ -1 +1 @@ -failure-threshold: "error" +failure-threshold: 'error' diff --git a/.github/linters/.yaml-lint.yml b/.github/linters/.yaml-lint.yml index c886e67b6ad..1fb6ff1a09b 100644 --- a/.github/linters/.yaml-lint.yml +++ b/.github/linters/.yaml-lint.yml @@ -1,5 +1,4 @@ --- - extends: default rules: diff --git a/.github/linters/editorconfig-checker.json b/.github/linters/editorconfig-checker.json index ce6e9ae3411..e8377c6bf56 100644 --- a/.github/linters/editorconfig-checker.json +++ b/.github/linters/editorconfig-checker.json @@ -1,22 +1,22 @@ { - "Verbose": false, - "Debug": false, - "IgnoreDefaults": false, - "SpacesAfterTabs": false, - "NoColor": false, - "Exclude": [ - "^web/public/vs/", - "^web/public/pdf.worker.min.mjs$", - "web/app/components/base/icons/src/vender/" - ], - "AllowedContentTypes": [], - "PassedFiles": [], - "Disable": { - "EndOfLine": false, - "Indentation": false, - "IndentSize": true, - "InsertFinalNewline": false, - "TrimTrailingWhitespace": false, - "MaxLineLength": false - } + "Verbose": false, + "Debug": false, + "IgnoreDefaults": false, + "SpacesAfterTabs": false, + "NoColor": false, + "Exclude": [ + "^web/public/vs/", + "^web/public/pdf.worker.min.mjs$", + "web/app/components/base/icons/src/vender/" + ], + "AllowedContentTypes": [], + "PassedFiles": [], + "Disable": { + "EndOfLine": false, + "Indentation": false, + "IndentSize": true, + "InsertFinalNewline": false, + "TrimTrailingWhitespace": false, + "MaxLineLength": false + } } diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 1e848612ec5..eecbb919659 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -12,8 +12,8 @@ ## Screenshots | Before | After | -|--------|-------| -| ... | ... | +| ------ | ----- | +| ... | ... | ## Checklist diff --git a/.github/scripts/generate-i18n-changes.mjs b/.github/scripts/generate-i18n-changes.mjs index 3d25115ac34..9fb3f979db9 100644 --- a/.github/scripts/generate-i18n-changes.mjs +++ b/.github/scripts/generate-i18n-changes.mjs @@ -8,31 +8,31 @@ const headSha = process.env.HEAD_SHA || '' const files = (process.env.CHANGED_FILES || '').split(/\s+/).filter(Boolean) const outputPath = process.env.I18N_CHANGES_OUTPUT_PATH || '/tmp/i18n-changes.json' -const englishPath = fileStem => path.join(repoRoot, 'web', 'i18n', 'en-US', `${fileStem}.json`) +const englishPath = (fileStem) => path.join(repoRoot, 'web', 'i18n', 'en-US', `${fileStem}.json`) const readCurrentJson = (fileStem) => { const filePath = englishPath(fileStem) - if (!fs.existsSync(filePath)) - return null + if (!fs.existsSync(filePath)) return null return JSON.parse(fs.readFileSync(filePath, 'utf8')) } const readBaseJson = (fileStem) => { - if (!baseSha) - return null + if (!baseSha) return null try { const relativePath = `web/i18n/en-US/${fileStem}.json` - const content = execFileSync('git', ['show', `${baseSha}:${relativePath}`], { encoding: 'utf8' }) + const content = execFileSync('git', ['show', `${baseSha}:${relativePath}`], { + encoding: 'utf8', + }) return JSON.parse(content) - } - catch { + } catch { return null } } -const compareJson = (beforeValue, afterValue) => JSON.stringify(beforeValue) === JSON.stringify(afterValue) +const compareJson = (beforeValue, afterValue) => + JSON.stringify(beforeValue) === JSON.stringify(afterValue) const changes = {} @@ -59,8 +59,7 @@ for (const fileStem of files) { } for (const key of Object.keys(beforeJson)) { - if (!(key in afterJson)) - deleted.push(key) + if (!(key in afterJson)) deleted.push(key) } changes[fileStem] = { @@ -78,5 +77,5 @@ fs.writeFileSync( headSha, files, changes, - }) + }), ) diff --git a/.github/workflows/accessibility-e2e.yml b/.github/workflows/accessibility-e2e.yml new file mode 100644 index 00000000000..03a41f587ea --- /dev/null +++ b/.github/workflows/accessibility-e2e.yml @@ -0,0 +1,101 @@ +name: Manual WCAG Accessibility Audit +run-name: WCAG ${{ inputs.level }} · ${{ inputs.page }} + +on: + workflow_dispatch: + inputs: + page: + description: Page to scan + required: true + default: home + type: choice + options: + - home + - sign-in + - studio + - agents + - knowledge + - integrations + - all + level: + description: WCAG level to scan + required: true + default: aa + type: choice + options: + - aa + - a + +permissions: + contents: read + +concurrency: + group: accessibility-e2e-${{ github.ref }}-${{ inputs.page }}-${{ inputs.level }} + cancel-in-progress: true + +jobs: + accessibility: + name: WCAG Level ${{ inputs.level == 'a' && 'A' || 'AA' }} · ${{ inputs.page }} + runs-on: depot-ubuntu-24.04-4 + timeout-minutes: 120 + defaults: + run: + shell: bash + + steps: + - name: Checkout code + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + with: + persist-credentials: false + + - name: Setup web dependencies + uses: ./.github/actions/setup-web + + - name: Setup UV and Python + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 + with: + enable-cache: true + python-version: '3.12' + cache-dependency-glob: | + api/uv.lock + + - name: Install API dependencies + run: uv sync --project api --dev + + - name: Install Chromium for accessibility E2E + timeout-minutes: 15 + working-directory: ./e2e + run: vp run e2e:install:ci:chromium + + - name: Run automated WCAG Level ${{ inputs.level == 'a' && 'A' || 'AA' }} checks + working-directory: ./e2e + env: + E2E_ADMIN_EMAIL: e2e-admin@example.com + E2E_ADMIN_NAME: E2E Admin + E2E_ADMIN_PASSWORD: E2eAdmin12345 + E2E_FORCE_WEB_BUILD: '1' + E2E_INIT_PASSWORD: E2eInit12345 + WCAG_PAGE: ${{ inputs.page }} + run: | + if [[ "$WCAG_PAGE" == "all" ]]; then + vp run e2e:accessibility:${{ inputs.level }} + else + pnpm exec tsx ./scripts/run-cucumber.ts --full -- --tags "@axe and @wcag-${{ inputs.level }} and @wcag-page-$WCAG_PAGE" + fi + + - name: Upload accessibility Cucumber report + if: ${{ !cancelled() }} + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 + with: + name: cucumber-report-accessibility-${{ inputs.level }}-${{ inputs.page }} + path: e2e/cucumber-report + retention-days: 7 + + - name: Upload accessibility E2E logs + if: ${{ !cancelled() }} + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 + with: + name: e2e-logs-accessibility-${{ inputs.level }}-${{ inputs.page }} + path: e2e/.logs/*.log + include-hidden-files: true + retention-days: 7 diff --git a/.github/workflows/api-tests.yml b/.github/workflows/api-tests.yml index e8bcd20cf66..54009a040e3 100644 --- a/.github/workflows/api-tests.yml +++ b/.github/workflows/api-tests.yml @@ -25,17 +25,17 @@ jobs: strategy: matrix: python-version: - - "3.12" + - '3.12' steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false - name: Setup UV and Python - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true python-version: ${{ matrix.python-version }} @@ -84,17 +84,17 @@ jobs: strategy: matrix: python-version: - - "3.12" + - '3.12' steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false - name: Setup UV and Python - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true python-version: ${{ matrix.python-version }} @@ -139,16 +139,16 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false - name: Setup UV and Python - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true - python-version: "3.12" + python-version: '3.12' cache-dependency-glob: api/uv.lock - name: Install dependencies diff --git a/.github/workflows/autofix.yml b/.github/workflows/autofix.yml index d05ebed87d5..6b3f297f482 100644 --- a/.github/workflows/autofix.yml +++ b/.github/workflows/autofix.yml @@ -1,12 +1,12 @@ name: autofix.ci on: pull_request: - branches: ["main"] + branches: ['main'] merge_group: - branches: ["main"] + branches: ['main'] types: [checks_requested] push: - branches: ["main"] + branches: ['main'] permissions: contents: read @@ -20,7 +20,7 @@ jobs: run: echo "autofix.ci updates pull request branches, not merge group refs." - if: github.event_name != 'merge_group' - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Check Docker Compose inputs if: github.event_name != 'merge_group' @@ -53,8 +53,7 @@ jobs: oxlint-suppressions.json eslint-suppressions.json .vscode/** - .github/workflows/autofix.yml - .github/workflows/style.yml + .github/** - name: Check api inputs if: github.event_name != 'merge_group' id: api-changes @@ -84,12 +83,12 @@ jobs: dify-agent/pyproject.toml dify-agent/uv.lock - if: github.event_name != 'merge_group' - uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0 + uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 with: - python-version: "3.11" + python-version: '3.11' - if: github.event_name != 'merge_group' - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 - name: Generate Docker Compose if: github.event_name != 'merge_group' && steps.docker-compose-changes.outputs.any_changed == 'true' diff --git a/.github/workflows/build-push.yml b/.github/workflows/build-push.yml index 545495ca712..61d654bd375 100644 --- a/.github/workflows/build-push.yml +++ b/.github/workflows/build-push.yml @@ -3,13 +3,13 @@ name: Build and Push API & Web on: push: branches: - - "main" - - "deploy/**" - - "build/**" - - "release/e-*" - - "hotfix/**" + - 'main' + - 'deploy/**' + - 'build/**' + - 'release/e-*' + - 'hotfix/**' tags: - - "*" + - '*' concurrency: group: build-push-${{ github.head_ref || github.run_id }} @@ -33,60 +33,60 @@ jobs: strategy: matrix: include: - - service_name: "build-api-amd64" - image_name_env: "DIFY_API_IMAGE_NAME" - artifact_context: "api" - build_context: "{{defaultContext}}" - file: "api/Dockerfile" + - service_name: 'build-api-amd64' + image_name_env: 'DIFY_API_IMAGE_NAME' + artifact_context: 'api' + build_context: '{{defaultContext}}' + file: 'api/Dockerfile' platform: linux/amd64 runs_on: depot-ubuntu-24.04-4 - - service_name: "build-api-arm64" - image_name_env: "DIFY_API_IMAGE_NAME" - artifact_context: "api" - build_context: "{{defaultContext}}" - file: "api/Dockerfile" + - service_name: 'build-api-arm64' + image_name_env: 'DIFY_API_IMAGE_NAME' + artifact_context: 'api' + build_context: '{{defaultContext}}' + file: 'api/Dockerfile' platform: linux/arm64 runs_on: depot-ubuntu-24.04-4 - - service_name: "build-web-amd64" - image_name_env: "DIFY_WEB_IMAGE_NAME" - artifact_context: "web" - build_context: "{{defaultContext}}" - file: "web/Dockerfile" + - service_name: 'build-web-amd64' + image_name_env: 'DIFY_WEB_IMAGE_NAME' + artifact_context: 'web' + build_context: '{{defaultContext}}' + file: 'web/Dockerfile' platform: linux/amd64 runs_on: depot-ubuntu-24.04-4 - - service_name: "build-web-arm64" - image_name_env: "DIFY_WEB_IMAGE_NAME" - artifact_context: "web" - build_context: "{{defaultContext}}" - file: "web/Dockerfile" + - service_name: 'build-web-arm64' + image_name_env: 'DIFY_WEB_IMAGE_NAME' + artifact_context: 'web' + build_context: '{{defaultContext}}' + file: 'web/Dockerfile' platform: linux/arm64 runs_on: depot-ubuntu-24.04-4 - - service_name: "build-agent-amd64" - image_name_env: "DIFY_AGENT_IMAGE_NAME" - artifact_context: "agent" - build_context: "{{defaultContext}}" - file: "dify-agent/Dockerfile" + - service_name: 'build-agent-amd64' + image_name_env: 'DIFY_AGENT_IMAGE_NAME' + artifact_context: 'agent' + build_context: '{{defaultContext}}' + file: 'dify-agent/Dockerfile' platform: linux/amd64 runs_on: depot-ubuntu-24.04-4 - - service_name: "build-agent-arm64" - image_name_env: "DIFY_AGENT_IMAGE_NAME" - artifact_context: "agent" - build_context: "{{defaultContext}}" - file: "dify-agent/Dockerfile" + - service_name: 'build-agent-arm64' + image_name_env: 'DIFY_AGENT_IMAGE_NAME' + artifact_context: 'agent' + build_context: '{{defaultContext}}' + file: 'dify-agent/Dockerfile' platform: linux/arm64 runs_on: depot-ubuntu-24.04-4 - - service_name: "build-agent-local-sandbox-amd64" - image_name_env: "DIFY_AGENT_LOCAL_SANDBOX_IMAGE_NAME" - artifact_context: "local-sandbox" - build_context: "{{defaultContext}}:dify-agent-runtime" - file: "docker/Dockerfile" + - service_name: 'build-agent-local-sandbox-amd64' + image_name_env: 'DIFY_AGENT_LOCAL_SANDBOX_IMAGE_NAME' + artifact_context: 'local-sandbox' + build_context: '{{defaultContext}}:dify-agent-runtime' + file: 'docker/Dockerfile' platform: linux/amd64 runs_on: depot-ubuntu-24.04-4 - - service_name: "build-agent-local-sandbox-arm64" - image_name_env: "DIFY_AGENT_LOCAL_SANDBOX_IMAGE_NAME" - artifact_context: "local-sandbox" - build_context: "{{defaultContext}}:dify-agent-runtime" - file: "docker/Dockerfile" + - service_name: 'build-agent-local-sandbox-arm64' + image_name_env: 'DIFY_AGENT_LOCAL_SANDBOX_IMAGE_NAME' + artifact_context: 'local-sandbox' + build_context: '{{defaultContext}}:dify-agent-runtime' + file: 'docker/Dockerfile' platform: linux/arm64 runs_on: depot-ubuntu-24.04-4 @@ -97,7 +97,7 @@ jobs: echo "PLATFORM_PAIR=${platform//\//-}" >> $GITHUB_ENV - name: Login to Docker Hub - uses: docker/login-action@af1e73f918a031802d376d3c8bbc3fe56130a9b0 # v4.4.0 + uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0 with: username: ${{ env.DOCKERHUB_USER }} password: ${{ env.DOCKERHUB_TOKEN }} @@ -146,18 +146,18 @@ jobs: strategy: matrix: include: - - service_name: "validate-api-amd64" - build_context: "{{defaultContext}}" - file: "api/Dockerfile" - - service_name: "validate-web-amd64" - build_context: "{{defaultContext}}" - file: "web/Dockerfile" - - service_name: "validate-agent-amd64" - build_context: "{{defaultContext}}" - file: "dify-agent/Dockerfile" - - service_name: "validate-agent-local-sandbox-amd64" - build_context: "{{defaultContext}}:dify-agent-runtime" - file: "docker/Dockerfile" + - service_name: 'validate-api-amd64' + build_context: '{{defaultContext}}' + file: 'api/Dockerfile' + - service_name: 'validate-web-amd64' + build_context: '{{defaultContext}}' + file: 'web/Dockerfile' + - service_name: 'validate-agent-amd64' + build_context: '{{defaultContext}}' + file: 'dify-agent/Dockerfile' + - service_name: 'validate-agent-local-sandbox-amd64' + build_context: '{{defaultContext}}:dify-agent-runtime' + file: 'docker/Dockerfile' steps: - name: Set up Docker Buildx uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0 @@ -178,18 +178,18 @@ jobs: strategy: matrix: include: - - service_name: "merge-api-images" - image_name_env: "DIFY_API_IMAGE_NAME" - context: "api" - - service_name: "merge-web-images" - image_name_env: "DIFY_WEB_IMAGE_NAME" - context: "web" - - service_name: "merge-agent-images" - image_name_env: "DIFY_AGENT_IMAGE_NAME" - context: "agent" - - service_name: "merge-agent-local-sandbox-images" - image_name_env: "DIFY_AGENT_LOCAL_SANDBOX_IMAGE_NAME" - context: "local-sandbox" + - service_name: 'merge-api-images' + image_name_env: 'DIFY_API_IMAGE_NAME' + context: 'api' + - service_name: 'merge-web-images' + image_name_env: 'DIFY_WEB_IMAGE_NAME' + context: 'web' + - service_name: 'merge-agent-images' + image_name_env: 'DIFY_AGENT_IMAGE_NAME' + context: 'agent' + - service_name: 'merge-agent-local-sandbox-images' + image_name_env: 'DIFY_AGENT_LOCAL_SANDBOX_IMAGE_NAME' + context: 'local-sandbox' steps: - name: Download digests uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8.0.1 @@ -199,7 +199,7 @@ jobs: merge-multiple: true - name: Login to Docker Hub - uses: docker/login-action@af1e73f918a031802d376d3c8bbc3fe56130a9b0 # v4.4.0 + uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0 with: username: ${{ env.DOCKERHUB_USER }} password: ${{ env.DOCKERHUB_TOKEN }} diff --git a/.github/workflows/cli-e2e.yml b/.github/workflows/cli-e2e.yml index b99853972dd..0c23cda5f58 100644 --- a/.github/workflows/cli-e2e.yml +++ b/.github/workflows/cli-e2e.yml @@ -4,19 +4,19 @@ on: workflow_dispatch: inputs: cli_ref: - description: "Git ref (default: current branch)" + description: 'Git ref (default: current branch)' type: string required: false edition: - description: "Dify edition" + description: 'Dify edition' type: choice required: false default: ee options: [ee, ce] test_scope: - description: "smoke = [P0] only / full = all cases" + description: 'smoke = [P0] only / full = all cases' type: choice required: false default: full @@ -24,23 +24,23 @@ on: # ── Suite on/off ──────────────────────────────────────────────────────── suite_framework_output_error: - description: "framework + output + error-handling suites" + description: 'framework + output + error-handling suites' type: boolean default: true suite_discovery: - description: "discovery suite (get app / describe app)" + description: 'discovery suite (get app / describe app)' type: boolean default: true suite_run: - description: "run suite (basic / streaming / conversation / file / hitl)" + description: 'run suite (basic / streaming / conversation / file / hitl)' type: boolean default: true suite_auth: - description: "auth suite (login / status / whoami / use / devices / logout)" + description: 'auth suite (login / status / whoami / use / devices / logout)' type: boolean default: true suite_agent: - description: "agent suite" + description: 'agent suite' type: boolean default: true @@ -51,35 +51,34 @@ permissions: # Each job reads DIFY_E2E_TOKEN + app IDs from the provision job outputs, # so global-setup skips minting and finds existing apps in < 10 s. env: - DIFY_E2E_NO_KEYRING: "1" # Linux CI has no keychain; skip probe - VITEST_RETRY: "2" # Retry flaky staging responses + DIFY_E2E_NO_KEYRING: '1' # Linux CI has no keychain; skip probe + VITEST_RETRY: '2' # Retry flaky staging responses jobs: - -# ════════════════════════════════════════════════════════════════════════════ -# 0. PROVISION — mint token + import DSL fixtures (runs once, outputs IDs) -# ════════════════════════════════════════════════════════════════════════════ + # ════════════════════════════════════════════════════════════════════════════ + # 0. PROVISION — mint token + import DSL fixtures (runs once, outputs IDs) + # ════════════════════════════════════════════════════════════════════════════ provision: - name: "Provision: mint token + DSL apps" + name: 'Provision: mint token + DSL apps' runs-on: ubuntu-latest timeout-minutes: 10 outputs: - token: ${{ steps.out.outputs.DIFY_E2E_TOKEN }} - workspace_id: ${{ steps.out.outputs.DIFY_E2E_WORKSPACE_ID }} - workspace_name: ${{ steps.out.outputs.DIFY_E2E_WORKSPACE_NAME }} - ws2_id: ${{ steps.out.outputs.DIFY_E2E_WS2_ID }} - chat_app_id: ${{ steps.out.outputs.DIFY_E2E_CHAT_APP_ID }} - workflow_app_id: ${{ steps.out.outputs.DIFY_E2E_WORKFLOW_APP_ID }} - file_app_id: ${{ steps.out.outputs.DIFY_E2E_FILE_APP_ID }} - file_chat_app_id: ${{ steps.out.outputs.DIFY_E2E_FILE_CHAT_APP_ID }} - hitl_app_id: ${{ steps.out.outputs.DIFY_E2E_HITL_APP_ID }} - hitl_external_app_id: ${{ steps.out.outputs.DIFY_E2E_HITL_EXTERNAL_APP_ID }} - hitl_single_action_app_id: ${{ steps.out.outputs.DIFY_E2E_HITL_SINGLE_ACTION_APP_ID }} - hitl_multi_node_app_id: ${{ steps.out.outputs.DIFY_E2E_HITL_MULTI_NODE_APP_ID }} - ws2_app_id: ${{ steps.out.outputs.DIFY_E2E_WS2_APP_ID }} + token: ${{ steps.out.outputs.DIFY_E2E_TOKEN }} + workspace_id: ${{ steps.out.outputs.DIFY_E2E_WORKSPACE_ID }} + workspace_name: ${{ steps.out.outputs.DIFY_E2E_WORKSPACE_NAME }} + ws2_id: ${{ steps.out.outputs.DIFY_E2E_WS2_ID }} + chat_app_id: ${{ steps.out.outputs.DIFY_E2E_CHAT_APP_ID }} + workflow_app_id: ${{ steps.out.outputs.DIFY_E2E_WORKFLOW_APP_ID }} + file_app_id: ${{ steps.out.outputs.DIFY_E2E_FILE_APP_ID }} + file_chat_app_id: ${{ steps.out.outputs.DIFY_E2E_FILE_CHAT_APP_ID }} + hitl_app_id: ${{ steps.out.outputs.DIFY_E2E_HITL_APP_ID }} + hitl_external_app_id: ${{ steps.out.outputs.DIFY_E2E_HITL_EXTERNAL_APP_ID }} + hitl_single_action_app_id: ${{ steps.out.outputs.DIFY_E2E_HITL_SINGLE_ACTION_APP_ID }} + hitl_multi_node_app_id: ${{ steps.out.outputs.DIFY_E2E_HITL_MULTI_NODE_APP_ID }} + ws2_app_id: ${{ steps.out.outputs.DIFY_E2E_WS2_APP_ID }} steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v4 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v4 with: ref: ${{ inputs.cli_ref || github.ref }} persist-credentials: false @@ -88,7 +87,7 @@ jobs: with: bun-version: latest - - uses: pnpm/action-setup@0ebf47130e4866e96fce0953f49152a61190b271 # v6.0.9 + - uses: pnpm/action-setup@0977fd99725f1db4007ccb2928dbb4e90d06cc86 # v6.0.10 with: package_json_field: packageManager run_install: false @@ -101,18 +100,18 @@ jobs: id: out working-directory: cli env: - DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} - DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} + DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} + DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} - DIFY_E2E_TOKEN: ${{ secrets.DIFY_E2E_TOKEN }} - DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} + DIFY_E2E_TOKEN: ${{ secrets.DIFY_E2E_TOKEN }} + DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} run: bun scripts/e2e-provision.ts -# ════════════════════════════════════════════════════════════════════════════ -# 1-B. framework + output + error-handling (parallel with run/discovery) -# ════════════════════════════════════════════════════════════════════════════ + # ════════════════════════════════════════════════════════════════════════════ + # 1-B. framework + output + error-handling (parallel with run/discovery) + # ════════════════════════════════════════════════════════════════════════════ suite-framework-output-error: - name: "Suite: framework + output + error-handling" + name: 'Suite: framework + output + error-handling' if: ${{ inputs.suite_framework_output_error != 'false' }} needs: provision runs-on: ubuntu-latest @@ -123,7 +122,7 @@ jobs: shell: bash steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v4 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v4 with: ref: ${{ inputs.cli_ref || github.ref }} persist-credentials: false @@ -131,23 +130,23 @@ jobs: - uses: ./.github/actions/setup-web - uses: oven-sh/setup-bun@v2 with: { bun-version: latest } - - uses: pnpm/action-setup@0ebf47130e4866e96fce0953f49152a61190b271 # v6.0.9 + - uses: pnpm/action-setup@0977fd99725f1db4007ccb2928dbb4e90d06cc86 # v6.0.10 with: { package_json_field: packageManager, run_install: false } - run: pnpm install --frozen-lockfile - run: pnpm tree:gen - name: Run framework + output + error-handling env: - DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} - DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} - DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} - DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} - DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} - DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} + DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} + DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} + DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} + DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} + DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} + DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} DIFY_E2E_WORKSPACE_NAME: ${{ needs.provision.outputs.workspace_name }} - DIFY_E2E_CHAT_APP_ID: ${{ needs.provision.outputs.chat_app_id }} + DIFY_E2E_CHAT_APP_ID: ${{ needs.provision.outputs.chat_app_id }} DIFY_E2E_WORKFLOW_APP_ID: ${{ needs.provision.outputs.workflow_app_id }} - DIFY_E2E_INCLUDE: "test/e2e/suites/framework/**/*.e2e.ts,test/e2e/suites/output/**/*.e2e.ts,test/e2e/suites/error-handling/**/*.e2e.ts" + DIFY_E2E_INCLUDE: 'test/e2e/suites/framework/**/*.e2e.ts,test/e2e/suites/output/**/*.e2e.ts,test/e2e/suites/error-handling/**/*.e2e.ts' run: | if [ "${{ inputs.test_scope }}" = "smoke" ]; then pnpm test:e2e -- -t "\[P0\]" @@ -155,11 +154,11 @@ jobs: pnpm test:e2e fi -# ════════════════════════════════════════════════════════════════════════════ -# 1-C. Discovery (parallel) -# ════════════════════════════════════════════════════════════════════════════ + # ════════════════════════════════════════════════════════════════════════════ + # 1-C. Discovery (parallel) + # ════════════════════════════════════════════════════════════════════════════ suite-discovery: - name: "Suite: discovery" + name: 'Suite: discovery' if: ${{ inputs.suite_discovery != 'false' }} needs: provision runs-on: ubuntu-latest @@ -170,7 +169,7 @@ jobs: shell: bash steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v4 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v4 with: ref: ${{ inputs.cli_ref || github.ref }} persist-credentials: false @@ -178,24 +177,24 @@ jobs: - uses: ./.github/actions/setup-web - uses: oven-sh/setup-bun@v2 with: { bun-version: latest } - - uses: pnpm/action-setup@0ebf47130e4866e96fce0953f49152a61190b271 # v6.0.9 + - uses: pnpm/action-setup@0977fd99725f1db4007ccb2928dbb4e90d06cc86 # v6.0.10 with: { package_json_field: packageManager, run_install: false } - run: pnpm install --frozen-lockfile - run: pnpm tree:gen - name: Run discovery suite env: - DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} - DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} - DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} - DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} - DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} - DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} + DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} + DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} + DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} + DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} + DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} + DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} DIFY_E2E_WORKSPACE_NAME: ${{ needs.provision.outputs.workspace_name }} - DIFY_E2E_WS2_ID: ${{ needs.provision.outputs.ws2_id }} - DIFY_E2E_CHAT_APP_ID: ${{ needs.provision.outputs.chat_app_id }} + DIFY_E2E_WS2_ID: ${{ needs.provision.outputs.ws2_id }} + DIFY_E2E_CHAT_APP_ID: ${{ needs.provision.outputs.chat_app_id }} DIFY_E2E_WORKFLOW_APP_ID: ${{ needs.provision.outputs.workflow_app_id }} - DIFY_E2E_INCLUDE: "test/e2e/suites/discovery/**/*.e2e.ts" + DIFY_E2E_INCLUDE: 'test/e2e/suites/discovery/**/*.e2e.ts' run: | if [ "${{ inputs.test_scope }}" = "smoke" ]; then pnpm test:e2e -- -t "\[P0\]" @@ -203,11 +202,11 @@ jobs: pnpm test:e2e fi -# ════════════════════════════════════════════════════════════════════════════ -# 1-D. Run suite — 5 files in matrix (parallel) -# ════════════════════════════════════════════════════════════════════════════ + # ════════════════════════════════════════════════════════════════════════════ + # 1-D. Run suite — 5 files in matrix (parallel) + # ════════════════════════════════════════════════════════════════════════════ suite-run: - name: "Suite: run / ${{ matrix.name }}" + name: 'Suite: run / ${{ matrix.name }}' if: ${{ inputs.suite_run != 'false' }} needs: provision runs-on: ubuntu-latest @@ -233,7 +232,7 @@ jobs: shell: bash steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v4 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v4 with: ref: ${{ inputs.cli_ref || github.ref }} persist-credentials: false @@ -241,30 +240,30 @@ jobs: - uses: ./.github/actions/setup-web - uses: oven-sh/setup-bun@v2 with: { bun-version: latest } - - uses: pnpm/action-setup@0ebf47130e4866e96fce0953f49152a61190b271 # v6.0.9 + - uses: pnpm/action-setup@0977fd99725f1db4007ccb2928dbb4e90d06cc86 # v6.0.10 with: { package_json_field: packageManager, run_install: false } - run: pnpm install --frozen-lockfile - run: pnpm tree:gen - - name: "Run run/${{ matrix.name }}" + - name: 'Run run/${{ matrix.name }}' env: - DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} - DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} - DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} - DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} - DIFY_E2E_SSO_TOKEN: ${{ secrets.DIFY_E2E_SSO_TOKEN }} - DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} - DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} - DIFY_E2E_WORKSPACE_NAME: ${{ needs.provision.outputs.workspace_name }} - DIFY_E2E_CHAT_APP_ID: ${{ needs.provision.outputs.chat_app_id }} - DIFY_E2E_WORKFLOW_APP_ID: ${{ needs.provision.outputs.workflow_app_id }} - DIFY_E2E_FILE_APP_ID: ${{ needs.provision.outputs.file_app_id }} - DIFY_E2E_FILE_CHAT_APP_ID: ${{ needs.provision.outputs.file_chat_app_id }} - DIFY_E2E_HITL_APP_ID: ${{ needs.provision.outputs.hitl_app_id }} - DIFY_E2E_HITL_EXTERNAL_APP_ID: ${{ needs.provision.outputs.hitl_external_app_id }} + DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} + DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} + DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} + DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} + DIFY_E2E_SSO_TOKEN: ${{ secrets.DIFY_E2E_SSO_TOKEN }} + DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} + DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} + DIFY_E2E_WORKSPACE_NAME: ${{ needs.provision.outputs.workspace_name }} + DIFY_E2E_CHAT_APP_ID: ${{ needs.provision.outputs.chat_app_id }} + DIFY_E2E_WORKFLOW_APP_ID: ${{ needs.provision.outputs.workflow_app_id }} + DIFY_E2E_FILE_APP_ID: ${{ needs.provision.outputs.file_app_id }} + DIFY_E2E_FILE_CHAT_APP_ID: ${{ needs.provision.outputs.file_chat_app_id }} + DIFY_E2E_HITL_APP_ID: ${{ needs.provision.outputs.hitl_app_id }} + DIFY_E2E_HITL_EXTERNAL_APP_ID: ${{ needs.provision.outputs.hitl_external_app_id }} DIFY_E2E_HITL_SINGLE_ACTION_APP_ID: ${{ needs.provision.outputs.hitl_single_action_app_id }} - DIFY_E2E_HITL_MULTI_NODE_APP_ID: ${{ needs.provision.outputs.hitl_multi_node_app_id }} - DIFY_E2E_INCLUDE: "test/e2e/suites/run/${{ matrix.file }}" + DIFY_E2E_HITL_MULTI_NODE_APP_ID: ${{ needs.provision.outputs.hitl_multi_node_app_id }} + DIFY_E2E_INCLUDE: 'test/e2e/suites/run/${{ matrix.file }}' run: | if [ "${{ inputs.test_scope }}" = "smoke" ]; then pnpm test:e2e -- -t "\[P0\]" @@ -280,11 +279,11 @@ jobs: path: cli/test-results/ retention-days: 3 -# ════════════════════════════════════════════════════════════════════════════ -# 1-E. auth/login + status + whoami (parallel, read-only, safe) -# ════════════════════════════════════════════════════════════════════════════ + # ════════════════════════════════════════════════════════════════════════════ + # 1-E. auth/login + status + whoami (parallel, read-only, safe) + # ════════════════════════════════════════════════════════════════════════════ suite-auth-safe: - name: "Suite: auth (login / status / whoami)" + name: 'Suite: auth (login / status / whoami)' if: ${{ inputs.suite_auth != 'false' }} needs: provision runs-on: ubuntu-latest @@ -295,7 +294,7 @@ jobs: shell: bash steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v4 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v4 with: ref: ${{ inputs.cli_ref || github.ref }} persist-credentials: false @@ -303,22 +302,22 @@ jobs: - uses: ./.github/actions/setup-web - uses: oven-sh/setup-bun@v2 with: { bun-version: latest } - - uses: pnpm/action-setup@0ebf47130e4866e96fce0953f49152a61190b271 # v6.0.9 + - uses: pnpm/action-setup@0977fd99725f1db4007ccb2928dbb4e90d06cc86 # v6.0.10 with: { package_json_field: packageManager, run_install: false } - run: pnpm install --frozen-lockfile - run: pnpm tree:gen - name: Run auth/login + status + whoami env: - DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} - DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} - DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} - DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} - DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} - DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} + DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} + DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} + DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} + DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} + DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} + DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} DIFY_E2E_WORKSPACE_NAME: ${{ needs.provision.outputs.workspace_name }} - DIFY_E2E_WS2_ID: ${{ needs.provision.outputs.ws2_id }} - DIFY_E2E_INCLUDE: "test/e2e/suites/auth/login.e2e.ts,test/e2e/suites/auth/status.e2e.ts,test/e2e/suites/auth/whoami.e2e.ts" + DIFY_E2E_WS2_ID: ${{ needs.provision.outputs.ws2_id }} + DIFY_E2E_INCLUDE: 'test/e2e/suites/auth/login.e2e.ts,test/e2e/suites/auth/status.e2e.ts,test/e2e/suites/auth/whoami.e2e.ts' run: | if [ "${{ inputs.test_scope }}" = "smoke" ]; then pnpm test:e2e -- -t "\[P0\]" @@ -326,13 +325,13 @@ jobs: pnpm test:e2e fi -# ════════════════════════════════════════════════════════════════════════════ -# 2. DESTRUCTIVE — auth/use + devices + logout + agent (serial, runs LAST) -# Must wait for ALL parallel suites to finish to avoid token revocation -# invalidating other in-flight requests. -# ════════════════════════════════════════════════════════════════════════════ + # ════════════════════════════════════════════════════════════════════════════ + # 2. DESTRUCTIVE — auth/use + devices + logout + agent (serial, runs LAST) + # Must wait for ALL parallel suites to finish to avoid token revocation + # invalidating other in-flight requests. + # ════════════════════════════════════════════════════════════════════════════ suite-last: - name: "Suite: auth-use + devices + logout + agent (last, serial)" + name: 'Suite: auth-use + devices + logout + agent (last, serial)' # Runs when auth is selected; also runs after all parallel jobs finish if: ${{ inputs.suite_auth != 'false' || inputs.suite_agent != 'false' }} needs: @@ -351,7 +350,7 @@ jobs: shell: bash steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v4 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v4 with: ref: ${{ inputs.cli_ref || github.ref }} persist-credentials: false @@ -359,27 +358,27 @@ jobs: - uses: ./.github/actions/setup-web - uses: oven-sh/setup-bun@v2 with: { bun-version: latest } - - uses: pnpm/action-setup@0ebf47130e4866e96fce0953f49152a61190b271 # v6.0.9 + - uses: pnpm/action-setup@0977fd99725f1db4007ccb2928dbb4e90d06cc86 # v6.0.10 with: { package_json_field: packageManager, run_install: false } - run: pnpm install --frozen-lockfile - run: pnpm tree:gen - name: Run use / devices / logout / agent (serial) env: - DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} - DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} - DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} - DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} - DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} - DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} - DIFY_E2E_WORKSPACE_NAME: ${{ needs.provision.outputs.workspace_name }} - DIFY_E2E_WS2_ID: ${{ needs.provision.outputs.ws2_id }} - DIFY_E2E_CHAT_APP_ID: ${{ needs.provision.outputs.chat_app_id }} - DIFY_E2E_WORKFLOW_APP_ID: ${{ needs.provision.outputs.workflow_app_id }} - DIFY_E2E_HITL_APP_ID: ${{ needs.provision.outputs.hitl_app_id }} - DIFY_E2E_HITL_EXTERNAL_APP_ID: ${{ needs.provision.outputs.hitl_external_app_id }} + DIFY_E2E_HOST: ${{ secrets.DIFY_E2E_HOST }} + DIFY_E2E_EMAIL: ${{ secrets.DIFY_E2E_EMAIL }} + DIFY_E2E_PASSWORD: ${{ secrets.DIFY_E2E_PASSWORD }} + DIFY_E2E_EDITION: ${{ inputs.edition || 'ee' }} + DIFY_E2E_TOKEN: ${{ needs.provision.outputs.token }} + DIFY_E2E_WORKSPACE_ID: ${{ needs.provision.outputs.workspace_id }} + DIFY_E2E_WORKSPACE_NAME: ${{ needs.provision.outputs.workspace_name }} + DIFY_E2E_WS2_ID: ${{ needs.provision.outputs.ws2_id }} + DIFY_E2E_CHAT_APP_ID: ${{ needs.provision.outputs.chat_app_id }} + DIFY_E2E_WORKFLOW_APP_ID: ${{ needs.provision.outputs.workflow_app_id }} + DIFY_E2E_HITL_APP_ID: ${{ needs.provision.outputs.hitl_app_id }} + DIFY_E2E_HITL_EXTERNAL_APP_ID: ${{ needs.provision.outputs.hitl_external_app_id }} DIFY_E2E_HITL_SINGLE_ACTION_APP_ID: ${{ needs.provision.outputs.hitl_single_action_app_id }} - DIFY_E2E_HITL_MULTI_NODE_APP_ID: ${{ needs.provision.outputs.hitl_multi_node_app_id }} + DIFY_E2E_HITL_MULTI_NODE_APP_ID: ${{ needs.provision.outputs.hitl_multi_node_app_id }} run: | # Collect files in safe order: use → devices → logout (revokes last) → agent FILES=() diff --git a/.github/workflows/cli-edge.yml b/.github/workflows/cli-edge.yml index d4d789d0b14..7d8c5290581 100644 --- a/.github/workflows/cli-edge.yml +++ b/.github/workflows/cli-edge.yml @@ -23,7 +23,7 @@ jobs: working-directory: ./cli steps: - name: Checkout - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false fetch-depth: 0 @@ -36,7 +36,7 @@ jobs: uses: ./.github/actions/setup-web - name: Setup Bun - uses: oven-sh/setup-bun@4bc047ad259df6fc24a6c9b0f9a0cb08cf17fbe5 # v2.0.2 + uses: oven-sh/setup-bun@0c5077e51419868618aeaa5fe8019c62421857d6 # v2.2.0 with: bun-version-file: cli/.bun-version diff --git a/.github/workflows/cli-release.yml b/.github/workflows/cli-release.yml index 788f0bfe21c..6e0670d11e7 100644 --- a/.github/workflows/cli-release.yml +++ b/.github/workflows/cli-release.yml @@ -35,7 +35,7 @@ jobs: dify_tag: ${{ steps.resolve.outputs.dify_tag }} steps: - name: Checkout - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -98,7 +98,7 @@ jobs: DIFY_TAG: ${{ needs.validate.outputs.dify_tag }} steps: - name: Checkout - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false fetch-depth: 1 @@ -114,7 +114,7 @@ jobs: run: node scripts/release-naming.mjs github-env >> "$GITHUB_ENV" - name: Setup Bun - uses: oven-sh/setup-bun@4bc047ad259df6fc24a6c9b0f9a0cb08cf17fbe5 # v2.0.2 + uses: oven-sh/setup-bun@0c5077e51419868618aeaa5fe8019c62421857d6 # v2.2.0 with: bun-version-file: cli/.bun-version diff --git a/.github/workflows/cli-smoke.yml b/.github/workflows/cli-smoke.yml index a46d93c2ac8..7c70a52428b 100644 --- a/.github/workflows/cli-smoke.yml +++ b/.github/workflows/cli-smoke.yml @@ -4,11 +4,11 @@ on: workflow_dispatch: inputs: dify_version: - description: "Dify image tag to test against (e.g. 1.7.0)" + description: 'Dify image tag to test against (e.g. 1.7.0)' type: string required: true cli_ref: - description: "Git ref to build the cli from (default: current branch)" + description: 'Git ref to build the cli from (default: current branch)' type: string required: false @@ -24,7 +24,7 @@ jobs: shell: bash steps: - name: Checkout cli ref - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: ref: ${{ inputs.cli_ref || github.ref }} persist-credentials: false diff --git a/.github/workflows/cli-tests.yml b/.github/workflows/cli-tests.yml index 3638f79cae6..39fb7647177 100644 --- a/.github/workflows/cli-tests.yml +++ b/.github/workflows/cli-tests.yml @@ -30,7 +30,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false diff --git a/.github/workflows/db-migration-test.yml b/.github/workflows/db-migration-test.yml index ae3d4f67c48..c68db72d226 100644 --- a/.github/workflows/db-migration-test.yml +++ b/.github/workflows/db-migration-test.yml @@ -13,16 +13,16 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false - name: Setup UV and Python - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true - python-version: "3.12" + python-version: '3.12' cache-dependency-glob: api/uv.lock - name: Install dependencies @@ -40,7 +40,7 @@ jobs: cp envs/middleware.env.example middleware.env - name: Set up Middlewares - uses: hoverkraft-tech/compose-action@11beaa1c2dae4e8ed7b1665aa074723b6cecb0e4 # v3.0.0 + uses: hoverkraft-tech/compose-action@ee6af68587292d72db67743171809c19787df4c9 # v3.1.0 with: compose-file: | docker/docker-compose.middleware.yaml @@ -63,16 +63,16 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false - name: Setup UV and Python - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true - python-version: "3.12" + python-version: '3.12' cache-dependency-glob: api/uv.lock - name: Install dependencies @@ -94,7 +94,7 @@ jobs: sed -i 's/DB_USERNAME=postgres/DB_USERNAME=mysql/' middleware.env - name: Set up Middlewares - uses: hoverkraft-tech/compose-action@11beaa1c2dae4e8ed7b1665aa074723b6cecb0e4 # v3.0.0 + uses: hoverkraft-tech/compose-action@ee6af68587292d72db67743171809c19787df4c9 # v3.1.0 with: compose-file: | docker/docker-compose.middleware.yaml diff --git a/.github/workflows/deploy-agent.yml b/.github/workflows/deploy-agent.yml index e244bb3f949..75844b1a1d1 100644 --- a/.github/workflows/deploy-agent.yml +++ b/.github/workflows/deploy-agent.yml @@ -5,9 +5,9 @@ permissions: on: workflow_run: - workflows: ["Build and Push API & Web"] + workflows: ['Build and Push API & Web'] branches: - - "deploy/agent" + - 'deploy/agent' types: - completed diff --git a/.github/workflows/deploy-dev.yml b/.github/workflows/deploy-dev.yml index c2ff8c63324..cddf5edc351 100644 --- a/.github/workflows/deploy-dev.yml +++ b/.github/workflows/deploy-dev.yml @@ -2,9 +2,9 @@ name: Deploy Dev on: workflow_run: - workflows: ["Build and Push API & Web"] + workflows: ['Build and Push API & Web'] branches: - - "deploy/dev" + - 'deploy/dev' types: - completed diff --git a/.github/workflows/deploy-enterprise.yml b/.github/workflows/deploy-enterprise.yml index 2740541f0ff..3d57efcbc2e 100644 --- a/.github/workflows/deploy-enterprise.yml +++ b/.github/workflows/deploy-enterprise.yml @@ -5,9 +5,9 @@ permissions: on: workflow_run: - workflows: ["Build and Push API & Web"] + workflows: ['Build and Push API & Web'] branches: - - "deploy/enterprise" + - 'deploy/enterprise' types: - completed diff --git a/.github/workflows/deploy-knowledge.yml b/.github/workflows/deploy-knowledge.yml index 26bc39d8bc0..0c85fe7dc49 100644 --- a/.github/workflows/deploy-knowledge.yml +++ b/.github/workflows/deploy-knowledge.yml @@ -1,13 +1,14 @@ name: Deploy Knowledge permissions: + actions: read contents: read on: workflow_run: - workflows: ["Build and Push API & Web"] + workflows: ['Build and Push API & Web'] branches: - - "deploy/konwledge" + - 'deploy/konwledge' types: - completed @@ -18,6 +19,48 @@ jobs: github.event.workflow_run.conclusion == 'success' && github.event.workflow_run.head_branch == 'deploy/konwledge' steps: + - name: Wait for KnowledgeFS CI + uses: actions/github-script@3a2844b7e9c422d3c10d287c895573f7108da1b3 # v9.0.0 + timeout-minutes: 35 + with: + github-token: ${{ secrets.GITHUB_TOKEN }} + script: | + const workflowId = "knowledge-fs-ci.yml"; + const headBranch = context.payload.workflow_run.head_branch; + const headSha = context.payload.workflow_run.head_sha; + const deadline = Date.now() + 30 * 60 * 1000; + const pollIntervalMs = 15 * 1000; + + while (Date.now() < deadline) { + const { data } = await github.rest.actions.listWorkflowRuns({ + owner: context.repo.owner, + repo: context.repo.repo, + workflow_id: workflowId, + branch: headBranch, + event: "push", + head_sha: headSha, + per_page: 10, + }); + const run = data.workflow_runs[0]; + + if (!run) { + core.info(`Waiting for ${workflowId} to start for ${headSha}.`); + } else if (run.status !== "completed") { + core.info(`Waiting for ${run.html_url}; current status is ${run.status}.`); + } else if (run.conclusion !== "success") { + throw new Error( + `${workflowId} did not succeed for ${headSha}: ${run.conclusion} (${run.html_url})`, + ); + } else { + core.info(`KnowledgeFS CI succeeded for ${headSha}: ${run.html_url}`); + return; + } + + await new Promise((resolve) => setTimeout(resolve, pollIntervalMs)); + } + + throw new Error(`Timed out waiting for ${workflowId} to succeed for ${headSha}.`); + - name: Deploy to server uses: appleboy/ssh-action@0ff4204d59e8e51228ff73bce53f80d53301dee2 # v1.2.5 with: diff --git a/.github/workflows/deploy-saas.yml b/.github/workflows/deploy-saas.yml index b00883c8c79..f0165c35ed4 100644 --- a/.github/workflows/deploy-saas.yml +++ b/.github/workflows/deploy-saas.yml @@ -5,9 +5,9 @@ permissions: on: workflow_run: - workflows: ["Build and Push API & Web"] + workflows: ['Build and Push API & Web'] branches: - - "deploy/saas" + - 'deploy/saas' types: - completed diff --git a/.github/workflows/docker-build.yml b/.github/workflows/docker-build.yml index 9fb71a42cfd..4e70609b7e1 100644 --- a/.github/workflows/docker-build.yml +++ b/.github/workflows/docker-build.yml @@ -3,7 +3,7 @@ name: Build docker image on: pull_request: branches: - - "main" + - 'main' paths: - api/Dockerfile - api/Dockerfile.dockerignore @@ -33,36 +33,36 @@ jobs: strategy: matrix: include: - - service_name: "api-amd64" + - service_name: 'api-amd64' platform: linux/amd64 runs_on: depot-ubuntu-24.04-4 - context: "{{defaultContext}}" - file: "api/Dockerfile" - - service_name: "api-arm64" + context: '{{defaultContext}}' + file: 'api/Dockerfile' + - service_name: 'api-arm64' platform: linux/arm64 runs_on: depot-ubuntu-24.04-4 - context: "{{defaultContext}}" - file: "api/Dockerfile" - - service_name: "web-amd64" + context: '{{defaultContext}}' + file: 'api/Dockerfile' + - service_name: 'web-amd64' platform: linux/amd64 runs_on: depot-ubuntu-24.04-4 - context: "{{defaultContext}}" - file: "web/Dockerfile" - - service_name: "web-arm64" + context: '{{defaultContext}}' + file: 'web/Dockerfile' + - service_name: 'web-arm64' platform: linux/arm64 runs_on: depot-ubuntu-24.04-4 - context: "{{defaultContext}}" - file: "web/Dockerfile" - - service_name: "local-sandbox-amd64" + context: '{{defaultContext}}' + file: 'web/Dockerfile' + - service_name: 'local-sandbox-amd64' platform: linux/amd64 runs_on: depot-ubuntu-24.04-4 - context: "{{defaultContext}}:dify-agent-runtime" - file: "docker/Dockerfile" - - service_name: "local-sandbox-arm64" + context: '{{defaultContext}}:dify-agent-runtime' + file: 'docker/Dockerfile' + - service_name: 'local-sandbox-arm64' platform: linux/arm64 runs_on: depot-ubuntu-24.04-4 - context: "{{defaultContext}}:dify-agent-runtime" - file: "docker/Dockerfile" + context: '{{defaultContext}}:dify-agent-runtime' + file: 'docker/Dockerfile' steps: - name: Set up Depot CLI uses: depot/setup-action@15c09a5f77a0840ad4bce955686522a257853461 # v1.7.1 @@ -85,15 +85,15 @@ jobs: strategy: matrix: include: - - service_name: "api-amd64" - context: "{{defaultContext}}" - file: "api/Dockerfile" - - service_name: "web-amd64" - context: "{{defaultContext}}" - file: "web/Dockerfile" - - service_name: "local-sandbox-amd64" - context: "{{defaultContext}}:dify-agent-runtime" - file: "docker/Dockerfile" + - service_name: 'api-amd64' + context: '{{defaultContext}}' + file: 'api/Dockerfile' + - service_name: 'web-amd64' + context: '{{defaultContext}}' + file: 'web/Dockerfile' + - service_name: 'local-sandbox-amd64' + context: '{{defaultContext}}:dify-agent-runtime' + file: 'docker/Dockerfile' steps: - name: Set up Docker Buildx uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0 diff --git a/.github/workflows/hotfix-cherry-pick.yml b/.github/workflows/hotfix-cherry-pick.yml index 55a1ac5f5ed..aff6c1c3664 100644 --- a/.github/workflows/hotfix-cherry-pick.yml +++ b/.github/workflows/hotfix-cherry-pick.yml @@ -24,7 +24,7 @@ jobs: name: Require cherry-pick provenance runs-on: depot-ubuntu-24.04 steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 diff --git a/.github/workflows/labeler.yml b/.github/workflows/labeler.yml index 65c972522e3..fb7fd284b65 100644 --- a/.github/workflows/labeler.yml +++ b/.github/workflows/labeler.yml @@ -1,4 +1,4 @@ -name: "Pull Request Labeler" +name: 'Pull Request Labeler' on: pull_request_target: @@ -9,6 +9,6 @@ jobs: pull-requests: write runs-on: depot-ubuntu-24.04 steps: - - uses: actions/labeler@b8dd2d9be0f68b860e7dae5dae7d772984eacd6d # v6.2.0 + - uses: actions/labeler@bf12e9b00b37c5c0ca2b87b79b2daf7891dbda13 # v7.0.0 with: sync-labels: true diff --git a/.github/workflows/main-ci.yml b/.github/workflows/main-ci.yml index aec8514a69d..72b1f496c71 100644 --- a/.github/workflows/main-ci.yml +++ b/.github/workflows/main-ci.yml @@ -2,9 +2,9 @@ name: Main CI Pipeline on: pull_request: - branches: ["main"] + branches: ['main'] merge_group: - branches: ["main"] + branches: ['main'] types: [checks_requested] permissions: @@ -15,7 +15,7 @@ permissions: statuses: write concurrency: - group: main-ci-${{ github.head_ref || github.run_id }} + group: main-ci-${{ github.event.pull_request.number || github.event.merge_group.head_sha || github.run_id }} cancel-in-progress: true jobs: @@ -47,8 +47,8 @@ jobs: migration-changed: ${{ steps.changes.outputs.migration }} sandbox-runtime-changed: ${{ steps.changes.outputs.sandbox-runtime }} steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 - - uses: dorny/paths-filter@7b450fff21473bca461d4b92ce414b9d0420d706 # v4.0.2 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: dorny/paths-filter@ceb8a2b8f2d89434be7ff52d3de7ec3738c5cc9d # v4.0.3 id: changes with: filters: | @@ -89,6 +89,7 @@ jobs: - 'pnpm-lock.yaml' - 'pnpm-workspace.yaml' - '.nvmrc' + - '.github/workflows/main-ci.yml' - '.github/workflows/web-tests.yml' - '.github/actions/setup-web/**' e2e: diff --git a/.github/workflows/post-merge.yml b/.github/workflows/post-merge.yml index c7ee9850b08..de1fce7977c 100644 --- a/.github/workflows/post-merge.yml +++ b/.github/workflows/post-merge.yml @@ -2,7 +2,7 @@ name: Post-Merge Checks on: push: - branches: ["main"] + branches: ['main'] permissions: contents: read @@ -18,8 +18,8 @@ jobs: outputs: external-e2e-changed: ${{ steps.changes.outputs.external_e2e }} steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 - - uses: dorny/paths-filter@7b450fff21473bca461d4b92ce414b9d0420d706 # v4.0.2 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: dorny/paths-filter@ceb8a2b8f2d89434be7ff52d3de7ec3738c5cc9d # v4.0.3 id: changes with: filters: | diff --git a/.github/workflows/pyrefly-diff.yml b/.github/workflows/pyrefly-diff.yml index b8ed10612c4..27b04f030e1 100644 --- a/.github/workflows/pyrefly-diff.yml +++ b/.github/workflows/pyrefly-diff.yml @@ -17,12 +17,12 @@ jobs: pull-requests: write steps: - name: Checkout PR branch - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 - name: Setup Python & UV - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true diff --git a/.github/workflows/pyrefly-type-coverage-comment.yml b/.github/workflows/pyrefly-type-coverage-comment.yml index 8fd1e0f788e..eacf485c7a1 100644 --- a/.github/workflows/pyrefly-type-coverage-comment.yml +++ b/.github/workflows/pyrefly-type-coverage-comment.yml @@ -21,10 +21,10 @@ jobs: if: ${{ github.event.workflow_run.conclusion == 'success' && github.event.workflow_run.pull_requests[0].head.repo.full_name != github.repository }} steps: - name: Checkout default branch (trusted code) - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Setup Python & UV - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true diff --git a/.github/workflows/pyrefly-type-coverage.yml b/.github/workflows/pyrefly-type-coverage.yml index fb223342701..19a2e18d48b 100644 --- a/.github/workflows/pyrefly-type-coverage.yml +++ b/.github/workflows/pyrefly-type-coverage.yml @@ -17,12 +17,12 @@ jobs: pull-requests: write steps: - name: Checkout PR branch - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 - name: Setup Python & UV - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true diff --git a/.github/workflows/sandbox-runtime-tests.yml b/.github/workflows/sandbox-runtime-tests.yml index 7e36067446c..45e26f06860 100644 --- a/.github/workflows/sandbox-runtime-tests.yml +++ b/.github/workflows/sandbox-runtime-tests.yml @@ -21,7 +21,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false @@ -45,7 +45,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false @@ -57,7 +57,7 @@ jobs: cache-dependency-path: dify-agent-runtime/go.sum - name: Run golangci-lint - uses: golangci/golangci-lint-action@ba0d7d2ec06a0ea1cb5fa41b2e4a3ab91d21278a # v6.5.0 + uses: golangci/golangci-lint-action@ba0d7d2ec06a0ea1cb5fa41b2e4a3ab91d21278a # v9.3.0 with: working-directory: dify-agent-runtime version: latest @@ -72,7 +72,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false diff --git a/.github/workflows/semantic-pull-request.yml b/.github/workflows/semantic-pull-request.yml index 6f3193bbf53..51f5b373593 100644 --- a/.github/workflows/semantic-pull-request.yml +++ b/.github/workflows/semantic-pull-request.yml @@ -8,7 +8,7 @@ on: - reopened - synchronize merge_group: - branches: ["main"] + branches: ['main'] types: [checks_requested] jobs: diff --git a/.github/workflows/stale.yml b/.github/workflows/stale.yml index 308ff84e6fa..b1c4e85322d 100644 --- a/.github/workflows/stale.yml +++ b/.github/workflows/stale.yml @@ -11,20 +11,19 @@ on: jobs: stale: - runs-on: depot-ubuntu-24.04 permissions: issues: write pull-requests: write steps: - - uses: actions/stale@1e223db275d687790206a7acac4d1a11bd6fe629 # v10.4.0 + - uses: actions/stale@4391f3da665fdf50b6810c1a66712fb9ba21aa93 # v11.0.0 with: days-before-issue-stale: 15 days-before-issue-close: 3 repo-token: ${{ secrets.GITHUB_TOKEN }} - stale-issue-message: "Closed due to inactivity. If you have any questions, you can reopen it." - stale-pr-message: "Closed due to inactivity. If you have any questions, you can reopen it." + stale-issue-message: 'Closed due to inactivity. If you have any questions, you can reopen it.' + stale-pr-message: 'Closed due to inactivity. If you have any questions, you can reopen it.' stale-issue-label: 'no-issue-activity' stale-pr-label: 'no-pr-activity' any-of-labels: '🌚 invalid,🙋‍♂️ question,wont-fix,no-issue-activity,no-pr-activity,💪 enhancement,🤔 cant-reproduce,🙏 help wanted' diff --git a/.github/workflows/style.yml b/.github/workflows/style.yml index ec7f5779083..5507baaba7b 100644 --- a/.github/workflows/style.yml +++ b/.github/workflows/style.yml @@ -7,10 +7,6 @@ on: required: true type: string -concurrency: - group: style-${{ github.head_ref || github.run_id }} - cancel-in-progress: true - permissions: checks: write statuses: write @@ -23,7 +19,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false fetch-depth: 0 @@ -45,10 +41,10 @@ jobs: - name: Setup UV and Python if: steps.changed-files.outputs.any_changed == 'true' - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: false - python-version: "3.12" + python-version: '3.12' cache-dependency-glob: api/uv.lock - name: Install dependencies @@ -93,7 +89,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -109,6 +105,8 @@ jobs: package.json pnpm-lock.yaml pnpm-workspace.yaml + knip.config.ts + scripts/check-web-production-unused-after-knip-fix.mjs .nvmrc .github/workflows/style.yml .github/actions/setup-web/** @@ -125,14 +123,17 @@ jobs: - name: Web dead code check if: steps.changed-files.outputs.any_changed == 'true' + working-directory: . run: vp run knip - name: Web dead code check production if: steps.changed-files.outputs.any_changed == 'true' + working-directory: . run: vp run knip:production - name: Web production unused declarations check if: steps.changed-files.outputs.any_changed == 'true' + working-directory: . run: vp run knip:production-unused-check ts-common-style: @@ -144,7 +145,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -168,9 +169,7 @@ jobs: oxlint-suppressions.json eslint-suppressions.json .vscode/** - .github/workflows/autofix.yml - .github/workflows/style.yml - .github/actions/setup-web/** + .github/** - name: Setup web environment if: steps.changed-files.outputs.any_changed == 'true' @@ -186,7 +185,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false diff --git a/.github/workflows/tool-test-sdks.yaml b/.github/workflows/tool-test-sdks.yaml index d474396a300..2d0133131eb 100644 --- a/.github/workflows/tool-test-sdks.yaml +++ b/.github/workflows/tool-test-sdks.yaml @@ -24,7 +24,7 @@ jobs: working-directory: sdks/nodejs-client steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false diff --git a/.github/workflows/translate-i18n-claude.yml b/.github/workflows/translate-i18n-claude.yml index 77702ffbaa9..b81d8c6876e 100644 --- a/.github/workflows/translate-i18n-claude.yml +++ b/.github/workflows/translate-i18n-claude.yml @@ -40,7 +40,7 @@ jobs: steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 token: ${{ secrets.GITHUB_TOKEN }} @@ -158,7 +158,7 @@ jobs: - name: Run Claude Code for Translation Sync if: steps.context.outputs.CHANGED_FILES != '' - uses: anthropics/claude-code-action@af0559ee4f514d1ef21826982bed13f7edc3c35e # v1.0.178 + uses: anthropics/claude-code-action@1623c36729ac1cd5895198cded705a287de7db79 # v1.0.187 with: anthropic_api_key: ${{ secrets.ANTHROPIC_API_KEY }} github_token: ${{ secrets.GITHUB_TOKEN }} diff --git a/.github/workflows/trigger-i18n-sync.yml b/.github/workflows/trigger-i18n-sync.yml index ad2a0675afa..6cb096562ff 100644 --- a/.github/workflows/trigger-i18n-sync.yml +++ b/.github/workflows/trigger-i18n-sync.yml @@ -21,7 +21,7 @@ jobs: steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 diff --git a/.github/workflows/vdb-tests-full.yml b/.github/workflows/vdb-tests-full.yml index 27923401e7c..e5680416058 100644 --- a/.github/workflows/vdb-tests-full.yml +++ b/.github/workflows/vdb-tests-full.yml @@ -20,11 +20,11 @@ jobs: strategy: matrix: python-version: - - "3.12" + - '3.12' steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -36,7 +36,7 @@ jobs: remove_tool_cache: true - name: Setup UV and Python - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true python-version: ${{ matrix.python-version }} @@ -48,16 +48,16 @@ jobs: - name: Install dependencies run: uv sync --project api --dev -# - name: Set up Vector Store (TiDB) -# uses: hoverkraft-tech/compose-action@v2.0.2 -# with: -# compose-file: docker/tidb/docker-compose.yaml -# services: | -# tidb -# tiflash + # - name: Set up Vector Store (TiDB) + # uses: hoverkraft-tech/compose-action@v2.0.2 + # with: + # compose-file: docker/tidb/docker-compose.yaml + # services: | + # tidb + # tiflash -# - name: Check VDB Ready (TiDB) -# run: uv run --project api python api/providers/vdb/tidb-vector/tests/integration_tests/check_tiflash_ready.py + # - name: Check VDB Ready (TiDB) + # run: uv run --project api python api/providers/vdb/tidb-vector/tests/integration_tests/check_tiflash_ready.py - name: Test Vector Stores run: | diff --git a/.github/workflows/vdb-tests.yml b/.github/workflows/vdb-tests.yml index 634fb1a4097..0fa646ad8e5 100644 --- a/.github/workflows/vdb-tests.yml +++ b/.github/workflows/vdb-tests.yml @@ -17,11 +17,11 @@ jobs: strategy: matrix: python-version: - - "3.12" + - '3.12' steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -33,7 +33,7 @@ jobs: remove_tool_cache: true - name: Setup UV and Python - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true python-version: ${{ matrix.python-version }} @@ -45,16 +45,16 @@ jobs: - name: Install dependencies run: uv sync --project api --dev -# - name: Set up Vector Store (TiDB) -# uses: hoverkraft-tech/compose-action@v2.0.2 -# with: -# compose-file: docker/tidb/docker-compose.yaml -# services: | -# tidb -# tiflash + # - name: Set up Vector Store (TiDB) + # uses: hoverkraft-tech/compose-action@v2.0.2 + # with: + # compose-file: docker/tidb/docker-compose.yaml + # services: | + # tidb + # tiflash -# - name: Check VDB Ready (TiDB) -# run: uv run --project api python api/providers/vdb/tidb-vector/tests/integration_tests/check_tiflash_ready.py + # - name: Check VDB Ready (TiDB) + # run: uv run --project api python api/providers/vdb/tidb-vector/tests/integration_tests/check_tiflash_ready.py - name: Test Vector Stores run: | diff --git a/.github/workflows/web-e2e.yml b/.github/workflows/web-e2e.yml index df7ed8d7c92..536fe160d74 100644 --- a/.github/workflows/web-e2e.yml +++ b/.github/workflows/web-e2e.yml @@ -11,10 +11,6 @@ on: permissions: contents: read -concurrency: - group: web-e2e-${{ github.head_ref || github.run_id }} - cancel-in-progress: true - jobs: test: name: Web Full-Stack E2E @@ -26,18 +22,23 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false + - name: Switch Ubuntu mirror + run: | + sudo sed -i 's|http://us-east-1.ec2.archive.ubuntu.com/ubuntu|http://azure.archive.ubuntu.com/ubuntu|g' /etc/apt/sources.list.d/ubuntu.sources + sudo apt-get update + - name: Setup web dependencies uses: ./.github/actions/setup-web - name: Setup UV and Python - uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2 + uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true - python-version: "3.12" + python-version: '3.12' cache-dependency-glob: | api/uv.lock dify-agent/uv.lock @@ -50,9 +51,17 @@ jobs: working-directory: ./e2e run: vp run test:unit - - name: Install Playwright browser + - name: Install Playwright browsers for core E2E + if: ${{ !inputs.run-external-runtime }} + timeout-minutes: 15 working-directory: ./e2e - run: vp run e2e:install + run: vp run e2e:install:ci + + - name: Install Chromium for external runtime E2E + if: ${{ inputs.run-external-runtime }} + timeout-minutes: 15 + working-directory: ./e2e + run: vp run e2e:install:ci:chromium - name: Run isolated source-api and built-web Cucumber E2E tests if: ${{ !inputs.run-external-runtime }} @@ -61,7 +70,7 @@ jobs: E2E_ADMIN_EMAIL: e2e-admin@example.com E2E_ADMIN_NAME: E2E Admin E2E_ADMIN_PASSWORD: E2eAdmin12345 - E2E_FORCE_WEB_BUILD: "1" + E2E_FORCE_WEB_BUILD: '1' E2E_INIT_PASSWORD: E2eInit12345 run: vp run e2e:full @@ -121,7 +130,7 @@ jobs: E2E_AGENT_DECISION_MODEL_NAME: ${{ vars.E2E_AGENT_DECISION_MODEL_NAME || 'gpt-5.5' }} E2E_AGENT_DECISION_MODEL_PROVIDER: ${{ vars.E2E_AGENT_DECISION_MODEL_PROVIDER || 'openai' }} E2E_AGENT_DECISION_MODEL_TYPE: ${{ vars.E2E_AGENT_DECISION_MODEL_TYPE || 'llm' }} - E2E_FORCE_WEB_BUILD: "1" + E2E_FORCE_WEB_BUILD: '1' E2E_INIT_PASSWORD: E2eInit12345 E2E_MARKETPLACE_API_URL: ${{ vars.E2E_MARKETPLACE_API_URL }} E2E_MARKETPLACE_PLUGIN_IDS: ${{ vars.E2E_MARKETPLACE_PLUGIN_IDS }} @@ -138,21 +147,6 @@ jobs: exit 1 fi - teardown_external_runtime() { - local run_status=$? - trap - EXIT - if ! vp run e2e:middleware:down; then - echo "::error title=E2E teardown failed::External runtime middleware did not shut down cleanly." - if [[ "$run_status" -eq 0 ]]; then - run_status=1 - fi - fi - exit "$run_status" - } - - trap teardown_external_runtime EXIT - vp run e2e:middleware:up - vp run e2e:post-merge:prepare vp run e2e:post-merge - name: Upload Cucumber report diff --git a/.github/workflows/web-tests.yml b/.github/workflows/web-tests.yml index ba61f53df7d..c57f2cfb56b 100644 --- a/.github/workflows/web-tests.yml +++ b/.github/workflows/web-tests.yml @@ -9,10 +9,6 @@ on: permissions: contents: read -concurrency: - group: web-tests-${{ github.head_ref || github.run_id }} - cancel-in-progress: true - jobs: test: name: Web Tests (${{ matrix.shardIndex }}/${{ matrix.shardTotal }}) @@ -29,7 +25,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -62,7 +58,7 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -91,6 +87,7 @@ jobs: dify-ui-test: name: dify-ui Tests runs-on: depot-ubuntu-24.04-4 + timeout-minutes: 20 env: CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} defaults: @@ -100,15 +97,21 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false + - name: Switch Ubuntu mirror + run: | + sudo sed -i 's|http://us-east-1.ec2.archive.ubuntu.com/ubuntu|http://azure.archive.ubuntu.com/ubuntu|g' /etc/apt/sources.list.d/ubuntu.sources + sudo apt-get update + - name: Setup web environment uses: ./.github/actions/setup-web - name: Install Chromium for Browser Mode - run: vp exec playwright install --with-deps chromium + timeout-minutes: 15 + run: vp exec playwright install --with-deps --only-shell chromium - name: Run dify-ui tests run: vp test run --project unit --coverage --silent=passed-only @@ -125,6 +128,7 @@ jobs: dify-ui-storybook-test: name: dify-ui Storybook Tests runs-on: depot-ubuntu-24.04-4 + timeout-minutes: 20 defaults: run: shell: bash @@ -132,15 +136,21 @@ jobs: steps: - name: Checkout code - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false + - name: Switch Ubuntu mirror + run: | + sudo sed -i 's|http://us-east-1.ec2.archive.ubuntu.com/ubuntu|http://azure.archive.ubuntu.com/ubuntu|g' /etc/apt/sources.list.d/ubuntu.sources + sudo apt-get update + - name: Setup web environment uses: ./.github/actions/setup-web - name: Install Chromium for Browser Mode - run: vp exec playwright install --with-deps chromium + timeout-minutes: 15 + run: vp exec playwright install --with-deps --only-shell chromium - name: Run dify-ui Storybook tests run: vp run test:storybook diff --git a/.gitignore b/.gitignore index 54fdb704d8c..183645694eb 100644 --- a/.gitignore +++ b/.gitignore @@ -140,6 +140,9 @@ dmypy.json pyrightconfig.json !api/pyrightconfig.json +# import-linter +.import_linter_cache/ + # Pyre type checker .pyre/ .idea/' diff --git a/AGENTS.md b/AGENTS.md index 0da1d9e9eb6..6f5673b316a 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1,46 +1,9 @@ # AGENTS.md -## Project Overview +Dify is an open-source platform for building LLM applications, agentic workflows, and RAG pipelines. This monorepo contains the backend API (`api/`), frontend application (`web/`), deployment assets (`docker/`), standalone agent backend (`dify-agent/`), CLI (`cli/`), and end-to-end suite (`e2e/`). Follow the nearest scoped `AGENTS.md` for the files being changed. -Dify is an open-source platform for developing LLM applications with an intuitive interface combining agentic AI workflows, RAG pipelines, agent capabilities, and model management. +## Repository Gotchas -The codebase is split into: - -- **Backend API** (`/api`): Python Flask application organized with Domain-Driven Design -- **Frontend Web** (`/web`): Next.js application using TypeScript and React -- **Docker deployment** (`/docker`): Containerized deployment configurations -- **Dify Agent Backend** (`/dify-agent`): Backend services for managing and executing agent - -## Backend Workflow - -- Read `api/AGENTS.md` for details -- Run backend CLI commands through `uv run --project api `. -- Integration tests are CI-only and are not expected to run in the local environment. - -## Frontend Workflow - -- Read `web/AGENTS.md` for details - -## Testing & Quality Practices - -- Follow TDD: red → green → refactor. -- Use `pytest` for backend tests with Arrange-Act-Assert structure. -- Enforce strong typing; avoid `Any` and prefer explicit type annotations. -- Write self-documenting code; only add comments that explain intent. - -## Language Style - -- **Python**: Keep type hints on functions and attributes, and implement relevant special methods (e.g., `__repr__`, `__str__`). Prefer `TypedDict` over `dict` or `Mapping` for type safety and better code documentation. -- **TypeScript**: Use the strict config, run `pnpm check` for formatting, Oxlint, ESLint non-code checks, and type checking, and avoid `any` types. - -## General Practices - -- Prefer editing existing files; add new documentation only when requested. -- Inject dependencies through constructors and preserve clean architecture boundaries. -- Handle errors with domain-specific exceptions at the correct layer. - -## Project Conventions - -- Backend architecture adheres to DDD and Clean Architecture principles. -- Async work runs through Celery with Redis as the broker. -- Frontend user-facing strings must use `web/i18n/en-US/`; avoid hardcoded text. +- Run backend commands through `uv run --project api `. +- Backend integration tests are CI-only and are not expected to run locally. +- Keep `docker/.env.example` limited to variables required for a default Docker Compose deployment to start. Put optional and provider-specific settings in the matching `docker/envs/*.env.example` file; `docker/.env` overrides those service-specific env files. diff --git a/README.md b/README.md index 7688ee889bd..72539199e50 100644 --- a/README.md +++ b/README.md @@ -209,12 +209,11 @@ At the same time, please consider supporting Dify by sharing it on social media ## Star History - - + - - - Star History Chart + + + Star History Chart diff --git a/RESOURCE_BOUNDARY_CHANGE_GUIDE.md b/RESOURCE_BOUNDARY_CHANGE_GUIDE.md new file mode 100644 index 00000000000..b22386f8e5b --- /dev/null +++ b/RESOURCE_BOUNDARY_CHANGE_GUIDE.md @@ -0,0 +1,12 @@ +# Resource Boundary Change Guide + +- Resolve the tenant-scoped parent at the request boundary, then pass the validated model, owner reference, and actor downstream. +- Put the complete owner tuple in the database query; do not load by a bare ID and check ownership afterward. +- Treat missing and foreign-owned resources alike as `404` before locks, rate limits, tasks, plugin calls, network calls, or writes. +- Reuse existing owner resolvers and trusted objects instead of adding parallel helpers or refetching the same resource. +- Pass tenant, actor, and session explicitly; authenticated code must not depend on ambient account or tenant fallbacks. +- Raise typed domain errors in services and translate them to HTTP errors in controllers; reserve `ValueError` for invalid values or state. +- Let RBAC own authorization when enabled, and run legacy dataset permission checks only when RBAC is disabled. +- Preserve successful HTTP responses and shared runtime contracts, especially Celery task names and argument shapes during rolling upgrades. +- Keep runtime validation and OpenAPI schemas aligned, then regenerate Markdown and TypeScript contracts after schema changes. +- Prove the boundary with a foreign-owner decoy and assert that rejected requests trigger no downstream side effects. diff --git a/api/.env.example b/api/.env.example index 2adde29d334..9d0f17fe664 100644 --- a/api/.env.example +++ b/api/.env.example @@ -39,6 +39,9 @@ ENABLE_COLLABORATION_MODE=true # Learn app feature toggle ENABLE_LEARN_APP=true +# Show the license expiry countdown badge in the console (enterprise only) +ENABLE_LICENSE_EXPIRY_NOTICE=true + # Access token expiration time in minutes ACCESS_TOKEN_EXPIRE_MINUTES=60 @@ -117,6 +120,19 @@ SQLALCHEMY_POOL_RESET_ON_RETURN=rollback # storage type: opendal, s3, aliyun-oss, azure-blob, baidu-obs, google-storage, huawei-obs, oci-storage, tencent-cos, volcengine-tos, supabase STORAGE_TYPE=opendal +# Key provider configuration, used to encrypt/decrypt tenant credentials (LLM/tool provider secrets) +# key provider type: local, azure-keyvault +KEY_PROVIDER_TYPE=local + +# Azure Key Vault configuration, required when KEY_PROVIDER_TYPE=azure-keyvault +# authentication uses DefaultAzureCredential (managed identity, environment variables, or Azure CLI login) +AZURE_KEYVAULT_VAULT_URL=https://.vault.azure.net +AZURE_KEYVAULT_KEY_SIZE=2048 +# optional: auto-rotate each tenant's key every N days. Leave empty to manage rotation manually. +# rotation is safe because old key versions are never given an expiry and stay usable forever -- +# do NOT set a rotation policy on this key in the Azure portal that expires old versions. +AZURE_KEYVAULT_ROTATION_INTERVAL_DAYS= + # Apache OpenDAL storage configuration, refer to https://github.com/apache/opendal OPENDAL_SCHEME=fs OPENDAL_FS_ROOT=storage @@ -326,6 +342,7 @@ TIDB_VECTOR_PORT=4000 TIDB_VECTOR_USER=xxx.root TIDB_VECTOR_PASSWORD=xxxxxx TIDB_VECTOR_DATABASE=dify +TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH=false # Tidb on qdrant configuration TIDB_ON_QDRANT_URL=http://127.0.0.1 @@ -333,6 +350,7 @@ TIDB_ON_QDRANT_API_KEY=dify TIDB_ON_QDRANT_CLIENT_TIMEOUT=20 TIDB_ON_QDRANT_GRPC_ENABLED=false TIDB_ON_QDRANT_GRPC_PORT=6334 +TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB=sandbox:60,professional:6400,team:25600 TIDB_PUBLIC_KEY=dify TIDB_PRIVATE_KEY=dify TIDB_API_URL=http://127.0.0.1 @@ -432,10 +450,12 @@ OPENGAUSS_MAX_CONNECTION=5 # Upload configuration UPLOAD_FILE_SIZE_LIMIT=15 +KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN=15 UPLOAD_FILE_BATCH_LIMIT=5 UPLOAD_IMAGE_FILE_SIZE_LIMIT=10 UPLOAD_VIDEO_FILE_SIZE_LIMIT=100 UPLOAD_AUDIO_FILE_SIZE_LIMIT=50 +UPLOAD_SKILL_FILE_SIZE_LIMIT=50 # Comma-separated list of file extensions blocked from upload for security reasons. # Extensions should be lowercase without dots (e.g., exe,bat,sh,dll). @@ -468,6 +488,11 @@ SENDGRID_API_KEY= # Sentry configuration SENTRY_DSN= +# Cloudflare Turnstile server-side verification for Dify Cloud sign-in +TURNSTILE_SECRET_KEY= +# Comma-separated parent or exact hostnames, for example: dify.ai,staging.dify.dev +TURNSTILE_ALLOWED_HOSTNAMES= + # DEBUG DEBUG=false ENABLE_REQUEST_LOGGING=False @@ -569,8 +594,8 @@ WORKFLOW_GENERATOR_NODE_BUILDER_MAX_WORKERS=6 GRAPH_ENGINE_MIN_WORKERS=3 # Maximum number of workers per GraphEngine instance (default: 10) GRAPH_ENGINE_MAX_WORKERS=10 -# Queue depth threshold that triggers worker scale up (default: 3) -GRAPH_ENGINE_SCALE_UP_THRESHOLD=3 +# Pending task threshold that triggers worker scale up (default: 0) +GRAPH_ENGINE_SCALE_UP_THRESHOLD=0 # Seconds of idle time before scaling down workers (default: 5.0) GRAPH_ENGINE_SCALE_DOWN_IDLE_TIME=5.0 @@ -666,6 +691,7 @@ PLUGIN_REMOTE_INSTALL_PORT=5003 PLUGIN_REMOTE_INSTALL_HOST=localhost PLUGIN_MAX_PACKAGE_SIZE=15728640 PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600 +PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED=true PLUGIN_MODEL_PROVIDERS_CACHE_TTL=86400 # Comma-separated marketplace plugin IDs whose latest versions are installed for newly registered users. # Example: langgenius/openai,langgenius/gemini @@ -677,6 +703,8 @@ INNER_API_KEY_FOR_PLUGIN=QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y # Dify Agent backend AGENT_BACKEND_BASE_URL=http://localhost:5050 +# Bearer token sent to the Agent backend /runs API. Must match DIFY_AGENT_API_TOKEN on the server side. +AGENT_BACKEND_API_TOKEN=dify-agent-run-token-for-dev-only AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS=30 AGENT_BACKEND_STREAM_MAX_RECONNECTS=3 AGENT_BACKEND_RUN_TIMEOUT_SECONDS=1200 @@ -729,7 +757,6 @@ OTEL_MAX_EXPORT_BATCH_SIZE=512 OTEL_METRIC_EXPORT_INTERVAL=60000 OTEL_BATCH_EXPORT_TIMEOUT=10000 OTEL_METRIC_EXPORT_TIMEOUT=30000 - # Prevent Clickjacking ALLOW_EMBED=false diff --git a/api/.importlinter b/api/.importlinter index 5e06947d941..a64da0ce5b6 100644 --- a/api/.importlinter +++ b/api/.importlinter @@ -8,7 +8,148 @@ root_packages = extensions factories libs + machinery models + repositories tasks services include_external_packages = True + +[importlinter:contract:no-direct-rsa-imports] +# Note: `libs` itself is deliberately excluded from source_modules -- import-linter's +# `forbidden` contract cannot check a package against its own descendant (libs.rsa lives +# inside libs). Sibling modules under libs/ importing libs.rsa are not covered by this +# contract; the realistic risk this guards against is application/service code reaching +# past the key provider abstraction, which is fully covered below. +name = Only the key provider abstraction may import libs.rsa directly +type = forbidden +source_modules = + core + constants + context + configs + controllers + extensions + factories + models + tasks + services +forbidden_modules = + libs.rsa +allow_indirect_imports = True + +[importlinter:contract:machinery-framework-boundary] +name = API machinery is framework neutral +type = forbidden +source_modules = + machinery +forbidden_modules = + controllers + extensions + flask + models + sqlalchemy + werkzeug + +[importlinter:contract:workspace-query-service-boundary] +name = Workspace query application service is framework and persistence neutral +type = forbidden +source_modules = + services.workspace_query_service +forbidden_modules = + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + +[importlinter:contract:init-validation-service-boundary] +name = Initialization validation application service is framework and persistence neutral +type = forbidden +source_modules = + services.init_validation_service +forbidden_modules = + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + +[importlinter:contract:explore-banner-query-service-boundary] +name = Explore banner query application service is framework and persistence neutral +type = forbidden +source_modules = + services.explore_banner_query_service +forbidden_modules = + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + +[importlinter:contract:feature-query-service-boundary] +name = Feature query application service is framework and persistence neutral +type = forbidden +source_modules = + services.feature_query_service +forbidden_modules = + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + +[importlinter:contract:workspace-member-query-service-boundary] +name = Workspace member query application service is framework and persistence neutral +type = forbidden +source_modules = + services.workspace_member_query_service +forbidden_modules = + configs + controllers + extensions + flask + models + repositories + services.enterprise + services.workspace_member_role_resolver + sqlalchemy + werkzeug + +[importlinter:contract:setup-service-boundary] +name = Setup application service is framework and persistence neutral +type = forbidden +source_modules = + services.setup_service +forbidden_modules = + configs + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug + +[importlinter:contract:schema-definition-service-boundary] +name = Schema definition application service is framework and persistence neutral +type = forbidden +source_modules = + services.schema_definition_service +forbidden_modules = + configs + controllers + extensions + flask + models + repositories + sqlalchemy + werkzeug diff --git a/api/AGENTS.md b/api/AGENTS.md index 474da7800b2..ac77301f8dd 100644 --- a/api/AGENTS.md +++ b/api/AGENTS.md @@ -1,232 +1,25 @@ # API Agent Guide -## Notes for Agent (must-check) +Read surrounding module, class, and function docstrings plus non-obvious comments before changing backend behavior. They are local contracts; update them only when their owned behavior changes, and keep them aligned with the current code. -Before changing any backend code under `api/`, you MUST read the surrounding docstrings and comments. These notes contain required context (invariants, edge cases, trade-offs) and are treated as part of the spec. +## Commands -Look for: +Run backend checks from the repository root: -- The module (file) docstring at the top of a source code file -- Docstrings on classes and functions/methods -- Paragraph/block comments for non-obvious logic - -### What to write where - -- Keep notes scoped: module notes cover module-wide context, class notes cover class-wide context, function/method notes cover behavioural contracts, and paragraph/block comments cover local “why”. Avoid duplicating the same content across scopes unless repetition prevents misuse. -- **Module (file) docstring**: purpose, boundaries, key invariants, and “gotchas” that a new reader must know before editing. - - Include cross-links to the key collaborators (modules/services) when discovery is otherwise hard. - - Prefer stable facts (invariants, contracts) over ephemeral “today we…” notes. -- **Class docstring**: responsibility, lifecycle, invariants, and how it should be used (or not used). - - If the class is intentionally stateful, note what state exists and what methods mutate it. - - If concurrency/async assumptions matter, state them explicitly. -- **Function/method docstring**: behavioural contract. - - Document arguments, return shape, side effects (DB writes, external I/O, task dispatch), and raised domain exceptions. - - Add examples only when they prevent misuse. -- **Paragraph/block comments**: explain *why* (trade-offs, historical constraints, surprising edge cases), not what the code already states. - - Keep comments adjacent to the logic they justify; delete or rewrite comments that no longer match reality. - -### Rules (must follow) - -In this section, “notes” means module/class/function docstrings plus any relevant paragraph/block comments. - -- **Before working** - - Read the notes in the area you’ll touch; treat them as part of the spec. - - If a docstring or comment conflicts with the current code, treat the **code as the single source of truth** and update the docstring or comment to match reality. - - If important intent/invariants/edge cases are missing, add them in the closest docstring or comment (module for overall scope, function for behaviour). -- **During working** - - Keep the notes in sync as you discover constraints, make decisions, or change approach. - - If you move/rename responsibilities across modules/classes, update the affected docstrings and comments so readers can still find the “why” and the invariants. - - Record non-obvious edge cases, trade-offs, and the test/verification plan in the nearest docstring or comment that will stay correct. - - Keep the notes **coherent**: integrate new findings into the relevant docstrings and comments; avoid append-only “recent fix” / changelog-style additions. -- **When finishing** - - Update the notes to reflect what changed, why, and any new edge cases/tests. - - Remove or rewrite any comments that could be mistaken as current guidance but no longer apply. - - Keep docstrings and comments concise and accurate; they are meant to prevent repeated rediscovery. - -## Coding Style - -This is the default standard for backend code in this repo. Follow it for new code and use it as the checklist when reviewing changes. - -### Linting & Formatting - -- Use Ruff for formatting and linting (follow `.ruff.toml`). -- Keep each line under 120 characters (including spaces). - -### Naming Conventions - -- Use `snake_case` for variables and functions. -- Use `PascalCase` for classes. -- Use `UPPER_CASE` for constants. - -### Typing & Class Layout - -- Code should usually include type annotations that match the repo’s current Python version (avoid untyped public APIs and “mystery” values). -- Prefer modern typing forms (e.g. `list[str]`, `dict[str, int]`) and avoid `Any` unless there’s a strong reason. -- For dictionary-like data with known keys and value types, prefer `TypedDict` over `dict[...]` or `Mapping[...]`. -- For optional keys in typed payloads, use `NotRequired[...]` (or `total=False` when most fields are optional). -- Keep `dict[...]` / `Mapping[...]` for truly dynamic key spaces where the key set is unknown. - -```python -from datetime import datetime -from typing import NotRequired, TypedDict - - -class UserProfile(TypedDict): - user_id: str - email: str - created_at: datetime - nickname: NotRequired[str] -``` - -- For classes, declare all member variables explicitly with types at the top of the class body (before `__init__`), even when the class is not a dataclass or Pydantic model, so the class shape is obvious at a glance: - -```python -from datetime import datetime - - -class Example: - user_id: str - created_at: datetime - - def __init__(self, user_id: str, created_at: datetime) -> None: - self.user_id = user_id - self.created_at = created_at -``` - -### General Rules - -- Use Pydantic v2 conventions. -- Use `uv` for Python package management in this repo (usually with `--project api`). -- Prefer simple functions over small “utility classes” for lightweight helpers. -- Avoid implementing dunder methods unless it’s clearly needed and matches existing patterns. -- Never start long-running services as part of agent work (`uv run app.py`, `flask run`, etc.); running tests is allowed. -- Keep files below ~800 lines; split when necessary. -- Keep code readable and explicit—avoid clever hacks. - -### Architecture & Boundaries - -- Mirror the layered architecture: controller → service → core/domain. -- Reuse existing helpers in `core/`, `services/`, and `libs/` before creating new abstractions. -- Optimise for observability: deterministic control flow, clear logging, actionable errors. - -### Owner-Bound Resource References - -- Resolve and validate the outer owner before binding a nested resource ID. -- For stable single-parent chains, use immutable nested `NamedTuple` refs. -- Root refs carry tenant plus root ID; child refs carry the parent ref. -- In production, construct refs through the domain ref service. -- Python allowing direct construction does not grant authorization. -- Scope every consuming query with complete owner predicates; refs are not security tokens. -- Keep polymorphic owners flat until explicit nominal owner types exist. -- Do not add generic ref bases or compatibility fields only for uniformity. -- Reconstruct internal refs from validated database state after payload or async boundaries. - -### Logging & Errors - -- Never use `print`; use a module-level logger: - - `logger = logging.getLogger(__name__)` -- Include tenant/app/workflow identifiers in log context when relevant. -- Raise domain-specific exceptions (`services/errors`, `core/errors`) and translate them into HTTP responses in controllers. -- Log retryable events at `warning`, terminal failures at `error`. - -### SQLAlchemy Patterns - -- Models inherit from `models.base.TypeBase`; do not create ad-hoc metadata or engines. -- Open sessions with context managers: - -```python -from sqlalchemy.orm import Session - -with Session(db.engine, expire_on_commit=False) as session: - stmt = select(Workflow).where( - Workflow.id == workflow_id, - Workflow.tenant_id == tenant_id, - ) - workflow = session.execute(stmt).scalar_one_or_none() -``` - -- Prefer SQLAlchemy expressions; avoid raw SQL unless necessary. -- Always scope queries by `tenant_id` and protect write paths with safeguards (`FOR UPDATE`, row counts, etc.). -- Introduce repository abstractions only for very large tables (e.g., workflow executions) or when alternative storage strategies are required. - -### Storage & External I/O - -- Access storage via `extensions.ext_storage.storage`. -- Use `core.helper.ssrf_proxy` for outbound HTTP fetches. -- Background tasks that touch storage must be idempotent, and should log relevant object identifiers. - -### Pydantic Usage - -- Define DTOs with Pydantic v2 models and forbid extras by default. -- Use `@field_validator` / `@model_validator` for domain rules. - -Example: - -```python -from pydantic import BaseModel, ConfigDict, HttpUrl, field_validator - - -class TriggerConfig(BaseModel): - endpoint: HttpUrl - secret: str - - model_config = ConfigDict(extra="forbid") - - @field_validator("secret") - def ensure_secret_prefix(cls, value: str) -> str: - if not value.startswith("dify_"): - raise ValueError("secret must start with dify_") - return value -``` - -### Generics & Protocols - -- Use `typing.Protocol` to define behavioural contracts (e.g., cache interfaces). -- Apply generics (`TypeVar`, `Generic`) for reusable utilities like caches or providers. -- Validate dynamic inputs at runtime when generics cannot enforce safety alone. - -### Tooling & Checks - -Quick checks while iterating: - -- Format: `make format` -- Lint (includes auto-fix): `make lint` +- Format and lint: `make lint` - Type check: `make type-check` - Unit tests: `make test` -- Full backend tests, including Docker-backed suites: `make test-all` -- Targeted tests: `make test TARGET_TESTS=./api/tests/` +- Targeted tests: `make test TARGET_TESTS=./api/tests/` -Before opening a PR / submitting: +Run direct Python commands through `uv run --project api`. Docker-backed integration suites are normally CI-owned. Do not start long-running services as part of routine agent work. -- `make lint` -- `make type-check` -- `make test` +## Architecture And Boundaries -### Controllers & Services - -- Controllers: parse input via Pydantic, invoke services, return serialised responses; no business logic. -- Services: coordinate repositories, providers, background tasks; keep side effects explicit. -- Document non-obvious behaviour with concise docstrings and comments. -- For `204 No Content` responses, return an empty body only; never return a dict, model, or other payload. -- For Flask-RESTX controller request, query, and response schemas, follow `controllers/API_SCHEMA_GUIDE.md`. - In short: use Pydantic models, document GET query params with `query_params_from_model(...)`, register response - DTOs with `register_response_schema_models(...)`, serialize response DTOs with `dump_response(...)`, - and avoid adding new legacy `ns.model(...)`, `@marshal_with(...)`, or GET `@ns.expect(...)` patterns. - -### System Features Contract - -- Treat the shared Console/Web `/system-features` response as a minimal unauthenticated bootstrap allowlist, not a - general configuration or feature-discovery endpoint. Existing fields do not establish precedent. -- Before adding a field, read `controllers/API_SCHEMA_GUIDE.md#public-system-features-contract` and provide evidence - that both Console and Web have production consumers that require it before authentication. -- Never place backend-only policy, surface-specific configuration, post-authentication state, speculative values, or - large/slow payloads in `SystemFeatureModel`. Use the consumer or domain owner described in the schema guide. -- Agents and reviewers must reject additions whose owner, public exposure, pre-authentication need, or root SSR cost - is not explicit. - -### Miscellaneous - -- Use `configs.dify_config` for configuration—never read environment variables directly. -- Maintain tenant awareness end-to-end; `tenant_id` must flow through every layer touching shared resources. -- Queue async work through `services/async_workflow_service`; implement tasks under `tasks/` with explicit queue selection. -- Keep experimental scripts under `dev/`; do not ship them in production builds. +- Keep transport parsing and serialization in controllers, orchestration in services, and domain policy in `core/` or its domain owner. Keep `libs/` business-agnostic and reuse existing owners before adding abstractions. +- Before changing controller schemas, generated API contracts, or `SystemFeatureModel`, read `controllers/API_SCHEMA_GUIDE.md`. Treat `/system-features` as a minimal unauthenticated bootstrap allowlist, not a general configuration registry. +- Scope tenant-owned reads and writes by the complete owner chain, and propagate `tenant_id` across every affected layer. Reconstruct trusted internal references from validated database state after payload or async boundaries. +- Keep write transactions explicit and bounded. Do not perform external I/O inside an open transaction unless a documented consistency contract requires it. +- Read configuration through `configs.dify_config`, access storage through `extensions.ext_storage.storage`, and route outbound HTTP through the existing SSRF-safe owner in `core.helper.ssrf_proxy`. +- Use Pydantic v2 for request and response models. Reuse domain-specific exceptions and translate them at the controller boundary. +- Use existing Celery task and queue owners for asynchronous work; do not route unrelated jobs through workflow-specific services. +- Celery tasks that may be retried or redelivered must keep side effects idempotent and log affected resource identifiers. diff --git a/api/Dockerfile b/api/Dockerfile index 1823a8f6a35..86eef5d329f 100644 --- a/api/Dockerfile +++ b/api/Dockerfile @@ -33,7 +33,6 @@ RUN uv sync --frozen --no-dev --no-editable FROM base AS production ENV FLASK_APP=app.py -ENV EDITION=SELF_HOSTED ENV DEPLOY_ENV=PRODUCTION ENV CONSOLE_API_URL=http://127.0.0.1:5001 ENV CONSOLE_WEB_URL=http://127.0.0.1:3000 @@ -99,9 +98,9 @@ ENV VIRTUAL_ENV=/app/api/.venv COPY --from=packages --chown=dify:dify ${VIRTUAL_ENV} ${VIRTUAL_ENV} ENV PATH="${VIRTUAL_ENV}/bin:${PATH}" -# Download nltk data RUN mkdir -p /usr/local/share/nltk_data \ - && NLTK_DATA=/usr/local/share/nltk_data python -c "import nltk; nltk.download('punkt'); nltk.download('averaged_perceptron_tagger'); nltk.download('stopwords')" \ + && NLTK_DATA=/usr/local/share/nltk_data python -m nltk.downloader punkt_tab averaged_perceptron_tagger_eng stopwords \ + && NLTK_DATA=/usr/local/share/nltk_data python -c "import nltk; nltk.data.find('tokenizers/punkt_tab'); nltk.data.find('taggers/averaged_perceptron_tagger_eng'); nltk.data.find('corpora/stopwords')" \ && chmod -R 755 /usr/local/share/nltk_data ENV TIKTOKEN_CACHE_DIR=/app/api/.tiktoken_cache diff --git a/api/app_factory.py b/api/app_factory.py index 2cea8cfb3f7..2a08aeeed3e 100644 --- a/api/app_factory.py +++ b/api/app_factory.py @@ -1,19 +1,23 @@ import logging import time +from collections.abc import Callable +from typing import NamedTuple import socketio from flask import request from opentelemetry.trace import get_current_span from opentelemetry.trace.span import INVALID_SPAN_ID, INVALID_TRACE_ID +from werkzeug.exceptions import Forbidden, HTTPException, ServiceUnavailable from configs import dify_config from contexts.wrapper import RecyclableContextVar from controllers.console.error import UnauthorizedAndForceLogout from core.logging.context import init_request_context from dify_app import DifyApp +from enums import DeploymentEdition from extensions.ext_socketio import sio from services.enterprise.enterprise_service import EnterpriseService -from services.feature_service import LicenseStatus +from services.entities.feature_entities import LicenseStatus logger = logging.getLogger(__name__) @@ -42,6 +46,53 @@ _CONSOLE_EXEMPT_PREFIXES = ( "/console/api/activate/check", ) +_WEBAPP_EXEMPT_PREFIXES = ("/api/system-features",) + +_INVALID_LICENSE_STATUSES = (LicenseStatus.INACTIVE, LicenseStatus.EXPIRED, LicenseStatus.LOST) + + +def _session_surface_error(license_status: LicenseStatus | None) -> HTTPException: + if license_status is None: + return UnauthorizedAndForceLogout("Unable to verify enterprise license. Please contact your administrator.") + return UnauthorizedAndForceLogout(f"Enterprise license is {license_status}. Please contact your administrator.") + + +def _bearer_surface_error(license_status: LicenseStatus | None) -> HTTPException: + """Token-authed: forcing a logout is meaningless and license state must not leak.""" + return Forbidden(description="license_required") + + +def _retryable_surface_error(license_status: LicenseStatus | None) -> HTTPException: + """Webhook senders retry on 5xx but treat 4xx as permanent, disabling the subscription.""" + return ServiceUnavailable(description="license_required") + + +class _LicenseGatedSurface(NamedTuple): + prefix: str + exempt_prefixes: tuple[str, ...] + build_error: Callable[[LicenseStatus | None], HTTPException] + + +# /files (plugin-daemon data plane), /inner/api (enterprise control plane) and /health +# stay ungated: blocking them breaks workflow execution or license recovery itself. +_LICENSE_GATED_SURFACES = ( + _LicenseGatedSurface("/console/api/", _CONSOLE_EXEMPT_PREFIXES, _session_surface_error), + _LicenseGatedSurface("/api/", _WEBAPP_EXEMPT_PREFIXES, _session_surface_error), + _LicenseGatedSurface("/v1", (), _bearer_surface_error), + _LicenseGatedSurface("/mcp", (), _bearer_surface_error), + _LicenseGatedSurface("/triggers", (), _retryable_surface_error), +) + + +def _match_license_gated_surface(path: str) -> _LicenseGatedSurface | None: + for surface in _LICENSE_GATED_SURFACES: + if not path.startswith(surface.prefix): + continue + if any(path.startswith(exempt) for exempt in surface.exempt_prefixes): + return None + return surface + return None + # ---------------------------- # Application Factory Function @@ -62,38 +113,17 @@ def create_flask_app_with_configs() -> DifyApp: init_request_context() RecyclableContextVar.increment_thread_recycles() - # Enterprise license validation for API endpoints (both console and webapp) - # When license expires, block all API access except bootstrap endpoints needed - # for the frontend to load the license expiration page without infinite reloads. - if dify_config.ENTERPRISE_ENABLED: - is_console_api = request.path.startswith("/console/api/") - is_webapp_api = request.path.startswith("/api/") + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE: + surface = _match_license_gated_surface(request.path) + if surface is not None: + try: + license_status = EnterpriseService.get_cached_license_status() + except Exception: + logger.exception("Failed to check enterprise license status") + license_status = None - if is_console_api or is_webapp_api: - if is_console_api: - is_exempt = any(request.path.startswith(p) for p in _CONSOLE_EXEMPT_PREFIXES) - else: # webapp API - is_exempt = request.path.startswith("/api/system-features") - - if not is_exempt: - try: - # Check license status (cached — see EnterpriseService for TTL details) - license_status = EnterpriseService.get_cached_license_status() - if license_status in (LicenseStatus.INACTIVE, LicenseStatus.EXPIRED, LicenseStatus.LOST): - raise UnauthorizedAndForceLogout( - f"Enterprise license is {license_status}. Please contact your administrator." - ) - if license_status is None: - raise UnauthorizedAndForceLogout( - "Unable to verify enterprise license. Please contact your administrator." - ) - except UnauthorizedAndForceLogout: - raise - except Exception: - logger.exception("Failed to check enterprise license status") - raise UnauthorizedAndForceLogout( - "Unable to verify enterprise license. Please contact your administrator." - ) + if license_status is None or license_status in _INVALID_LICENSE_STATUSES: + raise surface.build_error(license_status) # add after request hook for injecting trace headers from OpenTelemetry span context # Only adds headers when OTEL is enabled and has valid context @@ -143,6 +173,7 @@ def initialize_extensions(app: DifyApp): from context.flask_app_context import init_flask_context from extensions import ( ext_app_metrics, + ext_application_services, ext_blueprints, ext_celery, ext_code_based_extension, @@ -154,6 +185,7 @@ def initialize_extensions(app: DifyApp): ext_forward_refs, ext_hosting_provider, ext_import_modules, + ext_key_provider, ext_logging, ext_login, ext_logstore, @@ -189,6 +221,7 @@ def initialize_extensions(app: DifyApp): ext_migrate, ext_redis, ext_storage, + ext_key_provider, # Initialize after storage, since RSAKeyProvider reads private keys from it ext_set_secretkey, ext_logstore, # Initialize logstore after storage, before celery ext_celery, @@ -204,6 +237,7 @@ def initialize_extensions(app: DifyApp): ext_enterprise_telemetry, ext_request_logging, ext_session_factory, + ext_application_services, ext_oauth_bearer, ] for ext in extensions: diff --git a/api/clients/agent_backend/__init__.py b/api/clients/agent_backend/__init__.py index 67175c4795c..dbdc93cfbad 100644 --- a/api/clients/agent_backend/__init__.py +++ b/api/clients/agent_backend/__init__.py @@ -5,8 +5,6 @@ API adapters: request building from Dify product concepts, a thin client wrapper event adaptation for future workflow integration, and deterministic fakes. """ -from dify_agent.protocol import RuntimeLayerSpec, extract_runtime_layer_specs - from clients.agent_backend.client import AgentBackendRunClient, DifyAgentBackendRunClient from clients.agent_backend.errors import ( AgentBackendError, @@ -47,11 +45,6 @@ from clients.agent_backend.request_builder import ( AgentBackendWorkflowNodeRunInput, redact_for_agent_backend_log, ) -from clients.agent_backend.session_cleanup import ( - AgentBackendSessionCleanupPayload, - AgentBackendSessionCleanupResult, - cleanup_agent_backend_session, -) __all__ = [ "AGENT_SOUL_PROMPT_LAYER_ID", @@ -80,8 +73,6 @@ __all__ = [ "AgentBackendRunRequestBuilder", "AgentBackendRunStartedInternalEvent", "AgentBackendRunSucceededInternalEvent", - "AgentBackendSessionCleanupPayload", - "AgentBackendSessionCleanupResult", "AgentBackendStreamError", "AgentBackendStreamInternalEvent", "AgentBackendTransportError", @@ -90,9 +81,6 @@ __all__ = [ "DifyAgentBackendRunClient", "FakeAgentBackendRunClient", "FakeAgentBackendScenario", - "RuntimeLayerSpec", - "cleanup_agent_backend_session", "create_agent_backend_run_client", - "extract_runtime_layer_specs", "redact_for_agent_backend_log", ] diff --git a/api/clients/agent_backend/errors.py b/api/clients/agent_backend/errors.py index b48c4d08da1..99ef2fecf5d 100644 --- a/api/clients/agent_backend/errors.py +++ b/api/clients/agent_backend/errors.py @@ -10,6 +10,8 @@ from __future__ import annotations from typing import Any +from dify_agent.protocol import RunFailureType + class AgentBackendError(Exception): """Base error for API-side Agent backend integration failures.""" @@ -54,6 +56,7 @@ class AgentBackendRunFailedError(AgentBackendError): run_id: str detail: Any + error_type: RunFailureType | None reason: str | None source_event_id: str | None @@ -63,11 +66,13 @@ class AgentBackendRunFailedError(AgentBackendError): detail: Any, *, message: str | None = None, + error_type: RunFailureType | None = None, reason: str | None = None, source_event_id: str | None = None, ) -> None: self.run_id = run_id self.detail = detail + self.error_type = error_type self.reason = reason self.source_event_id = source_event_id display_message = message or f"Agent backend run failed: {run_id}" diff --git a/api/clients/agent_backend/event_adapter.py b/api/clients/agent_backend/event_adapter.py index 8fdc165ab3c..e72d2fbfb51 100644 --- a/api/clients/agent_backend/event_adapter.py +++ b/api/clients/agent_backend/event_adapter.py @@ -22,6 +22,7 @@ from dify_agent.protocol import ( RunCancelledEvent, RunEvent, RunFailedEvent, + RunFailureType, RunStartedEvent, RunSucceededEvent, ) @@ -96,6 +97,7 @@ class AgentBackendRunFailedInternalEvent(AgentBackendInternalEventBase): type: Literal[AgentBackendInternalEventType.RUN_FAILED] = AgentBackendInternalEventType.RUN_FAILED error: str + error_type: RunFailureType | None = None reason: str | None = None @@ -180,6 +182,7 @@ class AgentBackendRunEventAdapter: run_id=event.run_id, source_event_id=event.id, error=event.data.error, + error_type=event.data.error_type, reason=event.data.reason, ) ] diff --git a/api/clients/agent_backend/factory.py b/api/clients/agent_backend/factory.py index 0fcbf02bf70..2fd9c6faf4a 100644 --- a/api/clients/agent_backend/factory.py +++ b/api/clients/agent_backend/factory.py @@ -11,6 +11,7 @@ from clients.agent_backend.fake_client import FakeAgentBackendRunClient, FakeAge def create_agent_backend_run_client( *, base_url: str | None = None, + api_token: str | None = None, use_fake: bool = False, fake_scenario: str | FakeAgentBackendScenario = FakeAgentBackendScenario.SUCCESS, stream_read_timeout_seconds: float = 30, @@ -22,8 +23,11 @@ def create_agent_backend_run_client( return FakeAgentBackendRunClient(scenario=FakeAgentBackendScenario(fake_scenario)) if base_url is None: raise ValueError("base_url is required when creating a real Agent backend client") + headers: dict[str, str] = {} + if api_token: + headers["Authorization"] = f"Bearer {api_token}" return DifyAgentBackendRunClient( - Client(base_url=base_url, stream_timeout=stream_read_timeout_seconds), + Client(base_url=base_url, stream_timeout=stream_read_timeout_seconds, headers=headers), stream_max_reconnects=stream_max_reconnects, stream_timeout_seconds=stream_run_timeout_seconds, ) diff --git a/api/clients/agent_backend/request_builder.py b/api/clients/agent_backend/request_builder.py index 1fd7b2cd990..61d74eb4a4b 100644 --- a/api/clients/agent_backend/request_builder.py +++ b/api/clients/agent_backend/request_builder.py @@ -16,7 +16,6 @@ from collections.abc import Mapping from typing import ClassVar, Literal from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot from agenton.layers import ExitIntent from agenton_collections.layers.plain import PLAIN_PROMPT_LAYER_TYPE_ID, PromptLayerConfig from agenton_collections.layers.pydantic_ai import PYDANTIC_AI_HISTORY_LAYER_TYPE_ID @@ -37,6 +36,7 @@ from dify_agent.layers.execution_context import ( ) from dify_agent.layers.knowledge import DIFY_KNOWLEDGE_BASE_LAYER_TYPE_ID, DifyKnowledgeBaseLayerConfig from dify_agent.layers.output import DIFY_OUTPUT_LAYER_TYPE_ID, DifyOutputLayerConfig +from dify_agent.layers.runtime import DIFY_RUNTIME_LAYER_TYPE_ID, DifyRuntimeLayerConfig from dify_agent.layers.shell import DIFY_SHELL_LAYER_TYPE_ID, DifyShellLayerConfig from dify_agent.protocol import ( DIFY_AGENT_HISTORY_LAYER_ID, @@ -47,7 +47,6 @@ from dify_agent.protocol import ( LayerExitSignals, RunComposition, RunLayerSpec, - RuntimeLayerSpec, ) from pydantic import BaseModel, ConfigDict, Field, JsonValue, field_validator @@ -56,6 +55,7 @@ WORKFLOW_NODE_JOB_PROMPT_LAYER_ID = "workflow_node_job_prompt" WORKFLOW_USER_PROMPT_LAYER_ID = "workflow_user_prompt" AGENT_APP_USER_PROMPT_LAYER_ID = "agent_app_user_prompt" DIFY_EXECUTION_CONTEXT_LAYER_ID = "execution_context" +DIFY_RUNTIME_LAYER_ID = "runtime" DIFY_CONFIG_LAYER_ID = "config" DIFY_DRIVE_LAYER_ID = "drive" DIFY_PLUGIN_TOOLS_LAYER_ID = "tools" @@ -66,25 +66,11 @@ DIFY_SHELL_LAYER_ID = "shell" type AgentConfigVersionKind = Literal["snapshot", "draft", "build_draft"] -def _filter_snapshot_to_specs( - snapshot: CompositorSessionSnapshot, - specs: list[RuntimeLayerSpec], -) -> CompositorSessionSnapshot: - """Keep only snapshot layers whose names appear in the cleanup spec list. - - The agenton compositor rejects a snapshot whose layer-name sequence does - not match the active composition exactly. Cleanup-replay drops plugin - layers, so we must drop the matching snapshot entries here. - """ - kept_names = {spec.name for spec in specs} - filtered_layers: list[LayerSessionSnapshot] = [layer for layer in snapshot.layers if layer.name in kept_names] - if len(filtered_layers) == len(snapshot.layers): - return snapshot - return CompositorSessionSnapshot(schema_version=snapshot.schema_version, layers=filtered_layers) - - def _shell_layer_deps() -> dict[str, str]: - return {"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID} + return { + "execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID, + "runtime": DIFY_RUNTIME_LAYER_ID, + } def _drive_layer_deps() -> dict[str, str]: @@ -173,8 +159,12 @@ class AgentBackendModelConfig(BaseModel): # ``DifyPluginLLMLayerConfig.model_settings`` is pydantic_ai's ``ModelSettings`` # TypedDict (closed: unknown keys are rejected, explicit ``None`` values fail the # per-field type checks). Agent Soul model settings carry a wider, nullable shape -# (``stop`` / ``response_format`` plus null-padded fields), so the layer config -# only receives the keys the runtime contract accepts. +# (``stop`` / ``response_format`` plus null-padded fields, plus arbitrary +# plugin-declared parameters such as Qwen's ``enable_thinking``), so the layer +# config only receives the keys the runtime contract accepts directly; anything +# else is forwarded through ``extra_body``, the TypedDict's own escape hatch for +# provider-specific parameters (see +# ``dify_agent.adapters.llm.model._map_model_settings_to_parameters``). _AGENT_MODEL_SETTINGS_PASSTHROUGH_KEYS = ( "temperature", "top_p", @@ -182,6 +172,7 @@ _AGENT_MODEL_SETTINGS_PASSTHROUGH_KEYS = ( "frequency_penalty", "max_tokens", ) +_AGENT_MODEL_SETTINGS_KNOWN_KEYS = frozenset({*_AGENT_MODEL_SETTINGS_PASSTHROUGH_KEYS, "stop", "response_format"}) def _agent_model_settings(settings: Mapping[str, JsonValue]) -> dict[str, JsonValue] | None: @@ -191,6 +182,15 @@ def _agent_model_settings(settings: Mapping[str, JsonValue]) -> dict[str, JsonVa stop = settings.get("stop") if isinstance(stop, list) and stop: sanitized["stop_sequences"] = stop + + extra_body: dict[str, JsonValue] = { + key: value + for key, value in settings.items() + if key not in _AGENT_MODEL_SETTINGS_KNOWN_KEYS and value is not None + } + if extra_body: + sanitized["extra_body"] = extra_body + return sanitized or None @@ -214,6 +214,7 @@ class AgentBackendWorkflowNodeRunInput(BaseModel): model: AgentBackendModelConfig execution_context: DifyExecutionContextLayerConfig + backend_binding_ref: str = Field(min_length=1) workflow_node_job_prompt: str user_prompt: str agent_soul_prompt: str | None = None @@ -231,8 +232,8 @@ class AgentBackendWorkflowNodeRunInput(BaseModel): # the Agent Soul configures human involvement; a deferred call ends the run and # the workflow pauses via the existing HITL form mechanism (ENG-635). ask_human_config: DifyAskHumanLayerConfig | None = None - # Inject the sandboxed shell layer (dify.shell). Requires the agent backend - # to be wired with a shellctl entrypoint; see configs AGENT_SHELL_ENABLED. + # Inject the sandboxed shell graph. Requires a deployment-selected runtime + # backend plus the product-resolved persistent Binding. include_shell: bool = False shell_config: DifyShellLayerConfig | None = None session_snapshot: CompositorSessionSnapshot | None = None @@ -240,7 +241,6 @@ class AgentBackendWorkflowNodeRunInput(BaseModel): # (ENG-638). Keyed by the original deferred tool_call_id. deferred_tool_results: DeferredToolResultsPayload | None = None include_history: bool = True - suspend_on_exit: bool = True metadata: dict[str, JsonValue] = Field(default_factory=dict) model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid", arbitrary_types_allowed=True) @@ -264,6 +264,7 @@ class AgentBackendAgentAppRunInput(BaseModel): model: AgentBackendModelConfig execution_context: DifyExecutionContextLayerConfig + backend_binding_ref: str = Field(min_length=1) user_prompt: str agent_soul_prompt: str | None = None agent_config_version_kind: AgentConfigVersionKind = "snapshot" @@ -279,8 +280,8 @@ class AgentBackendAgentAppRunInput(BaseModel): # Human-in-the-loop ask_human deferred tool (dify.ask_human). Present only when # the Agent Soul configures human involvement (ENG-635). ask_human_config: DifyAskHumanLayerConfig | None = None - # Inject the sandboxed shell layer (dify.shell). Requires the agent backend - # to be wired with a shellctl entrypoint; see configs AGENT_SHELL_ENABLED. + # Inject the sandboxed shell graph. Requires a deployment-selected runtime + # backend plus the product-resolved persistent Binding. include_shell: bool = False shell_config: DifyShellLayerConfig | None = None session_snapshot: CompositorSessionSnapshot | None = None @@ -288,7 +289,6 @@ class AgentBackendAgentAppRunInput(BaseModel): # (ENG-638). Keyed by the original deferred tool_call_id. deferred_tool_results: DeferredToolResultsPayload | None = None include_history: bool = True - suspend_on_exit: bool = True metadata: dict[str, JsonValue] = Field(default_factory=dict) model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid", arbitrary_types_allowed=True) @@ -350,6 +350,14 @@ class AgentBackendRunRequestBuilder: run_input.include_shell or run_input.config_layer_config is not None or run_input.drive_config is not None ) if include_shell: + layers.append( + RunLayerSpec( + name=DIFY_RUNTIME_LAYER_ID, + type=DIFY_RUNTIME_LAYER_TYPE_ID, + metadata=run_input.metadata, + config=DifyRuntimeLayerConfig(backend_binding_ref=run_input.backend_binding_ref), + ) + ) # Sandboxed bash workspace (dify.shell). It enters before config/drive # so eager pulls materialize content in the same filesystem used by # model commands. @@ -481,53 +489,7 @@ class AgentBackendRunRequestBuilder: metadata=run_input.metadata, session_snapshot=run_input.session_snapshot, deferred_tool_results=run_input.deferred_tool_results, - on_exit=LayerExitSignals( - default=ExitIntent.SUSPEND if run_input.suspend_on_exit else ExitIntent.DELETE, - ), - ) - - def build_cleanup_request( - self, - *, - session_snapshot: CompositorSessionSnapshot, - runtime_layer_specs: list[RuntimeLayerSpec], - idempotency_key: str | None = None, - metadata: dict[str, JsonValue] | None = None, - ) -> CreateRunRequest: - """Build a lifecycle-only cleanup request that replays the prior layers. - - The agenton compositor enforces that the session snapshot's layer names - match the active composition in order, so cleanup must replay the same - non-plugin layer graph that produced the snapshot. Plugin layers - (``dify.plugin.llm``, ``dify.plugin.tools``) are excluded from both the - composition and the snapshot before submission because their configs - may carry credentials or runtime-only declarations that are not - persisted between runs. - """ - if not runtime_layer_specs: - raise ValueError( - "build_cleanup_request requires runtime_layer_specs; an empty " - "composition would fail the agent backend's snapshot validation." - ) - request_metadata = dict(metadata or {}) - request_metadata["agent_backend_lifecycle"] = "session_cleanup" - layers = [ - RunLayerSpec( - name=spec.name, - type=spec.type, - deps=dict(spec.deps), - metadata=dict(spec.metadata), - config=spec.config, - ) - for spec in runtime_layer_specs - ] - filtered_snapshot = _filter_snapshot_to_specs(session_snapshot, runtime_layer_specs) - return CreateRunRequest( - composition=RunComposition(layers=layers), - idempotency_key=idempotency_key, - metadata=request_metadata, - session_snapshot=filtered_snapshot, - on_exit=LayerExitSignals(default=ExitIntent.DELETE), + on_exit=LayerExitSignals(default=ExitIntent.SUSPEND), ) def build_for_workflow_node(self, run_input: AgentBackendWorkflowNodeRunInput) -> CreateRunRequest: @@ -580,6 +542,14 @@ class AgentBackendRunRequestBuilder: run_input.include_shell or run_input.config_layer_config is not None or run_input.drive_config is not None ) if include_shell: + layers.append( + RunLayerSpec( + name=DIFY_RUNTIME_LAYER_ID, + type=DIFY_RUNTIME_LAYER_TYPE_ID, + metadata=run_input.metadata, + config=DifyRuntimeLayerConfig(backend_binding_ref=run_input.backend_binding_ref), + ) + ) # Sandboxed bash workspace (dify.shell). It enters before drive so # drive can materialize mentioned targets with `dify-agent drive pull` # in the same shell-visible filesystem used by model commands. @@ -713,9 +683,7 @@ class AgentBackendRunRequestBuilder: metadata=run_input.metadata, session_snapshot=run_input.session_snapshot, deferred_tool_results=run_input.deferred_tool_results, - on_exit=LayerExitSignals( - default=ExitIntent.SUSPEND if run_input.suspend_on_exit else ExitIntent.DELETE, - ), + on_exit=LayerExitSignals(default=ExitIntent.SUSPEND), ) diff --git a/api/clients/agent_backend/session_cleanup.py b/api/clients/agent_backend/session_cleanup.py deleted file mode 100644 index 370ff544e68..00000000000 --- a/api/clients/agent_backend/session_cleanup.py +++ /dev/null @@ -1,100 +0,0 @@ -"""Shared API-side helper for Agent backend lifecycle-only session cleanup. - -Product code owns local row retirement and background-task dispatch. This module -only adapts persisted cleanup inputs into the public ``dify-agent`` run -protocol, performs the synchronous ``create_run + wait_run`` loop used by Celery -workers, and reports whether the backend cleanup succeeded, was skipped, or -failed. -""" - -from __future__ import annotations - -from dataclasses import dataclass -from typing import ClassVar, Literal - -from agenton.compositor import CompositorSessionSnapshot -from dify_agent.protocol import RuntimeLayerSpec -from pydantic import BaseModel, ConfigDict, Field, JsonValue - -from clients.agent_backend.client import AgentBackendRunClient -from clients.agent_backend.errors import AgentBackendError -from clients.agent_backend.request_builder import AgentBackendRunRequestBuilder - - -class AgentBackendSessionCleanupPayload(BaseModel): - """Serialized cleanup inputs preserved across API and Celery boundaries.""" - - session_snapshot: CompositorSessionSnapshot | None = None - runtime_layer_specs: list[RuntimeLayerSpec] = Field(default_factory=list) - idempotency_key: str | None = None - metadata: dict[str, JsonValue] = Field(default_factory=dict) - timeout_seconds: float = 30.0 - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -@dataclass(frozen=True, slots=True) -class AgentBackendSessionCleanupResult: - """Terminal outcome of one backend cleanup attempt.""" - - status: Literal["succeeded", "skipped", "failed"] - reason: str | None = None - cleanup_run_id: str | None = None - - @classmethod - def succeeded(cls, cleanup_run_id: str) -> AgentBackendSessionCleanupResult: - return cls(status="succeeded", cleanup_run_id=cleanup_run_id) - - @classmethod - def skipped(cls, reason: str) -> AgentBackendSessionCleanupResult: - return cls(status="skipped", reason=reason) - - @classmethod - def failed(cls, reason: str, cleanup_run_id: str | None = None) -> AgentBackendSessionCleanupResult: - return cls(status="failed", reason=reason, cleanup_run_id=cleanup_run_id) - - -def cleanup_agent_backend_session( - *, - payload: AgentBackendSessionCleanupPayload, - client: AgentBackendRunClient | None, - request_builder: AgentBackendRunRequestBuilder | None = None, -) -> AgentBackendSessionCleanupResult: - """Run lifecycle-only cleanup against the Agent backend and report status.""" - if client is None: - return AgentBackendSessionCleanupResult.skipped("no_agent_backend_client") - if payload.session_snapshot is None: - return AgentBackendSessionCleanupResult.skipped("missing_session_snapshot") - if not payload.runtime_layer_specs: - return AgentBackendSessionCleanupResult.skipped("missing_runtime_layer_specs") - - builder = request_builder or AgentBackendRunRequestBuilder() - request = builder.build_cleanup_request( - session_snapshot=payload.session_snapshot, - runtime_layer_specs=payload.runtime_layer_specs, - idempotency_key=payload.idempotency_key, - metadata=payload.metadata, - ) - - try: - response = client.create_run(request) - except AgentBackendError as exc: - return AgentBackendSessionCleanupResult.failed(str(exc)) - - try: - status_response = client.wait_run(response.run_id, timeout_seconds=payload.timeout_seconds) - except AgentBackendError as exc: - return AgentBackendSessionCleanupResult.failed(str(exc), cleanup_run_id=response.run_id) - - if status_response.status != "succeeded": - reason = status_response.error or f"cleanup run ended with status {status_response.status}" - return AgentBackendSessionCleanupResult.failed(reason, cleanup_run_id=response.run_id) - - return AgentBackendSessionCleanupResult.succeeded(response.run_id) - - -__all__ = [ - "AgentBackendSessionCleanupPayload", - "AgentBackendSessionCleanupResult", - "cleanup_agent_backend_session", -] diff --git a/api/commands/account.py b/api/commands/account.py index 9ea52dfd248..6751361a1a0 100644 --- a/api/commands/account.py +++ b/api/commands/account.py @@ -33,7 +33,7 @@ def reset_password(email, new_password, password_confirm): try: valid_password(new_password) - except: + except ValueError: click.echo(click.style(f"Invalid password. Must match {password_pattern}", fg="red")) return @@ -75,7 +75,7 @@ def reset_email(email, new_email, email_confirm): try: email_validate(normalized_new_email) - except: + except ValueError: click.echo(click.style(f"Invalid email: {new_email}", fg="red")) return diff --git a/api/commands/retention.py b/api/commands/retention.py index d03c9bcc6da..f94dbf25bb6 100644 --- a/api/commands/retention.py +++ b/api/commands/retention.py @@ -10,6 +10,8 @@ import click import sqlalchemy as sa from sqlalchemy.orm import Session, sessionmaker +from configs import dify_config +from enums import CloudPlan, DeploymentEdition from extensions.ext_database import db from libs.datetime_utils import naive_utc_now from services.clear_free_plan_tenant_expired_logs import ClearFreePlanTenantExpiredLogs @@ -126,14 +128,12 @@ def _get_archive_candidate_tenant_ids_by_prefix( def _filter_paid_workflow_archive_tenant_ids(tenant_ids: list[str]) -> tuple[list[str], list[str]]: - from configs import dify_config - from enums.cloud_plan import CloudPlan from services.billing_service import BillingService tenant_ids = sorted(set(tenant_ids)) if not tenant_ids: return [], [] - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return tenant_ids, [] plans = BillingService.get_plan_bulk_with_cache(tenant_ids) @@ -1462,7 +1462,7 @@ def cleanup_orphaned_draft_variables( "--graceful-period", default=21, show_default=True, - help="Graceful period in days after subscription expiration, will be ignored when billing is disabled.", + help="Graceful period in days after subscription expiration; ignored outside the Cloud edition.", ) @click.option("--dry-run", is_flag=True, default=False, help="Show messages logs would be cleaned without deleting") def clean_expired_messages( @@ -1514,8 +1514,8 @@ def clean_expired_messages( if from_days_ago <= before_days: raise click.UsageError("--from-days-ago must be greater than --before-days.") - # Create policy based on billing configuration - # NOTE: graceful_period will be ignored when billing is disabled. + # Create the policy for the configured deployment edition. + # NOTE: graceful_period is ignored outside the Cloud edition. policy = create_message_clean_policy(graceful_period_days=graceful_period) if from_days_ago is not None and before_days is not None: diff --git a/api/commands/system.py b/api/commands/system.py index 699e0c199c9..2dee41e4d3f 100644 --- a/api/commands/system.py +++ b/api/commands/system.py @@ -6,12 +6,12 @@ from sqlalchemy import delete, select, update from sqlalchemy.orm import sessionmaker from configs import dify_config -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from events.app_event import app_was_created from extensions.ext_database import db from extensions.ext_redis import redis_client from libs.db_migration_lock import DbMigrationAutoRenewLock -from libs.rsa import generate_key_pair +from libs.key_providers import generate_key_pair from models import Tenant from models.model import App, AppMode, Conversation from models.provider import Provider, ProviderModel diff --git a/api/configs/app_config.py b/api/configs/app_config.py index be2f3c7c0e5..29828a3be0d 100644 --- a/api/configs/app_config.py +++ b/api/configs/app_config.py @@ -5,7 +5,6 @@ from typing import Any, override from pydantic.fields import FieldInfo from pydantic_settings import BaseSettings, PydanticBaseSettingsSource, SettingsConfigDict, TomlConfigSettingsSource -from enums.deployment_edition import DeploymentEdition from libs.file_utils import search_file_upwards from .deploy import DeploymentConfig @@ -117,11 +116,3 @@ class DifyConfig( ), ), ) - - @property - def DEPLOYMENT_EDITION(self) -> DeploymentEdition: - if self.EDITION == "CLOUD": - return DeploymentEdition.CLOUD - if self.ENTERPRISE_ENABLED: - return DeploymentEdition.ENTERPRISE - return DeploymentEdition.COMMUNITY diff --git a/api/configs/deploy/__init__.py b/api/configs/deploy/__init__.py index 145e9fc5638..638b3dacfd5 100644 --- a/api/configs/deploy/__init__.py +++ b/api/configs/deploy/__init__.py @@ -1,8 +1,8 @@ -from typing import Literal - from pydantic import Field from pydantic_settings import BaseSettings +from enums import DeploymentEdition + class DeploymentConfig(BaseSettings): """ @@ -25,9 +25,14 @@ class DeploymentConfig(BaseSettings): default=False, ) - EDITION: Literal["SELF_HOSTED", "CLOUD"] = Field( - description="Deployment edition of the application (e.g., 'SELF_HOSTED', 'CLOUD')", - default="SELF_HOSTED", + DEPLOYMENT_EDITION: DeploymentEdition = Field( + description="Product edition of the application.", + default=DeploymentEdition.COMMUNITY, + ) + + INIT_PASSWORD: str = Field( + description="Password required before initializing a self-hosted deployment", + default="", ) DEPLOY_ENV: str = Field( diff --git a/api/configs/enterprise/__init__.py b/api/configs/enterprise/__init__.py index ecce58193b8..d839b840d1c 100644 --- a/api/configs/enterprise/__init__.py +++ b/api/configs/enterprise/__init__.py @@ -8,12 +8,6 @@ class EnterpriseFeatureConfig(BaseSettings): **Before using, please contact business@dify.ai by email to inquire about licensing matters.** """ - ENTERPRISE_ENABLED: bool = Field( - description="Enable or disable enterprise-level features." - "Before using, please contact business@dify.ai by email to inquire about licensing matters.", - default=False, - ) - WEBAPP_PUBLIC_ACCESS_ENABLED: bool = Field( description="Whether admins are allowed to set a webapp's access mode to public (anyone with the link, " "no auth). Disable in security-sensitive on-prem deployments.", @@ -25,6 +19,12 @@ class EnterpriseFeatureConfig(BaseSettings): default=False, ) + ENABLE_LICENSE_EXPIRY_NOTICE: bool = Field( + description="Show the license expiry countdown badge in the console when the license is expiring. " + "Disable to hide the badge; license status and all enforcement remain unaffected.", + default=True, + ) + ENTERPRISE_REQUEST_TIMEOUT: int = Field( ge=1, description="Maximum timeout in seconds for enterprise requests", default=5 ) @@ -53,7 +53,7 @@ class EnterpriseTelemetryConfig(BaseSettings): """ ENTERPRISE_TELEMETRY_ENABLED: bool = Field( - description="Enable enterprise telemetry collection (also requires ENTERPRISE_ENABLED=true).", + description="Enable enterprise telemetry collection for enterprise deployments.", default=False, ) diff --git a/api/configs/extra/__init__.py b/api/configs/extra/__init__.py index 3987f326f4b..a142dbd7988 100644 --- a/api/configs/extra/__init__.py +++ b/api/configs/extra/__init__.py @@ -3,6 +3,7 @@ from configs.extra.archive_config import ArchiveStorageConfig from configs.extra.knowledge_fs_config import KnowledgeFSConfig from configs.extra.notion_config import NotionConfig from configs.extra.sentry_config import SentryConfig +from configs.extra.turnstile_config import TurnstileConfig class ExtraServiceConfig( @@ -12,5 +13,6 @@ class ExtraServiceConfig( KnowledgeFSConfig, NotionConfig, SentryConfig, + TurnstileConfig, ): pass diff --git a/api/configs/extra/agent_backend_config.py b/api/configs/extra/agent_backend_config.py index 7baad3d0b44..fd412936390 100644 --- a/api/configs/extra/agent_backend_config.py +++ b/api/configs/extra/agent_backend_config.py @@ -12,6 +12,11 @@ class AgentBackendConfig(BaseSettings): default=None, ) + AGENT_BACKEND_API_TOKEN: str | None = Field( + description="Bearer token for authenticating with the Agent backend /runs API.", + default=None, + ) + AGENT_BACKEND_USE_FAKE: bool = Field( description="Use the deterministic in-process fake Agent backend client.", default=False, @@ -39,9 +44,8 @@ class AgentBackendConfig(BaseSettings): AGENT_SHELL_ENABLED: bool = Field( description=( - "Inject the dify.shell layer (sandboxed bash workspace) into Agent runs. " - "Requires the agent backend to be wired with a shellctl entrypoint before " - "shell-using Agent runs are executed." + "Inject the Home, Workspace, Sandbox, and Shell runtime layers into Agent runs. " + "Requires Dify Agent to have a deployment-selected runtime backend." ), default=True, ) diff --git a/api/configs/extra/turnstile_config.py b/api/configs/extra/turnstile_config.py new file mode 100644 index 00000000000..c4ae92924e3 --- /dev/null +++ b/api/configs/extra/turnstile_config.py @@ -0,0 +1,34 @@ +from pydantic import Field, SecretStr, field_validator +from pydantic_settings import BaseSettings + + +class TurnstileConfig(BaseSettings): + """Server-side Cloudflare Turnstile settings for Cloud sign-in.""" + + TURNSTILE_SECRET_KEY: SecretStr | None = Field( + default=None, + description="Secret key used to validate Cloudflare Turnstile tokens.", + ) + TURNSTILE_ALLOWED_HOSTNAMES: str = Field( + default="", + description="Comma-separated parent or exact hostnames accepted from Turnstile.", + ) + + @field_validator("TURNSTILE_SECRET_KEY", mode="before") + @classmethod + def normalize_secret_key(cls, value: object) -> object: + if isinstance(value, SecretStr): + normalized = value.get_secret_value().strip() + return SecretStr(normalized) if normalized else None + if isinstance(value, str): + normalized = value.strip() + return normalized or None + return value + + @property + def TURNSTILE_ALLOWED_HOSTNAME_SET(self) -> frozenset[str]: + return frozenset( + hostname.strip().lower().strip(".") + for hostname in self.TURNSTILE_ALLOWED_HOSTNAMES.split(",") + if hostname.strip().strip(".") + ) diff --git a/api/configs/feature/__init__.py b/api/configs/feature/__init__.py index c28716e3b0e..ea2b4cacf80 100644 --- a/api/configs/feature/__init__.py +++ b/api/configs/feature/__init__.py @@ -266,6 +266,12 @@ class PluginConfig(BaseSettings): default=60 * 60, ) + PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED: bool = Field( + description="Whether tenant plugin model providers are cached in Redis. Disable when plugins are installed " + "by a system other than this one, which cannot invalidate the cache when a tenant's plugins change.", + default=True, + ) + PLUGIN_MODEL_PROVIDERS_CACHE_TTL: PositiveInt = Field( description="TTL in seconds for caching tenant plugin model providers in Redis", default=60 * 60 * 24, @@ -444,6 +450,11 @@ class FileUploadConfig(BaseSettings): default=15, ) + KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN: NonNegativeInt = Field( + description="Maximum allowed file size for knowledge uploads on paid cloud plans in megabytes", + default=15, + ) + UPLOAD_FILE_BATCH_LIMIT: NonNegativeInt = Field( description="Maximum number of files allowed in a single upload batch", default=5, @@ -464,6 +475,11 @@ class FileUploadConfig(BaseSettings): default=50, ) + UPLOAD_SKILL_FILE_SIZE_LIMIT: NonNegativeInt = Field( + description="Maximum allowed Skill package size for uploads in megabytes", + default=50, + ) + BATCH_UPLOAD_LIMIT: NonNegativeInt = Field( description="Maximum number of files allowed in a batch upload operation", default=20, @@ -794,17 +810,6 @@ class ModelLoadBalanceConfig(BaseSettings): ) -class BillingConfig(BaseSettings): - """ - Configuration for platform billing features - """ - - BILLING_ENABLED: bool = Field( - description="Enable or disable billing functionality", - default=False, - ) - - class UpdateConfig(BaseSettings): """ Configuration for application update checks @@ -816,6 +821,41 @@ class UpdateConfig(BaseSettings): ) +class CommunityTelemetryConfig(BaseSettings): + """ + Configuration for anonymous self-hosted community telemetry. + """ + + DISABLE_TELEMETRY: bool = Field( + description="Disable anonymous community telemetry", + default=False, + ) + DO_NOT_TRACK: bool = Field( + description="Respect the standard do-not-track opt-out signal for telemetry", + default=False, + ) + TELEMETRY_ENDPOINT: str = Field( + description="Endpoint for anonymous community telemetry events", + default="https://otel.dify.ai/v1/events", + ) + TELEMETRY_FALLBACK_ENDPOINT: str = Field( + description="Fallback endpoint for anonymous community telemetry events", + default="https://otel.dify.cn/v1/events", + ) + TELEMETRY_TIMEOUT_SECONDS: PositiveInt = Field( + description="HTTP timeout in seconds for anonymous community telemetry requests", + default=3, + ) + TELEMETRY_HEARTBEAT_INTERVAL_MINUTES: PositiveInt = Field( + description="Celery beat interval in minutes for checking whether heartbeat telemetry is due", + default=30, + ) + CI: bool = Field( + description="Whether the process is running in CI; telemetry is skipped when true", + default=False, + ) + + class WorkflowVariableTruncationConfig(BaseSettings): WORKFLOW_VARIABLE_TRUNCATION_MAX_SIZE: PositiveInt = Field( # 1000 KiB @@ -824,7 +864,7 @@ class WorkflowVariableTruncationConfig(BaseSettings): ) WORKFLOW_VARIABLE_TRUNCATION_STRING_LENGTH: PositiveInt = Field( 100000, - description="maximum length for string to trigger tuncation, measure in number of characters", + description="maximum length for string to trigger truncation, measure in number of characters", ) WORKFLOW_VARIABLE_TRUNCATION_ARRAY_LENGTH: PositiveInt = Field( 1000, @@ -878,9 +918,9 @@ class WorkflowConfig(BaseSettings): default=10, ) - GRAPH_ENGINE_SCALE_UP_THRESHOLD: PositiveInt = Field( - description="Queue depth threshold that triggers worker scale up", - default=3, + GRAPH_ENGINE_SCALE_UP_THRESHOLD: NonNegativeInt = Field( + description="Pending task threshold that triggers worker scale up", + default=0, ) GRAPH_ENGINE_SCALE_DOWN_IDLE_TIME: float = Field( @@ -1570,7 +1610,6 @@ class FeatureConfig( # place the configs in alphabet order AppExecutionConfig, AuthConfig, # Changed from OAuthConfig to AuthConfig - BillingConfig, CodeExecutionSandboxConfig, CreatorsPlatformConfig, TriggerConfig, @@ -1599,6 +1638,7 @@ class FeatureConfig( TenantIsolatedTaskQueueConfig, ToolConfig, UpdateConfig, + CommunityTelemetryConfig, WorkflowConfig, WorkflowNodeExecutionConfig, WorkspaceConfig, diff --git a/api/configs/feature/hosted_service/__init__.py b/api/configs/feature/hosted_service/__init__.py index 42ede718c46..4cc6aa3ef89 100644 --- a/api/configs/feature/hosted_service/__init__.py +++ b/api/configs/feature/hosted_service/__init__.py @@ -451,6 +451,12 @@ class HostedFetchAppTemplateConfig(BaseSettings): default="https://tmpl.dify.ai", ) + HOSTED_FETCH_APP_TEMPLATES_CACHE_TTL: int = Field( + description="TTL in seconds for caching remote app template HTTP responses. 0 disables caching.", + default=600, + ge=0, + ) + class HostedFetchPipelineTemplateConfig(BaseSettings): """ diff --git a/api/configs/middleware/__init__.py b/api/configs/middleware/__init__.py index 865bb48c676..f442e913760 100644 --- a/api/configs/middleware/__init__.py +++ b/api/configs/middleware/__init__.py @@ -7,6 +7,7 @@ from pydantic_settings import BaseSettings from .cache.redis_config import RedisConfig from .cache.redis_pubsub_config import RedisPubSubConfig +from .key_provider.azure_keyvault_config import AzureKeyVaultConfig from .storage.aliyun_oss_storage_config import AliyunOSSStorageConfig from .storage.amazon_s3_storage_config import S3StorageConfig from .storage.azure_blob_storage_config import AzureBlobStorageConfig @@ -83,6 +84,22 @@ class StorageConfig(BaseSettings): ) +_VALID_KEY_PROVIDER_TYPE = Literal[ + "local", + "azure-keyvault", +] + + +class KeyProviderConfig(BaseSettings): + KEY_PROVIDER_TYPE: _VALID_KEY_PROVIDER_TYPE = Field( + description="Key provider used to encrypt/decrypt tenant credentials (LLM/tool provider secrets)." + " Options: 'local' (per-tenant RSA key pair, private key kept in the STORAGE_TYPE backend)," + " 'azure-keyvault' (per-tenant RSA key kept in Azure Key Vault, private key never leaves the vault)." + " Default is 'local'.", + default=cast(_VALID_KEY_PROVIDER_TYPE, "local"), + ) + + class VectorStoreConfig(BaseSettings): VECTOR_STORE: str | None = Field( description="Type of vector store to use for efficient similarity search." @@ -359,6 +376,9 @@ class MiddlewareConfig( KeywordStoreConfig, RedisConfig, RedisPubSubConfig, + # configs of the tenant credential encryption key provider + KeyProviderConfig, + AzureKeyVaultConfig, # configs of storage and storage providers StorageConfig, AliyunOSSStorageConfig, diff --git a/api/configs/middleware/key_provider/azure_keyvault_config.py b/api/configs/middleware/key_provider/azure_keyvault_config.py new file mode 100644 index 00000000000..0b90cd287d5 --- /dev/null +++ b/api/configs/middleware/key_provider/azure_keyvault_config.py @@ -0,0 +1,38 @@ +from pydantic import Field, field_validator +from pydantic_settings import BaseSettings + + +class AzureKeyVaultConfig(BaseSettings): + """ + Configuration settings for Azure Key Vault, used as a tenant credential encryption key provider + """ + + AZURE_KEYVAULT_VAULT_URL: str | None = Field( + description="URL of the Azure Key Vault instance (e.g., 'https://.vault.azure.net')." + " Required when KEY_PROVIDER_TYPE is set to 'azure-keyvault'. Authentication uses" + " DefaultAzureCredential (managed identity, environment variables, or Azure CLI login).", + default=None, + ) + + AZURE_KEYVAULT_KEY_SIZE: int = Field( + description="RSA key size (in bits) used when Dify provisions a new per-tenant key in Azure Key Vault.", + default=2048, + ) + + AZURE_KEYVAULT_ROTATION_INTERVAL_DAYS: int | None = Field( + description="If set, Dify configures each newly-created per-tenant Key Vault key to auto-rotate every" + " N days (using a 'time after create' trigger, with no expiry set on generated versions)." + " Old key versions are kept forever and remain usable for decrypting credentials encrypted" + " before the rotation, so rotation requires no manual re-encryption. Leave unset (default) to" + " not configure a rotation policy and manage rotation manually in Azure. Must be at least 7 days.", + default=None, + ge=7, + ) + + @field_validator("AZURE_KEYVAULT_ROTATION_INTERVAL_DAYS", mode="before") + @classmethod + def _empty_string_to_none_for_rotation_interval(cls, v): + """Allow empty string in env/.env (e.g. an unfilled template value) to mean 'unset'.""" + if isinstance(v, str) and v.strip() == "": + return None + return v diff --git a/api/configs/middleware/vdb/tidb_on_qdrant_config.py b/api/configs/middleware/vdb/tidb_on_qdrant_config.py index 9ca09551294..6fb7ff08efb 100644 --- a/api/configs/middleware/vdb/tidb_on_qdrant_config.py +++ b/api/configs/middleware/vdb/tidb_on_qdrant_config.py @@ -32,6 +32,11 @@ class TidbOnQdrantConfig(BaseSettings): default=6334, ) + TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB: str = Field( + description="Cloud pre-write thresholds for projected TiDB vector storage usage, in plan:MB pairs.", + default="sandbox:60,professional:6400,team:25600", + ) + TIDB_PUBLIC_KEY: str | None = Field( description="Tidb account public key", default=None, diff --git a/api/configs/middleware/vdb/tidb_vector_config.py b/api/configs/middleware/vdb/tidb_vector_config.py index 0ebf226bea6..172b5a7e56e 100644 --- a/api/configs/middleware/vdb/tidb_vector_config.py +++ b/api/configs/middleware/vdb/tidb_vector_config.py @@ -31,3 +31,8 @@ class TiDBVectorConfig(BaseSettings): description="Name of the TiDB Vector database to connect to", default=None, ) + + TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH: bool = Field( + description="Enable TiDB Vector full-text and hybrid search features", + default=False, + ) diff --git a/api/constants/model_template.py b/api/constants/model_template.py index 8a027f10e57..bc76c222dfd 100644 --- a/api/constants/model_template.py +++ b/api/constants/model_template.py @@ -88,8 +88,9 @@ default_app_templates: Mapping[AppMode, Mapping] = { AppMode.AGENT: { "app": { "mode": AppMode.AGENT, - "enable_site": True, - "enable_api": True, + # Public access is enabled atomically by the first successful publish. + "enable_site": False, + "enable_api": False, }, }, } diff --git a/api/controllers/API_SCHEMA_GUIDE.md b/api/controllers/API_SCHEMA_GUIDE.md index a1e412a3630..30e6ddc0a8d 100644 --- a/api/controllers/API_SCHEMA_GUIDE.md +++ b/api/controllers/API_SCHEMA_GUIDE.md @@ -162,6 +162,9 @@ That documents a GET request body and is not the expected contract. ## Responses +`204 No Content` responses must not serialize a response body. Return the status using the established controller pattern; +do not return a dictionary, response model, or other payload. + Response models should inherit from `ResponseModel`: ```python diff --git a/api/controllers/common/agent_app_parameters.py b/api/controllers/common/agent_app_parameters.py index c1c9fcdb23a..5bc41379444 100644 --- a/api/controllers/common/agent_app_parameters.py +++ b/api/controllers/common/agent_app_parameters.py @@ -3,6 +3,7 @@ from typing import Any from sqlalchemy import select from sqlalchemy.orm import Session +from core.agent.publish_visibility import agent_has_workflow_callable_active_snapshot from core.app.apps.agent_app.app_feature_projection import merge_agent_app_features from core.app.apps.agent_app.app_variable_projection import agent_app_variables_to_user_input_form from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError @@ -34,9 +35,7 @@ def get_published_agent_app_feature_dict_and_user_input_form( ) if agent is None: raise AgentAppGeneratorError("Agent App has no bound Agent") - # active_config_is_published means the draft has no unpublished edits; the public app - # can still read parameters from the active snapshot while a newer draft is pending. - if not agent.active_config_snapshot_id: + if not agent_has_workflow_callable_active_snapshot(session=session, agent=agent): raise AgentAppNotPublishedError("Agent has not been published") snapshot = session.scalar( diff --git a/api/controllers/common/fields.py b/api/controllers/common/fields.py index 7faa2cf8542..65d547a349f 100644 --- a/api/controllers/common/fields.py +++ b/api/controllers/common/fields.py @@ -126,12 +126,6 @@ class UsageCountResponse(ResponseModel): count: int -class IndexInfoResponse(ResponseModel): - welcome: str - api_version: str - server_version: str - - class AvatarUrlResponse(ResponseModel): avatar_url: str diff --git a/api/controllers/common/session.py b/api/controllers/common/session.py index 24b1a8729d3..fdffa46b189 100644 --- a/api/controllers/common/session.py +++ b/api/controllers/common/session.py @@ -52,7 +52,7 @@ def with_session[T, **P, R]( session.commit() return result except Exception: - session.rollback() # noqa: no-new-controller-sqlalchemy decorator owns transaction rollback + session.rollback() # guard-ignore: no-new-controller-sqlalchemy -- decorator owns rollback raise with session_factory.create_session() as session: diff --git a/api/controllers/common/wraps.py b/api/controllers/common/wraps.py index 34fe8cc96aa..236423dd351 100644 --- a/api/controllers/common/wraps.py +++ b/api/controllers/common/wraps.py @@ -152,7 +152,7 @@ def _extract_resource_id( if resource_type == RBACResourceScope.APP: app_id = matched_args.get("app_id") if app_id: - return str(app_id) + return str(app_id) # pyrefly: ignore[unnecessary-type-conversion] agent_id = matched_args.get("agent_id") if agent_id: @@ -161,7 +161,7 @@ def _extract_resource_id( resource_id = matched_args.get("resource_id") if resource_id: - return str(resource_id) + return str(resource_id) # pyrefly: ignore[unnecessary-type-conversion] raise ValueError("Missing app_id in request path") if resource_type == RBACResourceScope.DATASET: @@ -171,9 +171,11 @@ def _extract_resource_id( pipeline_id = matched_args.get("pipeline_id") if pipeline_id: - dataset = db.session.scalar(select(Dataset).where(Dataset.pipeline_id == str(pipeline_id))) + dataset = db.session.scalar( + select(Dataset).where(Dataset.pipeline_id == str(pipeline_id), Dataset.tenant_id == tenant_id) + ) if not dataset: raise NotFound("Dataset not found for pipeline") - return str(dataset.id) + return str(dataset.id) # pyrefly: ignore[unnecessary-type-conversion] raise ValueError("Missing dataset_id or pipeline_id in request path") raise ValueError(f"Unknown resource_type: {resource_type}") diff --git a/api/controllers/console/__init__.py b/api/controllers/console/__init__.py index 34869760f20..cd87e1f8243 100644 --- a/api/controllers/console/__init__.py +++ b/api/controllers/console/__init__.py @@ -41,10 +41,9 @@ from . import ( knowledge_fs_proxy, notification, onboarding, - ping, setup, spec, - version, + system, workflow_run_archive, ) from .agent import composer as agent_composer @@ -213,7 +212,6 @@ __all__ = [ "onboarding", "ops_trace", "parameter", - "ping", "plugin", "rag_pipeline", "rag_pipeline_datasets", @@ -231,11 +229,11 @@ __all__ = [ "socketio_workflow", "spec", "statistic", + "system", "tags", "tool_providers", "trial", "trigger_providers", - "version", "website", "workflow", "workflow_app_log", diff --git a/api/controllers/console/agent/composer.py b/api/controllers/console/agent/composer.py index c48818519ab..76d7863776e 100644 --- a/api/controllers/console/agent/composer.py +++ b/api/controllers/console/agent/composer.py @@ -1,6 +1,5 @@ from uuid import UUID -from flask import request from flask_restx import Resource from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound @@ -14,6 +13,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -64,8 +64,16 @@ class WorkflowAgentComposerApi(Resource): @with_current_tenant_id @with_session @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) - def get(self, session: Session, tenant_id: str, account_id: str, app_model: App, node_id: str): - query = WorkflowAgentComposerQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(WorkflowAgentComposerQuery) + def get( + self, + req_data: WorkflowAgentComposerQuery, + session: Session, + tenant_id: str, + account_id: str, + app_model: App, + node_id: str, + ): return dump_response( WorkflowAgentComposerResponse, AgentComposerService.load_workflow_composer( @@ -74,7 +82,7 @@ class WorkflowAgentComposerApi(Resource): app_id=app_model.id, node_id=node_id, account_id=account_id, - snapshot_id=query.snapshot_id, + snapshot_id=req_data.snapshot_id, ), ) @@ -91,8 +99,16 @@ class WorkflowAgentComposerApi(Resource): @with_current_tenant_id @with_session @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) - def put(self, session: Session, tenant_id: str, account_id: str, app_model: App, node_id: str): - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) + @model_validate(ComposerSavePayload) + def put( + self, + req_data: ComposerSavePayload, + session: Session, + tenant_id: str, + account_id: str, + app_model: App, + node_id: str, + ): return dump_response( WorkflowAgentComposerResponse, AgentComposerService.save_workflow_composer( @@ -101,7 +117,7 @@ class WorkflowAgentComposerApi(Resource): app_id=app_model.id, node_id=node_id, account_id=account_id, - payload=payload, + payload=req_data, ), ) @@ -123,8 +139,16 @@ class WorkflowAgentComposerCopyFromRosterApi(Resource): @with_current_tenant_id @with_session @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) - def post(self, session: Session, tenant_id: str, account_id: str, app_model: App, node_id: str): - payload = WorkflowComposerCopyFromRosterPayload.model_validate(console_ns.payload or {}) + @model_validate(WorkflowComposerCopyFromRosterPayload) + def post( + self, + req_data: WorkflowComposerCopyFromRosterPayload, + session: Session, + tenant_id: str, + account_id: str, + app_model: App, + node_id: str, + ): return dump_response( WorkflowAgentComposerResponse, AgentComposerService.copy_workflow_composer_from_roster( @@ -133,9 +157,9 @@ class WorkflowAgentComposerCopyFromRosterApi(Resource): app_id=app_model.id, node_id=node_id, account_id=account_id, - source_agent_id=payload.source_agent_id, - source_snapshot_id=payload.source_snapshot_id, - idempotency_key=payload.idempotency_key, + source_agent_id=req_data.source_agent_id, + source_snapshot_id=req_data.source_snapshot_id, + idempotency_key=req_data.idempotency_key, ), ) @@ -152,16 +176,16 @@ class WorkflowAgentComposerValidateApi(Resource): @with_current_tenant_id @with_session(write=False) @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) - def post(self, session: Session, tenant_id: str, app_model: App, node_id: str): - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) - ComposerConfigValidator.validate_publish_payload(payload) + @model_validate(ComposerSavePayload) + def post(self, req_data: ComposerSavePayload, session: Session, tenant_id: str, app_model: App, node_id: str): + ComposerConfigValidator.validate_publish_payload(req_data) AgentComposerService.validate_knowledge_datasets( - session=session, tenant_id=tenant_id, agent_soul=payload.agent_soul + session=session, tenant_id=tenant_id, agent_soul=req_data.agent_soul ) findings = AgentComposerService.collect_validation_findings( session=session, tenant_id=tenant_id, - payload=payload, + payload=req_data, agent_id=AgentComposerService.resolve_workflow_node_agent_id( session=session, tenant_id=tenant_id, app_id=app_model.id, node_id=node_id ), @@ -204,9 +228,9 @@ class WorkflowAgentComposerImpactApi(Resource): @with_current_tenant_id @with_session(write=False) @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) - def post(self, session: Session, tenant_id: str, app_model: App, node_id: str): - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) - current_snapshot_id = payload.binding.current_snapshot_id if payload.binding else None + @model_validate(ComposerSavePayload) + def post(self, req_data: ComposerSavePayload, session: Session, tenant_id: str, app_model: App, node_id: str): + current_snapshot_id = req_data.binding.current_snapshot_id if req_data.binding else None if not current_snapshot_id: return dump_response( AgentComposerImpactResponse, {"current_snapshot_id": None, "workflow_node_count": 0, "bindings": []} @@ -235,8 +259,16 @@ class WorkflowAgentComposerSaveToRosterApi(Resource): @with_current_tenant_id @with_session @get_app_model(mode=[AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) - def post(self, session: Session, tenant_id: str, account_id: str, app_model: App, node_id: str): - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) + @model_validate(ComposerSavePayload) + def post( + self, + req_data: ComposerSavePayload, + session: Session, + tenant_id: str, + account_id: str, + app_model: App, + node_id: str, + ): return dump_response( WorkflowAgentComposerResponse, AgentComposerService.save_workflow_composer( @@ -245,14 +277,14 @@ class WorkflowAgentComposerSaveToRosterApi(Resource): app_id=app_model.id, node_id=node_id, account_id=account_id, - payload=payload, + payload=req_data, ), ) def _require_snippet_app_id(*, session: Session, tenant_id: str, snippet_id: UUID) -> str: snippet = SnippetService(session=session).get_snippet_by_id( - snippet_id=str(snippet_id), + snippet_id=str(snippet_id), # pyrefly: ignore[unnecessary-type-conversion] tenant_id=tenant_id, ) if snippet is None: @@ -270,8 +302,16 @@ class SnippetAgentComposerApi(Resource): @with_current_user_id @with_current_tenant_id @with_session - def get(self, session: Session, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): - query = WorkflowAgentComposerQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(WorkflowAgentComposerQuery) + def get( + self, + req_data: WorkflowAgentComposerQuery, + session: Session, + tenant_id: str, + account_id: str, + snippet_id: UUID, + node_id: str, + ): return dump_response( WorkflowAgentComposerResponse, AgentComposerService.load_workflow_composer( @@ -280,7 +320,7 @@ class SnippetAgentComposerApi(Resource): app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, account_id=account_id, - snapshot_id=query.snapshot_id, + snapshot_id=req_data.snapshot_id, ), ) @@ -296,8 +336,16 @@ class SnippetAgentComposerApi(Resource): @with_current_user_id @with_current_tenant_id @with_session - def put(self, session: Session, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) + @model_validate(ComposerSavePayload) + def put( + self, + req_data: ComposerSavePayload, + session: Session, + tenant_id: str, + account_id: str, + snippet_id: UUID, + node_id: str, + ): return dump_response( WorkflowAgentComposerResponse, AgentComposerService.save_workflow_composer( @@ -306,7 +354,7 @@ class SnippetAgentComposerApi(Resource): app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, account_id=account_id, - payload=payload, + payload=req_data, ), ) @@ -327,8 +375,16 @@ class SnippetAgentComposerCopyFromRosterApi(Resource): @with_current_user_id @with_current_tenant_id @with_session - def post(self, session: Session, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): - payload = WorkflowComposerCopyFromRosterPayload.model_validate(console_ns.payload or {}) + @model_validate(WorkflowComposerCopyFromRosterPayload) + def post( + self, + req_data: WorkflowComposerCopyFromRosterPayload, + session: Session, + tenant_id: str, + account_id: str, + snippet_id: UUID, + node_id: str, + ): return dump_response( WorkflowAgentComposerResponse, AgentComposerService.copy_workflow_composer_from_roster( @@ -337,9 +393,9 @@ class SnippetAgentComposerCopyFromRosterApi(Resource): app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, account_id=account_id, - source_agent_id=payload.source_agent_id, - source_snapshot_id=payload.source_snapshot_id, - idempotency_key=payload.idempotency_key, + source_agent_id=req_data.source_agent_id, + source_snapshot_id=req_data.source_snapshot_id, + idempotency_key=req_data.idempotency_key, ), ) @@ -355,17 +411,17 @@ class SnippetAgentComposerValidateApi(Resource): @account_initialization_required @with_current_tenant_id @with_session(write=False) - def post(self, session: Session, tenant_id: str, snippet_id: UUID, node_id: str): + @model_validate(ComposerSavePayload) + def post(self, req_data: ComposerSavePayload, session: Session, tenant_id: str, snippet_id: UUID, node_id: str): app_id = _require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id) - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) - ComposerConfigValidator.validate_publish_payload(payload) + ComposerConfigValidator.validate_publish_payload(req_data) AgentComposerService.validate_knowledge_datasets( - session=session, tenant_id=tenant_id, agent_soul=payload.agent_soul + session=session, tenant_id=tenant_id, agent_soul=req_data.agent_soul ) findings = AgentComposerService.collect_validation_findings( session=session, tenant_id=tenant_id, - payload=payload, + payload=req_data, agent_id=AgentComposerService.resolve_workflow_node_agent_id( session=session, tenant_id=tenant_id, @@ -409,10 +465,10 @@ class SnippetAgentComposerImpactApi(Resource): @account_initialization_required @with_current_tenant_id @with_session(write=False) - def post(self, session: Session, tenant_id: str, snippet_id: UUID, node_id: str): + @model_validate(ComposerSavePayload) + def post(self, req_data: ComposerSavePayload, session: Session, tenant_id: str, snippet_id: UUID, node_id: str): _require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id) - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) - current_snapshot_id = payload.binding.current_snapshot_id if payload.binding else None + current_snapshot_id = req_data.binding.current_snapshot_id if req_data.binding else None if not current_snapshot_id: return dump_response( AgentComposerImpactResponse, {"current_snapshot_id": None, "workflow_node_count": 0, "bindings": []} @@ -444,8 +500,16 @@ class SnippetAgentComposerSaveToRosterApi(Resource): @with_current_user_id @with_current_tenant_id @with_session - def post(self, session: Session, tenant_id: str, account_id: str, snippet_id: UUID, node_id: str): - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) + @model_validate(ComposerSavePayload) + def post( + self, + req_data: ComposerSavePayload, + session: Session, + tenant_id: str, + account_id: str, + snippet_id: UUID, + node_id: str, + ): return dump_response( WorkflowAgentComposerResponse, AgentComposerService.save_workflow_composer( @@ -454,7 +518,7 @@ class SnippetAgentComposerSaveToRosterApi(Resource): app_id=_require_snippet_app_id(session=session, tenant_id=tenant_id, snippet_id=snippet_id), node_id=node_id, account_id=account_id, - payload=payload, + payload=req_data, ), ) @@ -484,8 +548,8 @@ class AgentComposerApi(Resource): @with_current_user_id @with_current_tenant_id @with_session - def put(self, session: Session, tenant_id: str, account_id: str, agent_id: UUID): - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) + @model_validate(ComposerSavePayload) + def put(self, req_data: ComposerSavePayload, session: Session, tenant_id: str, account_id: str, agent_id: UUID): return dump_response( AgentAppComposerResponse, AgentComposerService.save_agent_composer( @@ -493,7 +557,7 @@ class AgentComposerApi(Resource): tenant_id=tenant_id, agent_id=str(agent_id), account_id=account_id, - payload=payload, + payload=req_data, ), ) @@ -509,17 +573,17 @@ class AgentComposerValidateApi(Resource): @account_initialization_required @with_current_tenant_id @with_session - def post(self, session: Session, tenant_id: str, agent_id: UUID): + @model_validate(ComposerSavePayload) + def post(self, req_data: ComposerSavePayload, session: Session, tenant_id: str, agent_id: UUID): AgentComposerService.load_agent_composer(session=session, tenant_id=tenant_id, agent_id=str(agent_id)) - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) - ComposerConfigValidator.validate_publish_payload(payload) + ComposerConfigValidator.validate_publish_payload(req_data) AgentComposerService.validate_knowledge_datasets( - session=session, tenant_id=tenant_id, agent_soul=payload.agent_soul + session=session, tenant_id=tenant_id, agent_soul=req_data.agent_soul ) findings = AgentComposerService.collect_validation_findings( session=session, tenant_id=tenant_id, - payload=payload, + payload=req_data, agent_id=str(agent_id), ) return dump_response(AgentComposerValidateResponse, {"result": "success", "errors": [], **findings}) diff --git a/api/controllers/console/agent/roster.py b/api/controllers/console/agent/roster.py index 60fb018915b..e549d43bf27 100644 --- a/api/controllers/console/agent/roster.py +++ b/api/controllers/console/agent/roster.py @@ -3,7 +3,7 @@ from uuid import UUID from flask import abort, request from flask_restx import Resource from pydantic import AliasChoices, BaseModel, Field, field_validator -from sqlalchemy import func, select +from sqlalchemy import func, or_, select from sqlalchemy.orm import Session from controllers.common.schema import ( @@ -38,11 +38,13 @@ from controllers.console.wraps import ( edit_permission_required, enterprise_license_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, with_current_user, ) +from core.agent.publish_visibility import agent_has_workflow_callable_active_snapshot from fields.agent_fields import ( AgentConfigDraftSummaryResponse, AgentConfigSnapshotDetailResponse, @@ -62,7 +64,7 @@ from libs.datetime_utils import parse_time_range from libs.helper import dump_response from libs.login import login_required from models import Account -from models.agent import Agent, AgentConfigDraftType, AgentStatus +from models.agent import Agent, AgentStatus from models.agent_config_entities import AgentSoulConfig from models.enums import ApiTokenType from models.model import ApiToken, App, IconType @@ -139,6 +141,7 @@ class AgentApiStatusPayload(BaseModel): class AgentApiAccessResponse(BaseModel): + access_ready: bool enabled: bool service_api_base_url: str streaming_only: bool = True @@ -257,7 +260,7 @@ class AgentAppDetailWithSite(GenericAppDetailWithSite): debug_conversation_has_messages: bool = False debug_conversation_message_count: int = 0 role: str | None = None - active_config_is_published: bool = False + access_ready: bool = False class AgentDebugConversationRefreshResponse(BaseModel): @@ -266,13 +269,6 @@ class AgentDebugConversationRefreshResponse(BaseModel): debug_conversation_message_count: int = 0 -class AgentDebugConversationRefreshPayload(BaseModel): - draft_type: AgentConfigDraftType = Field( - default=AgentConfigDraftType.DEBUG_BUILD, - description="Agent draft surface whose conversation should be refreshed", - ) - - class AgentPublishPayload(BaseModel): version_note: str | None = Field(default=None, description="Optional note for this published Agent version") @@ -316,7 +312,6 @@ register_schema_models( AgentAppCopyPayload, AgentPublishPayload, AgentBuildDraftCheckoutPayload, - AgentDebugConversationRefreshPayload, ComposerSavePayload, AgentApiStatusPayload, AgentInviteOptionsQuery, @@ -396,11 +391,10 @@ def _serialize_agent_app_detail( payload["backing_app_id"] = roster_service.runtime_backing_app_id(agent) payload["hidden_app_backed"] = bool(agent.backing_app_id and agent.backing_app_id != agent.app_id) payload["id"] = agent.id - debug_conversation_id = roster_service.get_or_create_agent_app_debug_conversation_id( + debug_conversation_id = roster_service.get_or_create_build_conversation( 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( @@ -410,10 +404,7 @@ def _serialize_agent_app_detail( payload["debug_conversation_has_messages"] = message_count > 0 payload["debug_conversation_message_count"] = message_count payload["role"] = agent.role or "" - payload["active_config_is_published"] = roster_service.active_config_is_published( - tenant_id=app_model.tenant_id, - agent=agent, - ) + payload["access_ready"] = agent_has_workflow_callable_active_snapshot(session=session, agent=agent) return payload @@ -444,11 +435,10 @@ def _serialize_agent_app_pagination(session: Session, app_pagination, *, tenant_ tenant_id=tenant_id, agent_ids=[agent.id for agent in agents_by_app_id.values()], ) - debug_conversation_ids_by_agent_id = roster_service.load_or_create_agent_app_debug_conversation_ids_by_agent_id( + debug_conversation_ids_by_agent_id = roster_service.load_or_create_build_conversation_ids_by_agent_id( 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, @@ -494,22 +484,33 @@ def _resolve_agent_runtime_app_model(session: Session, *, tenant_id: str, agent_ return _agent_roster_service(session).get_agent_runtime_app_model(tenant_id=tenant_id, agent_id=str(agent_id)) -def _agent_api_key_count(session: Session, app_id: str) -> int: +def _agent_api_key_count(session: Session, app_model: App) -> int: return ( session.scalar( select(func.count(ApiToken.id)).where( + or_(ApiToken.tenant_id == app_model.tenant_id, ApiToken.tenant_id.is_(None)), ApiToken.type == ApiTokenType.APP, - ApiToken.app_id == app_id, + ApiToken.app_id == app_model.id, ) ) or 0 ) +def _agent_app_access_ready(session: Session, app_model: App) -> bool: + agent = _agent_roster_service(session).get_app_backing_agent( + tenant_id=app_model.tenant_id, + app_id=str(app_model.id), + ) + return bool(agent and agent_has_workflow_callable_active_snapshot(session=session, agent=agent)) + + def _serialize_agent_api_access(session: Session, app_model: App) -> dict: base_url = app_model.api_base_url + access_ready = _agent_app_access_ready(session, app_model) response = AgentApiAccessResponse( - enabled=bool(app_model.enable_api), + access_ready=access_ready, + enabled=bool(app_model.enable_api and access_ready), service_api_base_url=base_url, chat_endpoint=f"{base_url}/chat-messages", stop_endpoint=f"{base_url}/chat-messages/{{task_id}}/stop", @@ -521,7 +522,7 @@ def _serialize_agent_api_access(session: Session, app_model: App) -> dict: meta_endpoint=f"{base_url}/meta", api_rpm=app_model.api_rpm or 0, api_rph=app_model.api_rph or 0, - api_key_count=_agent_api_key_count(session, str(app_model.id)), + api_key_count=_agent_api_key_count(session, app_model), ) return response.model_dump(mode="json") @@ -608,16 +609,16 @@ class AgentAppListApi(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, current_tenant_id: str, current_user: Account): - args = AgentAppCreatePayload.model_validate(console_ns.payload) + @model_validate(AgentAppCreatePayload) + def post(self, req_data: AgentAppCreatePayload, session: Session, current_tenant_id: str, current_user: Account): params = CreateAppParams( - name=args.name, - description=args.description, + name=req_data.name, + description=req_data.description, mode="agent", - agent_role=args.role or "", - icon_type=args.icon_type, - icon=args.icon, - icon_background=args.icon_background, + agent_role=req_data.role or "", + icon_type=req_data.icon_type, + icon=req_data.icon, + icon_background=req_data.icon_background, ) app = AppService().create_app(current_tenant_id, params, current_user, session=session) @@ -650,18 +651,25 @@ class AgentAppApi(Resource): @with_current_user @with_current_tenant_id @with_session - def put(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + @model_validate(AgentAppUpdatePayload) + def put( + self, + req_data: AgentAppUpdatePayload, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + ): app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) - args = AgentAppUpdatePayload.model_validate(console_ns.payload) args_dict: AppService.ArgsDict = { - "name": args.name, - "description": args.description or "", - "icon_type": args.icon_type, - "icon": args.icon or "", - "icon_background": args.icon_background or "", - "use_icon_as_answer_icon": args.use_icon_as_answer_icon or False, - "max_active_requests": args.max_active_requests or 0, - "role": args.role, + "name": req_data.name, + "description": req_data.description or "", + "icon_type": req_data.icon_type, + "icon": req_data.icon or "", + "icon_background": req_data.icon_background or "", + "use_icon_as_answer_icon": req_data.use_icon_as_answer_icon or False, + "max_active_requests": req_data.max_active_requests or 0, + "role": req_data.role, } updated = AppService().update_app(app_model, args_dict, session=session) return _serialize_agent_app_detail(session, updated, current_user=current_user) @@ -683,16 +691,6 @@ class AgentAppApi(Resource): @console_ns.route("/agent//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", @@ -707,12 +705,10 @@ class AgentDebugConversationRefreshApi(Resource): @with_current_tenant_id @with_session def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): - args = AgentDebugConversationRefreshPayload.model_validate(request.get_json(silent=True) or {}) - debug_conversation_id = _agent_roster_service(session).refresh_agent_app_debug_conversation_id( + debug_conversation_id = _agent_roster_service(session).reset_build_conversation( 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, @@ -734,14 +730,21 @@ class AgentPublishApi(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): - args = AgentPublishPayload.model_validate(console_ns.payload or {}) + @model_validate(AgentPublishPayload) + def post( + self, + req_data: AgentPublishPayload, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + ): return AgentComposerService.publish_agent_app_draft( session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, - version_note=args.version_note, + version_note=req_data.version_note, ) @@ -756,15 +759,22 @@ class AgentBuildDraftCheckoutApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False) @with_current_user @with_current_tenant_id - @with_session - def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): - args = AgentBuildDraftCheckoutPayload.model_validate(console_ns.payload or {}) + @with_session(write=False) + @model_validate(AgentBuildDraftCheckoutPayload) + def post( + self, + req_data: AgentBuildDraftCheckoutPayload, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + ): return AgentComposerService.checkout_agent_app_build_draft( session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, - force=args.force, + force=req_data.force, ) @@ -796,14 +806,21 @@ class AgentBuildDraftApi(Resource): @with_current_user @with_current_tenant_id @with_session - def put(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): - payload = ComposerSavePayload.model_validate(console_ns.payload or {}) + @model_validate(ComposerSavePayload) + def put( + self, + req_data: ComposerSavePayload, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + ): return AgentComposerService.save_agent_app_build_draft( session=session, tenant_id=tenant_id, agent_id=str(agent_id), account_id=current_user.id, - payload=payload, + payload=req_data, ) @console_ns.response(200, "Agent build draft discarded", console_ns.models[AgentSimpleResultResponse.__name__]) @@ -813,7 +830,7 @@ class AgentBuildDraftApi(Resource): @edit_permission_required @with_current_user @with_current_tenant_id - @with_session + @with_session(write=False) def delete(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return AgentComposerService.discard_agent_app_build_draft( session=session, @@ -833,7 +850,7 @@ class AgentBuildDraftApplyApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False) @with_current_user @with_current_tenant_id - @with_session + @with_session(write=False) def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): return AgentComposerService.apply_agent_app_build_draft( session=session, @@ -857,18 +874,25 @@ class AgentAppCopyApi(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): - args = AgentAppCopyPayload.model_validate(console_ns.payload or {}) + @model_validate(AgentAppCopyPayload) + def post( + self, + req_data: AgentAppCopyPayload, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + ): copied_app = _agent_roster_service(session).duplicate_agent_app( tenant_id=tenant_id, agent_id=str(agent_id), account=current_user, - name=args.name, - description=args.description, - role=args.role, - icon_type=args.icon_type, - icon=args.icon, - icon_background=args.icon_background, + name=req_data.name, + description=req_data.description, + role=req_data.role, + icon_type=req_data.icon_type, + icon=req_data.icon, + icon_background=req_data.icon_background, ) return _serialize_agent_app_detail(session, copied_app, current_user=current_user), 201 @@ -900,10 +924,10 @@ class AgentApiStatusApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) @with_current_tenant_id @with_session - def post(self, session: Session, tenant_id: str, agent_id: UUID): + @model_validate(AgentApiStatusPayload) + def post(self, req_data: AgentApiStatusPayload, session: Session, tenant_id: str, agent_id: UUID): app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) - args = AgentApiStatusPayload.model_validate(console_ns.payload) - app_model = AppService().update_app_api_status(app_model, args.enable_api, session=session) + app_model = AppService().update_app_api_status(app_model, req_data.enable_api, session=session) return _serialize_agent_api_access(session, app_model) @@ -917,6 +941,8 @@ class AgentApiKeyListApi(BaseApiKeyListResource): @console_ns.response(200, "Agent service API keys", console_ns.models[ApiKeyList.__name__]) @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, resource_required=False) @with_current_tenant_id + @edit_permission_required + @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) @with_session(write=False) def get(self, session: Session, tenant_id: str, agent_id: UUID) -> dict[str, object]: app_model = _resolve_agent_app_model(session, tenant_id=tenant_id, agent_id=agent_id) @@ -971,16 +997,16 @@ class AgentInviteOptionsApi(Resource): @account_initialization_required @with_current_tenant_id @with_session(write=False) - def get(self, session: Session, tenant_id: str): - query = AgentInviteOptionsQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(AgentInviteOptionsQuery) + def get(self, req_data: AgentInviteOptionsQuery, session: Session, tenant_id: str): return dump_response( AgentInviteOptionsResponse, _agent_roster_service(session).list_invite_options( tenant_id=tenant_id, - page=query.page, - limit=query.limit, - keyword=query.keyword, - app_id=query.app_id, + page=req_data.page, + limit=req_data.limit, + keyword=req_data.keyword, + app_id=req_data.app_id, ), ) @@ -1046,7 +1072,7 @@ class AgentLogMessagesApi(Resource): payload = _agent_observability_service(session).list_log_messages( app=app_model, agent_id=str(agent_id), - conversation_id=str(conversation_id), + conversation_id=str(conversation_id), # pyrefly: ignore[unnecessary-type-conversion] params=AgentLogQueryParams( page=query.page, limit=query.limit, @@ -1095,16 +1121,23 @@ class AgentStatisticsSummaryApi(Resource): @with_current_user @with_current_tenant_id @with_session(write=False) - def get(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + @model_validate(AgentStatisticsQuery) + def get( + self, + req_data: AgentStatisticsQuery, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + ): app_model = _resolve_agent_runtime_app_model(session, tenant_id=tenant_id, agent_id=agent_id) - query = AgentStatisticsQuery.model_validate(request.args.to_dict(flat=True)) timezone = current_user.timezone or "UTC" - start, end = _parse_observability_time_range(query.start, query.end, current_user) + start, end = _parse_observability_time_range(req_data.start, req_data.end, current_user) try: payload = _agent_observability_service(session).get_statistics_summary( app=app_model, agent_id=str(agent_id), - params=AgentStatisticsQueryParams(source=query.source, start=start, end=end, timezone=timezone), + params=AgentStatisticsQueryParams(source=req_data.source, start=start, end=end, timezone=timezone), ) except ValueError as exc: abort(400, description=str(exc)) diff --git a/api/controllers/console/apikey.py b/api/controllers/console/apikey.py index 01df6af5cb9..660fb63dabb 100644 --- a/api/controllers/console/apikey.py +++ b/api/controllers/console/apikey.py @@ -5,7 +5,7 @@ import flask_restx from flask_restx import Resource from flask_restx._http import HTTPStatus from pydantic import field_validator -from sqlalchemy import delete, func, select +from sqlalchemy import func, or_, select from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden @@ -21,6 +21,7 @@ from models.dataset import Dataset from models.enums import ApiTokenType from models.model import ApiToken, App from services.api_token_service import ApiTokenCache +from services.app_service import AppService from . import console_ns from .wraps import ( @@ -88,7 +89,9 @@ class BaseApiKeyListResource(Resource): _get_resource(resource_id, current_tenant_id, self.resource_model, session=session) keys = session.scalars( select(ApiToken).where( - ApiToken.type == self.resource_type, getattr(ApiToken, self.resource_id_field) == resource_id + or_(ApiToken.tenant_id == current_tenant_id, ApiToken.tenant_id.is_(None)), + ApiToken.type == self.resource_type, + getattr(ApiToken, self.resource_id_field) == resource_id, ) ).all() return ApiKeyList.model_validate({"data": keys}, from_attributes=True) @@ -103,11 +106,15 @@ class BaseApiKeyListResource(Resource): def _create_api_key(self, resource_id: str, current_tenant_id: str, *, session: Session) -> ApiToken: assert self.resource_id_field is not None, "resource_id_field must be set" - _get_resource(resource_id, current_tenant_id, self.resource_model, session=session) + resource = _get_resource(resource_id, current_tenant_id, self.resource_model, session=session) + if isinstance(resource, App): + AppService.ensure_agent_app_access_ready(resource, session=session) current_key_count: int = ( session.scalar( select(func.count(ApiToken.id)).where( - ApiToken.type == self.resource_type, getattr(ApiToken, self.resource_id_field) == resource_id + or_(ApiToken.tenant_id == current_tenant_id, ApiToken.tenant_id.is_(None)), + ApiToken.type == self.resource_type, + getattr(ApiToken, self.resource_id_field) == resource_id, ) ) or 0 @@ -169,6 +176,7 @@ class BaseApiKeyResource(Resource): key = session.scalar( select(ApiToken) .where( + or_(ApiToken.tenant_id == current_tenant_id, ApiToken.tenant_id.is_(None)), getattr(ApiToken, self.resource_id_field) == resource_id, ApiToken.type == self.resource_type, ApiToken.id == api_key_id, @@ -184,7 +192,7 @@ class BaseApiKeyResource(Resource): assert key is not None # nosec - for type checker only ApiTokenCache.delete(key.token, key.type) - session.execute(delete(ApiToken).where(ApiToken.id == api_key_id)) + session.delete(key) session.commit() @@ -195,13 +203,15 @@ class AppApiKeyListResource(BaseApiKeyListResource): @console_ns.doc(params={"resource_id": "App ID"}) @console_ns.response(200, "API keys retrieved successfully", console_ns.models[ApiKeyList.__name__]) @with_current_tenant_id + @edit_permission_required + @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION) @agent_manage_required_for_agent_app @with_session(write=False) def get(self, session: Session, current_tenant_id: str, resource_id: UUID) -> dict[str, object]: """Get all API keys for an app""" return dump_response( ApiKeyList, - self._get_api_key_list(str(resource_id), current_tenant_id, session=session), + self._get_api_key_list(str(resource_id), current_tenant_id, session=session), # pyrefly: ignore[unnecessary-type-conversion] ) @console_ns.doc("create_app_api_key") @@ -218,7 +228,7 @@ class AppApiKeyListResource(BaseApiKeyListResource): """Create a new API key for an app""" return dump_response( ApiKeyItem, - self._create_api_key(str(resource_id), current_tenant_id, session=session), + self._create_api_key(str(resource_id), current_tenant_id, session=session), # pyrefly: ignore[unnecessary-type-conversion] ), 201 resource_type = ApiTokenType.APP @@ -248,7 +258,7 @@ class AppApiKeyResource(BaseApiKeyResource): ) -> tuple[str, int]: """Delete an API key for an app""" self._delete_api_key( - str(resource_id), + str(resource_id), # pyrefly: ignore[unnecessary-type-conversion] str(api_key_id), current_tenant_id, current_user, @@ -273,7 +283,7 @@ class DatasetApiKeyListResource(BaseApiKeyListResource): """Get all API keys for a dataset""" return dump_response( ApiKeyList, - self._get_api_key_list(str(resource_id), current_tenant_id, session=session), + self._get_api_key_list(str(resource_id), current_tenant_id, session=session), # pyrefly: ignore[unnecessary-type-conversion] ) @console_ns.doc("create_dataset_api_key") @@ -289,7 +299,7 @@ class DatasetApiKeyListResource(BaseApiKeyListResource): """Create a new API key for a dataset""" return dump_response( ApiKeyItem, - self._create_api_key(str(resource_id), current_tenant_id, session=session), + self._create_api_key(str(resource_id), current_tenant_id, session=session), # pyrefly: ignore[unnecessary-type-conversion] ), 201 resource_type = ApiTokenType.DATASET @@ -318,7 +328,7 @@ class DatasetApiKeyResource(BaseApiKeyResource): ) -> tuple[str, int]: """Delete an API key for a dataset""" self._delete_api_key( - str(resource_id), + str(resource_id), # pyrefly: ignore[unnecessary-type-conversion] str(api_key_id), current_tenant_id, current_user, diff --git a/api/controllers/console/app/advanced_prompt_template.py b/api/controllers/console/app/advanced_prompt_template.py index 90098739a45..8ded3ff7234 100644 --- a/api/controllers/console/app/advanced_prompt_template.py +++ b/api/controllers/console/app/advanced_prompt_template.py @@ -1,6 +1,5 @@ from typing import Any -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field @@ -10,7 +9,7 @@ from controllers.common.schema import ( register_response_schema_models, ) from controllers.console import console_ns -from controllers.console.wraps import account_initialization_required, setup_required +from controllers.console.wraps import account_initialization_required, model_validate, setup_required from fields.base import ResponseModel from libs.login import login_required from services.advanced_prompt_template_service import AdvancedPromptTemplateArgs, AdvancedPromptTemplateService @@ -49,12 +48,12 @@ class AdvancedPromptTemplateList(Resource): @setup_required @login_required @account_initialization_required - def get(self): - args = AdvancedPromptTemplateQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(AdvancedPromptTemplateQuery) + def get(self, req_data: AdvancedPromptTemplateQuery): prompt_args: AdvancedPromptTemplateArgs = { - "app_mode": args.app_mode, - "model_mode": args.model_mode, - "model_name": args.model_name, - "has_context": args.has_context, + "app_mode": req_data.app_mode, + "model_mode": req_data.model_mode, + "model_name": req_data.model_name, + "has_context": req_data.has_context, } return AdvancedPromptTemplateService.get_prompt(prompt_args) diff --git a/api/controllers/console/app/agent.py b/api/controllers/console/app/agent.py index 8be24acc8fb..325e747b3be 100644 --- a/api/controllers/console/app/agent.py +++ b/api/controllers/console/app/agent.py @@ -21,6 +21,7 @@ from controllers.console.wraps import ( RBACPermission, RBACResourceScope, account_initialization_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -211,12 +212,12 @@ def _upload_skill_for_app(*, session: Session, current_user: Account, app_model: def _commit_drive_file_for_app(*, session: Session, current_user: Account, app_model: App, allow_node_id: bool = True): + payload = AgentDriveFilePayload.model_validate(console_ns.payload or {}) query = query_params_from_request(AgentDriveMutationQuery) node_id = query.node_id if allow_node_id else None agent_id = _resolve_agent_id(session, app_model, node_id) if not agent_id: return _agent_not_bound() - payload = AgentDriveFilePayload.model_validate(console_ns.payload or {}) upload_file = session.scalar( select(UploadFile).where( @@ -341,11 +342,11 @@ class AgentLogApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @with_session(write=False) @get_app_model(mode=[AppMode.AGENT_CHAT]) - def get(self, session: Session, app_model: App): + @model_validate(AgentLogQuery) + def get(self, req_data: AgentLogQuery, session: Session, app_model: App): """Get agent logs""" - args = AgentLogQuery.model_validate(request.args.to_dict(flat=True)) - return AgentService.get_agent_logs(app_model, args.conversation_id, args.message_id, session) + return AgentService.get_agent_logs(app_model, req_data.conversation_id, req_data.message_id, session) @console_ns.route("/agent//skills/upload") diff --git a/api/controllers/console/app/agent_app_feature.py b/api/controllers/console/app/agent_app_feature.py index 99925727335..d88496cbb06 100644 --- a/api/controllers/console/app/agent_app_feature.py +++ b/api/controllers/console/app/agent_app_feature.py @@ -25,6 +25,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -87,14 +88,21 @@ class AgentAppFeatureConfigResource(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): + @model_validate(AgentAppFeaturesPayload) + def post( + self, + req_data: AgentAppFeaturesPayload, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + ): app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) - args = AgentAppFeaturesPayload.model_validate(console_ns.payload or {}) new_app_model_config = AgentAppFeatureConfigService.update_features( app_model=app_model, account=current_user, - config=args.model_dump(exclude_none=True), + config=req_data.model_dump(exclude_none=True), session=session, ) diff --git a/api/controllers/console/app/agent_app_sandbox.py b/api/controllers/console/app/agent_app_sandbox.py index 17b5bcd23ad..c1d6c3f0f46 100644 --- a/api/controllers/console/app/agent_app_sandbox.py +++ b/api/controllers/console/app/agent_app_sandbox.py @@ -1,8 +1,8 @@ """Console routes for Agent App and workflow Agent sandbox file access. -The API keeps product-facing locators (conversation or workflow node identity) -on this public boundary and proxies list/read/upload to the agent backend's new -``/sandbox`` contract. +The API accepts product-facing Conversation, Build Draft, or Workflow Node +Execution locators and proxies list/read/download to the agent backend's +``/execution-bindings/files`` contract. """ from __future__ import annotations @@ -11,10 +11,8 @@ from typing import Literal from uuid import UUID from dify_agent.client import DifyAgentClientError, DifyAgentHTTPError, DifyAgentTimeoutError -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field -from sqlalchemy.orm import Session from controllers.common.schema import ( query_params_from_model, @@ -22,14 +20,20 @@ from controllers.common.schema import ( register_response_schema_models, register_schema_models, ) -from controllers.common.session import with_session from controllers.console import console_ns -from controllers.console.agent.app_helpers import resolve_agent_runtime_app_model +from controllers.console.app.error import AppNotFoundError from controllers.console.app.wraps import get_app_model -from controllers.console.wraps import account_initialization_required, setup_required, with_current_tenant_id +from controllers.console.wraps import ( + account_initialization_required, + model_validate, + setup_required, + with_current_tenant_id, + with_current_user, +) from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required +from models import Account from models.model import App, AppMode from services.agent_app_sandbox_service import ( AgentAppSandboxService, @@ -37,52 +41,49 @@ from services.agent_app_sandbox_service import ( WorkflowAgentSandboxService, ) -_NODE_EXECUTION_ID_DESCRIPTION = ( - "Optional workflow node execution ID. When omitted, the latest active session for the node is used." +_BINDING_PATH_DESCRIPTION = ( + "Binding path: relative paths start in Workspace; exact `~` and paths beginning with `~/` start in Home; " + "`~user` is an ordinary relative path from Workspace; absolute paths remain absolute; `..` and paths outside " + "Workspace are governed by backend isolation, not a Workspace-root restriction" ) class AgentSandboxListQuery(BaseModel): - conversation_id: str = Field(min_length=1, description="Agent App conversation ID") - path: str = Field(default=".", description="Directory path relative to the sandbox workspace") + caller_type: Literal["conversation", "build_draft"] + caller_id: str = Field(min_length=1, description="Agent App caller ID") + path: str = Field(default=".", description=_BINDING_PATH_DESCRIPTION) class AgentSandboxInfoQuery(BaseModel): - conversation_id: str = Field(min_length=1, description="Agent App conversation ID") + caller_type: Literal["conversation", "build_draft"] + caller_id: str = Field(min_length=1, description="Agent App caller ID") class AgentSandboxFileQuery(BaseModel): - conversation_id: str = Field(min_length=1, description="Agent App conversation ID") - path: str = Field(min_length=1, description="File path relative to the sandbox workspace") + caller_type: Literal["conversation", "build_draft"] + caller_id: str = Field(min_length=1, description="Agent App caller ID") + path: str = Field(min_length=1, description=_BINDING_PATH_DESCRIPTION) -class AgentSandboxUploadPayload(BaseModel): - conversation_id: str = Field(min_length=1, description="Agent App conversation ID") - path: str = Field(min_length=1, description="File path relative to the sandbox workspace") +class AgentSandboxDownloadPayload(BaseModel): + caller_type: Literal["conversation", "build_draft"] + caller_id: str = Field(min_length=1, description="Agent App caller ID") + path: str = Field(min_length=1, description=_BINDING_PATH_DESCRIPTION) class WorkflowAgentSandboxListQuery(BaseModel): - path: str = Field(default=".", description="Directory path relative to the sandbox workspace") - node_execution_id: str | None = Field( - default=None, - description=_NODE_EXECUTION_ID_DESCRIPTION, - ) + node_execution_id: str = Field(min_length=1, description="Workflow node execution ID") + path: str = Field(default=".", description=_BINDING_PATH_DESCRIPTION) class WorkflowAgentSandboxFileQuery(BaseModel): - path: str = Field(min_length=1, description="File path relative to the sandbox workspace") - node_execution_id: str | None = Field( - default=None, - description=_NODE_EXECUTION_ID_DESCRIPTION, - ) + node_execution_id: str = Field(min_length=1, description="Workflow node execution ID") + path: str = Field(min_length=1, description=_BINDING_PATH_DESCRIPTION) -class WorkflowAgentSandboxUploadPayload(BaseModel): - path: str = Field(min_length=1, description="File path relative to the sandbox workspace") - node_execution_id: str | None = Field( - default=None, - description=_NODE_EXECUTION_ID_DESCRIPTION, - ) +class WorkflowAgentSandboxDownloadPayload(BaseModel): + node_execution_id: str = Field(min_length=1, description="Workflow node execution ID") + path: str = Field(min_length=1, description=_BINDING_PATH_DESCRIPTION) class SandboxFileEntryResponse(ResponseModel): @@ -99,7 +100,6 @@ class SandboxListResponse(ResponseModel): class SandboxInfoResponse(ResponseModel): - session_id: str workspace_cwd: str @@ -111,21 +111,21 @@ class SandboxReadResponse(ResponseModel): text: str | None = None -class SandboxUploadResponse(ResponseModel): +class SandboxDownloadResponse(ResponseModel): url: str register_schema_models( console_ns, - AgentSandboxUploadPayload, - WorkflowAgentSandboxUploadPayload, + AgentSandboxDownloadPayload, + WorkflowAgentSandboxDownloadPayload, ) register_response_schema_models( console_ns, SandboxInfoResponse, SandboxListResponse, SandboxReadResponse, - SandboxUploadResponse, + SandboxDownloadResponse, ) @@ -155,15 +155,19 @@ class AgentAppSandboxInfoResource(Resource): @login_required @account_initialization_required @with_current_tenant_id - @with_session(write=False) - def get(self, session: Session, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + @with_current_user + def get(self, current_user: Account, tenant_id: str, agent_id: UUID): + service = AgentAppSandboxService() + app_id = service.resolve_app_id(tenant_id=tenant_id, agent_id=str(agent_id)) query = query_params_from_request(AgentSandboxInfoQuery) try: - result = AgentAppSandboxService().get_info( + result = service.get_info( tenant_id=tenant_id, - app_id=app_model.id, - conversation_id=query.conversation_id, + app_id=app_id, + agent_id=str(agent_id), + caller_type=query.caller_type, + caller_id=query.caller_id, + account_id=current_user.id, ) except Exception as exc: return _handle(exc) @@ -180,15 +184,19 @@ class AgentAppSandboxListResource(Resource): @login_required @account_initialization_required @with_current_tenant_id - @with_session(write=False) - def get(self, session: Session, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + @with_current_user + def get(self, current_user: Account, tenant_id: str, agent_id: UUID): + service = AgentAppSandboxService() + app_id = service.resolve_app_id(tenant_id=tenant_id, agent_id=str(agent_id)) query = query_params_from_request(AgentSandboxListQuery) try: - result = AgentAppSandboxService().list_files( + result = service.list_files( tenant_id=tenant_id, - app_id=app_model.id, - conversation_id=query.conversation_id, + app_id=app_id, + agent_id=str(agent_id), + caller_type=query.caller_type, + caller_id=query.caller_id, + account_id=current_user.id, path=query.path, ) except Exception as exc: @@ -206,15 +214,19 @@ class AgentAppSandboxReadResource(Resource): @login_required @account_initialization_required @with_current_tenant_id - @with_session(write=False) - def get(self, session: Session, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) + @with_current_user + def get(self, current_user: Account, tenant_id: str, agent_id: UUID): + service = AgentAppSandboxService() + app_id = service.resolve_app_id(tenant_id=tenant_id, agent_id=str(agent_id)) query = query_params_from_request(AgentSandboxFileQuery) try: - result = AgentAppSandboxService().read_file( + result = service.read_file( tenant_id=tenant_id, - app_id=app_model.id, - conversation_id=query.conversation_id, + app_id=app_id, + agent_id=str(agent_id), + caller_type=query.caller_type, + caller_id=query.caller_id, + account_id=current_user.id, path=query.path, ) except Exception as exc: @@ -222,26 +234,36 @@ class AgentAppSandboxReadResource(Resource): return result.model_dump() -@console_ns.route("/agent//sandbox/files/upload") -class AgentAppSandboxUploadResource(Resource): - @console_ns.doc("upload_agent_app_sandbox_file") - @console_ns.doc(description="Upload one Agent App sandbox file and return a signed download URL") - @console_ns.expect(console_ns.models[AgentSandboxUploadPayload.__name__]) - @console_ns.response(200, "Uploaded", console_ns.models[SandboxUploadResponse.__name__]) +@console_ns.route("/agent//sandbox/files/download") +class AgentAppSandboxDownloadResource(Resource): + @console_ns.doc("download_agent_app_sandbox_file") + @console_ns.doc(description="Create a ToolFile from one Agent App Binding file and return its download URL") + @console_ns.expect(console_ns.models[AgentSandboxDownloadPayload.__name__]) + @console_ns.response(200, "Download URL returned", console_ns.models[SandboxDownloadResponse.__name__]) @setup_required @login_required @account_initialization_required @with_current_tenant_id - @with_session(write=False) - def post(self, session: Session, tenant_id: str, agent_id: UUID): - app_model = resolve_agent_runtime_app_model(session=session, tenant_id=tenant_id, agent_id=agent_id) - payload = AgentSandboxUploadPayload.model_validate(request.get_json(silent=True) or {}) + @with_current_user + @model_validate(AgentSandboxDownloadPayload) + def post( + self, + req_data: AgentSandboxDownloadPayload, + current_user: Account, + tenant_id: str, + agent_id: UUID, + ): + service = AgentAppSandboxService() + app_id = service.resolve_app_id(tenant_id=tenant_id, agent_id=str(agent_id)) try: - result = AgentAppSandboxService().upload_file( + result = service.download_file( tenant_id=tenant_id, - app_id=app_model.id, - conversation_id=payload.conversation_id, - path=payload.path, + app_id=app_id, + agent_id=str(agent_id), + caller_type=req_data.caller_type, + caller_id=req_data.caller_id, + account_id=current_user.id, + path=req_data.path, ) except Exception as exc: return _handle(exc) @@ -321,29 +343,41 @@ class WorkflowAgentSandboxReadResource(Resource): @console_ns.route( - "/apps//workflow-runs//agent-nodes//sandbox/files/upload" + "/apps//workflow-runs//agent-nodes//sandbox/files/download" ) -class WorkflowAgentSandboxUploadResource(Resource): - @console_ns.doc("upload_workflow_agent_sandbox_file") - @console_ns.doc(description="Upload one workflow Agent sandbox file and return a signed download URL") - @console_ns.expect(console_ns.models[WorkflowAgentSandboxUploadPayload.__name__]) - @console_ns.response(200, "Uploaded", console_ns.models[SandboxUploadResponse.__name__]) +class WorkflowAgentSandboxDownloadResource(Resource): + @console_ns.doc("download_workflow_agent_sandbox_file") + @console_ns.doc(description="Create a ToolFile from one workflow Agent Binding file and return its download URL") + @console_ns.expect(console_ns.models[WorkflowAgentSandboxDownloadPayload.__name__]) + @console_ns.response(200, "Download URL returned", console_ns.models[SandboxDownloadResponse.__name__]) @setup_required @login_required @account_initialization_required - @get_app_model(mode=[AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) + @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, app_model: App, workflow_run_id: UUID, node_id: str): - payload = WorkflowAgentSandboxUploadPayload.model_validate(request.get_json(silent=True) or {}) + @model_validate(WorkflowAgentSandboxDownloadPayload) + def post( + self, + req_data: WorkflowAgentSandboxDownloadPayload, + tenant_id: str, + current_user: Account, + app_id: UUID, + workflow_run_id: UUID, + node_id: str, + ): + service = WorkflowAgentSandboxService() + resolved_app_id = service.resolve_app_id(tenant_id=tenant_id, app_id=str(app_id)) + if resolved_app_id is None: + raise AppNotFoundError() try: - result = WorkflowAgentSandboxService().upload_file( + result = service.download_file( tenant_id=tenant_id, - app_id=app_model.id, + app_id=resolved_app_id, workflow_run_id=str(workflow_run_id), node_id=node_id, - node_execution_id=payload.node_execution_id, - path=payload.path, - session=db.session(), + node_execution_id=req_data.node_execution_id, + account_id=current_user.id, + path=req_data.path, ) except Exception as exc: return _handle(exc) diff --git a/api/controllers/console/app/agent_config_inspector.py b/api/controllers/console/app/agent_config_inspector.py index 6fd28bcbb29..fe211db7218 100644 --- a/api/controllers/console/app/agent_config_inspector.py +++ b/api/controllers/console/app/agent_config_inspector.py @@ -30,6 +30,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -810,14 +811,21 @@ class AgentConfigFilesByAgentApi(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, tenant_id: str, current_user: Account, agent_id: UUID): - payload = AgentConfigFileUploadPayload.model_validate(console_ns.payload or {}) + @model_validate(AgentConfigFileUploadPayload) + def post( + self, + req_data: AgentConfigFileUploadPayload, + session: Session, + tenant_id: str, + current_user: Account, + agent_id: UUID, + ): return _with_agent_route_target( session=session, tenant_id=tenant_id, agent_id=agent_id, current_user=current_user, - action=lambda target: _file_upload_response(target, payload), + action=lambda target: _file_upload_response(target, req_data), ) @@ -849,13 +857,13 @@ class AgentConfigFilesApi(Resource): @with_current_user @with_session @get_app_model(mode=_WORKFLOW_APP_MODES) - def post(self, session: Session, current_user: Account, app_model: App): - payload = AgentConfigFileUploadPayload.model_validate(console_ns.payload or {}) + @model_validate(AgentConfigFileUploadPayload) + def post(self, req_data: AgentConfigFileUploadPayload, session: Session, current_user: Account, app_model: App): return _with_app_route_target( session=session, app_model=app_model, current_user=current_user, - action=lambda target: _file_upload_response(target, payload), + action=lambda target: _file_upload_response(target, req_data), ) @@ -1324,4 +1332,5 @@ class AgentConfigFileApi(Resource): ) +# pyrefly: ignore [unresolvable-dunder-all] __all__ = [name for name, value in globals().items() if inspect.isclass(value) and issubclass(value, Resource)] diff --git a/api/controllers/console/app/annotation.py b/api/controllers/console/app/annotation.py index 6653a6e288c..dcad58593a0 100644 --- a/api/controllers/console/app/annotation.py +++ b/api/controllers/console/app/annotation.py @@ -20,6 +20,7 @@ from controllers.console.wraps import ( annotation_import_rate_limit, cloud_edition_billing_resource_check, edit_permission_required, + model_validate, rbac_permission_required, setup_required, ) @@ -180,14 +181,14 @@ class AnnotationReplyActionApi(Resource): @cloud_edition_billing_resource_check("annotation") @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) - def post(self, app_id: UUID, action: Literal["enable", "disable"]): - args = AnnotationReplyPayload.model_validate(console_ns.payload) + @model_validate(AnnotationReplyPayload) + def post(self, req_data: AnnotationReplyPayload, app_id: UUID, action: Literal["enable", "disable"]): match action: case "enable": enable_args: EnableAnnotationArgs = { - "score_threshold": args.score_threshold, - "embedding_provider_name": args.embedding_provider_name, - "embedding_model_name": args.embedding_model_name, + "score_threshold": req_data.score_threshold, + "embedding_provider_name": req_data.embedding_provider_name, + "embedding_model_name": req_data.embedding_model_name, } result = AppAnnotationService.enable_app_annotation(enable_args, str(app_id)) case "disable": @@ -231,12 +232,17 @@ class AppAnnotationSettingUpdateApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_session - def post(self, session: Session, app_id: UUID, annotation_setting_id: UUID): + @model_validate(AnnotationSettingUpdatePayload) + def post( + self, + req_data: AnnotationSettingUpdatePayload, + session: Session, + app_id: UUID, + annotation_setting_id: UUID, + ): annotation_setting_id_str = str(annotation_setting_id) - args = AnnotationSettingUpdatePayload.model_validate(console_ns.payload) - - setting_args: UpdateAnnotationSettingArgs = {"score_threshold": args.score_threshold} + setting_args: UpdateAnnotationSettingArgs = {"score_threshold": req_data.score_threshold} result = AppAnnotationService.update_app_annotation_setting( str(app_id), annotation_setting_id_str, setting_args, session ) @@ -290,11 +296,11 @@ class AnnotationApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @with_session(write=False) - def get(self, session: Session, app_id: UUID): - args = AnnotationListQuery.model_validate(request.args.to_dict(flat=True)) - page = args.page - limit = args.limit - keyword = args.keyword + @model_validate(AnnotationListQuery) + def get(self, req_data: AnnotationListQuery, session: Session, app_id: UUID): + page = req_data.page + limit = req_data.limit + keyword = req_data.keyword annotation_list, total = AppAnnotationService.get_annotation_list_by_app_id( str(app_id), page, limit, keyword, session @@ -317,17 +323,17 @@ class AnnotationApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_session - def post(self, session: Session, app_id: UUID): - args = CreateAnnotationPayload.model_validate(console_ns.payload) + @model_validate(CreateAnnotationPayload) + def post(self, req_data: CreateAnnotationPayload, session: Session, app_id: UUID): upsert_args: UpsertAnnotationArgs = {} - if args.answer is not None: - upsert_args["answer"] = args.answer - if args.content is not None: - upsert_args["content"] = args.content - if args.message_id is not None: - upsert_args["message_id"] = args.message_id - if args.question is not None: - upsert_args["question"] = args.question + if req_data.answer is not None: + upsert_args["answer"] = req_data.answer + if req_data.content is not None: + upsert_args["content"] = req_data.content + if req_data.message_id is not None: + upsert_args["message_id"] = req_data.message_id + if req_data.question is not None: + upsert_args["question"] = req_data.question annotation = AppAnnotationService.up_insert_app_annotation_from_message(upsert_args, str(app_id), session) return dump_response(Annotation, annotation), 201 @@ -407,13 +413,13 @@ class AnnotationUpdateDeleteApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_session - def post(self, session: Session, app_id: UUID, annotation_id: UUID): - args = UpdateAnnotationPayload.model_validate(console_ns.payload) + @model_validate(UpdateAnnotationPayload) + def post(self, req_data: UpdateAnnotationPayload, session: Session, app_id: UUID, annotation_id: UUID): update_args: UpdateAnnotationArgs = {} - if args.answer is not None: - update_args["answer"] = args.answer - if args.question is not None: - update_args["question"] = args.question + if req_data.answer is not None: + update_args["answer"] = req_data.answer + if req_data.question is not None: + update_args["question"] = req_data.question app_ref = _get_app_ref(session, str(app_id)) annotation_ref = AppRefService.create_annotation_ref(app_ref, str(annotation_id)) annotation = AppAnnotationService.update_app_annotation_directly(update_args, annotation_ref, session) diff --git a/api/controllers/console/app/app.py b/api/controllers/console/app/app.py index 701fa366d81..203c4a3dad5 100644 --- a/api/controllers/console/app/app.py +++ b/api/controllers/console/app/app.py @@ -4,7 +4,6 @@ from collections.abc import Sequence from datetime import datetime from typing import Any, Literal -from flask import request from flask_restx import Resource from pydantic import AliasChoices, BaseModel, Field, ValidationInfo, computed_field, field_validator, model_validator from sqlalchemy import select @@ -33,6 +32,7 @@ from controllers.console.wraps import ( edit_permission_required, enterprise_license_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -58,6 +58,7 @@ from services.app_service import ( AppResponseView, AppService, CreateAppParams, + RecentAppMode, StarredAppListParams, ) from services.enterprise import rbac_service as enterprise_rbac_service @@ -139,6 +140,10 @@ class AppListBaseQuery(BaseModel): raise ValueError("Invalid UUID format in creator_ids.") from exc +class RecentAppListQuery(BaseModel): + limit: int = Field(default=8, ge=1, le=8, description="Number of recently modified apps to return (1-8)") + + class AppListQuery(AppListBaseQuery): pass @@ -411,6 +416,33 @@ class AppPartial(AppResponseModel): return to_timestamp(value) +class RecentAppResponse(ResponseModel): + id: str + name: str + icon_type: IconType | None = None + icon: str | None = None + icon_background: str | None = None + mode: RecentAppMode + author_name: str | None = None + updated_at: int + permission_keys: list[str] = Field(default_factory=list) + maintainer: str | None = None + + @computed_field(return_type=str | None) # type: ignore[prop-decorator] + @property + def icon_url(self) -> str | None: + return build_icon_url(self.icon_type, self.icon) + + @field_validator("updated_at", mode="before") + @classmethod + def _normalize_timestamp(cls, value: datetime | int) -> int: + return to_timestamp(value) + + +class RecentAppListResponse(ResponseModel): + data: list[RecentAppResponse] + + class AppDetail(AppResponseModel): id: str name: str @@ -575,6 +607,8 @@ register_schema_models( register_response_schema_models( console_ns, AppPartial, + RecentAppResponse, + RecentAppListResponse, AppDetailWithSite, AppPagination, ) @@ -663,16 +697,16 @@ class AppListApi(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, current_tenant_id: str, current_user: Account): + @model_validate(CreateAppPayload) + def post(self, req_data: CreateAppPayload, session: Session, current_tenant_id: str, current_user: Account): """Create app""" - args = CreateAppPayload.model_validate(console_ns.payload) params = CreateAppParams( - name=args.name, - description=args.description, - mode=args.mode, - icon_type=args.icon_type, - icon=args.icon, - icon_background=args.icon_background, + name=req_data.name, + description=req_data.description, + mode=req_data.mode, + icon_type=req_data.icon_type, + icon=req_data.icon, + icon_background=req_data.icon_background, ) app_service = AppService() @@ -699,6 +733,49 @@ class AppListApi(Resource): return app_detail.model_dump(mode="json"), 201 +@console_ns.route("/apps/recent") +class RecentAppListApi(Resource): + @console_ns.doc("list_recent_apps") + @console_ns.doc(description="Get recently modified apps for the home Continue Work section") + @console_ns.doc(params=query_params_from_model(RecentAppListQuery)) + @console_ns.response(200, "Success", console_ns.models[RecentAppListResponse.__name__]) + @setup_required + @login_required + @account_initialization_required + @enterprise_license_required + @with_session(write=False) + @with_current_user_id + @with_current_tenant_id + def get(self, current_tenant_id: str, current_user_id: str, session: Session): + """Return the lightweight app cards needed by the Explore home page.""" + args = query_params_from_request(RecentAppListQuery) + params = AppListParams(limit=args.limit) + + permissions = enterprise_rbac_service.RBACService.MyPermissions.get( + current_tenant_id, + current_user_id, + session=session, + ) + if dify_config.RBAC_ENABLED: + access_filter = resolve_app_access_filter( + current_tenant_id, + current_user_id, + session=session, + permissions=permissions, + ) + access_filter.apply_to_params(params) + + recent_apps = AppService().get_recent_apps(current_user_id, current_tenant_id, params, session) + permission_keys_map = permissions.app.permission_keys_by_resource_ids([app.id for app in recent_apps]) + response_items = [ + RecentAppResponse.model_validate(app, from_attributes=True).model_copy( + update={"permission_keys": permission_keys_map.get(app.id, [])} + ) + for app in recent_apps + ] + return dump_response(RecentAppListResponse, {"data": response_items}), 200 + + @console_ns.route("/apps/starred") class StarredAppListApi(Resource): @console_ns.doc("list_starred_apps") @@ -831,20 +908,20 @@ class AppApi(Resource): @agent_manage_required_for_agent_app @with_session @get_app_model(mode=None) - def put(self, session: Session, app_model: App): + @model_validate(UpdateAppPayload) + def put(self, req_data: UpdateAppPayload, session: Session, app_model: App): """Update app""" - args = UpdateAppPayload.model_validate(console_ns.payload) app_service = AppService() args_dict: AppService.ArgsDict = { - "name": args.name, - "description": args.description or "", - "icon_type": args.icon_type, - "icon": args.icon or "", - "icon_background": args.icon_background or "", - "use_icon_as_answer_icon": args.use_icon_as_answer_icon or False, - "max_active_requests": args.max_active_requests or 0, + "name": req_data.name, + "description": req_data.description or "", + "icon_type": req_data.icon_type, + "icon": req_data.icon or "", + "icon_background": req_data.icon_background or "", + "use_icon_as_answer_icon": req_data.use_icon_as_answer_icon or False, + "max_active_requests": req_data.max_active_requests or 0, } app_model = app_service.update_app(app_model, args_dict, session=session) return AppDetailWithSite.model_validate( @@ -892,10 +969,10 @@ class AppCopyApi(Resource): @with_current_user @with_current_tenant_id @get_app_model(mode=None) - def post(self, current_tenant_id: str, current_user: Account, app_model: App): + @model_validate(CopyAppPayload) + def post(self, req_data: CopyAppPayload, current_tenant_id: str, current_user: Account, app_model: App): """Copy app""" # The role of the current user in the ta table must be admin, owner, or editor - args = CopyAppPayload.model_validate(console_ns.payload or {}) with Session(db.engine, expire_on_commit=False) as session: import_service = AppDslService(session) @@ -905,11 +982,11 @@ class AppCopyApi(Resource): account=current_user, import_mode=ImportMode.YAML_CONTENT, yaml_content=yaml_content, - name=args.name, - description=args.description, - icon_type=args.icon_type, - icon=args.icon, - icon_background=args.icon_background, + name=req_data.name, + description=req_data.description, + icon_type=req_data.icon_type, + icon=req_data.icon, + icon_background=req_data.icon_background, ) except NoPermissionError as e: raise Forbidden(str(e)) @@ -968,16 +1045,16 @@ class AppExportApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_IMPORT_EXPORT_DSL) @agent_manage_required_for_agent_app @get_app_model - def get(self, app_model: App): + @model_validate(AppExportQuery) + def get(self, req_data: AppExportQuery, app_model: App): """Export app""" - args = AppExportQuery.model_validate(request.args.to_dict(flat=True)) response = AppExportResponse( data=AppDslService.export_dsl( app_model=app_model, session=db.session(), - include_secret=args.include_secret, - workflow_id=args.workflow_id, + include_secret=req_data.include_secret, + workflow_id=req_data.workflow_id, ) ) return response.model_dump(mode="json") @@ -1025,11 +1102,11 @@ class AppNameApi(Resource): @agent_manage_required_for_agent_app @with_session @get_app_model(mode=None) - def post(self, session: Session, app_model: App): - args = AppNamePayload.model_validate(console_ns.payload) + @model_validate(AppNamePayload) + def post(self, req_data: AppNamePayload, session: Session, app_model: App): app_service = AppService() - app_model = app_service.update_app_name(app_model, args.name, session=session) + app_model = app_service.update_app_name(app_model, req_data.name, session=session) return AppDetail.model_validate( app_model, from_attributes=True, @@ -1053,15 +1130,15 @@ class AppIconApi(Resource): @agent_manage_required_for_agent_app @with_session @get_app_model(mode=None) - def post(self, session: Session, app_model: App): - args = AppIconPayload.model_validate(console_ns.payload or {}) + @model_validate(AppIconPayload) + def post(self, req_data: AppIconPayload, session: Session, app_model: App): app_service = AppService() app_model = app_service.update_app_icon( app_model, - args.icon or "", - args.icon_background or "", - args.icon_type, + req_data.icon or "", + req_data.icon_background or "", + req_data.icon_type, session=session, ) return AppDetail.model_validate( @@ -1087,11 +1164,11 @@ class AppSiteStatus(Resource): @agent_manage_required_for_agent_app @with_session @get_app_model(mode=None) - def post(self, session: Session, app_model: App): - args = AppSiteStatusPayload.model_validate(console_ns.payload) + @model_validate(AppSiteStatusPayload) + def post(self, req_data: AppSiteStatusPayload, session: Session, app_model: App): app_service = AppService() - app_model = app_service.update_app_site_status(app_model, args.enable_site, session=session) + app_model = app_service.update_app_site_status(app_model, req_data.enable_site, session=session) return AppDetail.model_validate( app_model, from_attributes=True, @@ -1115,11 +1192,11 @@ class AppApiStatus(Resource): @agent_manage_required_for_agent_app @with_session @get_app_model(mode=None) - def post(self, session: Session, app_model: App): - args = AppApiStatusPayload.model_validate(console_ns.payload) + @model_validate(AppApiStatusPayload) + def post(self, req_data: AppApiStatusPayload, session: Session, app_model: App): app_service = AppService() - app_model = app_service.update_app_api_status(app_model, args.enable_api, session=session) + app_model = app_service.update_app_api_status(app_model, req_data.enable_api, session=session) return AppDetail.model_validate( app_model, from_attributes=True, @@ -1165,14 +1242,14 @@ class AppTraceApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG) @get_app_model - def post(self, app_model: App): + @model_validate(AppTracePayload) + def post(self, req_data: AppTracePayload, app_model: App): # add app trace - args = AppTracePayload.model_validate(console_ns.payload) OpsTraceManager.update_app_tracing_config( app_id=app_model.id, - enabled=args.enabled, - tracing_provider=args.tracing_provider, + enabled=req_data.enabled, + tracing_provider=req_data.tracing_provider, ) return SimpleResultResponse(result="success").model_dump(mode="json") diff --git a/api/controllers/console/app/app_import.py b/api/controllers/console/app/app_import.py index c03e738bf8a..f2d3bace841 100644 --- a/api/controllers/console/app/app_import.py +++ b/api/controllers/console/app/app_import.py @@ -12,6 +12,7 @@ from controllers.console.wraps import ( account_initialization_required, cloud_edition_billing_resource_check, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_user, @@ -83,8 +84,8 @@ class AppImportApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_IMPORT_EXPORT_DSL, resource_required=False) @with_current_user - def post(self, current_user: Account | None = None): - args = AppImportPayload.model_validate(console_ns.payload) + @model_validate(AppImportPayload) + def post(self, req_data: AppImportPayload, current_user: Account | None = None): current_user = current_user if current_user is not None else _current_user_and_tenant_id(None)[0] # AppDslService performs internal commits for some creation paths, so use a plain @@ -96,15 +97,15 @@ class AppImportApi(Resource): try: result = import_service.import_app( account=account, - import_mode=args.mode, - yaml_content=args.yaml_content, - yaml_url=args.yaml_url, - name=args.name, - description=args.description, - icon_type=args.icon_type, - icon=args.icon, - icon_background=args.icon_background, - app_id=args.app_id, + import_mode=req_data.mode, + yaml_content=req_data.yaml_content, + yaml_url=req_data.yaml_url, + name=req_data.name, + description=req_data.description, + icon_type=req_data.icon_type, + icon=req_data.icon, + icon_background=req_data.icon_background, + app_id=req_data.app_id, ) except NoPermissionError as e: raise Forbidden(str(e)) @@ -113,7 +114,7 @@ class AppImportApi(Resource): else: session.commit() - is_created_app = args.app_id is None and result.status in { + is_created_app = req_data.app_id is None and result.status in { ImportStatus.COMPLETED, ImportStatus.COMPLETED_WITH_WARNINGS, } diff --git a/api/controllers/console/app/audio.py b/api/controllers/console/app/audio.py index 059cc96269d..b476817cf08 100644 --- a/api/controllers/console/app/audio.py +++ b/api/controllers/console/app/audio.py @@ -31,6 +31,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -276,15 +277,15 @@ class ChatMessageTextApi(Resource): @login_required @account_initialization_required @get_app_model - def post(self, app_model: App): + @model_validate(TextToSpeechPayload) + def post(self, req_data: TextToSpeechPayload, app_model: App): try: - payload = TextToSpeechPayload.model_validate(console_ns.payload) message_ref = None - if payload.message_id: + if req_data.message_id: app_ref = AppRefService.create_app_ref(app_model) message_ref = AppRefService.create_message_ref( app_ref, - payload.message_id, + req_data.message_id, account_id=current_user.id, ) @@ -292,8 +293,8 @@ class ChatMessageTextApi(Resource): return AudioService.transcript_tts( app_model=app_model, session=db.session(), - text=payload.text, - voice=payload.voice, + text=req_data.text, + voice=req_data.voice, message_ref=message_ref, is_draft=True, ) @@ -339,13 +340,12 @@ class TextModesApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @get_app_model - def get(self, app_model: App): + @model_validate(TextToSpeechVoiceQuery) + def get(self, req_data: TextToSpeechVoiceQuery, app_model: App): try: - args = TextToSpeechVoiceQuery.model_validate(request.args.to_dict(flat=True)) - response = AudioService.transcript_tts_voices( tenant_id=app_model.tenant_id, - language=args.language, + language=req_data.language, ) return dump_response(TextToSpeechVoiceListResponse, response) diff --git a/api/controllers/console/app/completion.py b/api/controllers/console/app/completion.py index beca00d621b..cfa18235cd0 100644 --- a/api/controllers/console/app/completion.py +++ b/api/controllers/console/app/completion.py @@ -29,6 +29,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -36,7 +37,7 @@ from controllers.console.wraps import ( with_current_user_id, ) from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError -from core.app.entities.app_invoke_entities import AGENT_RUNTIME_EXIT_INTENT_ARG, InvokeFrom +from core.app.entities.app_invoke_entities import InvokeFrom from core.app.features.rate_limiting.rate_limit import RateLimitGenerator from core.errors.error import ( ModelCurrentlyNotSupportError, @@ -160,11 +161,11 @@ class CompletionMessageApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TEST_AND_RUN) @with_session @get_app_model(mode=AppMode.COMPLETION) - def post(self, session: Session, current_user: Account, app_model: App): - args_model = CompletionMessagePayload.model_validate(console_ns.payload) - args = args_model.model_dump(exclude_none=True, by_alias=True) + @model_validate(CompletionMessagePayload) + def post(self, req_data: CompletionMessagePayload, session: Session, current_user: Account, app_model: App): + args = req_data.model_dump(exclude_none=True, by_alias=True) - streaming = args_model.response_mode != "blocking" + streaming = req_data.response_mode != "blocking" args["auto_generate_name"] = False try: @@ -353,12 +354,7 @@ def _resolve_current_user_agent_debug_conversation_id( draft_type: AgentConfigDraftType, start_new: bool = False, ) -> str: - """Resolve or rotate the current editor's conversation within one draft surface. - - ``start_new`` rotates the scoped mapping through ``AgentRosterService`` so - the old runtime session is retired before the new conversation is used. - Continuations and Build chat keep resolving the existing mapping. - """ + """Resolve the current editor's Build or Preview conversation.""" roster_service = AgentRosterService(session) resolved_agent_id = agent_id @@ -368,17 +364,26 @@ def _resolve_current_user_agent_debug_conversation_id( raise AgentNotFoundError() resolved_agent_id = agent.id - resolve_conversation = ( - roster_service.refresh_agent_app_debug_conversation_id - if start_new - else roster_service.get_or_create_agent_app_debug_conversation_id - ) - return resolve_conversation( + if draft_type == AgentConfigDraftType.DEBUG_BUILD: + return roster_service.get_or_create_build_conversation( + tenant_id=current_tenant_id, + agent_id=resolved_agent_id, + account_id=current_user.id, + ) + if start_new: + return roster_service.rotate_preview_conversation( + tenant_id=current_tenant_id, + agent_id=resolved_agent_id, + account_id=current_user.id, + ) + conversation_id = roster_service.get_current_preview_conversation( tenant_id=current_tenant_id, agent_id=resolved_agent_id, account_id=current_user.id, - draft_type=draft_type, ) + if conversation_id is None: + raise NotFound("Conversation Not Exists.") + return conversation_id def _create_chat_message( @@ -450,7 +455,6 @@ def _create_build_chat_finalization_message( "draft_type": "debug_build", "conversation_id": debug_conversation_id, "auto_generate_name": False, - AGENT_RUNTIME_EXIT_INTENT_ARG: "delete", } external_trace_id = get_external_trace_id(request) if external_trace_id: diff --git a/api/controllers/console/app/conversation.py b/api/controllers/console/app/conversation.py index 14ddb2446b8..325f3ef22d7 100644 --- a/api/controllers/console/app/conversation.py +++ b/api/controllers/console/app/conversation.py @@ -2,7 +2,7 @@ from typing import Literal from uuid import UUID import sqlalchemy as sa -from flask import abort, request +from flask import abort from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import func, or_ @@ -18,6 +18,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_user, @@ -108,17 +109,17 @@ class CompletionConversationApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @with_session(write=False) @get_app_model(mode=AppMode.COMPLETION) - def get(self, session: Session, current_user: Account, app_model: App): - args = CompletionConversationQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(CompletionConversationQuery) + def get(self, req_data: CompletionConversationQuery, session: Session, current_user: Account, app_model: App): query = sa.select(Conversation).where( Conversation.app_id == app_model.id, Conversation.mode == "completion", Conversation.is_deleted.is_(False) ) - if args.keyword: + if req_data.keyword: from libs.helper import escape_like_pattern - escaped_keyword = escape_like_pattern(args.keyword) + escaped_keyword = escape_like_pattern(req_data.keyword) query = query.join(Message, Message.conversation_id == Conversation.id).where( or_( Message.query.ilike(f"%{escaped_keyword}%", escape="\\"), @@ -130,7 +131,7 @@ class CompletionConversationApi(Resource): assert account.timezone is not None try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -142,7 +143,7 @@ class CompletionConversationApi(Resource): query = query.where(Conversation.created_at < end_datetime_utc) # FIXME, the type ignore in this file - if args.annotation_status == "annotated": + if req_data.annotation_status == "annotated": query = ( query.options(selectinload(Conversation.message_annotations)) # type: ignore[arg-type] .join( # type: ignore @@ -150,7 +151,7 @@ class CompletionConversationApi(Resource): ) .group_by(Conversation.id) ) - elif args.annotation_status == "not_annotated": + elif req_data.annotation_status == "not_annotated": query = ( query.outerjoin(MessageAnnotation, MessageAnnotation.conversation_id == Conversation.id) .group_by(Conversation.id) @@ -159,7 +160,7 @@ class CompletionConversationApi(Resource): query = query.order_by(Conversation.created_at.desc()) - conversations = paginate_query(query, session=session, page=args.page, per_page=args.limit) + conversations = paginate_query(query, session=session, page=req_data.page, per_page=req_data.limit) return dump_response( ConversationPaginationResponse, @@ -238,8 +239,8 @@ class ChatConversationApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @with_session(write=False) @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT]) - def get(self, session: Session, current_user: Account, app_model: App): - args = ChatConversationQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(ChatConversationQuery) + def get(self, req_data: ChatConversationQuery, session: Session, current_user: Account, app_model: App): subquery = ( sa.select(Conversation.id.label("conversation_id"), EndUser.session_id.label("from_end_user_session_id")) @@ -249,10 +250,10 @@ class ChatConversationApi(Resource): query = sa.select(Conversation).where(Conversation.app_id == app_model.id, Conversation.is_deleted.is_(False)) - if args.keyword: + if req_data.keyword: from libs.helper import escape_like_pattern - escaped_keyword = escape_like_pattern(args.keyword) + escaped_keyword = escape_like_pattern(req_data.keyword) keyword_filter = f"%{escaped_keyword}%" query = ( query.join( @@ -276,12 +277,12 @@ class ChatConversationApi(Resource): assert account.timezone is not None try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) if start_datetime_utc: - match args.sort_by: + match req_data.sort_by: case "updated_at" | "-updated_at": query = query.where(Conversation.updated_at >= start_datetime_utc) case "created_at" | "-created_at" | _: @@ -289,13 +290,13 @@ class ChatConversationApi(Resource): if end_datetime_utc: end_datetime_utc = end_datetime_utc.replace(second=59) - match args.sort_by: + match req_data.sort_by: case "updated_at" | "-updated_at": query = query.where(Conversation.updated_at <= end_datetime_utc) case "created_at" | "-created_at" | _: query = query.where(Conversation.created_at <= end_datetime_utc) - match args.annotation_status: + match req_data.annotation_status: case "annotated": query = ( query.options(selectinload(Conversation.message_annotations)) # type: ignore[arg-type] @@ -316,7 +317,7 @@ class ChatConversationApi(Resource): if app_model.mode == AppMode.ADVANCED_CHAT: query = query.where(Conversation.invoke_from != InvokeFrom.DEBUGGER) - match args.sort_by: + match req_data.sort_by: case "created_at": query = query.order_by(Conversation.created_at.asc()) case "-created_at": @@ -328,7 +329,7 @@ class ChatConversationApi(Resource): case _: query = query.order_by(Conversation.created_at.desc()) - conversations = paginate_query(query, session=session, page=args.page, per_page=args.limit) + conversations = paginate_query(query, session=session, page=req_data.page, per_page=req_data.limit) return dump_response( ConversationWithSummaryPaginationResponse, diff --git a/api/controllers/console/app/conversation_variables.py b/api/controllers/console/app/conversation_variables.py index aa8090f0440..ff0727039aa 100644 --- a/api/controllers/console/app/conversation_variables.py +++ b/api/controllers/console/app/conversation_variables.py @@ -3,7 +3,6 @@ from __future__ import annotations from datetime import datetime from typing import Any -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import select @@ -16,6 +15,7 @@ from controllers.console.wraps import ( RBACPermission, RBACResourceScope, account_initialization_required, + model_validate, rbac_permission_required, setup_required, ) @@ -101,15 +101,15 @@ class ConversationVariablesApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_CREATE_AND_MANAGEMENT) @get_app_model(mode=AppMode.ADVANCED_CHAT) - def get(self, app_model: App): - args = ConversationVariablesQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(ConversationVariablesQuery) + def get(self, req_data: ConversationVariablesQuery, app_model: App): stmt = ( select(ConversationVariable) .where(ConversationVariable.app_id == app_model.id) .order_by(ConversationVariable.created_at) ) - stmt = stmt.where(ConversationVariable.conversation_id == args.conversation_id) + stmt = stmt.where(ConversationVariable.conversation_id == req_data.conversation_id) # NOTE: This is a temporary solution to avoid performance issues. page = 1 diff --git a/api/controllers/console/app/generator.py b/api/controllers/console/app/generator.py index 495d1403d7e..a899e539d9e 100644 --- a/api/controllers/console/app/generator.py +++ b/api/controllers/console/app/generator.py @@ -17,7 +17,12 @@ from controllers.console.app.error import ( ProviderQuotaExceededError, ) from controllers.console.app.wraps import with_session -from controllers.console.wraps import account_initialization_required, setup_required, with_current_tenant_id +from controllers.console.wraps import ( + account_initialization_required, + model_validate, + setup_required, + with_current_tenant_id, +) from core.app.app_config.entities import ModelConfig from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from core.helper.code_executor.code_node_provider import CodeNodeProvider @@ -238,11 +243,11 @@ class RuleGenerateApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): - args = RuleGeneratePayload.model_validate(console_ns.payload) + @model_validate(RuleGeneratePayload) + def post(self, req_data: RuleGeneratePayload, current_tenant_id: str): try: - rules = LLMGenerator.generate_rule_config(tenant_id=current_tenant_id, args=args) + rules = LLMGenerator.generate_rule_config(tenant_id=current_tenant_id, args=req_data) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) except QuotaExceededError: @@ -267,13 +272,13 @@ class RuleCodeGenerateApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): - args = RuleCodeGeneratePayload.model_validate(console_ns.payload) + @model_validate(RuleCodeGeneratePayload) + def post(self, req_data: RuleCodeGeneratePayload, current_tenant_id: str): try: code_result = LLMGenerator.generate_code( tenant_id=current_tenant_id, - args=args, + args=req_data, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -299,13 +304,13 @@ class RuleStructuredOutputGenerateApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): - args = RuleStructuredOutputPayload.model_validate(console_ns.payload) + @model_validate(RuleStructuredOutputPayload) + def post(self, req_data: RuleStructuredOutputPayload, current_tenant_id: str): try: structured_output = LLMGenerator.generate_structured_output( tenant_id=current_tenant_id, - args=args, + args=req_data, ) except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) @@ -332,36 +337,36 @@ class InstructionGenerateApi(Resource): @account_initialization_required @with_current_tenant_id @with_session(write=False) - def post(self, session: Session, current_tenant_id: str): - args = InstructionGeneratePayload.model_validate(console_ns.payload) + @model_validate(InstructionGeneratePayload) + def post(self, req_data: InstructionGeneratePayload, session: Session, current_tenant_id: str): providers: list[type[CodeNodeProvider]] = [Python3CodeProvider, JavascriptCodeProvider] code_provider: type[CodeNodeProvider] | None = next( - (p for p in providers if p.is_accept_language(args.language)), None + (p for p in providers if p.is_accept_language(req_data.language)), None ) code_template = code_provider.get_default_code() if code_provider else "" try: # Generate from nothing for a workflow node - if (args.current in (code_template, "")) and args.node_id != "": + if (req_data.current in (code_template, "")) and req_data.node_id != "": app = session.scalar( - select(App).where(App.id == args.flow_id, App.tenant_id == current_tenant_id).limit(1) + select(App).where(App.id == req_data.flow_id, App.tenant_id == current_tenant_id).limit(1) ) if not app: - return {"error": f"app {args.flow_id} not found"}, 400 + return {"error": f"app {req_data.flow_id} not found"}, 400 workflow = WorkflowService().get_draft_workflow(app_model=app, session=session) if not workflow: - return {"error": f"workflow {args.flow_id} not found"}, 400 + return {"error": f"workflow {req_data.flow_id} not found"}, 400 nodes: Sequence = workflow.graph_dict["nodes"] - node = [node for node in nodes if node["id"] == args.node_id] + node = [node for node in nodes if node["id"] == req_data.node_id] if len(node) == 0: - return {"error": f"node {args.node_id} not found"}, 400 + return {"error": f"node {req_data.node_id} not found"}, 400 node_type = node[0]["data"]["type"] match node_type: case "llm": return LLMGenerator.generate_rule_config( current_tenant_id, args=RuleGeneratePayload( - instruction=args.instruction, - model_config=args.model_config_data, + instruction=req_data.instruction, + model_config=req_data.model_config_data, no_variable=True, ), ) @@ -369,8 +374,8 @@ class InstructionGenerateApi(Resource): return LLMGenerator.generate_rule_config( current_tenant_id, args=RuleGeneratePayload( - instruction=args.instruction, - model_config=args.model_config_data, + instruction=req_data.instruction, + model_config=req_data.model_config_data, no_variable=True, ), ) @@ -378,31 +383,31 @@ class InstructionGenerateApi(Resource): return LLMGenerator.generate_code( tenant_id=current_tenant_id, args=RuleCodeGeneratePayload( - instruction=args.instruction, - model_config=args.model_config_data, - code_language=args.language, + instruction=req_data.instruction, + model_config=req_data.model_config_data, + code_language=req_data.language, ), ) case _: return {"error": f"invalid node type: {node_type}"} - if args.node_id == "" and args.current != "": # For legacy app without a workflow + if req_data.node_id == "" and req_data.current != "": # For legacy app without a workflow return LLMGenerator.instruction_modify_legacy( tenant_id=current_tenant_id, - flow_id=args.flow_id, - current=args.current, - instruction=args.instruction, - model_config=args.model_config_data, - ideal_output=args.ideal_output, + flow_id=req_data.flow_id, + current=req_data.current, + instruction=req_data.instruction, + model_config=req_data.model_config_data, + ideal_output=req_data.ideal_output, ) - if args.node_id != "" and args.current != "": # For workflow node + if req_data.node_id != "" and req_data.current != "": # For workflow node return LLMGenerator.instruction_modify_workflow( tenant_id=current_tenant_id, - flow_id=args.flow_id, - node_id=args.node_id, - current=args.current, - instruction=args.instruction, - model_config=args.model_config_data, - ideal_output=args.ideal_output, + flow_id=req_data.flow_id, + node_id=req_data.node_id, + current=req_data.current, + instruction=req_data.instruction, + model_config=req_data.model_config_data, + ideal_output=req_data.ideal_output, workflow_service=WorkflowService(), ) return {"error": "incompatible parameters"}, 400 @@ -426,9 +431,9 @@ class InstructionGenerationTemplateApi(Resource): @setup_required @login_required @account_initialization_required - def post(self): - args = InstructionTemplatePayload.model_validate(console_ns.payload) - match args.type: + @model_validate(InstructionTemplatePayload) + def post(self, req_data: InstructionTemplatePayload): + match req_data.type: case "prompt": from core.llm_generator.prompts import INSTRUCTION_GENERATE_TEMPLATE_PROMPT @@ -438,7 +443,7 @@ class InstructionGenerationTemplateApi(Resource): return {"data": INSTRUCTION_GENERATE_TEMPLATE_CODE} case _: - raise ValueError(f"Invalid type: {args.type}") + raise ValueError(f"Invalid type: {req_data.type}") def _workflow_instruction_guard(args: WorkflowGeneratePayload) -> tuple[dict, int] | None: @@ -492,24 +497,24 @@ class WorkflowGenerateApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): - args = WorkflowGeneratePayload.model_validate(console_ns.payload) + @model_validate(WorkflowGeneratePayload) + def post(self, req_data: WorkflowGeneratePayload, current_tenant_id: str): # Reject empty / over-length instructions at the boundary (shared with # the streaming endpoint) before spending a planner+builder roundtrip. - guard = _workflow_instruction_guard(args) + guard = _workflow_instruction_guard(req_data) if guard is not None: return guard try: result = WorkflowGeneratorService.generate_workflow_graph( tenant_id=current_tenant_id, - mode=args.mode, - instruction=args.instruction, - model_config=args.model_config_data, - ideal_output=args.ideal_output, - current_graph=args.current_graph.model_dump(by_alias=True, exclude_none=True) - if args.current_graph + mode=req_data.mode, + instruction=req_data.instruction, + model_config=req_data.model_config_data, + ideal_output=req_data.ideal_output, + current_graph=req_data.current_graph.model_dump(by_alias=True, exclude_none=True) + if req_data.current_graph else None, ) except ProviderTokenNotInitError as ex: @@ -547,13 +552,13 @@ class WorkflowInstructionSuggestionsApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): - args = WorkflowInstructionSuggestionsPayload.model_validate(console_ns.payload) + @model_validate(WorkflowInstructionSuggestionsPayload) + def post(self, req_data: WorkflowInstructionSuggestionsPayload, current_tenant_id: str): suggestions = LLMGenerator.generate_workflow_instruction_suggestions( tenant_id=current_tenant_id, - mode=args.mode, - language=args.language, - count=args.count, + mode=req_data.mode, + language=req_data.language, + count=req_data.count, ) return dump_response(WorkflowInstructionSuggestionsResponse, {"suggestions": suggestions}) @@ -583,12 +588,12 @@ class WorkflowGenerateStreamApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): - args = WorkflowGeneratePayload.model_validate(console_ns.payload) + @model_validate(WorkflowGeneratePayload) + def post(self, req_data: WorkflowGeneratePayload, current_tenant_id: str): # Same boundary guards as the blocking endpoint — return a normal 400 # JSON for these BEFORE opening the stream. - guard = _workflow_instruction_guard(args) + guard = _workflow_instruction_guard(req_data) if guard is not None: return guard @@ -596,12 +601,14 @@ class WorkflowGenerateStreamApi(Resource): try: for event_name, payload in WorkflowGeneratorService.generate_workflow_graph_stream( tenant_id=current_tenant_id, - mode=args.mode, - instruction=args.instruction, - model_config=args.model_config_data, - ideal_output=args.ideal_output, + mode=req_data.mode, + instruction=req_data.instruction, + model_config=req_data.model_config_data, + ideal_output=req_data.ideal_output, current_graph=( - args.current_graph.model_dump(by_alias=True, exclude_none=True) if args.current_graph else None + req_data.current_graph.model_dump(by_alias=True, exclude_none=True) + if req_data.current_graph + else None ), ): body = {"event": event_name, **payload} diff --git a/api/controllers/console/app/mcp_server.py b/api/controllers/console/app/mcp_server.py index f5ab8b103ad..d03e27b7fae 100644 --- a/api/controllers/console/app/mcp_server.py +++ b/api/controllers/console/app/mcp_server.py @@ -1,7 +1,6 @@ import json from datetime import datetime from typing import Any -from uuid import UUID from flask_restx import Resource from pydantic import BaseModel, Field, field_validator @@ -16,6 +15,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -110,17 +110,16 @@ class AppMCPServerController(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_current_tenant_id @get_app_model - def post(self, current_tenant_id: str, app_model: App): - payload = MCPServerCreatePayload.model_validate(console_ns.payload or {}) - - description = payload.description + @model_validate(MCPServerCreatePayload) + def post(self, req_data: MCPServerCreatePayload, current_tenant_id: str, app_model: App): + description = req_data.description if not description: description = app_model.description or "" server = AppMCPServer( name=app_model.name, description=description, - parameters=json.dumps(payload.parameters, ensure_ascii=False), + parameters=json.dumps(req_data.parameters, ensure_ascii=False), status=AppMCPServerStatus.ACTIVE, app_id=app_model.id, tenant_id=current_tenant_id, @@ -145,10 +144,10 @@ class AppMCPServerController(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @get_app_model - def put(self, app_model: App): - payload = MCPServerUpdatePayload.model_validate(console_ns.payload or {}) + @model_validate(MCPServerUpdatePayload) + def put(self, req_data: MCPServerUpdatePayload, app_model: App): app_ref = AppRefService.create_app_ref(app_model) - server_ref = AppRefService.create_mcp_server_ref(app_ref, payload.id) + server_ref = AppRefService.create_mcp_server_ref(app_ref, req_data.id) server = db.session.scalar( select(AppMCPServer) .where( @@ -161,7 +160,7 @@ class AppMCPServerController(Resource): if not server: raise NotFound() - description = payload.description + description = req_data.description if description is None or not description: server.description = app_model.description or "" else: @@ -169,21 +168,21 @@ class AppMCPServerController(Resource): server.name = app_model.name - server.parameters = json.dumps(payload.parameters, ensure_ascii=False) - if payload.status: + server.parameters = json.dumps(req_data.parameters, ensure_ascii=False) + if req_data.status: try: - server.status = AppMCPServerStatus(payload.status) + server.status = AppMCPServerStatus(req_data.status) except ValueError: raise ValueError("Invalid status") db.session.commit() return dump_response(AppMCPServerResponse, server) -@console_ns.route("/apps//server/refresh") +@console_ns.route("/apps//server/refresh") class AppMCPServerRefreshController(Resource): @console_ns.doc("refresh_app_mcp_server") @console_ns.doc(description="Refresh MCP server configuration and regenerate server code") - @console_ns.doc(params={"server_id": "Server ID"}) + @console_ns.doc(params={"app_id": "App ID"}) @console_ns.response(200, "MCP server refreshed successfully", console_ns.models[AppMCPServerResponse.__name__]) @console_ns.response(403, "Insufficient permissions") @console_ns.response(404, "Server not found") @@ -193,10 +192,11 @@ class AppMCPServerRefreshController(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @with_current_tenant_id - def get(self, current_tenant_id: str, server_id: UUID): + @get_app_model + def post(self, current_tenant_id: str, app_model: App): server = db.session.scalar( select(AppMCPServer) - .where(AppMCPServer.id == server_id, AppMCPServer.tenant_id == current_tenant_id) + .where(AppMCPServer.app_id == app_model.id, AppMCPServer.tenant_id == current_tenant_id) .limit(1) ) if not server: diff --git a/api/controllers/console/app/message.py b/api/controllers/console/app/message.py index 080772107dc..8f6c2f54464 100644 --- a/api/controllers/console/app/message.py +++ b/api/controllers/console/app/message.py @@ -28,6 +28,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -320,8 +321,8 @@ class MessageFeedbackExportApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @get_app_model - def get(self, app_model: App): - args = FeedbackExportQuery.model_validate(request.args.to_dict()) + @model_validate(FeedbackExportQuery) + def get(self, req_data: FeedbackExportQuery, app_model: App): # Import the service function from services.feedback_service import FeedbackService @@ -330,12 +331,12 @@ class MessageFeedbackExportApi(Resource): export_data = FeedbackService.export_feedbacks( app_model.id, session=db.session(), - from_source=args.from_source, - rating=args.rating, - has_comment=args.has_comment, - start_date=args.start_date, - end_date=args.end_date, - format_type=args.format, + from_source=req_data.from_source, + rating=req_data.rating, + has_comment=req_data.has_comment, + start_date=req_data.start_date, + end_date=req_data.end_date, + format_type=req_data.format, ) return export_data diff --git a/api/controllers/console/app/ops_trace.py b/api/controllers/console/app/ops_trace.py index e86f65fc035..2332f164cd9 100644 --- a/api/controllers/console/app/ops_trace.py +++ b/api/controllers/console/app/ops_trace.py @@ -1,6 +1,5 @@ from typing import Any -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field from werkzeug.exceptions import BadRequest @@ -14,6 +13,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, ) @@ -74,12 +74,11 @@ class TraceAppConfigApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG) @get_app_model - def get(self, app_model: App): - args = TraceProviderQuery.model_validate(request.args.to_dict(flat=True)) # type: ignore - + @model_validate(TraceProviderQuery) + def get(self, req_data: TraceProviderQuery, app_model: App): try: trace_config = OpsService.get_tracing_app_config( - app_id=app_model.id, tracing_provider=args.tracing_provider, session=db.session() + app_id=app_model.id, tracing_provider=req_data.tracing_provider, session=db.session() ) if not trace_config: return {"has_not_configured": True} @@ -104,15 +103,14 @@ class TraceAppConfigApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG) @get_app_model - def post(self, app_model: App): + @model_validate(TraceConfigPayload) + def post(self, req_data: TraceConfigPayload, app_model: App): """Create a new trace app configuration""" - args = TraceConfigPayload.model_validate(console_ns.payload) - try: result = OpsService.create_tracing_app_config( app_id=app_model.id, - tracing_provider=args.tracing_provider, - tracing_config=args.tracing_config, + tracing_provider=req_data.tracing_provider, + tracing_config=req_data.tracing_config, session=db.session(), ) if not result: @@ -140,15 +138,14 @@ class TraceAppConfigApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG) @get_app_model - def patch(self, app_model: App): + @model_validate(TraceConfigPayload) + def patch(self, req_data: TraceConfigPayload, app_model: App): """Update an existing trace app configuration""" - args = TraceConfigPayload.model_validate(console_ns.payload) - try: result = OpsService.update_tracing_app_config( app_id=app_model.id, - tracing_provider=args.tracing_provider, - tracing_config=args.tracing_config, + tracing_provider=req_data.tracing_provider, + tracing_config=req_data.tracing_config, session=db.session(), ) if not result: @@ -170,13 +167,12 @@ class TraceAppConfigApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_TRACING_CONFIG) @get_app_model - def delete(self, app_model: App): + @model_validate(TraceProviderQuery) + def delete(self, req_data: TraceProviderQuery, app_model: App): """Delete an existing trace app configuration""" - args = TraceProviderQuery.model_validate(request.args.to_dict(flat=True)) - try: result = OpsService.delete_tracing_app_config( - app_id=app_model.id, tracing_provider=args.tracing_provider, session=db.session() + app_id=app_model.id, tracing_provider=req_data.tracing_provider, session=db.session() ) if not result: raise TracingConfigNotExist() diff --git a/api/controllers/console/app/site.py b/api/controllers/console/app/site.py index 669c59e6d53..01835769c19 100644 --- a/api/controllers/console/app/site.py +++ b/api/controllers/console/app/site.py @@ -17,6 +17,7 @@ from controllers.console.wraps import ( account_initialization_required, edit_permission_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_user, @@ -98,8 +99,8 @@ class AppSite(Resource): @with_current_user @with_session @get_app_model - def post(self, session: Session, current_user: Account, app_model: App): - args = AppSiteUpdatePayload.model_validate(console_ns.payload or {}) + @model_validate(AppSiteUpdatePayload) + def post(self, req_data: AppSiteUpdatePayload, session: Session, current_user: Account, app_model: App): site = session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1)) if not site: raise NotFound @@ -123,7 +124,7 @@ class AppSite(Resource): "show_workflow_steps", "use_icon_as_answer_icon", ]: - value = getattr(args, attr_name) + value = getattr(req_data, attr_name) if value is not None: setattr(site, attr_name, value) diff --git a/api/controllers/console/app/statistic.py b/api/controllers/console/app/statistic.py index c7c3c9e64b6..48635166e3a 100644 --- a/api/controllers/console/app/statistic.py +++ b/api/controllers/console/app/statistic.py @@ -1,7 +1,7 @@ from decimal import Decimal import sqlalchemy as sa -from flask import abort, request +from flask import abort from flask_restx import Resource from pydantic import BaseModel, Field, field_validator @@ -12,6 +12,7 @@ from controllers.console.wraps import ( RBACPermission, RBACResourceScope, account_initialization_required, + model_validate, rbac_permission_required, setup_required, with_current_user, @@ -157,8 +158,8 @@ class DailyMessageStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = StatisticTimeRangeQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(StatisticTimeRangeQuery) + def get(self, req_data: StatisticTimeRangeQuery, account: Account, app_model: App): converted_created_at = convert_datetime_to_date("created_at") sql_query = f"""SELECT @@ -177,7 +178,7 @@ WHERE } try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -217,8 +218,8 @@ class DailyConversationStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = StatisticTimeRangeQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(StatisticTimeRangeQuery) + def get(self, req_data: StatisticTimeRangeQuery, account: Account, app_model: App): converted_created_at = convert_datetime_to_date("created_at") sql_query = f"""SELECT @@ -237,7 +238,7 @@ WHERE } try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -276,8 +277,8 @@ class DailyTerminalsStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = StatisticTimeRangeQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(StatisticTimeRangeQuery) + def get(self, req_data: StatisticTimeRangeQuery, account: Account, app_model: App): converted_created_at = convert_datetime_to_date("created_at") sql_query = f"""SELECT @@ -296,7 +297,7 @@ WHERE } try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -336,8 +337,8 @@ class DailyTokenCostStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = StatisticTimeRangeQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(StatisticTimeRangeQuery) + def get(self, req_data: StatisticTimeRangeQuery, account: Account, app_model: App): converted_created_at = convert_datetime_to_date("created_at") sql_query = f"""SELECT @@ -357,7 +358,7 @@ WHERE } try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -399,8 +400,8 @@ class AverageSessionInteractionStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model(mode=[AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT]) - def get(self, account: Account, app_model: App): - args = StatisticTimeRangeQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(StatisticTimeRangeQuery) + def get(self, req_data: StatisticTimeRangeQuery, account: Account, app_model: App): converted_created_at = convert_datetime_to_date("c.created_at") sql_query = f"""SELECT @@ -427,7 +428,7 @@ FROM } try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -478,8 +479,8 @@ class UserSatisfactionRateStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = StatisticTimeRangeQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(StatisticTimeRangeQuery) + def get(self, req_data: StatisticTimeRangeQuery, account: Account, app_model: App): converted_created_at = convert_datetime_to_date("m.created_at") sql_query = f"""SELECT @@ -502,7 +503,7 @@ WHERE } try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -547,8 +548,8 @@ class AverageResponseTimeStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model(mode=AppMode.COMPLETION) - def get(self, account: Account, app_model: App): - args = StatisticTimeRangeQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(StatisticTimeRangeQuery) + def get(self, req_data: StatisticTimeRangeQuery, account: Account, app_model: App): converted_created_at = convert_datetime_to_date("created_at") sql_query = f"""SELECT @@ -567,7 +568,7 @@ WHERE } try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -607,8 +608,8 @@ class TokensPerSecondStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = StatisticTimeRangeQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(StatisticTimeRangeQuery) + def get(self, req_data: StatisticTimeRangeQuery, account: Account, app_model: App): converted_created_at = convert_datetime_to_date("created_at") sql_query = f"""SELECT @@ -630,7 +631,7 @@ WHERE } try: - start_datetime_utc, end_datetime_utc = parse_time_range(args.start, args.end, account.timezone) + start_datetime_utc, end_datetime_utc = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) diff --git a/api/controllers/console/app/workflow.py b/api/controllers/console/app/workflow.py index 6025d02fe39..d506a443a3a 100644 --- a/api/controllers/console/app/workflow.py +++ b/api/controllers/console/app/workflow.py @@ -2,11 +2,20 @@ import json import logging from collections.abc import Sequence from datetime import datetime -from typing import Any, NotRequired, TypedDict +from typing import Any, NotRequired, Self, TypedDict from flask import abort, request from flask_restx import Resource, fields -from pydantic import AliasChoices, BaseModel, ConfigDict, Field, RootModel, ValidationError, field_validator +from pydantic import ( + AliasChoices, + BaseModel, + ConfigDict, + Field, + RootModel, + ValidationError, + field_validator, + model_validator, +) from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import BadRequest, Forbidden, InternalServerError, NotFound @@ -54,6 +63,10 @@ from core.trigger.debug.event_selectors import ( create_event_poller, select_trigger_debug_events, ) +from core.workflow.llm_environment_variable import ( + LLM_ENVIRONMENT_VARIABLE_VALUE_TYPE, + LLMEnvironmentVariable, +) from extensions.ext_database import db from extensions.ext_redis import redis_client from factories import file_factory, variable_factory @@ -76,11 +89,13 @@ from models import Account, App from models.model import AppMode from models.workflow import Workflow from repositories.workflow_collaboration_repository import WORKFLOW_ONLINE_USERS_PREFIX +from services.agent.retirement_service import WorkflowAgentRetirementService from services.app_generate_service import AppGenerateService from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError from services.errors.llm import InvokeRateLimitError from services.workflow_ref_service import WorkflowRefService from services.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError, WorkflowService +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection logger = logging.getLogger(__name__) @@ -101,14 +116,35 @@ class EnvironmentVariableResponseDict(TypedDict): description: NotRequired[str | None] +class SyncEnvironmentVariablePatchPayload(BaseModel): + environment_variables: list[dict[str, Any]] = Field(default_factory=list) + deleted_environment_variable_ids: list[str] = Field(default_factory=list) + + @model_validator(mode="after") + def validate_patch(self) -> Self: + """Require stable, disjoint IDs so the service can merge the patch deterministically.""" + upsert_ids = [variable.get("id") for variable in self.environment_variables] + if any(not isinstance(variable_id, str) or not variable_id for variable_id in upsert_ids): + raise ValueError("patched environment variables require an id") + if len(set(upsert_ids)) != len(upsert_ids): + raise ValueError("patched environment variable ids must be unique") + if any(not variable_id for variable_id in self.deleted_environment_variable_ids): + raise ValueError("deleted environment variable ids must not be empty") + if len(set(self.deleted_environment_variable_ids)) != len(self.deleted_environment_variable_ids): + raise ValueError("deleted environment variable ids must be unique") + if set(upsert_ids).intersection(self.deleted_environment_variable_ids): + raise ValueError("an environment variable cannot be upserted and deleted in the same patch") + return self + + class SyncDraftWorkflowPayload(BaseModel): + model_config = ConfigDict(extra="forbid") + graph: dict[str, Any] features: dict[str, Any] hash: str | None = None is_collaborative: bool = Field(default=False, alias="_is_collaborative") - environment_variables: list[dict[str, Any]] = Field( - default_factory=list, - ) + environment_variable_patch: SyncEnvironmentVariablePatchPayload | None = None conversation_variables: list[dict[str, Any]] = Field( default_factory=list, ) @@ -279,6 +315,9 @@ class WorkflowResponse(ResponseModel): ) hash: str = Field(validation_alias=AliasChoices("unique_hash", "hash")) version: str + # NULL for drafts and for versions published before numbering was introduced; those + # render as "Untitled Version" instead of `#N`. Never 0, so clients must test for null. + version_number: int | None = None marked_name: str marked_comment: str created_by: SimpleAccount | None = Field( @@ -317,7 +356,7 @@ class _WorkflowResponseSource: self._session = session def __getattr__(self, name: str) -> object: - return getattr(self._workflow, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self._workflow, name) # guard-ignore: no-new-getattr -- delegates model fields @property def created_by_account(self) -> Account | None: @@ -481,6 +520,15 @@ def _parse_file(workflow: Workflow, files: list[dict] | None = None) -> Sequence def _serialize_environment_variable(value: Any) -> EnvironmentVariableResponseDict | Any: match value: + case LLMEnvironmentVariable(): + return { + "id": value.id, + "name": value.name, + "value": value.value, + "value_type": LLM_ENVIRONMENT_VARIABLE_VALUE_TYPE, + "description": value.description, + } + case SecretVariable(): return { "id": value.id, @@ -505,6 +553,8 @@ def _serialize_environment_variable(value: Any) -> EnvironmentVariableResponseDi raise TypeError( f"unexpected type for value_type field, value={value_type_str}, type={type(value_type_str)}" ) + if value_type_str == LLM_ENVIRONMENT_VARIABLE_VALUE_TYPE: + return value value_type = SegmentType(value_type_str).exposed_type() if value_type not in ENVIRONMENT_VARIABLE_SUPPORTED_TYPES: raise ValueError(f"Unsupported environment variable value type: {value_type}") @@ -588,30 +638,38 @@ class DraftWorkflowApi(Resource): return {"message": "Invalid JSON data"}, 400 else: abort(415) - args = args_model.model_dump() workflow_service = WorkflowService() try: - environment_variables_list = Workflow.normalize_environment_variable_mappings( - args.get("environment_variables") or [], - ) - environment_variables = [ - variable_factory.build_environment_variable_from_mapping(obj) for obj in environment_variables_list - ] - conversation_variables_list = args.get("conversation_variables") or [] + environment_variable_patch = args_model.environment_variable_patch + environment_variable_upserts: list[VariableBase] | None = None + deleted_environment_variable_ids: list[str] = [] + if environment_variable_patch is not None: + environment_variable_upsert_mappings = Workflow.normalize_environment_variable_mappings( + environment_variable_patch.environment_variables, + ) + environment_variable_upserts = [ + variable_factory.build_environment_variable_from_mapping(obj) + for obj in environment_variable_upsert_mappings + ] + deleted_environment_variable_ids = environment_variable_patch.deleted_environment_variable_ids conversation_variables = [ - variable_factory.build_conversation_variable_from_mapping(obj) for obj in conversation_variables_list + variable_factory.build_conversation_variable_from_mapping(obj) + for obj in args_model.conversation_variables ] workflow = workflow_service.sync_draft_workflow( app_model=app_model, - graph=args["graph"], - features=args["features"], - unique_hash=args.get("hash"), + graph=args_model.graph, + features=args_model.features, + unique_hash=args_model.hash, account=current_user, - environment_variables=environment_variables, + environment_variables=[], conversation_variables=conversation_variables, session=db.session(), - graph_only=args["is_collaborative"], + environment_variable_upserts=environment_variable_upserts, + deleted_environment_variable_ids=deleted_environment_variable_ids, + preserve_environment_variables=True, + graph_only=args_model.is_collaborative, ) except WorkflowHashNotEqualError: raise DraftWorkflowNotSync() @@ -1245,7 +1303,7 @@ class PublishedWorkflowApi(Resource): workflow_service = WorkflowService() with sessionmaker(db.engine).begin() as session: - workflow = workflow_service.publish_workflow( + workflow, retirement_candidates = workflow_service.publish_workflow( session=session, app_model=app_model, account=current_user, @@ -1262,6 +1320,16 @@ class PublishedWorkflowApi(Resource): workflow_created_at = TimestampField().format(workflow.created_at) + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=app_model.tenant_id, + agent_ids=retirement_candidates, + account_id=current_user.id, + ) + enqueue_agent_resource_collection( + tenant_id=app_model.tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) return { "result": "success", "created_at": workflow_created_at, diff --git a/api/controllers/console/app/workflow_app_log.py b/api/controllers/console/app/workflow_app_log.py index b3426e8f6ea..12a36939b8b 100644 --- a/api/controllers/console/app/workflow_app_log.py +++ b/api/controllers/console/app/workflow_app_log.py @@ -2,7 +2,6 @@ from datetime import datetime from typing import Any from dateutil.parser import isoparse -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy.orm import sessionmaker @@ -14,6 +13,7 @@ from controllers.console.wraps import ( RBACPermission, RBACResourceScope, account_initialization_required, + model_validate, rbac_permission_required, setup_required, ) @@ -183,11 +183,11 @@ class WorkflowAppLogApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_LOG_AND_ANNOTATION) @get_app_model(mode=[AppMode.WORKFLOW]) - def get(self, app_model: App): + @model_validate(WorkflowAppLogQuery) + def get(self, req_data: WorkflowAppLogQuery, app_model: App): """ Get workflow app logs """ - args = WorkflowAppLogQuery.model_validate(request.args.to_dict(flat=True)) # get paginate workflow app logs workflow_app_service = WorkflowAppService() @@ -195,15 +195,15 @@ class WorkflowAppLogApi(Resource): workflow_app_log_pagination = workflow_app_service.get_paginate_workflow_app_logs( session=session, app_model=app_model, - keyword=args.keyword, - status=args.status, - created_at_before=args.created_at__before, - created_at_after=args.created_at__after, - page=args.page, - limit=args.limit, - detail=args.detail, - created_by_end_user_session_id=args.created_by_end_user_session_id, - created_by_account=args.created_by_account, + keyword=req_data.keyword, + status=req_data.status, + created_at_before=req_data.created_at__before, + created_at_after=req_data.created_at__after, + page=req_data.page, + limit=req_data.limit, + detail=req_data.detail, + created_by_end_user_session_id=req_data.created_by_end_user_session_id, + created_by_account=req_data.created_by_account, ) return WorkflowAppLogPaginationResponse.model_validate( @@ -227,19 +227,19 @@ class WorkflowArchivedLogApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_LOG_AND_ANNOTATION) @get_app_model(mode=[AppMode.WORKFLOW]) - def get(self, app_model: App): + @model_validate(WorkflowAppLogQuery) + def get(self, req_data: WorkflowAppLogQuery, app_model: App): """ Get workflow archived logs """ - args = WorkflowAppLogQuery.model_validate(request.args.to_dict(flat=True)) workflow_app_service = WorkflowAppService() with sessionmaker(db.engine, expire_on_commit=False).begin() as session: workflow_app_log_pagination = workflow_app_service.get_paginate_workflow_archive_logs( session=session, app_model=app_model, - page=args.page, - limit=args.limit, + page=req_data.page, + limit=req_data.limit, ) return WorkflowArchivedLogPaginationResponse.model_validate( diff --git a/api/controllers/console/app/workflow_comment.py b/api/controllers/console/app/workflow_comment.py index 082b49b20f2..cc156294e86 100644 --- a/api/controllers/console/app/workflow_comment.py +++ b/api/controllers/console/app/workflow_comment.py @@ -10,6 +10,7 @@ from controllers.console.app.wraps import get_app_model from controllers.console.wraps import ( account_initialization_required, edit_permission_required, + model_validate, setup_required, with_current_tenant_id, with_current_user, @@ -239,18 +240,24 @@ class WorkflowCommentListApi(Resource): @with_current_user @with_current_tenant_id @get_app_model() - def post(self, current_tenant_id: str, current_user: Account, app_model: App): + @model_validate(WorkflowCommentCreatePayload) + def post( + self, + req_data: WorkflowCommentCreatePayload, + current_tenant_id: str, + current_user: Account, + app_model: App, + ): """Create a new workflow comment.""" - payload = WorkflowCommentCreatePayload.model_validate(console_ns.payload or {}) result = WorkflowCommentService.create_comment( tenant_id=current_tenant_id, app_id=app_model.id, created_by=current_user.id, - content=payload.content, - position_x=payload.position_x, - position_y=payload.position_y, - mentioned_user_ids=payload.mentioned_user_ids, + content=req_data.content, + position_x=req_data.position_x, + position_y=req_data.position_y, + mentioned_user_ids=req_data.mentioned_user_ids, ) return dump_response(WorkflowCommentCreate, result), 201 @@ -289,19 +296,26 @@ class WorkflowCommentDetailApi(Resource): @with_current_user @with_current_tenant_id @get_app_model() - def put(self, current_tenant_id: str, current_user: Account, app_model: App, comment_id: str): + @model_validate(WorkflowCommentUpdatePayload) + def put( + self, + req_data: WorkflowCommentUpdatePayload, + current_tenant_id: str, + current_user: Account, + app_model: App, + comment_id: str, + ): """Update a workflow comment.""" - payload = WorkflowCommentUpdatePayload.model_validate(console_ns.payload or {}) result = WorkflowCommentService.update_comment( tenant_id=current_tenant_id, app_id=app_model.id, comment_id=comment_id, user_id=current_user.id, - content=payload.content, - position_x=payload.position_x, - position_y=payload.position_y, - mentioned_user_ids=payload.mentioned_user_ids, + content=req_data.content, + position_x=req_data.position_x, + position_y=req_data.position_y, + mentioned_user_ids=req_data.mentioned_user_ids, ) return dump_response(WorkflowCommentUpdate, result) @@ -372,20 +386,26 @@ class WorkflowCommentReplyApi(Resource): @with_current_user @with_current_tenant_id @get_app_model() - def post(self, current_tenant_id: str, current_user: Account, app_model: App, comment_id: str): + @model_validate(WorkflowCommentReplyPayload) + def post( + self, + req_data: WorkflowCommentReplyPayload, + current_tenant_id: str, + current_user: Account, + app_model: App, + comment_id: str, + ): """Add a reply to a workflow comment.""" # Validate comment access first WorkflowCommentService.validate_comment_access( comment_id=comment_id, tenant_id=current_tenant_id, app_id=app_model.id ) - payload = WorkflowCommentReplyPayload.model_validate(console_ns.payload or {}) - result = WorkflowCommentService.create_reply( comment_id=comment_id, - content=payload.content, + content=req_data.content, created_by=current_user.id, - mentioned_user_ids=payload.mentioned_user_ids, + mentioned_user_ids=req_data.mentioned_user_ids, ) return dump_response(WorkflowCommentReplyCreate, result), 201 @@ -407,23 +427,30 @@ class WorkflowCommentReplyDetailApi(Resource): @with_current_user @with_current_tenant_id @get_app_model() - def put(self, current_tenant_id: str, current_user: Account, app_model: App, comment_id: str, reply_id: str): + @model_validate(WorkflowCommentReplyPayload) + def put( + self, + req_data: WorkflowCommentReplyPayload, + current_tenant_id: str, + current_user: Account, + app_model: App, + comment_id: str, + reply_id: str, + ): """Update a comment reply.""" # Validate comment access first WorkflowCommentService.validate_comment_access( comment_id=comment_id, tenant_id=current_tenant_id, app_id=app_model.id ) - payload = WorkflowCommentReplyPayload.model_validate(console_ns.payload or {}) - reply = WorkflowCommentService.update_reply( tenant_id=current_tenant_id, app_id=app_model.id, comment_id=comment_id, reply_id=reply_id, user_id=current_user.id, - content=payload.content, - mentioned_user_ids=payload.mentioned_user_ids, + content=req_data.content, + mentioned_user_ids=req_data.mentioned_user_ids, ) return dump_response(WorkflowCommentReplyUpdate, reply) diff --git a/api/controllers/console/app/workflow_draft_variable.py b/api/controllers/console/app/workflow_draft_variable.py index 1ffc01e3cef..6d52419f88c 100644 --- a/api/controllers/console/app/workflow_draft_variable.py +++ b/api/controllers/console/app/workflow_draft_variable.py @@ -1,12 +1,12 @@ import logging from collections.abc import Callable from functools import wraps -from typing import Any, Concatenate, TypedDict, override +from typing import Any, Concatenate, Self, TypedDict, override from uuid import UUID -from flask import Response, request +from flask import Response from flask_restx import Resource, fields, marshal, marshal_with -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, model_validator from sqlalchemy.orm import sessionmaker from controllers.common.errors import InvalidArgumentError, NotFoundError @@ -22,11 +22,13 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_user, ) from core.app.file_access import DatabaseFileAccessController +from core.workflow.llm_environment_variable import environment_variable_value_type from core.workflow.variable_prefixes import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID from extensions.ext_database import db from factories import variable_factory @@ -109,6 +111,35 @@ class EnvironmentVariableUpdatePayload(BaseModel): ..., description="Environment variables for the draft workflow", ) + patch: bool = Field( + default=False, + description="Treat environment_variables as per-ID upserts instead of replacing the full collection", + ) + deleted_environment_variable_ids: list[str] = Field( + default_factory=list, + description="Environment variable IDs to delete when patch is true", + ) + + @model_validator(mode="after") + def validate_patch(self) -> Self: + """Validate the per-variable patch contract without changing legacy replacement requests.""" + if not self.patch: + if self.deleted_environment_variable_ids: + raise ValueError("deleted_environment_variable_ids requires patch=true") + return self + + upsert_ids = [variable.id for variable in self.environment_variables] + if any(not variable_id for variable_id in upsert_ids): + raise ValueError("patched environment variables require an id") + if len(set(upsert_ids)) != len(upsert_ids): + raise ValueError("patched environment variable ids must be unique") + if any(not variable_id for variable_id in self.deleted_environment_variable_ids): + raise ValueError("deleted environment variable ids must not be empty") + if len(set(self.deleted_environment_variable_ids)) != len(self.deleted_environment_variable_ids): + raise ValueError("deleted environment variable ids must be unique") + if set(upsert_ids).intersection(self.deleted_environment_variable_ids): + raise ValueError("an environment variable cannot be upserted and deleted in the same patch") + return self class EnvironmentVariableItemResponse(ResponseModel): @@ -329,11 +360,11 @@ class WorkflowVariableCollectionApi(Resource): @_api_prerequisite @marshal_with(workflow_draft_variable_list_without_value_model) @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) - def get(self, current_user: Account, app_model: App): + @model_validate(WorkflowDraftVariableListQuery) + def get(self, req_data: WorkflowDraftVariableListQuery, current_user: Account, app_model: App): """ Get draft workflow """ - args = WorkflowDraftVariableListQuery.model_validate(request.args.to_dict(flat=True)) # fetch draft workflow by app_model workflow_service = WorkflowService() @@ -348,8 +379,8 @@ class WorkflowVariableCollectionApi(Resource): ) workflow_vars = draft_var_srv.list_variables_without_values( app_id=app_model.id, - page=args.page, - limit=args.limit, + page=req_data.page, + limit=req_data.limit, user_id=current_user.id, ) @@ -449,7 +480,14 @@ class VariableApi(Resource): @console_ns.response(404, "Variable not found") @_api_prerequisite @marshal_with(workflow_draft_variable_model) - def patch(self, current_user: Account, app_model: App, variable_id: UUID): + @model_validate(WorkflowDraftVariableUpdatePayload) + def patch( + self, + req_data: WorkflowDraftVariableUpdatePayload, + current_user: Account, + app_model: App, + variable_id: UUID, + ): # Request payload for file types: # # Local File: @@ -474,7 +512,6 @@ class VariableApi(Resource): draft_var_srv = WorkflowDraftVariableService( session=db.session(), ) - args_model = WorkflowDraftVariableUpdatePayload.model_validate(console_ns.payload or {}) variable_id_str = str(variable_id) variable = ensure_variable_access( @@ -484,8 +521,8 @@ class VariableApi(Resource): current_user_id=current_user.id, ) - new_name = args_model.name - raw_value = args_model.value + new_name = req_data.name + raw_value = req_data.value if new_name is None and raw_value is None: return variable @@ -630,13 +667,13 @@ class ConversationVariableCollectionApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_current_user @get_app_model(mode=AppMode.ADVANCED_CHAT) - def post(self, current_user: Account, app_model: App): - payload = ConversationVariableUpdatePayload.model_validate(console_ns.payload or {}) + @model_validate(ConversationVariableUpdatePayload) + def post(self, req_data: ConversationVariableUpdatePayload, current_user: Account, app_model: App): workflow_service = WorkflowService() conversation_variables_list = [ - variable.model_dump(mode="json", exclude_unset=True) for variable in payload.conversation_variables + variable.model_dump(mode="json", exclude_unset=True) for variable in req_data.conversation_variables ] conversation_variables = [ variable_factory.build_conversation_variable_from_mapping(obj) for obj in conversation_variables_list @@ -698,7 +735,7 @@ class EnvironmentVariableCollectionApi(Resource): "name": v.name, "description": v.description, "selector": v.selector, - "value_type": str(v.value_type.exposed_type()), + "value_type": environment_variable_value_type(v), "value": v.value, # Do not track edited for env vars. "edited": False, @@ -725,23 +762,32 @@ class EnvironmentVariableCollectionApi(Resource): @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_EDIT) @with_current_user @get_app_model(mode=[AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) - def post(self, current_user: Account, app_model: App): - payload = EnvironmentVariableUpdatePayload.model_validate(console_ns.payload or {}) + @model_validate(EnvironmentVariableUpdatePayload) + def post(self, req_data: EnvironmentVariableUpdatePayload, current_user: Account, app_model: App): workflow_service = WorkflowService() environment_variables_list = [ - variable.model_dump(mode="json", exclude_unset=True) for variable in payload.environment_variables + variable.model_dump(mode="json", exclude_unset=True) for variable in req_data.environment_variables ] environment_variables = [ variable_factory.build_environment_variable_from_mapping(obj) for obj in environment_variables_list ] - workflow_service.update_draft_workflow_environment_variables( - app_model=app_model, - account=current_user, - environment_variables=environment_variables, - session=db.session(), - ) + if req_data.patch: + workflow_service.patch_draft_workflow_environment_variables( + app_model=app_model, + account=current_user, + environment_variables=environment_variables, + deleted_environment_variable_ids=req_data.deleted_environment_variable_ids, + session=db.session(), + ) + else: + workflow_service.update_draft_workflow_environment_variables( + app_model=app_model, + account=current_user, + environment_variables=environment_variables, + session=db.session(), + ) return {"result": "success"} diff --git a/api/controllers/console/app/workflow_run.py b/api/controllers/console/app/workflow_run.py index a71aa444e27..bc842dbf6b4 100644 --- a/api/controllers/console/app/workflow_run.py +++ b/api/controllers/console/app/workflow_run.py @@ -2,7 +2,6 @@ from datetime import UTC, datetime, timedelta from typing import Literal from uuid import UUID -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import select @@ -17,6 +16,7 @@ from controllers.console.wraps import ( RBACPermission, RBACResourceScope, account_initialization_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -160,21 +160,21 @@ class AdvancedChatAppWorkflowRunListApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_CREATE_AND_MANAGEMENT) @get_app_model(mode=[AppMode.ADVANCED_CHAT]) - def get(self, app_model: App): + @model_validate(WorkflowRunListQuery) + def get(self, req_data: WorkflowRunListQuery, app_model: App): """ Get advanced chat app workflow run list """ - args_model = WorkflowRunListQuery.model_validate(request.args.to_dict(flat=True)) - args: WorkflowRunListArgs = {"limit": args_model.limit} - if args_model.last_id is not None: - args["last_id"] = args_model.last_id - if args_model.status is not None: - args["status"] = args_model.status + args: WorkflowRunListArgs = {"limit": req_data.limit} + if req_data.last_id is not None: + args["last_id"] = req_data.last_id + if req_data.status is not None: + args["status"] = req_data.status # Default to DEBUGGING if not specified triggered_from = ( - WorkflowRunTriggeredFrom(args_model.triggered_from) - if args_model.triggered_from + WorkflowRunTriggeredFrom(req_data.triggered_from) + if req_data.triggered_from else WorkflowRunTriggeredFrom.DEBUGGING ) @@ -258,17 +258,17 @@ class AdvancedChatAppWorkflowRunCountApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_CREATE_AND_MANAGEMENT) @get_app_model(mode=[AppMode.ADVANCED_CHAT]) - def get(self, app_model: App): + @model_validate(WorkflowRunCountQuery) + def get(self, req_data: WorkflowRunCountQuery, app_model: App): """ Get advanced chat workflow runs count statistics """ - args_model = WorkflowRunCountQuery.model_validate(request.args.to_dict(flat=True)) - args = args_model.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) # Default to DEBUGGING if not specified triggered_from = ( - WorkflowRunTriggeredFrom(args_model.triggered_from) - if args_model.triggered_from + WorkflowRunTriggeredFrom(req_data.triggered_from) + if req_data.triggered_from else WorkflowRunTriggeredFrom.DEBUGGING ) @@ -299,21 +299,21 @@ class WorkflowRunListApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_CREATE_AND_MANAGEMENT) @get_app_model(mode=[AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) - def get(self, app_model: App): + @model_validate(WorkflowRunListQuery) + def get(self, req_data: WorkflowRunListQuery, app_model: App): """ Get workflow run list """ - args_model = WorkflowRunListQuery.model_validate(request.args.to_dict(flat=True)) - args: WorkflowRunListArgs = {"limit": args_model.limit} - if args_model.last_id is not None: - args["last_id"] = args_model.last_id - if args_model.status is not None: - args["status"] = args_model.status + args: WorkflowRunListArgs = {"limit": req_data.limit} + if req_data.last_id is not None: + args["last_id"] = req_data.last_id + if req_data.status is not None: + args["status"] = req_data.status # Default to DEBUGGING for workflow if not specified (backward compatibility) triggered_from = ( - WorkflowRunTriggeredFrom(args_model.triggered_from) - if args_model.triggered_from + WorkflowRunTriggeredFrom(req_data.triggered_from) + if req_data.triggered_from else WorkflowRunTriggeredFrom.DEBUGGING ) @@ -341,17 +341,17 @@ class WorkflowRunCountApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_CREATE_AND_MANAGEMENT) @get_app_model(mode=[AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) - def get(self, app_model: App): + @model_validate(WorkflowRunCountQuery) + def get(self, req_data: WorkflowRunCountQuery, app_model: App): """ Get workflow runs count statistics """ - args_model = WorkflowRunCountQuery.model_validate(request.args.to_dict(flat=True)) - args = args_model.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) # Default to DEBUGGING for workflow if not specified (backward compatibility) triggered_from = ( - WorkflowRunTriggeredFrom(args_model.triggered_from) - if args_model.triggered_from + WorkflowRunTriggeredFrom(req_data.triggered_from) + if req_data.triggered_from else WorkflowRunTriggeredFrom.DEBUGGING ) diff --git a/api/controllers/console/app/workflow_statistic.py b/api/controllers/console/app/workflow_statistic.py index 0346d510fbc..e72cd057a65 100644 --- a/api/controllers/console/app/workflow_statistic.py +++ b/api/controllers/console/app/workflow_statistic.py @@ -1,4 +1,4 @@ -from flask import abort, jsonify, request +from flask import abort, jsonify from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy.orm import sessionmaker @@ -10,6 +10,7 @@ from controllers.console.wraps import ( RBACPermission, RBACResourceScope, account_initialization_required, + model_validate, rbac_permission_required, setup_required, with_current_user, @@ -104,13 +105,13 @@ class WorkflowDailyRunsStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = WorkflowStatisticQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(WorkflowStatisticQuery) + def get(self, req_data: WorkflowStatisticQuery, account: Account, app_model: App): assert account.timezone is not None try: - start_date, end_date = parse_time_range(args.start, args.end, account.timezone) + start_date, end_date = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -148,13 +149,13 @@ class WorkflowDailyTerminalsStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = WorkflowStatisticQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(WorkflowStatisticQuery) + def get(self, req_data: WorkflowStatisticQuery, account: Account, app_model: App): assert account.timezone is not None try: - start_date, end_date = parse_time_range(args.start, args.end, account.timezone) + start_date, end_date = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -192,13 +193,13 @@ class WorkflowDailyTokenCostStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model - def get(self, account: Account, app_model: App): - args = WorkflowStatisticQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(WorkflowStatisticQuery) + def get(self, req_data: WorkflowStatisticQuery, account: Account, app_model: App): assert account.timezone is not None try: - start_date, end_date = parse_time_range(args.start, args.end, account.timezone) + start_date, end_date = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) @@ -236,13 +237,13 @@ class WorkflowAverageAppInteractionStatistic(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_MONITOR) @get_app_model(mode=[AppMode.WORKFLOW]) - def get(self, account: Account, app_model: App): - args = WorkflowStatisticQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(WorkflowStatisticQuery) + def get(self, req_data: WorkflowStatisticQuery, account: Account, app_model: App): assert account.timezone is not None try: - start_date, end_date = parse_time_range(args.start, args.end, account.timezone) + start_date, end_date = parse_time_range(req_data.start, req_data.end, account.timezone) except ValueError as e: abort(400, description=str(e)) diff --git a/api/controllers/console/app/workflow_trigger.py b/api/controllers/console/app/workflow_trigger.py index 2f45d637256..4f544db7c46 100644 --- a/api/controllers/console/app/workflow_trigger.py +++ b/api/controllers/console/app/workflow_trigger.py @@ -1,7 +1,6 @@ import logging from datetime import datetime -from flask import request from flask_restx import Resource from pydantic import BaseModel, field_validator from sqlalchemy import select @@ -25,6 +24,7 @@ from ..wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -102,11 +102,11 @@ class WebhookTriggerApi(Resource): @console_ns.response(200, "Success", console_ns.models[WebhookTriggerResponse.__name__]) @rbac_permission_required(RBACResourceScope.APP, RBACPermission.APP_VIEW_LAYOUT) @get_app_model(mode=AppMode.WORKFLOW) - def get(self, app_model: App): + @model_validate(Parser) + def get(self, req_data: Parser, app_model: App): """Get webhook trigger for a node""" - args = Parser.model_validate(request.args.to_dict(flat=True)) - node_id = args.node_id + node_id = req_data.node_id with sessionmaker(db.engine, expire_on_commit=False).begin() as session: # Get webhook trigger for this app and node @@ -175,11 +175,11 @@ class AppTriggerEnableApi(Resource): @console_ns.response(200, "Success", console_ns.models[WorkflowTriggerResponse.__name__]) @with_current_tenant_id @get_app_model(mode=AppMode.WORKFLOW) - def post(self, current_tenant_id: str, app_model: App): + @model_validate(ParserEnable) + def post(self, req_data: ParserEnable, current_tenant_id: str, app_model: App): """Update app trigger (enable/disable)""" - args = ParserEnable.model_validate(console_ns.payload) - trigger_id = args.trigger_id + trigger_id = req_data.trigger_id with sessionmaker(db.engine, expire_on_commit=False).begin() as session: # Find the trigger using select trigger = session.execute( @@ -194,7 +194,7 @@ class AppTriggerEnableApi(Resource): raise NotFound("Trigger not found") # Update status based on enable_trigger boolean - trigger.status = AppTriggerStatus.ENABLED if args.enable_trigger else AppTriggerStatus.DISABLED + trigger.status = AppTriggerStatus.ENABLED if req_data.enable_trigger else AppTriggerStatus.DISABLED # Add computed icon field url_prefix = dify_config.CONSOLE_API_URL + "/console/api/workspaces/current/tool-provider/builtin/" diff --git a/api/controllers/console/auth/activate.py b/api/controllers/console/auth/activate.py index 3e9160f2bb0..07c9d30d997 100644 --- a/api/controllers/console/auth/activate.py +++ b/api/controllers/console/auth/activate.py @@ -9,6 +9,8 @@ from controllers.common.schema import query_params_from_model, register_schema_m from controllers.console import console_ns from controllers.console.auth.error import InvitationAccountMismatchError from controllers.console.error import AccountInFreezeError, AlreadyActivateError +from controllers.console.wraps import model_validate +from enums import DeploymentEdition from extensions.ext_database import db from libs.datetime_utils import naive_utc_now from libs.helper import EmailStr, timezone @@ -86,14 +88,14 @@ class ActivateCheckApi(Resource): "Success", console_ns.models[ActivationCheckResponse.__name__], ) - def get(self): - args = ActivateCheckQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(ActivateCheckQuery) + def get(self, req_data: ActivateCheckQuery): - workspaceId = args.workspace_id - token = args.token + workspaceId = req_data.workspace_id + token = req_data.token invitation = RegisterService.get_invitation_with_case_fallback( - workspaceId, args.email, token, session=db.session() + workspaceId, req_data.email, token, session=db.session() ) if invitation: data = invitation.get("data", {}) @@ -138,18 +140,18 @@ class ActivateApi(Resource): console_ns.models[ActivationResponse.__name__], ) @console_ns.response(400, "Already activated or invalid token") - def post(self): + @model_validate(ActivatePayload) + def post(self, req_data: ActivatePayload): """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 + normalized_request_email = req_data.email.lower() if req_data.email else None invitation = RegisterService.get_invitation_with_case_fallback( - args.workspace_id, args.email, args.token, session=db.session() + req_data.workspace_id, req_data.email, req_data.token, session=db.session() ) if invitation is None: raise AlreadyActivateError() @@ -160,7 +162,9 @@ class ActivateApi(Resource): if current_account.id != account.id: raise InvitationAccountMismatchError() - if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(account.email): + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and BillingService.is_email_in_freeze( + account.email + ): raise AccountInFreezeError() tenant = invitation["tenant"] @@ -185,11 +189,11 @@ class ActivateApi(Resource): setup_fields: tuple[str, str, str] | None = None if requires_setup: - if not args.name or not args.interface_language or not args.timezone: + if not req_data.name or not req_data.interface_language or not req_data.timezone: raise AlreadyActivateError() - setup_fields = (args.name, args.interface_language, args.timezone) + setup_fields = (req_data.name, req_data.interface_language, req_data.timezone) - RegisterService.revoke_token(args.workspace_id, normalized_request_email, args.token) + RegisterService.revoke_token(req_data.workspace_id, normalized_request_email, req_data.token) if membership_id is None: TenantService.create_tenant_member(tenant, account, db.session(), role=role) diff --git a/api/controllers/console/auth/data_source_bearer_auth.py b/api/controllers/console/auth/data_source_bearer_auth.py index fac725e8534..3f42b7a3ad7 100644 --- a/api/controllers/console/auth/data_source_bearer_auth.py +++ b/api/controllers/console/auth/data_source_bearer_auth.py @@ -17,6 +17,7 @@ from ..wraps import ( RBACResourceScope, account_initialization_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -87,10 +88,10 @@ class ApiKeyAuthDataSourceBinding(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_CREATE, resource_required=False) @console_ns.expect(console_ns.models[ApiKeyAuthBindingPayload.__name__]) @with_current_tenant_id - def post(self, current_tenant_id: str): + @model_validate(ApiKeyAuthBindingPayload) + def post(self, req_data: ApiKeyAuthBindingPayload, current_tenant_id: str): # The role of the current user in the table must be admin or owner - payload = ApiKeyAuthBindingPayload.model_validate(console_ns.payload) - data = payload.model_dump() + data = req_data.model_dump() ApiKeyAuthService.validate_api_key_auth_args(data) try: ApiKeyAuthService.create_provider_auth(current_tenant_id, data, session=db.session()) diff --git a/api/controllers/console/auth/email_register.py b/api/controllers/console/auth/email_register.py index a9b73ca4679..b3377bb7019 100644 --- a/api/controllers/console/auth/email_register.py +++ b/api/controllers/console/auth/email_register.py @@ -15,6 +15,7 @@ from controllers.console.auth.error import ( InvalidTokenError, PasswordMismatchError, ) +from enums import DeploymentEdition from extensions.ext_database import db from fields.base import ResponseModel from libs.helper import EmailStr, extract_remote_ip @@ -26,7 +27,7 @@ from services.billing_service import BillingService from services.errors.account import AccountRegisterError, SeatsLimitExceededError from ..error import AccountInFreezeError, EmailSendIpLimitError, SeatsLimitExceeded -from ..wraps import email_password_login_enabled, email_register_enabled, setup_required +from ..wraps import email_password_login_enabled, email_register_enabled, model_validate, setup_required class EmailRegisterSendPayload(BaseModel): @@ -87,21 +88,23 @@ class EmailRegisterSendEmailApi(Resource): @email_register_enabled @console_ns.expect(console_ns.models[EmailRegisterSendPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[SimpleResultDataResponse.__name__]) - def post(self): - args = EmailRegisterSendPayload.model_validate(console_ns.payload) - normalized_email = args.email.lower() + @model_validate(EmailRegisterSendPayload) + def post(self, req_data: EmailRegisterSendPayload): + normalized_email = req_data.email.lower() ip_address = extract_remote_ip(request) if AccountService.is_email_send_ip_limit(ip_address): raise EmailSendIpLimitError() language = "en-US" - if args.language is not None and args.language in languages: - language = args.language + if req_data.language is not None and req_data.language in languages: + language = req_data.language - if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(normalized_email): + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and BillingService.is_email_in_freeze( + normalized_email + ): raise AccountInFreezeError() - account = AccountService.get_account_by_email_with_case_fallback(args.email, session=db.session()) + account = AccountService.get_account_by_email_with_case_fallback(req_data.email, session=db.session()) token = AccountService.send_email_register_email(email=normalized_email, account=account, language=language) return {"result": "success", "data": token} @@ -113,16 +116,16 @@ class EmailRegisterCheckApi(Resource): @email_register_enabled @console_ns.expect(console_ns.models[EmailRegisterValidityPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[VerificationTokenResponse.__name__]) - def post(self): - args = EmailRegisterValidityPayload.model_validate(console_ns.payload) + @model_validate(EmailRegisterValidityPayload) + def post(self, req_data: EmailRegisterValidityPayload): - user_email = args.email.lower() + user_email = req_data.email.lower() is_email_register_error_rate_limit = AccountService.is_email_register_error_rate_limit(user_email) if is_email_register_error_rate_limit: raise EmailRegisterLimitError() - token_data = AccountService.get_email_register_data(args.token) + token_data = AccountService.get_email_register_data(req_data.token) if token_data is None: raise InvalidTokenError() @@ -132,16 +135,16 @@ class EmailRegisterCheckApi(Resource): if user_email != normalized_token_email: raise InvalidEmailError() - if args.code != token_data.get("code"): + if req_data.code != token_data.get("code"): AccountService.add_email_register_error_rate_limit(user_email) raise EmailCodeError() # Verified, revoke the first token - AccountService.revoke_email_register_token(args.token) + AccountService.revoke_email_register_token(req_data.token) # Refresh token data by generating a new token _, new_token = AccountService.generate_email_register_token( - user_email, code=args.code, additional_data={"phase": "register"} + user_email, code=req_data.code, additional_data={"phase": "register"} ) AccountService.reset_email_register_error_rate_limit(user_email) @@ -155,15 +158,15 @@ class EmailRegisterResetApi(Resource): @email_register_enabled @console_ns.expect(console_ns.models[EmailRegisterResetPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[EmailRegisterResetResponse.__name__]) - def post(self): - args = EmailRegisterResetPayload.model_validate(console_ns.payload) + @model_validate(EmailRegisterResetPayload) + def post(self, req_data: EmailRegisterResetPayload): # Validate passwords match - if args.new_password != args.password_confirm: + if req_data.new_password != req_data.password_confirm: raise PasswordMismatchError() # Validate token and get register data - register_data = AccountService.get_email_register_data(args.token) + register_data = AccountService.get_email_register_data(req_data.token) if not register_data: raise InvalidTokenError() # Must use token in reset phase @@ -171,7 +174,7 @@ class EmailRegisterResetApi(Resource): raise InvalidTokenError() # Revoke token to prevent reuse - AccountService.revoke_email_register_token(args.token) + AccountService.revoke_email_register_token(req_data.token) email = register_data.get("email", "") normalized_email = email.lower() @@ -183,9 +186,9 @@ class EmailRegisterResetApi(Resource): account = self._create_new_account( email=normalized_email, - password=args.password_confirm, - timezone=args.timezone, - language=args.language, + password=req_data.password_confirm, + timezone=req_data.timezone, + language=req_data.language, ) token_pair = AccountService.login(account=account, session=db.session(), ip_address=extract_remote_ip(request)) AccountService.reset_login_error_rate_limit(normalized_email) diff --git a/api/controllers/console/auth/error.py b/api/controllers/console/auth/error.py index 562de31270f..36bc34cc406 100644 --- a/api/controllers/console/auth/error.py +++ b/api/controllers/console/auth/error.py @@ -95,6 +95,18 @@ class EmailPasswordLoginLimitError(BaseHTTPException): code = 429 +class TurnstileVerificationFailedError(BaseHTTPException): + error_code = "turnstile_verification_failed" + description = "Turnstile verification failed. Please try again." + code = 400 + + +class TurnstileServiceUnavailableError(BaseHTTPException): + error_code = "turnstile_service_unavailable" + description = "Turnstile verification is temporarily unavailable. Please try again later." + code = 503 + + class EmailCodeLoginRateLimitExceededError(BaseHTTPException): error_code = "email_code_login_rate_limit_exceeded" description = "Too many login emails have been sent. Please try again in {minutes} minutes." diff --git a/api/controllers/console/auth/forgot_password.py b/api/controllers/console/auth/forgot_password.py index 8a46a2559cf..9a8784d543a 100644 --- a/api/controllers/console/auth/forgot_password.py +++ b/api/controllers/console/auth/forgot_password.py @@ -15,7 +15,7 @@ from controllers.console.auth.error import ( PasswordMismatchError, ) from controllers.console.error import AccountNotFound, EmailSendIpLimitError -from controllers.console.wraps import email_password_login_enabled, setup_required +from controllers.console.wraps import email_password_login_enabled, model_validate, setup_required from extensions.ext_database import db from libs.helper import EmailStr, extract_remote_ip from libs.password import hash_password @@ -68,20 +68,20 @@ class ForgotPasswordSendEmailApi(Resource): @console_ns.response(400, "Invalid email or rate limit exceeded") @setup_required @email_password_login_enabled - def post(self): - args = ForgotPasswordSendPayload.model_validate(console_ns.payload) - normalized_email = args.email.lower() + @model_validate(ForgotPasswordSendPayload) + def post(self, req_data: ForgotPasswordSendPayload): + normalized_email = req_data.email.lower() ip_address = extract_remote_ip(request) if AccountService.is_email_send_ip_limit(ip_address): raise EmailSendIpLimitError() - if args.language is not None and args.language == "zh-Hans": + if req_data.language is not None and req_data.language == "zh-Hans": language = "zh-Hans" else: language = "en-US" - account = AccountService.get_account_by_email_with_case_fallback(args.email, session=db.session()) + account = AccountService.get_account_by_email_with_case_fallback(req_data.email, session=db.session()) token = AccountService.send_reset_password_email( account=account, @@ -106,16 +106,16 @@ class ForgotPasswordCheckApi(Resource): @console_ns.response(400, "Invalid code or token") @setup_required @email_password_login_enabled - def post(self): - args = ForgotPasswordCheckPayload.model_validate(console_ns.payload) + @model_validate(ForgotPasswordCheckPayload) + def post(self, req_data: ForgotPasswordCheckPayload): - user_email = args.email.lower() + user_email = req_data.email.lower() is_forgot_password_error_rate_limit = AccountService.is_forgot_password_error_rate_limit(user_email) if is_forgot_password_error_rate_limit: raise EmailPasswordResetLimitError() - token_data = AccountService.get_reset_password_data(args.token) + token_data = AccountService.get_reset_password_data(req_data.token) if token_data is None: raise InvalidTokenError() @@ -127,16 +127,16 @@ class ForgotPasswordCheckApi(Resource): if user_email != normalized_token_email: raise InvalidEmailError() - if args.code != token_data.get("code"): + if req_data.code != token_data.get("code"): AccountService.add_forgot_password_error_rate_limit(user_email) raise EmailCodeError() # Verified, revoke the first token - AccountService.revoke_reset_password_token(args.token) + AccountService.revoke_reset_password_token(req_data.token) # Refresh token data by generating a new token _, new_token = AccountService.generate_reset_password_token( - token_email, code=args.code, additional_data={"phase": "reset"} + token_email, code=req_data.code, additional_data={"phase": "reset"} ) AccountService.reset_forgot_password_error_rate_limit(user_email) @@ -156,15 +156,15 @@ class ForgotPasswordResetApi(Resource): @console_ns.response(400, "Invalid token or password mismatch") @setup_required @email_password_login_enabled - def post(self): - args = ForgotPasswordResetPayload.model_validate(console_ns.payload) + @model_validate(ForgotPasswordResetPayload) + def post(self, req_data: ForgotPasswordResetPayload): # Validate passwords match - if args.new_password != args.password_confirm: + if req_data.new_password != req_data.password_confirm: raise PasswordMismatchError() # Validate token and get reset data - reset_data = AccountService.get_reset_password_data(args.token) + reset_data = AccountService.get_reset_password_data(req_data.token) if not reset_data: raise InvalidTokenError() # Must use token in reset phase @@ -172,11 +172,11 @@ class ForgotPasswordResetApi(Resource): raise InvalidTokenError() # Revoke token to prevent reuse - AccountService.revoke_reset_password_token(args.token) + AccountService.revoke_reset_password_token(req_data.token) # Generate secure salt and hash password salt = secrets.token_bytes(16) - password_hashed = hash_password(args.new_password, salt) + password_hashed = hash_password(req_data.new_password, salt) email = reset_data.get("email", "") account = AccountService.get_account_by_email_with_case_fallback(email, session=db.session()) diff --git a/api/controllers/console/auth/login.py b/api/controllers/console/auth/login.py index 49b248a1e48..78d7038cc76 100644 --- a/api/controllers/console/auth/login.py +++ b/api/controllers/console/auth/login.py @@ -25,6 +25,8 @@ from controllers.console.auth.error import ( EmailPasswordLoginLimitError, InvalidEmailError, InvalidTokenError, + TurnstileServiceUnavailableError, + TurnstileVerificationFailedError, ) from controllers.console.error import ( AccountBannedError, @@ -39,9 +41,11 @@ from controllers.console.wraps import ( decrypt_code_field, decrypt_password_field, email_password_login_enabled, + model_validate, setup_required, with_current_user, ) +from enums import DeploymentEdition from extensions.ext_database import db from libs.helper import EmailStr, extract_remote_ip from libs.helper import timezone as validate_timezone_string @@ -66,6 +70,11 @@ from services.errors.account import ( ) from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError from services.feature_service import FeatureService +from services.turnstile_service import ( + TurnstileChallengeRejectedError, + TurnstileService, + TurnstileUpstreamError, +) logger = logging.getLogger(__name__) @@ -80,6 +89,13 @@ class EmailPayload(BaseModel): language: str | None = Field(default=None) +class EmailCodeSendPayload(EmailPayload): + turnstile_token: str | None = Field( + default=None, + description="Cloudflare Turnstile token. Required at runtime for Dify Cloud.", + ) + + class EmailCodeLoginPayload(BaseModel): email: EmailStr = Field(...) code: str = Field(...) @@ -95,7 +111,7 @@ class EmailCodeLoginPayload(BaseModel): return validate_timezone_string(value) -register_schema_models(console_ns, LoginPayload, EmailPayload, EmailCodeLoginPayload) +register_schema_models(console_ns, LoginPayload, EmailPayload, EmailCodeSendPayload, EmailCodeLoginPayload) register_response_schema_models( console_ns, SimpleResultDataResponse, @@ -114,13 +130,15 @@ class LoginApi(Resource): @console_ns.expect(console_ns.models[LoginPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[SimpleResultOptionalDataResponse.__name__]) @decrypt_password_field - def post(self): + @model_validate(LoginPayload) + def post(self, req_data: LoginPayload): """Authenticate user and login.""" - args = LoginPayload.model_validate(console_ns.payload) - request_email = args.email + request_email = req_data.email normalized_email = request_email.lower() - if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(normalized_email): + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and BillingService.is_email_in_freeze( + normalized_email + ): _log_console_login_failure(email=normalized_email, reason=LoginFailureReason.ACCOUNT_IN_FREEZE) raise AccountInFreezeError() @@ -129,7 +147,7 @@ class LoginApi(Resource): _log_console_login_failure(email=normalized_email, reason=LoginFailureReason.LOGIN_RATE_LIMITED) raise EmailPasswordLoginLimitError() - invite_token = args.invite_token + invite_token = req_data.invite_token invitation_data: InvitationDetailDict | None = None if invite_token: invitation_data = RegisterService.get_invitation_with_case_fallback( @@ -150,7 +168,7 @@ class LoginApi(Resource): ) raise InvalidEmailError() account = _authenticate_account_with_case_fallback( - request_email, normalized_email, args.password, invite_token + request_email, normalized_email, req_data.password, invite_token ) except services.errors.account.AccountLoginError: _log_console_login_failure(email=normalized_email, reason=LoginFailureReason.ACCOUNT_BANNED) @@ -159,7 +177,6 @@ class LoginApi(Resource): AccountService.add_login_error_rate_limit(normalized_email) _log_console_login_failure(email=normalized_email, reason=LoginFailureReason.INVALID_CREDENTIALS) raise AuthenticationFailedError() from exc - # SELF_HOSTED only have one workspace tenants = TenantService.get_join_tenants(account, session=db.session()) if len(tenants) == 0: if ( @@ -215,16 +232,16 @@ class ResetPasswordSendEmailApi(Resource): @email_password_login_enabled @console_ns.expect(console_ns.models[EmailPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[SimpleResultDataResponse.__name__]) - def post(self): - args = EmailPayload.model_validate(console_ns.payload) - normalized_email = args.email.lower() + @model_validate(EmailPayload) + def post(self, req_data: EmailPayload): + normalized_email = req_data.email.lower() - if args.language is not None and args.language == "zh-Hans": + if req_data.language is not None and req_data.language == "zh-Hans": language = "zh-Hans" else: language = "en-US" try: - account = _get_account_with_case_fallback(args.email) + account = _get_account_with_case_fallback(req_data.email) except AccountRegisterError: raise AccountInFreezeError() @@ -241,22 +258,32 @@ class ResetPasswordSendEmailApi(Resource): @console_ns.route("/email-code-login") class EmailCodeLoginSendEmailApi(Resource): @setup_required - @console_ns.expect(console_ns.models[EmailPayload.__name__]) + @console_ns.expect(console_ns.models[EmailCodeSendPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[SimpleResultDataResponse.__name__]) - def post(self): - args = EmailPayload.model_validate(console_ns.payload) - normalized_email = args.email.lower() + @model_validate(EmailCodeSendPayload) + def post(self, req_data: EmailCodeSendPayload): + normalized_email = req_data.email.lower() ip_address = extract_remote_ip(request) if AccountService.is_email_send_ip_limit(ip_address): raise EmailSendIpLimitError() - if args.language is not None and args.language == "zh-Hans": + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: + try: + TurnstileService.verify(token=req_data.turnstile_token, remote_ip=ip_address) + except TurnstileChallengeRejectedError as exc: + logger.info("Turnstile rejected an email-code login challenge") + raise TurnstileVerificationFailedError() from exc + except TurnstileUpstreamError as exc: + logger.warning("Turnstile verification is unavailable", exc_info=True) + raise TurnstileServiceUnavailableError() from exc + + if req_data.language is not None and req_data.language == "zh-Hans": language = "zh-Hans" else: language = "en-US" try: - account = _get_account_with_case_fallback(args.email) + account = _get_account_with_case_fallback(req_data.email) except AccountRegisterError: raise AccountInFreezeError() @@ -277,14 +304,14 @@ class EmailCodeLoginApi(Resource): @console_ns.expect(console_ns.models[EmailCodeLoginPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) @decrypt_code_field - def post(self): - args = EmailCodeLoginPayload.model_validate(console_ns.payload) + @model_validate(EmailCodeLoginPayload) + def post(self, req_data: EmailCodeLoginPayload): - original_email = args.email + original_email = req_data.email user_email = original_email.lower() - language = args.language + language = req_data.language - token_data = AccountService.get_email_code_login_data(args.token) + token_data = AccountService.get_email_code_login_data(req_data.token) if token_data is None: _log_console_login_failure(email=user_email, reason=LoginFailureReason.INVALID_EMAIL_CODE_TOKEN) raise InvalidTokenError() @@ -295,11 +322,11 @@ class EmailCodeLoginApi(Resource): _log_console_login_failure(email=user_email, reason=LoginFailureReason.EMAIL_CODE_EMAIL_MISMATCH) raise InvalidEmailError() - if token_data["code"] != args.code: + if token_data["code"] != req_data.code: _log_console_login_failure(email=user_email, reason=LoginFailureReason.INVALID_EMAIL_CODE) raise EmailCodeError() - AccountService.revoke_email_code_login_token(args.token) + AccountService.revoke_email_code_login_token(req_data.token) try: account = _get_account_with_case_fallback(original_email) except Unauthorized as exc: @@ -325,7 +352,7 @@ class EmailCodeLoginApi(Resource): email=user_email, name=user_email, interface_language=get_valid_language(language), - timezone=args.timezone, + timezone=req_data.timezone, session=db.session(), ) except WorkSpaceNotAllowedCreateError: diff --git a/api/controllers/console/auth/oauth.py b/api/controllers/console/auth/oauth.py index c6b80fbb392..0e1b28c58f7 100644 --- a/api/controllers/console/auth/oauth.py +++ b/api/controllers/console/auth/oauth.py @@ -12,6 +12,7 @@ from configs import dify_config from constants.languages import languages from controllers.common.fields import RedirectResponse from controllers.common.schema import query_params_from_model, register_response_schema_model, register_schema_models +from enums import DeploymentEdition from extensions.ext_database import db from libs.datetime_utils import naive_utc_now from libs.helper import extract_remote_ip @@ -301,7 +302,9 @@ def _generate_account( normalized_email = user_info.email.lower() oauth_new_user = True if not FeatureService.get_system_features().is_allow_register: - if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(normalized_email): + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and BillingService.is_email_in_freeze( + normalized_email + ): raise AccountRegisterError( description=( "This email account has been deleted within the past " diff --git a/api/controllers/console/auth/oauth_server.py b/api/controllers/console/auth/oauth_server.py index 1c2cd07ff5a..7eb043b2fd1 100644 --- a/api/controllers/console/auth/oauth_server.py +++ b/api/controllers/console/auth/oauth_server.py @@ -9,7 +9,7 @@ from pydantic import BaseModel from werkzeug.exceptions import BadRequest, NotFound from controllers.common.schema import register_response_schema_models, register_schema_models -from controllers.console.wraps import account_initialization_required, setup_required, with_current_user +from controllers.console.wraps import account_initialization_required, model_validate, setup_required, with_current_user from core.db.session_factory import session_factory from graphon.model_runtime.utils.encoders import jsonable_encoder from libs.login import login_required @@ -154,8 +154,8 @@ class OAuthServerAppApi(Resource): @console_ns.expect(console_ns.models[OAuthProviderRequest.__name__]) @console_ns.response(200, "Success", console_ns.models[OAuthProviderAppResponse.__name__]) @oauth_server_client_id_required - def post(self, oauth_provider_app: OAuthProviderApp): - payload = OAuthProviderRequest.model_validate(request.get_json()) + @model_validate(OAuthProviderRequest) + def post(self, payload: OAuthProviderRequest, oauth_provider_app: OAuthProviderApp): redirect_uri = payload.redirect_uri # check if redirect_uri is valid @@ -196,9 +196,8 @@ class OAuthServerUserTokenApi(Resource): @console_ns.expect(console_ns.models[OAuthTokenRequest.__name__]) @console_ns.response(200, "Success", console_ns.models[OAuthProviderTokenResponse.__name__]) @oauth_server_client_id_required - def post(self, oauth_provider_app: OAuthProviderApp): - payload = OAuthTokenRequest.model_validate(request.get_json()) - + @model_validate(OAuthTokenRequest) + def post(self, payload: OAuthTokenRequest, oauth_provider_app: OAuthProviderApp): try: grant_type = OAuthGrantType(payload.grant_type) except ValueError: diff --git a/api/controllers/console/billing/billing.py b/api/controllers/console/billing/billing.py index 3a983b50176..da162ac49a8 100644 --- a/api/controllers/console/billing/billing.py +++ b/api/controllers/console/billing/billing.py @@ -1,7 +1,6 @@ import base64 from typing import Any, Literal -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, RootModel from werkzeug.exceptions import BadRequest @@ -10,12 +9,13 @@ from controllers.common.schema import query_params_from_model, register_response from controllers.console import console_ns from controllers.console.wraps import ( account_initialization_required, + model_validate, only_edition_cloud, setup_required, with_current_tenant_id, with_current_user, ) -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from extensions.ext_database import db from fields.base import ResponseModel from libs.login import login_required @@ -40,24 +40,28 @@ class BillingInvoiceResponse(ResponseModel): url: str +class BillingSubscriptionResponse(ResponseModel): + url: str + + register_schema_models(console_ns, SubscriptionQuery, PartnerTenantsPayload) -register_response_schema_models(console_ns, BillingResponse, BillingInvoiceResponse) +register_response_schema_models(console_ns, BillingResponse, BillingInvoiceResponse, BillingSubscriptionResponse) @console_ns.route("/billing/subscription") class Subscription(Resource): @console_ns.doc(params=query_params_from_model(SubscriptionQuery)) - @console_ns.response(200, "Success", console_ns.models[BillingResponse.__name__]) + @console_ns.response(200, "Success", console_ns.models[BillingSubscriptionResponse.__name__]) @setup_required @login_required @account_initialization_required @only_edition_cloud @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account): - args = SubscriptionQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(SubscriptionQuery) + def get(self, req_data: SubscriptionQuery, current_tenant_id: str, current_user: Account): BillingService.is_tenant_owner_or_admin(current_user, session=db.session()) - return BillingService.get_subscription(args.plan, args.interval, current_user.email, current_tenant_id) + return BillingService.get_subscription(req_data.plan, req_data.interval, current_user.email, current_tenant_id) @console_ns.route("/billing/invoices") @@ -87,10 +91,10 @@ class PartnerTenants(Resource): @account_initialization_required @only_edition_cloud @with_current_user - def put(self, current_user: Account, partner_key: str): + @model_validate(PartnerTenantsPayload) + def put(self, req_data: PartnerTenantsPayload, current_user: Account, partner_key: str): try: - args = PartnerTenantsPayload.model_validate(console_ns.payload or {}) - click_id = args.click_id + click_id = req_data.click_id decoded_partner_key = base64.b64decode(partner_key).decode("utf-8") except Exception: raise BadRequest("Invalid partner_key") diff --git a/api/controllers/console/billing/compliance.py b/api/controllers/console/billing/compliance.py index ea5852586a9..f62d6f8e417 100644 --- a/api/controllers/console/billing/compliance.py +++ b/api/controllers/console/billing/compliance.py @@ -14,6 +14,7 @@ from ...common.schema import DEFAULT_REF_TEMPLATE_OPENAPI_3_0 from .. import console_ns from ..wraps import ( account_initialization_required, + model_validate, only_edition_cloud, setup_required, with_current_tenant_id, @@ -48,13 +49,13 @@ class ComplianceApi(Resource): @only_edition_cloud @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account): - args = ComplianceDownloadQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(ComplianceDownloadQuery) + def get(self, req_data: ComplianceDownloadQuery, current_tenant_id: str, current_user: Account): ip_address = extract_remote_ip(request) device_info = request.headers.get("User-Agent", "Unknown device") return BillingService.get_compliance_download_link( - doc_name=args.doc_name, + doc_name=req_data.doc_name, account_id=current_user.id, tenant_id=current_tenant_id, ip=ip_address, diff --git a/api/controllers/console/datasets/data_source.py b/api/controllers/console/datasets/data_source.py index 590fbdbb87d..9cf8420c251 100644 --- a/api/controllers/console/datasets/data_source.py +++ b/api/controllers/console/datasets/data_source.py @@ -36,6 +36,7 @@ from ..wraps import ( RBACPermission, RBACResourceScope, account_initialization_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -229,12 +230,18 @@ class DataSourceNotionListApi(Resource): @with_current_user @with_current_tenant_id @with_session(write=False) - def get(self, session: Session, current_tenant_id: str, current_user: Account) -> tuple[dict[str, Any], int]: - query = DataSourceNotionListQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(DataSourceNotionListQuery) + def get( + self, + req_data: DataSourceNotionListQuery, + session: Session, + current_tenant_id: str, + current_user: Account, + ) -> tuple[dict[str, Any], int]: datasource_provider_service = DatasourceProviderService() credential = datasource_provider_service.get_datasource_credentials( tenant_id=current_tenant_id, - credential_id=query.credential_id, + credential_id=req_data.credential_id, provider="notion_datasource", plugin_id="langgenius/notion_datasource", ) @@ -242,8 +249,8 @@ class DataSourceNotionListApi(Resource): raise NotFound("Credential not found.") exist_page_ids = [] # import notion in the exist dataset - if query.dataset_id: - dataset = DatasetService.get_dataset(query.dataset_id, session) + if req_data.dataset_id: + dataset = DatasetService.get_dataset(req_data.dataset_id, session) if not dataset: raise NotFound("Dataset not found.") if dataset.data_source_type != "notion_import": @@ -251,7 +258,7 @@ class DataSourceNotionListApi(Resource): documents = session.scalars( select(Document).where( - Document.dataset_id == query.dataset_id, + Document.dataset_id == req_data.dataset_id, Document.tenant_id == current_tenant_id, Document.data_source_type == "notion_import", Document.enabled.is_(True), @@ -318,13 +325,19 @@ class DataSourceNotionPreviewApi(Resource): @console_ns.doc(params=query_params_from_model(DataSourceNotionPreviewQuery)) @console_ns.response(200, "Success", console_ns.models[TextContentResponse.__name__]) @with_current_tenant_id - def get(self, current_tenant_id: str, page_id: UUID, page_type: str) -> tuple[dict[str, str], int]: - query = DataSourceNotionPreviewQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(DataSourceNotionPreviewQuery) + def get( + self, + req_data: DataSourceNotionPreviewQuery, + current_tenant_id: str, + page_id: UUID, + page_type: str, + ) -> tuple[dict[str, str], int]: datasource_provider_service = DatasourceProviderService() credential = datasource_provider_service.get_datasource_credentials( tenant_id=current_tenant_id, - credential_id=query.credential_id, + credential_id=req_data.credential_id, provider="notion_datasource", plugin_id="langgenius/notion_datasource", ) @@ -354,12 +367,17 @@ class DataSourceNotionIndexingEstimateApi(Resource): @console_ns.response(200, "Success", console_ns.models[IndexingEstimate.__name__]) @with_current_tenant_id @with_session - def post(self, session: Session, current_tenant_id: str) -> tuple[dict[str, Any], int]: - payload = NotionEstimatePayload.model_validate(console_ns.payload or {}) - args = payload.model_dump() + @model_validate(NotionEstimatePayload) + def post( + self, + req_data: NotionEstimatePayload, + session: Session, + current_tenant_id: str, + ) -> tuple[dict[str, Any], int]: + args = req_data.model_dump() # validate args DocumentService.estimate_args_validate(args) - notion_info_list = payload.notion_info_list + notion_info_list = req_data.notion_info_list extract_settings = [] for notion_info in notion_info_list: workspace_id = notion_info["workspace_id"] diff --git a/api/controllers/console/datasets/datasets.py b/api/controllers/console/datasets/datasets.py index 2ec5a9fe103..aedc9cf8940 100644 --- a/api/controllers/console/datasets/datasets.py +++ b/api/controllers/console/datasets/datasets.py @@ -26,6 +26,7 @@ from controllers.console.wraps import ( cloud_edition_billing_rate_limit_check, enterprise_license_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -52,6 +53,7 @@ from models.enums import ApiTokenType, SegmentStatus from models.provider_ids import ModelProviderID from services.api_token_service import ApiTokenCache from services.app_service import AppService +from services.dataset_ref_service import DatasetRefService from services.dataset_service import DatasetPermissionService, DatasetService, DocumentService from services.enterprise import rbac_service as enterprise_rbac_service from services.enterprise.rbac_service import RBACResourceWhitelistScope, ReplaceMemberBindings @@ -66,6 +68,18 @@ def _has_dataset_list_permission(permission_keys: list[str]) -> bool: return any(permission_key in DATASET_LIST_PERMISSION_KEYS for permission_key in permission_keys) +def _get_accessible_dataset(dataset_id: UUID, tenant_id: str, current_user: Account, session: Session) -> Dataset: + dataset = DatasetService.get_dataset_for_tenant(str(dataset_id), tenant_id, session=session) + if dataset is None: + raise NotFound("Dataset not found.") + if not dify_config.RBAC_ENABLED: + try: + DatasetService.check_dataset_permission(dataset, current_user, session) + except services.errors.account.NoPermissionError as e: + raise Forbidden(str(e)) + return dataset + + def _validate_indexing_technique(value: str | None) -> str | None: if value is None: return value @@ -218,7 +232,7 @@ class _DatasetQueryResponseSource: return self.query.get_queries(session=self.session) def __getattr__(self, name: str) -> Any: - return getattr(self.query, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self.query, name) # guard-ignore: no-new-getattr -- delegates model fields class DatasetQueryListResponse(ResponseModel): @@ -257,7 +271,7 @@ class _RelatedAppResponseSource: return self.app.mode_compatible_with_agent_with_session(session=self.session) def __getattr__(self, name: str) -> Any: - return getattr(self.app, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self.app, name) # guard-ignore: no-new-getattr -- delegates model fields class RelatedAppListResponse(ResponseModel): @@ -360,7 +374,6 @@ def _get_retrieval_methods_by_vector_type(vector_type: str | None, is_mock: bool # Define vector database types that only support semantic search semantic_only_types = { VectorType.RELYT, - VectorType.TIDB_VECTOR, VectorType.CHROMA, VectorType.PGVECTO_RS, VectorType.VIKINGDB, @@ -408,6 +421,9 @@ def _get_retrieval_methods_by_vector_type(vector_type: str | None, is_mock: bool if vector_type == VectorType.MILVUS: return semantic_methods if is_mock else full_methods + if vector_type == VectorType.TIDB_VECTOR: + return full_methods if dify_config.TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH else semantic_methods + if vector_type in semantic_only_types: return semantic_methods elif vector_type in full_search_types: @@ -570,9 +586,8 @@ class DatasetListApi(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, current_tenant_id: str, current_user: Account): - payload = DatasetCreatePayload.model_validate(console_ns.payload or {}) - + @model_validate(DatasetCreatePayload) + def post(self, req_data: DatasetCreatePayload, session: Session, current_tenant_id: str, current_user: Account): # The role of the current user in the ta table must be admin, owner, or editor, or dataset_operator if not current_user.is_dataset_editor: raise Forbidden() @@ -580,20 +595,20 @@ class DatasetListApi(Resource): if dify_config.RBAC_ENABLED: permission = DatasetPermissionEnum.ALL_TEAM else: - permission = payload.permission or DatasetPermissionEnum.ONLY_ME + permission = req_data.permission or DatasetPermissionEnum.ONLY_ME try: dataset = DatasetService.create_empty_dataset( session=session, tenant_id=current_tenant_id, - name=payload.name, - description=payload.description, - indexing_technique=payload.indexing_technique, + name=req_data.name, + description=req_data.description, + indexing_technique=req_data.indexing_technique, account=current_user, permission=permission, - provider=payload.provider, - external_knowledge_api_id=payload.external_knowledge_api_id, - external_knowledge_id=payload.external_knowledge_id, + provider=req_data.provider, + external_knowledge_api_id=req_data.external_knowledge_api_id, + external_knowledge_id=req_data.external_knowledge_id, ) except services.errors.dataset.DatasetNameDuplicateError: raise DatasetNameDuplicateError() @@ -713,28 +728,35 @@ class DatasetApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def patch(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): + @model_validate(DatasetUpdatePayload) + def patch( + self, + req_data: DatasetUpdatePayload, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + ): dataset_id_str = str(dataset_id) dataset = DatasetService.get_dataset(dataset_id_str, session) if dataset is None: raise NotFound("Dataset not found.") - payload = DatasetUpdatePayload.model_validate(console_ns.payload or {}) # check embedding model setting if ( - payload.indexing_technique == IndexTechniqueType.HIGH_QUALITY - and payload.embedding_model_provider is not None - and payload.embedding_model is not None + req_data.indexing_technique == IndexTechniqueType.HIGH_QUALITY + and req_data.embedding_model_provider is not None + and req_data.embedding_model is not None ): is_multimodal = DatasetService.check_is_multimodal_model( - dataset.tenant_id, payload.embedding_model_provider, payload.embedding_model + dataset.tenant_id, req_data.embedding_model_provider, req_data.embedding_model ) - payload.is_multimodal = is_multimodal - payload_data = payload.model_dump(exclude_unset=True) + req_data.is_multimodal = is_multimodal + payload_data = req_data.model_dump(exclude_unset=True) # The role of the current user in the ta table must be admin, owner, editor, or dataset_operator if not dify_config.RBAC_ENABLED: DatasetPermissionService.check_permission( - current_user, dataset, payload.permission, payload.partial_member_list, session=session + current_user, dataset, req_data.permission, req_data.partial_member_list, session=session ) dataset = DatasetService.update_dataset(dataset_id_str, payload_data, current_user, session=session) @@ -752,12 +774,12 @@ class DatasetApi(Resource): result_data["permission_keys"] = permission_keys_map.get(dataset_id_str, []) tenant_id = current_tenant_id - if payload.partial_member_list is not None and payload.permission == DatasetPermissionEnum.PARTIAL_TEAM: + if req_data.partial_member_list is not None and req_data.permission == DatasetPermissionEnum.PARTIAL_TEAM: DatasetPermissionService.update_partial_member_list( - tenant_id, dataset_id_str, payload.partial_member_list, session + tenant_id, dataset_id_str, req_data.partial_member_list, session ) # clear partial member list when permission is only_me or all_team_members - elif payload.permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.ALL_TEAM}: + elif req_data.permission in {DatasetPermissionEnum.ONLY_ME, DatasetPermissionEnum.ALL_TEAM}: DatasetPermissionService.clear_partial_member_list(dataset_id_str, session) partial_member_list = DatasetPermissionService.get_dataset_partial_member_list(dataset_id_str, session) @@ -802,12 +824,13 @@ class DatasetUseCheckApi(Resource): @setup_required @login_required @account_initialization_required + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) @with_session(write=False) - def get(self, session: Session, dataset_id: UUID): - dataset_id_str = str(dataset_id) - - dataset_is_using = DatasetService.dataset_use_check(dataset_id_str, session) + def get(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): + dataset = _get_accessible_dataset(dataset_id, current_tenant_id, current_user, session) + dataset_is_using = DatasetService.dataset_use_check(DatasetRefService.create_dataset_ref(dataset), session) return UsageCheckResponse(is_using=dataset_is_using).model_dump(mode="json"), 200 @@ -870,9 +893,9 @@ class DatasetIndexingEstimateApi(Resource): @console_ns.expect(console_ns.models[IndexingEstimatePayload.__name__]) @with_current_tenant_id @with_session - def post(self, session: Session, current_tenant_id: str): - payload = IndexingEstimatePayload.model_validate(console_ns.payload or {}) - args = payload.model_dump() + @model_validate(IndexingEstimatePayload) + def post(self, req_data: IndexingEstimatePayload, session: Session, current_tenant_id: str): + args = req_data.model_dump() # validate args DocumentService.estimate_args_validate(args) extract_settings = [] @@ -1018,13 +1041,15 @@ class DatasetIndexingStatusApi(Resource): @setup_required @login_required @account_initialization_required + @with_current_user @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) @with_session(write=False) - def get(self, session: Session, current_tenant_id: str, dataset_id: UUID): - dataset_id_str = str(dataset_id) + def get(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): + dataset = _get_accessible_dataset(dataset_id, current_tenant_id, current_user, session) + documents = session.scalars( - select(Document).where(Document.dataset_id == dataset_id_str, Document.tenant_id == current_tenant_id) + select(Document).where(Document.dataset_id == dataset.id, Document.tenant_id == dataset.tenant_id) ).all() documents_status = [] for document in documents: @@ -1032,6 +1057,8 @@ class DatasetIndexingStatusApi(Resource): session.scalar( select(func.count(DocumentSegment.id)).where( DocumentSegment.completed_at.isnot(None), + DocumentSegment.tenant_id == dataset.tenant_id, + DocumentSegment.dataset_id == dataset.id, DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT, ) @@ -1041,6 +1068,8 @@ class DatasetIndexingStatusApi(Resource): total_segments = ( session.scalar( select(func.count(DocumentSegment.id)).where( + DocumentSegment.tenant_id == dataset.tenant_id, + DocumentSegment.dataset_id == dataset.id, DocumentSegment.document_id == str(document.id), DocumentSegment.status != SegmentStatus.RE_SEGMENT, ) @@ -1077,6 +1106,8 @@ class DatasetApiKeyApi(Resource): @console_ns.response(200, "API keys retrieved successfully", console_ns.models[ApiKeyList.__name__]) @setup_required @login_required + @is_admin_or_owner_required + @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_API_KEY_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id @with_session(write=False) @@ -1168,12 +1199,16 @@ class DatasetEnableApiApi(Resource): @login_required @account_initialization_required @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def post(self, session: Session, dataset_id: UUID, status: str): - dataset_id_str = str(dataset_id) + def post(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID, status: str): + dataset = _get_accessible_dataset(dataset_id, current_tenant_id, current_user, session) + if not current_user.is_dataset_editor: + raise Forbidden() - DatasetService.update_dataset_api_status(dataset_id_str, status == "enable", session) + DatasetService.update_dataset_api_status(dataset, status == "enable", current_user, session) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -1239,14 +1274,15 @@ class DatasetErrorDocs(Resource): @setup_required @login_required @account_initialization_required + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) @with_session(write=False) - def get(self, session: Session, dataset_id: UUID): - dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) - if dataset is None: - raise NotFound("Dataset not found.") - results = DocumentService.get_error_documents_by_dataset_id(dataset_id_str, session) + def get(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): + dataset = _get_accessible_dataset(dataset_id, current_tenant_id, current_user, session) + results = DocumentService.get_error_documents_by_dataset_ref( + DatasetRefService.create_dataset_ref(dataset), session + ) return dump_response(ErrorDocsResponse, {"data": results, "total": len(results)}), 200 @@ -1298,12 +1334,13 @@ class DatasetAutoDisableLogApi(Resource): @setup_required @login_required @account_initialization_required + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_READONLY) @with_session(write=False) - def get(self, session: Session, dataset_id: UUID): - dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) - if dataset is None: - raise NotFound("Dataset not found.") - auto_disable_logs = DatasetService.get_dataset_auto_disable_logs(dataset_id_str, session) + def get(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): + dataset = _get_accessible_dataset(dataset_id, current_tenant_id, current_user, session) + auto_disable_logs = DatasetService.get_dataset_auto_disable_logs( + DatasetRefService.create_dataset_ref(dataset), session + ) return dump_response(AutoDisableLogsResponse, auto_disable_logs), 200 diff --git a/api/controllers/console/datasets/datasets_document.py b/api/controllers/console/datasets/datasets_document.py index 3d79af3ffc9..4e5d0efe668 100644 --- a/api/controllers/console/datasets/datasets_document.py +++ b/api/controllers/console/datasets/datasets_document.py @@ -16,12 +16,13 @@ from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services +from configs import dify_config from controllers.common.controller_schemas import DocumentBatchDownloadZipPayload from controllers.common.fields import SimpleResultMessageResponse, SimpleResultResponse, UrlResponse from controllers.common.schema import register_response_schema_models, register_schema_models from controllers.common.session import with_session from controllers.console import console_ns -from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required +from controllers.console.wraps import RBACPermission, RBACResourceScope, model_validate, rbac_permission_required from core.entities.knowledge_entities import IndexingEstimate from core.errors.error import ( LLMBadRequestError, @@ -54,13 +55,17 @@ from libs.helper import dump_response, to_timestamp from libs.login import login_required from libs.pagination import paginate_query from models import Account, Document, DocumentSegment, UploadFile -from models.dataset import DocumentPipelineExecutionLog +from models.dataset import DatasetPermissionEnum, DocumentPipelineExecutionLog from models.enums import IndexingStatus, ProcessRuleMode, SegmentStatus from services.dataset_ref_service import DatasetRefService from services.dataset_service import DatasetService, DocumentService +from services.enterprise import rbac_service as enterprise_rbac_service +from services.enterprise.rbac_service import RBACResourceWhitelistScope, ReplaceMemberBindings from services.entities.knowledge_entities.knowledge_entities import KnowledgeConfig, ProcessRule, RetrievalModel from services.file_service import FileService +from services.vector_space_admission_service import get_vector_space_admission_error_fields from tasks.generate_summary_index_task import generate_summary_index_task +from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task from ..app.error import ( ProviderModelCurrentlyNotSupportError, @@ -77,6 +82,7 @@ from ..datasets.error import ( ) from ..wraps import ( account_initialization_required, + check_knowledge_rate_limit, cloud_edition_billing_rate_limit_check, cloud_edition_billing_resource_check, setup_required, @@ -293,23 +299,23 @@ class DocumentResource(Resource): def get_document( self, session: Session, dataset_id: str, document_id: str, current_user: Account, current_tenant_id: str ) -> Document: - dataset = DatasetService.get_dataset(dataset_id, session) + dataset = DatasetService.get_dataset_for_tenant(dataset_id, current_tenant_id, session=session) if not dataset: raise NotFound("Dataset not found.") - try: - DatasetService.check_dataset_permission(dataset, current_user, session) - except services.errors.account.NoPermissionError as e: - raise Forbidden(str(e)) + if not dify_config.RBAC_ENABLED: + try: + DatasetService.check_dataset_permission(dataset, current_user, session) + except services.errors.account.NoPermissionError as e: + raise Forbidden(str(e)) - document = DocumentService.get_document(dataset_id, document_id, session=session) + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + document_ref = DatasetRefService.create_document_ref_from_id(dataset_ref, document_id) + document = DatasetRefService.get_document_by_ref(document_ref, session=session) if not document: raise NotFound("Document not found.") - if document.tenant_id != current_tenant_id: - raise Forbidden("No permission.") - return document def get_batch_documents( @@ -577,18 +583,33 @@ class DatasetDocumentListApi(Resource): @setup_required @login_required @account_initialization_required - @cloud_edition_billing_rate_limit_check("knowledge") @console_ns.response(204, "Documents deleted successfully") + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def delete(self, session: Session, dataset_id: UUID): + def delete( + self, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + ): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, current_tenant_id, session=session) if dataset is None: raise NotFound("Dataset not found.") - # check user's model setting - DatasetService.check_dataset_model_setting(dataset) + if not current_user.is_dataset_editor: + raise Forbidden() + + if not dify_config.RBAC_ENABLED: + try: + DatasetService.check_dataset_permission(dataset, current_user, session) + except services.errors.account.NoPermissionError as e: + raise Forbidden(str(e)) + + check_knowledge_rate_limit() try: document_ids = request.args.getlist("document_id") dataset_ref = DatasetRefService.create_dataset_ref(dataset) @@ -661,6 +682,21 @@ class DatasetInitApi(Resource): except ModelCurrentlyNotSupportError: raise ProviderModelCurrentlyNotSupportError() + if dify_config.RBAC_ENABLED: + dataset.permission = DatasetPermissionEnum.ALL_TEAM + else: + dataset.permission = DatasetPermissionEnum.ONLY_ME + session.flush() + + if dify_config.RBAC_ENABLED: + enterprise_rbac_service.RBACService.DatasetAccess.replace_whitelist( + current_tenant_id, + current_user.id, + dataset.id, + ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL), + ) + initialize_created_app_rbac_access_task.delay(current_tenant_id, current_user.id, dataset_id=dataset.id) + return dump_response( DatasetAndDocumentResponse, {"dataset": dataset, "documents": document_responses(documents, session=session), "batch": batch}, @@ -935,6 +971,7 @@ class DocumentBatchIndexingStatusApi(DocumentResource): "completed_at": document.completed_at, "paused_at": document.paused_at, "error": document.error, + **get_vector_space_admission_error_fields(document.error), "stopped_at": document.stopped_at, "completed_segments": completed_segments, "total_segments": total_segments, @@ -995,6 +1032,7 @@ class DocumentIndexingStatusApi(DocumentResource): "completed_at": document.completed_at, "paused_at": document.paused_at, "error": document.error, + **get_vector_space_admission_error_fields(document.error), "stopped_at": document.stopped_at, "completed_segments": completed_segments, "total_segments": total_segments, @@ -1264,13 +1302,20 @@ class DocumentMetadataApi(DocumentResource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def put(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID, document_id: UUID): + @model_validate(DocumentMetadataUpdatePayload) + def put( + self, + req_data: DocumentMetadataUpdatePayload, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + document_id: UUID, + ): dataset_id_str = str(dataset_id) document_id_str = str(document_id) document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) - req_data = DocumentMetadataUpdatePayload.model_validate(request.get_json() or {}) - doc_type = req_data.doc_type doc_metadata = req_data.doc_metadata @@ -1357,29 +1402,32 @@ class DocumentPauseApi(DocumentResource): @setup_required @login_required @account_initialization_required - @cloud_edition_billing_rate_limit_check("knowledge") @console_ns.response(204, "Document paused successfully") + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def patch(self, session: Session, dataset_id: UUID, document_id: UUID): + def patch( + self, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + document_id: UUID, + ): """pause document.""" dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) - if not dataset: - raise NotFound("Dataset not found.") - - document = DocumentService.get_document(dataset.id, document_id_str, session=session) - - # 404 if document not found - if document is None: - raise NotFound("Document Not Exists.") + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) + if not current_user.is_dataset_editor: + raise Forbidden() # 403 if document is archived if DocumentService.check_archived(document): raise ArchivedDocumentImmutableError() + check_knowledge_rate_limit() try: # pause document DocumentService.pause_document(document, session) @@ -1394,26 +1442,31 @@ class DocumentRecoverApi(DocumentResource): @setup_required @login_required @account_initialization_required - @cloud_edition_billing_rate_limit_check("knowledge") @console_ns.response(204, "Document resumed successfully") + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def patch(self, session: Session, dataset_id: UUID, document_id: UUID): + def patch( + self, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + document_id: UUID, + ): """recover document.""" dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) - if not dataset: - raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str, session=session) - # 404 if document not found - if document is None: - raise NotFound("Document Not Exists.") + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) + if not current_user.is_dataset_editor: + raise Forbidden() # 403 if document is archived if DocumentService.check_archived(document): raise ArchivedDocumentImmutableError() + check_knowledge_rate_limit() try: # pause document DocumentService.recover_document(document, session) @@ -1428,22 +1481,43 @@ class DocumentRetryApi(DocumentResource): @setup_required @login_required @account_initialization_required - @cloud_edition_billing_rate_limit_check("knowledge") @console_ns.expect(console_ns.models[DocumentRetryPayload.__name__]) @console_ns.response(204, "Documents retry started successfully") + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def post(self, session: Session, dataset_id: UUID): + def post( + self, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + ): """retry document.""" - payload = DocumentRetryPayload.model_validate(console_ns.payload or {}) dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) - retry_documents = [] + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, current_tenant_id, session=session) if not dataset: raise NotFound("Dataset not found.") + + if not current_user.is_dataset_editor: + raise Forbidden() + + if not dify_config.RBAC_ENABLED: + try: + DatasetService.check_dataset_permission(dataset, current_user, session) + except services.errors.account.NoPermissionError as e: + raise Forbidden(str(e)) + + payload = DocumentRetryPayload.model_validate(console_ns.payload or {}) + documents = DocumentService.get_documents_by_ids( + DatasetRefService.create_dataset_ref(dataset), payload.document_ids, session + ) + documents_by_id = {document.id: document for document in documents} + retry_documents = [] for document_id in payload.document_ids: try: - document = DocumentService.get_document(dataset.id, document_id, session=session) + document = documents_by_id.get(document_id) # 404 if document not found if document is None: @@ -1461,6 +1535,7 @@ class DocumentRetryApi(DocumentResource): logger.exception("Failed to retry document, document id: %s", document_id) continue # retry document + check_knowledge_rate_limit() DocumentService.retry_document(dataset_id_str, retry_documents, session) return "", 204 @@ -1500,28 +1575,34 @@ class WebsiteDocumentSyncApi(DocumentResource): @login_required @account_initialization_required @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) + @with_current_user @with_current_tenant_id - @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) + @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def get(self, session: Session, current_tenant_id: str, dataset_id: UUID, document_id: UUID): + def get( + self, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + document_id: UUID, + ): """sync website document.""" dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, current_tenant_id, session=session) if not dataset: raise NotFound("Dataset not found.") + if not current_user.is_dataset_editor: + raise Forbidden() document_id_str = str(document_id) - document = DocumentService.get_document(dataset.id, document_id_str, session=session) - if not document: - raise NotFound("Document not found.") - if document.tenant_id != current_tenant_id: - raise Forbidden("No permission.") + document = self.get_document(session, dataset.id, document_id_str, current_user, current_tenant_id) if document.data_source_type != "website_crawl": raise ValueError("Document is not a website document.") # 403 if document is archived if DocumentService.check_archived(document): raise ArchivedDocumentImmutableError() # sync document - DocumentService.sync_website_document(dataset_id_str, document, session) + DocumentService.sync_website_document(dataset, document, session) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -1536,21 +1617,25 @@ class DocumentPipelineExecutionLogApi(DocumentResource): @setup_required @login_required @account_initialization_required + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) @with_session(write=False) - def get(self, session: Session, dataset_id: UUID, document_id: UUID): + def get( + self, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + document_id: UUID, + ): dataset_id_str = str(dataset_id) document_id_str = str(document_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) - if not dataset: - raise NotFound("Dataset not found.") - document = DocumentService.get_document(dataset.id, document_id_str, session=session) - if not document: - raise NotFound("Document not found.") + document = self.get_document(session, dataset_id_str, document_id_str, current_user, current_tenant_id) log = session.scalar( select(DocumentPipelineExecutionLog) - .where(DocumentPipelineExecutionLog.document_id == document_id_str) + .where(DocumentPipelineExecutionLog.document_id == document.id) .order_by(DocumentPipelineExecutionLog.created_at.desc()) .limit(1) ) @@ -1633,7 +1718,9 @@ class DocumentGenerateSummaryApi(Resource): raise ValueError("Summary index is not enabled for this dataset. Please enable it in the dataset settings.") # Verify all documents exist and belong to the dataset - documents = DocumentService.get_documents_by_ids(dataset_id_str, document_list, session) + documents = DocumentService.get_documents_by_ids( + DatasetRefService.create_dataset_ref(dataset), document_list, session + ) if len(documents) != len(document_list): found_ids = {doc.id for doc in documents} diff --git a/api/controllers/console/datasets/datasets_segments.py b/api/controllers/console/datasets/datasets_segments.py index 33a5a1f752a..d1fa14902e6 100644 --- a/api/controllers/console/datasets/datasets_segments.py +++ b/api/controllers/console/datasets/datasets_segments.py @@ -36,6 +36,7 @@ from controllers.console.wraps import ( cloud_edition_billing_knowledge_limit_check, cloud_edition_billing_rate_limit_check, cloud_edition_billing_resource_check, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -414,8 +415,10 @@ class DatasetDocumentSegmentAddApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session + @model_validate(SegmentCreatePayload) def post( self, + req_data: SegmentCreatePayload, session: Session, current_tenant_id: str, current_user: Account, @@ -455,8 +458,7 @@ class DatasetDocumentSegmentAddApi(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # validate args - payload = SegmentCreatePayload.model_validate(console_ns.payload or {}) - payload_dict = payload.model_dump(exclude_none=True) + payload_dict = req_data.model_dump(exclude_none=True) SegmentService.segment_create_args_validate(payload_dict, document) segment = type_cast(DocumentSegment, SegmentService.create_segment(payload_dict, document, dataset, session)) summary = SummaryIndexService.get_segment_summary( @@ -485,8 +487,10 @@ class DatasetDocumentSegmentUpdateApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session + @model_validate(SegmentUpdatePayload) def patch( self, + req_data: SegmentUpdatePayload, session: Session, current_tenant_id: str, current_user: Account, @@ -532,13 +536,12 @@ class DatasetDocumentSegmentUpdateApi(Resource): segment_id_str = str(segment_id) _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) # validate args - payload = SegmentUpdatePayload.model_validate(console_ns.payload or {}) - payload_dict = payload.model_dump(exclude_none=True) + payload_dict = req_data.model_dump(exclude_none=True) SegmentService.segment_create_args_validate(payload_dict, document) # Update segment (summary update with change detection is handled in SegmentService.update_segment) segment = SegmentService.update_segment( - SegmentUpdateArgs.model_validate(payload.model_dump(exclude_none=True)), + SegmentUpdateArgs.model_validate(req_data.model_dump(exclude_none=True)), segment, document, dataset, @@ -616,8 +619,10 @@ class DatasetDocumentSegmentBatchImportApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session + @model_validate(BatchImportPayload) def post( self, + req_data: BatchImportPayload, session: Session, current_tenant_id: str, current_user: Account, @@ -635,8 +640,7 @@ class DatasetDocumentSegmentBatchImportApi(Resource): if not document: raise NotFound("Document not found.") - payload = BatchImportPayload.model_validate(console_ns.payload or {}) - upload_file_id = payload.upload_file_id + upload_file_id = req_data.upload_file_id upload_file = session.scalar(select(UploadFile).where(UploadFile.id == upload_file_id).limit(1)) if not upload_file: @@ -697,8 +701,10 @@ class ChildChunkAddApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session + @model_validate(ChildChunkCreatePayload) def post( self, + req_data: ChildChunkCreatePayload, session: Session, current_tenant_id: str, current_user: Account, @@ -742,8 +748,7 @@ class ChildChunkAddApi(Resource): _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) # validate args try: - payload = ChildChunkCreatePayload.model_validate(console_ns.payload or {}) - child_chunk = SegmentService.create_child_chunk(payload.content, segment, document, dataset, session) + child_chunk = SegmentService.create_child_chunk(req_data.content, segment, document, dataset, session) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkDetailResponse, {"data": child_chunk}), 200 @@ -812,8 +817,10 @@ class ChildChunkAddApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session + @model_validate(ChildChunkBatchUpdatePayload) def patch( self, + req_data: ChildChunkBatchUpdatePayload, session: Session, current_tenant_id: str, current_user: Account, @@ -843,9 +850,8 @@ class ChildChunkAddApi(Resource): segment_id_str = str(segment_id) _, segment = _get_segment_for_document(session, dataset, document, segment_id_str) # validate args - payload = ChildChunkBatchUpdatePayload.model_validate(console_ns.payload or {}) try: - child_chunks = SegmentService.update_child_chunks(payload.chunks, segment, document, dataset, session) + child_chunks = SegmentService.update_child_chunks(req_data.chunks, segment, document, dataset, session) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) return dump_response(ChildChunkBatchUpdateResponse, {"data": child_chunks}), 200 @@ -918,8 +924,10 @@ class ChildChunkUpdateApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session + @model_validate(ChildChunkUpdatePayload) def patch( self, + req_data: ChildChunkUpdatePayload, session: Session, current_tenant_id: str, current_user: Account, @@ -955,9 +963,8 @@ class ChildChunkUpdateApi(Resource): raise NotFound("Child chunk not found.") # validate args try: - payload = ChildChunkUpdatePayload.model_validate(console_ns.payload or {}) child_chunk = SegmentService.update_child_chunk( - payload.content, child_chunk, segment, document, dataset, session + req_data.content, child_chunk, segment, document, dataset, session ) except ChildChunkIndexingServiceError as e: raise ChildChunkIndexingError(str(e)) diff --git a/api/controllers/console/datasets/external.py b/api/controllers/console/datasets/external.py index 94efe388561..26d958e15d3 100644 --- a/api/controllers/console/datasets/external.py +++ b/api/controllers/console/datasets/external.py @@ -1,9 +1,10 @@ +from __future__ import annotations + from dataclasses import dataclass from datetime import datetime from typing import Any from uuid import UUID -from flask import request from flask_restx import Resource from pydantic import AliasChoices, BaseModel, Field, field_validator from sqlalchemy.orm import Session @@ -24,6 +25,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -37,6 +39,7 @@ from models import Account from models.dataset import ExternalKnowledgeApis from services.dataset_service import DatasetService from services.enterprise import rbac_service as enterprise_rbac_service +from services.entities.external_knowledge_entities.external_knowledge_entities import ExternalDatasetCreatePayload from services.external_knowledge_service import ExternalDatasetService from services.hit_testing_service import HitTestingService from services.knowledge_service import BedrockRetrievalSetting, ExternalDatasetTestService @@ -47,14 +50,6 @@ class ExternalKnowledgeApiPayload(BaseModel): settings: dict[str, Any] -class ExternalDatasetCreatePayload(BaseModel): - external_knowledge_api_id: str - external_knowledge_id: str - name: str = Field(..., min_length=1, max_length=100) - description: str | None = Field(None, max_length=400) - external_retrieval_model: dict[str, Any] | None = None - - class ExternalHitTestingPayload(BaseModel): query: str external_retrieval_model: dict[str, Any] | None = None @@ -62,7 +57,7 @@ class ExternalHitTestingPayload(BaseModel): class BedrockRetrievalPayload(BaseModel): - retrieval_setting: "BedrockRetrievalSetting" + retrieval_setting: BedrockRetrievalSetting query: str knowledge_id: str @@ -106,7 +101,7 @@ class ExternalKnowledgeApiResponseSource: return self.external_knowledge_api.get_dataset_bindings(session=self.session) def __getattr__(self, name: str) -> Any: - return getattr(self.external_knowledge_api, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self.external_knowledge_api, name) # guard-ignore: no-new-getattr -- delegates model fields def external_knowledge_api_response( @@ -191,18 +186,18 @@ class ExternalApiTemplateListApi(Resource): @with_current_tenant_id @account_initialization_required @with_session(write=False) - def get(self, session: Session, current_tenant_id: str): - query = ExternalApiTemplateListQuery.model_validate(request.args.to_dict()) + @model_validate(ExternalApiTemplateListQuery) + def get(self, req_data: ExternalApiTemplateListQuery, session: Session, current_tenant_id: str): external_knowledge_apis, total = ExternalDatasetService.get_external_knowledge_apis( - query.page, query.limit, current_tenant_id, query.keyword, session=session + req_data.page, req_data.limit, current_tenant_id, req_data.keyword, session=session ) return ExternalKnowledgeApiListResponse( data=[external_knowledge_api_response(item, session=session) for item in external_knowledge_apis], - has_more=len(external_knowledge_apis) == query.limit, - limit=query.limit, + has_more=len(external_knowledge_apis) == req_data.limit, + limit=req_data.limit, total=total, - page=query.page, + page=req_data.page, ).model_dump(mode="json"), 200 @console_ns.doc("create_external_api_template") @@ -220,10 +215,16 @@ class ExternalApiTemplateListApi(Resource): @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, current_tenant_id: str, current_user: Account): - payload = ExternalKnowledgeApiPayload.model_validate(console_ns.payload or {}) + @model_validate(ExternalKnowledgeApiPayload) + def post( + self, + req_data: ExternalKnowledgeApiPayload, + session: Session, + current_tenant_id: str, + current_user: Account, + ): - ExternalDatasetService.validate_api_list(payload.settings) + ExternalDatasetService.validate_api_list(req_data.settings) # The role of the current user in the ta table must be admin, owner, or editor, or dataset_operator if not current_user.is_dataset_editor: @@ -233,7 +234,7 @@ class ExternalApiTemplateListApi(Resource): external_knowledge_api = ExternalDatasetService.create_external_knowledge_api( tenant_id=current_tenant_id, user_id=current_user.id, - args=payload.model_dump(), + args=req_data.model_dump(), session=session, ) except services.errors.dataset.DatasetNameDuplicateError: @@ -284,17 +285,24 @@ class ExternalApiTemplateApi(Resource): @with_current_user @with_current_tenant_id @with_session - def patch(self, session: Session, current_tenant_id: str, current_user: Account, external_knowledge_api_id: UUID): + @model_validate(ExternalKnowledgeApiPayload) + def patch( + self, + req_data: ExternalKnowledgeApiPayload, + session: Session, + current_tenant_id: str, + current_user: Account, + external_knowledge_api_id: UUID, + ): external_knowledge_api_id_str = str(external_knowledge_api_id) - payload = ExternalKnowledgeApiPayload.model_validate(console_ns.payload or {}) - ExternalDatasetService.validate_api_list(payload.settings) + ExternalDatasetService.validate_api_list(req_data.settings) external_knowledge_api = ExternalDatasetService.update_external_knowledge_api( tenant_id=current_tenant_id, user_id=current_user.id, external_knowledge_api_id=external_knowledge_api_id_str, - args=payload.model_dump(), + args=req_data.model_dump(), session=session, ) @@ -353,14 +361,21 @@ class ExternalDatasetCreateApi(Resource): @login_required @account_initialization_required @edit_permission_required - @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EXTERNAL_CONNECT) + @rbac_permission_required( + RBACResourceScope.DATASET, RBACPermission.DATASET_EXTERNAL_CONNECT, resource_required=False + ) @with_current_user @with_current_tenant_id @with_session - def post(self, session: Session, current_tenant_id: str, current_user: Account): + @model_validate(ExternalDatasetCreatePayload) + def post( + self, + req_data: ExternalDatasetCreatePayload, + session: Session, + current_tenant_id: str, + current_user: Account, + ): # The role of the current user in the ta table must be admin, owner, or editor - payload = ExternalDatasetCreatePayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) # The role of the current user in the ta table must be admin, owner, or editor, or dataset_operator if not current_user.is_dataset_editor: @@ -370,7 +385,7 @@ class ExternalDatasetCreateApi(Resource): dataset = ExternalDatasetService.create_external_dataset( tenant_id=current_tenant_id, user_id=current_user.id, - args=args, + args=req_data, session=session, ) except services.errors.dataset.DatasetNameDuplicateError: @@ -409,7 +424,8 @@ class ExternalKnowledgeHitTestingApi(Resource): @with_current_user @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_PIPELINE_TEST) @with_session - def post(self, session: Session, current_user: Account, dataset_id: UUID): + @model_validate(ExternalHitTestingPayload) + def post(self, req_data: ExternalHitTestingPayload, session: Session, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) dataset = DatasetService.get_dataset(dataset_id_str, session) if dataset is None: @@ -420,17 +436,16 @@ class ExternalKnowledgeHitTestingApi(Resource): except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) - payload = ExternalHitTestingPayload.model_validate(console_ns.payload or {}) - HitTestingService.hit_testing_args_check(payload.model_dump()) + HitTestingService.hit_testing_args_check(req_data.model_dump()) try: response = HitTestingService.external_retrieve( session=session, dataset=dataset, - query=payload.query, + query=req_data.query, account=current_user, - external_retrieval_model=payload.external_retrieval_model, - metadata_filtering_conditions=payload.metadata_filtering_conditions, + external_retrieval_model=req_data.external_retrieval_model, + metadata_filtering_conditions=req_data.metadata_filtering_conditions, ) return dump_response(ExternalHitTestingResponse, response) @@ -445,11 +460,11 @@ class BedrockRetrievalApi(Resource): @console_ns.doc(description="Bedrock retrieval test (internal use only)") @console_ns.expect(console_ns.models[BedrockRetrievalPayload.__name__]) @console_ns.response(200, "Bedrock retrieval test completed", console_ns.models[BedrockRetrievalResponse.__name__]) - def post(self): - payload = BedrockRetrievalPayload.model_validate(console_ns.payload or {}) + @model_validate(BedrockRetrievalPayload) + def post(self, req_data: BedrockRetrievalPayload): # Call the knowledge retrieval service result = ExternalDatasetTestService.knowledge_retrieval( - payload.retrieval_setting, payload.query, payload.knowledge_id + req_data.retrieval_setting, req_data.query, req_data.knowledge_id ) return dump_response(BedrockRetrievalResponse, result), 200 diff --git a/api/controllers/console/datasets/metadata.py b/api/controllers/console/datasets/metadata.py index 32c4151a017..c076d455f49 100644 --- a/api/controllers/console/datasets/metadata.py +++ b/api/controllers/console/datasets/metadata.py @@ -3,8 +3,10 @@ from uuid import UUID from flask_restx import Resource from sqlalchemy.orm import Session -from werkzeug.exceptions import NotFound +from werkzeug.exceptions import Forbidden, NotFound +import services +from configs import dify_config from controllers.common.controller_schemas import MetadataUpdatePayload from controllers.common.schema import register_response_schema_models, register_schema_models from controllers.common.session import with_session @@ -14,6 +16,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, enterprise_license_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -34,6 +37,7 @@ from services.entities.knowledge_entities.knowledge_entities import ( MetadataDetail, MetadataOperationData, ) +from services.errors.metadata import MetadataResourceNotFoundError from services.metadata_service import MetadataService register_schema_models( @@ -59,9 +63,15 @@ class DatasetMetadataCreateApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def post(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): - metadata_args = MetadataArgs.model_validate(console_ns.payload or {}) - + @model_validate(MetadataArgs) + def post( + self, + req_data: MetadataArgs, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + ): dataset_id_str = str(dataset_id) dataset = DatasetService.get_dataset(dataset_id_str, session) if dataset is None: @@ -69,7 +79,7 @@ class DatasetMetadataCreateApi(Resource): DatasetService.check_dataset_permission(dataset, current_user, session) metadata = MetadataService.create_metadata( - dataset_id_str, metadata_args, current_user, current_tenant_id, session=session + dataset_id_str, req_data, current_user, current_tenant_id, session=session ) return dump_response(DatasetMetadataResponse, metadata), 201 @@ -80,13 +90,20 @@ class DatasetMetadataCreateApi(Resource): @console_ns.response( 200, "Metadata retrieved successfully", console_ns.models[DatasetMetadataListResponse.__name__] ) + @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT) @with_session(write=False) - def get(self, session: Session, dataset_id: UUID): + def get(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, current_tenant_id, session=session) if dataset is None: raise NotFound("Dataset not found.") + if not dify_config.RBAC_ENABLED: + try: + DatasetService.check_dataset_permission(dataset, current_user, session) + except services.errors.account.NoPermissionError as e: + raise Forbidden(str(e)) metadata = MetadataService.get_dataset_metadatas(dataset, session) return dump_response(DatasetMetadataListResponse, metadata), 200 @@ -103,27 +120,26 @@ class DatasetMetadataApi(Resource): @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session + @model_validate(MetadataUpdatePayload) def patch( self, + req_data: MetadataUpdatePayload, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID, metadata_id: UUID, ): - payload = MetadataUpdatePayload.model_validate(console_ns.payload or {}) - name = payload.name + name = req_data.name dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, current_tenant_id, session=session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, session) - metadata = MetadataService.update_metadata_name( - dataset_id_str, metadata_id_str, name, current_user, current_tenant_id, session=session - ) + metadata = MetadataService.update_metadata_name(dataset, metadata_id_str, name, current_user, session=session) return dump_response(DatasetMetadataResponse, metadata), 200 @setup_required @@ -132,17 +148,25 @@ class DatasetMetadataApi(Resource): @enterprise_license_required @console_ns.response(204, "Metadata deleted successfully") @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def delete(self, session: Session, current_user: Account, dataset_id: UUID, metadata_id: UUID): + def delete( + self, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + metadata_id: UUID, + ): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, current_tenant_id, session=session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, session) - MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session) + MetadataService.delete_metadata(dataset, metadata_id_str, session) # Frontend callers only await success and invalidate metadata caches; no response body is consumed. return "", 204 @@ -200,19 +224,29 @@ class DocumentMetadataEditApi(Resource): 204, "Documents metadata updated successfully", ) + @console_ns.response(404, "Dataset, document, or metadata not found") @with_current_user + @with_current_tenant_id @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def post(self, session: Session, current_user: Account, dataset_id: UUID): - dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + @model_validate(MetadataOperationData) + def post( + self, + req_data: MetadataOperationData, + session: Session, + current_tenant_id: str, + current_user: Account, + dataset_id: UUID, + ): + dataset = DatasetService.get_dataset_for_tenant(str(dataset_id), current_tenant_id, session=session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, session) - metadata_args = MetadataOperationData.model_validate(console_ns.payload or {}) - - MetadataService.update_documents_metadata(dataset, metadata_args, current_user, session=session) + try: + MetadataService.update_documents_metadata(dataset, req_data, current_user, session=session) + except MetadataResourceNotFoundError as exc: + raise NotFound(str(exc)) from exc # Frontend callers only await success and invalidate caches; no response body is consumed. return "", 204 diff --git a/api/controllers/console/datasets/rag_pipeline/datasource_auth.py b/api/controllers/console/datasets/rag_pipeline/datasource_auth.py index 57d6b628d4b..c8bfc5bae86 100644 --- a/api/controllers/console/datasets/rag_pipeline/datasource_auth.py +++ b/api/controllers/console/datasets/rag_pipeline/datasource_auth.py @@ -14,6 +14,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -275,8 +276,8 @@ class DatasourceAuth(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.CREDENTIAL_CREATE, resource_required=False) @with_current_tenant_id - def post(self, current_tenant_id: str, provider_id: str): - payload = DatasourceCredentialPayload.model_validate(console_ns.payload or {}) + @model_validate(DatasourceCredentialPayload) + def post(self, req_data: DatasourceCredentialPayload, current_tenant_id: str, provider_id: str): datasource_provider_id = DatasourceProviderID(provider_id) datasource_provider_service = DatasourceProviderService() @@ -284,8 +285,8 @@ class DatasourceAuth(Resource): datasource_provider_service.add_datasource_api_key_provider( tenant_id=current_tenant_id, provider_id=datasource_provider_id, - credentials=payload.credentials, - name=payload.name, + credentials=req_data.credentials, + name=req_data.name, ) except CredentialsValidateFailedError as ex: raise ValueError(str(ex)) @@ -325,16 +326,16 @@ class DatasourceAuthDeleteApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @with_current_tenant_id - def post(self, current_tenant_id: str, provider_id: str): + @model_validate(DatasourceCredentialDeletePayload) + def post(self, req_data: DatasourceCredentialDeletePayload, current_tenant_id: str, provider_id: str): datasource_provider_id = DatasourceProviderID(provider_id) plugin_id = datasource_provider_id.plugin_id provider_name = datasource_provider_id.provider_name - payload = DatasourceCredentialDeletePayload.model_validate(console_ns.payload or {}) datasource_provider_service = DatasourceProviderService() datasource_provider_service.remove_datasource_credentials( tenant_id=current_tenant_id, - auth_id=payload.credential_id, + auth_id=req_data.credential_id, provider=provider_name, plugin_id=plugin_id, session=db.session(), @@ -354,18 +355,18 @@ class DatasourceAuthUpdateApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @with_current_tenant_id - def post(self, current_tenant_id: str, provider_id: str): + @model_validate(DatasourceCredentialUpdatePayload) + def post(self, req_data: DatasourceCredentialUpdatePayload, current_tenant_id: str, provider_id: str): datasource_provider_id = DatasourceProviderID(provider_id) - payload = DatasourceCredentialUpdatePayload.model_validate(console_ns.payload or {}) datasource_provider_service = DatasourceProviderService() datasource_provider_service.update_datasource_credentials( tenant_id=current_tenant_id, - auth_id=payload.credential_id, + auth_id=req_data.credential_id, provider=datasource_provider_id.provider_name, plugin_id=datasource_provider_id.plugin_id, - credentials=payload.credentials or {}, - name=payload.name, + credentials=req_data.credentials or {}, + name=req_data.name, ) return SimpleResultResponse(result="success").model_dump(mode="json"), 201 @@ -420,15 +421,15 @@ class DatasourceAuthOauthCustomClient(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @with_current_tenant_id - def post(self, current_tenant_id: str, provider_id: str): - payload = DatasourceCustomClientPayload.model_validate(console_ns.payload or {}) + @model_validate(DatasourceCustomClientPayload) + def post(self, req_data: DatasourceCustomClientPayload, current_tenant_id: str, provider_id: str): datasource_provider_id = DatasourceProviderID(provider_id) datasource_provider_service = DatasourceProviderService() datasource_provider_service.setup_oauth_custom_client_params( tenant_id=current_tenant_id, datasource_provider_id=datasource_provider_id, - client_params=payload.client_params or {}, - enabled=payload.enable_oauth_custom_client or False, + client_params=req_data.client_params or {}, + enabled=req_data.enable_oauth_custom_client or False, ) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -457,14 +458,14 @@ class DatasourceAuthDefaultApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @with_current_tenant_id - def post(self, current_tenant_id: str, provider_id: str): - payload = DatasourceDefaultPayload.model_validate(console_ns.payload or {}) + @model_validate(DatasourceDefaultPayload) + def post(self, req_data: DatasourceDefaultPayload, current_tenant_id: str, provider_id: str): datasource_provider_id = DatasourceProviderID(provider_id) datasource_provider_service = DatasourceProviderService() datasource_provider_service.set_default_datasource_provider( tenant_id=current_tenant_id, datasource_provider_id=datasource_provider_id, - credential_id=payload.id, + credential_id=req_data.id, ) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -479,14 +480,14 @@ class DatasourceUpdateProviderNameApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @with_current_tenant_id - def post(self, current_tenant_id: str, provider_id: str): - payload = DatasourceUpdateNamePayload.model_validate(console_ns.payload or {}) + @model_validate(DatasourceUpdateNamePayload) + def post(self, req_data: DatasourceUpdateNamePayload, current_tenant_id: str, provider_id: str): datasource_provider_id = DatasourceProviderID(provider_id) datasource_provider_service = DatasourceProviderService() datasource_provider_service.update_datasource_provider_name( tenant_id=current_tenant_id, datasource_provider_id=datasource_provider_id, - name=payload.name, - credential_id=payload.credential_id, + name=req_data.name, + credential_id=req_data.credential_id, ) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 diff --git a/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py b/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py index 873ba130064..07ae94d3924 100644 --- a/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py +++ b/api/controllers/console/datasets/rag_pipeline/datasource_content_preview.py @@ -8,7 +8,7 @@ from pydantic import BaseModel from controllers.common.schema import register_schema_models from controllers.console import console_ns from controllers.console.datasets.wraps import get_rag_pipeline -from controllers.console.wraps import account_initialization_required, setup_required, with_current_user +from controllers.console.wraps import account_initialization_required, model_validate, setup_required, with_current_user from extensions.ext_database import db from libs.login import login_required from models import Account @@ -34,14 +34,14 @@ class DataSourceContentPreviewApi(Resource): @account_initialization_required @get_rag_pipeline @with_current_user - def post(self, current_user: Account, pipeline: Pipeline, node_id: str): + @model_validate(Parser) + def post(self, req_data: Parser, current_user: Account, pipeline: Pipeline, node_id: str): """ Run datasource content preview """ - args = Parser.model_validate(console_ns.payload) - inputs = args.inputs - datasource_type = args.datasource_type + inputs = req_data.inputs + datasource_type = req_data.datasource_type rag_pipeline_service = RagPipelineService(db.session()) preview_content = rag_pipeline_service.run_datasource_node_preview( pipeline=pipeline, @@ -50,6 +50,6 @@ class DataSourceContentPreviewApi(Resource): account=current_user, datasource_type=datasource_type, is_published=True, - credential_id=args.credential_id, + credential_id=req_data.credential_id, ) return preview_content, 200 diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py index 2d824afb6ef..d2a80cadf26 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline.py @@ -1,13 +1,12 @@ import logging from typing import Any -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field -from sqlalchemy import select -from sqlalchemy.orm import Session, sessionmaker -from werkzeug.exceptions import NotFound +from sqlalchemy.orm import Session +from werkzeug.exceptions import Forbidden, NotFound +from configs import dify_config from controllers.common.fields import SimpleDataResponse from controllers.common.schema import ( JsonResponseWithStatus, @@ -17,10 +16,15 @@ from controllers.common.schema import ( ) from controllers.console import console_ns from controllers.console.app.wraps import with_session +from controllers.console.datasets.wraps import get_rag_pipeline from controllers.console.wraps import ( + RBACPermission, + RBACResourceScope, account_initialization_required, enterprise_license_required, knowledge_pipeline_publish_enabled, + model_validate, + rbac_permission_required, setup_required, with_current_tenant_id, with_current_user, @@ -30,8 +34,11 @@ from fields.base import ResponseModel from libs.helper import dump_response from libs.login import login_required from models.account import Account -from models.dataset import PipelineCustomizedTemplate +from models.dataset import Pipeline +from services.dataset_service import DatasetService from services.entities.knowledge_entities.rag_pipeline_entities import IconInfo, PipelineTemplateInfoEntity +from services.errors.account import NoPermissionError +from services.errors.rag_pipeline import RagPipelineResourceNotFoundError from services.rag_pipeline.rag_pipeline import RagPipelineService logger: logging.Logger = logging.getLogger(__name__) @@ -104,12 +111,17 @@ class PipelineTemplateListApi(Resource): @enterprise_license_required @with_current_tenant_id @with_session - def get(self, session: Session, current_tenant_id: str) -> JsonResponseWithStatus: - query = PipelineTemplateListQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(PipelineTemplateListQuery) + def get( + self, + req_data: PipelineTemplateListQuery, + session: Session, + current_tenant_id: str, + ) -> JsonResponseWithStatus: # get pipeline templates pipeline_templates = RagPipelineService.get_pipeline_templates( - type=query.type, - language=query.language, + type=req_data.type, + language=req_data.language, current_tenant_id=current_tenant_id, session=session, ) @@ -120,16 +132,25 @@ class PipelineTemplateListApi(Resource): class PipelineTemplateDetailApi(Resource): @console_ns.doc(params=query_params_from_model(PipelineTemplateDetailQuery)) @console_ns.response(200, "Pipeline template", console_ns.models[PipelineTemplateDetailResponse.__name__]) + @console_ns.response(404, "Pipeline template not found") @setup_required @login_required @account_initialization_required @enterprise_license_required - @with_session - def get(self, session: Session, template_id: str) -> JsonResponseWithStatus: - query = PipelineTemplateDetailQuery.model_validate(request.args.to_dict(flat=True)) + @with_current_tenant_id + @with_session(write=False) + @model_validate(PipelineTemplateDetailQuery) + def get( + self, + req_data: PipelineTemplateDetailQuery, + session: Session, + current_tenant_id: str, + template_id: str, + ) -> JsonResponseWithStatus: pipeline_template = RagPipelineService.get_pipeline_template_detail( template_id, - type=query.type, + current_tenant_id, + type=req_data.type, session=session, ) if pipeline_template is None: @@ -147,9 +168,15 @@ class CustomizedPipelineTemplateApi(Resource): @enterprise_license_required @with_current_user @with_current_tenant_id - def patch(self, current_tenant_id: str, current_user: Account, template_id: str) -> tuple[str, int]: - payload = CustomizedPipelineTemplatePayload.model_validate(console_ns.payload or {}) - pipeline_template_info = PipelineTemplateInfoEntity.model_validate(payload.model_dump()) + @model_validate(CustomizedPipelineTemplatePayload) + def patch( + self, + req_data: CustomizedPipelineTemplatePayload, + current_tenant_id: str, + current_user: Account, + template_id: str, + ) -> tuple[str, int]: + pipeline_template_info = PipelineTemplateInfoEntity.model_validate(req_data.model_dump()) RagPipelineService.update_customized_pipeline_template( template_id, pipeline_template_info, current_user, current_tenant_id, session=db.session() ) @@ -170,32 +197,57 @@ class CustomizedPipelineTemplateApi(Resource): @account_initialization_required @enterprise_license_required @console_ns.response(200, "Success", console_ns.models[SimpleDataResponse.__name__]) - def post(self, template_id: str) -> JsonResponseWithStatus: - with sessionmaker(db.engine, expire_on_commit=False).begin() as session: - template = session.scalar( - select(PipelineCustomizedTemplate).where(PipelineCustomizedTemplate.id == template_id).limit(1) + @console_ns.response(404, "Customized pipeline template not found") + @with_current_tenant_id + @with_session(write=False) + def post(self, session: Session, current_tenant_id: str, template_id: str) -> JsonResponseWithStatus: + try: + yaml_content = RagPipelineService.get_customized_pipeline_template_yaml( + template_id, current_tenant_id, session=session ) - if not template: - raise ValueError("Customized pipeline template not found.") + except RagPipelineResourceNotFoundError as exc: + raise NotFound(str(exc)) from exc - return dump_response(SimpleDataResponse, {"data": template.yaml_content}), 200 + return dump_response(SimpleDataResponse, {"data": yaml_content}), 200 @console_ns.route("/rag/pipelines//customized/publish") class PublishCustomizedPipelineTemplateApi(Resource): @console_ns.expect(console_ns.models[CustomizedPipelineTemplatePayload.__name__]) @console_ns.response(204, "Pipeline template published") + @console_ns.response(404, "Pipeline, workflow, or dataset not found") @setup_required @login_required @account_initialization_required @enterprise_license_required @knowledge_pipeline_publish_enabled @with_current_user - @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account, pipeline_id: str) -> tuple[str, int]: - payload = CustomizedPipelineTemplatePayload.model_validate(console_ns.payload or {}) - rag_pipeline_service = RagPipelineService(db.session()) - rag_pipeline_service.publish_customized_pipeline_template( - pipeline_id, payload.model_dump(), current_user, current_tenant_id, session=db.session() - ) + @get_rag_pipeline + @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_PIPELINE_RELEASE) + @model_validate(CustomizedPipelineTemplatePayload) + def post( + self, + req_data: CustomizedPipelineTemplatePayload, + current_user: Account, + pipeline: Pipeline, + ) -> tuple[str, int]: + session = db.session() + dataset = pipeline.retrieve_dataset(session=session) + if dataset is None: + raise NotFound("Dataset not found") + + if not dify_config.RBAC_ENABLED: + if not current_user.is_dataset_editor: + raise Forbidden() + try: + DatasetService.check_dataset_permission(dataset, current_user, session) + except NoPermissionError as exc: + raise Forbidden(str(exc)) from exc + + try: + RagPipelineService.publish_customized_pipeline_template( + pipeline, dataset, req_data.model_dump(), current_user, session=session + ) + except RagPipelineResourceNotFoundError as exc: + raise NotFound(str(exc)) from exc return "", 204 diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py index 5ad764871e4..847aa902010 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_datasets.py @@ -10,6 +10,7 @@ from controllers.console.datasets.rag_pipeline.rag_pipeline_import import RagPip from controllers.console.wraps import ( account_initialization_required, cloud_edition_billing_rate_limit_check, + model_validate, setup_required, with_current_tenant_id, with_current_user, @@ -47,8 +48,13 @@ class CreateRagPipelineDatasetApi(Resource): @cloud_edition_billing_rate_limit_check("knowledge") @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account) -> JsonResponseWithStatus: - payload = RagPipelineDatasetImportPayload.model_validate(console_ns.payload or {}) + @model_validate(RagPipelineDatasetImportPayload) + def post( + self, + req_data: RagPipelineDatasetImportPayload, + current_tenant_id: str, + current_user: Account, + ) -> JsonResponseWithStatus: # The role of the current user in the ta table must be admin, owner, or editor, or dataset_operator if not current_user.is_dataset_editor: raise Forbidden() @@ -62,7 +68,7 @@ class CreateRagPipelineDatasetApi(Resource): ), permission=DatasetPermissionEnum.ONLY_ME, partial_member_list=None, - yaml_content=payload.yaml_content, + yaml_content=req_data.yaml_content, ) try: rag_pipeline_dsl_service = RagPipelineDslService(db.session()) diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py index 25628a67177..b7d470ca13e 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_draft_variable.py @@ -4,7 +4,7 @@ from functools import wraps from typing import Any, Concatenate, NoReturn from uuid import UUID -from flask import Response, request +from flask import Response from flask_restx import Resource, marshal, marshal_with from pydantic import BaseModel, Field from sqlalchemy.orm import sessionmaker @@ -24,8 +24,9 @@ from controllers.console.app.workflow_draft_variable import ( workflow_draft_variable_model, ) from controllers.console.datasets.wraps import get_rag_pipeline -from controllers.console.wraps import account_initialization_required, setup_required, with_current_user +from controllers.console.wraps import account_initialization_required, model_validate, setup_required, with_current_user from core.app.file_access import DatabaseFileAccessController +from core.workflow.llm_environment_variable import LLMEnvironmentVariable, environment_variable_value_type from core.workflow.variable_prefixes import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID from extensions.ext_database import db from factories.file_factory import build_from_mapping, build_from_mappings @@ -91,11 +92,11 @@ class RagPipelineVariableCollectionApi(Resource): ) @_api_prerequisite @marshal_with(workflow_draft_variable_list_without_value_model) - def get(self, current_user: Account, pipeline: Pipeline): + @model_validate(PaginationQuery) + def get(self, req_data: PaginationQuery, current_user: Account, pipeline: Pipeline): """ Get draft workflow """ - query = PaginationQuery.model_validate(request.args.to_dict()) # fetch draft workflow by app_model rag_pipeline_service = RagPipelineService(db.session()) @@ -110,8 +111,8 @@ class RagPipelineVariableCollectionApi(Resource): ) workflow_vars = draft_var_srv.list_variables_without_values( app_id=pipeline.id, - page=query.page, - limit=query.limit, + page=req_data.page, + limit=req_data.limit, user_id=current_user.id, ) @@ -195,7 +196,14 @@ class RagPipelineVariableApi(Resource): @_api_prerequisite @marshal_with(workflow_draft_variable_model) @console_ns.expect(console_ns.models[WorkflowDraftVariablePatchPayload.__name__]) - def patch(self, _current_user: Account, pipeline: Pipeline, variable_id: UUID): + @model_validate(WorkflowDraftVariablePatchPayload) + def patch( + self, + req_data: WorkflowDraftVariablePatchPayload, + _current_user: Account, + pipeline: Pipeline, + variable_id: UUID, + ): # Request payload for file types: # # Local File: @@ -220,8 +228,7 @@ class RagPipelineVariableApi(Resource): draft_var_srv = WorkflowDraftVariableService( session=db.session(), ) - payload = WorkflowDraftVariablePatchPayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) variable_id_str = str(variable_id) variable = draft_var_srv.get_variable(variable_id=variable_id_str) @@ -362,7 +369,11 @@ class RagPipelineEnvironmentVariableCollectionApi(Resource): "name": v.name, "description": v.description, "selector": v.selector, - "value_type": v.value_type.value, + "value_type": ( + environment_variable_value_type(v) + if isinstance(v, LLMEnvironmentVariable) + else v.value_type.value + ), "value": v.value, # Do not track edited for env vars. "edited": False, diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_import.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_import.py index 0b83dc9312d..d55b5b6a366 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_import.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_import.py @@ -1,4 +1,3 @@ -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field from sqlalchemy.orm import Session @@ -17,6 +16,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_user, @@ -85,9 +85,9 @@ class RagPipelineImportApi(Resource): RBACResourceScope.DATASET, RBACPermission.DATASET_CREATE_AND_MANAGEMENT, resource_required=False ) @with_current_user - def post(self, current_user: Account) -> JsonResponseWithStatus: + @model_validate(RagPipelineImportPayload) + def post(self, req_data: RagPipelineImportPayload, current_user: Account) -> JsonResponseWithStatus: # Check user role first - payload = RagPipelineImportPayload.model_validate(console_ns.payload or {}) # Use a plain Session so that caught exceptions inside the service # (which return FAILED status instead of re-raising) do not leave the @@ -98,11 +98,11 @@ class RagPipelineImportApi(Resource): account = current_user result = import_service.import_rag_pipeline( account=account, - import_mode=payload.mode, - yaml_content=payload.yaml_content, - yaml_url=payload.yaml_url, - pipeline_id=payload.pipeline_id, - dataset_name=payload.name, + import_mode=req_data.mode, + yaml_content=req_data.yaml_content, + yaml_url=req_data.yaml_url, + pipeline_id=req_data.pipeline_id, + dataset_name=req_data.name, ) if result.status == ImportStatus.FAILED: session.rollback() @@ -179,14 +179,14 @@ class RagPipelineExportApi(Resource): @account_initialization_required @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_IMPORT_EXPORT_DSL) - def get(self, pipeline: Pipeline) -> JsonResponseWithStatus: + @model_validate(IncludeSecretQuery) + def get(self, req_data: IncludeSecretQuery, pipeline: Pipeline) -> JsonResponseWithStatus: # Add include_secret params - query = IncludeSecretQuery.model_validate(request.args.to_dict()) with Session(db.engine, expire_on_commit=False) as session: export_service = RagPipelineDslService(session) result = export_service.export_rag_pipeline_dsl( - pipeline=pipeline, include_secret=query.include_secret == "true" + pipeline=pipeline, include_secret=req_data.include_secret == "true" ) return dump_response(SimpleDataResponse, {"data": result}), 200 diff --git a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py index a264450078c..63a14083c85 100644 --- a/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/console/datasets/rag_pipeline/rag_pipeline_workflow.py @@ -10,6 +10,7 @@ from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import BadRequest, Forbidden, InternalServerError, NotFound import services +from configs import dify_config from controllers.common.controller_schemas import DefaultBlockConfigQuery, WorkflowListQuery, WorkflowUpdatePayload from controllers.common.fields import SimpleResultResponse from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models @@ -33,6 +34,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -59,8 +61,10 @@ from models import Account from models.dataset import Pipeline from models.model import EndUser from models.workflow import Workflow +from services.dataset_service import DatasetService from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError from services.errors.llm import InvokeRateLimitError +from services.errors.rag_pipeline import RagPipelineResourceNotFoundError from services.rag_pipeline.pipeline_generate_service import PipelineGenerateService from services.rag_pipeline.rag_pipeline import RagPipelineService from services.rag_pipeline.rag_pipeline_manage_service import RagPipelineManageService @@ -274,12 +278,12 @@ class RagPipelineDraftRunIterationNodeApi(Resource): @get_rag_pipeline @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_user: Account, pipeline: Pipeline, node_id: str): + @model_validate(NodeRunPayload) + def post(self, req_data: NodeRunPayload, current_user: Account, pipeline: Pipeline, node_id: str): """ Run draft workflow iteration node """ - payload = NodeRunPayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) try: response = PipelineGenerateService.generate_single_iteration( @@ -309,12 +313,12 @@ class RagPipelineDraftRunLoopNodeApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @get_rag_pipeline - def post(self, current_user: Account, pipeline: Pipeline, node_id: str): + @model_validate(NodeRunPayload) + def post(self, req_data: NodeRunPayload, current_user: Account, pipeline: Pipeline, node_id: str): """ Run draft workflow loop node """ - payload = NodeRunPayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) try: response = PipelineGenerateService.generate_single_loop( @@ -344,13 +348,13 @@ class DraftRagPipelineRunApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @with_session - def post(self, session: Session, current_user: Account, pipeline_id: UUID): + @model_validate(DraftWorkflowRunPayload) + def post(self, req_data: DraftWorkflowRunPayload, session: Session, current_user: Account, pipeline_id: UUID): """ Run draft workflow """ pipeline = load_rag_pipeline(session, str(pipeline_id)) - payload = DraftWorkflowRunPayload.model_validate(console_ns.payload or {}) - args = payload.model_dump() + args = req_data.model_dump() try: response = PipelineGenerateService.generate( @@ -378,14 +382,14 @@ class PublishedRagPipelineRunApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @with_session - def post(self, session: Session, current_user: Account, pipeline_id: UUID): + @model_validate(PublishedWorkflowRunPayload) + def post(self, req_data: PublishedWorkflowRunPayload, session: Session, current_user: Account, pipeline_id: UUID): """ Run published workflow """ pipeline = load_rag_pipeline(session, str(pipeline_id)) - payload = PublishedWorkflowRunPayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) - streaming = payload.response_mode == "streaming" + args = req_data.model_dump(exclude_none=True) + streaming = req_data.response_mode == "streaming" try: response = PipelineGenerateService.generate( @@ -393,7 +397,7 @@ class PublishedRagPipelineRunApi(Resource): pipeline=pipeline, user=current_user, args=args, - invoke_from=InvokeFrom.DEBUGGER if payload.is_preview else InvokeFrom.PUBLISHED_PIPELINE, + invoke_from=InvokeFrom.DEBUGGER if req_data.is_preview else InvokeFrom.PUBLISHED_PIPELINE, streaming=streaming, ) @@ -413,11 +417,11 @@ class RagPipelinePublishedDatasourceNodeRunApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @get_rag_pipeline - def post(self, current_user: Account, pipeline: Pipeline, node_id: str): + @model_validate(DatasourceNodeRunPayload) + def post(self, req_data: DatasourceNodeRunPayload, current_user: Account, pipeline: Pipeline, node_id: str): """ Run rag pipeline datasource """ - payload = DatasourceNodeRunPayload.model_validate(console_ns.payload or {}) rag_pipeline_service = RagPipelineService(db.session()) return helper.compact_generate_response( @@ -425,11 +429,11 @@ class RagPipelinePublishedDatasourceNodeRunApi(Resource): rag_pipeline_service.run_datasource_workflow_node( pipeline=pipeline, node_id=node_id, - user_inputs=payload.inputs, + user_inputs=req_data.inputs, account=current_user, - datasource_type=payload.datasource_type, + datasource_type=req_data.datasource_type, is_published=False, - credential_id=payload.credential_id, + credential_id=req_data.credential_id, ) ) ) @@ -446,11 +450,11 @@ class RagPipelineDraftDatasourceNodeRunApi(Resource): @account_initialization_required @with_current_user @get_rag_pipeline - def post(self, current_user: Account, pipeline: Pipeline, node_id: str): + @model_validate(DatasourceNodeRunPayload) + def post(self, req_data: DatasourceNodeRunPayload, current_user: Account, pipeline: Pipeline, node_id: str): """ Run rag pipeline datasource """ - payload = DatasourceNodeRunPayload.model_validate(console_ns.payload or {}) rag_pipeline_service = RagPipelineService(db.session()) return helper.compact_generate_response( @@ -458,11 +462,11 @@ class RagPipelineDraftDatasourceNodeRunApi(Resource): rag_pipeline_service.run_datasource_workflow_node( pipeline=pipeline, node_id=node_id, - user_inputs=payload.inputs, + user_inputs=req_data.inputs, account=current_user, - datasource_type=payload.datasource_type, + datasource_type=req_data.datasource_type, is_published=False, - credential_id=payload.credential_id, + credential_id=req_data.credential_id, ) ) ) @@ -483,12 +487,12 @@ class RagPipelineDraftNodeRunApi(Resource): @account_initialization_required @with_current_user @get_rag_pipeline - def post(self, current_user: Account, pipeline: Pipeline, node_id: str): + @model_validate(NodeRunRequiredPayload) + def post(self, req_data: NodeRunRequiredPayload, current_user: Account, pipeline: Pipeline, node_id: str): """ Run draft workflow node """ - payload = NodeRunRequiredPayload.model_validate(console_ns.payload or {}) - inputs = payload.inputs + inputs = req_data.inputs rag_pipeline_service = RagPipelineService(db.session()) workflow_node_execution = rag_pipeline_service.run_draft_workflow_node( @@ -617,16 +621,16 @@ class DefaultRagPipelineBlockConfigApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @get_rag_pipeline - def get(self, pipeline: Pipeline, block_type: str): + @model_validate(DefaultBlockConfigQuery) + def get(self, req_data: DefaultBlockConfigQuery, pipeline: Pipeline, block_type: str): """ Get default block config """ - query = DefaultBlockConfigQuery.model_validate(request.args.to_dict()) filters = None - if query.q: + if req_data.q: try: - filters = json.loads(query.q) + filters = json.loads(req_data.q) except json.JSONDecodeError: raise ValueError("Invalid filters") @@ -651,16 +655,16 @@ class PublishedAllRagPipelineApi(Resource): @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_current_user @get_rag_pipeline - def get(self, current_user: Account, pipeline: Pipeline): + @model_validate(WorkflowListQuery) + def get(self, req_data: WorkflowListQuery, current_user: Account, pipeline: Pipeline): """ Get published workflows """ - query = WorkflowListQuery.model_validate(request.args.to_dict()) - page = query.page - limit = query.limit - user_id = query.user_id - named_only = query.named_only + page = req_data.page + limit = req_data.limit + user_id = req_data.user_id + named_only = req_data.named_only if user_id: if user_id != current_user.id: @@ -733,12 +737,12 @@ class RagPipelineByIdApi(Resource): @with_current_user @get_rag_pipeline @console_ns.expect(console_ns.models[WorkflowUpdatePayload.__name__]) - def patch(self, current_user: Account, pipeline: Pipeline, workflow_id: str): + @model_validate(WorkflowUpdatePayload) + def patch(self, req_data: WorkflowUpdatePayload, current_user: Account, pipeline: Pipeline, workflow_id: str): """ Update workflow attributes """ - payload = WorkflowUpdatePayload.model_validate(console_ns.payload or {}) - update_data = payload.model_dump(exclude_unset=True) + update_data = req_data.model_dump(exclude_unset=True) if not update_data: return {"message": "No valid fields to update"}, 400 @@ -803,12 +807,12 @@ class PublishedRagPipelineSecondStepApi(Resource): @get_rag_pipeline @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def get(self, pipeline: Pipeline): + @model_validate(NodeIdQuery) + def get(self, req_data: NodeIdQuery, pipeline: Pipeline): """ Get second step parameters of rag pipeline """ - query = NodeIdQuery.model_validate(request.args.to_dict()) - node_id = query.node_id + node_id = req_data.node_id rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_second_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=False) return { @@ -826,12 +830,12 @@ class PublishedRagPipelineFirstStepApi(Resource): @get_rag_pipeline @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def get(self, pipeline: Pipeline): + @model_validate(NodeIdQuery) + def get(self, req_data: NodeIdQuery, pipeline: Pipeline): """ Get first step parameters of rag pipeline """ - query = NodeIdQuery.model_validate(request.args.to_dict()) - node_id = query.node_id + node_id = req_data.node_id rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_first_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=False) return { @@ -849,12 +853,12 @@ class DraftRagPipelineFirstStepApi(Resource): @get_rag_pipeline @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def get(self, pipeline: Pipeline): + @model_validate(NodeIdQuery) + def get(self, req_data: NodeIdQuery, pipeline: Pipeline): """ Get first step parameters of rag pipeline """ - query = NodeIdQuery.model_validate(request.args.to_dict()) - node_id = query.node_id + node_id = req_data.node_id rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_first_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=True) return { @@ -872,12 +876,12 @@ class DraftRagPipelineSecondStepApi(Resource): @get_rag_pipeline @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def get(self, pipeline: Pipeline): + @model_validate(NodeIdQuery) + def get(self, req_data: NodeIdQuery, pipeline: Pipeline): """ Get second step parameters of rag pipeline """ - query = NodeIdQuery.model_validate(request.args.to_dict()) - node_id = query.node_id + node_id = req_data.node_id rag_pipeline_service = RagPipelineService(db.session()) variables = rag_pipeline_service.get_second_step_parameters(pipeline=pipeline, node_id=node_id, is_draft=True) @@ -1015,19 +1019,31 @@ class RagPipelineWorkflowLastRunApi(Resource): @console_ns.route("/rag/pipelines/transform/datasets/") class RagPipelineTransformApi(Resource): @console_ns.response(200, "Success", console_ns.models[RagPipelineOpaqueResponse.__name__]) + @console_ns.response(404, "Dataset or pipeline not found") @setup_required @login_required @account_initialization_required @with_current_user + @with_current_tenant_id + @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) @with_session - def post(self, session: Session, current_user: Account, dataset_id: UUID): - if not (current_user.has_edit_permission or current_user.is_dataset_operator): - raise Forbidden() + def post(self, session: Session, current_tenant_id: str, current_user: Account, dataset_id: UUID): + dataset = DatasetService.get_dataset_for_tenant(str(dataset_id), current_tenant_id, session=session) + if dataset is None: + raise NotFound("Dataset not found.") - dataset_id_str = str(dataset_id) - rag_pipeline_transform_service = RagPipelineTransformService() - result = rag_pipeline_transform_service.transform_dataset(dataset_id_str, session) - return result + if not dify_config.RBAC_ENABLED: + if not (current_user.has_edit_permission or current_user.is_dataset_operator): + raise Forbidden() + try: + DatasetService.check_dataset_permission(dataset, current_user, session) + except services.errors.account.NoPermissionError as exc: + raise Forbidden(str(exc)) from exc + + try: + return RagPipelineTransformService().transform_dataset(dataset, current_user.id, session) + except RagPipelineResourceNotFoundError as exc: + raise NotFound(str(exc)) from exc @console_ns.route("/rag/pipelines//workflows/draft/datasource/variables-inspect") @@ -1045,11 +1061,12 @@ class RagPipelineDatasourceVariableApi(Resource): @get_rag_pipeline @edit_permission_required @rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT) - def post(self, current_user: Account, pipeline: Pipeline): + @model_validate(DatasourceVariablesPayload) + def post(self, req_data: DatasourceVariablesPayload, current_user: Account, pipeline: Pipeline): """ Set datasource variables """ - args = DatasourceVariablesPayload.model_validate(console_ns.payload or {}).model_dump() + args = req_data.model_dump() rag_pipeline_service = RagPipelineService(db.session()) workflow_node_execution = rag_pipeline_service.set_datasource_variables( @@ -1071,9 +1088,11 @@ class RagPipelineRecommendedPluginApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account): - query = RagPipelineRecommendedPluginQuery.model_validate(request.args.to_dict()) + @model_validate(RagPipelineRecommendedPluginQuery) + def get(self, req_data: RagPipelineRecommendedPluginQuery, current_tenant_id: str, current_user: Account): rag_pipeline_service = RagPipelineService(db.session()) - recommended_plugins = rag_pipeline_service.get_recommended_plugins(query.type, current_user, current_tenant_id) + recommended_plugins = rag_pipeline_service.get_recommended_plugins( + req_data.type, current_user, current_tenant_id + ) return recommended_plugins diff --git a/api/controllers/console/datasets/website.py b/api/controllers/console/datasets/website.py index 39be3f3ce5f..9411d128195 100644 --- a/api/controllers/console/datasets/website.py +++ b/api/controllers/console/datasets/website.py @@ -1,13 +1,12 @@ from typing import Any, Literal -from flask import request from flask_restx import Resource from pydantic import BaseModel, RootModel from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models from controllers.console import console_ns from controllers.console.datasets.error import WebsiteCrawlError -from controllers.console.wraps import account_initialization_required, setup_required +from controllers.console.wraps import account_initialization_required, model_validate, setup_required from libs.login import login_required from services.website_service import WebsiteCrawlApiRequest, WebsiteCrawlStatusApiRequest, WebsiteService @@ -40,12 +39,11 @@ class WebsiteCrawlApi(Resource): @setup_required @login_required @account_initialization_required - def post(self): - payload = WebsiteCrawlPayload.model_validate(console_ns.payload or {}) - + @model_validate(WebsiteCrawlPayload) + def post(self, req_data: WebsiteCrawlPayload): # Create typed request and validate try: - api_request = WebsiteCrawlApiRequest.from_args(payload.model_dump()) + api_request = WebsiteCrawlApiRequest.from_args(req_data.model_dump()) except ValueError as e: raise WebsiteCrawlError(str(e)) @@ -69,12 +67,11 @@ class WebsiteCrawlStatusApi(Resource): @setup_required @login_required @account_initialization_required - def get(self, job_id: str): - args = WebsiteCrawlStatusQuery.model_validate(request.args.to_dict()) - + @model_validate(WebsiteCrawlStatusQuery) + def get(self, req_data: WebsiteCrawlStatusQuery, job_id: str): # Create typed request and validate try: - api_request = WebsiteCrawlStatusApiRequest.from_args(args.model_dump(), job_id) + api_request = WebsiteCrawlStatusApiRequest.from_args(req_data.model_dump(), job_id) except ValueError as e: raise WebsiteCrawlError(str(e)) diff --git a/api/controllers/console/error.py b/api/controllers/console/error.py index e4352f92f88..638af52e284 100644 --- a/api/controllers/console/error.py +++ b/api/controllers/console/error.py @@ -109,6 +109,12 @@ class EducationActivateLimitError(BaseHTTPException): code = 429 +class EducationDiscountTemporarilyPausedError(BaseHTTPException): + error_code = "education_discount_temporarily_paused" + description = "Education discount temporarily paused, while we upgrade our security measures." + code = 503 + + class ComplianceRateLimitError(BaseHTTPException): error_code = "compliance_rate_limit" description = "Rate limit exceeded for downloading compliance report." diff --git a/api/controllers/console/explore/audio.py b/api/controllers/console/explore/audio.py index 6219571b2d6..e50fe8b16d3 100644 --- a/api/controllers/console/explore/audio.py +++ b/api/controllers/console/explore/audio.py @@ -20,6 +20,7 @@ from controllers.console.app.error import ( UnsupportedAudioTypeError, ) from controllers.console.explore.wraps import InstalledAppResource +from controllers.console.wraps import model_validate from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from extensions.ext_database import db from graphon.model_runtime.errors.invoke import InvokeError @@ -100,16 +101,15 @@ class ChatAudioApi(InstalledAppResource): class ChatTextApi(InstalledAppResource): @console_ns.expect(console_ns.models[TextToAudioPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[AudioBinaryResponse.__name__]) - def post(self, installed_app: InstalledApp): + @model_validate(TextToAudioPayload) + def post(self, req_data: TextToAudioPayload, installed_app: InstalledApp): app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() try: - payload = TextToAudioPayload.model_validate(console_ns.payload or {}) - - message_id = payload.message_id - text = payload.text - voice = payload.voice + message_id = req_data.message_id + text = req_data.text + voice = req_data.voice message_ref = None if message_id: current_user, _ = current_account_with_tenant() diff --git a/api/controllers/console/explore/banner.py b/api/controllers/console/explore/banner.py index a6321bb797c..9b1e3a106a5 100644 --- a/api/controllers/console/explore/banner.py +++ b/api/controllers/console/explore/banner.py @@ -1,37 +1,59 @@ -from typing import Any, cast +from datetime import datetime +from typing import cast -from flask import request from flask_restx import Namespace, Resource -from pydantic import BaseModel, Field, RootModel -from sqlalchemy import select +from pydantic import BaseModel, Field, RootModel, field_validator from controllers.common.schema import query_params_from_model, register_response_schema_models from controllers.console import api -from controllers.console.explore.wraps import explore_banner_enabled -from extensions.ext_database import db +from controllers.console.wraps import model_validate +from extensions.ext_application_services import application_services from fields.base import ResponseModel +from libs.helper import dump_response from models.enums import BannerStatus -from models.model import ExporleBanner class BannerListQuery(BaseModel): language: str = Field(default="en-US", description="Banner language") +class BannerContentResponse(ResponseModel): + category: str + title: str = Field(min_length=1) + description: str + image_source: str = Field( + min_length=1, + validation_alias="img-src", + serialization_alias="img-src", + ) + + class BannerResponse(ResponseModel): id: str - content: Any - link: str | None = None + content: BannerContentResponse + link: str sort: int - status: str - created_at: str | None = None + status: BannerStatus + created_at: str + + @field_validator("created_at", mode="before") + @classmethod + def serialize_created_at(cls, value: datetime | str) -> str: + if isinstance(value, datetime): + return value.isoformat() + return value class BannerListResponse(RootModel[list[BannerResponse]]): root: list[BannerResponse] -register_response_schema_models(cast(Namespace, api), BannerListResponse) +register_response_schema_models( + cast(Namespace, api), + BannerContentResponse, + BannerResponse, + BannerListResponse, +) class BannerApi(Resource): @@ -39,38 +61,11 @@ class BannerApi(Resource): @api.doc(params=query_params_from_model(BannerListQuery)) @api.response(200, "Success", api.models[BannerListResponse.__name__]) - @explore_banner_enabled - def get(self): + @model_validate(BannerListQuery) + def get(self, req_data: BannerListQuery): """Get banner list.""" - language = request.args.get("language", "en-US") - - # Build base query for enabled banners - base_query = select(ExporleBanner).where(ExporleBanner.status == BannerStatus.ENABLED) - - # Try to get banners in the requested language - banners = db.session.scalars( - base_query.where(ExporleBanner.language == language).order_by(ExporleBanner.sort) - ).all() - - # Fallback to en-US if no banners found and language is not en-US - if not banners and language != "en-US": - banners = db.session.scalars( - base_query.where(ExporleBanner.language == "en-US").order_by(ExporleBanner.sort) - ).all() - # Convert banners to serializable format - result = [] - for banner in banners: - banner_data = { - "id": banner.id, - "content": banner.content, # Already parsed as JSON by SQLAlchemy - "link": banner.link, - "sort": banner.sort, - "status": banner.status, - "created_at": banner.created_at.isoformat() if banner.created_at else None, - } - result.append(banner_data) - - return result + banners = application_services().explore_banner_queries.list_for_language(req_data.language) + return dump_response(BannerListResponse, banners) api.add_resource(BannerApi, "/explore/banners") diff --git a/api/controllers/console/explore/completion.py b/api/controllers/console/explore/completion.py index 5f034fabb81..cad51cd0454 100644 --- a/api/controllers/console/explore/completion.py +++ b/api/controllers/console/explore/completion.py @@ -20,7 +20,7 @@ from controllers.console.app.error import ( from controllers.console.app.wraps import with_session from controllers.console.explore.error import NotChatAppError, NotCompletionAppError from controllers.console.explore.wraps import InstalledAppResource -from controllers.console.wraps import with_current_user, with_current_user_id +from controllers.console.wraps import model_validate, with_current_user, with_current_user_id from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError from core.app.entities.app_invoke_entities import InvokeFrom from core.errors.error import ( @@ -89,17 +89,23 @@ class CompletionApi(InstalledAppResource): @console_ns.response(200, "Success") @with_current_user @with_session - def post(self, session: Session, current_user: Account, installed_app: InstalledApp): + @model_validate(CompletionMessageExplorePayload) + def post( + self, + req_data: CompletionMessageExplorePayload, + session: Session, + current_user: Account, + installed_app: InstalledApp, + ): app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() if app_model.mode != AppMode.COMPLETION: raise NotCompletionAppError() - payload = CompletionMessageExplorePayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) - streaming = payload.response_mode == "streaming" + streaming = req_data.response_mode == "streaming" args["auto_generate_name"] = False installed_app.last_used_at = naive_utc_now() @@ -173,7 +179,8 @@ class ChatApi(InstalledAppResource): @console_ns.response(200, "Success") @with_current_user @with_session - def post(self, session: Session, current_user: Account, installed_app: InstalledApp): + @model_validate(ChatMessagePayload) + def post(self, req_data: ChatMessagePayload, session: Session, current_user: Account, installed_app: InstalledApp): app_model = installed_app.app_with_session(session=session) if app_model is None: raise AppUnavailableError() @@ -181,8 +188,7 @@ class ChatApi(InstalledAppResource): if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT}: raise NotChatAppError() - payload = ChatMessagePayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) args["auto_generate_name"] = False @@ -191,10 +197,10 @@ class ChatApi(InstalledAppResource): try: # Eagerly validate conversation to avoid hanging on invalid conversation_id - if payload.conversation_id: + if req_data.conversation_id: ConversationService.get_conversation( app_model=app_model, - conversation_id=payload.conversation_id, + conversation_id=req_data.conversation_id, user=current_user, session=session, ) diff --git a/api/controllers/console/explore/conversation.py b/api/controllers/console/explore/conversation.py index 9e21fd496af..90b0d97e488 100644 --- a/api/controllers/console/explore/conversation.py +++ b/api/controllers/console/explore/conversation.py @@ -11,7 +11,7 @@ from controllers.common.schema import query_params_from_model, register_response from controllers.console.app.error import AppUnavailableError from controllers.console.explore.error import NotChatAppError from controllers.console.explore.wraps import InstalledAppResource -from controllers.console.wraps import with_current_user +from controllers.console.wraps import model_validate, with_current_user from core.app.entities.app_invoke_entities import InvokeFrom from extensions.ext_database import db from fields.conversation_fields import ( @@ -133,7 +133,8 @@ class ConversationRenameApi(InstalledAppResource): @console_ns.expect(console_ns.models[ConversationRenamePayload.__name__]) @console_ns.response(200, "Conversation renamed successfully", console_ns.models[SimpleConversation.__name__]) @with_current_user - def post(self, current_user: Account, installed_app: InstalledApp, c_id: UUID): + @model_validate(ConversationRenamePayload) + def post(self, req_data: ConversationRenamePayload, current_user: Account, installed_app: InstalledApp, c_id: UUID): app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() @@ -143,12 +144,10 @@ class ConversationRenameApi(InstalledAppResource): conversation_id = str(c_id) - payload = ConversationRenamePayload.model_validate(console_ns.payload or {}) - try: session = db.session() conversation = ConversationService.rename( - app_model, conversation_id, current_user, payload.name, payload.auto_generate, session=session + app_model, conversation_id, current_user, req_data.name, req_data.auto_generate, session=session ) return ( TypeAdapter(SimpleConversation) diff --git a/api/controllers/console/explore/installed_app.py b/api/controllers/console/explore/installed_app.py index 1fe1201bab7..e1b8295a983 100644 --- a/api/controllers/console/explore/installed_app.py +++ b/api/controllers/console/explore/installed_app.py @@ -1,11 +1,11 @@ +import base64 +import binascii import logging -from datetime import datetime -from typing import Any from flask import request from flask_restx import Resource -from pydantic import BaseModel, Field, computed_field, field_validator -from sqlalchemy import and_, exists, or_, select +from pydantic import BaseModel, Field, computed_field +from sqlalchemy import and_, select from werkzeug.exceptions import BadRequest, Forbidden, NotFound from controllers.common.fields import SimpleMessageResponse, SimpleResultMessageResponse @@ -15,6 +15,7 @@ from controllers.console.explore.wraps import InstalledAppResource from controllers.console.wraps import ( account_initialization_required, cloud_edition_billing_resource_check, + model_validate, with_current_tenant_id, with_current_user, ) @@ -22,13 +23,12 @@ from extensions.ext_database import db from fields.base import ResponseModel from graphon.file import helpers as file_helpers from libs.datetime_utils import naive_utc_now -from libs.helper import to_timestamp +from libs.helper import dump_response, to_timestamp from libs.login import login_required -from models import Account, App, AppModelConfig, InstalledApp, RecommendedApp, Workflow +from models import Account, App, InstalledApp, RecommendedApp from models.model import AppMode, IconType from services.account_service import TenantService -from services.enterprise.enterprise_service import EnterpriseService -from services.feature_service import FeatureService +from services.installed_app_service import InstalledAppCursor, InstalledAppService class InstalledAppCreatePayload(BaseModel): @@ -41,65 +41,53 @@ class InstalledAppUpdatePayload(BaseModel): class InstalledAppsListQuery(BaseModel): app_id: str | None = Field(default=None, description="App ID to filter by") + name: str | None = Field(default=None, max_length=100, description="App name to search for") + cursor: str | None = Field(default=None, description="Opaque cursor returned by the previous page") + limit: int = Field( + default=20, + ge=1, + le=100, + description="Number of installed apps to return", + ) logger = logging.getLogger(__name__) -def _build_icon_url(icon_type: str | IconType | None, icon: str | None) -> str | None: +def _build_icon_url(icon_type: IconType | None, icon: str | None) -> str | None: if icon is None or icon_type is None: return None - icon_type_value = icon_type.value if isinstance(icon_type, IconType) else str(icon_type) - if icon_type_value.lower() != IconType.IMAGE: + if icon_type != IconType.IMAGE: return None return file_helpers.get_signed_file_url(icon) -def _safe_primitive(value: Any) -> Any: - if value is None or isinstance(value, (str, int, float, bool, datetime)): - return value - return None +def _encode_installed_app_cursor(cursor: InstalledAppCursor) -> str: + payload = cursor.model_dump_json().encode() + return base64.urlsafe_b64encode(payload).decode().rstrip("=") -def _published_app_filter(): - """Return the SQL predicate for installed-app web API availability. +def _decode_installed_app_cursor(cursor: str | None) -> InstalledAppCursor | None: + if cursor is None: + return None - The installed-app parameters endpoint reads the published workflow for - workflow-style apps and the published app model config for easy UI apps. - Keep the list endpoint aligned in SQL so it does not return entries that - will immediately fail with app_unavailable when opened. - """ - workflow_app_modes = (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW) - has_published_workflow = exists(select(Workflow.id).where(Workflow.id == App.workflow_id)) - has_published_model_config = exists(select(AppModelConfig.id).where(AppModelConfig.id == App.app_model_config_id)) - - return and_( - App.mode != AppMode.AGENT, - or_( - and_(App.mode.in_(workflow_app_modes), App.workflow_id.isnot(None), has_published_workflow), - and_(~App.mode.in_(workflow_app_modes), App.app_model_config_id.isnot(None), has_published_model_config), - ), - ) + try: + padded_cursor = cursor + "=" * (-len(cursor) % 4) + payload = base64.b64decode(padded_cursor, altchars=b"-_", validate=True) + return InstalledAppCursor.model_validate_json(payload) + except (binascii.Error, UnicodeDecodeError, ValueError): + raise BadRequest("Invalid cursor") from None class InstalledAppInfoResponse(ResponseModel): id: str - name: str | None = None - description: str | None = None - mode: str | None = None - icon_type: str | None = None - icon: str | None = None - icon_background: str | None = None - use_icon_as_answer_icon: bool | None = None - - @field_validator("mode", "icon_type", mode="before") - @classmethod - def _normalize_enum_like(cls, value: Any) -> str | None: - if value is None: - return None - if isinstance(value, str): - return value - return str(getattr(value, "value", value)) + name: str + description: str + mode: AppMode + icon_type: IconType | None + icon: str | None + icon_background: str | None + use_icon_as_answer_icon: bool @computed_field(return_type=str | None) # type: ignore[prop-decorator] @property @@ -112,34 +100,33 @@ class InstalledAppResponse(ResponseModel): app: InstalledAppInfoResponse app_owner_tenant_id: str is_pinned: bool - last_used_at: int | None = None + last_used_at: int | None editable: bool uninstallable: bool - @field_validator("app", mode="before") - @classmethod - def _normalize_app(cls, value: Any) -> Any: - if isinstance(value, dict): - return value - return { - "id": _safe_primitive(getattr(value, "id", "")) or "", - "name": _safe_primitive(getattr(value, "name", None)), - "description": _safe_primitive(getattr(value, "description", None)), - "mode": _safe_primitive(getattr(value, "mode", None)), - "icon_type": _safe_primitive(getattr(value, "icon_type", None)), - "icon": _safe_primitive(getattr(value, "icon", None)), - "icon_background": _safe_primitive(getattr(value, "icon_background", None)), - "use_icon_as_answer_icon": _safe_primitive(getattr(value, "use_icon_as_answer_icon", None)), - } - - @field_validator("last_used_at", mode="before") - @classmethod - def _normalize_timestamp(cls, value: datetime | int | None) -> int | None: - return to_timestamp(value) - class InstalledAppListResponse(ResponseModel): installed_apps: list[InstalledAppResponse] + has_more: bool + next_cursor: str | None + + +def _installed_app_response_data( + installed_app: InstalledApp, + app_model: App, + *, + current_tenant_id: str, + current_user: Account, +) -> InstalledAppResponse: + return InstalledAppResponse( + id=installed_app.id, + app=InstalledAppInfoResponse.model_validate(app_model), + app_owner_tenant_id=installed_app.app_owner_tenant_id, + is_pinned=installed_app.is_pinned, + last_used_at=to_timestamp(installed_app.last_used_at), + editable=current_user.role in {"owner", "admin"}, + uninstallable=current_tenant_id == installed_app.app_owner_tenant_id, + ) register_schema_models( @@ -168,78 +155,40 @@ class InstalledAppsListApi(Resource): @with_current_tenant_id def get(self, current_tenant_id: str, current_user: Account): query = InstalledAppsListQuery.model_validate(request.args.to_dict()) - - stmt = ( - select(InstalledApp, App) - .join(App, App.id == InstalledApp.app_id) - .where(InstalledApp.tenant_id == current_tenant_id, _published_app_filter()) - ) - if query.app_id: - stmt = stmt.where(InstalledApp.app_id == query.app_id) - - installed_apps = db.session.execute(stmt).all() - + cursor = _decode_installed_app_cursor(query.cursor) if current_user.current_tenant is None: raise ValueError("current_user.current_tenant must not be None") - current_user.role = TenantService.get_user_role(current_user, current_user.current_tenant, session=db.session()) - installed_app_list: list[dict[str, Any]] = [] - for installed_app, app_model in installed_apps: - installed_app_list.append( - { - "id": installed_app.id, - "app": app_model, - "app_owner_tenant_id": installed_app.app_owner_tenant_id, - "is_pinned": installed_app.is_pinned, - "last_used_at": installed_app.last_used_at, - "editable": current_user.role in {"owner", "admin"}, - "uninstallable": current_tenant_id == installed_app.app_owner_tenant_id, - } - ) - # filter out apps that user doesn't have access to - if FeatureService.get_system_features().webapp_auth.enabled: - user_id = current_user.id - app_ids = [installed_app["app"].id for installed_app in installed_app_list] - webapp_settings = EnterpriseService.WebAppAuth.batch_get_app_access_mode_by_id(app_ids) - - # Pre-filter out apps without setting or with sso_verified - filtered_installed_apps = [] - - for installed_app in installed_app_list: - app_id = installed_app["app"].id - webapp_setting = webapp_settings.get(app_id) - if not webapp_setting or webapp_setting.access_mode == "sso_verified": - continue - filtered_installed_apps.append(installed_app) - - # Batch permission check - app_ids = [installed_app["app"].id for installed_app in filtered_installed_apps] - permissions = EnterpriseService.WebAppAuth.batch_is_user_allowed_to_access_webapps( - user_id=user_id, - app_ids=app_ids, - ) - - # Keep only allowed apps - res = [] - for installed_app in filtered_installed_apps: - app_id = installed_app["app"].id - if permissions.get(app_id): - res.append(installed_app) - - installed_app_list = res - logger.debug("installed_app_list: %s, user_id: %s", installed_app_list, user_id) - - installed_app_list.sort( - key=lambda app: ( - -app["is_pinned"], - app["last_used_at"] is None, - -app["last_used_at"].timestamp() if app["last_used_at"] is not None else 0, - ) + installed_apps, has_more, next_cursor = InstalledAppService.get_visible_page( + tenant_id=current_tenant_id, + user_id=str(current_user.id), + cursor=cursor, + limit=query.limit, + app_id=query.app_id, + name=query.name, + session=db.session, ) - return InstalledAppListResponse.model_validate( - {"installed_apps": installed_app_list}, from_attributes=True - ).model_dump(mode="json") + current_user.role = TenantService.get_user_role(current_user, current_user.current_tenant, session=db.session()) + installed_app_list = [ + _installed_app_response_data( + installed_app, + app_model, + current_tenant_id=current_tenant_id, + current_user=current_user, + ) + for installed_app, app_model in installed_apps + ] + + logger.debug("installed_app_list: %s, user_id: %s", installed_app_list, current_user.id) + return dump_response( + InstalledAppListResponse, + { + "installed_apps": installed_app_list, + "has_more": has_more, + "next_cursor": _encode_installed_app_cursor(next_cursor) if next_cursor else None, + }, + ) @login_required @account_initialization_required @@ -247,16 +196,15 @@ class InstalledAppsListApi(Resource): @console_ns.expect(console_ns.models[InstalledAppCreatePayload.__name__]) @console_ns.response(200, "Success", console_ns.models[SimpleMessageResponse.__name__]) @with_current_tenant_id - def post(self, current_tenant_id: str): - payload = InstalledAppCreatePayload.model_validate(console_ns.payload or {}) - + @model_validate(InstalledAppCreatePayload) + def post(self, req_data: InstalledAppCreatePayload, current_tenant_id: str): recommended_app = db.session.scalar( - select(RecommendedApp).where(RecommendedApp.app_id == payload.app_id).limit(1) + select(RecommendedApp).where(RecommendedApp.app_id == req_data.app_id).limit(1) ) if recommended_app is None: raise NotFound("Recommended app not found") - app = db.session.get(App, payload.app_id) + app = db.session.get(App, req_data.app_id) if app is None: raise NotFound("App entity not found") @@ -266,7 +214,7 @@ class InstalledAppsListApi(Resource): installed_app = db.session.scalar( select(InstalledApp) - .where(and_(InstalledApp.app_id == payload.app_id, InstalledApp.tenant_id == current_tenant_id)) + .where(and_(InstalledApp.app_id == req_data.app_id, InstalledApp.tenant_id == current_tenant_id)) .limit(1) ) @@ -275,7 +223,7 @@ class InstalledAppsListApi(Resource): recommended_app.install_count += 1 new_installed_app = InstalledApp( - app_id=payload.app_id, + app_id=req_data.app_id, tenant_id=current_tenant_id, app_owner_tenant_id=app.tenant_id, is_pinned=False, @@ -290,10 +238,36 @@ class InstalledAppsListApi(Resource): @console_ns.route("/installed-apps/") class InstalledAppApi(InstalledAppResource): """ - update and delete an installed app + get, update, and delete an installed app use InstalledAppResource to apply default decorators and get installed_app """ + @console_ns.response(200, "Success", console_ns.models[InstalledAppResponse.__name__]) + @with_current_user + @with_current_tenant_id + def get( + self, + current_tenant_id: str, + current_user: Account, + installed_app: InstalledApp, + ): + app_model = InstalledAppService.get_published_app(installed_app.app_id, session=db.session) + if app_model is None: + raise NotFound("Installed app not found") + if current_user.current_tenant is None: + raise ValueError("current_user.current_tenant must not be None") + + current_user.role = TenantService.get_user_role(current_user, current_user.current_tenant, session=db.session()) + return dump_response( + InstalledAppResponse, + _installed_app_response_data( + installed_app, + app_model, + current_tenant_id=current_tenant_id, + current_user=current_user, + ), + ) + @console_ns.response(204, "App uninstalled successfully") @with_current_tenant_id def delete(self, current_tenant_id: str, installed_app: InstalledApp): @@ -307,12 +281,11 @@ class InstalledAppApi(InstalledAppResource): @console_ns.response(200, "Success", console_ns.models[SimpleResultMessageResponse.__name__]) @console_ns.expect(console_ns.models[InstalledAppUpdatePayload.__name__]) - def patch(self, installed_app: InstalledApp): - payload = InstalledAppUpdatePayload.model_validate(console_ns.payload or {}) - + @model_validate(InstalledAppUpdatePayload) + def patch(self, req_data: InstalledAppUpdatePayload, installed_app: InstalledApp): commit_args = False - if payload.is_pinned is not None: - installed_app.is_pinned = payload.is_pinned + if req_data.is_pinned is not None: + installed_app.is_pinned = req_data.is_pinned commit_args = True if commit_args: diff --git a/api/controllers/console/explore/message.py b/api/controllers/console/explore/message.py index fb78ddede70..9e38c44c2c4 100644 --- a/api/controllers/console/explore/message.py +++ b/api/controllers/console/explore/message.py @@ -24,7 +24,7 @@ from controllers.console.explore.error import ( NotCompletionAppError, ) from controllers.console.explore.wraps import InstalledAppResource -from controllers.console.wraps import with_current_user +from controllers.console.wraps import model_validate, with_current_user from core.app.entities.app_invoke_entities import InvokeFrom from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from extensions.ext_database import db @@ -119,22 +119,23 @@ class MessageFeedbackApi(InstalledAppResource): @console_ns.expect(console_ns.models[MessageFeedbackPayload.__name__]) @console_ns.response(200, "Feedback submitted successfully", console_ns.models[ResultResponse.__name__]) @with_current_user - def post(self, current_user: Account, installed_app: InstalledApp, message_id: UUID): + @model_validate(MessageFeedbackPayload) + def post( + self, req_data: MessageFeedbackPayload, current_user: Account, installed_app: InstalledApp, message_id: UUID + ): app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() message_id_str = str(message_id) - payload = MessageFeedbackPayload.model_validate(console_ns.payload or {}) - try: MessageService.create_feedback( app_model=app_model, message_id=message_id_str, user=current_user, - rating=FeedbackRating(payload.rating) if payload.rating else None, - content=payload.content, + rating=FeedbackRating(req_data.rating) if req_data.rating else None, + content=req_data.content, session=db.session(), ) except MessageNotExistsError: diff --git a/api/controllers/console/explore/recommended_app.py b/api/controllers/console/explore/recommended_app.py index 79eaa305d61..aa91cd2d6c9 100644 --- a/api/controllers/console/explore/recommended_app.py +++ b/api/controllers/console/explore/recommended_app.py @@ -1,17 +1,16 @@ from typing import Any from uuid import UUID -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, RootModel, computed_field, field_validator from constants.languages import languages from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models from controllers.console import console_ns -from controllers.console.wraps import account_initialization_required, with_current_user +from controllers.console.wraps import account_initialization_required, model_validate, with_current_user from extensions.ext_database import db from fields.base import ResponseModel -from libs.helper import build_icon_url +from libs.helper import build_icon_url, dump_response from libs.login import login_required from models import Account from services.recommended_app_service import RecommendedAppService @@ -58,7 +57,7 @@ class RecommendedAppResponse(ResponseModel): categories: list[str] = Field(default_factory=list) position: int | None = None is_listed: bool | None = None - can_trial: bool | None = None + can_trial: bool class RecommendedAppListResponse(ResponseModel): @@ -77,7 +76,7 @@ class RecommendedAppDetailResponse(ResponseModel): icon_background: str | None = None mode: str export_data: str - can_trial: bool | None = None + can_trial: bool class RecommendedAppDetailNullableResponse(RootModel[RecommendedAppDetailResponse | None]): @@ -114,15 +113,15 @@ class RecommendedAppListApi(Resource): @login_required @account_initialization_required @with_current_user - def get(self, current_user: Account): + @model_validate(RecommendedAppsQuery) + def get(self, req_data: RecommendedAppsQuery, current_user: Account): # language args - args = RecommendedAppsQuery.model_validate(request.args.to_dict(flat=True)) - language_prefix = _resolve_language(args.language, current_user) + language_prefix = _resolve_language(req_data.language, current_user) - return RecommendedAppListResponse.model_validate( + return dump_response( + RecommendedAppListResponse, RecommendedAppService.get_recommended_apps_and_categories(language_prefix, session=db.session()), - from_attributes=True, - ).model_dump(mode="json") + ) @console_ns.route("/explore/apps/learn-dify") @@ -132,14 +131,14 @@ class LearnDifyAppListApi(Resource): @login_required @account_initialization_required @with_current_user - def get(self, current_user: Account): - args = RecommendedAppsQuery.model_validate(request.args.to_dict(flat=True)) - language_prefix = _resolve_language(args.language, current_user) + @model_validate(RecommendedAppsQuery) + def get(self, req_data: RecommendedAppsQuery, current_user: Account): + language_prefix = _resolve_language(req_data.language, current_user) - return LearnDifyAppListResponse.model_validate( + return dump_response( + LearnDifyAppListResponse, RecommendedAppService.get_learn_dify_apps(language_prefix, session=db.session()), - from_attributes=True, - ).model_dump(mode="json") + ) @console_ns.route("/explore/apps/") @@ -148,4 +147,5 @@ class RecommendedAppApi(Resource): @login_required @account_initialization_required def get(self, app_id: UUID): - return RecommendedAppService.get_recommend_app_detail(str(app_id), session=db.session()) + result = RecommendedAppService.get_recommend_app_detail(str(app_id), session=db.session()) + return RecommendedAppDetailNullableResponse.model_validate(result).model_dump(mode="json") diff --git a/api/controllers/console/explore/saved_message.py b/api/controllers/console/explore/saved_message.py index d057c49bff7..7662f9e3ac3 100644 --- a/api/controllers/console/explore/saved_message.py +++ b/api/controllers/console/explore/saved_message.py @@ -10,7 +10,7 @@ from controllers.console import console_ns from controllers.console.app.error import AppUnavailableError from controllers.console.explore.error import NotCompletionAppError from controllers.console.explore.wraps import InstalledAppResource -from controllers.console.wraps import with_current_user +from controllers.console.wraps import model_validate, with_current_user from extensions.ext_database import db from fields.conversation_fields import MessageResponseSource, ResultResponse from fields.message_fields import SavedMessageInfiniteScrollPagination, SavedMessageItem @@ -55,17 +55,16 @@ class SavedMessageListApi(InstalledAppResource): @console_ns.expect(console_ns.models[SavedMessageCreatePayload.__name__]) @console_ns.response(200, "Success", console_ns.models[ResultResponse.__name__]) @with_current_user - def post(self, current_user: Account, installed_app: InstalledApp): + @model_validate(SavedMessageCreatePayload) + def post(self, req_data: SavedMessageCreatePayload, current_user: Account, installed_app: InstalledApp): app_model = installed_app.app_with_session(session=db.session()) if app_model is None: raise AppUnavailableError() if app_model.mode != "completion": raise NotCompletionAppError() - payload = SavedMessageCreatePayload.model_validate(console_ns.payload or {}) - try: - SavedMessageService.save(app_model, current_user, str(payload.message_id), session=db.session()) + SavedMessageService.save(app_model, current_user, str(req_data.message_id), session=db.session()) except MessageNotExistsError: raise NotFound("Message Not Exists.") diff --git a/api/controllers/console/explore/trial.py b/api/controllers/console/explore/trial.py index 67ff9708959..620f3d3f85e 100644 --- a/api/controllers/console/explore/trial.py +++ b/api/controllers/console/explore/trial.py @@ -46,10 +46,10 @@ from controllers.console.explore.error import ( NotCompletionAppError, NotWorkflowAppError, ) -from controllers.console.explore.wraps import TrialAppResource, trial_feature_enable +from controllers.console.explore.wraps import TrialAppResource from controllers.console.files import FILE_UPLOAD_PARAMS, upload_file_from_request from controllers.console.remote_files import RemoteFileUploadPayload, upload_remote_file_from_request -from controllers.console.wraps import cloud_edition_billing_resource_check, with_current_user +from controllers.console.wraps import cloud_edition_billing_resource_check, model_validate, with_current_user from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.base_app_queue_manager import AppQueueManager @@ -59,6 +59,8 @@ from core.errors.error import ( ProviderTokenNotInitError, QuotaExceededError, ) +from core.helper import encrypter +from core.workflow.llm_environment_variable import LLMEnvironmentVariable, dump_environment_variable from extensions.ext_database import db from extensions.ext_redis import redis_client from fields.base import ResponseModel @@ -67,6 +69,7 @@ from fields.file_fields import FileResponse, FileWithSignedUrl from fields.message_fields import SuggestedQuestionsResponse from graphon.graph_engine.manager import GraphEngineManager from graphon.model_runtime.errors.invoke import InvokeError +from graphon.variables import SecretVariable, VariableBase from libs import helper from libs.helper import dump_response, to_timestamp, uuid_value from models import Account, App @@ -385,6 +388,26 @@ class TrialWorkflowResponse(ResponseModel): def _normalize_timestamp(cls, value: datetime | int | None) -> int | None: return to_timestamp(value) + @field_validator("environment_variables", mode="before") + @classmethod + def _serialize_environment_variables(cls, value: Any) -> list[Any]: + if value is None: + return [] + + result: list[Any] = [] + for item in value: + if isinstance(item, SecretVariable): + serialized = item.model_dump(mode="json") + serialized["value"] = encrypter.full_mask_token() + result.append(serialized) + elif isinstance(item, LLMEnvironmentVariable): + result.append(dump_environment_variable(item, mode="json")) + elif isinstance(item, VariableBase): + result.append(item.model_dump(mode="json")) + else: + result.append(item) + return result + @dataclass(frozen=True) class TrialWorkflowResponseSource: @@ -404,7 +427,7 @@ class TrialWorkflowResponseSource: return self.workflow.get_tool_published(session=self.session) def __getattr__(self, name: str) -> Any: - return getattr(self.workflow, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self.workflow, name) # guard-ignore: no-new-getattr -- delegates model fields register_schema_models( @@ -432,7 +455,6 @@ simple_account_model = console_ns.models[TrialSimpleAccount.__name__] class TrialAppFileUploadApi(TrialAppResource): - @trial_feature_enable @cloud_edition_billing_resource_check("documents") @console_ns.doc(consumes=["multipart/form-data"], params=FILE_UPLOAD_PARAMS) @console_ns.response(201, "File uploaded successfully", console_ns.models[FileResponse.__name__]) @@ -447,7 +469,6 @@ class TrialAppFileUploadApi(TrialAppResource): class TrialAppRemoteFileUploadApi(TrialAppResource): - @trial_feature_enable @cloud_edition_billing_resource_check("documents") @console_ns.expect(console_ns.models[RemoteFileUploadPayload.__name__]) @console_ns.response(201, "File uploaded successfully", console_ns.models[FileWithSignedUrl.__name__]) @@ -462,12 +483,12 @@ class TrialAppRemoteFileUploadApi(TrialAppResource): class TrialAppWorkflowRunApi(TrialAppResource): - @trial_feature_enable @console_ns.expect(console_ns.models[WorkflowRunRequest.__name__]) @console_ns.response(200, "Success") @with_current_user @with_session - def post(self, session: Session, current_user: Account, trial_app): + @model_validate(WorkflowRunRequest) + def post(self, req_data: WorkflowRunRequest, session: Session, current_user: Account, trial_app): """ Run workflow """ @@ -478,8 +499,7 @@ class TrialAppWorkflowRunApi(TrialAppResource): if app_mode != AppMode.WORKFLOW: raise NotWorkflowAppError() - request_data = WorkflowRunRequest.model_validate(console_ns.payload) - args = request_data.model_dump() + args = req_data.model_dump() try: app_id = app_model.id user_id = current_user.id @@ -513,7 +533,6 @@ class TrialAppWorkflowRunApi(TrialAppResource): class TrialAppWorkflowTaskStopApi(TrialAppResource): @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) - @trial_feature_enable def post(self, trial_app, task_id: str): """ Stop workflow task @@ -538,17 +557,16 @@ class TrialAppWorkflowTaskStopApi(TrialAppResource): class TrialChatApi(TrialAppResource): @console_ns.expect(console_ns.models[ChatRequest.__name__]) @console_ns.response(200, "Success") - @trial_feature_enable @with_current_user @with_session - def post(self, session: Session, current_user: Account, trial_app): + @model_validate(ChatRequest) + def post(self, req_data: ChatRequest, session: Session, current_user: Account, trial_app): app_model = trial_app app_mode = AppMode.value_of(app_model.mode) if app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT}: raise NotChatAppError() - request_data = ChatRequest.model_validate(console_ns.payload) - args = request_data.model_dump() + args = req_data.model_dump() # Validate UUID values if provided if args.get("conversation_id"): @@ -640,7 +658,6 @@ class TrialMessageSuggestedQuestionApi(TrialAppResource): class TrialChatAudioApi(TrialAppResource): @console_ns.response(200, "Success", console_ns.models[AudioTranscriptResponse.__name__]) - @trial_feature_enable @with_current_user def post(self, current_user: Account, trial_app): app_model = trial_app @@ -691,16 +708,14 @@ class TrialChatAudioApi(TrialAppResource): class TrialChatTextApi(TrialAppResource): @console_ns.expect(console_ns.models[TextToSpeechRequest.__name__]) @console_ns.response(200, "Success", console_ns.models[AudioBinaryResponse.__name__]) - @trial_feature_enable @with_current_user - def post(self, current_user: Account, trial_app): + @model_validate(TextToSpeechRequest) + def post(self, req_data: TextToSpeechRequest, current_user: Account, trial_app): app_model = trial_app try: - request_data = TextToSpeechRequest.model_validate(console_ns.payload) - - message_id = request_data.message_id - text = request_data.text - voice = request_data.voice + message_id = req_data.message_id + text = req_data.text + voice = req_data.voice message_ref = None if message_id: app_ref = AppRefService.create_app_ref(app_model) @@ -752,16 +767,15 @@ class TrialChatTextApi(TrialAppResource): class TrialCompletionApi(TrialAppResource): @console_ns.expect(console_ns.models[CompletionRequest.__name__]) @console_ns.response(200, "Success") - @trial_feature_enable @with_current_user @with_session - def post(self, session: Session, current_user: Account, trial_app): + @model_validate(CompletionRequest) + def post(self, req_data: CompletionRequest, session: Session, current_user: Account, trial_app): app_model = trial_app if app_model.mode != "completion": raise NotCompletionAppError() - request_data = CompletionRequest.model_validate(console_ns.payload) - args = request_data.model_dump() + args = req_data.model_dump() streaming = args["response_mode"] == "streaming" args["auto_generate_name"] = False diff --git a/api/controllers/console/explore/workflow.py b/api/controllers/console/explore/workflow.py index a8c176c6778..29e76d81498 100644 --- a/api/controllers/console/explore/workflow.py +++ b/api/controllers/console/explore/workflow.py @@ -15,7 +15,7 @@ from controllers.console.app.error import ( from controllers.console.app.wraps import with_session from controllers.console.explore.error import NotWorkflowAppError from controllers.console.explore.wraps import InstalledAppResource -from controllers.console.wraps import with_current_user +from controllers.console.wraps import model_validate, with_current_user from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError from core.app.apps.base_app_queue_manager import AppQueueManager from core.app.entities.app_invoke_entities import InvokeFrom @@ -47,7 +47,14 @@ class InstalledAppWorkflowRunApi(InstalledAppResource): @console_ns.response(200, "Success") @with_current_user @with_session - def post(self, session: Session, current_user: Account, installed_app: InstalledApp): + @model_validate(WorkflowRunPayload) + def post( + self, + req_data: WorkflowRunPayload, + session: Session, + current_user: Account, + installed_app: InstalledApp, + ): """ Run workflow """ @@ -58,8 +65,7 @@ class InstalledAppWorkflowRunApi(InstalledAppResource): if app_mode != AppMode.WORKFLOW: raise NotWorkflowAppError() - payload = WorkflowRunPayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) try: response = AppGenerateService.generate( session=session, diff --git a/api/controllers/console/explore/wraps.py b/api/controllers/console/explore/wraps.py index d67f3e18d53..1f4da57f9aa 100644 --- a/api/controllers/console/explore/wraps.py +++ b/api/controllers/console/explore/wraps.py @@ -14,6 +14,7 @@ from libs.login import current_account_with_tenant, login_required from models import AccountTrialAppRecord, App, InstalledApp, TrialApp from services.enterprise.enterprise_service import EnterpriseService from services.feature_service import FeatureService +from services.recommended_app_service import RecommendedAppService def installed_app_required[**P, R](view: Callable[Concatenate[InstalledApp, P], R] | None = None): @@ -106,25 +107,13 @@ def trial_app_required[**P, R](view: Callable[Concatenate[App, P], R] | None = N def trial_feature_enable[**P, R](view: Callable[P, R]): @wraps(view) def decorated(*args: P.args, **kwargs: P.kwargs): - features = FeatureService.get_system_features() - if not features.enable_trial_app: + if not RecommendedAppService.is_trial_app_enabled(): abort(403, "Trial app feature is not enabled.") return view(*args, **kwargs) return decorated -def explore_banner_enabled[**P, R](view: Callable[P, R]): - @wraps(view) - def decorated(*args: P.args, **kwargs: P.kwargs): - features = FeatureService.get_system_features() - if not features.enable_explore_banner: - abort(403, "Explore banner feature is not enabled.") - return view(*args, **kwargs) - - return decorated - - class InstalledAppResource(Resource): # must be reversed if there are multiple decorators @@ -141,6 +130,7 @@ class TrialAppResource(Resource): method_decorators = [ trial_app_required, + trial_feature_enable, account_initialization_required, login_required, ] diff --git a/api/controllers/console/extension.py b/api/controllers/console/extension.py index cc06204a905..b7ba7bf8075 100644 --- a/api/controllers/console/extension.py +++ b/api/controllers/console/extension.py @@ -17,7 +17,7 @@ from services.code_based_extension_service import CodeBasedExtensionService from ..common.schema import query_params_from_model, register_response_schema_models, register_schema_models from . import console_ns -from .wraps import account_initialization_required, setup_required, with_current_tenant_id +from .wraps import account_initialization_required, model_validate, setup_required, with_current_tenant_id class CodeBasedExtensionQuery(BaseModel): @@ -123,14 +123,13 @@ class APIBasedExtensionAPI(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): - payload = APIBasedExtensionPayload.model_validate(console_ns.payload or {}) - + @model_validate(APIBasedExtensionPayload) + def post(self, req_data: APIBasedExtensionPayload, current_tenant_id: str): extension_data = APIBasedExtension( tenant_id=current_tenant_id, - name=payload.name, - api_endpoint=payload.api_endpoint, - api_key=payload.api_key, + name=req_data.name, + api_endpoint=req_data.api_endpoint, + api_key=req_data.api_key, ) extension = APIBasedExtensionService.save(extension_data, session=db.session()) @@ -138,7 +137,7 @@ class APIBasedExtensionAPI(Resource): id=extension.id, name=extension.name, api_endpoint=extension.api_endpoint, - api_key=payload.api_key, + api_key=req_data.api_key, created_at=to_timestamp(extension.created_at), ).model_dump(mode="json"), 201 @@ -172,22 +171,22 @@ class APIBasedExtensionDetailAPI(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str, id: UUID): + @model_validate(APIBasedExtensionPayload) + def post(self, req_data: APIBasedExtensionPayload, current_tenant_id: str, id: UUID): api_based_extension_id = str(id) extension_data_from_db = APIBasedExtensionService.get_with_tenant_id( current_tenant_id, api_based_extension_id, session=db.session() ) - payload = APIBasedExtensionPayload.model_validate(console_ns.payload or {}) api_key_for_response = extension_data_from_db.api_key - extension_data_from_db.name = payload.name - extension_data_from_db.api_endpoint = payload.api_endpoint + extension_data_from_db.name = req_data.name + extension_data_from_db.api_endpoint = req_data.api_endpoint - if payload.api_key != HIDDEN_VALUE: - extension_data_from_db.api_key = payload.api_key - api_key_for_response = payload.api_key + if req_data.api_key != HIDDEN_VALUE: + extension_data_from_db.api_key = req_data.api_key + api_key_for_response = req_data.api_key APIBasedExtensionService.save(extension_data_from_db, session=db.session()) return APIBasedExtensionResponse( diff --git a/api/controllers/console/feature.py b/api/controllers/console/feature.py index 594a8431513..41aa6513a63 100644 --- a/api/controllers/console/feature.py +++ b/api/controllers/console/feature.py @@ -1,24 +1,21 @@ from flask_restx import Resource from controllers.common.schema import register_response_schema_models +from controllers.console.flask_admission import console_account_admission +from extensions.ext_application_services import application_services from fields.base import ResponseModel from libs.helper import dump_response -from libs.login import login_required -from services.feature_service import ( +from machinery.context import RequestContext +from services.entities.feature_entities import ( FeatureModel, - FeatureService, LicenseModel, LimitationModel, SystemFeatureModel, + VectorSpaceLimitationModel, ) from . import console_ns -from .wraps import ( - account_initialization_required, - cloud_utm_record, - setup_required, - with_current_tenant_id, -) +from .wraps import cloud_utm_record class TrialModelsResponse(ResponseModel): @@ -37,6 +34,7 @@ register_response_schema_models( LimitationModel, SystemFeatureModel, TrialModelsResponse, + VectorSpaceLimitationModel, ) @@ -49,17 +47,11 @@ class FeatureApi(Resource): "Success", console_ns.models[FeatureModel.__name__], ) - @setup_required - @login_required - @account_initialization_required + @console_account_admission() @cloud_utm_record - @with_current_tenant_id - def get(self, current_tenant_id: str): + def get(self, request_context: RequestContext): """Get feature configuration for current tenant""" - payload = FeatureService.get_features( - current_tenant_id, - exclude_vector_space=True, - ).model_dump() + payload = application_services().feature_queries.get_features(request_context).model_dump() payload.pop("vector_space", None) return payload @@ -71,16 +63,13 @@ class FeatureVectorSpaceApi(Resource): @console_ns.response( 200, "Success", - console_ns.models[LimitationModel.__name__], + console_ns.models[VectorSpaceLimitationModel.__name__], ) - @setup_required - @login_required - @account_initialization_required + @console_account_admission() @cloud_utm_record - @with_current_tenant_id - def get(self, current_tenant_id: str): + def get(self, request_context: RequestContext): """Get vector-space usage and limit for current tenant""" - return FeatureService.get_vector_space(current_tenant_id).model_dump() + return application_services().feature_queries.get_vector_space(request_context).model_dump() @console_ns.route("/trial-models") @@ -92,14 +81,12 @@ class TrialModelsApi(Resource): "Success", console_ns.models[TrialModelsResponse.__name__], ) - @setup_required - @login_required - @account_initialization_required - def get(self): + @console_account_admission() + def get(self, _request_context: RequestContext): """Get hosted trial model provider configuration for model-provider pages.""" return dump_response( TrialModelsResponse, - {"trial_models": FeatureService.get_trial_models()}, + {"trial_models": application_services().feature_queries.get_trial_models()}, ) @@ -116,7 +103,7 @@ class AppDslVersionApi(Resource): """Get current app DSL version for workflow clipboard compatibility.""" return dump_response( AppDslVersionResponse, - {"app_dsl_version": FeatureService.get_app_dsl_version()}, + {"app_dsl_version": application_services().feature_queries.get_app_dsl_version()}, ) @@ -138,7 +125,7 @@ class SystemFeatureApi(Resource): Authentication configuration must be available before the authentication flow can be selected. Authenticated license detail is served separately by SystemFeatureLicenseApi. """ - return dump_response(SystemFeatureModel, FeatureService.get_system_features()) + return dump_response(SystemFeatureModel, application_services().feature_queries.get_system_features()) @console_ns.route("/system-features/license") @@ -150,13 +137,11 @@ class SystemFeatureLicenseApi(Resource): "Success", console_ns.models[LicenseModel.__name__], ) - @setup_required - @login_required - @account_initialization_required - def get(self): + @console_account_admission() + def get(self, _request_context: RequestContext): """Get full license detail (status, expiry, workspace/seat usage). Authenticated counterpart to the license *status* exposed on the public system-features endpoint. """ - return FeatureService.get_license().model_dump() + return application_services().feature_queries.get_license().model_dump() diff --git a/api/controllers/console/files.py b/api/controllers/console/files.py index a2df1a07845..2f166340846 100644 --- a/api/controllers/console/files.py +++ b/api/controllers/console/files.py @@ -30,6 +30,7 @@ from fields.file_fields import FileResponse, UploadConfig from libs.helper import dump_response from libs.login import login_required from models import Account, UploadFile +from services.feature_service import FeatureService from services.file_service import FileService from . import console_ns @@ -58,7 +59,7 @@ FILE_UPLOAD_PARAMS = { def upload_file_from_request(*, current_user: Account, resource_tenant_id: str | None = None) -> UploadFile: """Validate the multipart request and persist the file under the requested resource tenant.""" - source_str = request.form.get("source") + source_str = request.args.get("source") or request.form.get("source") source: Literal["datasets"] | None = "datasets" if source_str == "datasets" else None if "file" not in request.files: @@ -76,6 +77,12 @@ def upload_file_from_request(*, current_user: Account, resource_tenant_id: str | if source not in ("datasets", None): source = None + default_file_size_limit = ( + FeatureService.get_knowledge_file_size_limit(resource_tenant_id or current_user.current_tenant_id) + if source == "datasets" + else None + ) + try: return FileService(db.engine).upload_file( filename=file.filename, @@ -84,6 +91,7 @@ def upload_file_from_request(*, current_user: Account, resource_tenant_id: str | user=current_user, tenant_id=resource_tenant_id, source=source, + default_file_size_limit=default_file_size_limit, ) except services.errors.file.FileTooLargeError as file_too_large_error: raise FileTooLargeError(file_too_large_error.description) @@ -99,14 +107,17 @@ class FileApi(Resource): @login_required @account_initialization_required @console_ns.response(200, "Success", console_ns.models[UploadConfig.__name__]) - def get(self): + @with_current_tenant_id + def get(self, current_tenant_id: str): config = UploadConfig( file_size_limit=dify_config.UPLOAD_FILE_SIZE_LIMIT, + knowledge_file_size_limit=FeatureService.get_knowledge_file_size_limit(current_tenant_id), batch_count_limit=dify_config.UPLOAD_FILE_BATCH_LIMIT, file_upload_limit=dify_config.BATCH_UPLOAD_LIMIT, image_file_size_limit=dify_config.UPLOAD_IMAGE_FILE_SIZE_LIMIT, video_file_size_limit=dify_config.UPLOAD_VIDEO_FILE_SIZE_LIMIT, audio_file_size_limit=dify_config.UPLOAD_AUDIO_FILE_SIZE_LIMIT, + skill_file_size_limit=dify_config.UPLOAD_SKILL_FILE_SIZE_LIMIT, workflow_file_upload_limit=dify_config.WORKFLOW_FILE_UPLOAD_LIMIT, image_file_batch_limit=dify_config.IMAGE_FILE_BATCH_LIMIT, single_chunk_attachment_limit=dify_config.SINGLE_CHUNK_ATTACHMENT_LIMIT, diff --git a/api/controllers/console/flask_admission.py b/api/controllers/console/flask_admission.py new file mode 100644 index 00000000000..d7ce23b67e1 --- /dev/null +++ b/api/controllers/console/flask_admission.py @@ -0,0 +1,64 @@ +"""Flask adapter for Console API admission.""" + +from collections.abc import Callable +from functools import wraps +from typing import Concatenate + +from flask import Response, abort, request + +from configs import dify_config +from controllers.console.wraps import account_initialization_required, enterprise_license_required, setup_required +from core.logging.context import get_request_id, get_trace_id +from enums import DeploymentEdition +from libs.login import current_account_with_tenant, login_required +from machinery.context import RequestContext + + +def console_account_admission[T, **P, R]( + *, + editions: frozenset[DeploymentEdition] | None = None, + require_valid_enterprise_license: bool = False, +) -> Callable[ + [Callable[Concatenate[T, RequestContext, P], R]], + Callable[Concatenate[T, P], R | Response], +]: + """Declare Console account admission and inject a stable RequestContext. + + All combinations use this decorator factory. Requirements are data, while + the execution order stays fixed: edition, setup, login/CSRF, account + initialization, optional enterprise license, then context construction. + """ + + def decorator( + view: Callable[Concatenate[T, RequestContext, P], R], + ) -> Callable[Concatenate[T, P], R | Response]: + @wraps(view) + def inject_request_context(self: T, /, *args: P.args, **kwargs: P.kwargs) -> R: + account_with_tenant = current_account_with_tenant() + request_context = RequestContext( + account_id=account_with_tenant.account.id, + active_workspace_id=account_with_tenant.tenant_id, + request_id=get_request_id(), + trace_id=get_trace_id() or request.headers.get("X-Trace-Id"), + ) + return view(self, request_context, *args, **kwargs) + + admitted: Callable[Concatenate[T, P], R | Response] = inject_request_context + if require_valid_enterprise_license: + admitted = enterprise_license_required(admitted) + admitted = account_initialization_required(admitted) + admitted = login_required(admitted) + admitted = setup_required(admitted) + + if editions is None: + return admitted + + @wraps(view) + def enforce_edition(self: T, /, *args: P.args, **kwargs: P.kwargs) -> R | Response: + if dify_config.DEPLOYMENT_EDITION not in editions: + abort(404) + return admitted(self, *args, **kwargs) + + return enforce_edition + + return decorator diff --git a/api/controllers/console/human_input_form.py b/api/controllers/console/human_input_form.py index 9700667f4aa..4cbf29a89f1 100644 --- a/api/controllers/console/human_input_form.py +++ b/api/controllers/console/human_input_form.py @@ -137,7 +137,7 @@ class ConsoleHumanInputFormApi(Resource): self._ensure_console_access(form, current_tenant_id) self._ensure_console_recipient_type(form) recipient_type = form.recipient_type - # The type checker is not smart enought to validate the following invariant. + # The type checker is not smart enough to validate the following invariant. # So we need to assert it manually. assert recipient_type is not None, "recipient_type cannot be None here." @@ -215,6 +215,7 @@ class ConsoleWorkflowEventsApi(Resource): raise InvalidArgumentError(f"cannot subscribe to workflow run, workflow_run_id={workflow_run.id}") include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true" + continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true" def _generate_stream_events(): if include_state_snapshot: @@ -225,6 +226,8 @@ class ConsoleWorkflowEventsApi(Resource): tenant_id=workflow_run.tenant_id, app_id=workflow_run.app_id, session_maker=session_maker, + human_input_surface=HumanInputSurface.CONSOLE, + close_on_pause=not continue_on_pause, ) ) return generator.convert_to_event_stream( diff --git a/api/controllers/console/init_validate.py b/api/controllers/console/init_validate.py index aca22c67f8e..171d58c1caa 100644 --- a/api/controllers/console/init_validate.py +++ b/api/controllers/console/init_validate.py @@ -1,17 +1,11 @@ -import os from typing import Literal from flask import session from pydantic import BaseModel, Field -from sqlalchemy import select -from sqlalchemy.orm import Session -from configs import dify_config from controllers.fastopenapi import console_router -from enums.deployment_edition import DeploymentEdition -from extensions.ext_database import db -from models.model import DifySetup -from services.account_service import TenantService +from extensions.ext_application_services import application_services +from services.init_validation_service import AlreadyInitializedError, InvalidInitializationPasswordError from .error import AlreadySetupError, InitValidateFailedError from .wraps import only_edition_self_hosted @@ -36,7 +30,7 @@ class InitValidateResponse(BaseModel): ) def get_init_status() -> InitStatusResponse: """Get initialization validation status.""" - init_status = get_init_validate_status() + init_status = is_init_validated() if init_status: return InitStatusResponse(status="finished") return InitStatusResponse(status="not_started") @@ -51,25 +45,19 @@ def get_init_status() -> InitStatusResponse: @only_edition_self_hosted def validate_init_password(payload: InitValidatePayload) -> InitValidateResponse: """Validate initialization password.""" - tenant_count = TenantService.get_tenant_count(session=db.session()) - if tenant_count > 0: - raise AlreadySetupError() - - if payload.password != os.environ.get("INIT_PASSWORD"): + try: + application_services().init_validation.validate_password(payload.password) + except AlreadyInitializedError: + raise AlreadySetupError() from None + except InvalidInitializationPasswordError: session["is_init_validated"] = False - raise InitValidateFailedError() + raise InitValidateFailedError() from None session["is_init_validated"] = True return InitValidateResponse(result="success") -def get_init_validate_status() -> bool: - if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: - if os.environ.get("INIT_PASSWORD"): - if session.get("is_init_validated"): - return True - - with Session(db.engine) as db_session: - return db_session.execute(select(DifySetup)).scalar_one_or_none() is not None - - return True +def is_init_validated() -> bool: + return application_services().init_validation.is_validated( + session_validated=bool(session.get("is_init_validated")), + ) diff --git a/api/controllers/console/notification.py b/api/controllers/console/notification.py index 798d674f9e9..3e58f598bf7 100644 --- a/api/controllers/console/notification.py +++ b/api/controllers/console/notification.py @@ -1,7 +1,6 @@ from collections.abc import Mapping from typing import TypedDict -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field @@ -10,6 +9,7 @@ from controllers.common.schema import register_response_schema_models, register_ from controllers.console import console_ns from controllers.console.wraps import ( account_initialization_required, + model_validate, only_edition_cloud, setup_required, with_current_user, @@ -141,8 +141,8 @@ class NotificationDismissApi(Resource): @only_edition_cloud @console_ns.expect(console_ns.models[DismissNotificationPayload.__name__]) @console_ns.response(200, "Success", console_ns.models[SimpleResultResponse.__name__]) - def post(self, current_user: Account): - payload = DismissNotificationPayload.model_validate(request.get_json()) + @model_validate(DismissNotificationPayload) + def post(self, payload: DismissNotificationPayload, current_user: Account): BillingService.dismiss_notification( notification_id=payload.notification_id, account_id=str(current_user.id), diff --git a/api/controllers/console/onboarding.py b/api/controllers/console/onboarding.py index 7458bd46a3d..f26e2d539e4 100644 --- a/api/controllers/console/onboarding.py +++ b/api/controllers/console/onboarding.py @@ -21,7 +21,13 @@ from models import Account from services.step_by_step_tour_service import StepByStepTourPatch, StepByStepTourService from . import console_ns -from .wraps import account_initialization_required, setup_required, with_current_tenant_id, with_current_user +from .wraps import ( + account_initialization_required, + model_validate, + setup_required, + with_current_tenant_id, + with_current_user, +) StepByStepTourAction = Literal[ "skip", @@ -92,9 +98,9 @@ class StepByStepTourStateApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def patch(self, current_tenant_id: str, current_user: Account): - payload = StepByStepTourStatePatchPayload.model_validate(console_ns.payload or {}) - patch = cast(StepByStepTourPatch, payload.model_dump(exclude_unset=True, exclude_none=True)) + @model_validate(StepByStepTourStatePatchPayload) + def patch(self, req_data: StepByStepTourStatePatchPayload, current_tenant_id: str, current_user: Account): + patch = cast(StepByStepTourPatch, req_data.model_dump(exclude_unset=True, exclude_none=True)) return dump_response( StepByStepTourStateResponse, StepByStepTourService.patch_state( diff --git a/api/controllers/console/ping.py b/api/controllers/console/ping.py deleted file mode 100644 index d480af312b1..00000000000 --- a/api/controllers/console/ping.py +++ /dev/null @@ -1,17 +0,0 @@ -from pydantic import BaseModel, Field - -from controllers.fastopenapi import console_router - - -class PingResponse(BaseModel): - result: str = Field(description="Health check result", examples=["pong"]) - - -@console_router.get( - "/ping", - response_model=PingResponse, - tags=["console"], -) -def ping() -> PingResponse: - """Health check endpoint for connection testing.""" - return PingResponse(result="pong") diff --git a/api/controllers/console/remote_files.py b/api/controllers/console/remote_files.py index c771f4489ac..e3aefb72595 100644 --- a/api/controllers/console/remote_files.py +++ b/api/controllers/console/remote_files.py @@ -51,8 +51,8 @@ def upload_remote_file_from_request( current_user: Account, resource_tenant_id: str | None = None, ) -> FileWithSignedUrl: - """Validate the JSON request, fetch its remote file, and persist it under the requested tenant.""" payload = RemoteFileUploadPayload.model_validate(console_ns.payload) + """Validate the JSON request, fetch its remote file, and persist it under the requested tenant.""" url = payload.url # Try to fetch remote file metadata/content first diff --git a/api/controllers/console/setup.py b/api/controllers/console/setup.py index 825cfb02c3e..20fa631d057 100644 --- a/api/controllers/console/setup.py +++ b/api/controllers/console/setup.py @@ -2,18 +2,19 @@ from typing import Literal from flask import request from pydantic import BaseModel, Field, field_validator -from sqlalchemy import select -from configs import dify_config from controllers.fastopenapi import console_router -from enums.deployment_edition import DeploymentEdition +from extensions.ext_application_services import application_services from libs.helper import EmailStr, extract_remote_ip from libs.password import valid_password -from models.model import DifySetup, db -from services.account_service import RegisterService, TenantService +from services.setup_service import ( + InitializationValidationRequiredError, + SetupAlreadyCompletedError, + SetupInput, +) from .error import AlreadySetupError, NotInitValidateError -from .init_validate import get_init_validate_status +from .init_validate import is_init_validated from .wraps import mark_setup_completed, only_edition_self_hosted @@ -53,14 +54,12 @@ def get_setup_status_api() -> SetupStatusResponse: Only bootstrap-safe status information should be returned by this endpoint. """ - if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: - setup_status = get_setup_status() - if setup_status and not isinstance(setup_status, bool): - return SetupStatusResponse(step="finished", setup_at=setup_status.setup_at.isoformat()) - if setup_status: - return SetupStatusResponse(step="finished") + setup_status = application_services().setup.get_status() + if not setup_status.completed: return SetupStatusResponse(step="not_started") - return SetupStatusResponse(step="finished") + + setup_at = setup_status.setup_at.isoformat() if setup_status.setup_at is not None else None + return SetupStatusResponse(step="finished", setup_at=setup_at) @console_router.post( @@ -74,36 +73,25 @@ def setup_system(payload: SetupRequestPayload) -> SetupResponse: """Initialize system setup with admin account. NOTE: This endpoint is unauthenticated by design for first-time bootstrap. - Access is restricted by deployment mode (`SELF_HOSTED`), one-time setup guards, + Access is restricted to self-hosted editions (`COMMUNITY` and `ENTERPRISE`), one-time setup guards, and init-password validation rather than user session authentication. """ - if get_setup_status(): - raise AlreadySetupError() + try: + application_services().setup.initialize( + SetupInput( + email=payload.email, + name=payload.name, + password=payload.password, + ip_address=extract_remote_ip(request), + language=payload.language, + ), + initialization_validated=is_init_validated(), + ) + except SetupAlreadyCompletedError: + raise AlreadySetupError() from None + except InitializationValidationRequiredError: + raise NotInitValidateError() from None - tenant_count = TenantService.get_tenant_count(session=db.session()) - if tenant_count > 0: - raise AlreadySetupError() - - if not get_init_validate_status(): - raise NotInitValidateError() - - normalized_email = payload.email.lower() - - RegisterService.setup( - email=normalized_email, - name=payload.name, - password=payload.password, - ip_address=extract_remote_ip(request), - language=payload.language, - session=db.session(), - ) mark_setup_completed() return SetupResponse(result="success") - - -def get_setup_status() -> DifySetup | bool | None: - if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: - return db.session.scalar(select(DifySetup).limit(1)) - - return True diff --git a/api/controllers/console/snippets/snippet_workflow.py b/api/controllers/console/snippets/snippet_workflow.py index 558d50930c5..d1399801a06 100644 --- a/api/controllers/console/snippets/snippet_workflow.py +++ b/api/controllers/console/snippets/snippet_workflow.py @@ -36,6 +36,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_user, @@ -56,10 +57,12 @@ from libs.helper import TimestampField from libs.login import current_account_with_tenant, login_required from models import Account from models.snippet import CustomizedSnippet +from services.agent.retirement_service import WorkflowAgentRetirementService from services.agent.workflow_publish_service import WorkflowAgentPublishService from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError from services.snippet_generate_service import SnippetGenerateService from services.snippet_service import SnippetService +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection logger = logging.getLogger(__name__) @@ -202,18 +205,18 @@ class SnippetDraftWorkflowApi(Resource): @rbac_permission_required( RBACResourceScope.WORKSPACE, RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False ) - def post(self, current_user: Account, snippet: CustomizedSnippet): + @model_validate(SnippetDraftSyncPayload) + def post(self, req_data: SnippetDraftSyncPayload, current_user: Account, snippet: CustomizedSnippet): """Sync draft workflow for snippet.""" - payload = SnippetDraftSyncPayload.model_validate(console_ns.payload or {}) try: snippet_service = _snippet_service() workflow = snippet_service.sync_draft_workflow( snippet=snippet, - graph=payload.graph, - unique_hash=payload.hash, + graph=req_data.graph, + unique_hash=req_data.hash, account=current_user, - input_fields=payload.input_fields, + input_fields=req_data.input_fields, ) except WorkflowHashNotEqualError: raise DraftWorkflowNotSync() @@ -295,8 +298,9 @@ class SnippetPublishedWorkflowApi(Resource): with Session(db.engine) as session: snippet = session.merge(snippet) + tenant_id = snippet.tenant_id try: - workflow = snippet_service.publish_workflow( + workflow, retirement_candidates = snippet_service.publish_workflow( session=session, snippet=snippet, account=current_user, @@ -306,6 +310,16 @@ class SnippetPublishedWorkflowApi(Resource): except ValueError as e: return {"message": str(e)}, 400 + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=tenant_id, + agent_ids=retirement_candidates, + account_id=current_user.id, + ) + enqueue_agent_resource_collection( + tenant_id=tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) return { "result": "success", "created_at": workflow_created_at, @@ -350,24 +364,24 @@ class SnippetPublishedAllWorkflowApi(Resource): @rbac_permission_required( RBACResourceScope.WORKSPACE, RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False ) - def get(self, snippet: CustomizedSnippet): + @model_validate(SnippetWorkflowListQuery) + def get(self, req_data: SnippetWorkflowListQuery, snippet: CustomizedSnippet): """Get all published workflow versions for snippet.""" - args = SnippetWorkflowListQuery.model_validate(request.args.to_dict(flat=True)) snippet_service = _snippet_service() with Session(db.engine) as session: workflows, has_more = snippet_service.get_all_published_workflows( session=session, snippet=snippet, - page=args.page, - limit=args.limit, + page=req_data.page, + limit=req_data.limit, ) response = SnippetWorkflowPaginationResponse.model_validate( { "items": workflows, - "page": args.page, - "limit": args.limit, + "page": req_data.page, + "limit": req_data.limit, "has_more": has_more, }, from_attributes=True, @@ -436,10 +450,16 @@ class SnippetWorkflowByIdApi(Resource): @rbac_permission_required( RBACResourceScope.WORKSPACE, RBACPermission.SNIPPETS_CREATE_AND_MODIFY, resource_required=False ) - def patch(self, current_user: Account, snippet: CustomizedSnippet, workflow_id: str): + @model_validate(WorkflowUpdatePayload) + def patch( + self, + req_data: WorkflowUpdatePayload, + current_user: Account, + snippet: CustomizedSnippet, + workflow_id: str, + ): """Update a published snippet workflow version's display metadata.""" - payload = WorkflowUpdatePayload.model_validate(console_ns.payload or {}) - update_data = payload.model_dump(exclude_unset=True) + update_data = req_data.model_dump(exclude_unset=True) if not update_data: return {"message": "No valid fields to update"}, 400 @@ -562,16 +582,22 @@ class SnippetDraftNodeRunApi(Resource): @with_current_user @get_snippet @edit_permission_required - def post(self, current_user: Account, snippet: CustomizedSnippet, node_id: str): + @model_validate(SnippetDraftNodeRunPayload) + def post( + self, + req_data: SnippetDraftNodeRunPayload, + current_user: Account, + snippet: CustomizedSnippet, + node_id: str, + ): """ Run a single node in snippet draft workflow. Executes a specific node with provided inputs for single-step debugging. Returns the node execution result including status, outputs, and timing. """ - payload = SnippetDraftNodeRunPayload.model_validate(console_ns.payload or {}) - user_inputs = payload.inputs + user_inputs = req_data.inputs # Get draft workflow for file parsing snippet_service = _snippet_service() @@ -579,14 +605,14 @@ class SnippetDraftNodeRunApi(Resource): if not draft_workflow: raise NotFound("Draft workflow not found") - files = SnippetGenerateService.parse_files(draft_workflow, payload.files) + files = SnippetGenerateService.parse_files(draft_workflow, req_data.files) workflow_node_execution = SnippetGenerateService.run_draft_node( snippet=snippet, node_id=node_id, user_inputs=user_inputs, account=current_user, - query=payload.query, + query=req_data.query, files=files, session_maker=_snippet_session_maker(), ) @@ -650,14 +676,21 @@ class SnippetDraftRunIterationNodeApi(Resource): @with_current_user @get_snippet @edit_permission_required - def post(self, current_user: Account, snippet: CustomizedSnippet, node_id: str): + @model_validate(SnippetIterationNodeRunPayload) + def post( + self, + req_data: SnippetIterationNodeRunPayload, + current_user: Account, + snippet: CustomizedSnippet, + node_id: str, + ): """ Run a draft workflow iteration node for snippet. Iteration nodes execute their internal sub-graph multiple times over an input list. Returns an SSE event stream with iteration progress and results. """ - args = SnippetIterationNodeRunPayload.model_validate(console_ns.payload or {}).model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) try: response = SnippetGenerateService.generate_single_iteration( @@ -695,21 +728,27 @@ class SnippetDraftRunLoopNodeApi(Resource): @with_current_user @get_snippet @edit_permission_required - def post(self, current_user: Account, snippet: CustomizedSnippet, node_id: str): + @model_validate(SnippetLoopNodeRunPayload) + def post( + self, + req_data: SnippetLoopNodeRunPayload, + current_user: Account, + snippet: CustomizedSnippet, + node_id: str, + ): """ Run a draft workflow loop node for snippet. Loop nodes execute their internal sub-graph repeatedly until a condition is met. Returns an SSE event stream with loop progress and results. """ - args = SnippetLoopNodeRunPayload.model_validate(console_ns.payload or {}) try: response = SnippetGenerateService.generate_single_loop( snippet=snippet, user=current_user, node_id=node_id, - args=args, + args=req_data, streaming=True, session_maker=_snippet_session_maker(), ) @@ -738,15 +777,15 @@ class SnippetDraftWorkflowRunApi(Resource): @with_current_user @get_snippet @edit_permission_required - def post(self, current_user: Account, snippet: CustomizedSnippet): + @model_validate(SnippetDraftRunPayload) + def post(self, req_data: SnippetDraftRunPayload, current_user: Account, snippet: CustomizedSnippet): """ Run draft workflow for snippet. Executes the snippet's draft workflow with the provided inputs and returns an SSE event stream with execution progress and results. """ - payload = SnippetDraftRunPayload.model_validate(console_ns.payload or {}) - args = payload.model_dump(exclude_none=True) + args = req_data.model_dump(exclude_none=True) try: response = SnippetGenerateService.generate( diff --git a/api/controllers/console/snippets/snippet_workflow_draft_variable.py b/api/controllers/console/snippets/snippet_workflow_draft_variable.py index a28ba07b5dd..725c168ea60 100644 --- a/api/controllers/console/snippets/snippet_workflow_draft_variable.py +++ b/api/controllers/console/snippets/snippet_workflow_draft_variable.py @@ -36,10 +36,12 @@ from controllers.console.snippets.snippet_workflow import get_snippet from controllers.console.wraps import ( account_initialization_required, edit_permission_required, + model_validate, setup_required, with_current_user, ) from core.app.file_access import DatabaseFileAccessController +from core.workflow.llm_environment_variable import environment_variable_value_type from core.workflow.variable_prefixes import CONVERSATION_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID from extensions.ext_database import db from factories.file_factory import build_from_mapping, build_from_mappings @@ -185,9 +187,15 @@ class SnippetVariableApi(Resource): @console_ns.response(404, "Variable not found") @_snippet_draft_var_prerequisite @marshal_with(workflow_draft_variable_model) - def patch(self, current_user: Account, snippet: CustomizedSnippet, variable_id: str) -> WorkflowDraftVariable: + @model_validate(WorkflowDraftVariableUpdatePayload) + def patch( + self, + req_data: WorkflowDraftVariableUpdatePayload, + current_user: Account, + snippet: CustomizedSnippet, + variable_id: str, + ) -> WorkflowDraftVariable: draft_var_srv = WorkflowDraftVariableService(session=db.session()) - args_model = WorkflowDraftVariableUpdatePayload.model_validate(console_ns.payload or {}) variable = ensure_variable_access( variable=draft_var_srv.get_variable(variable_id=variable_id), @@ -197,8 +205,8 @@ class SnippetVariableApi(Resource): ) _ensure_snippet_draft_variable_row_allowed(variable=variable, variable_id=variable_id) - new_name = args_model.name - raw_value = args_model.value + new_name = req_data.name + raw_value = req_data.value if new_name is None and raw_value is None: return variable @@ -329,7 +337,7 @@ class SnippetEnvironmentVariableCollectionApi(Resource): "name": v.name, "description": v.description, "selector": v.selector, - "value_type": v.value_type.exposed_type().value, + "value_type": environment_variable_value_type(v), "value": v.value, "edited": False, "visible": True, diff --git a/api/controllers/console/spec.py b/api/controllers/console/spec.py index 70e0d1d14ae..e2140744910 100644 --- a/api/controllers/console/spec.py +++ b/api/controllers/console/spec.py @@ -1,23 +1,18 @@ -import logging from collections.abc import Mapping +from http import HTTPStatus from typing import Any from flask_restx import Resource from pydantic import Field, RootModel from controllers.common.schema import register_response_schema_models -from controllers.console.wraps import ( - account_initialization_required, - setup_required, -) -from core.schemas.schema_manager import SchemaManager +from controllers.console.flask_admission import console_account_admission +from extensions.ext_application_services import application_services from fields.base import ResponseModel -from libs.login import login_required +from machinery.context import RequestContext from . import console_ns -logger = logging.getLogger(__name__) - class SchemaDefinitionItemResponse(ResponseModel): name: str @@ -34,20 +29,13 @@ register_response_schema_models(console_ns, SchemaDefinitionItemResponse, Schema @console_ns.route("/spec/schema-definitions") class SpecSchemaDefinitionsApi(Resource): - @console_ns.response(200, "Success", console_ns.models[SchemaDefinitionsResponse.__name__]) - @setup_required - @login_required - @account_initialization_required - def get(self): + @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[SchemaDefinitionsResponse.__name__]) + @console_account_admission() + def get(self, _request_context: RequestContext): """ Get system JSON Schema definitions specification Used for frontend component type mapping """ - try: - schema_manager = SchemaManager() - schema_definitions = schema_manager.get_all_schema_definitions() - return schema_definitions, 200 - except Exception: - logger.exception("Failed to get schema definitions from local registry") - # Return empty array as fallback - return [], 200 + schema_definitions = application_services().schema_definitions.list() + response = SchemaDefinitionsResponse.model_validate(schema_definitions).model_dump(mode="json") + return response, HTTPStatus.OK diff --git a/api/controllers/console/version.py b/api/controllers/console/system.py similarity index 68% rename from api/controllers/console/version.py rename to api/controllers/console/system.py index fdb23acf52a..9b1d70a0709 100644 --- a/api/controllers/console/version.py +++ b/api/controllers/console/system.py @@ -10,21 +10,27 @@ from controllers.fastopenapi import console_router logger = logging.getLogger(__name__) +class PingResponse(BaseModel): + result: str = Field(description="Health check result", examples=["pong"]) + + class VersionQuery(BaseModel): current_version: str = Field(..., description="Current application version") -class VersionFeatures(BaseModel): - can_replace_logo: bool = Field(description="Whether logo replacement is supported") - model_load_balancing_enabled: bool = Field(description="Whether model load balancing is enabled") - - class VersionResponse(BaseModel): version: str = Field(description="Latest version number") - release_date: str = Field(description="Release date of latest version") release_notes: str = Field(description="Release notes for latest version") - can_auto_update: bool = Field(description="Whether auto-update is supported") - features: VersionFeatures = Field(description="Feature flags and capabilities") + + +@console_router.get( + "/ping", + response_model=PingResponse, + tags=["console"], +) +def ping() -> PingResponse: + """Health check endpoint for connection testing.""" + return PingResponse(result="pong") @console_router.get( @@ -38,13 +44,7 @@ def check_version_update(query: VersionQuery) -> VersionResponse: result = VersionResponse( version=dify_config.project.version, - release_date="", release_notes="", - can_auto_update=False, - features=VersionFeatures( - can_replace_logo=dify_config.CAN_REPLACE_LOGO, - model_load_balancing_enabled=dify_config.MODEL_LB_ENABLED, - ), ) if not check_update_url: @@ -62,11 +62,9 @@ def check_version_update(query: VersionQuery) -> VersionResponse: result.version = query.current_version return result latest_version = content.get("version", result.version) - if _has_new_version(latest_version=latest_version, current_version=f"{query.current_version}"): + if _has_new_version(latest_version=latest_version, current_version=query.current_version): result.version = latest_version - result.release_date = content.get("releaseDate", "") result.release_notes = content.get("releaseNotes", "") - result.can_auto_update = content.get("canAutoUpdate", False) return result @@ -75,7 +73,6 @@ def _has_new_version(*, latest_version: str, current_version: str) -> bool: latest = version.parse(latest_version) current = version.parse(current_version) - # Compare versions return latest > current except version.InvalidVersion: logger.warning("Invalid version format: latest=%s, current=%s", latest_version, current_version) diff --git a/api/controllers/console/tag/tags.py b/api/controllers/console/tag/tags.py index 86c1ad9c54c..0084e614c5b 100644 --- a/api/controllers/console/tag/tags.py +++ b/api/controllers/console/tag/tags.py @@ -1,7 +1,6 @@ from typing import Literal from uuid import UUID -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, RootModel, field_validator from sqlalchemy import select @@ -17,6 +16,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, setup_required, with_current_tenant_id, with_current_user, @@ -134,10 +134,9 @@ class TagListApi(Resource): @console_ns.doc(params=query_params_from_model(TagListQueryParam)) @console_ns.response(200, "Success", console_ns.models[TagListResponse.__name__]) @with_current_tenant_id - def get(self, current_tenant_id: str): - raw_args = request.args.to_dict() - param = TagListQueryParam.model_validate(raw_args) - tags = TagService.get_tags(param.type, current_tenant_id, param.keyword, session=db.session()) + @model_validate(TagListQueryParam) + def get(self, req_data: TagListQueryParam, current_tenant_id: str): + tags = TagService.get_tags(req_data.type, current_tenant_id, req_data.keyword, session=db.session()) return dump_response(TagListResponse, tags), 200 @@ -147,14 +146,14 @@ class TagListApi(Resource): @login_required @account_initialization_required @with_current_user - def post(self, current_user: Account): + @model_validate(TagBasePayload) + def post(self, req_data: TagBasePayload, current_user: Account): # Allow users with edit permission, or dataset editors (including dataset operators). if not (current_user.has_edit_permission or current_user.is_dataset_editor): raise Forbidden() - payload = TagBasePayload.model_validate(console_ns.payload or {}) - _enforce_snippet_tag_rbac_if_needed(payload.type) - tag = TagService.save_tags(SaveTagPayload(name=payload.name, type=payload.type), db.session()) + _enforce_snippet_tag_rbac_if_needed(req_data.type) + tag = TagService.save_tags(SaveTagPayload(name=req_data.name, type=req_data.type), db.session()) return dump_response(TagResponse, {"id": tag.id, "name": tag.name, "type": tag.type, "binding_count": 0}), 200 @@ -167,15 +166,15 @@ class TagUpdateDeleteApi(Resource): @login_required @account_initialization_required @with_current_user - def patch(self, current_user: Account, tag_id: UUID): + @model_validate(TagUpdateRequestPayload) + def patch(self, req_data: TagUpdateRequestPayload, current_user: Account, tag_id: UUID): tag_id_str = str(tag_id) # The role of the current user in the ta table must be admin, owner, or editor if not (current_user.has_edit_permission or current_user.is_dataset_editor): raise Forbidden() - payload = TagUpdateRequestPayload.model_validate(console_ns.payload or {}) _enforce_snippet_tag_rbac_by_tag_id(tag_id_str) - tag = TagService.update_tags(UpdateTagPayload(name=payload.name), tag_id_str, db.session()) + tag = TagService.update_tags(UpdateTagPayload(name=req_data.name), tag_id_str, db.session()) binding_count = TagService.get_tag_binding_count(tag_id_str, db.session()) @@ -212,10 +211,9 @@ def _require_tag_binding_edit_permission(current_user: Account) -> None: raise Forbidden() -def _create_tag_bindings(current_user: Account) -> tuple[dict[str, str], int]: +def _create_tag_bindings(current_user: Account, payload: TagBindingPayload) -> tuple[dict[str, str], int]: _require_tag_binding_edit_permission(current_user) - payload = TagBindingPayload.model_validate(console_ns.payload or {}) _enforce_snippet_tag_rbac_if_needed(payload.type) TagService.save_tag_binding( TagBindingCreatePayload( @@ -228,10 +226,9 @@ def _create_tag_bindings(current_user: Account) -> tuple[dict[str, str], int]: return {"result": "success"}, 200 -def _remove_tag_bindings(current_user: Account) -> tuple[dict[str, str], int]: +def _remove_tag_bindings(current_user: Account, payload: TagBindingRemovePayload) -> tuple[dict[str, str], int]: _require_tag_binding_edit_permission(current_user) - payload = TagBindingRemovePayload.model_validate(console_ns.payload or {}) _enforce_snippet_tag_rbac_if_needed(payload.type) TagService.delete_tag_binding( TagBindingDeletePayload( @@ -255,8 +252,9 @@ class TagBindingCollectionApi(Resource): @login_required @account_initialization_required @with_current_user - def post(self, current_user: Account): - return _create_tag_bindings(current_user) + @model_validate(TagBindingPayload) + def post(self, req_data: TagBindingPayload, current_user: Account): + return _create_tag_bindings(current_user, req_data) @console_ns.route("/tag-bindings/remove") @@ -271,5 +269,6 @@ class TagBindingRemoveApi(Resource): @login_required @account_initialization_required @with_current_user - def post(self, current_user: Account): - return _remove_tag_bindings(current_user) + @model_validate(TagBindingRemovePayload) + def post(self, req_data: TagBindingRemovePayload, current_user: Account): + return _remove_tag_bindings(current_user, req_data) diff --git a/api/controllers/console/workflow_run_archive.py b/api/controllers/console/workflow_run_archive.py index e15b897d5b5..170d5268a6f 100644 --- a/api/controllers/console/workflow_run_archive.py +++ b/api/controllers/console/workflow_run_archive.py @@ -11,8 +11,8 @@ from controllers.common.schema import register_response_schema_models, register_ from controllers.console import console_ns from controllers.console.wraps import ( account_initialization_required, - cloud_edition_billing_enabled, cloud_edition_billing_paid_plan_required, + model_validate, only_edition_cloud, setup_required, ) @@ -121,7 +121,6 @@ class WorkflowRunArchivesApi(Resource): @login_required @account_initialization_required @only_edition_cloud - @cloud_edition_billing_enabled @cloud_edition_billing_paid_plan_required def get(self): tenant_id, _ = _current_owner_or_admin_ids() @@ -142,18 +141,17 @@ class WorkflowRunArchiveDownloadsApi(Resource): @login_required @account_initialization_required @only_edition_cloud - @cloud_edition_billing_enabled @cloud_edition_billing_paid_plan_required - def post(self): + @model_validate(WorkflowRunArchiveDownloadPayload) + def post(self, req_data: WorkflowRunArchiveDownloadPayload): tenant_id, account_id = _current_owner_or_admin_ids() - payload = WorkflowRunArchiveDownloadPayload.model_validate(console_ns.payload or {}) try: task = create_workflow_run_archive_download_task( db.session(), tenant_id=tenant_id, requested_by=account_id, - year=payload.year, - month=payload.month, + year=req_data.year, + month=req_data.month, ) except WorkflowRunArchiveNotFoundError as exc: raise NotFound(str(exc)) from exc @@ -169,7 +167,6 @@ class WorkflowRunArchiveDownloadApi(Resource): @login_required @account_initialization_required @only_edition_cloud - @cloud_edition_billing_enabled @cloud_edition_billing_paid_plan_required def get(self, download_id: str): tenant_id, _ = _current_owner_or_admin_ids() @@ -194,7 +191,6 @@ class WorkflowRunArchiveDownloadFileApi(Resource): @login_required @account_initialization_required @only_edition_cloud - @cloud_edition_billing_enabled @cloud_edition_billing_paid_plan_required def get(self, download_id: str): tenant_id, _ = _current_owner_or_admin_ids() diff --git a/api/controllers/console/workspace/account.py b/api/controllers/console/workspace/account.py index a37476953c4..548243e43a9 100644 --- a/api/controllers/console/workspace/account.py +++ b/api/controllers/console/workspace/account.py @@ -28,7 +28,12 @@ from controllers.console.auth.error import ( InvalidEmailError, InvalidTokenError, ) -from controllers.console.error import AccountInFreezeError, AccountNotFound, EmailSendIpLimitError +from controllers.console.error import ( + AccountInFreezeError, + AccountNotFound, + EducationDiscountTemporarilyPausedError, + EmailSendIpLimitError, +) from controllers.console.workspace.error import ( AccountAlreadyInitedError, CurrentPasswordIncorrectError, @@ -38,15 +43,14 @@ from controllers.console.workspace.error import ( ) from controllers.console.wraps import ( account_initialization_required, - cloud_edition_billing_enabled, enable_change_email, enterprise_license_required, + model_validate, only_edition_cloud, setup_required, - with_current_tenant_id, with_current_user, ) -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from extensions.ext_database import db from fields.base import ResponseModel from fields.member_fields import AccountResponse @@ -333,10 +337,9 @@ class AccountAvatarApi(Resource): @login_required @account_initialization_required @with_current_user - @with_current_tenant_id - def get(self, current_tenant_id: str, current_user: Account): - args = AccountAvatarQuery.model_validate(request.args.to_dict(flat=True)) - avatar = args.avatar + @model_validate(AccountAvatarQuery) + def get(self, req_data: AccountAvatarQuery, current_user: Account): + avatar = req_data.avatar if avatar.startswith(("http://", "https://")): return AvatarUrlResponse(avatar_url=avatar).model_dump(mode="json") @@ -345,9 +348,6 @@ class AccountAvatarApi(Resource): if upload_file is None: raise NotFound("Avatar file not found") - if upload_file.tenant_id != current_tenant_id: - raise NotFound("Avatar file not found") - if upload_file.created_by_role != CreatorUserRole.ACCOUNT or upload_file.created_by != current_user.id: raise NotFound("Avatar file not found") @@ -540,7 +540,6 @@ class EducationVerifyApi(Resource): @login_required @account_initialization_required @only_edition_cloud - @cloud_edition_billing_enabled @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[EducationVerifyResponse.__name__]) @with_current_user def get(self, account: Account): @@ -558,20 +557,14 @@ class EducationApi(Resource): @login_required @account_initialization_required @only_edition_cloud - @cloud_edition_billing_enabled @with_current_user def post(self, account: Account): - payload = console_ns.payload or {} - args = EducationActivatePayload.model_validate(payload) - - result = BillingService.EducationIdentity.activate(account, args.token, args.institution, args.role) - return result + raise EducationDiscountTemporarilyPausedError() @setup_required @login_required @account_initialization_required @only_edition_cloud - @cloud_edition_billing_enabled @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[EducationStatusResponse.__name__]) @with_current_user def get(self, account: Account): @@ -589,7 +582,6 @@ class EducationAutoCompleteApi(Resource): @login_required @account_initialization_required @only_edition_cloud - @cloud_edition_billing_enabled @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[EducationAutocompleteResponse.__name__]) def get(self): payload = request.args.to_dict(flat=True) diff --git a/api/controllers/console/workspace/endpoint.py b/api/controllers/console/workspace/endpoint.py index 781b4b6c1a2..f02ae4a73fb 100644 --- a/api/controllers/console/workspace/endpoint.py +++ b/api/controllers/console/workspace/endpoint.py @@ -11,7 +11,6 @@ from enum import StrEnum from http import HTTPStatus from typing import Any -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field @@ -23,6 +22,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -166,43 +166,38 @@ register_response_schema_models( ) -def _create_endpoint(tenant_id: str, user_id: str) -> bool: +def _create_endpoint(tenant_id: str, user_id: str, req_data: EndpointCreatePayload) -> bool: """Create a plugin endpoint for the injected workspace and user.""" - args = EndpointCreatePayload.model_validate(console_ns.payload) - try: return EndpointService.create_endpoint( tenant_id=tenant_id, user_id=user_id, - plugin_unique_identifier=args.plugin_unique_identifier, - name=args.name, - settings=args.settings, + plugin_unique_identifier=req_data.plugin_unique_identifier, + name=req_data.name, + settings=req_data.settings, ) except PluginPermissionDeniedError as e: raise ValueError(e.description) from e -def _update_endpoint(tenant_id: str, user_id: str, endpoint_id: str) -> bool: +def _update_endpoint(tenant_id: str, user_id: str, endpoint_id: str, req_data: EndpointUpdatePayload) -> bool: """Update a plugin endpoint identified by the canonical path parameter.""" - args = EndpointUpdatePayload.model_validate(console_ns.payload) - return EndpointService.update_endpoint( tenant_id=tenant_id, user_id=user_id, endpoint_id=endpoint_id, - name=args.name, - settings=args.settings, + name=req_data.name, + settings=req_data.settings, ) -def _legacy_update_endpoint(tenant_id: str, user_id: str) -> bool: - args = LegacyEndpointUpdatePayload.model_validate(console_ns.payload) +def _legacy_update_endpoint(tenant_id: str, user_id: str, req_data: LegacyEndpointUpdatePayload) -> bool: return EndpointService.update_endpoint( tenant_id=tenant_id, user_id=user_id, - endpoint_id=args.endpoint_id, - name=args.name, - settings=args.settings, + endpoint_id=req_data.endpoint_id, + name=req_data.name, + settings=req_data.settings, ) @@ -215,15 +210,13 @@ def _delete_endpoint(tenant_id: str, user_id: str, endpoint_id: str) -> bool: ) -def _delete_endpoint_from_payload(tenant_id: str, user_id: str) -> bool: - args = EndpointIdPayload.model_validate(console_ns.payload) - return _delete_endpoint(tenant_id=tenant_id, user_id=user_id, endpoint_id=args.endpoint_id) +def _delete_endpoint_from_payload(tenant_id: str, user_id: str, req_data: EndpointIdPayload) -> bool: + return _delete_endpoint(tenant_id=tenant_id, user_id=user_id, endpoint_id=req_data.endpoint_id) -def _set_endpoint_enabled(tenant_id: str, user_id: str, *, enabled: bool) -> bool: - args = EndpointIdPayload.model_validate(console_ns.payload) +def _set_endpoint_enabled(tenant_id: str, user_id: str, req_data: EndpointIdPayload, *, enabled: bool) -> bool: action = EndpointService.enable_endpoint if enabled else EndpointService.disable_endpoint - return action(tenant_id=tenant_id, user_id=user_id, endpoint_id=args.endpoint_id) + return action(tenant_id=tenant_id, user_id=user_id, endpoint_id=req_data.endpoint_id) @console_ns.route("/workspaces/current/endpoints") @@ -246,8 +239,11 @@ class EndpointCollectionApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, user_id: str): - return SuccessResponse(success=_create_endpoint(tenant_id=tenant_id, user_id=user_id)).model_dump(mode="json") + @model_validate(EndpointCreatePayload) + def post(self, req_data: EndpointCreatePayload, tenant_id: str, user_id: str): + return SuccessResponse( + success=_create_endpoint(tenant_id=tenant_id, user_id=user_id, req_data=req_data) + ).model_dump(mode="json") @console_ns.route("/workspaces/current/endpoints/create") @@ -275,8 +271,11 @@ class DeprecatedEndpointCreateApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, user_id: str): - return SuccessResponse(success=_create_endpoint(tenant_id=tenant_id, user_id=user_id)).model_dump(mode="json") + @model_validate(EndpointCreatePayload) + def post(self, req_data: EndpointCreatePayload, tenant_id: str, user_id: str): + return SuccessResponse( + success=_create_endpoint(tenant_id=tenant_id, user_id=user_id, req_data=req_data) + ).model_dump(mode="json") @console_ns.route("/workspaces/current/endpoints/list") @@ -294,14 +293,14 @@ class EndpointListApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def get(self, tenant_id: str, user_id: str): - args = EndpointListQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(EndpointListQuery) + def get(self, req_data: EndpointListQuery, tenant_id: str, user_id: str): endpoints = EndpointService.list_endpoints( tenant_id=tenant_id, user_id=user_id, - page=args.page, - page_size=args.page_size, + page=req_data.page, + page_size=req_data.page_size, ) return EndpointListResponse(endpoints=endpoints).model_dump(mode="json") @@ -322,15 +321,15 @@ class EndpointListForSinglePluginApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def get(self, tenant_id: str, user_id: str): - args = EndpointListForPluginQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(EndpointListForPluginQuery) + def get(self, req_data: EndpointListForPluginQuery, tenant_id: str, user_id: str): endpoints = EndpointService.list_endpoints_for_single_plugin( tenant_id=tenant_id, user_id=user_id, - plugin_id=args.plugin_id, - page=args.page, - page_size=args.page_size, + plugin_id=req_data.plugin_id, + page=req_data.page, + page_size=req_data.page_size, ) return EndpointListResponse(endpoints=endpoints).model_dump(mode="json") @@ -378,9 +377,10 @@ class EndpointItemApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def patch(self, tenant_id: str, user_id: str, id: str): + @model_validate(EndpointUpdatePayload) + def patch(self, req_data: EndpointUpdatePayload, tenant_id: str, user_id: str, id: str): return SuccessResponse( - success=_update_endpoint(tenant_id=tenant_id, user_id=user_id, endpoint_id=id) + success=_update_endpoint(tenant_id=tenant_id, user_id=user_id, endpoint_id=id, req_data=req_data) ).model_dump(mode="json") @@ -410,10 +410,11 @@ class DeprecatedEndpointDeleteApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, user_id: str): - return SuccessResponse(success=_delete_endpoint_from_payload(tenant_id=tenant_id, user_id=user_id)).model_dump( - mode="json" - ) + @model_validate(EndpointIdPayload) + def post(self, req_data: EndpointIdPayload, tenant_id: str, user_id: str): + return SuccessResponse( + success=_delete_endpoint_from_payload(tenant_id=tenant_id, user_id=user_id, req_data=req_data) + ).model_dump(mode="json") @console_ns.route("/workspaces/current/endpoints/update") @@ -442,10 +443,11 @@ class DeprecatedEndpointUpdateApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, user_id: str): - return SuccessResponse(success=_legacy_update_endpoint(tenant_id=tenant_id, user_id=user_id)).model_dump( - mode="json" - ) + @model_validate(LegacyEndpointUpdatePayload) + def post(self, req_data: LegacyEndpointUpdatePayload, tenant_id: str, user_id: str): + return SuccessResponse( + success=_legacy_update_endpoint(tenant_id=tenant_id, user_id=user_id, req_data=req_data) + ).model_dump(mode="json") @console_ns.route("/workspaces/current/endpoints/enable") @@ -466,9 +468,10 @@ class EndpointEnableApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, user_id: str): + @model_validate(EndpointIdPayload) + def post(self, req_data: EndpointIdPayload, tenant_id: str, user_id: str): return SuccessResponse( - success=_set_endpoint_enabled(tenant_id=tenant_id, user_id=user_id, enabled=True) + success=_set_endpoint_enabled(tenant_id=tenant_id, user_id=user_id, req_data=req_data, enabled=True) ).model_dump(mode="json") @@ -490,7 +493,8 @@ class EndpointDisableApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def post(self, tenant_id: str, user_id: str): + @model_validate(EndpointIdPayload) + def post(self, req_data: EndpointIdPayload, tenant_id: str, user_id: str): return SuccessResponse( - success=_set_endpoint_enabled(tenant_id=tenant_id, user_id=user_id, enabled=False) + success=_set_endpoint_enabled(tenant_id=tenant_id, user_id=user_id, req_data=req_data, enabled=False) ).model_dump(mode="json") diff --git a/api/controllers/console/workspace/error.py b/api/controllers/console/workspace/error.py index 567009a15cb..a19528eb9a8 100644 --- a/api/controllers/console/workspace/error.py +++ b/api/controllers/console/workspace/error.py @@ -1,6 +1,12 @@ from libs.exception import BaseHTTPException +class CurrentWorkspaceArchivedError(BaseHTTPException): + error_code = "current_workspace_archived" + description = "The current workspace has been archived." + code = 409 + + class RepeatPasswordNotMatchError(BaseHTTPException): error_code = "repeat_password_not_match" description = "New password and repeat password does not match." diff --git a/api/controllers/console/workspace/load_balancing_config.py b/api/controllers/console/workspace/load_balancing_config.py index abeb691be03..2960db6c067 100644 --- a/api/controllers/console/workspace/load_balancing_config.py +++ b/api/controllers/console/workspace/load_balancing_config.py @@ -6,6 +6,7 @@ from controllers.common.schema import register_response_schema_models, register_ from controllers.console import console_ns from controllers.console.wraps import ( account_initialization_required, + model_validate, setup_required, with_current_tenant_id, with_current_user, @@ -49,14 +50,15 @@ class LoadBalancingCredentialsValidateApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account, provider: str): + @model_validate(LoadBalancingCredentialPayload) + def post( + self, req_data: LoadBalancingCredentialPayload, current_tenant_id: str, current_user: Account, provider: str + ): if not TenantAccountRole.is_privileged_role(current_user.current_role): raise Forbidden() tenant_id = current_tenant_id - payload = LoadBalancingCredentialPayload.model_validate(console_ns.payload or {}) - # validate model load balancing credentials model_load_balancing_service = ModelLoadBalancingService() @@ -67,9 +69,9 @@ class LoadBalancingCredentialsValidateApi(Resource): model_load_balancing_service.validate_load_balancing_credentials( tenant_id=tenant_id, provider=provider, - model=payload.model, - model_type=payload.model_type, - credentials=payload.credentials, + model=req_data.model, + model_type=req_data.model_type, + credentials=req_data.credentials, session=db.session(), ) except CredentialsValidateFailedError as ex: @@ -99,14 +101,20 @@ class LoadBalancingConfigCredentialsValidateApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account, provider: str, config_id: str): + @model_validate(LoadBalancingCredentialPayload) + def post( + self, + req_data: LoadBalancingCredentialPayload, + current_tenant_id: str, + current_user: Account, + provider: str, + config_id: str, + ): if not TenantAccountRole.is_privileged_role(current_user.current_role): raise Forbidden() tenant_id = current_tenant_id - payload = LoadBalancingCredentialPayload.model_validate(console_ns.payload or {}) - # validate model load balancing config credentials model_load_balancing_service = ModelLoadBalancingService() @@ -117,9 +125,9 @@ class LoadBalancingConfigCredentialsValidateApi(Resource): model_load_balancing_service.validate_load_balancing_credentials( tenant_id=tenant_id, provider=provider, - model=payload.model, - model_type=payload.model_type, - credentials=payload.credentials, + model=req_data.model, + model_type=req_data.model_type, + credentials=req_data.credentials, session=db.session(), config_id=config_id, ) diff --git a/api/controllers/console/workspace/members.py b/api/controllers/console/workspace/members.py index 9d40d32aaa2..ebc30c9fa59 100644 --- a/api/controllers/console/workspace/members.py +++ b/api/controllers/console/workspace/members.py @@ -24,6 +24,7 @@ from controllers.console.auth.error import ( OwnerTransferLimitError, ) from controllers.console.error import EmailSendIpLimitError, SeatsLimitExceeded, WorkspaceMembersLimitExceeded +from controllers.console.flask_admission import console_account_admission from controllers.console.workspace.error import InvalidMemberRoleError from controllers.console.wraps import ( account_initialization_required, @@ -31,15 +32,17 @@ from controllers.console.wraps import ( setup_required, with_current_user, ) +from enums import DeploymentEdition +from extensions.ext_application_services import application_services from extensions.ext_database import db from extensions.ext_redis import redis_client from fields.base import ResponseModel from fields.member_fields import AccountWithRoleListResponse, AccountWithRoleResponse from libs.helper import dump_response, extract_remote_ip -from libs.login import current_account_with_tenant, login_required +from libs.login import login_required +from machinery.context import RequestContext from models.account import Account, TenantAccountJoin, TenantAccountRole from services.account_service import AccountService, RegisterService, TenantService -from services.enterprise import rbac_service as enterprise_rbac_service from services.errors.account import AccountAlreadyInTenantError from services.feature_service import FeatureService @@ -144,22 +147,6 @@ def _is_role_enabled(role: TenantAccountRole | str, tenant_id: str) -> bool: return FeatureService.get_features(tenant_id=tenant_id, exclude_vector_space=True).dataset_operator_enabled -def _serialize_member_roles( - current_role: str | None, member_roles: list[enterprise_rbac_service.RBACRole] -) -> list[dict[str, str]]: - if dify_config.RBAC_ENABLED: - return [{"id": role.id, "name": role.name} for role in member_roles] - else: - if current_role: - return [{"id": current_role, "name": current_role}] - return [] - - -def _normalize_enum_value(value: object) -> str: - normalized = getattr(value, "value", value) - return str(normalized) if normalized is not None else "" - - def _count_new_member_invites(tenant_id: str, emails: list[str]) -> tuple[int, int]: new_member_count = 0 new_account_count = 0 @@ -193,7 +180,7 @@ def _check_member_invite_limits(tenant_id: str, new_member_count: int, new_accou features = FeatureService.get_features(tenant_id=tenant_id, exclude_vector_space=True) - if dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE: workspace_members = features.workspace_members if workspace_members.enabled is True and not workspace_members.is_available(new_member_count): raise WorkspaceMembersLimitExceeded() @@ -203,7 +190,7 @@ def _check_member_invite_limits(tenant_id: str, new_member_count: int, new_accou raise SeatsLimitExceeded() return - if dify_config.BILLING_ENABLED and features.billing.enabled is True: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: members = features.members current_member_count = _count_current_members(tenant_id) if 0 < members.limit < current_member_count + new_member_count: @@ -214,46 +201,25 @@ def _check_member_invite_limits(tenant_id: str, new_member_count: int, new_accou class MemberListApi(Resource): """List all members of current tenant.""" - @setup_required - @login_required - @account_initialization_required @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[AccountWithRoleListResponse.__name__]) - @with_current_user - def get(self, current_user: Account | None = None): - if current_user is None: - current_user, _ = current_account_with_tenant() - if not current_user.current_tenant: - raise ValueError("No current tenant") - members = TenantService.get_tenant_members(current_user.current_tenant, session=db.session()) - if dify_config.RBAC_ENABLED: - member_ids = [member.id for member in members] - member_roles = enterprise_rbac_service.RBACService.MemberRoles.batch_get( - str(current_user.current_tenant.id), - current_user.id, - member_ids, - ) - roles_map = {item.account_id: item.roles for item in member_roles} - else: - roles_map = {} - - serialized_members = [] - for member in members: - current_role = _normalize_enum_value(member.current_role) - serialized_members.append( - { - "id": member.id, - "name": member.name, - "email": member.email, - "avatar": member.avatar, - "last_login_at": member.last_login_at, - "last_active_at": member.last_active_at, - "created_at": member.created_at, - "role": current_role, - "roles": _serialize_member_roles(current_role, roles_map.get(member.id, [])), - "status": _normalize_enum_value(member.status), - } - ) - + @console_account_admission() + def get(self, request_context: RequestContext): + members = application_services().workspace_member_queries.list_current(request_context) + serialized_members = [ + { + "id": member.id, + "name": member.name, + "email": member.email, + "avatar": member.avatar, + "last_login_at": member.last_login_at, + "last_active_at": member.last_active_at, + "created_at": member.created_at, + "role": member.role, + "roles": [{"id": role.id, "name": role.name} for role in member.roles], + "status": member.status, + } + for member in members + ] return dump_response(AccountWithRoleListResponse, {"accounts": serialized_members}), HTTPStatus.OK @@ -300,7 +266,10 @@ class MemberInviteEmailApi(Resource): tenant_id = inviter.current_tenant.id with redis_client.lock(f"workspace_member_invite:{tenant_id}", timeout=60): - if dify_config.ENTERPRISE_ENABLED is True or dify_config.BILLING_ENABLED is True: + if dify_config.DEPLOYMENT_EDITION in { + DeploymentEdition.CLOUD, + DeploymentEdition.ENTERPRISE, + }: new_member_count, new_account_count = _count_new_member_invites(tenant_id, invitee_emails) _check_member_invite_limits(tenant_id, new_member_count, new_account_count) diff --git a/api/controllers/console/workspace/model_providers.py b/api/controllers/console/workspace/model_providers.py index adebe1a4a91..5ff114aa4e0 100644 --- a/api/controllers/console/workspace/model_providers.py +++ b/api/controllers/console/workspace/model_providers.py @@ -4,9 +4,11 @@ from typing import Any, Literal from flask import request, send_file from flask_restx import Resource from pydantic import BaseModel, Field, field_validator +from sqlalchemy.orm import Session from controllers.common.fields import SimpleResultResponse, ValidationResultResponse from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models +from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.wraps import ( RBACPermission, @@ -26,8 +28,13 @@ from libs.helper import dump_response, uuid_value from libs.login import login_required from models import Account from services.billing_service import BillingService -from services.entities.model_provider_entities import ProviderResponse +from services.entities.model_provider_entities import ( + ModelProviderPluginSummaryResponse, + ModelProviderSummaryResponse, + ProviderResponse, +) from services.model_provider_service import ModelProviderService +from services.workspace_service import WorkspaceService class ParserModelList(BaseModel): @@ -91,6 +98,22 @@ class ModelProviderListResponse(ResponseModel): data: list[ProviderResponse] +class ModelProviderSummaryListResponse(ResponseModel): + data: list[ModelProviderSummaryResponse] + plugins: dict[str, ModelProviderPluginSummaryResponse] + + +class ModelProviderCreditsResponse(ResponseModel): + pool_type: Literal["paid", "trial"] | None + quota_limit: int | None = Field(description="Credit limit for the effective pool; -1 means unlimited.") + quota_used: int | None + remaining_credits: int | None = Field(description="Remaining credits; -1 means unlimited.") + is_unlimited: bool + is_exhausted: bool + exhausted_at: int | None + next_credit_reset_date: int | None + + class ProviderCredentialsResponse(ResponseModel): credentials: dict[str, Any] | None = None @@ -114,6 +137,8 @@ register_response_schema_models( console_ns, SimpleResultResponse, ModelProviderListResponse, + ModelProviderSummaryListResponse, + ModelProviderCreditsResponse, ProviderCredentialsResponse, ValidationResultResponse, ModelProviderPaymentCheckoutUrlResponse, @@ -140,6 +165,40 @@ class ModelProviderListApi(Resource): return ModelProviderListResponse(data=provider_list).model_dump(mode="json") +@console_ns.route("/workspaces/current/model-providers/summary") +class ModelProviderSummaryListApi(Resource): + @console_ns.response( + 200, + "Model provider summaries retrieved successfully", + console_ns.models[ModelProviderSummaryListResponse.__name__], + ) + @setup_required + @login_required + @account_initialization_required + @with_current_tenant_id + def get(self, tenant_id: str): + providers, plugins = ModelProviderService().get_provider_summary_list(tenant_id=tenant_id) + return dump_response( + ModelProviderSummaryListResponse, + {"data": providers, "plugins": plugins}, + ) + + +@console_ns.route("/workspaces/current/model-providers/credits") +class ModelProviderCreditsApi(Resource): + @console_ns.response( + 200, "Model provider credits retrieved successfully", console_ns.models[ModelProviderCreditsResponse.__name__] + ) + @setup_required + @login_required + @account_initialization_required + @with_current_tenant_id + @with_session(write=False) + def get(self, session: Session, tenant_id: str): + credit_pool = WorkspaceService.get_effective_credit_pool(tenant_id, session=session) + return dump_response(ModelProviderCreditsResponse, credit_pool) + + @console_ns.route("/workspaces/current/model-providers//credentials") class ModelProviderCredentialApi(Resource): @console_ns.doc(params=query_params_from_model(ParserCredentialId)) diff --git a/api/controllers/console/workspace/models.py b/api/controllers/console/workspace/models.py index b3c7833a070..54d021251fd 100644 --- a/api/controllers/console/workspace/models.py +++ b/api/controllers/console/workspace/models.py @@ -1,7 +1,6 @@ import logging from typing import Any, cast -from flask import request from flask_restx import Resource from pydantic import BaseModel, Field, field_validator @@ -18,6 +17,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -212,12 +212,12 @@ class DefaultModelApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str): - args = ParserGetDefault.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserGetDefault) + def get(self, req_data: ParserGetDefault, tenant_id: str): model_provider_service = ModelProviderService() default_model_entity = model_provider_service.get_default_model_of_model_type( - tenant_id=tenant_id, model_type=args.model_type + tenant_id=tenant_id, model_type=req_data.model_type ) return DefaultModelDataResponse(data=default_model_entity).model_dump(mode="json") @@ -230,10 +230,10 @@ class DefaultModelApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserPostDefault.model_validate(console_ns.payload) + @model_validate(ParserPostDefault) + def post(self, req_data: ParserPostDefault, tenant_id: str): model_provider_service = ModelProviderService() - model_settings = args.model_settings + model_settings = req_data.model_settings for model_setting in model_settings: if model_setting.provider is None: continue @@ -279,43 +279,43 @@ class ModelProviderModelApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, provider: str): + @model_validate(ParserPostModels) + def post(self, req_data: ParserPostModels, tenant_id: str, provider: str): # To save the model's load balance configs - args = ParserPostModels.model_validate(console_ns.payload) - if args.config_from == "custom-model": - if not args.credential_id: + if req_data.config_from == "custom-model": + if not req_data.credential_id: raise ValueError("credential_id is required when configuring a custom-model") service = ModelProviderService() service.switch_active_custom_model_credential( tenant_id=tenant_id, provider=provider, - model_type=args.model_type, - model=args.model, - credential_id=args.credential_id, + model_type=req_data.model_type, + model=req_data.model, + credential_id=req_data.credential_id, ) model_load_balancing_service = ModelLoadBalancingService() - if args.load_balancing and args.load_balancing.configs: + if req_data.load_balancing and req_data.load_balancing.configs: # save load balancing configs model_load_balancing_service.update_load_balancing_configs( tenant_id=tenant_id, provider=provider, - model=args.model, - model_type=args.model_type, - configs=args.load_balancing.configs, - config_from=args.config_from or "", + model=req_data.model, + model_type=req_data.model_type, + configs=req_data.load_balancing.configs, + config_from=req_data.config_from or "", session=db.session(), ) - if args.load_balancing.enabled: + if req_data.load_balancing.enabled: model_load_balancing_service.enable_model_load_balancing( - tenant_id=tenant_id, provider=provider, model=args.model, model_type=args.model_type + tenant_id=tenant_id, provider=provider, model=req_data.model, model_type=req_data.model_type ) else: model_load_balancing_service.disable_model_load_balancing( - tenant_id=tenant_id, provider=provider, model=args.model, model_type=args.model_type + tenant_id=tenant_id, provider=provider, model=req_data.model, model_type=req_data.model_type ) return SimpleResultResponse(result="success").model_dump(mode="json"), 200 @@ -328,12 +328,12 @@ class ModelProviderModelApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @account_initialization_required @with_current_tenant_id - def delete(self, tenant_id: str, provider: str): - args = ParserDeleteModels.model_validate(console_ns.payload) + @model_validate(ParserDeleteModels) + def delete(self, req_data: ParserDeleteModels, tenant_id: str, provider: str): model_provider_service = ModelProviderService() model_provider_service.remove_model( - tenant_id=tenant_id, provider=provider, model=args.model, model_type=args.model_type + tenant_id=tenant_id, provider=provider, model=req_data.model, model_type=req_data.model_type ) return "", 204 @@ -352,29 +352,29 @@ class ModelProviderModelCredentialApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, user: Account, provider: str): - args = ParserGetCredentials.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserGetCredentials) + def get(self, req_data: ParserGetCredentials, tenant_id: str, user: Account, provider: str): model_provider_service = ModelProviderService() current_credential = model_provider_service.get_model_credential( tenant_id=tenant_id, provider=provider, - model_type=args.model_type, - model=args.model, - credential_id=args.credential_id, + model_type=req_data.model_type, + model=req_data.model, + credential_id=req_data.credential_id, ) model_load_balancing_service = ModelLoadBalancingService() is_load_balancing_enabled, load_balancing_configs = model_load_balancing_service.get_load_balancing_configs( tenant_id=tenant_id, provider=provider, - model=args.model, - model_type=args.model_type, + model=req_data.model, + model_type=req_data.model_type, session=db.session(), - config_from=args.config_from or "", + config_from=req_data.config_from or "", ) - if args.config_from == "predefined-model": + if req_data.config_from == "predefined-model": # Only the predefined-model branch needs visibility filtering by user. # The account is injected once by the handler and only passed into the # service branch that needs user-scoped credential visibility. @@ -387,8 +387,8 @@ class ModelProviderModelCredentialApi(Resource): available_credentials = model_provider_service.get_provider_model_available_credentials( tenant_id=tenant_id, provider=provider, - model_type=args.model_type, - model=args.model, + model_type=req_data.model_type, + model=req_data.model, ) credentials: dict[str, Any] = {} @@ -414,8 +414,8 @@ class ModelProviderModelCredentialApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_CREATE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, provider: str): - args = ParserCreateCredential.model_validate(console_ns.payload) + @model_validate(ParserCreateCredential) + def post(self, req_data: ParserCreateCredential, tenant_id: str, provider: str): model_provider_service = ModelProviderService() @@ -423,17 +423,17 @@ class ModelProviderModelCredentialApi(Resource): model_provider_service.create_model_credential( tenant_id=tenant_id, provider=provider, - model=args.model, - model_type=args.model_type, - credentials=args.credentials, - credential_name=args.name, + model=req_data.model, + model_type=req_data.model_type, + credentials=req_data.credentials, + credential_name=req_data.name, ) except CredentialsValidateFailedError as ex: logger.exception( "Failed to save model credentials, tenant_id: %s, model: %s, model_type: %s", tenant_id, - args.model, - args.model_type, + req_data.model, + req_data.model_type, ) raise ValueError(str(ex)) @@ -447,8 +447,8 @@ class ModelProviderModelCredentialApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def put(self, current_tenant_id: str, provider: str): - args = ParserUpdateCredential.model_validate(console_ns.payload) + @model_validate(ParserUpdateCredential) + def put(self, req_data: ParserUpdateCredential, current_tenant_id: str, provider: str): model_provider_service = ModelProviderService() @@ -456,11 +456,11 @@ class ModelProviderModelCredentialApi(Resource): model_provider_service.update_model_credential( tenant_id=current_tenant_id, provider=provider, - model_type=args.model_type, - model=args.model, - credentials=args.credentials, - credential_id=args.credential_id, - credential_name=args.name, + model_type=req_data.model_type, + model=req_data.model, + credentials=req_data.credentials, + credential_id=req_data.credential_id, + credential_name=req_data.name, ) except CredentialsValidateFailedError as ex: raise ValueError(str(ex)) @@ -475,16 +475,16 @@ class ModelProviderModelCredentialApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def delete(self, current_tenant_id: str, provider: str): - args = ParserDeleteCredential.model_validate(console_ns.payload) + @model_validate(ParserDeleteCredential) + def delete(self, req_data: ParserDeleteCredential, current_tenant_id: str, provider: str): model_provider_service = ModelProviderService() model_provider_service.remove_model_credential( tenant_id=current_tenant_id, provider=provider, - model_type=args.model_type, - model=args.model, - credential_id=args.credential_id, + model_type=req_data.model_type, + model=req_data.model, + credential_id=req_data.credential_id, ) return "", 204 @@ -500,16 +500,16 @@ class ModelProviderModelCredentialSwitchApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_USE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str, provider: str): - args = ParserSwitch.model_validate(console_ns.payload) + @model_validate(ParserSwitch) + def post(self, req_data: ParserSwitch, current_tenant_id: str, provider: str): service = ModelProviderService() service.add_model_credential_to_model_list( tenant_id=current_tenant_id, provider=provider, - model_type=args.model_type, - model=args.model, - credential_id=args.credential_id, + model_type=req_data.model_type, + model=req_data.model, + credential_id=req_data.credential_id, ) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -525,12 +525,12 @@ class ModelProviderModelEnableApi(Resource): @account_initialization_required @with_current_tenant_id @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) - def patch(self, tenant_id: str, provider: str): - args = ParserDeleteModels.model_validate(console_ns.payload) + @model_validate(ParserDeleteModels) + def patch(self, req_data: ParserDeleteModels, tenant_id: str, provider: str): model_provider_service = ModelProviderService() model_provider_service.enable_model( - tenant_id=tenant_id, provider=provider, model=args.model, model_type=args.model_type + tenant_id=tenant_id, provider=provider, model=req_data.model, model_type=req_data.model_type ) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -547,12 +547,12 @@ class ModelProviderModelDisableApi(Resource): @account_initialization_required @with_current_tenant_id @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) - def patch(self, tenant_id: str, provider: str): - args = ParserDeleteModels.model_validate(console_ns.payload) + @model_validate(ParserDeleteModels) + def patch(self, req_data: ParserDeleteModels, tenant_id: str, provider: str): model_provider_service = ModelProviderService() model_provider_service.disable_model( - tenant_id=tenant_id, provider=provider, model=args.model, model_type=args.model_type + tenant_id=tenant_id, provider=provider, model=req_data.model, model_type=req_data.model_type ) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -579,8 +579,8 @@ class ModelProviderModelValidateApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, provider: str): - args = ParserValidate.model_validate(console_ns.payload) + @model_validate(ParserValidate) + def post(self, req_data: ParserValidate, tenant_id: str, provider: str): model_provider_service = ModelProviderService() @@ -591,9 +591,9 @@ class ModelProviderModelValidateApi(Resource): model_provider_service.validate_model_credentials( tenant_id=tenant_id, provider=provider, - model=args.model, - model_type=args.model_type, - credentials=args.credentials, + model=req_data.model, + model_type=req_data.model_type, + credentials=req_data.credentials, ) except CredentialsValidateFailedError as ex: result = False @@ -617,12 +617,12 @@ class ModelProviderModelParameterRuleApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str, provider: str): - args = ParserParameter.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserParameter) + def get(self, req_data: ParserParameter, tenant_id: str, provider: str): model_provider_service = ModelProviderService() parameter_rules = model_provider_service.get_model_parameter_rules( - tenant_id=tenant_id, provider=provider, model=args.model + tenant_id=tenant_id, provider=provider, model=req_data.model ) return ModelParameterRuleListResponse(data=parameter_rules).model_dump(mode="json") diff --git a/api/controllers/console/workspace/plugin.py b/api/controllers/console/workspace/plugin.py index 98eb1192027..236d6ac1f24 100644 --- a/api/controllers/console/workspace/plugin.py +++ b/api/controllers/console/workspace/plugin.py @@ -1,5 +1,5 @@ import io -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from datetime import datetime from typing import Any, Literal, TypedDict @@ -24,6 +24,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -43,6 +44,7 @@ from core.plugin.entities.plugin_daemon import PluginDecodeResponse, PluginInsta from core.plugin.impl.exc import PluginDaemonClientSideError from core.plugin.plugin_service import PluginService from core.tools.builtin_tool.providers._positions import BuiltinToolProviderSort +from core.tools.entities.api_entities import ToolProviderApiEntity from core.tools.entities.common_entities import I18nObject from core.tools.entities.tool_entities import ToolProviderType from core.tools.tool_manager import ToolManager @@ -90,9 +92,21 @@ class ParserList(BaseModel): page_size: int = Field(default=256, ge=1, le=256, description="Page size (1-256)") +type PluginCategoryListLanguage = Literal["en_US", "zh_Hans", "ja_JP", "pt_BR"] + + class PluginCategoryListQuery(BaseModel): page: int = Field(default=1, ge=1, description="Page number") page_size: int = Field(default=256, ge=1, le=256, description="Page size (1-256)") + query: str = Field(default="", max_length=256, description="Case-insensitive search query") + tags: list[str] = Field(default_factory=list, max_length=128, description="Match any plugin tag") + language: Literal["en_US", "zh_Hans", "ja_JP", "pt_BR"] = Field( + default="en_US", description="Language used for localized label and description search" + ) + + +class PluginInstalledIdsQuery(BaseModel): + category: PluginCategory = Field(description="Plugin category to include") class ParserLatest(BaseModel): @@ -150,6 +164,7 @@ class ParserGithubUpgrade(BaseModel): class ParserUninstall(BaseModel): plugin_installation_id: str + preserve_credentials: bool = False class ParserPermissionChange(BaseModel): @@ -325,6 +340,10 @@ class PluginListResponse(ResponseModel): total: int +class PluginInstalledIdsResponse(ResponseModel): + plugin_ids: list[str] + + class PluginVersionsResponse(ResponseModel): versions: Mapping[str, PluginService.LatestPluginCache | None] @@ -384,6 +403,7 @@ register_schema_models( console_ns, ParserList, PluginCategoryListQuery, + PluginInstalledIdsQuery, PluginAutoUpgradeSettingsPayload, PluginPermissionSettingsPayload, ParserLatest, @@ -420,6 +440,7 @@ register_response_schema_models( PluginDebuggingKeyResponse, PluginDynamicOptionsResponse, PluginInstallationsResponse, + PluginInstalledIdsResponse, PluginInstallTaskStartResponse, PluginListResponse, PluginManifestResponse, @@ -477,7 +498,39 @@ def _read_upload_content(file: FileStorage, max_size: int) -> bytes: return content -def _list_hardcoded_builtin_tool_providers(tenant_id: str) -> list[dict[str, Any]]: +def _localized_builtin_tool_text(value: I18nObject, language: PluginCategoryListLanguage) -> str: + return value.to_dict()[language] or value.en_US + + +def _builtin_tool_provider_matches_filters( + provider: ToolProviderApiEntity, + *, + query: str, + tags: Sequence[str], + language: PluginCategoryListLanguage, +) -> bool: + if tags and not any(tag in provider.labels for tag in tags): + return False + if not query: + return True + + lower_query = query.lower() + candidates = ( + provider.name, + _localized_builtin_tool_text(provider.label, language), + _localized_builtin_tool_text(provider.description, language), + ) + return any(lower_query in candidate.lower() for candidate in candidates) + + +def _list_hardcoded_builtin_tool_providers( + tenant_id: str, + *, + query: str = "", + tags: Sequence[str] = (), + language: PluginCategoryListLanguage = "en_US", +) -> list[dict[str, Any]]: + """List builtin providers using the same search and tag semantics as category plugins.""" db_builtin_providers = { str(ToolProviderID(provider.provider)): provider for provider in ToolManager.list_default_builtin_providers(tenant_id) @@ -498,6 +551,13 @@ def _list_hardcoded_builtin_tool_providers(tenant_id: str) -> list[dict[str, Any db_provider=db_builtin_providers.get(provider.entity.identity.name), decrypt_credentials=False, ) + if not _builtin_tool_provider_matches_filters( + user_provider, + query=query, + tags=tags, + language=language, + ): + continue ToolTransformService.repack_provider(tenant_id=tenant_id, provider=user_provider) builtin_providers.append(user_provider) @@ -533,10 +593,10 @@ class PluginListApi(Resource): @account_initialization_required @with_current_user_id @with_current_tenant_id - def get(self, tenant_id: str, user_id: str): - args = ParserList.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserList) + def get(self, req_data: ParserList, tenant_id: str, user_id: str): try: - plugins_with_total = PluginService.list_with_total(tenant_id, user_id, args.page, args.page_size) + plugins_with_total = PluginService.list_with_total(tenant_id, user_id, req_data.page, req_data.page_size) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -552,7 +612,9 @@ class PluginCategoryListApi(Resource): @account_initialization_required @with_current_tenant_id def get(self, tenant_id: str, category: str): - args = PluginCategoryListQuery.model_validate(request.args.to_dict(flat=True)) + args = PluginCategoryListQuery.model_validate( + {**request.args.to_dict(flat=True), "tags": request.args.getlist("tags")} + ) try: plugin_category = PluginCategory(category) @@ -560,13 +622,26 @@ class PluginCategoryListApi(Resource): return {"code": "invalid_param", "message": "invalid plugin category"}, 400 try: - plugins = PluginService.list_by_category(tenant_id, plugin_category, args.page, args.page_size) + plugins = PluginService.list_by_category( + tenant_id, + plugin_category, + args.page, + args.page_size, + query=args.query, + tags=args.tags, + language=args.language, + ) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 builtin_tools = [] if plugin_category == PluginCategory.Tool: - builtin_tools = _list_hardcoded_builtin_tool_providers(tenant_id) + builtin_tools = _list_hardcoded_builtin_tool_providers( + tenant_id, + query=args.query, + tags=args.tags, + language=args.language, + ) return dump_response( PluginCategoryListResponse, @@ -578,6 +653,24 @@ class PluginCategoryListApi(Resource): ) +@console_ns.route("/workspaces/current/plugin/installed-ids") +class PluginInstalledIdsApi(Resource): + @console_ns.doc(params=query_params_from_model(PluginInstalledIdsQuery)) + @console_ns.response(200, "Success", console_ns.models[PluginInstalledIdsResponse.__name__]) + @setup_required + @login_required + @account_initialization_required + @with_current_tenant_id + @model_validate(PluginInstalledIdsQuery) + def get(self, req_data: PluginInstalledIdsQuery, tenant_id: str): + try: + plugin_ids = PluginService.list_installed_plugin_ids(tenant_id, req_data.category) + except PluginDaemonClientSideError as e: + return {"code": "plugin_error", "message": e.description}, 400 + + return dump_response(PluginInstalledIdsResponse, {"plugin_ids": plugin_ids}) + + @console_ns.route("/workspaces/current/plugin/list/latest-versions") class PluginListLatestVersionsApi(Resource): @console_ns.expect(console_ns.models[ParserLatest.__name__]) @@ -585,11 +678,11 @@ class PluginListLatestVersionsApi(Resource): @setup_required @login_required @account_initialization_required - def post(self): - args = ParserLatest.model_validate(console_ns.payload) + @model_validate(ParserLatest) + def post(self, req_data: ParserLatest): try: - versions = PluginService.list_latest_versions(args.plugin_ids) + versions = PluginService.list_latest_versions(req_data.plugin_ids) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -604,11 +697,11 @@ class PluginListInstallationsFromIdsApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserLatest.model_validate(console_ns.payload) + @model_validate(ParserLatest) + def post(self, req_data: ParserLatest, tenant_id: str): try: - plugins = PluginService.list_installations_from_ids(tenant_id, args.plugin_ids) + plugins = PluginService.list_installations_from_ids(tenant_id, req_data.plugin_ids) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -620,11 +713,11 @@ class PluginIconApi(Resource): @console_ns.doc(params=query_params_from_model(ParserIcon)) @console_ns.response(200, "Success", console_ns.models[BinaryFileResponse.__name__]) @setup_required - def get(self): - args = ParserIcon.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserIcon) + def get(self, req_data: ParserIcon): try: - icon_bytes, mimetype = PluginService.get_asset(args.tenant_id, args.filename) + icon_bytes, mimetype = PluginService.get_asset(req_data.tenant_id, req_data.filename) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -640,11 +733,11 @@ class PluginAssetApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str): - args = ParserAsset.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserAsset) + def get(self, req_data: ParserAsset, tenant_id: str): try: - binary = PluginService.extract_asset(tenant_id, args.plugin_unique_identifier, args.file_name) + binary = PluginService.extract_asset(tenant_id, req_data.plugin_unique_identifier, req_data.file_name) return send_file(io.BytesIO(binary), mimetype="application/octet-stream") except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -681,11 +774,13 @@ class PluginUploadFromGithubApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_INSTALL, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserGithubUpload.model_validate(console_ns.payload) + @model_validate(ParserGithubUpload) + def post(self, req_data: ParserGithubUpload, tenant_id: str): try: - response = PluginService.upload_pkg_from_github(tenant_id, args.repo, args.version, args.package) + response = PluginService.upload_pkg_from_github( + tenant_id, req_data.repo, req_data.version, req_data.package + ) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -722,11 +817,11 @@ class PluginInstallFromPkgApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_INSTALL, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserPluginIdentifiers.model_validate(console_ns.payload) + @model_validate(ParserPluginIdentifiers) + def post(self, req_data: ParserPluginIdentifiers, tenant_id: str): try: - response = PluginService.install_from_local_pkg(tenant_id, args.plugin_unique_identifiers) + response = PluginService.install_from_local_pkg(tenant_id, req_data.plugin_unique_identifiers) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -743,16 +838,16 @@ class PluginInstallFromGithubApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_INSTALL, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserGithubInstall.model_validate(console_ns.payload) + @model_validate(ParserGithubInstall) + def post(self, req_data: ParserGithubInstall, tenant_id: str): try: response = PluginService.install_from_github( tenant_id, - args.plugin_unique_identifier, - args.repo, - args.version, - args.package, + req_data.plugin_unique_identifier, + req_data.repo, + req_data.version, + req_data.package, ) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -770,11 +865,11 @@ class PluginInstallFromMarketplaceApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_INSTALL, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserPluginIdentifiers.model_validate(console_ns.payload) + @model_validate(ParserPluginIdentifiers) + def post(self, req_data: ParserPluginIdentifiers, tenant_id: str): try: - response = PluginService.install_from_marketplace_pkg(tenant_id, args.plugin_unique_identifiers) + response = PluginService.install_from_marketplace_pkg(tenant_id, req_data.plugin_unique_identifiers) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -791,15 +886,15 @@ class PluginFetchMarketplacePkgApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_INSTALL, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def get(self, tenant_id: str): - args = ParserPluginIdentifierQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserPluginIdentifierQuery) + def get(self, req_data: ParserPluginIdentifierQuery, tenant_id: str): try: return jsonable_encoder( { "manifest": PluginService.fetch_marketplace_pkg( tenant_id, - args.plugin_unique_identifier, + req_data.plugin_unique_identifier, ) } ) @@ -817,12 +912,16 @@ class PluginFetchManifestApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_INSTALL, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def get(self, tenant_id: str): - args = ParserPluginIdentifierQuery.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserPluginIdentifierQuery) + def get(self, req_data: ParserPluginIdentifierQuery, tenant_id: str): try: return jsonable_encoder( - {"manifest": PluginService.fetch_plugin_manifest(tenant_id, args.plugin_unique_identifier).model_dump()} + { + "manifest": PluginService.fetch_plugin_manifest( + tenant_id, req_data.plugin_unique_identifier + ).model_dump() + }, ) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -837,11 +936,13 @@ class PluginFetchInstallTasksApi(Resource): @account_initialization_required @plugin_permission_required(install_required=True) @with_current_tenant_id - def get(self, tenant_id: str): - args = ParserTasks.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserTasks) + def get(self, req_data: ParserTasks, tenant_id: str): try: - return jsonable_encoder({"tasks": PluginService.fetch_install_tasks(tenant_id, args.page, args.page_size)}) + return jsonable_encoder( + {"tasks": PluginService.fetch_install_tasks(tenant_id, req_data.page, req_data.page_size)} + ) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -916,13 +1017,13 @@ class PluginUpgradeFromMarketplaceApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_MODEL_CONFIG, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserMarketplaceUpgrade.model_validate(console_ns.payload) + @model_validate(ParserMarketplaceUpgrade) + def post(self, req_data: ParserMarketplaceUpgrade, tenant_id: str): try: return jsonable_encoder( PluginService.upgrade_plugin_with_marketplace( - tenant_id, args.original_plugin_unique_identifier, args.new_plugin_unique_identifier + tenant_id, req_data.original_plugin_unique_identifier, req_data.new_plugin_unique_identifier ) ) except PluginDaemonClientSideError as e: @@ -939,18 +1040,18 @@ class PluginUpgradeFromGithubApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_MODEL_CONFIG, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserGithubUpgrade.model_validate(console_ns.payload) + @model_validate(ParserGithubUpgrade) + def post(self, req_data: ParserGithubUpgrade, tenant_id: str): try: return jsonable_encoder( PluginService.upgrade_plugin_with_github( tenant_id, - args.original_plugin_unique_identifier, - args.new_plugin_unique_identifier, - args.repo, - args.version, - args.package, + req_data.original_plugin_unique_identifier, + req_data.new_plugin_unique_identifier, + req_data.repo, + req_data.version, + req_data.package, ) ) except PluginDaemonClientSideError as e: @@ -967,11 +1068,17 @@ class PluginUninstallApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_DELETE, resource_required=False) @plugin_permission_required(install_required=True) @with_current_tenant_id - def post(self, tenant_id: str): - args = ParserUninstall.model_validate(console_ns.payload) + @model_validate(ParserUninstall) + def post(self, req_data: ParserUninstall, tenant_id: str): try: - return {"success": PluginService.uninstall(tenant_id, args.plugin_installation_id)} + return { + "success": PluginService.uninstall( + tenant_id, + req_data.plugin_installation_id, + preserve_credentials=req_data.preserve_credentials, + ) + } except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -985,14 +1092,13 @@ class PluginChangePermissionApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): + @model_validate(ParserPermissionChange) + def post(self, req_data: ParserPermissionChange, tenant_id: str, user: Account): if not user.is_admin_or_owner: raise Forbidden() - args = ParserPermissionChange.model_validate(console_ns.payload) - set_permission_result = PluginPermissionService.change_permission( - tenant_id, args.install_permission, args.debug_permission, session=db.session() + tenant_id, req_data.install_permission, req_data.debug_permission, session=db.session() ) if not set_permission_result: return jsonable_encoder({"success": False, "message": "Failed to set permission"}) @@ -1036,19 +1142,19 @@ class PluginFetchDynamicSelectOptionsApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def get(self, tenant_id: str, current_user: Account): - args = ParserDynamicOptions.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserDynamicOptions) + def get(self, req_data: ParserDynamicOptions, tenant_id: str, current_user: Account): try: options = PluginParameterService.get_dynamic_select_options( tenant_id=tenant_id, user_id=current_user.id, - plugin_id=args.plugin_id, - provider=args.provider, - action=args.action, - parameter=args.parameter, - credential_id=args.credential_id, - provider_type=args.provider_type, + plugin_id=req_data.plugin_id, + provider=req_data.provider, + action=req_data.action, + parameter=req_data.parameter, + credential_id=req_data.credential_id, + provider_type=req_data.provider_type, ) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -1067,20 +1173,20 @@ class PluginFetchDynamicSelectOptionsWithCredentialsApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, current_user: Account): + @model_validate(ParserDynamicOptionsWithCredentials) + def post(self, req_data: ParserDynamicOptionsWithCredentials, tenant_id: str, current_user: Account): """Fetch dynamic options using credentials directly (for edit mode).""" - args = ParserDynamicOptionsWithCredentials.model_validate(console_ns.payload) try: options = PluginParameterService.get_dynamic_select_options_with_credentials( tenant_id=tenant_id, user_id=current_user.id, - plugin_id=args.plugin_id, - provider=args.provider, - action=args.action, - parameter=args.parameter, - credential_id=args.credential_id, - credentials=args.credentials, + plugin_id=req_data.plugin_id, + provider=req_data.provider, + action=req_data.action, + parameter=req_data.parameter, + credential_id=req_data.credential_id, + credentials=req_data.credentials, ) except PluginDaemonClientSideError as e: return {"code": "plugin_error", "message": e.description}, 400 @@ -1098,13 +1204,12 @@ class PluginChangeAutoUpgradeApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): + @model_validate(ParserAutoUpgradeChange) + def post(self, req_data: ParserAutoUpgradeChange, tenant_id: str, user: Account): if not dify_config.RBAC_ENABLED and not user.is_admin_or_owner: raise Forbidden() - args = ParserAutoUpgradeChange.model_validate(console_ns.payload) - - auto_upgrade = args.auto_upgrade + auto_upgrade = req_data.auto_upgrade set_auto_upgrade_strategy_result = PluginAutoUpgradeService.change_strategy( tenant_id, auto_upgrade.strategy_setting, @@ -1112,7 +1217,7 @@ class PluginChangeAutoUpgradeApi(Resource): auto_upgrade.upgrade_mode, auto_upgrade.exclude_plugins, auto_upgrade.include_plugins, - category=args.category, + category=req_data.category, session=db.session(), ) if not set_auto_upgrade_strategy_result: @@ -1129,16 +1234,16 @@ class PluginFetchAutoUpgradeApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str): - args = ParserAutoUpgradeFetch.model_validate(request.args.to_dict(flat=True)) - auto_upgrade = PluginAutoUpgradeService.get_strategy(tenant_id, args.category, session=db.session()) + @model_validate(ParserAutoUpgradeFetch) + def get(self, req_data: ParserAutoUpgradeFetch, tenant_id: str): + auto_upgrade = PluginAutoUpgradeService.get_strategy(tenant_id, req_data.category, session=db.session()) auto_upgrade_dict = ( _auto_upgrade_settings_to_dict(auto_upgrade) if auto_upgrade else _missing_auto_upgrade_settings(tenant_id) ) return jsonable_encoder( { - "category": args.category, + "category": req_data.category, "auto_upgrade": auto_upgrade_dict, } ) @@ -1153,14 +1258,14 @@ class PluginAutoUpgradeExcludePluginApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @with_current_tenant_id - def post(self, tenant_id: str): + @model_validate(ParserExcludePlugin) + def post(self, req_data: ParserExcludePlugin, tenant_id: str): # exclude one single plugin - args = ParserExcludePlugin.model_validate(console_ns.payload) return jsonable_encoder( { "success": PluginAutoUpgradeService.exclude_plugin( - tenant_id, args.plugin_id, args.category, session=db.session() + tenant_id, req_data.plugin_id, req_data.category, session=db.session() ) } ) @@ -1174,8 +1279,12 @@ class PluginReadmeApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def get(self, tenant_id: str): - args = ParserReadme.model_validate(request.args.to_dict(flat=True)) + @model_validate(ParserReadme) + def get(self, req_data: ParserReadme, tenant_id: str): return jsonable_encoder( - {"readme": PluginService.fetch_plugin_readme(tenant_id, args.plugin_unique_identifier, args.language)} + { + "readme": PluginService.fetch_plugin_readme( + tenant_id, req_data.plugin_unique_identifier, req_data.language + ) + }, ) diff --git a/api/controllers/console/workspace/rbac.py b/api/controllers/console/workspace/rbac.py index 878f91fe168..9aaf2e59e3d 100644 --- a/api/controllers/console/workspace/rbac.py +++ b/api/controllers/console/workspace/rbac.py @@ -11,9 +11,10 @@ from werkzeug.exceptions import NotFound from configs import dify_config from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models from controllers.console import console_ns -from controllers.console.wraps import RBACPermission, RBACResourceScope, rbac_permission_required +from controllers.console.wraps import RBACPermission, RBACResourceScope, model_validate, rbac_permission_required from core.db.session_factory import session_factory from core.rbac import RBACResourceWhitelistScope +from enums import DeploymentEdition from extensions.ext_database import db from libs.login import current_account_with_tenant, login_required from models import Account @@ -196,7 +197,7 @@ def _pagination_options() -> svc.ListOption: def _legacy_workspace_roles( - options: svc.ListOption | None = None, *, include_owner: int = 0 + options: svc.ListOption | None = None, *, include_owner: int = 0, billing_enabled: bool = True ) -> svc.Paginated[svc.RBACRole]: """Return the built-in legacy workspace roles in the RBAC list shape. @@ -207,6 +208,14 @@ def _legacy_workspace_roles( for role_name in ("owner", "admin", "editor", "normal", "dataset_operator"): if not dify_config.DATASET_OPERATOR_ENABLED and role_name == "dataset_operator": continue + + permission_keys = _LEGACY_ROLE_PERMISSION_KEYS[role_name] + valid_permission_keys = [] + for permission_key in permission_keys: + if not billing_enabled and "billing" in permission_key: + continue + valid_permission_keys.append(permission_key) + legacy_roles.append( svc.RBACRole( id=role_name, @@ -216,7 +225,7 @@ def _legacy_workspace_roles( name=role_name, description="", is_builtin=True, - permission_keys=list(dict.fromkeys(_LEGACY_ROLE_PERMISSION_KEYS[role_name])), + permission_keys=valid_permission_keys, role_tag="owner" if role_name == "owner" else "", ) ) @@ -302,15 +311,23 @@ class RBACRolesApi(Resource): RBACResourceScope.WORKSPACE, RBACPermission.WORKSPACE_ROLE_MANAGE, resource_required=False ) @console_ns.response(200, "Success", console_ns.models[_RBACRoleList.__name__]) - def get(self): + @model_validate(_RolesListQuery) + def get(self, req_data: _RolesListQuery): tenant_id, account_id = _current_ids() - query = _RolesListQuery.model_validate(request.args.to_dict(flat=True)) - options = query.to_inner_options() + options = req_data.to_inner_options() if not dify_config.RBAC_ENABLED: - result = _legacy_workspace_roles(options, include_owner=query.include_owner) + result = _legacy_workspace_roles( + options, + include_owner=req_data.include_owner, + billing_enabled=dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD, + ) else: result = svc.RBACService.Roles.list( - tenant_id, account_id, include_owner=query.include_owner, options=options + tenant_id, + account_id, + include_owner=req_data.include_owner, + biiling_enabled=dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD, + options=options, ) return _dump(result) @@ -336,7 +353,14 @@ class RBACRoleItemApi(Resource): @console_ns.response(200, "Success", console_ns.models[svc.RBACRole.__name__]) def get(self, role_id): tenant_id, account_id = _current_ids() - return _dump(svc.RBACService.Roles.get(tenant_id, account_id, str(role_id))) + return _dump( + svc.RBACService.Roles.get( + tenant_id, + account_id, + role_id, + billing_enabled=dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD, + ) + ) @login_required @rbac_permission_required( @@ -703,14 +727,17 @@ class RBACDatasetWhitelistApi(Resource): def put(self, dataset_id): tenant_id, account_id = _current_ids() request = _payload(_ResourceAccessScopeRequest) - return _dump( - svc.RBACService.DatasetAccess.replace_whitelist( - tenant_id, - account_id, - str(dataset_id), - svc.ReplaceMemberBindings(scope=request.scope.value), - ) + result = svc.RBACService.DatasetAccess.replace_whitelist( + tenant_id, + account_id, + str(dataset_id), + svc.ReplaceMemberBindings(scope=request.scope.value), ) + # Widening the scope only records it: the members still need the default access policy + # before they can reach the dataset, same as the app whitelist route above. + if dify_config.RBAC_ENABLED and request.scope is RBACResourceWhitelistScope.ALL: + initialize_created_app_rbac_access_task.delay(tenant_id, account_id, dataset_id=str(dataset_id)) + return _dump(result) @console_ns.route("/workspaces/current/rbac/datasets//user-access-policies") diff --git a/api/controllers/console/workspace/snippets.py b/api/controllers/console/workspace/snippets.py index faf56f8e715..8c254dfa58b 100644 --- a/api/controllers/console/workspace/snippets.py +++ b/api/controllers/console/workspace/snippets.py @@ -27,6 +27,7 @@ from controllers.console.wraps import ( RBACResourceScope, account_initialization_required, edit_permission_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -151,27 +152,26 @@ class CustomizedSnippetsApi(Resource): ) @with_current_user @with_current_tenant_id - def post(self, current_tenant_id: str, current_user: Account): + @model_validate(CreateSnippetPayload) + def post(self, req_data: CreateSnippetPayload, current_tenant_id: str, current_user: Account): """Create a new customized snippet.""" - payload = CreateSnippetPayload.model_validate(console_ns.payload or {}) - try: - snippet_type = SnippetType(payload.type) + snippet_type = SnippetType(req_data.type) except ValueError: snippet_type = SnippetType.NODE try: - if payload.graph is not None: - SnippetService.validate_snippet_graph_forbidden_nodes(payload.graph) + if req_data.graph is not None: + SnippetService.validate_snippet_graph_forbidden_nodes(req_data.graph) snippet_service = _snippet_service() snippet = snippet_service.create_snippet( tenant_id=current_tenant_id, - name=payload.name, - description=payload.description, + name=req_data.name, + description=req_data.description, snippet_type=snippet_type, - icon_info=payload.icon_info.model_dump() if payload.icon_info else None, - input_fields=[f.model_dump() for f in payload.input_fields] if payload.input_fields else None, + icon_info=req_data.icon_info.model_dump() if req_data.icon_info else None, + input_fields=[f.model_dump() for f in req_data.input_fields] if req_data.input_fields else None, account=current_user, ) except ValueError as e: @@ -216,7 +216,8 @@ class CustomizedSnippetDetailApi(Resource): ) @with_current_user @with_current_tenant_id - def patch(self, current_tenant_id: str, current_user: Account, snippet_id: str): + @model_validate(UpdateSnippetPayload) + def patch(self, req_data: UpdateSnippetPayload, current_tenant_id: str, current_user: Account, snippet_id: str): """Update customized snippet.""" snippet_service = _snippet_service() snippet = snippet_service.get_snippet_by_id( @@ -227,11 +228,10 @@ class CustomizedSnippetDetailApi(Resource): if not snippet: raise NotFound("Snippet not found") - payload = UpdateSnippetPayload.model_validate(console_ns.payload or {}) - update_data = payload.model_dump(exclude_unset=True) + update_data = req_data.model_dump(exclude_unset=True) if "icon_info" in update_data and update_data["icon_info"] is not None: - update_data["icon_info"] = payload.icon_info.model_dump() if payload.icon_info else None + update_data["icon_info"] = req_data.icon_info.model_dump() if req_data.icon_info else None if not update_data: return {"message": "No valid fields to update"}, 400 @@ -349,19 +349,18 @@ class CustomizedSnippetImportApi(Resource): ) @with_current_user @with_session - def post(self, session: Session, current_user: Account): + @model_validate(SnippetImportPayload) + def post(self, req_data: SnippetImportPayload, session: Session, current_user: Account): """Import snippet from DSL.""" - payload = SnippetImportPayload.model_validate(console_ns.payload or {}) - import_service = SnippetDslService(session) result = import_service.import_snippet( account=current_user, - import_mode=payload.mode, - yaml_content=payload.yaml_content, - yaml_url=payload.yaml_url, - snippet_id=payload.snippet_id, - name=payload.name, - description=payload.description, + import_mode=req_data.mode, + yaml_content=req_data.yaml_content, + yaml_url=req_data.yaml_url, + snippet_id=req_data.snippet_id, + name=req_data.name, + description=req_data.description, ) # Return appropriate status code based on result diff --git a/api/controllers/console/workspace/tool_providers.py b/api/controllers/console/workspace/tool_providers.py index 6e9b4be8ac4..68daf17d535 100644 --- a/api/controllers/console/workspace/tool_providers.py +++ b/api/controllers/console/workspace/tool_providers.py @@ -33,6 +33,7 @@ from controllers.console.wraps import ( account_initialization_required, enterprise_license_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -60,6 +61,7 @@ from core.tools.entities.tool_entities import ( ToolProviderType, WorkflowToolParameterConfiguration, ) +from enums import DeploymentEdition from extensions.ext_database import db from fields.base import ResponseModel from libs.helper import alphanumeric, dump_response, uuid_value @@ -281,10 +283,10 @@ def _resolve_identity_mode(requested: IdentityMode | None, *, current: IdentityM can never imply forwarding that the runtime won't perform. This gates the API surface to match the backend gate in ``MCPTool._forwarding_requested`` — both the API and the backend - invocation must be gated on ``dify_config.ENTERPRISE_ENABLED``. + invocation must be gated on the Enterprise deployment edition. """ mode = current if requested is None else requested - if mode != IdentityMode.OFF and not dify_config.ENTERPRISE_ENABLED: + if mode != IdentityMode.OFF and dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return IdentityMode.OFF return mode @@ -555,16 +557,15 @@ class ToolBuiltinProviderDeleteApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, provider: str): - - payload = BuiltinToolCredentialDeletePayload.model_validate(console_ns.payload or {}) + @model_validate(BuiltinToolCredentialDeletePayload) + def post(self, req_data: BuiltinToolCredentialDeletePayload, tenant_id: str, provider: str): return dump_response( SimpleResultResponse, BuiltinToolManageService.delete_builtin_tool_provider( tenant_id, provider, - payload.credential_id, + req_data.credential_id, ), ) @@ -582,19 +583,18 @@ class ToolBuiltinProviderAddApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account, provider: str): - payload = BuiltinToolAddPayload.model_validate(console_ns.payload or {}) - + @model_validate(BuiltinToolAddPayload) + def post(self, req_data: BuiltinToolAddPayload, tenant_id: str, user: Account, provider: str): return dump_response( SimpleResultResponse, BuiltinToolManageService.add_builtin_tool_provider( user_id=user.id, tenant_id=tenant_id, provider=provider, - credentials=payload.credentials, - name=payload.name, - api_type=CredentialType.of(payload.type), - visibility=payload.visibility, + credentials=req_data.credentials, + name=req_data.name, + api_type=CredentialType.of(req_data.type), + visibility=req_data.visibility, ), ) @@ -614,16 +614,15 @@ class ToolBuiltinProviderUpdateApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account, provider: str): - payload = BuiltinToolUpdatePayload.model_validate(console_ns.payload or {}) - + @model_validate(BuiltinToolUpdatePayload) + def post(self, req_data: BuiltinToolUpdatePayload, tenant_id: str, user: Account, provider: str): result = BuiltinToolManageService.update_builtin_tool_provider( user_id=user.id, tenant_id=tenant_id, provider=provider, - credential_id=payload.credential_id, - credentials=payload.credentials, - name=payload.name or "", + credential_id=req_data.credential_id, + credentials=req_data.credentials, + name=req_data.name or "", ) return dump_response(SimpleResultResponse, result) @@ -683,22 +682,21 @@ class ToolApiProviderAddApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): - payload = ApiToolProviderAddPayload.model_validate(console_ns.payload or {}) - + @model_validate(ApiToolProviderAddPayload) + def post(self, req_data: ApiToolProviderAddPayload, tenant_id: str, user: Account): return dump_response( SimpleResultResponse, ApiToolManageService.create_api_tool_provider( user.id, tenant_id, - payload.provider, - payload.icon.model_dump(mode="json"), - payload.credentials, - payload.schema_type, - payload.schema_, - payload.privacy_policy or "", - payload.custom_disclaimer or "", - payload.labels or [], + req_data.provider, + req_data.icon.model_dump(mode="json"), + req_data.credentials, + req_data.schema_type, + req_data.schema_, + req_data.privacy_policy or "", + req_data.custom_disclaimer or "", + req_data.labels or [], ), ) @@ -764,23 +762,22 @@ class ToolApiProviderUpdateApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): - payload = ApiToolProviderUpdatePayload.model_validate(console_ns.payload or {}) - + @model_validate(ApiToolProviderUpdatePayload) + def post(self, req_data: ApiToolProviderUpdatePayload, tenant_id: str, user: Account): return dump_response( SimpleResultResponse, ApiToolManageService.update_api_tool_provider( user.id, tenant_id, - payload.provider, - payload.original_provider, - payload.icon.model_dump(mode="json"), - payload.credentials, - payload.schema_type, - payload.schema_, - payload.privacy_policy, - payload.custom_disclaimer, - payload.labels or [], + req_data.provider, + req_data.original_provider, + req_data.icon.model_dump(mode="json"), + req_data.credentials, + req_data.schema_type, + req_data.schema_, + req_data.privacy_policy, + req_data.custom_disclaimer, + req_data.labels or [], ), ) @@ -796,15 +793,14 @@ class ToolApiProviderDeleteApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): - payload = ApiToolProviderDeletePayload.model_validate(console_ns.payload or {}) - + @model_validate(ApiToolProviderDeletePayload) + def post(self, req_data: ApiToolProviderDeletePayload, tenant_id: str, user: Account): return dump_response( SimpleResultResponse, ApiToolManageService.delete_api_tool_provider( user.id, tenant_id, - payload.provider, + req_data.provider, ), ) @@ -861,10 +857,9 @@ class ToolApiProviderSchemaApi(Resource): @setup_required @login_required @account_initialization_required - def post(self): - payload = ApiToolSchemaPayload.model_validate(console_ns.payload or {}) - - return dump_response(ApiSchemaParseResponse, ApiToolManageService.parser_api_schema(schema=payload.schema_)) + @model_validate(ApiToolSchemaPayload) + def post(self, req_data: ApiToolSchemaPayload): + return dump_response(ApiSchemaParseResponse, ApiToolManageService.parser_api_schema(schema=req_data.schema_)) @console_ns.route("/workspaces/current/tool-provider/api/test/pre") @@ -879,18 +874,18 @@ class ToolApiProviderPreviousTestApi(Resource): @login_required @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str): - payload = ApiToolTestPayload.model_validate(console_ns.payload or {}) + @model_validate(ApiToolTestPayload) + def post(self, req_data: ApiToolTestPayload, current_tenant_id: str): return dump_response( ApiToolPreviewResponse, ApiToolManageService.test_api_tool_preview( current_tenant_id, - payload.provider_name or "", - payload.tool_name, - payload.credentials, - payload.parameters, - payload.schema_type, - payload.schema_, + req_data.provider_name or "", + req_data.tool_name, + req_data.credentials, + req_data.parameters, + req_data.schema_type, + req_data.schema_, ), ) @@ -906,22 +901,21 @@ class ToolWorkflowProviderCreateApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): - payload = WorkflowToolCreatePayload.model_validate(console_ns.payload or {}) - + @model_validate(WorkflowToolCreatePayload) + def post(self, req_data: WorkflowToolCreatePayload, tenant_id: str, user: Account): return dump_response( SimpleResultResponse, WorkflowToolManageService.create_workflow_tool( user_id=user.id, tenant_id=tenant_id, - workflow_app_id=payload.workflow_app_id, - name=payload.name, - label=payload.label, - icon=payload.icon.model_dump(mode="json"), - description=payload.description, - parameters=payload.parameters, - privacy_policy=payload.privacy_policy or "", - labels=payload.labels or [], + workflow_app_id=req_data.workflow_app_id, + name=req_data.name, + label=req_data.label, + icon=req_data.icon.model_dump(mode="json"), + description=req_data.description, + parameters=req_data.parameters, + privacy_policy=req_data.privacy_policy or "", + labels=req_data.labels or [], ), ) @@ -937,22 +931,21 @@ class ToolWorkflowProviderUpdateApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): - payload = WorkflowToolUpdatePayload.model_validate(console_ns.payload or {}) - + @model_validate(WorkflowToolUpdatePayload) + def post(self, req_data: WorkflowToolUpdatePayload, tenant_id: str, user: Account): return dump_response( SimpleResultResponse, WorkflowToolManageService.update_workflow_tool( user.id, tenant_id, - payload.workflow_tool_id, - payload.name, - payload.label, - payload.icon.model_dump(mode="json"), - payload.description, - payload.parameters, - payload.privacy_policy or "", - payload.labels or [], + req_data.workflow_tool_id, + req_data.name, + req_data.label, + req_data.icon.model_dump(mode="json"), + req_data.description, + req_data.parameters, + req_data.privacy_policy or "", + req_data.labels or [], ), ) @@ -968,15 +961,14 @@ class ToolWorkflowProviderDeleteApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): - payload = WorkflowToolDeletePayload.model_validate(console_ns.payload or {}) - + @model_validate(WorkflowToolDeletePayload) + def post(self, req_data: WorkflowToolDeletePayload, tenant_id: str, user: Account): return dump_response( SimpleResultResponse, WorkflowToolManageService.delete_workflow_tool( user.id, tenant_id, - payload.workflow_tool_id, + req_data.workflow_tool_id, ), ) @@ -1224,12 +1216,12 @@ class ToolBuiltinProviderSetDefaultApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_USE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, current_tenant_id: str, provider: str): - payload = BuiltinProviderDefaultCredentialPayload.model_validate(console_ns.payload or {}) + @model_validate(BuiltinProviderDefaultCredentialPayload) + def post(self, req_data: BuiltinProviderDefaultCredentialPayload, current_tenant_id: str, provider: str): return dump_response( SimpleResultResponse, BuiltinToolManageService.set_default_provider( - tenant_id=current_tenant_id, provider=provider, id=payload.id + tenant_id=current_tenant_id, provider=provider, id=req_data.id ), ) @@ -1246,17 +1238,16 @@ class ToolOAuthCustomClient(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, provider: str): - payload = ToolOAuthCustomClientPayload.model_validate(console_ns.payload or {}) - + @model_validate(ToolOAuthCustomClientPayload) + def post(self, req_data: ToolOAuthCustomClientPayload, tenant_id: str, provider: str): return dump_response( SimpleResultResponse, BuiltinToolManageService.save_custom_oauth_client_params( tenant_id=tenant_id, provider=provider, - client_params=payload.client_params or {}, - enable_oauth_custom_client=payload.enable_oauth_custom_client - if payload.enable_oauth_custom_client is not None + client_params=req_data.client_params or {}, + enable_oauth_custom_client=req_data.enable_oauth_custom_client + if req_data.enable_oauth_custom_client is not None else True, ), ) @@ -1349,11 +1340,10 @@ class ToolProviderMCPApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.MCP_MANAGE, resource_required=False) @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account): - payload = MCPProviderCreatePayload.model_validate(console_ns.payload or {}) - - configuration = payload.configuration or MCPConfiguration() - authentication = payload.authentication + @model_validate(MCPProviderCreatePayload) + def post(self, req_data: MCPProviderCreatePayload, tenant_id: str, user: Account): + configuration = req_data.configuration or MCPConfiguration() + authentication = req_data.authentication # 1) Create provider in a short transaction (no network I/O inside) with session_factory.create_session() as session, session.begin(): @@ -1361,24 +1351,24 @@ class ToolProviderMCPApi(Resource): result = service.create_provider( tenant_id=tenant_id, user_id=user.id, - server_url=payload.server_url, - name=payload.name, - icon=payload.icon, - icon_type=payload.icon_type, - icon_background=payload.icon_background, - server_identifier=payload.server_identifier, - headers=payload.headers or {}, + server_url=req_data.server_url, + name=req_data.name, + icon=req_data.icon, + icon_type=req_data.icon_type, + icon_background=req_data.icon_background, + server_identifier=req_data.server_identifier, + headers=req_data.headers or {}, configuration=configuration, authentication=authentication, - identity_mode=_resolve_identity_mode(payload.identity_mode, current=IdentityMode.OFF), + identity_mode=_resolve_identity_mode(req_data.identity_mode, current=IdentityMode.OFF), ) # 2) Try to fetch tools immediately after creation so they appear without a second save. # Perform network I/O outside any DB session to avoid holding locks. try: reconnect = MCPToolManageService.reconnect_with_url( - server_url=payload.server_url, - headers=payload.headers or {}, + server_url=req_data.server_url, + headers=req_data.headers or {}, timeout=configuration.timeout, sse_read_timeout=configuration.sse_read_timeout, ) @@ -1403,24 +1393,24 @@ class ToolProviderMCPApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.MCP_MANAGE, resource_required=False) @with_current_tenant_id - def put(self, current_tenant_id: str): - payload = MCPProviderUpdatePayload.model_validate(console_ns.payload or {}) - configuration = payload.configuration or MCPConfiguration() - authentication = payload.authentication + @model_validate(MCPProviderUpdatePayload) + def put(self, req_data: MCPProviderUpdatePayload, current_tenant_id: str): + configuration = req_data.configuration or MCPConfiguration() + authentication = req_data.authentication # Step 1: Get provider data for URL validation (short-lived session, no network I/O) validation_data = None with sessionmaker(db.engine).begin() as session: service = MCPToolManageService(session=session) validation_data = service.get_provider_for_url_validation( - tenant_id=current_tenant_id, provider_id=payload.provider_id + tenant_id=current_tenant_id, provider_id=req_data.provider_id ) # Step 2: Perform URL validation with network I/O OUTSIDE of any database session # This prevents holding database locks during potentially slow network operations validation_result = MCPToolManageService.validate_server_url_standalone( tenant_id=current_tenant_id, - new_server_url=payload.server_url, + new_server_url=req_data.server_url, validation_data=validation_data, ) @@ -1428,20 +1418,20 @@ class ToolProviderMCPApi(Resource): with sessionmaker(db.engine).begin() as session: service = MCPToolManageService(session=session) # Resolve "leave unchanged" (None) against the stored value, and gate - # the result on ENTERPRISE_ENABLED — both are API-layer concerns, so + # the result on the Enterprise edition — both are API-layer concerns, so # the service receives a concrete IdentityMode. - existing = service.get_provider(provider_id=payload.provider_id, tenant_id=current_tenant_id) - identity_mode = _resolve_identity_mode(payload.identity_mode, current=IdentityMode(existing.identity_mode)) + existing = service.get_provider(provider_id=req_data.provider_id, tenant_id=current_tenant_id) + identity_mode = _resolve_identity_mode(req_data.identity_mode, current=IdentityMode(existing.identity_mode)) service.update_provider( tenant_id=current_tenant_id, - provider_id=payload.provider_id, - server_url=payload.server_url, - name=payload.name, - icon=payload.icon, - icon_type=payload.icon_type, - icon_background=payload.icon_background, - server_identifier=payload.server_identifier, - headers=payload.headers or {}, + provider_id=req_data.provider_id, + server_url=req_data.server_url, + name=req_data.name, + icon=req_data.icon, + icon_type=req_data.icon_type, + icon_background=req_data.icon_background, + server_identifier=req_data.server_identifier, + headers=req_data.headers or {}, configuration=configuration, authentication=authentication, validation_result=validation_result, @@ -1457,12 +1447,11 @@ class ToolProviderMCPApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.MCP_MANAGE, resource_required=False) @with_current_tenant_id - def delete(self, current_tenant_id: str): - payload = MCPProviderDeletePayload.model_validate(console_ns.payload or {}) - + @model_validate(MCPProviderDeletePayload) + def delete(self, req_data: MCPProviderDeletePayload, current_tenant_id: str): with sessionmaker(db.engine).begin() as session: service = MCPToolManageService(session=session) - service.delete_provider(tenant_id=current_tenant_id, provider_id=payload.provider_id) + service.delete_provider(tenant_id=current_tenant_id, provider_id=req_data.provider_id) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -1476,9 +1465,9 @@ class ToolMCPAuthApi(Resource): @account_initialization_required @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.MCP_MANAGE, resource_required=False) @with_current_tenant_id - def post(self, tenant_id: str): - payload = MCPAuthPayload.model_validate(console_ns.payload or {}) - provider_id = payload.provider_id + @model_validate(MCPAuthPayload) + def post(self, req_data: MCPAuthPayload, tenant_id: str): + provider_id = req_data.provider_id with sessionmaker(db.engine).begin() as session: service = MCPToolManageService(session=session) @@ -1515,7 +1504,7 @@ class ToolMCPAuthApi(Resource): # Pass the extracted OAuth metadata hints to auth() auth_result = auth( provider_entity, - payload.authorization_code, + req_data.authorization_code, resource_metadata_url=e.resource_metadata_url, scope_hint=e.scope_hint, ) diff --git a/api/controllers/console/workspace/trigger_providers.py b/api/controllers/console/workspace/trigger_providers.py index aff800a5bc9..8b329e13536 100644 --- a/api/controllers/console/workspace/trigger_providers.py +++ b/api/controllers/console/workspace/trigger_providers.py @@ -39,6 +39,7 @@ from ..wraps import ( account_initialization_required, edit_permission_required, is_admin_or_owner_required, + model_validate, rbac_permission_required, setup_required, with_current_tenant_id, @@ -233,13 +234,12 @@ class TriggerSubscriptionBuilderCreateApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account, provider: str): + @model_validate(TriggerSubscriptionBuilderCreatePayload) + def post(self, req_data: TriggerSubscriptionBuilderCreatePayload, tenant_id: str, user: Account, provider: str): """Add a new subscription instance for a trigger provider""" - payload = TriggerSubscriptionBuilderCreatePayload.model_validate(console_ns.payload or {}) - try: - credential_type = CredentialType.of(payload.credential_type) + credential_type = CredentialType.of(req_data.credential_type) subscription_builder = TriggerSubscriptionBuilderService.create_trigger_subscription_builder( tenant_id=tenant_id, user_id=user.id, @@ -268,9 +268,16 @@ class TriggerSubscriptionBuilderGetApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @account_initialization_required - def get(self, provider: str, subscription_builder_id: str): + @with_current_user + @with_current_tenant_id + def get(self, tenant_id: str, user: Account, provider: str, subscription_builder_id: str): """Get a subscription instance for a trigger provider""" - subscription_builder = TriggerSubscriptionBuilderService.get_subscription_builder_by_id(subscription_builder_id) + subscription_builder = TriggerSubscriptionBuilderService.get_subscription_builder_by_id( + tenant_id=tenant_id, + user_id=user.id, + provider_id=TriggerProviderID(provider), + subscription_builder_id=subscription_builder_id, + ) return subscription_builder.model_dump(mode="json") @@ -291,11 +298,17 @@ class TriggerSubscriptionBuilderVerifyApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account, provider: str, subscription_builder_id: str): + @model_validate(TriggerSubscriptionBuilderVerifyPayload) + def post( + self, + req_data: TriggerSubscriptionBuilderVerifyPayload, + tenant_id: str, + user: Account, + provider: str, + subscription_builder_id: str, + ): """Verify and update a subscription instance for a trigger provider""" - payload = TriggerSubscriptionBuilderVerifyPayload.model_validate(console_ns.payload or {}) - try: # Use atomic update_and_verify to prevent race conditions result = TriggerSubscriptionBuilderService.update_and_verify_builder( @@ -304,7 +317,7 @@ class TriggerSubscriptionBuilderVerifyApi(Resource): provider_id=TriggerProviderID(provider), subscription_builder_id=subscription_builder_id, subscription_builder_updater=SubscriptionBuilderUpdater( - credentials=payload.credentials, + credentials=req_data.credentials, ), ) return dump_response(TriggerVerificationResponse, result) @@ -328,21 +341,30 @@ class TriggerSubscriptionBuilderUpdateApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.CREDENTIAL_MANAGE, resource_required=False) @account_initialization_required + @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, provider: str, subscription_builder_id: str): + @model_validate(TriggerSubscriptionBuilderUpdatePayload) + def post( + self, + req_data: TriggerSubscriptionBuilderUpdatePayload, + tenant_id: str, + user: Account, + provider: str, + subscription_builder_id: str, + ): """Update a subscription instance for a trigger provider""" - payload = TriggerSubscriptionBuilderUpdatePayload.model_validate(console_ns.payload or {}) try: return TriggerSubscriptionBuilderService.update_trigger_subscription_builder( tenant_id=tenant_id, + user_id=user.id, provider_id=TriggerProviderID(provider), subscription_builder_id=subscription_builder_id, subscription_builder_updater=SubscriptionBuilderUpdater( - name=payload.name, - parameters=payload.parameters, - properties=payload.properties, - credentials=payload.credentials, + name=req_data.name, + parameters=req_data.parameters, + properties=req_data.properties, + credentials=req_data.credentials, ), ).model_dump(mode="json") except Exception as e: @@ -364,11 +386,18 @@ class TriggerSubscriptionBuilderLogsApi(Resource): @edit_permission_required @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @account_initialization_required - def get(self, provider: str, subscription_builder_id: str): + @with_current_user + @with_current_tenant_id + def get(self, tenant_id: str, user: Account, provider: str, subscription_builder_id: str): """Get the request logs for a subscription instance for a trigger provider""" try: - logs = TriggerSubscriptionBuilderService.list_logs(subscription_builder_id) + logs = TriggerSubscriptionBuilderService.list_logs( + tenant_id=tenant_id, + user_id=user.id, + provider_id=TriggerProviderID(provider), + subscription_builder_id=subscription_builder_id, + ) return dump_response(TriggerSubscriptionBuilderLogsResponse, {"logs": logs}) except Exception as e: logger.exception("Error getting request logs for subscription builder", exc_info=e) @@ -390,9 +419,16 @@ class TriggerSubscriptionBuilderBuildApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account, provider: str, subscription_builder_id: str): + @model_validate(TriggerSubscriptionBuilderUpdatePayload) + def post( + self, + req_data: TriggerSubscriptionBuilderUpdatePayload, + tenant_id: str, + user: Account, + provider: str, + subscription_builder_id: str, + ): """Build a subscription instance for a trigger provider""" - payload = TriggerSubscriptionBuilderUpdatePayload.model_validate(console_ns.payload or {}) try: # Use atomic update_and_build to prevent race conditions TriggerSubscriptionBuilderService.update_and_build_builder( @@ -401,9 +437,9 @@ class TriggerSubscriptionBuilderBuildApi(Resource): provider_id=TriggerProviderID(provider), subscription_builder_id=subscription_builder_id, subscription_builder_updater=SubscriptionBuilderUpdater( - name=payload.name, - parameters=payload.parameters, - properties=payload.properties, + name=req_data.name, + parameters=req_data.parameters, + properties=req_data.properties, ), ) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -426,11 +462,10 @@ class TriggerSubscriptionUpdateApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, subscription_id: str): + @model_validate(TriggerSubscriptionBuilderUpdatePayload) + def post(self, req_data: TriggerSubscriptionBuilderUpdatePayload, tenant_id: str, subscription_id: str): """Update a subscription instance""" - request = TriggerSubscriptionBuilderUpdatePayload.model_validate(console_ns.payload or {}) - subscription = TriggerProviderService.get_subscription_by_id( tenant_id=tenant_id, subscription_id=subscription_id, @@ -442,7 +477,9 @@ class TriggerSubscriptionUpdateApi(Resource): try: # For rename only, just update the name - rename = request.name is not None and not any((request.credentials, request.parameters, request.properties)) + rename = req_data.name is not None and not any( + (req_data.credentials, req_data.parameters, req_data.properties) + ) # When credential type is UNAUTHORIZED, it indicates the subscription was manually created # For Manually created subscription, they dont have credentials, parameters # They only have name and properties(which is input by user) @@ -451,8 +488,8 @@ class TriggerSubscriptionUpdateApi(Resource): TriggerProviderService.update_trigger_subscription( tenant_id=tenant_id, subscription_id=subscription_id, - name=request.name, - properties=request.properties, + name=req_data.name, + properties=req_data.properties, ) return SimpleResultResponse(result="success").model_dump(mode="json") @@ -460,11 +497,11 @@ class TriggerSubscriptionUpdateApi(Resource): # we need to call third party provider(e.g. GitHub) to rebuild the subscription TriggerProviderService.rebuild_trigger_subscription( tenant_id=tenant_id, - name=request.name, + name=req_data.name, provider_id=provider_id, subscription_id=subscription_id, - credentials=request.credentials or subscription.credentials, - parameters=request.parameters or subscription.parameters, + credentials=req_data.credentials or subscription.credentials, + parameters=req_data.parameters or subscription.parameters, ) return SimpleResultResponse(result="success").model_dump(mode="json") except ValueError as e: @@ -652,6 +689,7 @@ class TriggerOAuthCallbackApi(Resource): # Update subscription builder TriggerSubscriptionBuilderService.update_trigger_subscription_builder( tenant_id=tenant_id, + user_id=user_id, provider_id=provider_id, subscription_builder_id=subscription_builder_id, subscription_builder_updater=SubscriptionBuilderUpdater( @@ -723,18 +761,17 @@ class TriggerOAuthClientManageApi(Resource): @rbac_permission_required(RBACResourceScope.WORKSPACE, RBACPermission.PLUGIN_PREFERENCES, resource_required=False) @account_initialization_required @with_current_tenant_id - def post(self, tenant_id: str, provider: str): + @model_validate(TriggerOAuthClientPayload) + def post(self, req_data: TriggerOAuthClientPayload, tenant_id: str, provider: str): """Configure custom OAuth client for a provider""" - payload = TriggerOAuthClientPayload.model_validate(console_ns.payload or {}) - try: provider_id = TriggerProviderID(provider) result = TriggerProviderService.save_custom_oauth_client_params( tenant_id=tenant_id, provider_id=provider_id, - client_params=payload.client_params, - enabled=payload.enabled, + client_params=req_data.client_params, + enabled=req_data.enabled, ) return dump_response(SimpleResultResponse, result) @@ -788,18 +825,24 @@ class TriggerSubscriptionVerifyApi(Resource): @account_initialization_required @with_current_user @with_current_tenant_id - def post(self, tenant_id: str, user: Account, provider: str, subscription_id: str): + @model_validate(TriggerSubscriptionBuilderVerifyPayload) + def post( + self, + req_data: TriggerSubscriptionBuilderVerifyPayload, + tenant_id: str, + user: Account, + provider: str, + subscription_id: str, + ): """Verify credentials for an existing subscription (edit mode only)""" - verify_request = TriggerSubscriptionBuilderVerifyPayload.model_validate(console_ns.payload or {}) - try: result = TriggerProviderService.verify_subscription_credentials( tenant_id=tenant_id, user_id=user.id, provider_id=TriggerProviderID(provider), subscription_id=subscription_id, - credentials=verify_request.credentials, + credentials=req_data.credentials, ) return dump_response(TriggerVerificationResponse, result) except ValueError as e: diff --git a/api/controllers/console/workspace/workspace.py b/api/controllers/console/workspace/workspace.py index dbb9047e0ec..4ded2d8fd85 100644 --- a/api/controllers/console/workspace/workspace.py +++ b/api/controllers/console/workspace/workspace.py @@ -7,7 +7,7 @@ from flask_restx import Resource from pydantic import BaseModel, Field, field_validator from sqlalchemy import select from sqlalchemy.orm import Session -from werkzeug.exceptions import NotFound, Unauthorized +from werkzeug.exceptions import NotFound import services from configs import dify_config @@ -28,6 +28,8 @@ from controllers.common.session import with_session from controllers.console import console_ns from controllers.console.admin import admin_required from controllers.console.error import AccountNotLinkTenantError +from controllers.console.flask_admission import console_account_admission +from controllers.console.workspace.error import CurrentWorkspaceArchivedError from controllers.console.wraps import ( account_initialization_required, cloud_edition_billing_resource_check, @@ -36,18 +38,17 @@ from controllers.console.wraps import ( with_current_tenant_id, with_current_user, ) -from enums.cloud_plan import CloudPlan -from enums.deployment_edition import DeploymentEdition +from enums import CloudPlan +from extensions.ext_application_services import application_services from extensions.ext_database import db from fields.base import ResponseModel from libs.helper import dump_response, to_timestamp from libs.login import login_required from libs.pagination import paginate_query -from models.account import Account, Tenant, TenantAccountJoin, TenantCustomConfigDict, TenantStatus +from machinery.context import RequestContext +from models.account import Account, Tenant, TenantAccountRole, TenantCustomConfigDict, TenantStatus from services.account_service import TenantService -from services.billing_service import BillingService, SubscriptionPlan from services.enterprise.enterprise_service import EnterpriseService -from services.feature_service import FeatureService from services.file_service import FileService from services.workspace_service import WorkspaceService @@ -80,7 +81,7 @@ class WorkspaceInfoPayload(BaseModel): class TenantInfoResponse(ResponseModel): id: str name: str | None = None - plan: str | None = None + plan: CloudPlan | None = None status: str | None = None created_at: int | None = None role: str | None = None @@ -92,7 +93,7 @@ class TenantInfoResponse(ResponseModel): trial_credits_exhausted_at: int | None = None next_credit_reset_date: int | None = None - @field_validator("plan", "status", "trial_end_reason", mode="before") + @field_validator("status", "trial_end_reason", mode="before") @classmethod def _normalize_enum_like(cls, value): if value is None: @@ -107,16 +108,24 @@ class TenantInfoResponse(ResponseModel): return to_timestamp(value) +class CurrentWorkspaceSummaryResponse(ResponseModel): + id: str + name: str + role: TenantAccountRole + plan: CloudPlan | None + credits: int | None = Field(description="Remaining credits in the effective pool; -1 means unlimited.") + + class TenantListItemResponse(ResponseModel): id: str name: str | None = None - plan: str | None = None + plan: CloudPlan | None = None status: str | None = None created_at: int | None = None last_opened_at: int | None = None current: bool - @field_validator("plan", "status", mode="before") + @field_validator("status", mode="before") @classmethod def _normalize_enum_like(cls, value): if value is None: @@ -203,6 +212,7 @@ register_schema_models( ) register_response_schema_models( console_ns, + CurrentWorkspaceSummaryResponse, TenantInfoResponse, TenantListItemResponse, TenantListResponse, @@ -219,58 +229,10 @@ register_response_schema_models( @console_ns.route("/workspaces") class TenantListApi(Resource): @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[TenantListResponse.__name__]) - @setup_required - @login_required - @account_initialization_required - @with_current_user - @with_current_tenant_id - @with_session(write=False) - def get(self, session: Session, current_tenant_id: str, current_user: Account): - tenant_rows: list[tuple[Tenant, TenantAccountJoin]] = [ - (tenant, membership) - for tenant, membership in TenantService.get_workspaces_for_account(current_user.id, session=session) - if tenant.status == TenantStatus.NORMAL - ] - tenants = [tenant for tenant, _ in tenant_rows] - tenant_dicts = [] - is_enterprise_only = dify_config.ENTERPRISE_ENABLED and not dify_config.BILLING_ENABLED - is_saas = dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and dify_config.BILLING_ENABLED - tenant_plans: dict[str, SubscriptionPlan] = {} - - if is_saas: - tenant_ids = [tenant.id for tenant in tenants] - if tenant_ids: - tenant_plans = BillingService.get_plan_bulk(tenant_ids) - if not tenant_plans: - logger.warning("get_plan_bulk returned empty result, falling back to legacy feature path") - - for tenant, membership in tenant_rows: - plan: str = CloudPlan.SANDBOX - if is_saas: - tenant_plan = tenant_plans.get(tenant.id) - if tenant_plan: - plan = tenant_plan["plan"] or CloudPlan.SANDBOX - else: - features = FeatureService.get_features(tenant.id, exclude_vector_space=True) - plan = features.billing.subscription.plan or CloudPlan.SANDBOX - elif not is_enterprise_only: - features = FeatureService.get_features(tenant.id, exclude_vector_space=True) - plan = features.billing.subscription.plan or CloudPlan.SANDBOX - - # Create a dictionary with tenant attributes - tenant_dict = { - "id": tenant.id, - "name": tenant.name, - "status": tenant.status, - "created_at": tenant.created_at, - "last_opened_at": membership.last_opened_at, - "plan": plan, - "current": tenant.id == current_tenant_id if current_tenant_id else False, - } - - tenant_dicts.append(tenant_dict) - - return dump_response(TenantListResponse, {"workspaces": tenant_dicts}), HTTPStatus.OK + @console_account_admission() + def get(self, request_context: RequestContext): + workspaces = application_services().workspace_queries.list_for_account(request_context) + return dump_response(TenantListResponse, {"workspaces": workspaces}), HTTPStatus.OK @console_ns.route("/all-workspaces") @@ -295,35 +257,31 @@ class WorkspaceListApi(Resource): ).model_dump(mode="json"), HTTPStatus.OK -@console_ns.route("/workspaces/current", endpoint="workspaces_current") -@console_ns.route("/info", endpoint="info") # Deprecated -class TenantApi(Resource): +@console_ns.route("/workspaces/current/summary") +class CurrentWorkspaceSummaryApi(Resource): + @console_ns.response( + HTTPStatus.OK, + "Success", + console_ns.models[CurrentWorkspaceSummaryResponse.__name__], + ) + @console_ns.response(HTTPStatus.CONFLICT, "Current workspace is archived") @setup_required @login_required @account_initialization_required - @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[TenantInfoResponse.__name__]) @with_current_user - @with_session - def post(self, session: Session, current_user: Account): - if request.path == "/info": - logger.warning("Deprecated URL /info was used.") - + @with_session(write=False) + def get(self, session: Session, current_user: Account): tenant = current_user.current_tenant if not tenant: raise ValueError("No current tenant") - if tenant.status == TenantStatus.ARCHIVE: - tenants = TenantService.get_join_tenants(current_user, session=session) - # if there is any tenant, switch to the first one - if len(tenants) > 0: - TenantService.switch_tenant(current_user, tenants[0].id, session=session) - tenant = tenants[0] - # else, raise Unauthorized - else: - raise Unauthorized("workspace is archived") + raise CurrentWorkspaceArchivedError() return ( - dump_response(TenantInfoResponse, WorkspaceService.get_tenant_info(tenant, session=session)), + dump_response( + CurrentWorkspaceSummaryResponse, + WorkspaceService.get_current_workspace_summary(tenant, current_user.id, session=session), + ), HTTPStatus.OK, ) @@ -358,6 +316,31 @@ class SwitchWorkspaceApi(Resource): @console_ns.route("/workspaces/custom-config") class CustomConfigWorkspaceApi(Resource): + @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[WorkspaceCustomConfigResponse.__name__]) + @setup_required + @login_required + @account_initialization_required + @with_current_tenant_id + @with_session(write=False) + def get(self, session: Session, current_tenant_id: str): + tenant = TenantService.get_tenant_by_id(current_tenant_id, session=session) + if tenant is None: + raise NotFound() + + custom_config = tenant.custom_config_dict + replace_webapp_logo = ( + f"{dify_config.FILES_URL}/files/workspaces/{tenant.id}/webapp-logo" + if custom_config.get("replace_webapp_logo") + else None + ) + return dump_response( + WorkspaceCustomConfigResponse, + { + "remove_webapp_brand": custom_config.get("remove_webapp_brand", False), + "replace_webapp_logo": replace_webapp_logo, + }, + ) + @console_ns.expect(console_ns.models[WorkspaceCustomConfigPayload.__name__]) @console_ns.response(HTTPStatus.OK, "Success", console_ns.models[WorkspaceTenantResultResponse.__name__]) @setup_required diff --git a/api/controllers/console/wraps.py b/api/controllers/console/wraps.py index b77c78761a0..8e353c6c980 100644 --- a/api/controllers/console/wraps.py +++ b/api/controllers/console/wraps.py @@ -1,6 +1,5 @@ import contextlib import json -import os import time from collections.abc import Callable from functools import wraps @@ -19,8 +18,7 @@ from controllers.common.wraps import ( ) from controllers.console.auth.error import AuthenticationFailedError, EmailCodeError from controllers.console.workspace.error import AccountNotInitializedError -from enums.cloud_plan import CloudPlan -from enums.deployment_edition import DeploymentEdition +from enums import CloudPlan, DeploymentEdition from extensions.ext_database import db from extensions.ext_redis import redis_client from libs.encryption import FieldEncryption @@ -30,7 +28,8 @@ from models.account import AccountStatus from models.dataset import RateLimitLog from models.model import DifySetup from services.billing_service import BillingService -from services.feature_service import FeatureService, LicenseStatus +from services.entities.feature_entities import LicenseStatus +from services.feature_service import FeatureService from services.operation_service import OperationService, UtmInfo from .error import NotInitValidateError, NotSetupError, UnauthorizedAndForceLogout @@ -141,7 +140,7 @@ def only_edition_cloud[**P, R](view: Callable[P, R]) -> Callable[P, R]: def only_edition_enterprise[**P, R](view: Callable[P, R]) -> Callable[P, R]: @wraps(view) def decorated(*args: P.args, **kwargs: P.kwargs): - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: abort(404) return view(*args, **kwargs) @@ -160,16 +159,6 @@ def only_edition_self_hosted[**P, R](view: Callable[P, R]) -> Callable[P, R]: return decorated -def cloud_edition_billing_enabled[**P, R](view: Callable[P, R]) -> Callable[P, R]: - @wraps(view) - def decorated(*args: P.args, **kwargs: P.kwargs): - if not dify_config.BILLING_ENABLED: - abort(403, "Billing feature is not enabled.") - return view(*args, **kwargs) - - return decorated - - def cloud_edition_billing_paid_plan_required[**P, R](view: Callable[P, R]) -> Callable[P, R]: @wraps(view) def decorated(*args: P.args, **kwargs: P.kwargs): @@ -191,7 +180,7 @@ def cloud_edition_billing_resource_check[**P, R](resource: str) -> Callable[[Cal def decorated(*args: P.args, **kwargs: P.kwargs): _, current_tenant_id = current_account_with_tenant() if resource == "vector_space": - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return view(*args, **kwargs) vector_space = FeatureService.get_vector_space(current_tenant_id) @@ -215,7 +204,7 @@ def cloud_edition_billing_resource_check[**P, R](resource: str) -> Callable[[Cal elif resource == "documents" and 0 < documents_upload_quota.limit <= documents_upload_quota.size: # The api of file upload is used in the multiple places, # so we need to check the source of the request from datasets - source = request.args.get("source") + source = request.args.get("source") or request.form.get("source") if source == "datasets": abort(403, "The number of documents has reached the limit of your subscription.") else: @@ -259,35 +248,35 @@ def cloud_edition_billing_knowledge_limit_check[**P, R]( return interceptor +def check_knowledge_rate_limit() -> None: + _, current_tenant_id = current_account_with_tenant() + knowledge_rate_limit = FeatureService.get_knowledge_rate_limit(current_tenant_id) + if not knowledge_rate_limit.enabled: + return + + current_time = int(time.time() * 1000) + key = f"rate_limit_{current_tenant_id}" + redis_client.zadd(key, {current_time: current_time}) + redis_client.zremrangebyscore(key, 0, current_time - 60000) + + if redis_client.zcard(key) > knowledge_rate_limit.limit: + db.session.add( # guard-ignore: no-new-controller-sqlalchemy -- existing decorator audit write + RateLimitLog( + tenant_id=current_tenant_id, + subscription_plan=knowledge_rate_limit.subscription_plan, + operation="knowledge", + ) + ) + db.session.commit() + abort(403, "Sorry, you have reached the knowledge base request rate limit of your subscription.") + + def cloud_edition_billing_rate_limit_check[**P, R](resource: str) -> Callable[[Callable[P, R]], Callable[P, R]]: def interceptor(view: Callable[P, R]): @wraps(view) def decorated(*args: P.args, **kwargs: P.kwargs): if resource == "knowledge": - _, current_tenant_id = current_account_with_tenant() - knowledge_rate_limit = FeatureService.get_knowledge_rate_limit(current_tenant_id) - if knowledge_rate_limit.enabled: - current_time = int(time.time() * 1000) - key = f"rate_limit_{current_tenant_id}" - - redis_client.zadd(key, {current_time: current_time}) - - redis_client.zremrangebyscore(key, 0, current_time - 60000) - - request_count = redis_client.zcard(key) - - if request_count > knowledge_rate_limit.limit: - # add ratelimit record - rate_limit_log = RateLimitLog( - tenant_id=current_tenant_id, - subscription_plan=knowledge_rate_limit.subscription_plan, - operation="knowledge", - ) - db.session.add(rate_limit_log) - db.session.commit() - abort( - 403, "Sorry, you have reached the knowledge base request rate limit of your subscription." - ) + check_knowledge_rate_limit() return view(*args, **kwargs) return decorated @@ -300,7 +289,7 @@ def cloud_utm_record[**P, R](view: Callable[P, R]) -> Callable[P, R]: def decorated(*args: P.args, **kwargs: P.kwargs): with contextlib.suppress(Exception): utm_info = request.cookies.get("utm_info") - if dify_config.BILLING_ENABLED and utm_info: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and utm_info: _, current_tenant_id = current_account_with_tenant() utm_info_dict: UtmInfo = json.loads(utm_info) OperationService.record_utm(current_tenant_id, utm_info_dict) @@ -329,7 +318,7 @@ def setup_required[R](view: Callable[..., R]) -> Callable[..., R]: # preserving support for plain functions used in tests and utilities. # check setup if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD and not _is_setup_completed(): - if os.environ.get("INIT_PASSWORD"): + if dify_config.INIT_PASSWORD: raise NotInitValidateError() raise NotSetupError() diff --git a/api/controllers/files/upload.py b/api/controllers/files/upload.py index 3de92589f6a..82a7ae4fb65 100644 --- a/api/controllers/files/upload.py +++ b/api/controllers/files/upload.py @@ -1,3 +1,5 @@ +from typing import Literal + from flask import request from flask_restx import Resource from flask_restx.api import HTTPStatus @@ -5,10 +7,12 @@ from pydantic import BaseModel, Field from werkzeug.exceptions import Forbidden import services +from core.db.session_factory import session_factory from core.tools.signature import verify_plugin_file_signature from core.tools.tool_file_manager import ToolFileManager, resolve_extension from core.workflow.file_reference import build_file_reference from fields.file_fields import FileResponse +from services.account_service import TenantService from ..common.errors import ( FileTooLargeError, @@ -26,6 +30,7 @@ class PluginUploadQuery(BaseModel): sign: str = Field(..., description="HMAC signature") tenant_id: str = Field(..., description="Tenant identifier") user_id: str | None = Field(default=None, description="User identifier") + user_from: Literal["account", "end-user"] | None = Field(default=None, description="User identity type") conversation_id: str | None = Field(default=None, description="Conversation identifier") @@ -77,8 +82,20 @@ class PluginUploadFileApi(Resource): nonce = args.nonce sign = args.sign tenant_id = args.tenant_id - user_id = args.user_id - user = get_user(tenant_id, user_id) + if args.user_from == "account": + if args.user_id is None: + raise Forbidden("Invalid request.") + with session_factory.create_session() as session: + is_tenant_member = TenantService.account_belongs_to_tenant( + args.user_id, + tenant_id, + session=session, + ) + if not is_tenant_member: + raise Forbidden("Invalid request.") + owner_id = args.user_id + else: + owner_id = get_user(tenant_id, args.user_id).id filename = file.filename mimetype = file.mimetype @@ -90,8 +107,9 @@ class PluginUploadFileApi(Resource): filename=filename, mimetype=mimetype, tenant_id=tenant_id, - user_id=user.id, + user_id=owner_id, conversation_id=args.conversation_id, + user_from=args.user_from, timestamp=timestamp, nonce=nonce, sign=sign, @@ -100,7 +118,7 @@ class PluginUploadFileApi(Resource): try: tool_file = ToolFileManager().create_file_by_raw( - user_id=user.id, + user_id=owner_id, tenant_id=tenant_id, file_binary=file.stream.read(), mimetype=mimetype, diff --git a/api/controllers/inner_api/__init__.py b/api/controllers/inner_api/__init__.py index f47861cf274..277e0f5ec0e 100644 --- a/api/controllers/inner_api/__init__.py +++ b/api/controllers/inner_api/__init__.py @@ -17,6 +17,7 @@ inner_api_ns = Namespace("inner_api", description="Internal API operations", pat from . import mail as _mail from . import runtime_credentials as _runtime_credentials +from .agent import files as _agent_files from .agent import tools as _agent_tools from .app import dsl as _app_dsl from .knowledge import retrieval as _knowledge_retrieval @@ -30,6 +31,7 @@ api.add_namespace(inner_api_ns) __all__ = [ "_agent_config", "_agent_drive", + "_agent_files", "_agent_tools", "_app_dsl", "_knowledge_retrieval", diff --git a/api/controllers/inner_api/agent/files.py b/api/controllers/inner_api/agent/files.py new file mode 100644 index 00000000000..8838ba2d4fe --- /dev/null +++ b/api/controllers/inner_api/agent/files.py @@ -0,0 +1,205 @@ +"""Agent-owned inner endpoints for CLI file URL allocation.""" + +from __future__ import annotations + +from typing import Literal + +from flask_restx import Resource +from pydantic import BaseModel, ConfigDict, ValidationError +from sqlalchemy.orm import Session + +from configs import dify_config +from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.common.session import with_session +from controllers.console.wraps import setup_required +from controllers.inner_api import inner_api_ns +from controllers.inner_api.plugin.wraps import get_user +from controllers.inner_api.wraps import plugin_inner_api_only +from core.plugin.entities.request import RequestDownloadFileMapping, RequestRequestUploadFile +from core.tools.signature import bind_file_uri, get_signed_file_uri_for_plugin +from fields.base import ResponseModel +from libs.exception import BaseHTTPException +from services.account_service import TenantService +from services.file_request_service import FileRequestService + + +class AgentFileRequestHttpError(BaseHTTPException): + error_code = "agent_file_request_failed" + description = "Agent file request failed." + code = 500 + + def __init__(self, *, error_code: str, description: str, status_code: int) -> None: + self.error_code = error_code + self.description = description + self.code = status_code + super().__init__(description) + + +class AgentFileUploadRequestPayload(RequestRequestUploadFile): + tenant_id: str + user_id: str + user_from: Literal["account", "end-user"] | None = None + + model_config = ConfigDict(extra="forbid") + + +class AgentFileDownloadRequestPayload(BaseModel): + tenant_id: str + user_id: str + user_from: Literal["account", "end-user"] + invoke_from: Literal[ + "service-api", + "openapi", + "web-app", + "trigger", + "explore", + "debugger", + "published", + "validation", + ] + file: RequestDownloadFileMapping + for_frontend: bool = True + + model_config = ConfigDict(extra="forbid") + + +class AgentFileUploadRequestResponse(ResponseModel): + upload_uri: str + + +class AgentFileDownloadRequestResponse(ResponseModel): + filename: str + mime_type: str | None = None + size: int + download_uri: str + + +register_schema_models(inner_api_ns, AgentFileUploadRequestPayload, AgentFileDownloadRequestPayload) +register_response_schema_models( + inner_api_ns, + AgentFileUploadRequestResponse, + AgentFileDownloadRequestResponse, +) + + +@inner_api_ns.route("/agent/files/upload-request") +class AgentFileUploadRequestApi(Resource): + """Allocate an origin-free signed upload URI for the Agent CLI.""" + + @setup_required + @plugin_inner_api_only + @inner_api_ns.doc("inner_agent_file_upload_request") + @inner_api_ns.expect(inner_api_ns.models[AgentFileUploadRequestPayload.__name__]) + @inner_api_ns.response( + 200, + "Upload URI allocated", + inner_api_ns.models[AgentFileUploadRequestResponse.__name__], + ) + @with_session(write=False) + def post(self, session: Session) -> dict[str, object]: + try: + payload = AgentFileUploadRequestPayload.model_validate(inner_api_ns.payload or {}) + except ValidationError as exc: + raise AgentFileRequestHttpError( + error_code="invalid_request", + description=str(exc), + status_code=400, + ) from exc + + tenant = TenantService.get_tenant_by_id(payload.tenant_id, session=session) + if tenant is None: + raise AgentFileRequestHttpError( + error_code="tenant_not_found", + description="tenant not found", + status_code=404, + ) + try: + if payload.user_from == "account": + if not TenantService.account_belongs_to_tenant(payload.user_id, tenant.id, session=session): + raise ValueError("account not found") + owner_id = payload.user_id + else: + owner_id = get_user(tenant.id, payload.user_id).id + upload_uri = get_signed_file_uri_for_plugin( + filename=payload.filename, + mimetype=payload.mimetype, + tenant_id=tenant.id, + user_id=owner_id, + conversation_id=payload.conversation_id, + user_from=payload.user_from, + ) + except ValueError as exc: + raise AgentFileRequestHttpError( + error_code="user_not_found", + description=str(exc), + status_code=404, + ) from exc + + return AgentFileUploadRequestResponse(upload_uri=upload_uri).model_dump(mode="json") + + +@inner_api_ns.route("/agent/files/download-request") +class AgentFileDownloadRequestApi(Resource): + """Allocate a transfer URI or frontend URL for one Agent CLI file.""" + + @setup_required + @plugin_inner_api_only + @inner_api_ns.doc("inner_agent_file_download_request") + @inner_api_ns.expect(inner_api_ns.models[AgentFileDownloadRequestPayload.__name__]) + @inner_api_ns.response( + 200, + "Download URI allocated", + inner_api_ns.models[AgentFileDownloadRequestResponse.__name__], + ) + @with_session(write=False) + def post(self, session: Session) -> dict[str, object]: + try: + payload = AgentFileDownloadRequestPayload.model_validate(inner_api_ns.payload or {}) + except ValidationError as exc: + raise AgentFileRequestHttpError( + error_code="invalid_request", + description=str(exc), + status_code=400, + ) from exc + + if TenantService.get_tenant_by_id(payload.tenant_id, session=session) is None: + raise AgentFileRequestHttpError( + error_code="tenant_not_found", + description="tenant not found", + status_code=404, + ) + try: + result = FileRequestService().request_download( + tenant_id=payload.tenant_id, + user_id=payload.user_id, + user_from=payload.user_from, + invoke_from=payload.invoke_from, + file_mapping=payload.file.model_dump(mode="python", exclude_none=True), + ) + except ValueError as exc: + raise AgentFileRequestHttpError( + error_code="file_not_accessible", + description=str(exc), + status_code=404, + ) from exc + + download_uri = result.download_uri + if payload.for_frontend: + download_uri = bind_file_uri(download_uri, dify_config.FILES_URL) + + return AgentFileDownloadRequestResponse( + filename=result.filename, + mime_type=result.mime_type, + size=result.size, + download_uri=download_uri, + ).model_dump(mode="json") + + +__all__ = [ + "AgentFileDownloadRequestApi", + "AgentFileDownloadRequestPayload", + "AgentFileDownloadRequestResponse", + "AgentFileUploadRequestApi", + "AgentFileUploadRequestPayload", + "AgentFileUploadRequestResponse", +] diff --git a/api/controllers/inner_api/app/dsl.py b/api/controllers/inner_api/app/dsl.py index 8b206cc36be..e4276d44f20 100644 --- a/api/controllers/inner_api/app/dsl.py +++ b/api/controllers/inner_api/app/dsl.py @@ -5,13 +5,15 @@ to attribute the created app; workspace/membership validation is done by the Go admin-api caller. """ +from uuid import UUID + from flask import request from flask_restx import Resource -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, ValidationError, field_validator from sqlalchemy import select from sqlalchemy.orm import Session -from controllers.common.schema import register_schema_model +from controllers.common.schema import query_params_from_model, register_schema_model from controllers.console.wraps import setup_required from controllers.inner_api import inner_api_ns from controllers.inner_api.wraps import enterprise_inner_api_only @@ -20,6 +22,7 @@ from models import Account, App from models.account import AccountStatus from services.app_dsl_service import AppDslService from services.entities.dsl_entities import ImportMode, ImportStatus +from services.errors.app import IsDraftWorkflowError, WorkflowNotFoundError class InnerAppDSLImportPayload(BaseModel): @@ -29,6 +32,18 @@ class InnerAppDSLImportPayload(BaseModel): description: str | None = Field(default=None, description="Override app description from DSL") +class EnterpriseAppDSLExportQuery(BaseModel): + include_secret: bool = Field(default=False, description="Whether to include secret values in the exported DSL") + workflow_id: UUID | None = Field(default=None, description="Published workflow version ID to export") + + @field_validator("include_secret", mode="before") + @classmethod + def parse_include_secret(cls, value: object) -> bool: + if isinstance(value, str): + return value.lower() == "true" + return bool(value) + + register_schema_model(inner_api_ns, InnerAppDSLImportPayload) @@ -82,24 +97,48 @@ class EnterpriseAppDSLExport(Resource): @enterprise_inner_api_only @inner_api_ns.doc( "enterprise_app_dsl_export", + params=query_params_from_model(EnterpriseAppDSLExportQuery), responses={ 200: "Export successful", - 404: "App not found", + 400: "Invalid workflow ID or unpublished workflow version", + 404: "App or workflow version not found", }, ) def get(self, app_id: str): """Export an app's DSL as YAML.""" - include_secret = request.args.get("include_secret", "false").lower() == "true" + try: + query = EnterpriseAppDSLExportQuery.model_validate(request.args.to_dict(flat=True)) + except ValidationError: + return { + "code": "invalid_workflow_id", + "message": "workflow_id must be a valid UUID", + "status": 400, + }, 400 + + workflow_id = str(query.workflow_id) if query.workflow_id else None app_model = db.session.get(App, app_id) if not app_model: return {"message": "app not found"}, 404 - data = AppDslService.export_dsl( - app_model=app_model, - session=db.session(), - include_secret=include_secret, - ) + if not workflow_id: + data = AppDslService.export_dsl( + app_model=app_model, + session=db.session(), + include_secret=query.include_secret, + ) + else: + try: + data = AppDslService.export_dsl( + app_model=app_model, + session=db.session(), + include_secret=query.include_secret, + workflow_id=workflow_id, + ) + except WorkflowNotFoundError as exc: + return {"code": "workflow_version_not_found", "message": str(exc), "status": 404}, 404 + except IsDraftWorkflowError as exc: + return {"code": "workflow_version_not_published", "message": str(exc), "status": 400}, 400 return {"data": data}, 200 diff --git a/api/controllers/inner_api/plugin/agent_config.py b/api/controllers/inner_api/plugin/agent_config.py index 6f12f6e5d02..196909196bd 100644 --- a/api/controllers/inner_api/plugin/agent_config.py +++ b/api/controllers/inner_api/plugin/agent_config.py @@ -2,20 +2,24 @@ These endpoints are called by the dify-agent server with the inner API key. They resolve the requested Agent config version directly from Agent Soul JSON -and never expose signed download URLs or drive-owned metadata. +and authorize downloads with the existing Config target, tenant, and source +ownership semantics. Download requests return metadata plus a short-lived, +origin-free ``/files/*`` URI, never file bytes; the Sandbox fetches those bytes +directly from the Dify API data plane. """ from __future__ import annotations -import io +from typing import Literal -from flask import request, send_file +from flask import request from flask_restx import Resource -from pydantic import BaseModel, ValidationError +from pydantic import BaseModel, ConfigDict, ValidationError, model_validator from controllers.console.wraps import setup_required from controllers.inner_api import inner_api_ns from controllers.inner_api.wraps import plugin_inner_api_only +from models.agent_config_entities import validate_config_name, validate_config_skill_name from services.agent_config_service import ( AgentConfigService, AgentConfigServiceError, @@ -38,6 +42,24 @@ class _ConfigMutationRequest(BaseModel): config_version_kind: AgentConfigVersionKind +class _ConfigDownloadSource(BaseModel): + model_config = ConfigDict(extra="forbid") + + kind: Literal["file", "skill"] + name: str + + @model_validator(mode="after") + def validate_name(self) -> _ConfigDownloadSource: + self.name = validate_config_skill_name(self.name) if self.kind == "skill" else validate_config_name(self.name) + return self + + +class _ConfigDownloadRequest(_ConfigTargetQuery): + model_config = ConfigDict(extra="forbid") + + config: _ConfigDownloadSource + + class _ConfigPushRequest(_ConfigMutationRequest): files: list[dict] = [] skills: list[dict] = [] @@ -99,28 +121,29 @@ class AgentConfigManifestApi(Resource): return _error_response(exc) -@inner_api_ns.route("/agent-config//skills//pull") -class AgentConfigSkillPullApi(Resource): +@inner_api_ns.route("/agent-config//download-request") +class AgentConfigDownloadRequestApi(Resource): @setup_required @plugin_inner_api_only - @inner_api_ns.doc("agent_config_skill_pull") - def get(self, agent_id: str, name: str): + @inner_api_ns.doc("agent_config_download_request") + def post(self, agent_id: str): try: - query = _target_query_from_request() - result = AgentConfigService().pull_skill( - tenant_id=query.tenant_id, + body = _ConfigDownloadRequest.model_validate(request.get_json(silent=True) or {}) + result = AgentConfigService().request_download( + tenant_id=body.tenant_id, agent_id=agent_id, - user_id=query.user_id, - config_version_id=query.config_version_id, - config_version_kind=query.config_version_kind, - name=name, - ) - return send_file( - io.BytesIO(result.payload), - mimetype=result.mime_type, - as_attachment=True, - download_name=result.filename, + user_id=body.user_id, + config_version_id=body.config_version_id, + config_version_kind=body.config_version_kind, + kind=body.config.kind, + name=body.config.name, ) + return { + "filename": result.filename, + "mime_type": result.mime_type, + "size": result.size, + "download_uri": result.download_uri, + } except ValidationError as exc: return {"code": "invalid_request", "message": str(exc)}, 400 except AgentConfigServiceError as exc: @@ -149,34 +172,6 @@ class AgentConfigSkillInspectApi(Resource): return _error_response(exc) -@inner_api_ns.route("/agent-config//files//pull") -class AgentConfigFilePullApi(Resource): - @setup_required - @plugin_inner_api_only - @inner_api_ns.doc("agent_config_file_pull") - def get(self, agent_id: str, name: str): - try: - query = _target_query_from_request() - result = AgentConfigService().pull_file( - tenant_id=query.tenant_id, - agent_id=agent_id, - user_id=query.user_id, - config_version_id=query.config_version_id, - config_version_kind=query.config_version_kind, - name=name, - ) - return send_file( - io.BytesIO(result.payload), - mimetype=result.mime_type, - as_attachment=True, - download_name=result.filename, - ) - except ValidationError as exc: - return {"code": "invalid_request", "message": str(exc)}, 400 - except AgentConfigServiceError as exc: - return _error_response(exc) - - @inner_api_ns.route("/agent-config//push") class AgentConfigPushApi(Resource): @setup_required diff --git a/api/controllers/inner_api/plugin/plugin.py b/api/controllers/inner_api/plugin/plugin.py index 1171761fc53..221887f73c7 100644 --- a/api/controllers/inner_api/plugin/plugin.py +++ b/api/controllers/inner_api/plugin/plugin.py @@ -1,6 +1,7 @@ from flask_restx import Resource from sqlalchemy.orm import Session +from configs import dify_config from controllers.console.app.wraps import with_session from controllers.console.wraps import setup_required from controllers.inner_api import inner_api_ns @@ -31,7 +32,7 @@ from core.plugin.entities.request import ( RequestRequestUploadFile, ) from core.tools.entities.tool_entities import ToolProviderType -from core.tools.signature import get_signed_file_url_for_plugin +from core.tools.signature import bind_file_uri, get_signed_file_uri_for_plugin from extensions.ext_database import db from graphon.model_runtime.utils.encoders import jsonable_encoder from libs.helper import length_prefixed_response @@ -429,13 +430,14 @@ class PluginUploadFileRequestApi(Resource): ) def post(self, user_model: Account | EndUser, tenant_model: Tenant, payload: RequestRequestUploadFile): # generate signed url - url = get_signed_file_url_for_plugin( + uri = get_signed_file_uri_for_plugin( filename=payload.filename, mimetype=payload.mimetype, tenant_id=tenant_model.id, user_id=user_model.id, conversation_id=payload.conversation_id, ) + url = bind_file_uri(uri, dify_config.INTERNAL_FILES_URL or dify_config.FILES_URL) return BaseBackwardsInvocationResponse(data={"url": url}).model_dump() @@ -454,36 +456,33 @@ class PluginDownloadFileRequestApi(Resource): } ) def post(self, payload: RequestRequestDownloadFile): - """Resolve signed download metadata for trusted external runtimes. + """Adapt a shared file request result to the Plugin backward contract. - Unlike end-user-facing upload/download APIs, this inner endpoint serves - trusted callers such as the ``dify-agent`` back proxy. The caller sends - flattened ``tenant_id`` / ``user_id`` / ``user_from`` / ``invoke_from`` - context explicitly in the body, and ``FileRequestService`` rebuilds the - corresponding ``FileAccessScope`` before resolving the signed URL. - - The response is control-plane metadata only: filename, mime type, size, - and the signed download URL. File bytes still flow through the existing - signed file endpoints rather than through this inner API. + ``FileRequestService`` rebuilds the caller's ``FileAccessScope`` and + resolves one origin-free signed URI. This controller binds that URI to + the Plugin-selected external or internal files base URL, then returns + the existing backward-invocation envelope. """ tenant_model = db.session.get(Tenant, payload.tenant_id) if tenant_model is None: raise ValueError("tenant not found") - result = FileRequestService().request_download_url( + result = FileRequestService().request_download( tenant_id=tenant_model.id, user_id=payload.user_id, user_from=payload.user_from, invoke_from=payload.invoke_from, file_mapping=payload.file.model_dump(mode="python", exclude_none=True), - for_external=payload.for_external, + ) + base_url = ( + dify_config.FILES_URL if payload.for_external else (dify_config.INTERNAL_FILES_URL or dify_config.FILES_URL) ) return BaseBackwardsInvocationResponse( data={ "filename": result.filename, "mime_type": result.mime_type, "size": result.size, - "download_url": result.download_url, + "download_url": bind_file_uri(result.download_uri, base_url), } ).model_dump() diff --git a/api/controllers/mcp/mcp.py b/api/controllers/mcp/mcp.py index 45a5a4e6899..95aa2be8218 100644 --- a/api/controllers/mcp/mcp.py +++ b/api/controllers/mcp/mcp.py @@ -240,9 +240,9 @@ class MCPAppApi(Resource): with sessionmaker(db.engine, expire_on_commit=False).begin() as session: return session.scalar( select(EndUser) - .where(EndUser.tenant_id == tenant_id) - .where(EndUser.session_id == mcp_server_id) - .where(EndUser.type == EndUserType.MCP) + .where( + EndUser.tenant_id == tenant_id, EndUser.session_id == mcp_server_id, EndUser.type == EndUserType.MCP + ) .limit(1) ) diff --git a/api/controllers/openapi/__init__.py b/api/controllers/openapi/__init__.py index 0260422ec1c..681f93e49ba 100644 --- a/api/controllers/openapi/__init__.py +++ b/api/controllers/openapi/__init__.py @@ -139,7 +139,6 @@ register_response_schema_models( register_enum_models(openapi_ns, OpenApiErrorCode) from . import ( - _meta, account, app_dsl, app_run, @@ -157,7 +156,6 @@ from . import ( # Request models are imported from _models.py and registered above. __all__ = [ - "_meta", "account", "app_dsl", "app_run", diff --git a/api/controllers/openapi/_meta.py b/api/controllers/openapi/_meta.py deleted file mode 100644 index c49f7526acc..00000000000 --- a/api/controllers/openapi/_meta.py +++ /dev/null @@ -1,24 +0,0 @@ -"""Meta endpoint: `GET /openapi/v1/_version` — no auth. - -Returns the server's project version and edition so the difyctl CLI can probe -compatibility without needing to be logged in. Mirrors the `_health` endpoint -in `index.py`. -""" - -from flask_restx import Resource - -from configs import dify_config -from controllers.openapi import openapi_ns -from controllers.openapi._contract import returns -from controllers.openapi._models import ServerVersionResponse - - -@openapi_ns.route("/_version") -class VersionApi(Resource): - @returns(200, ServerVersionResponse, description="Server version") - def get(self): - edition = dify_config.EDITION if dify_config.EDITION in ("SELF_HOSTED", "CLOUD") else "SELF_HOSTED" - return ServerVersionResponse( - version=dify_config.project.version, - edition=edition, - ) diff --git a/api/controllers/openapi/_models.py b/api/controllers/openapi/_models.py index 5337612e7b6..c5e8a22d466 100644 --- a/api/controllers/openapi/_models.py +++ b/api/controllers/openapi/_models.py @@ -7,6 +7,7 @@ from typing import Any, Final, Literal from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator +from enums import DeploymentEdition from libs.helper import EmailStr, UUIDStr, UUIDStrOrEmpty, uuid_value from models.model import AppMode @@ -258,7 +259,7 @@ class ServerVersionResponse(BaseModel): """Meta endpoint payload for `GET /openapi/v1/_version` — no auth required.""" version: str - edition: Literal["SELF_HOSTED", "CLOUD"] + edition: DeploymentEdition class HealthResponse(BaseModel): diff --git a/api/controllers/openapi/apps.py b/api/controllers/openapi/apps.py index 9a3e609b70a..f6dace6fd56 100644 --- a/api/controllers/openapi/apps.py +++ b/api/controllers/openapi/apps.py @@ -3,7 +3,7 @@ from __future__ import annotations import uuid as _uuid -from typing import Any, cast +from typing import Any from flask_restx import Resource from sqlalchemy.orm import Session @@ -247,8 +247,8 @@ class AppListApi(Resource): env = AppListResponse( page=query.page, limit=query.limit, - total=cast(int, pagination.total), - has_more=query.page * query.limit < cast(int, pagination.total), + total=pagination.total, + has_more=query.page * query.limit < pagination.total, data=items, ) return env diff --git a/api/controllers/openapi/apps_permitted_external.py b/api/controllers/openapi/apps_permitted_external.py index e00ec5c87f7..a56ab4ae621 100644 --- a/api/controllers/openapi/apps_permitted_external.py +++ b/api/controllers/openapi/apps_permitted_external.py @@ -23,7 +23,8 @@ from controllers.openapi._models import ( ) from controllers.openapi.apps import build_app_describe_response from controllers.openapi.auth.composition import auth_router -from controllers.openapi.auth.data import AuthData, Edition +from controllers.openapi.auth.data import AuthData +from enums import DeploymentEdition from libs.oauth_bearer import Scope, TokenType from models import App from models.enums import AppStatus @@ -37,7 +38,7 @@ class PermittedExternalAppsListApi(Resource): @auth_router.guard( scope=Scope.APPS_READ_PERMITTED_EXTERNAL, allowed_token_types=frozenset({TokenType.OAUTH_EXTERNAL_SSO}), - edition=frozenset({Edition.EE}), + edition=frozenset({DeploymentEdition.ENTERPRISE}), ) @returns(200, PermittedExternalAppsListResponse, description="Permitted external apps list") @accepts(query=PermittedExternalAppsListQuery) @@ -94,7 +95,7 @@ class PermittedExternalAppDescribeApi(Resource): @auth_router.guard( scope=Scope.APPS_READ_PERMITTED_EXTERNAL, allowed_token_types=frozenset({TokenType.OAUTH_EXTERNAL_SSO}), - edition=frozenset({Edition.EE}), + edition=frozenset({DeploymentEdition.ENTERPRISE}), ) @returns(200, AppDescribeResponse, description="Permitted external app description") @accepts(query=AppDescribeQuery) diff --git a/api/controllers/openapi/auth/composition.py b/api/controllers/openapi/auth/composition.py index 67f7001c080..9040058b85f 100644 --- a/api/controllers/openapi/auth/composition.py +++ b/api/controllers/openapi/auth/composition.py @@ -1,7 +1,7 @@ from __future__ import annotations from controllers.openapi.auth.conditions import ( - EDITION_EE, + EDITION_ENTERPRISE, HAS_ALLOWED_ROLES, HAS_RBAC, LOADED_APP_IS_PRIVATE, @@ -11,7 +11,6 @@ from controllers.openapi.auth.conditions import ( WORKSPACE_MEMBERSHIP_REQUIRED, WORKSPACE_SCOPED, ) -from controllers.openapi.auth.data import Edition from controllers.openapi.auth.flow import When from controllers.openapi.auth.pipeline import AuthPipeline, PipelineRoute, PipelineRouter from controllers.openapi.auth.prepare import ( @@ -33,6 +32,7 @@ from controllers.openapi.auth.verify import ( check_workspace_mismatch, check_workspace_role, ) +from enums import DeploymentEdition from libs.oauth_bearer import TokenType account_pipeline = AuthPipeline( @@ -42,7 +42,7 @@ account_pipeline = AuthPipeline( When(WORKSPACE_MEMBERSHIP_REQUIRED, then=load_tenant_from_request), load_account, When(WORKSPACE_SCOPED, then=load_workspace_role), - When(PATH_HAS_APP_ID & EDITION_EE, then=load_app_access_mode), + When(PATH_HAS_APP_ID & EDITION_ENTERPRISE, then=load_app_access_mode), ], auth=[ When(PATH_HAS_APP_ID, then=check_app_api_enabled), @@ -51,8 +51,8 @@ account_pipeline = AuthPipeline( When(PATH_HAS_APP_ID, then=check_workspace_mismatch), When(HAS_ALLOWED_ROLES, then=check_workspace_role), When(HAS_RBAC, then=check_rbac_permission), - When(PATH_HAS_APP_ID & EDITION_EE & WEBAPP_AUTH_ENABLED & WEBAPP_RUN_SCOPED, then=check_acl), - When(EDITION_EE & LOADED_APP_IS_PRIVATE & WEBAPP_RUN_SCOPED, then=check_private_app_permission), + When(PATH_HAS_APP_ID & EDITION_ENTERPRISE & WEBAPP_AUTH_ENABLED & WEBAPP_RUN_SCOPED, then=check_acl), + When(EDITION_ENTERPRISE & LOADED_APP_IS_PRIVATE & WEBAPP_RUN_SCOPED, then=check_private_app_permission), ], ) @@ -74,6 +74,9 @@ external_sso_pipeline = AuthPipeline( auth_router = PipelineRouter( { TokenType.OAUTH_ACCOUNT: PipelineRoute(account_pipeline), - TokenType.OAUTH_EXTERNAL_SSO: PipelineRoute(external_sso_pipeline, required_edition=frozenset({Edition.EE})), + TokenType.OAUTH_EXTERNAL_SSO: PipelineRoute( + external_sso_pipeline, + required_edition=frozenset({DeploymentEdition.ENTERPRISE}), + ), } ) diff --git a/api/controllers/openapi/auth/conditions.py b/api/controllers/openapi/auth/conditions.py index 73a767b8d8e..a25eaf78aa1 100644 --- a/api/controllers/openapi/auth/conditions.py +++ b/api/controllers/openapi/auth/conditions.py @@ -2,7 +2,9 @@ from __future__ import annotations from collections.abc import Callable -from controllers.openapi.auth.data import AuthData, Edition, RequestContext, current_edition +from configs import dify_config +from controllers.openapi.auth.data import AuthData, RequestContext +from enums import DeploymentEdition from libs.oauth_bearer import Scope, TokenType from services.enterprise.enterprise_service import WebAppAccessMode from services.feature_service import FeatureService @@ -44,9 +46,9 @@ TOKEN_IS_OAUTH_EXTERNAL_SSO = request_cond(lambda ctx: ctx.token_type == TokenTy PATH_HAS_APP_ID = request_cond(lambda ctx: "app_id" in ctx.path_params) -EDITION_CE = config_cond(lambda: current_edition() == Edition.CE) -EDITION_EE = config_cond(lambda: current_edition() == Edition.EE) -EDITION_SAAS = config_cond(lambda: current_edition() == Edition.SAAS) +EDITION_COMMUNITY = config_cond(lambda: dify_config.DEPLOYMENT_EDITION == DeploymentEdition.COMMUNITY) +EDITION_ENTERPRISE = config_cond(lambda: dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE) +EDITION_CLOUD = config_cond(lambda: dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD) WEBAPP_AUTH_ENABLED = config_cond(lambda: FeatureService.get_system_features().webapp_auth.enabled) diff --git a/api/controllers/openapi/auth/data.py b/api/controllers/openapi/auth/data.py index 898a9eb8f87..79d139841aa 100644 --- a/api/controllers/openapi/auth/data.py +++ b/api/controllers/openapi/auth/data.py @@ -6,34 +6,18 @@ from enum import StrEnum from pydantic import BaseModel, ConfigDict, Field from werkzeug.exceptions import InternalServerError -from configs import dify_config from core.rbac import RBACPermission, RBACResourceScope -from enums.deployment_edition import DeploymentEdition from libs.oauth_bearer import Scope, TokenType from models.account import Account, Tenant, TenantAccountRole from models.model import App, EndUser from services.enterprise.enterprise_service import WebAppAccessMode -class Edition(StrEnum): - CE = "ce" - EE = "ee" - SAAS = "saas" - - class CallerKind(StrEnum): ACCOUNT = "account" END_USER = "end_user" -def current_edition() -> Edition: - if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: - return Edition.SAAS - if dify_config.ENTERPRISE_ENABLED: - return Edition.EE - return Edition.CE - - class ExternalIdentity(BaseModel): model_config = ConfigDict(frozen=True) diff --git a/api/controllers/openapi/auth/pipeline.py b/api/controllers/openapi/auth/pipeline.py index 3e0aca53d3c..f27064eda96 100644 --- a/api/controllers/openapi/auth/pipeline.py +++ b/api/controllers/openapi/auth/pipeline.py @@ -16,16 +16,16 @@ from flask import current_app, request from flask_login import user_logged_in from werkzeug.exceptions import Forbidden, NotFound, Unauthorized +from configs import dify_config from controllers.openapi._audit import emit_wrong_surface from controllers.openapi.auth.data import ( AuthData, - Edition, ExternalIdentity, RBACRequirement, RequestContext, - current_edition, ) from controllers.openapi.auth.flow import When +from enums import DeploymentEdition from libs.oauth_bearer import ( AuthContext, Scope, @@ -36,7 +36,8 @@ from libs.oauth_bearer import ( set_auth_ctx, ) from models.account import TenantAccountRole -from services.feature_service import FeatureService, LicenseStatus +from services.entities.feature_entities import LicenseStatus +from services.feature_service import FeatureService class AuthPipeline: @@ -111,7 +112,7 @@ class AuthPipeline: @dataclass(frozen=True) class PipelineRoute: pipeline: AuthPipeline - required_edition: frozenset[Edition] | None = None + required_edition: frozenset[DeploymentEdition] | None = None class PipelineRouter: @@ -130,7 +131,7 @@ class PipelineRouter: *, scope: Scope | None = None, allowed_token_types: frozenset[TokenType] | None = None, - edition: frozenset[Edition] | None = None, + edition: frozenset[DeploymentEdition] | None = None, workspace_membership: bool = False, allowed_roles: frozenset[TenantAccountRole] | None = None, rbac: RBACRequirement | None = None, @@ -149,7 +150,7 @@ class PipelineRouter: *, scope: Scope | None = None, allowed_token_types: frozenset[TokenType] | None = None, - edition: frozenset[Edition] | None = None, + edition: frozenset[DeploymentEdition] | None = None, allowed_roles: frozenset[TenantAccountRole] | None = None, rbac: RBACRequirement | None = None, ) -> Callable: @@ -167,7 +168,7 @@ class PipelineRouter: *, scope: Scope | None, allowed_token_types: frozenset[TokenType] | None, - edition: frozenset[Edition] | None, + edition: frozenset[DeploymentEdition] | None, workspace_membership: bool, allowed_roles: frozenset[TenantAccountRole] | None, rbac: RBACRequirement | None, @@ -199,17 +200,17 @@ class PipelineRouter: *, scope: Scope | None, allowed_token_types: frozenset[TokenType] | None, - edition: frozenset[Edition] | None, + edition: frozenset[DeploymentEdition] | None, workspace_membership: bool = False, allowed_roles: frozenset[TenantAccountRole] | None = None, rbac: RBACRequirement | None = None, ) -> Any: # 404 not 403 — this edition doesn't expose the feature at all - if edition is not None and current_edition() not in edition: + if edition is not None and dify_config.DEPLOYMENT_EDITION not in edition: raise NotFound() license_checked = False - if edition is not None and Edition.EE in edition: + if edition is not None and DeploymentEdition.ENTERPRISE in edition: _check_license() license_checked = True @@ -233,9 +234,9 @@ class PipelineRouter: raise Forbidden("unsupported_token_type") if route.required_edition is not None: - if current_edition() not in route.required_edition: + if dify_config.DEPLOYMENT_EDITION not in route.required_edition: raise Forbidden("external_sso_requires_ee") - if not license_checked and Edition.EE in route.required_edition: + if not license_checked and DeploymentEdition.ENTERPRISE in route.required_edition: _check_license() return route.pipeline._run( diff --git a/api/controllers/openapi/index.py b/api/controllers/openapi/index.py index 97e9c6e75d9..6f63232ad58 100644 --- a/api/controllers/openapi/index.py +++ b/api/controllers/openapi/index.py @@ -1,8 +1,11 @@ +"""Unauthenticated health and version probes for the OpenAPI surface.""" + from flask_restx import Resource +from configs import dify_config from controllers.openapi import openapi_ns from controllers.openapi._contract import returns -from controllers.openapi._models import HealthResponse +from controllers.openapi._models import HealthResponse, ServerVersionResponse @openapi_ns.route("/_health") @@ -10,3 +13,13 @@ class HealthApi(Resource): @returns(200, HealthResponse, description="Health check") def get(self): return HealthResponse(ok=True) + + +@openapi_ns.route("/_version") +class VersionApi(Resource): + @returns(200, ServerVersionResponse, description="Server version") + def get(self): + return ServerVersionResponse( + version=dify_config.project.version, + edition=dify_config.DEPLOYMENT_EDITION, + ) diff --git a/api/controllers/service_api/app/completion.py b/api/controllers/service_api/app/completion.py index a75e4b391ec..c16118f7799 100644 --- a/api/controllers/service_api/app/completion.py +++ b/api/controllers/service_api/app/completion.py @@ -42,7 +42,7 @@ from core.errors.error import ( QuotaExceededError, ) from core.helper.trace_id_helper import get_external_trace_id, get_trace_session_id, omit_trace_session_id_from_payload -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from graphon.model_runtime.errors.invoke import InvokeError from libs import helper from libs.helper import UUIDStrOrEmpty @@ -376,7 +376,11 @@ class ChatApi(Resource): payload = ChatRequestPayload.model_validate(omit_trace_session_id_from_payload(service_api_ns.payload) or {}) - if app_mode == AppMode.ADVANCED_CHAT and payload.workflow_id and dify_config.BILLING_ENABLED: + if ( + app_mode == AppMode.ADVANCED_CHAT + and payload.workflow_id + and dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD + ): billing_info = BillingService.get_info(app_model.tenant_id, exclude_vector_space=True) if billing_info["enabled"] and billing_info["subscription"]["plan"] == CloudPlan.SANDBOX: raise WorkflowVersionExecutionNotAllowedError() diff --git a/api/controllers/service_api/app/workflow.py b/api/controllers/service_api/app/workflow.py index 1041c038d72..846a500adc3 100644 --- a/api/controllers/service_api/app/workflow.py +++ b/api/controllers/service_api/app/workflow.py @@ -45,7 +45,7 @@ from core.errors.error import ( QuotaExceededError, ) from core.helper.trace_id_helper import get_external_trace_id, get_trace_session_id, omit_trace_session_id_from_payload -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from extensions.ext_database import db from extensions.ext_redis import redis_client from fields.base import ResponseModel @@ -451,7 +451,7 @@ class WorkflowRunByIdApi(Resource): if app_mode != AppMode.WORKFLOW: raise NotWorkflowAppError() - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: billing_info = BillingService.get_info(app_model.tenant_id, exclude_vector_space=True) if billing_info["enabled"] and billing_info["subscription"]["plan"] == CloudPlan.SANDBOX: raise WorkflowVersionExecutionNotAllowedError() diff --git a/api/controllers/service_api/dataset/document.py b/api/controllers/service_api/dataset/document.py index 44a7169fc91..2047246f52d 100644 --- a/api/controllers/service_api/dataset/document.py +++ b/api/controllers/service_api/dataset/document.py @@ -85,6 +85,7 @@ from services.entities.knowledge_entities.knowledge_entities import ( ProcessRule, RetrievalModel, ) +from services.feature_service import FeatureService from services.file_service import FileService from services.summary_index_service import SummaryIndexService @@ -699,9 +700,10 @@ class DocumentAddByFileApi(DatasetApiResource): "- `provider_not_initialize` : No valid model provider credentials found. Please go to " "Settings -> Model Provider to complete your provider credentials.\n" "- `invalid_param` : Knowledge base does not exist, external datasets not supported, " - "file too large, unsupported file type, missing required fields, or invalid doc_form " + "unsupported file type, missing required fields, or invalid doc_form " "(must be `text_model`, `hierarchical_model`, or `qa_model`)." ), + 413: "`file_too_large` : File size exceeded.", }, ) @service_api_ns.doc("create_document_by_file") @@ -712,6 +714,7 @@ class DocumentAddByFileApi(DatasetApiResource): 200: "Document created successfully", 401: "Unauthorized - invalid API token", 400: "Bad request - invalid file or parameters", + 413: "File too large", } ) @service_api_ns.response( @@ -778,13 +781,17 @@ class DocumentAddByFileApi(DatasetApiResource): if not current_user: raise ValueError("current_user is required") - upload_file = FileService(db.engine).upload_file( - filename=file.filename, - content=file.stream.read(), - mimetype=file.mimetype, - user=current_user, - source="datasets", - ) + try: + upload_file = FileService(db.engine).upload_file( + filename=file.filename, + content=file.stream.read(), + mimetype=file.mimetype, + user=current_user, + source="datasets", + default_file_size_limit=FeatureService.get_knowledge_file_size_limit(tenant_id), + ) + except services.errors.file.FileTooLargeError as file_too_large_error: + raise FileTooLargeError(file_too_large_error.description) data_source = { "type": "upload_file", "info_list": {"data_source_type": "upload_file", "file_info_list": {"file_ids": [upload_file.id]}}, @@ -859,6 +866,7 @@ def _update_document_by_file( mimetype=file.mimetype, user=current_user, source="datasets", + default_file_size_limit=FeatureService.get_knowledge_file_size_limit(tenant_id), ) except services.errors.file.FileTooLargeError as file_too_large_error: raise FileTooLargeError(file_too_large_error.description) @@ -916,9 +924,10 @@ class DeprecatedDocumentUpdateByFileApi(DatasetApiResource): "- `provider_not_initialize` : No valid model provider credentials found. Please go to " "Settings -> Model Provider to complete your provider credentials.\n" "- `invalid_param` : Knowledge base does not exist, external datasets not supported, " - "file too large, unsupported file type, or invalid doc_form (must be `text_model`, " - "`hierarchical_model`, or `qa_model`)." + "unsupported file type, or invalid doc_form (must be `text_model`, `hierarchical_model`, " + "or `qa_model`)." ), + 413: "`file_too_large` : File size exceeded.", }, ) @service_api_ns.doc("update_document_by_file_deprecated") @@ -935,6 +944,7 @@ class DeprecatedDocumentUpdateByFileApi(DatasetApiResource): 200: "Document updated successfully", 401: "Unauthorized - invalid API token", 404: "Document not found", + 413: "File too large", } ) @service_api_ns.response( @@ -1400,9 +1410,10 @@ class DocumentApi(DatasetApiResource): "- `provider_not_initialize` : No valid model provider credentials found. Please go to " "Settings -> Model Provider to complete your provider credentials.\n" "- `invalid_param` : Knowledge base does not exist, external datasets not supported, " - "file too large, unsupported file type, or invalid doc_form (must be `text_model`, " - "`hierarchical_model`, or `qa_model`)." + "unsupported file type, or invalid doc_form (must be `text_model`, `hierarchical_model`, " + "or `qa_model`)." ), + 413: "`file_too_large` : File size exceeded.", }, ) @service_api_ns.doc("update_document_by_file") @@ -1413,6 +1424,7 @@ class DocumentApi(DatasetApiResource): 200: "Document updated successfully", 401: "Unauthorized - invalid API token", 404: "Document not found", + 413: "File too large", } ) @service_api_ns.response( diff --git a/api/controllers/service_api/dataset/metadata.py b/api/controllers/service_api/dataset/metadata.py index 912071806a5..262ae176b7e 100644 --- a/api/controllers/service_api/dataset/metadata.py +++ b/api/controllers/service_api/dataset/metadata.py @@ -1,4 +1,4 @@ -from typing import Literal +from typing import Literal, cast from uuid import UUID from flask_login import current_user @@ -17,6 +17,7 @@ from fields.dataset_fields import ( DatasetMetadataResponse, ) from libs.helper import dump_response +from models import Account from services.dataset_service import DatasetService from services.entities.knowledge_entities.knowledge_entities import ( DocumentMetadataOperation, @@ -24,6 +25,7 @@ from services.entities.knowledge_entities.knowledge_entities import ( MetadataDetail, MetadataOperationData, ) +from services.errors.metadata import MetadataResourceNotFoundError from services.metadata_service import MetadataService BUILT_IN_METADATA_ACTION_PARAM = { @@ -158,12 +160,14 @@ class DatasetMetadataServiceApi(DatasetApiResource): dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, tenant_id, session=session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, session) - metadata = MetadataService.update_metadata_name(dataset_id_str, metadata_id_str, payload.name, session=session) + metadata = MetadataService.update_metadata_name( + dataset, metadata_id_str, payload.name, cast(Account, current_user), session=session + ) return dump_response(DatasetMetadataResponse, metadata), 200 @service_api_ns.doc( @@ -194,12 +198,12 @@ class DatasetMetadataServiceApi(DatasetApiResource): """Delete metadata.""" dataset_id_str = str(dataset_id) metadata_id_str = str(metadata_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + dataset = DatasetService.get_dataset_for_tenant(dataset_id_str, tenant_id, session=session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, session) - MetadataService.delete_metadata(dataset_id_str, metadata_id_str, session) + MetadataService.delete_metadata(dataset, metadata_id_str, session) return "", 204 @@ -297,7 +301,7 @@ class DocumentMetadataEditServiceApi(DatasetApiResource): responses={ 200: "Documents metadata updated successfully", 401: "Unauthorized - invalid API token", - 404: "Dataset not found", + 404: "Dataset, document, or metadata not found", } ) @service_api_ns.response( @@ -309,14 +313,18 @@ class DocumentMetadataEditServiceApi(DatasetApiResource): @with_session def post(self, session: Session, tenant_id, dataset_id: UUID): """Update metadata for multiple documents.""" - dataset_id_str = str(dataset_id) - dataset = DatasetService.get_dataset(dataset_id_str, session) + dataset = DatasetService.get_dataset_for_tenant(str(dataset_id), str(tenant_id), session=session) if dataset is None: raise NotFound("Dataset not found.") DatasetService.check_dataset_permission(dataset, current_user, session) metadata_args = MetadataOperationData.model_validate(service_api_ns.payload or {}) - MetadataService.update_documents_metadata(dataset, metadata_args, session=session) + try: + MetadataService.update_documents_metadata( + dataset, metadata_args, cast(Account, current_user), session=session + ) + except MetadataResourceNotFoundError as exc: + raise NotFound(str(exc)) from exc return dump_response(DatasetMetadataActionResponse, {"result": "success"}), 200 diff --git a/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py b/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py index 35f3a4c01a0..a2a1ae40bf1 100644 --- a/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py +++ b/api/controllers/service_api/dataset/rag_pipeline/rag_pipeline_workflow.py @@ -10,7 +10,12 @@ from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services -from controllers.common.errors import FilenameNotExistsError, NoFileUploadedError, TooManyFilesError +from controllers.common.errors import ( + FilenameNotExistsError, + FileTooLargeError, + NoFileUploadedError, + TooManyFilesError, +) from controllers.common.fields import GeneratedAppResponse from controllers.common.schema import ( query_params_from_model, @@ -32,7 +37,8 @@ from libs.login import current_user from models import Account from models.dataset import Dataset, Pipeline from models.engine import db -from services.errors.file import FileTooLargeError, UnsupportedFileTypeError +from services.errors.file import UnsupportedFileTypeError +from services.feature_service import FeatureService from services.file_service import FileService from services.rag_pipeline.entity.pipeline_service_api_entities import ( DatasourceNodeRunApiEntity, @@ -363,6 +369,7 @@ class KnowledgebasePipelineFileUploadApi(DatasetApiResource): content=file.stream.read(), mimetype=file.mimetype, user=current_user, + default_file_size_limit=FeatureService.get_knowledge_file_size_limit(tenant_id), ) except services.errors.file.FileTooLargeError as file_too_large_error: raise FileTooLargeError(file_too_large_error.description) diff --git a/api/controllers/service_api/index.py b/api/controllers/service_api/index.py index 41f8ef53a5b..365a7fa252f 100644 --- a/api/controllers/service_api/index.py +++ b/api/controllers/service_api/index.py @@ -1,9 +1,16 @@ from flask_restx import Resource from configs import dify_config -from controllers.common.fields import IndexInfoResponse from controllers.common.schema import register_response_schema_models from controllers.service_api import service_api_ns +from fields.base import ResponseModel + + +class IndexInfoResponse(ResponseModel): + welcome: str + api_version: str + server_version: str + register_response_schema_models(service_api_ns, IndexInfoResponse) @@ -12,8 +19,8 @@ register_response_schema_models(service_api_ns, IndexInfoResponse) class IndexApi(Resource): @service_api_ns.response(200, "Success", service_api_ns.models[IndexInfoResponse.__name__]) def get(self) -> dict[str, str]: - return { - "welcome": "Dify OpenAPI", - "api_version": "v1", - "server_version": dify_config.project.version, - } + return IndexInfoResponse( + welcome="Dify OpenAPI", + api_version="v1", + server_version=dify_config.project.version, + ).model_dump(mode="json") diff --git a/api/controllers/service_api/wraps.py b/api/controllers/service_api/wraps.py index c3c8e02e438..6d74c93ed59 100644 --- a/api/controllers/service_api/wraps.py +++ b/api/controllers/service_api/wraps.py @@ -13,7 +13,7 @@ from flask_restx.utils import merge from pydantic import BaseModel from sqlalchemy import select from sqlalchemy.orm import sessionmaker -from werkzeug.exceptions import Forbidden, NotFound, Unauthorized +from werkzeug.exceptions import Forbidden, NotFound, ServiceUnavailable, Unauthorized from configs import dify_config from controllers.service_api.schema import ( @@ -22,7 +22,7 @@ from controllers.service_api.schema import ( USER_QUERY_PARAM, USER_REQUIRED_ATTR, ) -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from extensions.ext_database import db from extensions.ext_redis import redis_client from libs.login import current_user @@ -186,10 +186,16 @@ def cloud_edition_billing_resource_check[**P, R]( def decorated(*args: P.args, **kwargs: P.kwargs): api_token = validate_and_get_api_token(api_token_type) if resource == "vector_space": - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return view(*args, **kwargs) vector_space = FeatureService.get_vector_space(api_token.tenant_id) + if vector_space.usage_unknown: + features = FeatureService.get_features(api_token.tenant_id, exclude_vector_space=True) + if features.billing.enabled and features.billing.subscription.plan == CloudPlan.SANDBOX: + raise ServiceUnavailable( + "Unable to verify vector space usage right now. Please try again later." + ) if 0 < vector_space.limit <= vector_space.size: raise Forbidden("The capacity of the vector space has reached the limit of your subscription.") return view(*args, **kwargs) @@ -322,11 +328,12 @@ def validate_dataset_token[R](view: Callable[..., R]) -> Callable[..., R]: raise Forbidden("Dataset api access is not enabled.") tenant_account_join = db.session.execute( - select(Tenant, TenantAccountJoin) - .where(Tenant.id == api_token.tenant_id) - .where(TenantAccountJoin.tenant_id == Tenant.id) - .where(TenantAccountJoin.role.in_(["owner"])) - .where(Tenant.status == TenantStatus.NORMAL) + select(Tenant, TenantAccountJoin).where( + Tenant.id == api_token.tenant_id, + TenantAccountJoin.tenant_id == Tenant.id, + TenantAccountJoin.role.in_(["owner"]), + Tenant.status == TenantStatus.NORMAL, + ) ).one_or_none() # TODO: only owner information is required, so only one is returned. if tenant_account_join: tenant, ta = tenant_account_join diff --git a/api/controllers/trigger/webhook.py b/api/controllers/trigger/webhook.py index 04b2b50bc4a..7715090b967 100644 --- a/api/controllers/trigger/webhook.py +++ b/api/controllers/trigger/webhook.py @@ -7,7 +7,7 @@ from werkzeug.exceptions import NotFound, RequestEntityTooLarge from controllers.trigger import bp from core.trigger.debug.event_bus import TriggerDebugEventBus from core.trigger.debug.events import WebhookDebugEvent, build_webhook_pool_key -from enums.quota_type import QuotaType +from enums import QuotaType from services.errors.app import QuotaExceededError from services.trigger.webhook_service import RawWebhookDataDict, WebhookService diff --git a/api/controllers/web/feature.py b/api/controllers/web/feature.py index 919788687dc..fcaaac98e28 100644 --- a/api/controllers/web/feature.py +++ b/api/controllers/web/feature.py @@ -2,8 +2,9 @@ from flask_restx import Resource from controllers.common.schema import register_response_schema_models from controllers.web import web_ns +from extensions.ext_application_services import application_services from libs.helper import dump_response -from services.feature_service import FeatureService, SystemFeatureModel +from services.entities.feature_entities import SystemFeatureModel register_response_schema_models(web_ns, SystemFeatureModel) @@ -29,4 +30,7 @@ class SystemFeatureApi(Resource): Authentication configuration must be available before the authentication flow can be selected. """ - return dump_response(SystemFeatureModel, FeatureService.get_system_features()) + return dump_response( + SystemFeatureModel, + application_services().feature_queries.get_system_features(), + ) diff --git a/api/controllers/web/human_input_form.py b/api/controllers/web/human_input_form.py index 775e802a70c..a897508e798 100644 --- a/api/controllers/web/human_input_form.py +++ b/api/controllers/web/human_input_form.py @@ -16,6 +16,7 @@ from configs import dify_config from controllers.common.errors import NotFoundError from controllers.common.human_input import HumanInputFormSubmitPayload, stringify_form_default_values from controllers.common.schema import register_response_schema_models, register_schema_models +from controllers.console.wraps import model_validate from controllers.web import web_ns from controllers.web.error import WebFormRateLimitExceededError from controllers.web.site import WebAppSiteResponse @@ -24,7 +25,7 @@ from extensions.ext_database import db from fields.base import ResponseModel from libs.helper import RateLimiter, dump_response, extract_remote_ip, to_timestamp from models.account import TenantStatus -from models.model import App, Site +from models.model import App, AppMode, Site from repositories.factory import DifyAPIRepositoryFactory from services.feature_service import FeatureService from services.human_input_file_upload_service import HumanInputFileUploadService @@ -207,6 +208,7 @@ class HumanInputFormApi(Resource): site=WebAppSiteResponse.from_app_site( tenant=tenant, app_model=app_model, + mode=AppMode.value_of(app_model.mode), site=site, end_user_id=None, features=features, @@ -234,7 +236,8 @@ class HumanInputFormApi(Resource): "Form submitted successfully", web_ns.models[HumanInputFormSubmitResponse.__name__], ) - def post(self, form_token: str): + @model_validate(HumanInputFormSubmitPayload) + def post(self, payload: HumanInputFormSubmitPayload, form_token: str): """ Submit human input form by token. @@ -248,8 +251,6 @@ class HumanInputFormApi(Resource): "action": "Approve" } """ - payload = HumanInputFormSubmitPayload.model_validate(request.get_json()) - ip_address = extract_remote_ip(request) if _FORM_SUBMIT_RATE_LIMITER.is_rate_limited(ip_address): raise WebFormRateLimitExceededError() diff --git a/api/controllers/web/login.py b/api/controllers/web/login.py index 0aa42f43687..b841056743d 100644 --- a/api/controllers/web/login.py +++ b/api/controllers/web/login.py @@ -30,6 +30,7 @@ from controllers.console.wraps import ( ) from controllers.web import web_ns from controllers.web.wraps import decode_jwt_token +from enums import DeploymentEdition from extensions.ext_database import db from libs.helper import EmailStr, extract_remote_ip from libs.passport import PassportService @@ -146,8 +147,9 @@ class LoginStatusApi(Resource): if not app_code: return LoginStatusResponse(logged_in=bool(token), app_logged_in=False).model_dump(mode="json") app_id = AppService.get_app_id_by_code(app_code, session=db.session()) - is_public = not dify_config.ENTERPRISE_ENABLED or not WebAppAuthService.is_app_require_permission_check( - app_id=app_id, session=db.session() + is_public = ( + dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE + or not WebAppAuthService.is_app_require_permission_check(app_id=app_id, session=db.session()) ) user_logged_in = False diff --git a/api/controllers/web/site.py b/api/controllers/web/site.py index f6c4af013a1..827fdce72b0 100644 --- a/api/controllers/web/site.py +++ b/api/controllers/web/site.py @@ -8,14 +8,15 @@ from configs import dify_config from controllers.common.schema import register_response_schema_models from controllers.web import web_ns from controllers.web.wraps import WebApiResource -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from extensions.ext_database import db from extensions.storage.storage_type import StorageType from fields.base import ResponseModel from libs.helper import build_icon_url from models.account import Tenant, TenantStatus -from models.model import App, EndUser, IconType, Site -from services.feature_service import FeatureModel, FeatureService +from models.model import App, AppMode, EndUser, IconType, Site +from services.entities.feature_entities import FeatureModel +from services.feature_service import FeatureService from services.file_service import FileService @@ -67,6 +68,7 @@ class WebAppCustomConfigResponse(ResponseModel): class WebAppSiteResponse(ResponseModel): app_id: str + mode: AppMode end_user_id: str | None = None enable_site: bool site: WebSiteResponse @@ -83,6 +85,7 @@ class WebAppSiteResponse(ResponseModel): *, tenant: Tenant, app_model: App, + mode: AppMode, site: Site, end_user_id: str | None, features: FeatureModel, @@ -109,6 +112,7 @@ class WebAppSiteResponse(ResponseModel): return cls( app_id=app_model.id, + mode=mode, end_user_id=end_user_id, enable_site=app_model.enable_site, site=site_response, @@ -167,6 +171,7 @@ class AppSiteApi(WebApiResource): return WebAppSiteResponse.from_app_site( tenant=tenant, app_model=app_model, + mode=AppMode.value_of(app_model.mode_compatible_with_agent_with_session(session=db.session())), site=site, end_user_id=end_user.id, features=features, diff --git a/api/controllers/web/workflow_events.py b/api/controllers/web/workflow_events.py index 48eba33f04c..2954b9a81de 100644 --- a/api/controllers/web/workflow_events.py +++ b/api/controllers/web/workflow_events.py @@ -86,6 +86,7 @@ class WorkflowEventsApi(WebApiResource): raise InvalidArgumentError(f"cannot subscribe to workflow run, workflow_run_id={workflow_run.id}") include_state_snapshot = request.args.get("include_state_snapshot", "false").lower() == "true" + continue_on_pause = request.args.get("continue_on_pause", "false").lower() == "true" def _generate_stream_events(): if include_state_snapshot: @@ -96,6 +97,7 @@ class WorkflowEventsApi(WebApiResource): tenant_id=app_model.tenant_id, app_id=app_model.id, session_maker=session_maker, + close_on_pause=not continue_on_pause, ) ) return generator.convert_to_event_stream( diff --git a/api/core/agent/base_agent_runner.py b/api/core/agent/base_agent_runner.py index 2d20ca75a5e..806f7c6590f 100644 --- a/api/core/agent/base_agent_runner.py +++ b/api/core/agent/base_agent_runner.py @@ -122,7 +122,8 @@ class BaseAgentRunner(AppRunner): model_schema = llm_model.get_model_schema(model_instance.model_name, model_instance.credentials) features = model_schema.features if model_schema and model_schema.features else [] self.stream_tool_call = ModelFeature.STREAM_TOOL_CALL in features - self.files = application_generate_entity.files if ModelFeature.VISION in features else [] + self.vision_enabled = ModelFeature.VISION in features + self.files = application_generate_entity.files if self.vision_enabled else [] self.query: str = "" self._current_thoughts: list[PromptMessage] = [] diff --git a/api/core/agent/fc_agent_runner.py b/api/core/agent/fc_agent_runner.py index 5bffa0002bf..78980f0d943 100644 --- a/api/core/agent/fc_agent_runner.py +++ b/api/core/agent/fc_agent_runner.py @@ -1,19 +1,25 @@ import json import logging +import re from collections.abc import Generator from copy import deepcopy from typing import Any, Union +from sqlalchemy import select from sqlalchemy.orm import Session from core.agent.base_agent_runner import BaseAgentRunner from core.agent.errors import AgentMaxIterationError from core.app.apps.base_app_queue_manager import PublishFrom from core.app.entities.queue_entities import QueueAgentThoughtEvent, QueueMessageEndEvent, QueueMessageFileEvent +from core.app.file_access import grant_upload_file_access from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform from core.tools.entities.tool_entities import ToolInvokeMeta +from core.tools.signature import sign_upload_file_preview_url from core.tools.tool_engine import ToolEngine -from graphon.file import file_manager +from core.tools.utils.dataset_retriever_tool import DatasetRetrieverTool +from core.workflow.file_reference import build_file_reference +from graphon.file import File, FileTransferMethod, FileType, file_manager from graphon.model_runtime.entities import ( AssistantPromptMessage, LLMResult, @@ -28,12 +34,70 @@ from graphon.model_runtime.entities import ( UserPromptMessage, ) from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes +from models import UploadFile from models.model import Message logger = logging.getLogger(__name__) +_FILE_PREVIEW_ID_PATTERN = re.compile(r"/files/([a-fA-F0-9-]{36})/file-preview") +_KNOWLEDGE_RETRIEVAL_PROMPT_NAME = "knowledge_retrieval" + class FunctionCallAgentRunner(BaseAgentRunner): + def _build_dataset_tool_image_contents( + self, session: Session, tool_response: str, tool_instance: Any + ) -> list[PromptMessageContentUnionTypes]: + if not self.vision_enabled or not isinstance(tool_instance, DatasetRetrieverTool): + return [] + + upload_file_ids = list(dict.fromkeys(_FILE_PREVIEW_ID_PATTERN.findall(tool_response))) + if not upload_file_ids: + return [] + + upload_files = session.scalars(select(UploadFile).where(UploadFile.id.in_(upload_file_ids))).all() + upload_file_map = {str(upload_file.id): upload_file for upload_file in upload_files} + ordered_upload_files = [ + upload_file_map[upload_file_id] for upload_file_id in upload_file_ids if upload_file_id in upload_file_map + ] + image_upload_files = [ + upload_file for upload_file in ordered_upload_files if (upload_file.mime_type or "").startswith("image/") + ] + if not image_upload_files: + return [] + + grant_upload_file_access(str(upload_file.id) for upload_file in image_upload_files) + + image_detail_config = ( + self.application_generate_entity.file_upload_config.image_config.detail + if ( + self.application_generate_entity.file_upload_config + and self.application_generate_entity.file_upload_config.image_config + ) + else None + ) + image_detail_config = image_detail_config or ImagePromptMessageContent.DETAIL.LOW + + prompt_message_contents: list[PromptMessageContentUnionTypes] = [] + for upload_file in image_upload_files: + prompt_file = File( + file_id=upload_file.id, + filename=upload_file.name, + extension="." + upload_file.extension, + mime_type=upload_file.mime_type, + file_type=FileType.IMAGE, + transfer_method=FileTransferMethod.LOCAL_FILE, + remote_url=upload_file.source_url, + reference=build_file_reference(record_id=str(upload_file.id)), + size=upload_file.size, + storage_key=upload_file.key, + url=sign_upload_file_preview_url(upload_file.id, upload_file.extension), + ) + prompt_message_contents.append( + file_manager.to_prompt_message_content(prompt_file, image_detail_config=image_detail_config) + ) + + return prompt_message_contents + def run( self, session: Session, message: Message, query: str, **kwargs: Any ) -> Generator[LLMResultChunk, None, None]: @@ -285,13 +349,29 @@ class FunctionCallAgentRunner(BaseAgentRunner): tool_responses.append(tool_response) if tool_response["tool_response"] is not None: + tool_response_text = str(tool_response["tool_response"]) + dataset_image_contents = self._build_dataset_tool_image_contents( + session=session, + tool_response=tool_response_text, + tool_instance=tool_instance, + ) self._current_thoughts.append( ToolPromptMessage( - content=str(tool_response["tool_response"]), + content=tool_response_text, tool_call_id=tool_call_id, name=tool_call_name, ) ) + if dataset_image_contents: + self._current_thoughts.append( + UserPromptMessage( + name=_KNOWLEDGE_RETRIEVAL_PROMPT_NAME, + content=[ + *dataset_image_contents, + TextPromptMessageContent(data=self.query or tool_response_text), + ], + ) + ) if len(tool_responses) > 0: # save agent thought @@ -453,6 +533,8 @@ class FunctionCallAgentRunner(BaseAgentRunner): for prompt_message in prompt_messages: if isinstance(prompt_message, UserPromptMessage): + if prompt_message.name == _KNOWLEDGE_RETRIEVAL_PROMPT_NAME: + continue if isinstance(prompt_message.content, list): prompt_message.content = "\n".join( [ diff --git a/api/core/app/apps/advanced_chat/app_generator.py b/api/core/app/apps/advanced_chat/app_generator.py index 3c677804028..aa963683734 100644 --- a/api/core/app/apps/advanced_chat/app_generator.py +++ b/api/core/app/apps/advanced_chat/app_generator.py @@ -616,23 +616,34 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): message_snapshot = MessageSnapshot.from_message(message) session.close() - # return response or stream generator - response = self._handle_advanced_chat_response( - application_generate_entity=application_generate_entity, - workflow=workflow_snapshot, - queue_manager=queue_manager, - conversation=conversation_snapshot, - message=message_snapshot, - user=user, - stream=stream, - draft_var_saver_factory=self._get_draft_var_saver_factory( - invoke_from, - account=user, - tenant_id=application_generate_entity.app_config.tenant_id, - ), - ) + try: + response = self._handle_advanced_chat_response( + application_generate_entity=application_generate_entity, + workflow=workflow_snapshot, + queue_manager=queue_manager, + conversation=conversation_snapshot, + message=message_snapshot, + user=user, + stream=stream, + draft_var_saver_factory=self._get_draft_var_saver_factory( + invoke_from, + account=user, + tenant_id=application_generate_entity.app_config.tenant_id, + ), + ) + converted_response = AdvancedChatAppGenerateResponseConverter.convert( + response=response, + invoke_from=invoke_from, + ) + except BaseException: + self._join_worker_thread(worker_thread) + raise - return AdvancedChatAppGenerateResponseConverter.convert(response=response, invoke_from=invoke_from) + if isinstance(converted_response, Generator): + return self._wrap_stream_with_worker_thread_join(converted_response, worker_thread) + + self._join_worker_thread(worker_thread) + return converted_response def _generate_worker( self, @@ -674,6 +685,12 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator): ) if workflow is None: raise ValueError("Workflow not found") + if graph_runtime_state is not None: + self._restore_workflow_run_graph( + session=session, + workflow=workflow, + workflow_run_id=application_generate_entity.workflow_run_id, + ) # Determine system_user_id based on invocation source is_external_api_call = application_generate_entity.invoke_from in { diff --git a/api/core/app/apps/advanced_chat/app_runner.py b/api/core/app/apps/advanced_chat/app_runner.py index 31a65578d49..cf3943a6ccc 100644 --- a/api/core/app/apps/advanced_chat/app_runner.py +++ b/api/core/app/apps/advanced_chat/app_runner.py @@ -16,6 +16,7 @@ from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner from core.app.entities.app_invoke_entities import ( AdvancedChatAppGenerateEntity, AppGenerateEntity, + DifyRunContext, InvokeFrom, ) from core.app.entities.queue_entities import ( @@ -31,7 +32,7 @@ from core.moderation.base import ModerationError from core.moderation.input_moderation import InputModeration from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository from core.workflow.node_factory import get_default_root_node_id -from core.workflow.nodes.agent_v2.session_cleanup_layer import build_workflow_agent_session_cleanup_layer +from core.workflow.nodes.agent_v2.workspace_retirement_layer import build_workflow_agent_workspace_retirement_layer from core.workflow.system_variables import ( build_bootstrap_variables, build_system_variables, @@ -267,7 +268,18 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner): ) workflow_entry.graph_engine.layer(persistence_layer) - workflow_entry.graph_engine.layer(build_workflow_agent_session_cleanup_layer()) + workflow_entry.graph_engine.layer( + build_workflow_agent_workspace_retirement_layer( + dify_run_context=DifyRunContext( + tenant_id=self._workflow.tenant_id, + app_id=self._workflow.app_id, + user_id=self.application_generate_entity.user_id, + user_from=user_from, + invoke_from=invoke_from, + trace_session_id=self.application_generate_entity.extras.get("trace_session_id"), + ) + ) + ) conversation_variable_layer = ConversationVariablePersistenceLayer( ConversationVariableUpdater(session_factory.get_session_maker()) ) diff --git a/api/core/app/apps/agent_app/app_generator.py b/api/core/app/apps/agent_app/app_generator.py index 26e5d0cdcb6..a291b84b39b 100644 --- a/api/core/app/apps/agent_app/app_generator.py +++ b/api/core/app/apps/agent_app/app_generator.py @@ -27,21 +27,20 @@ from clients.agent_backend import AgentBackendRunEventAdapter from clients.agent_backend.factory import create_agent_backend_run_client from configs import dify_config from constants import UUID_NIL +from core.agent.publish_visibility import agent_has_workflow_callable_active_snapshot from core.app.app_config.easy_ui_based_app.model_config.converter import ModelConfigConverter from core.app.apps.agent_app.app_config_manager import AgentAppConfigManager from core.app.apps.agent_app.app_runner import AgentAppRunner from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError from core.app.apps.agent_app.generate_response_converter import AgentAppGenerateResponseConverter from core.app.apps.agent_app.runtime_request_builder import AgentAppRuntimeRequestBuilder -from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore +from core.app.apps.agent_app.session_store import AgentAppWorkspaceStore from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.message_based_app_generator import MessageBasedAppGenerator from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager from core.app.entities.app_invoke_entities import ( - AGENT_RUNTIME_EXIT_INTENT_ARG, AgentAppGenerateEntity, - AgentRuntimeExitIntent, DifyRunContext, InvokeFrom, UserFrom, @@ -51,19 +50,23 @@ from core.db.session_factory import session_factory from core.ops.ops_trace_manager import TraceQueueManager from core.workflow.file_reference import build_file_reference, is_canonical_file_reference from extensions.ext_database import db -from models import Account, App, AppModelConfig, EndUser, Message, MessageAnnotation +from models import Account, App, AppModelConfig, Conversation, EndUser, Message, MessageAnnotation from models.agent import ( APP_BACKED_AGENT_SOURCES, Agent, AgentConfigDraft, AgentConfigDraftType, AgentConfigSnapshot, + AgentConfigVersionKind, AgentScope, - AgentSource, AgentStatus, + AgentWorkingResourceStatus, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, ) from models.agent_config_entities import AgentSoulConfig from models.model import load_annotation_reply_config +from services.agent.workspace_service import AgentWorkspaceService, WorkspaceOwnerScope from services.conversation_service import ConversationService logger = logging.getLogger(__name__) @@ -150,19 +153,6 @@ class AgentAppGenerator(MessageBasedAppGenerator): inputs = args["inputs"] prompt_file_mappings = args.get("files") or [] - # Resolve the bound roster Agent + its current Agent Soul snapshot. - agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent( - app_model, - invoke_from=invoke_from, - draft_type=args.get("draft_type"), - user=user, - session=session, - ) - runtime_session_snapshot_id = self._runtime_session_snapshot_id( - invoke_from=invoke_from, - snapshot_id=agent_config_id, - ) - conversation = None conversation_id = args.get("conversation_id") if conversation_id: @@ -170,6 +160,21 @@ class AgentAppGenerator(MessageBasedAppGenerator): app_model=app_model, conversation_id=conversation_id, user=user, session=session ) + # New conversations use the current Agent generation. Existing + # conversations use the immutable generation named by their Binding. + agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent( + app_model, + invoke_from=invoke_from, + draft_type=args.get("draft_type"), + user=user, + session=session, + conversation=conversation, + ) + session_scope_config_version_id = self._session_scope_config_version_id( + invoke_from=invoke_from, + config_version_id=agent_config_id, + ) + # Build the EasyUI-shaped config from the Agent Soul so the chat pipeline # can persist usage; the answer itself comes from the agent backend. app_model_config = ( @@ -186,8 +191,6 @@ class AgentAppGenerator(MessageBasedAppGenerator): model_conf = ModelConfigConverter.convert(app_config) trace_manager = TraceQueueManager(app_model.id, user.id if isinstance(user, Account) else user.session_id) - agent_runtime_exit_intent = self._resolve_agent_runtime_exit_intent(args) - application_generate_entity = AgentAppGenerateEntity( task_id=str(uuid.uuid4()), app_config=app_config, @@ -215,8 +218,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): agent_id=agent.id, agent_config_snapshot_id=agent_config_id, agent_config_version_kind=agent_config_version_kind, - agent_runtime_session_snapshot_id=runtime_session_snapshot_id, - agent_runtime_exit_intent=agent_runtime_exit_intent, + agent_session_scope_config_version_id=session_scope_config_version_id, ) conversation, message = self._init_generate_records( @@ -265,6 +267,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): app_model: App, user: Account | EndUser, conversation_id: str, + form_id: str, invoke_from: InvokeFrom, session: Session, ) -> None: @@ -279,14 +282,21 @@ class AgentAppGenerator(MessageBasedAppGenerator): conversation = ConversationService.get_conversation( app_model=app_model, conversation_id=conversation_id, user=user, session=session ) + draft_type, draft_id = self._resolve_resume_draft( + app_model=app_model, + conversation=conversation, + user=user, + form_id=form_id, + session=session, + ) agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent( app_model, invoke_from=invoke_from, - draft_type=self._resume_draft_type( - app_model=app_model, conversation=conversation, user=user, session=session - ), + draft_type=draft_type, + draft_id=draft_id, user=user, session=session, + conversation=conversation, ) app_model_config = ( @@ -384,30 +394,37 @@ class AgentAppGenerator(MessageBasedAppGenerator): ) @staticmethod - def _resume_draft_type( - *, app_model: App, conversation: Any, user: Account | EndUser, session: Session - ) -> str | None: + def _resolve_resume_draft( + *, + app_model: App, + conversation: Any, + user: Account | EndUser, + form_id: str, + session: Session, + ) -> tuple[str | None, str | None]: if conversation.invoke_from != InvokeFrom.DEBUGGER: - return None - active_session = AgentAppRuntimeSessionStore().load_active_session_for_conversation( - tenant_id=app_model.tenant_id, - app_id=app_model.id, - conversation_id=conversation.id, - ) - snapshot_id = active_session.scope.agent_config_snapshot_id if active_session is not None else None - if snapshot_id and isinstance(user, Account): - draft = session.scalar( - select(AgentConfigDraft).where( - AgentConfigDraft.tenant_id == app_model.tenant_id, - AgentConfigDraft.id == snapshot_id, - ) + return None, None + if not isinstance(user, Account): + return AgentConfigDraftType.DRAFT.value, None + + build_draft = session.scalar( + select(AgentConfigDraft) + .join( + AgentWorkspaceBinding, + AgentWorkspaceBinding.id == AgentConfigDraft.agent_workspace_binding_id, ) - if draft is not None: - if draft.draft_type == AgentConfigDraftType.DEBUG_BUILD and draft.account_id == user.id: - return AgentConfigDraftType.DEBUG_BUILD.value - if draft.draft_type == AgentConfigDraftType.DRAFT and draft.account_id is None: - return AgentConfigDraftType.DRAFT.value - return AgentConfigDraftType.DRAFT.value + .where( + AgentConfigDraft.tenant_id == app_model.tenant_id, + AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD, + AgentConfigDraft.account_id == user.id, + AgentWorkspaceBinding.tenant_id == app_model.tenant_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + AgentWorkspaceBinding.pending_form_id == form_id, + ) + ) + if build_draft is not None: + return AgentConfigDraftType.DEBUG_BUILD.value, build_draft.id + return AgentConfigDraftType.DRAFT.value, None def _generate_worker( self, @@ -482,7 +499,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): invoke_from=application_generate_entity.invoke_from, ) with session_factory.create_session() as session: - _, _, agent_soul = self._resolve_agent_by_id( + agent, config_version, agent_soul = self._resolve_agent_by_id( tenant_id=app_config.tenant_id, agent_id=application_generate_entity.agent_id, snapshot_id=application_generate_entity.agent_config_snapshot_id, @@ -496,13 +513,18 @@ class AgentAppGenerator(MessageBasedAppGenerator): agent_config_snapshot_id=application_generate_entity.agent_config_snapshot_id, agent_config_version_kind=application_generate_entity.agent_config_version_kind, agent_soul=agent_soul, + home_snapshot_id=config_version.home_snapshot_id, conversation_id=conversation.id, query=query, message_id=message.id, model_name=application_generate_entity.model_conf.model, queue_manager=queue_manager, - session_scope_snapshot_id=application_generate_entity.agent_runtime_session_snapshot_id, - agent_runtime_exit_intent=application_generate_entity.agent_runtime_exit_intent, + session_scope_snapshot_id=application_generate_entity.agent_session_scope_config_version_id, + build_draft_id=( + application_generate_entity.agent_config_snapshot_id + if application_generate_entity.agent_config_version_kind == AgentConfigVersionKind.BUILD_DRAFT + else None + ), ) except GenerateTaskStoppedError: pass @@ -519,18 +541,6 @@ class AgentAppGenerator(MessageBasedAppGenerator): raise AgentAppGeneratorError("query is required") return query.replace("\x00", "") - @staticmethod - def _resolve_agent_runtime_exit_intent(args: Mapping[str, Any]) -> AgentRuntimeExitIntent: - """Resolve API-internal runtime exit policy from controller-owned args. - - Only the private controller-injected "delete" value changes behavior. - Normal chat and resume flows default/fallback to "suspend" so public - payloads and invalid internal values preserve existing semantics. - """ - if args.get(AGENT_RUNTIME_EXIT_INTENT_ARG) == "delete": - return "delete" - return "suspend" - @staticmethod def _build_runner(dify_context: DifyRunContext) -> AgentAppRunner: credentials_provider, _ = build_dify_model_access(dify_context) @@ -538,6 +548,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): request_builder=AgentAppRuntimeRequestBuilder(credentials_provider=credentials_provider), agent_backend_client=create_agent_backend_run_client( base_url=dify_config.AGENT_BACKEND_BASE_URL, + api_token=dify_config.AGENT_BACKEND_API_TOKEN, use_fake=dify_config.AGENT_BACKEND_USE_FAKE, fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO, stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS, @@ -545,7 +556,7 @@ class AgentAppGenerator(MessageBasedAppGenerator): stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS, ), event_adapter=AgentBackendRunEventAdapter(), - session_store=AgentAppRuntimeSessionStore(), + session_store=AgentAppWorkspaceStore(), text_delta_debounce_seconds=dify_config.AGENT_APP_TEXT_DELTA_DEBOUNCE_SECONDS, ) @@ -609,8 +620,10 @@ class AgentAppGenerator(MessageBasedAppGenerator): *, invoke_from: InvokeFrom, draft_type: Any, + draft_id: str | None = None, user: Account | EndUser, session: Session, + conversation: Conversation | None = None, ) -> tuple[Agent, str, Literal["snapshot", "draft", "build_draft"], AgentSoulConfig]: agent = session.scalar( select(Agent) @@ -631,17 +644,12 @@ class AgentAppGenerator(MessageBasedAppGenerator): ) if agent is None: raise AgentAppGeneratorError("Agent App has no bound Agent") - if ( - agent.source == AgentSource.IMPORTED - and not agent.active_config_is_published - and invoke_from != InvokeFrom.DEBUGGER - ): - raise AgentAppNotPublishedError("Agent has not been published") if invoke_from == InvokeFrom.DEBUGGER: draft = self._resolve_debug_draft( tenant_id=app_model.tenant_id, agent=agent, draft_type=draft_type, + draft_id=draft_id, account_id=user.id if isinstance(user, Account) else None, session=session, ) @@ -650,32 +658,86 @@ class AgentAppGenerator(MessageBasedAppGenerator): "build_draft" if draft.draft_type == AgentConfigDraftType.DEBUG_BUILD else "draft" ) return agent, draft.id, config_version_kind, agent_soul - # active_config_is_published tracks whether the editable draft matches the active snapshot. - # Public runtime must keep serving the active snapshot even when unpublished draft edits exist. - if not agent.active_config_snapshot_id: + # Dirty drafts do not revoke a published snapshot, while the seeded + # create/import snapshot must never become public runtime configuration. + if not agent_has_workflow_callable_active_snapshot(session=session, agent=agent): raise AgentAppNotPublishedError("Agent has not been published") + conversation_binding = self._resolve_conversation_binding( + session=session, + tenant_id=app_model.tenant_id, + app_id=app_model.id, + agent_id=agent.id, + conversation=conversation, + ) + snapshot_id = ( + conversation_binding.agent_config_version_id + if conversation_binding is not None + else agent.active_config_snapshot_id + ) _, snapshot, agent_soul = self._resolve_agent_by_id( tenant_id=app_model.tenant_id, agent_id=agent.id, - snapshot_id=agent.active_config_snapshot_id, + snapshot_id=snapshot_id, session=session, ) + if conversation_binding is not None: + AgentWorkspaceService.validate_binding_generation( + conversation_binding, + base_home_snapshot_id=snapshot.home_snapshot_id, + agent_config_version_id=snapshot.id, + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) return agent, snapshot.id, "snapshot", agent_soul @staticmethod - def _runtime_session_snapshot_id(*, invoke_from: InvokeFrom, snapshot_id: str) -> str | None: - """Return the session scope snapshot id for Agent App runtime state. + def _resolve_conversation_binding( + *, + session: Session, + tenant_id: str, + app_id: str, + agent_id: str, + conversation: Conversation | None, + ) -> AgentWorkspaceBinding | None: + """Resolve the exact participant generation owned by an existing conversation.""" + + if conversation is None or conversation.agent_workspace_binding_id is None: + return None + binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=tenant_id, + binding_id=conversation.agent_workspace_binding_id, + expected_owner_scope=WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=app_id, + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id=conversation.id, + ), + ) + if binding is None or binding.agent_id != agent_id: + raise AgentAppGeneratorError("Conversation participant Binding is unavailable") + return binding + + @staticmethod + def _session_scope_config_version_id(*, invoke_from: InvokeFrom, config_version_id: str) -> str | None: + """Return the config version id that scopes Agent App session reuse. Console preview/debug chat uses a stable Agent draft row id; build mode uses the current user's build-draft row id. Published/web/API runs use - immutable published snapshot ids. This keeps runtime session continuity + immutable published snapshot ids. This keeps Workspace Binding continuity inside one editable surface without mixing draft/build/published state. """ - return snapshot_id + del invoke_from + return config_version_id @staticmethod def _resolve_debug_draft( - *, tenant_id: str, agent: Agent, draft_type: Any, account_id: str | None, session: Session + *, + tenant_id: str, + agent: Agent, + draft_type: Any, + account_id: str | None, + session: Session, + draft_id: str | None = None, ) -> AgentConfigDraft: effective_draft_type = ( AgentConfigDraftType.DEBUG_BUILD @@ -699,6 +761,8 @@ class AgentAppGenerator(MessageBasedAppGenerator): AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD, AgentConfigDraft.account_id == account_id, ) + if draft_id is not None: + stmt = stmt.where(AgentConfigDraft.id == draft_id) draft = session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1)) if draft is not None: return draft diff --git a/api/core/app/apps/agent_app/app_runner.py b/api/core/app/apps/agent_app/app_runner.py index f844068073d..da1008ddfe8 100644 --- a/api/core/app/apps/agent_app/app_runner.py +++ b/api/core/app/apps/agent_app/app_runner.py @@ -3,8 +3,7 @@ Unlike the legacy ``AgentChatAppRunner`` (which runs an in-process ReAct loop), this runner delegates to the Agent backend, consumes the streamed event flow, republishes the assistant answer through the existing EasyUI chat task -pipeline, and then either saves or retires the conversation-owned runtime -session depending on the turn's exit policy. +pipeline, and saves the latest Agenton snapshot on the persistent Binding. """ from __future__ import annotations @@ -31,22 +30,20 @@ from clients.agent_backend import ( AgentBackendRunFailedInternalEvent, AgentBackendRunSucceededInternalEvent, AgentBackendStreamInternalEvent, - extract_runtime_layer_specs, ) -from clients.agent_backend.session_cleanup import AgentBackendSessionCleanupPayload from core.app.apps.agent_app.runtime_request_builder import ( AgentAppRuntimeBuildContext, AgentAppRuntimeRequest, AgentAppRuntimeRequestBuilder, ) from core.app.apps.agent_app.session_store import ( - AgentAppRuntimeSessionStore, AgentAppSessionScope, + AgentAppWorkspaceStore, StoredAgentAppSession, ) from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.exc import GenerateTaskStoppedError -from core.app.entities.app_invoke_entities import AgentRuntimeExitIntent, DifyRunContext +from core.app.entities.app_invoke_entities import DifyRunContext from core.app.entities.queue_entities import ( QueueAgentMessageEvent, QueueAgentThoughtEvent, @@ -67,10 +64,10 @@ from graphon.model_runtime.errors.invoke import ( InvokeRateLimitError, InvokeServerUnavailableError, ) +from models.agent import AgentConfigVersionKind from models.agent_config_entities import AgentSoulConfig from models.enums import CreatorUserRole from models.model import MessageAgentThought -from tasks.agent_backend_session_cleanup_task import cleanup_conversation_agent_runtime_session logger = logging.getLogger(__name__) @@ -92,9 +89,10 @@ _AGENT_BACKEND_INVOKE_ERROR_BY_REASON: Mapping[str, type[InvokeError]] = { def _agent_backend_failure_to_exception(event: AgentBackendRunFailedInternalEvent) -> Exception: - err_cls = _AGENT_BACKEND_INVOKE_ERROR_BY_REASON.get(event.reason or "") - if err_cls is not None: - return err_cls(event.error) + if event.error_type is None: + err_cls = _AGENT_BACKEND_INVOKE_ERROR_BY_REASON.get(event.reason or "") + if err_cls is not None: + return err_cls(event.error) message = event.error or "Agent backend run did not complete successfully." return AgentBackendRunFailedError( event.run_id, @@ -104,6 +102,7 @@ def _agent_backend_failure_to_exception(event: AgentBackendRunFailedInternalEven "source_event_id": event.source_event_id, }, message=message, + error_type=event.error_type, reason=event.reason, source_event_id=event.source_event_id, ) @@ -620,7 +619,7 @@ class AgentAppRunner: request_builder: AgentAppRuntimeRequestBuilder, agent_backend_client: AgentBackendRunClient, event_adapter: AgentBackendRunEventAdapter, - session_store: AgentAppRuntimeSessionStore, + session_store: AgentAppWorkspaceStore, text_delta_debounce_seconds: float, ) -> None: self._request_builder = request_builder @@ -637,37 +636,41 @@ class AgentAppRunner: agent_config_snapshot_id: str, agent_config_version_kind: Literal["snapshot", "draft", "build_draft"] = "snapshot", agent_soul: AgentSoulConfig, + home_snapshot_id: str | None, conversation_id: str, query: str, message_id: str, model_name: str, queue_manager: AppQueueManager, session_scope_snapshot_id: str | None | _DefaultSessionScopeSnapshotId = _DEFAULT_SESSION_SCOPE_SNAPSHOT_ID, - agent_runtime_exit_intent: AgentRuntimeExitIntent = "suspend", + build_draft_id: str | None = None, ) -> None: - preserve_session = agent_runtime_exit_intent == "suspend" scope = self._build_session_scope( dify_context=dify_context, agent_id=agent_id, agent_config_snapshot_id=agent_config_snapshot_id, + home_snapshot_id=home_snapshot_id, conversation_id=conversation_id, session_scope_snapshot_id=session_scope_snapshot_id, + agent_config_version_kind=AgentConfigVersionKind(agent_config_version_kind), + build_draft_id=build_draft_id, ) # ENG-638: if a prior turn paused on ask_human and the form is now answered, # resume by threading the human's reply into this run as deferred_tool_results. - stored = self._session_store.load_active_session(scope) + stored = self._session_store.load_or_create(scope) runtime = self._build_runtime( dify_context=dify_context, agent_id=agent_id, agent_config_snapshot_id=agent_config_snapshot_id, agent_config_version_kind=agent_config_version_kind, agent_soul=agent_soul, + binding_id=stored.binding_id, + backend_binding_ref=stored.backend_binding_ref, conversation_id=conversation_id, query=query, idempotency_key=message_id, stored=stored, message_id=message_id, - suspend_on_exit=preserve_session, ) create_response = self._agent_backend_client.create_run(runtime.request) @@ -681,9 +684,6 @@ class AgentAppRunner: ) if isinstance(terminal, AgentBackendDeferredToolCallInternalEvent): - if not preserve_session: - self._mark_session_cleaned(scope=scope, backend_run_id=terminal.run_id) - raise AgentBackendError("Agent App finalization cannot pause for human input.") # ENG-635: the agent asked a human. End this turn with the question and # a conversation-owned HITL form; a form submission resumes the run. self._pause_for_ask_human( @@ -703,8 +703,8 @@ class AgentAppRunner: if not isinstance(terminal, AgentBackendRunSucceededInternalEvent): if isinstance(terminal, AgentBackendRunFailedInternalEvent): reason = terminal.reason - if reason == "sandbox_expired": - raise AgentBackendError("The agent session sandbox has expired. Please start a new conversation.") + if terminal.error_type is None and reason == "binding_lost": + raise AgentBackendError("The retained agent working environment is no longer available.") raise _agent_backend_failure_to_exception(terminal) raise AgentBackendError("Agent backend run did not complete successfully.") @@ -719,38 +719,18 @@ class AgentAppRunner: message_id, exc_info=True, ) - if preserve_session: - superseded_sessions = self._load_superseded_sessions(scope=scope) - self._publish_terminal_answer( - queue_manager=queue_manager, - model_name=model_name, - answer=answer, - query=query, - usage=_llm_usage_from_agent_backend(terminal.usage), - ) - session_saved = self._save_session( - scope=scope, - backend_run_id=terminal.run_id, - snapshot=terminal.session_snapshot, - runtime_layer_specs=extract_runtime_layer_specs(runtime.request.composition), - ) - if session_saved: - self._cleanup_superseded_sessions(superseded_sessions) - else: - # The backend has already accepted a terminal success with - # delete-on-exit semantics. Local publish/persistence errors must - # not keep the API-side session row active, and cleanup failures - # must not replace the original publish/error outcome. - try: - self._publish_terminal_answer( - queue_manager=queue_manager, - model_name=model_name, - answer=answer, - query=query, - usage=_llm_usage_from_agent_backend(terminal.usage), - ) - finally: - self._mark_session_cleaned(scope=scope, backend_run_id=terminal.run_id) + self._publish_terminal_answer( + queue_manager=queue_manager, + model_name=model_name, + answer=answer, + query=query, + usage=_llm_usage_from_agent_backend(terminal.usage), + ) + self._save_session( + scope=scope, + binding_id=runtime.binding_id, + snapshot=terminal.session_snapshot, + ) def _build_session_scope( self, @@ -758,8 +738,11 @@ class AgentAppRunner: dify_context: DifyRunContext, agent_id: str, agent_config_snapshot_id: str, + home_snapshot_id: str | None, conversation_id: str, session_scope_snapshot_id: str | None | _DefaultSessionScopeSnapshotId, + agent_config_version_kind: AgentConfigVersionKind, + build_draft_id: str | None = None, ) -> AgentAppSessionScope: if isinstance(session_scope_snapshot_id, _DefaultSessionScopeSnapshotId): effective_session_scope_snapshot_id: str | None = agent_config_snapshot_id @@ -770,7 +753,10 @@ class AgentAppRunner: app_id=dify_context.app_id, conversation_id=conversation_id, agent_id=agent_id, - agent_config_snapshot_id=effective_session_scope_snapshot_id, + agent_config_snapshot_id=effective_session_scope_snapshot_id or agent_config_snapshot_id, + home_snapshot_id=home_snapshot_id, + agent_config_version_kind=agent_config_version_kind, + build_draft_id=build_draft_id, ) def _build_runtime( @@ -781,14 +767,15 @@ class AgentAppRunner: agent_config_snapshot_id: str, agent_config_version_kind: Literal["snapshot", "draft", "build_draft"], agent_soul: AgentSoulConfig, + binding_id: str, + backend_binding_ref: str, conversation_id: str, query: str, idempotency_key: str, - stored: StoredAgentAppSession | None, + stored: StoredAgentAppSession, message_id: str | None, - suspend_on_exit: bool, ) -> AgentAppRuntimeRequest: - session_snapshot = stored.session_snapshot if stored is not None else None + session_snapshot = stored.session_snapshot deferred_tool_results = ( self._resolve_pending_ask_human(stored=stored, dify_context=dify_context, message_id=message_id) if message_id is not None @@ -804,9 +791,10 @@ class AgentAppRunner: conversation_id=conversation_id, user_query=query, idempotency_key=idempotency_key, + binding_id=binding_id, + backend_binding_ref=backend_binding_ref, session_snapshot=session_snapshot, deferred_tool_results=deferred_tool_results, - suspend_on_exit=suspend_on_exit, ) ) @@ -843,9 +831,8 @@ class AgentAppRunner: # second run with the human's answer (ENG-637/638 columns, conversation owner). self._save_session( scope=scope, - backend_run_id=terminal.run_id, + binding_id=runtime.binding_id, snapshot=terminal.session_snapshot, - runtime_layer_specs=extract_runtime_layer_specs(runtime.request.composition), pending_form_id=created.form_id, pending_tool_call_id=terminal.deferred_tool_call.tool_call_id, ) @@ -862,12 +849,12 @@ class AgentAppRunner: def _resolve_pending_ask_human( self, *, - stored: StoredAgentAppSession | None, + stored: StoredAgentAppSession, dify_context: DifyRunContext, message_id: str, ) -> DeferredToolResultsPayload | None: """Build deferred_tool_results when a pending ask_human form is answered.""" - if stored is None or stored.pending_form_id is None or stored.pending_tool_call_id is None: + if stored.pending_form_id is None or stored.pending_tool_call_id is None: return None outcome = resolve_ask_human_form( form_id=stored.pending_form_id, @@ -1036,18 +1023,16 @@ class AgentAppRunner: self, *, scope: AgentAppSessionScope, - backend_run_id: str, + binding_id: str, snapshot: Any, - runtime_layer_specs: Any, pending_form_id: str | None = None, pending_tool_call_id: str | None = None, ) -> bool: try: self._session_store.save_active_snapshot( scope=scope, - backend_run_id=backend_run_id, + binding_id=binding_id, snapshot=snapshot, - runtime_layer_specs=runtime_layer_specs, pending_form_id=pending_form_id, pending_tool_call_id=pending_tool_call_id, ) @@ -1064,87 +1049,6 @@ class AgentAppRunner: ) return False - def _load_superseded_sessions(self, *, scope: AgentAppSessionScope) -> list[StoredAgentAppSession]: - try: - stored_sessions = self._session_store.list_active_sessions_for_conversation( - tenant_id=scope.tenant_id, - app_id=scope.app_id, - conversation_id=scope.conversation_id, - ) - except Exception: - logger.warning( - "Failed to load existing Agent App conversation sessions before snapshot save: " - "tenant_id=%s app_id=%s conversation_id=%s agent_id=%s", - scope.tenant_id, - scope.app_id, - scope.conversation_id, - scope.agent_id, - exc_info=True, - ) - return [] - - return [stored for stored in stored_sessions if stored.scope != scope] - - def _cleanup_superseded_sessions(self, stored_sessions: list[StoredAgentAppSession]) -> None: - for stored_session in stored_sessions: - try: - if stored_session.runtime_layer_specs: - payload = AgentBackendSessionCleanupPayload( - session_snapshot=stored_session.session_snapshot, - runtime_layer_specs=stored_session.runtime_layer_specs, - idempotency_key=( - f"{stored_session.scope.tenant_id}:{stored_session.scope.app_id}:" - f"{stored_session.scope.conversation_id}:{stored_session.scope.agent_id}:" - f"{stored_session.scope.agent_config_snapshot_id or 'no-config'}:" - f"superseded-session-cleanup:{stored_session.backend_run_id or 'no-run'}" - ), - metadata={ - "tenant_id": stored_session.scope.tenant_id, - "app_id": stored_session.scope.app_id, - "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, - "previous_agent_backend_run_id": stored_session.backend_run_id, - }, - ) - cleanup_conversation_agent_runtime_session.delay(payload.model_dump(mode="json")) - except Exception: - logger.warning( - "Failed to enqueue Agent backend cleanup for superseded Agent App session: " - "tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s", - stored_session.scope.tenant_id, - stored_session.scope.app_id, - stored_session.scope.conversation_id, - stored_session.scope.agent_id, - stored_session.backend_run_id, - exc_info=True, - ) - - def _mark_session_cleaned( - self, - *, - scope: AgentAppSessionScope, - backend_run_id: str, - ) -> None: - """Best-effort delete-on-exit cleanup for the API-side session row. - - Once the Agent backend reaches a terminal event, cleanup persistence - must not replace the original publish/error outcome for that turn. - """ - try: - self._session_store.mark_cleaned(scope=scope, backend_run_id=backend_run_id) - except Exception: - logger.warning( - "Failed to retire Agent App conversation session after delete-on-exit: " - "tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s", - scope.tenant_id, - scope.app_id, - scope.conversation_id, - scope.agent_id, - backend_run_id, - exc_info=True, - ) - @staticmethod def _terminal_output_to_answer(output: JsonValue) -> str: """Normalize the backend's terminal output to assistant text. diff --git a/api/core/app/apps/agent_app/runtime_request_builder.py b/api/core/app/apps/agent_app/runtime_request_builder.py index c90ab77d681..ccad454d8a8 100644 --- a/api/core/app/apps/agent_app/runtime_request_builder.py +++ b/api/core/app/apps/agent_app/runtime_request_builder.py @@ -70,11 +70,12 @@ class AgentAppRuntimeBuildContext: conversation_id: str user_query: str idempotency_key: str + binding_id: str + backend_binding_ref: str agent_config_version_kind: Literal["snapshot", "draft", "build_draft"] = "snapshot" session_snapshot: CompositorSessionSnapshot | None = None # ENG-638: set when resuming a chat turn after a submitted ask_human form. deferred_tool_results: DeferredToolResultsPayload | None = None - suspend_on_exit: bool = True @dataclass(frozen=True, slots=True) @@ -82,6 +83,7 @@ class AgentAppRuntimeRequest: request: CreateRunRequest redacted_request: dict[str, Any] metadata: dict[str, Any] + binding_id: str class AgentAppRuntimeRequestBuilder: @@ -160,6 +162,7 @@ class AgentAppRuntimeRequestBuilder: invoke_from=cast(DifyExecutionContextInvokeFrom, context.dify_context.invoke_from.value), agent_mode="agent_app", ), + backend_binding_ref=context.backend_binding_ref, # ENG-616: expand slash-menu mention tokens to canonical names so # no frontend-internal {{#…#}} marker ever reaches the model. agent_soul_prompt=expand_prompt_mentions(agent_soul.prompt.system_prompt, soul_prompt_resolver).strip() @@ -175,13 +178,17 @@ class AgentAppRuntimeRequestBuilder: shell_config=build_shell_layer_config(agent_soul), session_snapshot=context.session_snapshot, deferred_tool_results=context.deferred_tool_results, - suspend_on_exit=context.suspend_on_exit, idempotency_key=context.idempotency_key, metadata=metadata, ) ) redacted = cast(dict[str, Any], redact_for_agent_backend_log(request)) - return AgentAppRuntimeRequest(request=request, redacted_request=redacted, metadata=metadata) + return AgentAppRuntimeRequest( + request=request, + redacted_request=redacted, + metadata=metadata, + binding_id=context.binding_id, + ) def _build_tool_layers( self, diff --git a/api/core/app/apps/agent_app/session_store.py b/api/core/app/apps/agent_app/session_store.py index 7696155a1c4..07478da10f1 100644 --- a/api/core/app/apps/agent_app/session_store.py +++ b/api/core/app/apps/agent_app/session_store.py @@ -1,255 +1,176 @@ -"""Conversation-keyed Agent backend session store for the Agent App type. - -Shares the unified ``agent_runtime_sessions`` table with the workflow Agent -Node store, but owns rows with ``owner_type = conversation``: one Agent App -conversation maps to one Agent session, so multi-turn chat re-enters the same -``session_snapshot``. Cross-conversation memory (PRD Global / Per app) is a -phase-2 concern and not modeled here. -""" +"""Persist and resolve the exact participant owned by an Agent App caller.""" from __future__ import annotations -from dataclasses import dataclass, field +from dataclasses import dataclass from agenton.compositor import CompositorSessionSnapshot -from dify_agent.protocol import RuntimeLayerSpec -from pydantic import TypeAdapter from sqlalchemy import select +from sqlalchemy.orm import Session from core.db.session_factory import session_factory -from libs.datetime_utils import naive_utc_now from models.agent import ( - AgentRuntimeSession, - AgentRuntimeSessionOwnerType, - AgentRuntimeSessionStatus, + AgentConfigDraft, + AgentConfigDraftType, + AgentConfigVersionKind, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, +) +from models.model import App, Conversation +from services.agent.workspace_service import ( + AgentWorkspaceNotFoundError, + AgentWorkspaceService, + WorkspaceOwnerScope, ) - -_RUNTIME_LAYER_SPECS_ADAPTER: TypeAdapter[list[RuntimeLayerSpec]] = TypeAdapter(list[RuntimeLayerSpec]) - - -def _serialize_runtime_layer_specs(specs: list[RuntimeLayerSpec]) -> str: - return _RUNTIME_LAYER_SPECS_ADAPTER.dump_json(specs).decode() - - -def _deserialize_runtime_layer_specs(value: str | None) -> list[RuntimeLayerSpec]: - if not value: - return [] - return _RUNTIME_LAYER_SPECS_ADAPTER.validate_json(value) @dataclass(frozen=True, slots=True) class AgentAppSessionScope: - """Identity of one Agent App conversation session.""" - tenant_id: str app_id: str conversation_id: str agent_id: str - agent_config_snapshot_id: str | None + agent_config_snapshot_id: str + home_snapshot_id: str | None + agent_config_version_kind: AgentConfigVersionKind = AgentConfigVersionKind.SNAPSHOT + build_draft_id: str | None = None + + @property + def workspace_owner(self) -> WorkspaceOwnerScope: + owner_type = ( + AgentWorkspaceOwnerType.BUILD_DRAFT if self.build_draft_id else AgentWorkspaceOwnerType.CONVERSATION + ) + return WorkspaceOwnerScope( + tenant_id=self.tenant_id, + app_id=self.app_id, + owner_type=owner_type, + owner_id=self.build_draft_id or self.conversation_id, + ) @dataclass(frozen=True, slots=True) class StoredAgentAppSession: - """Persisted Agent App conversation session with reusable runtime specs.""" - scope: AgentAppSessionScope - session_snapshot: CompositorSessionSnapshot - backend_run_id: str | None - runtime_layer_specs: list[RuntimeLayerSpec] = field(default_factory=list) - # ENG-635: set while the conversation turn is paused on a dify.ask_human - # deferred call, awaiting a HITL form submission. + binding_id: str + workspace_id: str + backend_binding_ref: str + session_snapshot: CompositorSessionSnapshot | None pending_form_id: str | None = None pending_tool_call_id: str | None = None -class AgentAppRuntimeSessionStore: - """Persists Agent backend session snapshots for Agent App conversations.""" +class AgentAppWorkspaceStore: + """Resolve Agent App sessions through a caller-owned Binding pointer.""" - def load_active_snapshot(self, scope: AgentAppSessionScope) -> CompositorSessionSnapshot | None: - stored = self.load_active_session(scope) - return stored.session_snapshot if stored is not None else None - - def load_active_session(self, scope: AgentAppSessionScope) -> StoredAgentAppSession | None: + def load_or_create(self, scope: AgentAppSessionScope) -> StoredAgentAppSession: with session_factory.create_session() as session: - row = session.scalar(self._active_stmt(scope)) - if row is None: - return None - return StoredAgentAppSession( - scope=scope, - session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot), - backend_run_id=row.backend_run_id, - runtime_layer_specs=_deserialize_runtime_layer_specs(row.composition_layer_specs), - pending_form_id=row.pending_form_id, - pending_tool_call_id=row.pending_tool_call_id, - ) - - def load_active_session_for_conversation( - self, *, tenant_id: str, app_id: str, conversation_id: str - ) -> StoredAgentAppSession | None: - """Load the latest ACTIVE session for one conversation-level sandbox lookup. - - Sandbox inspection only knows the product locator - ``tenant_id + app_id + conversation_id``; it does not know which - ``agent_id`` or Agent Soul snapshot produced the active shell session. - This method therefore resolves the newest ACTIVE conversation-owned row - for that conversation and returns both the resumable snapshot and the - persisted non-sensitive runtime layer specs needed to build a - ``SandboxLocator``. - """ - stmt = ( - select(AgentRuntimeSession) - .where( - AgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION, - AgentRuntimeSession.tenant_id == tenant_id, - AgentRuntimeSession.app_id == app_id, - AgentRuntimeSession.conversation_id == conversation_id, - AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE, - ) - .order_by(AgentRuntimeSession.updated_at.desc()) - ) - with session_factory.create_session() as session: - row = session.scalar(stmt) - if row is None: - return None - return StoredAgentAppSession( - scope=AgentAppSessionScope( - tenant_id=row.tenant_id, - app_id=row.app_id, - conversation_id=row.conversation_id or "", - agent_id=row.agent_id, - agent_config_snapshot_id=row.agent_config_snapshot_id or "", - ), - session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot), - backend_run_id=row.backend_run_id, - runtime_layer_specs=_deserialize_runtime_layer_specs(row.composition_layer_specs), - ) - - def list_active_sessions_for_conversation( - self, *, tenant_id: str, app_id: str, conversation_id: str - ) -> list[StoredAgentAppSession]: - """List all ACTIVE conversation-owned sessions for lifecycle cleanup.""" - stmt = ( - select(AgentRuntimeSession) - .where( - AgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION, - AgentRuntimeSession.tenant_id == tenant_id, - AgentRuntimeSession.app_id == app_id, - AgentRuntimeSession.conversation_id == conversation_id, - AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE, - ) - .order_by(AgentRuntimeSession.updated_at.desc()) - ) - with session_factory.create_session() as session: - rows = session.scalars(stmt).all() - return [ - StoredAgentAppSession( - scope=AgentAppSessionScope( - tenant_id=row.tenant_id, - app_id=row.app_id, - conversation_id=row.conversation_id or "", - agent_id=row.agent_id, - agent_config_snapshot_id=row.agent_config_snapshot_id, - ), - session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot), - backend_run_id=row.backend_run_id, - runtime_layer_specs=_deserialize_runtime_layer_specs(row.composition_layer_specs), - pending_form_id=row.pending_form_id, - pending_tool_call_id=row.pending_tool_call_id, + caller = self._load_caller(session=session, scope=scope) + binding_id = caller.agent_workspace_binding_id + if binding_id is None: + binding = AgentWorkspaceService.create_binding( + session=session, + scope=scope.workspace_owner, + agent_id=scope.agent_id, + base_home_snapshot_id=scope.home_snapshot_id, + agent_config_version_id=scope.agent_config_snapshot_id, + agent_config_version_kind=scope.agent_config_version_kind, ) - for row in rows - ] + caller.agent_workspace_binding_id = binding.id + session.commit() + else: + binding = self._get_binding(session=session, scope=scope, binding_id=binding_id) + return self._stored(scope, binding) + + @staticmethod + def _load_caller(*, session: Session, scope: AgentAppSessionScope) -> Conversation | AgentConfigDraft: + if scope.build_draft_id is not None: + if scope.agent_config_version_kind != AgentConfigVersionKind.BUILD_DRAFT: + raise AgentWorkspaceNotFoundError("Build Draft caller requires build_draft generation") + draft = session.scalar( + select(AgentConfigDraft).where( + AgentConfigDraft.id == scope.build_draft_id, + AgentConfigDraft.tenant_id == scope.tenant_id, + AgentConfigDraft.agent_id == scope.agent_id, + AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD, + ) + ) + if draft is None: + raise AgentWorkspaceNotFoundError("Build Draft caller is unavailable") + return draft + if scope.agent_config_version_kind == AgentConfigVersionKind.BUILD_DRAFT: + raise AgentWorkspaceNotFoundError("Build Draft caller ID is required") + conversation = session.scalar( + select(Conversation) + .join(App, App.id == Conversation.app_id) + .where( + App.tenant_id == scope.tenant_id, + Conversation.id == scope.conversation_id, + Conversation.app_id == scope.app_id, + Conversation.is_deleted.is_(False), + ) + ) + if conversation is None: + raise AgentWorkspaceNotFoundError("Conversation caller is unavailable") + return conversation + + @staticmethod + def _get_binding( + *, + session: Session, + scope: AgentAppSessionScope, + binding_id: str, + ) -> AgentWorkspaceBinding: + binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=scope.tenant_id, + binding_id=binding_id, + expected_owner_scope=scope.workspace_owner, + ) + if binding is None or binding.agent_id != scope.agent_id: + raise AgentWorkspaceNotFoundError("Caller participant Binding is unavailable") + AgentWorkspaceService.validate_binding_generation( + binding, + base_home_snapshot_id=scope.home_snapshot_id, + agent_config_version_id=scope.agent_config_snapshot_id, + agent_config_version_kind=scope.agent_config_version_kind, + ) + return binding def save_active_snapshot( self, *, scope: AgentAppSessionScope, - backend_run_id: str, + binding_id: str, snapshot: CompositorSessionSnapshot | None, - runtime_layer_specs: list[RuntimeLayerSpec], pending_form_id: str | None = None, pending_tool_call_id: str | None = None, ) -> None: - """Persist the current conversation snapshot and enforce one ACTIVE row. - - Agent App chat treats one conversation as one resumable runtime shell. - Saving the latest snapshot therefore upserts the scoped row back to - ACTIVE and retires any other ACTIVE conversation-owned rows for the - same ``tenant_id + app_id + conversation_id`` so later lookups see a - single active session. - """ if snapshot is None: return - snapshot_json = snapshot.model_dump_json() - runtime_layer_specs_json = _serialize_runtime_layer_specs(runtime_layer_specs) - with session_factory.create_session() as session: - row = session.scalar(self._scope_stmt(scope)) - if row is None: - row = AgentRuntimeSession( - tenant_id=scope.tenant_id, - app_id=scope.app_id, - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id=scope.agent_id, - agent_config_snapshot_id=scope.agent_config_snapshot_id, - conversation_id=scope.conversation_id, - backend_run_id=backend_run_id, - session_snapshot=snapshot_json, - composition_layer_specs=runtime_layer_specs_json, - status=AgentRuntimeSessionStatus.ACTIVE, - pending_form_id=pending_form_id, - pending_tool_call_id=pending_tool_call_id, - ) - session.add(row) - else: - row.backend_run_id = backend_run_id - row.session_snapshot = snapshot_json - row.composition_layer_specs = runtime_layer_specs_json - row.status = AgentRuntimeSessionStatus.ACTIVE - row.cleaned_at = None - # Set (or clear, when omitted) the ask_human pause correlation. - row.pending_form_id = pending_form_id - row.pending_tool_call_id = pending_tool_call_id - session.flush() - other_rows = session.scalars( - select(AgentRuntimeSession).where( - AgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION, - AgentRuntimeSession.tenant_id == scope.tenant_id, - AgentRuntimeSession.app_id == scope.app_id, - AgentRuntimeSession.conversation_id == scope.conversation_id, - AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE, - AgentRuntimeSession.id != row.id, - ) - ).all() - for other_row in other_rows: - other_row.status = AgentRuntimeSessionStatus.CLEANED - other_row.cleaned_at = naive_utc_now() - session.commit() - - def mark_cleaned(self, *, scope: AgentAppSessionScope, backend_run_id: str | None = None) -> None: - with session_factory.create_session() as session: - row = session.scalar(self._active_stmt(scope)) - if row is None: - return - if backend_run_id is not None: - row.backend_run_id = backend_run_id - row.status = AgentRuntimeSessionStatus.CLEANED - row.cleaned_at = naive_utc_now() - session.commit() + AgentWorkspaceService.save_binding_session_snapshot( + tenant_id=scope.tenant_id, + binding_id=binding_id, + session_snapshot=snapshot.model_dump_json(), + pending_form_id=pending_form_id, + pending_tool_call_id=pending_tool_call_id, + ) @staticmethod - def _scope_stmt(scope: AgentAppSessionScope): - stmt = select(AgentRuntimeSession).where( - AgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION, - AgentRuntimeSession.tenant_id == scope.tenant_id, - AgentRuntimeSession.conversation_id == scope.conversation_id, - AgentRuntimeSession.agent_id == scope.agent_id, + def _stored(scope: AgentAppSessionScope, binding: AgentWorkspaceBinding) -> StoredAgentAppSession: + snapshot = ( + CompositorSessionSnapshot.model_validate_json(binding.session_snapshot) + if binding.session_snapshot + else None + ) + return StoredAgentAppSession( + scope=scope, + binding_id=binding.id, + workspace_id=binding.workspace_id, + backend_binding_ref=binding.backend_binding_ref, + session_snapshot=snapshot, + pending_form_id=binding.pending_form_id, + pending_tool_call_id=binding.pending_tool_call_id, ) - if scope.agent_config_snapshot_id is None: - return stmt.where(AgentRuntimeSession.agent_config_snapshot_id.is_(None)) - return stmt.where(AgentRuntimeSession.agent_config_snapshot_id == scope.agent_config_snapshot_id) - - @classmethod - def _active_stmt(cls, scope: AgentAppSessionScope): - return cls._scope_stmt(scope).where(AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE) -__all__ = ["AgentAppRuntimeSessionStore", "AgentAppSessionScope", "StoredAgentAppSession"] +__all__ = ["AgentAppSessionScope", "AgentAppWorkspaceStore", "StoredAgentAppSession"] diff --git a/api/core/app/apps/base_app_generate_response_converter.py b/api/core/app/apps/base_app_generate_response_converter.py index d30969d3714..aef54cc049e 100644 --- a/api/core/app/apps/base_app_generate_response_converter.py +++ b/api/core/app/apps/base_app_generate_response_converter.py @@ -3,9 +3,10 @@ from abc import ABC, abstractmethod from collections.abc import Generator, Mapping from typing import Any, Union, cast +from dify_agent.protocol import RunFailureType from pydantic import JsonValue -from clients.agent_backend.errors import AgentBackendError +from clients.agent_backend.errors import AgentBackendError, AgentBackendRunFailedError from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.task_entities import AppBlockingResponse, AppStreamResponse from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError @@ -117,6 +118,13 @@ class AppGenerateResponseConverter[TBlockingResponse: AppBlockingResponse](ABC): :param e: exception :return: """ + if isinstance(e, AgentBackendRunFailedError) and e.error_type == RunFailureType.AGENT_RUN_LIMIT_EXCEEDED: + return { + "code": RunFailureType.AGENT_RUN_LIMIT_EXCEEDED.value, + "status": 400, + "message": str(e), + } + error_responses: dict[type[Exception], dict[str, JsonValue]] = { ValueError: {"code": "invalid_param", "status": 400}, ProviderTokenNotInitError: {"code": "provider_not_initialize", "status": 400}, diff --git a/api/core/app/apps/base_app_generator.py b/api/core/app/apps/base_app_generator.py index 2762f99301d..aa49d90d1d9 100644 --- a/api/core/app/apps/base_app_generator.py +++ b/api/core/app/apps/base_app_generator.py @@ -1,8 +1,11 @@ +import logging +import threading from collections.abc import Generator, Mapping, Sequence from contextlib import AbstractContextManager, nullcontext from typing import TYPE_CHECKING, Any, Union, final from sqlalchemy.orm import Session +from sqlalchemy.orm.attributes import set_committed_value from core.app.apps.draft_variable_saver import ( DraftVariableSaver, @@ -17,12 +20,16 @@ from graphon.enums import NodeType from graphon.file import File, FileUploadConfig from graphon.variables.input_entities import VariableEntityType from libs.orjson import orjson_dumps -from models import Account, EndUser +from models import Account, EndUser, Workflow, WorkflowRun from services.workflow_draft_variable_service import DraftVariableSaver as DraftVariableSaverImpl if TYPE_CHECKING: from graphon.variables.input_entities import VariableEntity +logger = logging.getLogger(__name__) + +_WORKER_THREAD_JOIN_TIMEOUT_SECONDS = 300 + @final class _DebuggerDraftVariableSaver: @@ -64,6 +71,38 @@ class _DebuggerDraftVariableSaver: class BaseAppGenerator: _file_access_controller: DatabaseFileAccessController = DatabaseFileAccessController() + @staticmethod + def _restore_workflow_run_graph(*, session: Session, workflow: Workflow, workflow_run_id: str | None) -> None: + if workflow_run_id is None: + raise ValueError("Workflow run id is required when resuming") + workflow_run = session.get(WorkflowRun, workflow_run_id) + if workflow_run is None or workflow_run.graph is None: + raise ValueError(f"Workflow run graph not found: {workflow_run_id}") + set_committed_value(workflow, "graph", workflow_run.graph) + + @staticmethod + def _join_worker_thread(worker_thread: threading.Thread) -> None: + # Bound the wait so a leaked app worker cannot occupy an execution slot indefinitely. + worker_thread.join(timeout=_WORKER_THREAD_JOIN_TIMEOUT_SECONDS) + if worker_thread.is_alive(): + logger.warning( + "Possible app worker thread leak: thread_name=%s timeout_seconds=%s; " + "continuing without waiting further to avoid occupying an execution slot indefinitely", + worker_thread.name, + _WORKER_THREAD_JOIN_TIMEOUT_SECONDS, + ) + + @staticmethod + def _wrap_stream_with_worker_thread_join[ResponseT]( + response_stream: Generator[ResponseT, None, None], + worker_thread: threading.Thread, + ) -> Generator[ResponseT, None, None]: + """Keep the producer owned by the response stream until both finish.""" + try: + yield from response_stream + finally: + BaseAppGenerator._join_worker_thread(worker_thread) + @staticmethod def _bind_file_access_scope( *, diff --git a/api/core/app/apps/common/workflow_response_converter.py b/api/core/app/apps/common/workflow_response_converter.py index 236b8c58d3a..e95a68c0b44 100644 --- a/api/core/app/apps/common/workflow_response_converter.py +++ b/api/core/app/apps/common/workflow_response_converter.py @@ -76,7 +76,7 @@ from graphon.runtime import GraphRuntimeState from graphon.variables.segments import ArrayFileSegment, FileSegment, Segment from graphon.variables.variables import Variable from graphon.workflow_type_encoder import WorkflowRuntimeTypeConverter -from libs.datetime_utils import naive_utc_now +from libs.datetime_utils import naive_utc_now, to_utc_timestamp from models import Account, EndUser from models.human_input import HumanInputForm from models.workflow import WorkflowRun @@ -371,7 +371,7 @@ class WorkflowResponseConverter: pause_reasons, dispositions_by_form_id=dispositions_by_form_id, expiration_times_by_form_id={ - form_id: int(expiration_time.timestamp()) + form_id: to_utc_timestamp(expiration_time) for form_id, expiration_time in expiration_times_by_form_id.items() }, ) @@ -399,7 +399,7 @@ class WorkflowResponseConverter: form_token=disposition.form_token if disposition else None, approval_channels=list(disposition.approval_channels) if disposition else [], resolved_default_values=reason.resolved_default_values, - expiration_time=int(expiration_time.timestamp()), + expiration_time=to_utc_timestamp(expiration_time), ), ) ) @@ -452,7 +452,7 @@ class WorkflowResponseConverter: data=HumanInputFormTimeoutResponse.Data( node_id=event.node_id, node_title=event.node_title, - expiration_time=int(event.expiration_time.timestamp()), + expiration_time=to_utc_timestamp(event.expiration_time), ), ) @@ -856,7 +856,9 @@ class WorkflowResponseConverter: return tuple(flattened_files) @classmethod - def _fetch_files_from_variable_value(cls, value: Union[dict, list, Segment]) -> Sequence[Mapping[str, Any]]: + def _fetch_files_from_variable_value( + cls, value: Union[dict, list, Segment, File, None] + ) -> Sequence[Mapping[str, Any]]: """ Fetch files from variable value :param value: variable value diff --git a/api/core/app/apps/completion/app_generator.py b/api/core/app/apps/completion/app_generator.py index 54634fe2664..3d1338a8bf6 100644 --- a/api/core/app/apps/completion/app_generator.py +++ b/api/core/app/apps/completion/app_generator.py @@ -283,6 +283,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator): """ Generate App response. + :param session: caller-owned database session used for message and historical model-config reads :param app_model: App :param message_id: message ID :param user: account or end user diff --git a/api/core/app/apps/message_based_app_queue_manager.py b/api/core/app/apps/message_based_app_queue_manager.py index b253d93ee52..1be1e532694 100644 --- a/api/core/app/apps/message_based_app_queue_manager.py +++ b/api/core/app/apps/message_based_app_queue_manager.py @@ -10,6 +10,7 @@ from core.app.entities.queue_entities import ( QueueErrorEvent, QueueMessageEndEvent, QueueStopEvent, + QueueWorkflowPausedEvent, ) from models.model import AppMode @@ -43,7 +44,12 @@ class MessageBasedAppQueueManager(AppQueueManager): self._q.put(message) if isinstance( - event, QueueStopEvent | QueueErrorEvent | QueueMessageEndEvent | QueueAdvancedChatMessageEndEvent + event, + QueueStopEvent + | QueueErrorEvent + | QueueMessageEndEvent + | QueueAdvancedChatMessageEndEvent + | QueueWorkflowPausedEvent, ): self.stop_listen(execution_terminal=True) diff --git a/api/core/app/apps/pipeline/pipeline_generator.py b/api/core/app/apps/pipeline/pipeline_generator.py index 3eb93e7c08a..1c3787518d4 100644 --- a/api/core/app/apps/pipeline/pipeline_generator.py +++ b/api/core/app/apps/pipeline/pipeline_generator.py @@ -188,7 +188,7 @@ class PipelineGenerator(BaseAppGenerator): datasource_type=datasource_type, datasource_info=datasource_info, dataset_id=dataset.id, - original_document_id=args.get("original_document_id"), + original_document_id=None if is_retry else args.get("original_document_id"), start_node_id=start_node_id, batch=batch, document_id=document_id, @@ -351,17 +351,28 @@ class PipelineGenerator(BaseAppGenerator): user, tenant_id=pipeline.tenant_id, ) - # return response or stream generator - response = self._handle_response( - application_generate_entity=application_generate_entity, - workflow=workflow, - queue_manager=queue_manager, - user=user, - stream=streaming, - draft_var_saver_factory=draft_var_saver_factory, - ) + try: + response = self._handle_response( + application_generate_entity=application_generate_entity, + workflow=workflow, + queue_manager=queue_manager, + user=user, + stream=streaming, + draft_var_saver_factory=draft_var_saver_factory, + ) + converted_response = WorkflowAppGenerateResponseConverter.convert( + response=response, + invoke_from=invoke_from, + ) + except BaseException: + self._join_worker_thread(worker_thread) + raise - return WorkflowAppGenerateResponseConverter.convert(response=response, invoke_from=invoke_from) + if isinstance(converted_response, Generator): + return self._wrap_stream_with_worker_thread_join(converted_response, worker_thread) + + self._join_worker_thread(worker_thread) + return converted_response def single_iteration_generate( self, diff --git a/api/core/app/apps/workflow/active_workflow_tasks.py b/api/core/app/apps/workflow/active_workflow_tasks.py index 4aa23ad1ccf..ad2d8cfb5e0 100644 --- a/api/core/app/apps/workflow/active_workflow_tasks.py +++ b/api/core/app/apps/workflow/active_workflow_tasks.py @@ -1,7 +1,7 @@ """In-process registry for workflow application task IDs.""" import threading -from collections.abc import Iterator +from collections.abc import Generator from contextlib import contextmanager _active_task_ids: set[str] = set() @@ -9,7 +9,7 @@ _active_task_ids_lock = threading.RLock() @contextmanager -def active_workflow_task(task_id: str) -> Iterator[None]: +def active_workflow_task(task_id: str) -> Generator[None]: """Register a workflow application task ID for the duration of a workflow run.""" if not task_id: raise ValueError("task_id must not be empty") diff --git a/api/core/app/apps/workflow/app_generator.py b/api/core/app/apps/workflow/app_generator.py index fb5393d7730..2d6c1512c48 100644 --- a/api/core/app/apps/workflow/app_generator.py +++ b/api/core/app/apps/workflow/app_generator.py @@ -405,17 +405,28 @@ class WorkflowAppGenerator(BaseAppGenerator): tenant_id=app_model.tenant_id, ) - # return response or stream generator - response = self._handle_response( - application_generate_entity=application_generate_entity, - workflow=workflow, - queue_manager=queue_manager, - user=user, - draft_var_saver_factory=draft_var_saver_factory, - stream=streaming, - ) + try: + response = self._handle_response( + application_generate_entity=application_generate_entity, + workflow=workflow, + queue_manager=queue_manager, + user=user, + draft_var_saver_factory=draft_var_saver_factory, + stream=streaming, + ) + converted_response = WorkflowAppGenerateResponseConverter.convert( + response=response, + invoke_from=invoke_from, + ) + except BaseException: + self._join_worker_thread(worker_thread) + raise - return WorkflowAppGenerateResponseConverter.convert(response=response, invoke_from=invoke_from) + if isinstance(converted_response, Generator): + return self._wrap_stream_with_worker_thread_join(converted_response, worker_thread) + + self._join_worker_thread(worker_thread) + return converted_response def single_iteration_generate( self, @@ -637,6 +648,12 @@ class WorkflowAppGenerator(BaseAppGenerator): raise ValueError("Workflow not found") workflow = self._ensure_snippet_start_node_in_worker(session=session, workflow=workflow) + if graph_runtime_state is not None: + self._restore_workflow_run_graph( + session=session, + workflow=workflow, + workflow_run_id=application_generate_entity.workflow_execution_id, + ) # Determine system_user_id based on invocation source is_external_api_call = application_generate_entity.invoke_from in { diff --git a/api/core/app/apps/workflow/app_queue_manager.py b/api/core/app/apps/workflow/app_queue_manager.py index 67df3044fb2..da99d034477 100644 --- a/api/core/app/apps/workflow/app_queue_manager.py +++ b/api/core/app/apps/workflow/app_queue_manager.py @@ -9,6 +9,7 @@ from core.app.entities.queue_entities import ( QueueStopEvent, QueueWorkflowFailedEvent, QueueWorkflowPartialSuccessEvent, + QueueWorkflowPausedEvent, QueueWorkflowSucceededEvent, WorkflowQueueMessage, ) @@ -32,6 +33,11 @@ class WorkflowAppQueueManager(AppQueueManager): self._q.put(message) + # A pause ends only the current listener segment; the workflow stays PAUSED and + # resumes with the same task ID. Without this marker, listen() cleanup calls + # _abort_execution(), whose stop flag and abort command can stop the resumed run. + # This is a compatibility workaround: cancellation policy belongs to the execution + # owner, not the response-stream listener. if isinstance( event, QueueStopEvent @@ -39,6 +45,7 @@ class WorkflowAppQueueManager(AppQueueManager): | QueueMessageEndEvent | QueueWorkflowSucceededEvent | QueueWorkflowFailedEvent + | QueueWorkflowPausedEvent | QueueWorkflowPartialSuccessEvent, ): self.stop_listen(execution_terminal=True) diff --git a/api/core/app/apps/workflow/app_runner.py b/api/core/app/apps/workflow/app_runner.py index 95c9d777ebd..d7427408792 100644 --- a/api/core/app/apps/workflow/app_runner.py +++ b/api/core/app/apps/workflow/app_runner.py @@ -10,11 +10,11 @@ from core.app.apps.workflow.command_channels import ( CombinedCommandChannel, ) from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner -from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity +from core.app.entities.app_invoke_entities import DifyRunContext, InvokeFrom, WorkflowAppGenerateEntity from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository from core.workflow.node_factory import get_default_root_node_id -from core.workflow.nodes.agent_v2.session_cleanup_layer import build_workflow_agent_session_cleanup_layer +from core.workflow.nodes.agent_v2.workspace_retirement_layer import build_workflow_agent_workspace_retirement_layer from core.workflow.snippet_start import get_compatible_start_aliases from core.workflow.system_variables import build_bootstrap_variables, build_system_variables from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool @@ -197,7 +197,18 @@ class WorkflowAppRunner(WorkflowBasedAppRunner): ) workflow_entry.graph_engine.layer(persistence_layer) - workflow_entry.graph_engine.layer(build_workflow_agent_session_cleanup_layer()) + workflow_entry.graph_engine.layer( + build_workflow_agent_workspace_retirement_layer( + dify_run_context=DifyRunContext( + tenant_id=self._workflow.tenant_id, + app_id=self._workflow.app_id, + user_id=self.application_generate_entity.user_id, + user_from=user_from, + invoke_from=invoke_from, + trace_session_id=self.application_generate_entity.extras.get("trace_session_id"), + ) + ) + ) for layer in self._graph_engine_layers: workflow_entry.graph_engine.layer(layer) diff --git a/api/core/app/apps/workflow_app_runner.py b/api/core/app/apps/workflow_app_runner.py index 3d2857f130a..f8cf6e4fc9c 100644 --- a/api/core/app/apps/workflow_app_runner.py +++ b/api/core/app/apps/workflow_app_runner.py @@ -42,6 +42,7 @@ from core.workflow.node_factory import ( get_default_root_node_id, resolve_workflow_node_class, ) +from core.workflow.nodes.agent.events import NodeRunAgentLogEvent from core.workflow.nodes.human_input.boundary import enrich_graph_pause_reasons from core.workflow.nodes.human_input.pause_reason import HumanInputRequired from core.workflow.system_variables import ( @@ -55,6 +56,7 @@ from core.workflow.variable_pool_initializer import add_variables_to_pool from core.workflow.workflow_entry import WorkflowEntry from core.workflow.workflow_run_outputs import project_node_outputs_for_workflow_run from graphon.entities.graph_config import NodeConfigDictAdapter +from graphon.entities.pause_reason import HitlRequired from graphon.graph import Graph from graphon.graph_engine.layers import GraphEngineLayer from graphon.graph_events import ( @@ -65,7 +67,6 @@ from graphon.graph_events import ( GraphRunPausedEvent, GraphRunStartedEvent, GraphRunSucceededEvent, - NodeRunAgentLogEvent, NodeRunExceptionEvent, NodeRunFailedEvent, NodeRunHumanInputFormFilledEvent, @@ -433,7 +434,9 @@ class WorkflowBasedAppRunner: ) case GraphRunPausedEvent(): runtime_state = workflow_entry.graph_engine.graph_runtime_state - paused_nodes = runtime_state.get_paused_nodes() + paused_nodes = list( + dict.fromkeys(reason.node_id for reason in event.reasons if isinstance(reason, HitlRequired)) + ) enriched_reasons = enrich_graph_pause_reasons( reasons=event.reasons, form_repository=HumanInputFormSubmissionRepository(), diff --git a/api/core/app/entities/app_invoke_entities.py b/api/core/app/entities/app_invoke_entities.py index b5f7515cd9f..b4e2ced08aa 100644 --- a/api/core/app/entities/app_invoke_entities.py +++ b/api/core/app/entities/app_invoke_entities.py @@ -15,8 +15,6 @@ if TYPE_CHECKING: DIFY_RUN_CONTEXT_KEY = "_dify" -AGENT_RUNTIME_EXIT_INTENT_ARG = "_agent_runtime_exit_intent" -type AgentRuntimeExitIntent = Literal["suspend", "delete"] class UserFrom(StrEnum): @@ -227,12 +225,8 @@ class AgentAppGenerateEntity(ChatAppGenerateEntity): backend should read from: immutable snapshot, shared draft, or per-user build draft. - ``agent_runtime_session_snapshot_id`` carries the runtime session scope - used to resume or suspend within the same editable config surface. - - ``agent_runtime_exit_intent`` is API-internal lifecycle policy for the - Agent backend session after this turn finishes. Normal chat/resume turns - suspend on exit; build-chat finalization deletes the backend runtime. + ``agent_session_scope_config_version_id`` identifies the draft or immutable + config version whose Workspace Binding should be reused for this session. ``prompt_file_mappings`` preserves the raw request ``files`` array for the Agent backend prompt. These references are appended to the backend prompt @@ -242,8 +236,7 @@ class AgentAppGenerateEntity(ChatAppGenerateEntity): agent_id: str agent_config_snapshot_id: str agent_config_version_kind: Literal["snapshot", "draft", "build_draft"] = "snapshot" - agent_runtime_session_snapshot_id: str | None = None - agent_runtime_exit_intent: AgentRuntimeExitIntent = "suspend" + agent_session_scope_config_version_id: str | None = None prompt_file_mappings: Sequence[JsonValue] = Field(default_factory=list) diff --git a/api/core/app/llm/__init__.py b/api/core/app/llm/__init__.py index d20a5b2344d..6f4e6909e18 100644 --- a/api/core/app/llm/__init__.py +++ b/api/core/app/llm/__init__.py @@ -5,6 +5,7 @@ from .quota import ( deduct_llm_quota_for_model, ensure_llm_quota_available, ensure_llm_quota_available_for_model, + reserve_llm_quota_for_model, ) __all__ = [ @@ -12,4 +13,5 @@ __all__ = [ "deduct_llm_quota_for_model", "ensure_llm_quota_available", "ensure_llm_quota_available_for_model", + "reserve_llm_quota_for_model", ] diff --git a/api/core/app/llm/model_access.py b/api/core/app/llm/model_access.py index 765268f7a0e..d2b8e3539fa 100644 --- a/api/core/app/llm/model_access.py +++ b/api/core/app/llm/model_access.py @@ -151,6 +151,9 @@ def fetch_model_config( credentials_provider: CredentialsProvider, model_factory: DifyModelFactory, ) -> tuple[ModelInstance, ModelConfigWithCredentialsEntity]: + if not node_data_model.provider or not node_data_model.name: + raise ValueError("LLM provider and model are required.") + if not node_data_model.mode: raise LLMModeRequiredError("LLM mode is required.") diff --git a/api/core/app/llm/quota.py b/api/core/app/llm/quota.py index d26d5d8a998..e4e84502cfa 100644 --- a/api/core/app/llm/quota.py +++ b/api/core/app/llm/quota.py @@ -7,6 +7,10 @@ with a non-LLM model. """ import warnings +from dataclasses import dataclass, field +from enum import StrEnum, auto +from typing import Any +from uuid import uuid4 from sqlalchemy import select from sqlalchemy.orm import sessionmaker @@ -23,6 +27,68 @@ from graphon.model_runtime.entities.model_entities import ModelType from libs.datetime_utils import naive_utc_now from models.provider import Provider, ProviderType from models.provider_ids import ModelProviderID +from services.credit_pool_service import CreditPoolReservation, CreditPoolService + + +class LLMQuotaReservationState(StrEnum): + RESERVED = auto() + COMMITTED = auto() + RELEASED = auto() + + +@dataclass +class LLMQuotaReservation: + """Quota reserved for one system-hosted LLM invocation.""" + + tenant_id: str + provider: str + model: str + provider_configuration: Any + quota_unit: QuotaUnit | None = None + credit_pool_reservation: CreditPoolReservation | None = None + requires_usage: bool = False + _state: LLMQuotaReservationState = field(default=LLMQuotaReservationState.RESERVED, init=False, repr=False) + + @property + def state(self) -> LLMQuotaReservationState: + return self._state + + @property + def commit_before_delivery(self) -> bool: + return self.credit_pool_reservation is not None + + def commit(self, usage: LLMUsage | None = None) -> None: + if self._state == LLMQuotaReservationState.COMMITTED: + return + if self._state == LLMQuotaReservationState.RELEASED: + raise RuntimeError("Cannot commit a released LLM quota reservation.") + + if self.credit_pool_reservation is not None: + self.credit_pool_reservation.commit() + elif self.requires_usage: + if usage is None: + raise ValueError("Accurate terminal usage is required for token-based LLM quota settlement.") + used_quota = _resolve_llm_used_quota( + system_configuration=self.provider_configuration.system_configuration, + model=self.model, + usage=usage, + ) + _deduct_used_llm_quota( + tenant_id=self.tenant_id, + provider=self.provider, + provider_configuration=self.provider_configuration, + used_quota=used_quota, + ) + + self._state = LLMQuotaReservationState.COMMITTED + + def release(self) -> None: + if self._state in {LLMQuotaReservationState.COMMITTED, LLMQuotaReservationState.RELEASED}: + return + + if self.credit_pool_reservation is not None: + self.credit_pool_reservation.release() + self._state = LLMQuotaReservationState.RELEASED def _get_provider_configuration(*, tenant_id: str, provider: str): @@ -34,6 +100,67 @@ def _get_provider_configuration(*, tenant_id: str, provider: str): return provider_configuration +def _get_current_quota_configuration(system_configuration): + return next( + ( + quota_configuration + for quota_configuration in system_configuration.quota_configurations + if quota_configuration.quota_type == system_configuration.current_quota_type + ), + None, + ) + + +def reserve_llm_quota_for_model(*, tenant_id: str, provider: str, model: str) -> LLMQuotaReservation: + """Reserve system-hosted LLM quota before invoking the provider.""" + provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider) + reservation = LLMQuotaReservation( + tenant_id=tenant_id, + provider=provider, + model=model, + provider_configuration=provider_configuration, + ) + if provider_configuration.using_provider_type != ProviderType.SYSTEM: + return reservation + + provider_model = provider_configuration.get_provider_model(model_type=ModelType.LLM, model=model) + if provider_model and provider_model.status == ModelStatus.QUOTA_EXCEEDED: + raise QuotaExceededError(f"Model provider {provider} quota exceeded.") + + system_configuration = provider_configuration.system_configuration + quota_configuration = _get_current_quota_configuration(system_configuration) + if quota_configuration is None or quota_configuration.quota_limit == -1: + return reservation + + reservation.quota_unit = quota_configuration.quota_unit + quota_type = system_configuration.current_quota_type + if quota_type in {ProviderQuotaType.TRIAL, ProviderQuotaType.PAID}: + match quota_configuration.quota_unit: + case QuotaUnit.CREDITS: + amount = dify_config.get_model_credits(model) + case QuotaUnit.TIMES: + amount = 1 + case QuotaUnit.TOKENS: + # Token usage is unknown before invocation. Enabling TOKENS for a hosted + # credit pool requires accurate terminal usage and an upper-bound reservation strategy. + raise ValueError("Token-based hosted credit pools do not support pre-invocation reservation.") + case _: + raise ValueError(f"Unsupported hosted credit pool quota unit: {quota_configuration.quota_unit}") + + reservation.credit_pool_reservation = CreditPoolService.reserve_credits( + tenant_id=tenant_id, + credits_required=amount, + pool_type="paid" if quota_type == ProviderQuotaType.PAID else "trial", + request_id=str(uuid4()), + session_factory=db.session, + meta={"source": "llm.invoke", "provider": provider, "model": model}, + ) + elif quota_type == ProviderQuotaType.FREE: + reservation.requires_usage = True + + return reservation + + def ensure_llm_quota_available_for_model(*, tenant_id: str, provider: str, model: str) -> None: """Raise when a tenant-bound LLM model is already out of quota.""" provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider) diff --git a/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py b/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py index 7b38d973943..3c519b3e108 100644 --- a/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py +++ b/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py @@ -373,7 +373,7 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat delta_text = "" # EasyUI streams text only; structured multimodal chunks contribute their text parts. - for content in delta_content: + for content in cast(list[object], delta_content): logger.debug("The content type %s in LLM chunk delta message content.: %r", type(content), content) match content: case TextPromptMessageContent(): diff --git a/api/core/app/task_pipeline/message_cycle_manager.py b/api/core/app/task_pipeline/message_cycle_manager.py index 3bffe83ee28..91b67a03e47 100644 --- a/api/core/app/task_pipeline/message_cycle_manager.py +++ b/api/core/app/task_pipeline/message_cycle_manager.py @@ -66,11 +66,7 @@ class MessageCycleManager: # Use SQLAlchemy 2.x style session.scalar(select(...)) with session_factory.create_session() as session: message_file = session.scalar( - select(MessageFile) - .where( - MessageFile.message_id == message_id, - ) - .where(MessageFile.belongs_to == "assistant") + select(MessageFile).where(MessageFile.message_id == message_id, MessageFile.belongs_to == "assistant") ) if message_file: diff --git a/api/core/app/workflow/file_runtime.py b/api/core/app/workflow/file_runtime.py index dec92c03b61..7bdabd93717 100644 --- a/api/core/app/workflow/file_runtime.py +++ b/api/core/app/workflow/file_runtime.py @@ -13,7 +13,7 @@ from configs import dify_config from core.app.file_access import DatabaseFileAccessController, FileAccessControllerProtocol from core.db.session_factory import session_factory from core.file import remote_fetcher -from core.tools.signature import sign_tool_file +from core.tools.signature import bind_file_uri, sign_tool_file_uri from core.workflow.file_reference import parse_file_reference from extensions.ext_storage import storage from graphon.file import FileTransferMethod @@ -62,32 +62,42 @@ class DifyWorkflowFileRuntime(WorkflowFileRuntimeProtocol): @override def resolve_file_url(self, *, file: File, for_external: bool = True) -> str | None: + uri = self.resolve_file_uri(file=file) + if uri is None or file.transfer_method == FileTransferMethod.REMOTE_URL: + return uri + return bind_file_uri(uri, self._base_url(for_external=for_external)) + + def resolve_file_uri(self, *, file: File) -> str | None: + """Resolve a signed file URI without binding Dify-owned files to an origin. + + Remote URLs retain their external absolute URL. Dify-owned files return + a signed ``/files/...`` URI that callers can bind to their own network + audience without exposing ``FILES_URL`` or ``INTERNAL_FILES_URL``. + """ + if file.transfer_method == FileTransferMethod.REMOTE_URL: return file.remote_url parsed_reference = parse_file_reference(file.reference) if parsed_reference is None: raise ValueError("Missing file reference") if file.transfer_method == FileTransferMethod.LOCAL_FILE: - return self.resolve_upload_file_url( + return self.resolve_upload_file_uri( upload_file_id=parsed_reference.record_id, - for_external=for_external, ) if file.transfer_method == FileTransferMethod.DATASOURCE_FILE: if file.extension is None: raise ValueError("Missing file extension") self._assert_upload_file_access(upload_file_id=parsed_reference.record_id) - return sign_tool_file( + return sign_tool_file_uri( tool_file_id=parsed_reference.record_id, extension=file.extension, - for_external=for_external, ) if file.transfer_method == FileTransferMethod.TOOL_FILE: if file.extension is None: raise ValueError("Missing file extension") - return self.resolve_tool_file_url( + return self.resolve_tool_file_uri( tool_file_id=parsed_reference.record_id, extension=file.extension, - for_external=for_external, ) return None @@ -99,18 +109,34 @@ class DifyWorkflowFileRuntime(WorkflowFileRuntimeProtocol): as_attachment: bool = False, for_external: bool = True, ) -> str: + uri = self.resolve_upload_file_uri(upload_file_id=upload_file_id, as_attachment=as_attachment) + return bind_file_uri(uri, self._base_url(for_external=for_external)) + + def resolve_upload_file_uri( + self, + *, + upload_file_id: str, + as_attachment: bool = False, + ) -> str: + """Resolve a signed UploadFile URI without selecting an origin.""" + self._assert_upload_file_access(upload_file_id=upload_file_id) - base_url = self._base_url(for_external=for_external) - url = f"{base_url}/files/{upload_file_id}/file-preview" + uri = f"/files/{upload_file_id}/file-preview" query = self._sign_query(payload=f"file-preview|{upload_file_id}") if as_attachment: query["as_attachment"] = "true" - return f"{url}?{urllib.parse.urlencode(query)}" + return f"{uri}?{urllib.parse.urlencode(query)}" @override def resolve_tool_file_url(self, *, tool_file_id: str, extension: str, for_external: bool = True) -> str: + uri = self.resolve_tool_file_uri(tool_file_id=tool_file_id, extension=extension) + return bind_file_uri(uri, self._base_url(for_external=for_external)) + + def resolve_tool_file_uri(self, *, tool_file_id: str, extension: str) -> str: + """Resolve a signed ToolFile URI without selecting an origin.""" + self._assert_tool_file_access(tool_file_id=tool_file_id) - return sign_tool_file(tool_file_id=tool_file_id, extension=extension, for_external=for_external) + return sign_tool_file_uri(tool_file_id=tool_file_id, extension=extension) @override def verify_preview_signature( diff --git a/api/core/app/workflow/layers/__init__.py b/api/core/app/workflow/layers/__init__.py index 7d5841275db..945f75303c7 100644 --- a/api/core/app/workflow/layers/__init__.py +++ b/api/core/app/workflow/layers/__init__.py @@ -1,11 +1,9 @@ """Workflow-level GraphEngine layers that depend on outer infrastructure.""" -from .llm_quota import LLMQuotaLayer from .observability import ObservabilityLayer from .persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer __all__ = [ - "LLMQuotaLayer", "ObservabilityLayer", "PersistenceWorkflowInfo", "WorkflowPersistenceLayer", diff --git a/api/core/app/workflow/layers/llm_quota.py b/api/core/app/workflow/layers/llm_quota.py deleted file mode 100644 index 2422eed5a70..00000000000 --- a/api/core/app/workflow/layers/llm_quota.py +++ /dev/null @@ -1,194 +0,0 @@ -""" -LLM quota deduction layer for GraphEngine. - -This layer centralizes model-quota handling outside node implementations. - -Graphon LLM-backed nodes expose provider/model identity through public node -configuration and, after execution, through ``node_run_result.inputs``. Resolve -quota billing from that public identity instead of depending on -``ModelInstance`` reconstruction inside the workflow layer. Missing identity on -quota-tracked nodes is treated as a workflow bug and aborts execution so quota -handling is never silently skipped. -""" - -import logging -from typing import final, override - -from core.app.llm import deduct_llm_quota_for_model, ensure_llm_quota_available_for_model -from core.errors.error import QuotaExceededError -from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus -from graphon.graph_engine.entities.commands import AbortCommand, CommandType -from graphon.graph_engine.layers import GraphEngineLayer -from graphon.graph_events import GraphEngineEvent, GraphNodeEventBase, NodeRunSucceededEvent -from graphon.node_events import NodeRunResult -from graphon.nodes.base.node import Node - -logger = logging.getLogger(__name__) -_QUOTA_NODE_TYPES = frozenset( - [ - BuiltinNodeTypes.LLM, - BuiltinNodeTypes.PARAMETER_EXTRACTOR, - BuiltinNodeTypes.QUESTION_CLASSIFIER, - ] -) - - -@final -class LLMQuotaLayer(GraphEngineLayer): - """Graph layer that applies tenant-scoped quota checks to LLM-backed nodes.""" - - tenant_id: str - _abort_sent: bool - - def __init__(self, tenant_id: str) -> None: - super().__init__() - self.tenant_id = tenant_id - self._abort_sent = False - - @override - def on_graph_start(self) -> None: - self._abort_sent = False - - @override - def on_event(self, event: GraphEngineEvent) -> None: - _ = event - - @override - def on_graph_end(self, error: Exception | None) -> None: - _ = error - - @override - def on_node_run_start(self, node: Node) -> None: - if self._abort_sent: - return - - if not self._supports_quota(node): - return - - model_identity = self._extract_model_identity_from_node(node) - if model_identity is None: - reason = "LLM quota check requires public node model identity before execution." - self._abort_before_node_run(node=node, reason=reason, error_type="LLMQuotaIdentityError") - logger.error("LLM quota handling aborted, node_id=%s, reason=%s", node.id, reason) - return - - provider, model_name = model_identity - try: - ensure_llm_quota_available_for_model( - tenant_id=self.tenant_id, - provider=provider, - model=model_name, - ) - except QuotaExceededError as exc: - self._abort_before_node_run(node=node, reason=str(exc), error_type=QuotaExceededError.__name__) - logger.warning("LLM quota check failed, node_id=%s, error=%s", node.id, exc) - - @override - def on_node_run_end( - self, node: Node, error: Exception | None, result_event: GraphNodeEventBase | None = None - ) -> None: - if error is not None or not isinstance(result_event, NodeRunSucceededEvent) or not self._supports_quota(node): - return - - model_identity = self._extract_model_identity_from_result_event(result_event) - if model_identity is None: - self._abort_for_missing_model_identity( - node=node, - reason="LLM quota deduction requires model identity in the node result event.", - ) - return - - provider, model_name = model_identity - - try: - deduct_llm_quota_for_model( - tenant_id=self.tenant_id, - provider=provider, - model=model_name, - usage=result_event.node_run_result.llm_usage, - ) - except QuotaExceededError as exc: - self._set_stop_event(node) - self._send_abort_command(reason=str(exc)) - logger.warning("LLM quota deduction exceeded, node_id=%s, error=%s", node.id, exc) - except Exception: - logger.exception("LLM quota deduction failed, node_id=%s", node.id) - - @staticmethod - def _set_stop_event(node: Node) -> None: - stop_event = getattr(node.graph_runtime_state, "stop_event", None) - if stop_event is not None: - stop_event.set() - - def _abort_before_node_run(self, *, node: Node, reason: str, error_type: str) -> None: - self._set_stop_event(node) - node.node_data.error_strategy = None - node.node_data.retry_config.retry_enabled = False - - def quota_aborted_run() -> NodeRunResult: - return NodeRunResult( - status=WorkflowNodeExecutionStatus.FAILED, - error=reason, - error_type=error_type, - ) - - # TODO: Push Graphon to expose a public pre-run failure/skip hook, then replace this private _run override. - node._run = quota_aborted_run # type: ignore[method-assign] - self._send_abort_command(reason=reason) - - def _abort_for_missing_model_identity(self, *, node: Node, reason: str) -> None: - self._set_stop_event(node) - self._send_abort_command(reason=reason) - logger.error("LLM quota handling aborted, node_id=%s, reason=%s", node.id, reason) - - def _send_abort_command(self, *, reason: str) -> None: - if not self.command_channel or self._abort_sent: - return - - try: - self.command_channel.send_command( - AbortCommand( - command_type=CommandType.ABORT, - reason=reason, - ) - ) - self._abort_sent = True - except Exception: - logger.exception("Failed to send quota abort command") - - @staticmethod - def _supports_quota(node: Node) -> bool: - return node.node_type in _QUOTA_NODE_TYPES - - @staticmethod - def _extract_model_identity_from_result_event(result_event: NodeRunSucceededEvent) -> tuple[str, str] | None: - provider = result_event.node_run_result.inputs.get("model_provider") - model_name = result_event.node_run_result.inputs.get("model_name") - if isinstance(provider, str) and provider and isinstance(model_name, str) and model_name: - return provider, model_name - return None - - @staticmethod - def _extract_model_identity_from_node(node: Node) -> tuple[str, str] | None: - node_data = getattr(node, "node_data", None) - if node_data is None: - node_data = getattr(node, "data", None) - - model_config = getattr(node_data, "model", None) - if model_config is None: - logger.warning( - "LLMQuotaLayer skipped quota handling because node model config is missing, node_id=%s", - node.id, - ) - return None - - provider = getattr(model_config, "provider", None) - model_name = getattr(model_config, "name", None) - if isinstance(provider, str) and provider and isinstance(model_name, str) and model_name: - return provider, model_name - - logger.warning( - "LLMQuotaLayer skipped quota handling because node model identity is invalid, node_id=%s", - node.id, - ) - return None diff --git a/api/core/app/workflow/layers/persistence.py b/api/core/app/workflow/layers/persistence.py index 52887765c02..8b04661deb3 100644 --- a/api/core/app/workflow/layers/persistence.py +++ b/api/core/app/workflow/layers/persistence.py @@ -20,11 +20,13 @@ from core.helper.trace_id_helper import ParentTraceContext from core.ops.entities.trace_entity import TraceTaskName from core.ops.ops_trace_manager import TraceQueueManager, TraceTask from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository +from core.workflow.node_execution_process_data import preserve_workflow_agent_binding_id from core.workflow.system_variables import SystemVariableKey from core.workflow.variable_prefixes import SYSTEM_VARIABLE_NODE_ID from core.workflow.workflow_run_outputs import project_node_outputs_for_workflow_run -from graphon.entities import WorkflowExecution, WorkflowNodeExecution +from graphon.entities import WorkflowExecution, WorkflowNodeExecution, WorkflowStartReason from graphon.enums import ( + BuiltinNodeTypes, WorkflowExecutionStatus, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus, @@ -116,7 +118,7 @@ class WorkflowPersistenceLayer(GraphEngineLayer): def on_event(self, event: GraphEngineEvent) -> None: match event: case GraphRunStartedEvent(): - self._handle_graph_run_started() + self._handle_graph_run_started(event) case GraphRunSucceededEvent(): self._handle_graph_run_succeeded(event) case GraphRunPartialSucceededEvent(): @@ -147,7 +149,7 @@ class WorkflowPersistenceLayer(GraphEngineLayer): # ------------------------------------------------------------------ # Graph-level handlers # ------------------------------------------------------------------ - def _handle_graph_run_started(self) -> None: + def _handle_graph_run_started(self, event: GraphRunStartedEvent | None = None) -> None: execution_id = self._get_execution_id() workflow_execution = WorkflowExecution.new( id_=execution_id, @@ -161,6 +163,10 @@ class WorkflowPersistenceLayer(GraphEngineLayer): self._workflow_execution_repository.save(workflow_execution) self._workflow_execution = workflow_execution + if event is not None and event.reason == WorkflowStartReason.RESUMPTION: + node_executions = self._workflow_node_execution_repository.get_by_workflow_execution(execution_id) + self._node_execution_cache = {execution.id: execution for execution in node_executions} + self._node_sequence = max((execution.index for execution in node_executions), default=0) def _handle_graph_run_succeeded(self, event: GraphRunSucceededEvent) -> None: execution = self._get_workflow_execution() @@ -241,7 +247,10 @@ class WorkflowPersistenceLayer(GraphEngineLayer): ) self._node_execution_cache[event.id] = domain_execution - self._workflow_node_execution_repository.save(domain_execution) + if event.node_type == BuiltinNodeTypes.AGENT and event.node_version == "2": + self._workflow_node_execution_repository.save_synchronously(domain_execution) + else: + self._workflow_node_execution_repository.save(domain_execution) snapshot = _NodeRuntimeSnapshot( node_id=event.node_id, @@ -361,7 +370,11 @@ class WorkflowPersistenceLayer(GraphEngineLayer): def _append_retry_history(self, execution: WorkflowNodeExecution, event: NodeRunRetryEvent) -> None: """Append a validated full attempt before repository truncation or offload.""" finished_at = naive_utc_now() - process_data = dict(execution.process_data or {}) + process_data = preserve_workflow_agent_binding_id( + event.node_run_result.process_data, + execution.process_data, + ) + process_data = dict(process_data or {}) raw_history = process_data.get(RETRY_HISTORY_PROCESS_DATA_KEY) history = list(raw_history) if isinstance(raw_history, list) else [] projected_outputs = project_node_outputs_for_workflow_run( @@ -390,11 +403,12 @@ class WorkflowPersistenceLayer(GraphEngineLayer): next_process_data: Mapping[str, Any] | None, ) -> Mapping[str, Any] | None: """Keep internal retry history while replacing node-specific Process Data.""" + merged_process_data = preserve_workflow_agent_binding_id(existing_process_data, next_process_data) raw_history = (existing_process_data or {}).get(RETRY_HISTORY_PROCESS_DATA_KEY) if not isinstance(raw_history, list) or not raw_history: - return next_process_data + return merged_process_data - merged_process_data = dict(next_process_data or {}) + merged_process_data = dict(merged_process_data or {}) merged_process_data[RETRY_HISTORY_PROCESS_DATA_KEY] = raw_history return merged_process_data @@ -440,6 +454,11 @@ class WorkflowPersistenceLayer(GraphEngineLayer): outputs=projected_outputs, metadata=node_result.metadata, ) + else: + domain_execution.process_data = preserve_workflow_agent_binding_id( + node_result.process_data, + domain_execution.process_data, + ) self._workflow_node_execution_repository.save(domain_execution) self._workflow_node_execution_repository.save_execution_data(domain_execution) diff --git a/api/core/datasource/datasource_file_manager.py b/api/core/datasource/datasource_file_manager.py index 521791206db..317f362d662 100644 --- a/api/core/datasource/datasource_file_manager.py +++ b/api/core/datasource/datasource_file_manager.py @@ -12,8 +12,8 @@ from uuid import uuid4 import httpx from configs import dify_config +from core.db.session_factory import session_factory from core.file import remote_fetcher -from extensions.ext_database import db from extensions.ext_storage import storage from extensions.storage.storage_type import StorageType from models.enums import CreatorUserRole @@ -54,6 +54,7 @@ class DatasourceFileManager: mimetype: str, filename: str | None = None, ) -> UploadFile: + """Persist an uploaded datasource file and its storage payload.""" extension = guess_extension(mimetype) or ".bin" unique_name = uuid4().hex unique_filename = f"{unique_name}{extension}" @@ -82,9 +83,10 @@ class DatasourceFileManager: created_at=datetime.now(), ) - db.session.add(upload_file) - db.session.commit() - db.session.refresh(upload_file) + with session_factory.create_session() as session: + session.add(upload_file) + session.commit() + session.refresh(upload_file) return upload_file @@ -95,6 +97,7 @@ class DatasourceFileManager: file_url: str, conversation_id: str | None = None, ) -> ToolFile: + """Download a remote file and persist its tool-file metadata.""" # try to download image try: response = remote_fetcher.make_request("GET", file_url) @@ -125,8 +128,9 @@ class DatasourceFileManager: size=len(blob), ) - db.session.add(tool_file) - db.session.commit() + with session_factory.create_session() as session: + session.add(tool_file) + session.commit() return tool_file @@ -139,7 +143,8 @@ class DatasourceFileManager: :return: the binary of the file, mime type """ - upload_file: UploadFile | None = db.session.get(UploadFile, id) + with session_factory.create_session() as session: + upload_file: UploadFile | None = session.get(UploadFile, id) if not upload_file: return None @@ -157,21 +162,24 @@ class DatasourceFileManager: :return: the binary of the file, mime type """ - message_file: MessageFile | None = db.session.get(MessageFile, id) + with session_factory.create_session() as session: + message_file: MessageFile | None = session.get(MessageFile, id) - # Check if message_file is not None - if message_file is not None: - # get tool file id - if message_file.url is not None: - tool_file_id = message_file.url.split("/")[-1] - # trim extension - tool_file_id = tool_file_id.split(".")[0] + # Check if message_file is not None + if message_file is not None: + # get tool file id + if message_file.url is not None: + tool_file_id = message_file.url.split("/")[-1] + # trim extension + tool_file_id = tool_file_id.split(".")[0] + else: + tool_file_id = None else: tool_file_id = None - else: - tool_file_id = None - tool_file: ToolFile | None = db.session.get(ToolFile, tool_file_id) + if not tool_file_id: + return None + tool_file: ToolFile | None = session.get(ToolFile, tool_file_id) if not tool_file: return None @@ -185,11 +193,12 @@ class DatasourceFileManager: """ get file binary - :param tool_file_id: the id of the tool file + :param upload_file_id: the id of the upload file :return: the binary of the file, mime type """ - upload_file: UploadFile | None = db.session.get(UploadFile, upload_file_id) + with session_factory.create_session() as session: + upload_file: UploadFile | None = session.get(UploadFile, upload_file_id) if not upload_file: return None, None diff --git a/api/core/entities/parameter_entities.py b/api/core/entities/parameter_entities.py index b61c4ad4bb5..dd25e667a5e 100644 --- a/api/core/entities/parameter_entities.py +++ b/api/core/entities/parameter_entities.py @@ -16,6 +16,8 @@ class CommonParameterType(StrEnum): TOOLS_SELECTOR = "array[tools]" CHECKBOX = "checkbox" ANY = auto() + DATE = "date" + DATE_RANGE = "date-range" # Dynamic select parameter # Once you are not sure about the available options until authorization is done diff --git a/api/core/entities/provider_configuration.py b/api/core/entities/provider_configuration.py index 95b8b686e05..a73160ed1de 100644 --- a/api/core/entities/provider_configuration.py +++ b/api/core/entities/provider_configuration.py @@ -9,8 +9,9 @@ from json import JSONDecodeError from typing import Any, override from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, model_validator -from sqlalchemy import func, select +from sqlalchemy import String, func, literal, select from sqlalchemy.orm import Session +from sqlalchemy.sql.elements import BindParameter from constants import HIDDEN_VALUE from core.entities import PluginCredentialType @@ -57,6 +58,29 @@ logger = logging.getLogger(__name__) original_provider_configurate_methods: dict[str, list[ConfigurateMethod]] = {} +def _model_type_db_values(model_type: ModelType) -> tuple[str, ...]: + """Return DB values that may represent ``model_type`` after pre-1.15 upgrades. + + SQL equality against the canonical value misses unmigrated rows, so legacy + lookups need to match both the current and provider-native spellings. + """ + values = [model_type.value] + origin = model_type.to_origin_model_type() + if origin not in values: + values.append(origin) + return tuple(values) + + +def _model_type_db_literals(model_type: ModelType) -> tuple[BindParameter[str], ...]: + """Return string-typed literals for legacy-aware model type filters. + + ``EnumText`` rejects legacy spellings during normal binding so they cannot + be written back. Explicit string literals bypass that bind processor only + for compatibility lookups of rows that predate the canonical enum values. + """ + return tuple(literal(value, type_=String) for value in _model_type_db_values(model_type)) + + class ProviderConfiguration(BaseModel): """ Provider configuration entity for managing model provider settings. @@ -846,7 +870,7 @@ class ProviderConfiguration(BaseModel): ProviderModel.tenant_id == self.tenant_id, ProviderModel.provider_name.in_(provider_names), ProviderModel.model_name == model, - ProviderModel.model_type == model_type, + ProviderModel.model_type.in_(_model_type_db_literals(model_type)), ) return session.execute(stmt).scalar_one_or_none() @@ -871,7 +895,7 @@ class ProviderConfiguration(BaseModel): ProviderModelCredential.tenant_id == self.tenant_id, ProviderModelCredential.provider_name.in_(self._get_provider_names()), ProviderModelCredential.model_name == model, - ProviderModelCredential.model_type == model_type, + ProviderModelCredential.model_type.in_(_model_type_db_literals(model_type)), ) credential_record = session.execute(stmt).scalar_one_or_none() @@ -1184,7 +1208,7 @@ class ProviderConfiguration(BaseModel): ProviderModelCredential.tenant_id == self.tenant_id, ProviderModelCredential.provider_name.in_(self._get_provider_names()), ProviderModelCredential.model_name == model, - ProviderModelCredential.model_type == model_type, + ProviderModelCredential.model_type.in_(_model_type_db_literals(model_type)), ) credential_record = session.execute(stmt).scalar_one_or_none() if not credential_record: @@ -1228,7 +1252,7 @@ class ProviderConfiguration(BaseModel): ProviderModelCredential.tenant_id == self.tenant_id, ProviderModelCredential.provider_name.in_(self._get_provider_names()), ProviderModelCredential.model_name == model, - ProviderModelCredential.model_type == model_type, + ProviderModelCredential.model_type.in_(_model_type_db_literals(model_type)), ) available_credentials_count = session.execute(count_stmt).scalar() or 0 session.delete(credential_record) @@ -1394,7 +1418,7 @@ class ProviderConfiguration(BaseModel): stmt = select(ProviderModelSetting).where( ProviderModelSetting.tenant_id == self.tenant_id, ProviderModelSetting.provider_name.in_(self._get_provider_names()), - ProviderModelSetting.model_type == model_type, + ProviderModelSetting.model_type.in_(_model_type_db_literals(model_type)), ProviderModelSetting.model_name == model, ) return session.execute(stmt).scalars().first() @@ -1833,10 +1857,10 @@ class ProviderConfiguration(BaseModel): ) ) - # if llm name not in restricted llm list, remove it + # Hosted allowlists currently use exact model names across model types. restrict_model_names = [rm.model for rm in restrict_models] for provider_model in provider_models: - if provider_model.model_type == ModelType.LLM and provider_model.model not in restrict_model_names: + if provider_model.model not in restrict_model_names: provider_model.status = ModelStatus.NO_PERMISSION elif not quota_configuration.is_valid: provider_model.status = ModelStatus.QUOTA_EXCEEDED diff --git a/api/core/helper/code_executor/jinja2/jinja2_transformer.py b/api/core/helper/code_executor/jinja2/jinja2_transformer.py index d1c75c981b6..1277c9b331f 100644 --- a/api/core/helper/code_executor/jinja2/jinja2_transformer.py +++ b/api/core/helper/code_executor/jinja2/jinja2_transformer.py @@ -39,15 +39,16 @@ class Jinja2TemplateTransformer(TemplateTransformer): @override def get_runner_script(cls) -> str: runner_script = dedent(f""" - import jinja2 import json from base64 import b64decode + from jinja2.sandbox import SandboxedEnvironment # declare main function def main(**inputs): # Decode base64-encoded template to handle special characters safely template_code = b64decode('{cls._template_b64_placeholder}').decode('utf-8') - template = jinja2.Template(template_code) + env = SandboxedEnvironment() + template = env.from_string(template_code) return template.render(**inputs) # decode and prepare input dict @@ -67,12 +68,13 @@ class Jinja2TemplateTransformer(TemplateTransformer): @override def get_preload_script(cls) -> str: preload_script = dedent(""" - import jinja2 + from jinja2.sandbox import SandboxedEnvironment from base64 import b64decode def _jinja2_preload_(): - # prepare jinja2 environment, load template and render before to avoid sandbox issue - template = jinja2.Template('{{s}}') + # prepare jinja2 sandboxed environment, load template and render + env = SandboxedEnvironment() + template = env.from_string('{{s}}') template.render(s='a') if __name__ == '__main__': diff --git a/api/core/helper/encrypter.py b/api/core/helper/encrypter.py index f72f6c6be9f..acb949e3982 100644 --- a/api/core/helper/encrypter.py +++ b/api/core/helper/encrypter.py @@ -1,8 +1,5 @@ import base64 - -from Crypto.PublicKey import RSA - -from libs import rsa +from typing import Any def obfuscated_token(token: str) -> str: @@ -18,29 +15,35 @@ def full_mask_token(token_length: int = 20) -> str: def encrypt_token(tenant_id: str, token: str) -> str: - from models.account import Tenant - from models.engine import db + from extensions.ext_key_provider import key_provider_manager - if not (tenant := db.session.get(Tenant, tenant_id)): - raise ValueError(f"Tenant with id {tenant_id} not found") - assert tenant.encrypt_public_key is not None - encrypted_token = rsa.encrypt(token, tenant.encrypt_public_key) + encrypted_token = key_provider_manager.provider.encrypt(tenant_id, token) return base64.b64encode(encrypted_token).decode() def decrypt_token(tenant_id: str, token: str) -> str: - return rsa.decrypt(base64.b64decode(token), tenant_id) + from extensions.ext_key_provider import key_provider_manager + + return key_provider_manager.provider.decrypt(tenant_id, base64.b64decode(token)) def batch_decrypt_token(tenant_id: str, tokens: list[str]) -> list[str]: - rsa_key, cipher_rsa = rsa.get_decrypt_decoding(tenant_id) - - return [rsa.decrypt_token_with_decoding(base64.b64decode(token), rsa_key, cipher_rsa) for token in tokens] + decoding = get_decrypt_decoding(tenant_id) + return [decrypt_token_with_decoding(token, decoding) for token in tokens] -def get_decrypt_decoding(tenant_id: str) -> tuple[RSA.RsaKey, object]: - return rsa.get_decrypt_decoding(tenant_id) +def get_decrypt_decoding(tenant_id: str) -> Any: + """ + Return a reusable decoding context for batch/repeated decryption of a tenant's credentials + (e.g. across many provider/model configs in the same request). The returned object is opaque + and must only be passed back into decrypt_token_with_decoding. + """ + from extensions.ext_key_provider import key_provider_manager + + return key_provider_manager.provider.get_decrypt_decoding(tenant_id) -def decrypt_token_with_decoding(token: str, rsa_key: RSA.RsaKey, cipher_rsa: object) -> str: - return rsa.decrypt_token_with_decoding(base64.b64decode(token), rsa_key, cipher_rsa) +def decrypt_token_with_decoding(token: str, decoding: Any) -> str: + from extensions.ext_key_provider import key_provider_manager + + return key_provider_manager.provider.decrypt_with_decoding(base64.b64decode(token), decoding) diff --git a/api/core/helper/marketplace.py b/api/core/helper/marketplace.py index e6e1d565769..e14c7185858 100644 --- a/api/core/helper/marketplace.py +++ b/api/core/helper/marketplace.py @@ -19,7 +19,7 @@ MARKETPLACE_TIMEOUT = 30 def get_plugin_pkg_url(plugin_unique_identifier: str) -> str: query = urlencode({"unique_identifier": plugin_unique_identifier}) - return f"{marketplace_api_url / 'api/v1/plugins/download'}?{query}" + return f"{marketplace_api_url / 'api/v1/plugins/download-url'}?{query}" def download_plugin_pkg(plugin_unique_identifier: str) -> bytes: diff --git a/api/core/helper/ssrf_proxy.py b/api/core/helper/ssrf_proxy.py index 9f0bc17f0f2..95b07c36903 100644 --- a/api/core/helper/ssrf_proxy.py +++ b/api/core/helper/ssrf_proxy.py @@ -246,10 +246,21 @@ def make_request( # Squid typically identifies itself in Server or Via headers if "squid" in server_header or "squid" in via_header: + # The deny ACL is usually ``to_private_networks`` (RFC1918 + + # loopback / link-local / CGN / IPv6 ULA, etc.). We don't know + # which specific ACL tripped from Squid's response alone, but + # the actionable remediation is the same in every case: + # allowlist the destination in the SSRF proxy. Tell the user + # exactly which env var to set so they don't have to grep the + # squid config. Mention a concrete example CIDR (e.g. the + # 172.21.0.0/16 from the bug report) so they can copy-paste it. response.close() raise ToolSSRFError( - f"Access to '{url}' was blocked by SSRF protection. " - f"The URL may point to a private or local network address. " + f"Access to '{url}' was blocked by SSRF protection " + f"(e.g. SSRF_PROXY_ALLOW_PRIVATE_IPS=172.21.0.0/16 to " + f"allow 172.21.0.0/16). The URL resolves to a private, " + f"loopback, link-local, or otherwise non-public network " + f"address. See https://github.com/infiniflow/ragflow/issues/38443." ) if response.status_code not in STATUS_FORCELIST or max_retries == 0: diff --git a/api/core/hosting_configuration.py b/api/core/hosting_configuration.py index 09473b8b78f..90c11376ec1 100644 --- a/api/core/hosting_configuration.py +++ b/api/core/hosting_configuration.py @@ -6,7 +6,7 @@ from pydantic import BaseModel from configs import dify_config from core.entities import DEFAULT_PLUGIN_ID from core.entities.provider_entities import ProviderQuotaType, QuotaUnit, RestrictModel -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from graphon.model_runtime.entities.model_entities import ModelType diff --git a/api/core/indexing_runner.py b/api/core/indexing_runner.py index 92246b6614c..7526bc533cc 100644 --- a/api/core/indexing_runner.py +++ b/api/core/indexing_runner.py @@ -34,6 +34,7 @@ from core.rag.splitter.fixed_text_splitter import ( ) from core.rag.splitter.text_splitter import TextSplitter from core.tools.utils.web_reader_tool import get_image_upload_file_ids +from enums import DeploymentEdition from extensions.ext_redis import redis_client from extensions.ext_storage import storage from graphon.model_runtime.entities.model_entities import ModelType @@ -44,13 +45,19 @@ from models.dataset import AutomaticRulesConfig, ChildChunk, Dataset, DatasetPro from models.dataset import Document as DatasetDocument from models.enums import DataSourceType, IndexingStatus, ProcessRuleMode, SegmentStatus from models.model import UploadFile +from services.vector_space_admission_service import VectorSpaceAdmissionService logger = logging.getLogger(__name__) class IndexingRunner: - def __init__(self): + def __init__( + self, + *, + enforce_vector_space_admission: bool = False, + ): self.storage = storage + self.enforce_vector_space_admission = enforce_vector_space_admission @staticmethod def _get_model_manager(tenant_id: str) -> ModelManager: @@ -73,6 +80,7 @@ class IndexingRunner: The phase commits keep document locks short and make newly created segments visible to the worker sessions used for keyword and vector indexing. """ + vector_space_admission = VectorSpaceAdmissionService() for dataset_document in dataset_documents: document_id = dataset_document.id try: @@ -114,6 +122,15 @@ class IndexingRunner: current_user=current_user, session=session, ) + if self.enforce_vector_space_admission: + vector_space_admission.ensure_document_can_be_indexed( + dataset=dataset, + document_id=requeried_document.id, + doc_form=requeried_document.doc_form, + documents=documents, + include_summaries=bool(requeried_document.need_summary), + session=session, + ) token_counts = calculate_segment_token_counts(dataset=dataset, documents=documents) total_tokens = sum(token_counts) # save segment @@ -314,7 +331,7 @@ class IndexingRunner: Estimate the indexing for the document. """ # check document limit - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: count = len(extract_settings) batch_upload_limit = dify_config.BATCH_UPLOAD_LIMIT if count > batch_upload_limit: @@ -519,7 +536,7 @@ class IndexingRunner: def filter_string(text): text = re.sub(r"<\|", "<", text) text = re.sub(r"\|>", ">", text) - text = re.sub(r"[\x00-\x08\x0B\x0C\x0E-\x1F\x7F\xEF\xBF\xBE]", "", text) + text = re.sub(r"[\x00-\x08\x0B\x0C\x0E-\x1F\x7F]", "", text) # Unicode U+FFFE text = re.sub("\ufffe", "", text) return text diff --git a/api/core/llm_generator/prompts.py b/api/core/llm_generator/prompts.py index 3c6f8c468a0..97cd8a812a7 100644 --- a/api/core/llm_generator/prompts.py +++ b/api/core/llm_generator/prompts.py @@ -254,7 +254,7 @@ Your task is to convert simple user descriptions into properly formatted JSON Sc } ### Example 4: -**User Input:** I need album schema, the ablum has songs, and each song has name, duration, and artist. +**User Input:** I need album schema, the album has songs, and each song has name, duration, and artist. **JSON Schema Output:** { "type": "object", @@ -273,7 +273,7 @@ Your task is to convert simple user descriptions into properly formatted JSON Sc "duration": { "type": "string" }, - "aritst": { + "artist": { "type": "string" } }, @@ -281,7 +281,7 @@ Your task is to convert simple user descriptions into properly formatted JSON Sc "name", "id", "duration", - "aritst" + "artist" ] } } diff --git a/api/core/mcp/auth_client_comparison.md b/api/core/mcp/auth_client_comparison.md deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/api/core/mcp/session/base_session.py b/api/core/mcp/session/base_session.py index 00f6ac8fa3e..26c1fc1700a 100644 --- a/api/core/mcp/session/base_session.py +++ b/api/core/mcp/session/base_session.py @@ -135,7 +135,7 @@ class BaseSession[ messages when entered. """ - _response_streams: dict[RequestId, queue.Queue[JSONRPCResponse | JSONRPCError | HTTPStatusError]] + _response_streams: dict[RequestId, queue.Queue[JSONRPCResponse | JSONRPCError | HTTPStatusError | None]] _request_id: int _in_flight: dict[RequestId, RequestResponder[ReceiveRequestT, SendResultT]] _receive_request_type: type[ReceiveRequestT] @@ -216,7 +216,7 @@ class BaseSession[ request_id = self._request_id self._request_id = request_id + 1 - response_queue: queue.Queue[JSONRPCResponse | JSONRPCError | HTTPStatusError] = queue.Queue() + response_queue: queue.Queue[JSONRPCResponse | JSONRPCError | HTTPStatusError | None] = queue.Queue() self._response_streams[request_id] = response_queue try: @@ -229,9 +229,9 @@ class BaseSession[ self._write_stream.put(SessionMessage(message=JSONRPCMessage(jsonrpc_request), metadata=metadata)) timeout = DEFAULT_RESPONSE_READ_TIMEOUT if request_read_timeout_seconds is not None: - timeout = float(request_read_timeout_seconds.total_seconds()) + timeout = request_read_timeout_seconds.total_seconds() elif self._session_read_timeout_seconds is not None: - timeout = float(self._session_read_timeout_seconds.total_seconds()) + timeout = self._session_read_timeout_seconds.total_seconds() while True: try: response_or_error = response_queue.get(timeout=timeout) diff --git a/api/core/model_manager.py b/api/core/model_manager.py index 29113ac6b2c..c07cc74583d 100644 --- a/api/core/model_manager.py +++ b/api/core/model_manager.py @@ -1,19 +1,19 @@ import logging from collections.abc import Callable, Generator, Iterable, Mapping, Sequence from copy import deepcopy -from typing import IO, Any, Literal, Optional, ParamSpec, TypeVar, Union, cast, overload +from typing import IO, Any, Literal, Optional, ParamSpec, TypeVar, Union, cast, overload, override from configs import dify_config from core.entities import PluginCredentialType from core.entities.embedding_type import EmbeddingInputType from core.entities.provider_configuration import ProviderConfiguration, ProviderModelBundle from core.entities.provider_entities import ModelLoadBalancingConfiguration -from core.errors.error import ProviderTokenNotInitError +from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError from core.plugin.impl.model_runtime_factory import create_plugin_provider_manager from core.provider_manager import ProviderManager from extensions.ext_redis import redis_client from graphon.model_runtime.callbacks.base_callback import Callback -from graphon.model_runtime.entities.llm_entities import LLMResult +from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage from graphon.model_runtime.entities.message_entities import PromptMessage, PromptMessageTool from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelFeature, ModelType from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult @@ -442,6 +442,149 @@ class ModelInstance: ) +class QuotaManagedModelInstance(ModelInstance): + """A system-hosted LLM instance that owns quota settlement per invocation.""" + + def reserve_quota(self): + from core.app.llm.quota import reserve_llm_quota_for_model + + return reserve_llm_quota_for_model( + tenant_id=self.provider_model_bundle.configuration.tenant_id, + provider=self.provider, + model=self.model_name, + ) + + @staticmethod + def release_quota_safely(reservation) -> None: + try: + reservation.release() + except Exception: + logger.exception("Failed to release LLM quota reservation") + + @overload + def invoke_llm( + self, + prompt_messages: Sequence[PromptMessage], + model_parameters: dict[str, Any] | None = None, + tools: Sequence[PromptMessageTool] | None = None, + stop: list[str] | None = None, + stream: Literal[True] = True, + callbacks: list[Callback] | None = None, + request_metadata: Mapping[str, object] | None = None, + ) -> Generator: ... + + @overload + def invoke_llm( + self, + prompt_messages: list[PromptMessage], + model_parameters: dict[str, Any] | None = None, + tools: Sequence[PromptMessageTool] | None = None, + stop: list[str] | None = None, + stream: Literal[False] = False, + callbacks: list[Callback] | None = None, + request_metadata: Mapping[str, object] | None = None, + ) -> LLMResult: ... + + @overload + def invoke_llm( + self, + prompt_messages: list[PromptMessage], + model_parameters: dict[str, Any] | None = None, + tools: Sequence[PromptMessageTool] | None = None, + stop: list[str] | None = None, + stream: bool = True, + callbacks: list[Callback] | None = None, + request_metadata: Mapping[str, object] | None = None, + ) -> Union[LLMResult, Generator]: ... + + @override + def invoke_llm( + self, + prompt_messages: Sequence[PromptMessage], + model_parameters: dict[str, Any] | None = None, + tools: Sequence[PromptMessageTool] | None = None, + stop: Sequence[str] | None = None, + stream: bool = True, + callbacks: list[Callback] | None = None, + request_metadata: Mapping[str, object] | None = None, + ) -> Union[LLMResult, Generator]: + normalized_prompt_messages = list(prompt_messages) + normalized_stop = list(stop) if stop else None + if stream: + return self._invoke_llm_stream( + prompt_messages=normalized_prompt_messages, + model_parameters=model_parameters, + tools=tools, + stop=normalized_stop, + callbacks=callbacks, + request_metadata=request_metadata, + ) + + reservation = self.reserve_quota() + try: + response = super().invoke_llm( + prompt_messages=normalized_prompt_messages, + model_parameters=model_parameters, + tools=tools, + stop=normalized_stop, + stream=False, + callbacks=callbacks, + request_metadata=request_metadata, + ) + if isinstance(response, Generator): + raise TypeError("Non-streaming LLM invocation returned a generator.") + reservation.commit(response.usage) + return response + finally: + self.release_quota_safely(reservation) + + def _invoke_llm_stream( + self, + *, + prompt_messages: list[PromptMessage], + model_parameters: dict[str, Any] | None, + tools: Sequence[PromptMessageTool] | None, + stop: list[str] | None, + callbacks: list[Callback] | None, + request_metadata: Mapping[str, object] | None, + ) -> Generator: + reservation = self.reserve_quota() + usage: LLMUsage | None = None + try: + response = super().invoke_llm( + prompt_messages=prompt_messages, + model_parameters=model_parameters, + tools=tools, + stop=stop, + stream=True, + callbacks=callbacks, + request_metadata=request_metadata, + ) + if not isinstance(response, Generator): + raise TypeError("Streaming LLM invocation did not return a generator.") + + if reservation.commit_before_delivery: + for chunk in response: + chunk_usage = chunk.delta.usage + if chunk_usage is not None: + usage = chunk_usage + reservation.commit(usage) + yield chunk + return + + buffered_chunks = [] + for chunk in response: + chunk_usage = chunk.delta.usage + if chunk_usage is not None: + usage = chunk_usage + buffered_chunks.append(chunk) + + reservation.commit(usage) + yield from buffered_chunks + finally: + self.release_quota_safely(reservation) + + class ModelManager: """Resolves :class:`ModelInstance` objects for a tenant and provider. @@ -472,6 +615,43 @@ class ModelManager: def for_tenant(cls, tenant_id: str, user_id: str | None = None) -> "ModelManager": return cls(provider_manager=create_plugin_provider_manager(tenant_id=tenant_id, user_id=user_id)) + @staticmethod + def _validate_system_model_access( + provider_model_bundle: ProviderModelBundle, + *, + model_type: ModelType, + model: str, + ) -> None: + configuration = provider_model_bundle.configuration + if configuration.using_provider_type != ProviderType.SYSTEM: + return + + # Hosted allowlists retain the existing comma-separated format. Model names + # are matched exactly; model-type-specific entries will be introduced later. + quota_configuration = next( + ( + quota + for quota in configuration.system_configuration.quota_configurations + if quota.quota_type == configuration.system_configuration.current_quota_type + ), + None, + ) + if quota_configuration is None or not quota_configuration.restrict_models: + return + if any(restricted_model.model == model for restricted_model in quota_configuration.restrict_models): + return + + raise ModelCurrentlyNotSupportError(f"System model {model_type.value}/{model} is not allowed.") + + @staticmethod + def _model_instance_class(provider_model_bundle: ProviderModelBundle, model_type: ModelType) -> type[ModelInstance]: + if ( + model_type == ModelType.LLM + and provider_model_bundle.configuration.using_provider_type == ProviderType.SYSTEM + ): + return QuotaManagedModelInstance + return ModelInstance + def get_model_instance( self, tenant_id: str, @@ -493,17 +673,19 @@ class ModelManager: provider_model_bundle = self._provider_manager.get_provider_model_bundle( tenant_id=tenant_id, provider=provider, model_type=model_type ) + self._validate_system_model_access(provider_model_bundle, model_type=model_type, model=model) + model_instance_class = self._model_instance_class(provider_model_bundle, model_type) cred_cache_key = (tenant_id, provider, model_type.value, model) if cred_cache_key in self._credentials_cache: - return ModelInstance( + return model_instance_class( provider_model_bundle, model, deepcopy(self._credentials_cache[cred_cache_key]), ) - ret = ModelInstance(provider_model_bundle, model) + ret = model_instance_class(provider_model_bundle, model) if self._enable_credentials_cache: self._credentials_cache[cred_cache_key] = deepcopy(ret.credentials) return ret diff --git a/api/core/moderation/api/api.py b/api/core/moderation/api/api.py index ec9a1906f8a..361edf5c574 100644 --- a/api/core/moderation/api/api.py +++ b/api/core/moderation/api/api.py @@ -2,6 +2,7 @@ from typing import Any, override from pydantic import BaseModel, Field from sqlalchemy import select +from sqlalchemy.orm import scoped_session from core.extension.api_based_extension_requestor import APIBasedExtensionPoint, APIBasedExtensionRequestor from core.helper.encrypter import decrypt_token @@ -40,7 +41,7 @@ class ApiModeration(Moderation): if not api_based_extension_id: raise ValueError("api_based_extension_id is required") - extension = cls._get_api_based_extension(tenant_id, api_based_extension_id) + extension = cls._get_api_based_extension(tenant_id, api_based_extension_id, db.session) if not extension: raise ValueError("API-based Extension not found. Please check it again.") @@ -81,7 +82,9 @@ class ApiModeration(Moderation): def _get_config_by_requestor(self, extension_point: APIBasedExtensionPoint, params: dict[str, Any]): if self.config is None: raise ValueError("The config is not set.") - extension = self._get_api_based_extension(self.tenant_id, self.config.get("api_based_extension_id", "")) + extension = self._get_api_based_extension( + self.tenant_id, self.config.get("api_based_extension_id", ""), db.session + ) if not extension: raise ValueError("API-based Extension not found. Please check it again.") requestor = APIBasedExtensionRequestor(extension.api_endpoint, decrypt_token(self.tenant_id, extension.api_key)) @@ -90,10 +93,12 @@ class ApiModeration(Moderation): return result @staticmethod - def _get_api_based_extension(tenant_id: str, api_based_extension_id: str) -> APIBasedExtension | None: + def _get_api_based_extension( + tenant_id: str, api_based_extension_id: str, session: scoped_session + ) -> APIBasedExtension | None: stmt = select(APIBasedExtension).where( APIBasedExtension.tenant_id == tenant_id, APIBasedExtension.id == api_based_extension_id ) - extension = db.session.scalar(stmt) + extension = session.scalar(stmt) return extension diff --git a/api/core/ops/utils.py b/api/core/ops/utils.py index 50cccd9d088..5e99ed89a9d 100644 --- a/api/core/ops/utils.py +++ b/api/core/ops/utils.py @@ -34,6 +34,7 @@ def measure_time(): try: yield timing_info finally: + # pyrefly: ignore [bad-assignment] timing_info["end"] = datetime.now() diff --git a/api/core/plugin/backwards_invocation/model.py b/api/core/plugin/backwards_invocation/model.py index c03665272b9..df2dd7c795f 100644 --- a/api/core/plugin/backwards_invocation/model.py +++ b/api/core/plugin/backwards_invocation/model.py @@ -3,7 +3,6 @@ from binascii import hexlify, unhexlify from collections.abc import Generator from typing import Any -from core.app.llm import deduct_llm_quota from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output from core.model_manager import ModelManager from core.plugin.backwards_invocation.base import BaseBackwardsInvocation @@ -80,15 +79,11 @@ class PluginModelBackwardsInvocation(BaseBackwardsInvocation): def handle() -> Generator[LLMResultChunk, None, None]: for chunk in response: - if chunk.delta.usage: - deduct_llm_quota(tenant_id=tenant.id, model_instance=model_instance, usage=chunk.delta.usage) chunk.prompt_messages = [] yield chunk return handle() else: - if response.usage: - deduct_llm_quota(tenant_id=tenant.id, model_instance=model_instance, usage=response.usage) def handle_non_streaming(response: LLMResult) -> Generator[LLMResultChunk, None, None]: yield LLMResultChunk( @@ -141,15 +136,11 @@ class PluginModelBackwardsInvocation(BaseBackwardsInvocation): def handle() -> Generator[LLMResultChunkWithStructuredOutput, None, None]: for chunk in response: - if chunk.delta.usage: - deduct_llm_quota(tenant_id=tenant.id, model_instance=model_instance, usage=chunk.delta.usage) chunk.prompt_messages = [] yield chunk return handle() else: - if response.usage: - deduct_llm_quota(tenant_id=tenant.id, model_instance=model_instance, usage=response.usage) def handle_non_streaming( response: LLMResultWithStructuredOutput, diff --git a/api/core/plugin/entities/parameters.py b/api/core/plugin/entities/parameters.py index 4c0abe4ed4a..a3e59b739ad 100644 --- a/api/core/plugin/entities/parameters.py +++ b/api/core/plugin/entities/parameters.py @@ -1,4 +1,5 @@ import json +from datetime import date from enum import StrEnum, auto from typing import Any, Union @@ -46,6 +47,8 @@ class PluginParameterType(StrEnum): # MCP object and array type parameters ARRAY = CommonParameterType.ARRAY OBJECT = CommonParameterType.OBJECT + DATE = CommonParameterType.DATE + DATE_RANGE = CommonParameterType.DATE_RANGE class MCPServerParameterType(StrEnum): @@ -92,15 +95,32 @@ class PluginParameter(BaseModel): def as_normal_type(typ: StrEnum): + if typ.value == PluginParameterType.DATE_RANGE: + return "object" if typ.value in { PluginParameterType.SECRET_INPUT, PluginParameterType.SELECT, PluginParameterType.CHECKBOX, + PluginParameterType.DATE, }: return "string" return typ.value +def _validate_date(value: Any, name: str = "date") -> str: + if not isinstance(value, str): + raise ValueError(f"The {name} parameter must be a string in YYYY-MM-DD format.") + + try: + parsed = date.fromisoformat(value) + except ValueError as exc: + raise ValueError(f"The {name} parameter must be a valid date in YYYY-MM-DD format.") from exc + + if parsed.isoformat() != value: + raise ValueError(f"The {name} parameter must use YYYY-MM-DD format.") + return value + + def cast_parameter_value(typ: StrEnum, value: Any, /): try: match typ.value: @@ -115,6 +135,34 @@ def cast_parameter_value(typ: StrEnum, value: Any, /): return "" else: return value if isinstance(value, str) else str(value) + case PluginParameterType.DATE: + if value is None or value == "": + return "" + return _validate_date(value) + case PluginParameterType.DATE_RANGE: + if value is None or value == "": + return {} + if isinstance(value, dict): + out: dict[str, str] = {} + for key in ("start", "end"): + if key not in value or value[key] is None or value[key] == "": + continue + out[key] = _validate_date(value[key], f"date-range {key}") + if out.get("start") and out.get("end") and out["start"] > out["end"]: + raise ValueError("The date-range start date must not be after the end date.") + return out + if isinstance(value, str): + try: + parsed_value = json.loads(value) + if isinstance(parsed_value, dict): + return cast_parameter_value(typ, parsed_value) + except json.JSONDecodeError: + pass + stripped = value.strip() + if not stripped: + return {} + return {"start": _validate_date(stripped, "date-range start")} + raise ValueError("The date-range parameter must be a JSON object, JSON string, or empty.") case PluginParameterType.BOOLEAN: match value: @@ -202,7 +250,10 @@ def init_frontend_parameter(rule: PluginParameter, type: StrEnum, value: Any): init frontend parameter by rule """ parameter_value = value - if not parameter_value and parameter_value != 0: + is_empty_tools_selection = ( + type == PluginParameterType.TOOLS_SELECTOR and isinstance(parameter_value, list) and not parameter_value + ) + if not is_empty_tools_selection and not parameter_value and parameter_value != 0: # get default value parameter_value = rule.default if not parameter_value and rule.required: diff --git a/api/core/plugin/entities/plugin_daemon.py b/api/core/plugin/entities/plugin_daemon.py index 4cf55ef8e3c..521884e21c0 100644 --- a/api/core/plugin/entities/plugin_daemon.py +++ b/api/core/plugin/entities/plugin_daemon.py @@ -12,7 +12,7 @@ from core.agent.plugin_entities import AgentProviderEntityWithPlugin from core.datasource.entities.datasource_entities import DatasourceProviderEntityWithPlugin from core.plugin.entities.base import BasePluginEntity from core.plugin.entities.parameters import PluginParameterOption -from core.plugin.entities.plugin import PluginDeclaration, PluginEntity +from core.plugin.entities.plugin import PluginDeclaration, PluginEntity, PluginInstallationSource from core.tools.entities.common_entities import I18nObject from core.tools.entities.tool_entities import ToolProviderEntityWithPlugin from core.trigger.entities.entities import TriggerProviderEntity @@ -83,6 +83,11 @@ class PluginModelSchemaEntity(BaseModel): model_config = ConfigDict(protected_namespaces=()) +class PluginModelProviderDeclaration(ProviderEntity): + plugin_unique_identifier: str = Field(description="The plugin unique identifier.") + installation_source: PluginInstallationSource | None = Field(description="The plugin installation source.") + + class PluginModelProviderEntity(BaseModel): id: str = Field(description="ID") created_at: datetime = Field(description="The created at time of the model provider.") @@ -91,9 +96,25 @@ class PluginModelProviderEntity(BaseModel): tenant_id: str = Field(description="The tenant ID.") plugin_unique_identifier: str = Field(description="The plugin unique identifier.") plugin_id: str = Field(description="The plugin ID.") + installation_source: PluginInstallationSource | None = Field( + default=None, description="The plugin installation source." + ) declaration: ProviderEntity = Field(description="The declaration of the model provider.") +class PluginModelProviderBinding(BaseModel): + """Lightweight installation metadata for one model provider.""" + + provider: str + installation_id: str + plugin_id: str + plugin_unique_identifier: str + runtime_type: str + source: PluginInstallationSource + version: str + verified: bool = False + + class PluginTextEmbeddingNumTokensResponse(BaseModel): """ Response for number of tokens. @@ -207,6 +228,10 @@ class PluginListResponse(BaseModel): total: int +class PluginInstalledIdsDaemonResponse(BaseModel): + plugin_ids: list[str] + + class PluginListWithoutTotalResponse(BaseModel): list: list[PluginEntity] has_more: bool diff --git a/api/core/plugin/impl/model.py b/api/core/plugin/impl/model.py index c69be8a3933..c38399bdcff 100644 --- a/api/core/plugin/impl/model.py +++ b/api/core/plugin/impl/model.py @@ -6,6 +6,7 @@ from core.plugin.entities.plugin_daemon import ( PluginBasicBooleanResponse, PluginDaemonInnerError, PluginLLMNumTokensResponse, + PluginModelProviderBinding, PluginModelProviderEntity, PluginModelSchemaEntity, PluginStringResultResponse, @@ -19,6 +20,7 @@ from graphon.model_runtime.entities.message_entities import PromptMessage, Promp from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelType from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult from graphon.model_runtime.entities.text_embedding_entities import EmbeddingResult +from graphon.model_runtime.protocols.tts_runtime import TTSModelVoice from graphon.model_runtime.utils.encoders import jsonable_encoder _POLLING_UNSUPPORTED_INVOKE_ERROR_TYPES = frozenset((NotImplementedError.__name__,)) @@ -47,6 +49,14 @@ class PluginModelClient(BasePluginClient): ) return response + def fetch_model_provider_bindings(self, tenant_id: str) -> Sequence[PluginModelProviderBinding]: + """Fetch only model-provider installation identities from the daemon.""" + return self._request_with_plugin_daemon_response( + "GET", + f"plugin/{tenant_id}/management/models/bindings", + list[PluginModelProviderBinding], + ) + def get_model_schema( self, tenant_id: str, @@ -614,7 +624,7 @@ class PluginModelClient(BasePluginClient): model: str, credentials: dict[str, Any], language: str | None = None, - ): + ) -> list[TTSModelVoice]: """ Get tts model voices """ @@ -641,7 +651,7 @@ class PluginModelClient(BasePluginClient): ) for resp in response: - voices = [] + voices: list[TTSModelVoice] = [] for voice in resp.voices: voices.append({"name": voice.name, "value": voice.value}) diff --git a/api/core/plugin/impl/model_runtime.py b/api/core/plugin/impl/model_runtime.py index 454bd38958a..2957906bd1d 100644 --- a/api/core/plugin/impl/model_runtime.py +++ b/api/core/plugin/impl/model_runtime.py @@ -13,6 +13,7 @@ from configs import dify_config from core.llm_generator.output_parser.structured_output import ( invoke_llm_with_structured_output as invoke_llm_with_structured_output_helper, ) +from core.plugin.entities.plugin_daemon import PluginModelProviderDeclaration from core.plugin.impl.asset import PluginAssetManager from core.plugin.impl.model import PluginModelClient from core.plugin.plugin_service import PluginService @@ -31,6 +32,7 @@ from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput from graphon.model_runtime.entities.text_embedding_entities import EmbeddingInputType, EmbeddingResult from graphon.model_runtime.model_providers.base.large_language_model import normalize_non_stream_runtime_result from graphon.model_runtime.protocols.runtime import ModelRuntime +from graphon.model_runtime.protocols.tts_runtime import TTSModelVoice from models.provider_ids import ModelProviderID logger = logging.getLogger(__name__) @@ -130,7 +132,7 @@ class PluginModelRuntime(ModelRuntime): self._plugin_service = plugin_service @override - def fetch_model_providers(self) -> Sequence[ProviderEntity]: + def fetch_model_providers(self) -> Sequence[PluginModelProviderDeclaration]: return self._plugin_service.fetch_plugin_model_providers(tenant_id=self.tenant_id, client=self.client) @override @@ -658,7 +660,7 @@ class PluginModelRuntime(ModelRuntime): model: str, credentials: dict[str, Any], language: str | None, - ) -> Any: + ) -> list[TTSModelVoice]: plugin_id, provider_name = self._split_provider(provider) return self.client.get_tts_model_voices( tenant_id=self.tenant_id, diff --git a/api/core/plugin/impl/plugin.py b/api/core/plugin/impl/plugin.py index 34e8d315d86..8ac33b50297 100644 --- a/api/core/plugin/impl/plugin.py +++ b/api/core/plugin/impl/plugin.py @@ -14,6 +14,7 @@ from core.plugin.entities.plugin import ( ) from core.plugin.entities.plugin_daemon import ( PluginDecodeResponse, + PluginInstalledIdsDaemonResponse, PluginInstallTask, PluginInstallTaskStartResponse, PluginListResponse, @@ -68,6 +69,16 @@ class PluginInstaller(BasePluginClient): ) return result.list + def list_installed_plugin_ids(self, tenant_id: str, category: PluginCategory) -> list[str]: + """List all currently installed plugin IDs in one category.""" + result = self._request_with_plugin_daemon_response( + "GET", + f"plugin/{tenant_id}/management/installation/ids", + PluginInstalledIdsDaemonResponse, + params={"category": category.value}, + ) + return result.plugin_ids + def list_plugins_with_total(self, tenant_id: str, page: int, page_size: int) -> PluginListResponse: return self._request_with_plugin_daemon_response( "GET", @@ -77,13 +88,28 @@ class PluginInstaller(BasePluginClient): ) def list_plugins_by_category( - self, tenant_id: str, category: PluginCategory, page: int, page_size: int + self, + tenant_id: str, + category: PluginCategory, + page: int, + page_size: int, + *, + query: str = "", + tags: Sequence[str] = (), + language: str = "en_US", ) -> PluginListWithoutTotalResponse: return self._request_with_plugin_daemon_response( "GET", f"plugin/{tenant_id}/management/{category.value}/list", PluginListWithoutTotalResponse, - params={"page": page, "page_size": page_size, "response_type": "paged"}, + params={ + "page": page, + "page_size": page_size, + "response_type": "paged", + "query": query, + "tags": list(tags), + "language": language, + }, ) def upload_pkg( diff --git a/api/core/plugin/plugin_service.py b/api/core/plugin/plugin_service.py index 89274b635ac..0fdfd4d27b3 100644 --- a/api/core/plugin/plugin_service.py +++ b/api/core/plugin/plugin_service.py @@ -16,7 +16,7 @@ metadata. import logging import time -from collections.abc import Iterator, Mapping, Sequence +from collections.abc import Generator, Mapping, Sequence from contextlib import contextmanager from mimetypes import guess_type from typing import Literal, Protocol @@ -48,6 +48,8 @@ from core.plugin.entities.plugin_daemon import ( PluginInstallTaskStatus, PluginListResponse, PluginListWithoutTotalResponse, + PluginModelProviderBinding, + PluginModelProviderDeclaration, PluginModelProviderEntity, PluginVerification, ) @@ -56,20 +58,23 @@ from core.plugin.impl.debugging import PluginDebuggingClient from core.plugin.impl.endpoint import PluginEndpointClient from core.plugin.impl.model import PluginModelClient from core.plugin.impl.plugin import PluginInstaller +from enums import DeploymentEdition from extensions.ext_database import db from extensions.ext_redis import redis_client -from graphon.model_runtime.entities.provider_entities import ProviderEntity from models.provider import Provider, ProviderCredential, TenantPreferredModelProvider from models.provider_ids import GenericProviderID, ModelProviderID from services.enterprise.plugin_manager_service import ( PluginManagerService, PreUninstallPluginRequest, ) +from services.entities.feature_entities import PluginInstallationPermissionModel, PluginInstallationScope from services.errors.plugin import PluginInstallationForbiddenError -from services.feature_service import FeatureService, PluginInstallationScope +from services.feature_service import FeatureService logger = logging.getLogger(__name__) -_provider_entities_adapter: TypeAdapter[list[ProviderEntity]] = TypeAdapter(list[ProviderEntity]) +_provider_entities_adapter: TypeAdapter[list[PluginModelProviderDeclaration]] = TypeAdapter( + list[PluginModelProviderDeclaration] +) class _RedisLock(Protocol): @@ -78,6 +83,12 @@ class _RedisLock(Protocol): def release(self) -> None: ... +class _ModelPluginIdentity(Protocol): + plugin_id: str + plugin_unique_identifier: str + source: PluginInstallationSource + + class PluginService: class LatestPluginCache(BaseModel): plugin_id: str @@ -146,11 +157,44 @@ class PluginService: return provider.provider @classmethod - def _to_provider_entity(cls, provider: PluginModelProviderEntity) -> ProviderEntity: - declaration = provider.declaration.model_copy(deep=True) - declaration.provider = f"{provider.plugin_id}/{provider.provider}" - declaration.provider_name = cls._get_provider_short_name_alias(provider) - return declaration + def _to_provider_entity( + cls, + provider: PluginModelProviderEntity, + installation_source: PluginInstallationSource | None, + ) -> PluginModelProviderDeclaration: + return PluginModelProviderDeclaration.model_validate( + { + **provider.declaration.model_dump(), + "provider": f"{provider.plugin_id}/{provider.provider}", + "provider_name": cls._get_provider_short_name_alias(provider), + "plugin_unique_identifier": provider.plugin_unique_identifier, + "installation_source": installation_source, + } + ) + + @classmethod + def _resolve_model_provider_installation_sources( + cls, + tenant_id: str, + providers: Sequence[PluginModelProviderEntity], + ) -> Mapping[str, PluginInstallationSource]: + unresolved_plugin_ids = list( + dict.fromkeys(provider.plugin_id for provider in providers if provider.installation_source is None) + ) + if not unresolved_plugin_ids: + return {} + + try: + installations = cls.list_installations_from_ids(tenant_id, unresolved_plugin_ids) + except Exception: + logger.warning( + "Failed to resolve model provider installation sources for tenant %s.", + tenant_id, + exc_info=True, + ) + return {} + + return {installation.plugin_unique_identifier: installation.source for installation in installations} @classmethod def _encode_plugin_model_providers_cache_payload(cls, payload: bytes) -> bytes: @@ -206,7 +250,7 @@ class PluginService: @classmethod def _load_cached_plugin_model_providers_for_generation( cls, tenant_id: str, generation: int | None - ) -> tuple[tuple[ProviderEntity, ...] | None, bool]: + ) -> tuple[tuple[PluginModelProviderDeclaration, ...] | None, bool]: if generation is None: return None, False @@ -253,7 +297,7 @@ class PluginService: @classmethod def _store_cached_plugin_model_providers( - cls, tenant_id: str, generation: int, providers: Sequence[ProviderEntity] + cls, tenant_id: str, generation: int, providers: Sequence[PluginModelProviderDeclaration] ) -> None: cache_key = cls._get_plugin_model_providers_cache_key(tenant_id, generation) try: @@ -265,7 +309,7 @@ class PluginService: logger.warning("Failed to cache plugin model providers for tenant %s.", tenant_id, exc_info=True) @classmethod - def _get_remote_model_plugin_cache_marker(cls, plugins: Sequence[PluginEntity]) -> str | None: + def _get_remote_model_plugin_cache_marker(cls, plugins: Sequence[_ModelPluginIdentity]) -> str | None: remote_model_plugins = sorted( f"{plugin.plugin_id}:{plugin.plugin_unique_identifier}" for plugin in plugins @@ -350,7 +394,7 @@ class PluginService: def _should_invalidate_model_provider_cache_for_remote_model_plugins( cls, tenant_id: str, - plugins: Sequence[PluginEntity], + plugins: Sequence[_ModelPluginIdentity], ) -> bool: remote_model_plugin_marker = cls._get_remote_model_plugin_cache_marker(plugins) cached_remote_model_plugin_marker = cls._load_cached_remote_model_plugin_marker(tenant_id) @@ -373,7 +417,7 @@ class PluginService: @contextmanager def _plugin_model_providers_refresh_lock( cls, tenant_id: str, generation: int, *, wait_timeout: float - ) -> Iterator[bool]: + ) -> Generator[bool]: lock_key = cls._get_plugin_model_providers_lock_key(tenant_id, generation) try: refresh_lock: _RedisLock = redis_client.lock( @@ -434,14 +478,26 @@ class PluginService: exc_info=True, ) + @classmethod + def _fetch_plugin_model_providers_uncached( + cls, tenant_id: str, client: PluginModelClient | None + ) -> tuple[PluginModelProviderDeclaration, ...]: + model_client = client or PluginModelClient() + providers = model_client.fetch_model_providers(tenant_id) + installation_sources = cls._resolve_model_provider_installation_sources(tenant_id, providers) + return tuple( + cls._to_provider_entity( + provider, + provider.installation_source or installation_sources.get(provider.plugin_unique_identifier), + ) + for provider in providers + ) + @classmethod def _fetch_and_cache_plugin_model_providers( cls, tenant_id: str, client: PluginModelClient | None, *, refresh_generation: int | None - ) -> tuple[ProviderEntity, ...]: - model_client = client or PluginModelClient() - providers = tuple( - cls._to_provider_entity(provider) for provider in model_client.fetch_model_providers(tenant_id) - ) + ) -> tuple[PluginModelProviderDeclaration, ...]: + providers = cls._fetch_plugin_model_providers_uncached(tenant_id, client) generation = cls._load_plugin_model_providers_generation(tenant_id) if generation is not None and generation == refresh_generation: cls._store_cached_plugin_model_providers(tenant_id, generation, providers) @@ -463,7 +519,7 @@ class PluginService: @classmethod def fetch_plugin_model_providers( cls, *, tenant_id: str, client: PluginModelClient | None = None - ) -> Sequence[ProviderEntity]: + ) -> Sequence[PluginModelProviderDeclaration]: """ Fetch plugin model providers through the tenant-scoped plugin cache. @@ -471,6 +527,9 @@ class PluginService: are intentionally owned by this service so tenant isolation and cache expiry are handled in one place. """ + if not dify_config.PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED: + return cls._fetch_plugin_model_providers_uncached(tenant_id, client) + deadline = time.monotonic() + cls.PLUGIN_MODEL_PROVIDERS_LOCK_WAIT_TIMEOUT while True: @@ -597,22 +656,30 @@ class PluginService: return result @staticmethod - def _check_marketplace_only_permission(): + def _check_marketplace_only_permission() -> None: """ Check if the marketplace only permission is enabled """ - features = FeatureService.get_system_features() - if features.plugin_installation_permission.restrict_to_marketplace_only: + permission = PluginService._get_plugin_installation_permission() + if permission.restrict_to_marketplace_only: raise PluginInstallationForbiddenError("Plugin installation is restricted to marketplace only") @staticmethod - def _check_plugin_installation_scope(plugin_verification: PluginVerification | None): + def _get_plugin_installation_permission() -> PluginInstallationPermissionModel: + """Resolve the validated policy and reject deny-all before any installation side effect.""" + permission = FeatureService.get_plugin_installation_permission() + if permission.plugin_installation_scope == PluginInstallationScope.NONE: + raise PluginInstallationForbiddenError("Installing plugins is not allowed") + return permission + + @staticmethod + def _check_plugin_installation_scope(plugin_verification: PluginVerification | None) -> None: """ Check the plugin installation scope """ - features = FeatureService.get_system_features() + permission = PluginService._get_plugin_installation_permission() - match features.plugin_installation_permission.plugin_installation_scope: + match permission.plugin_installation_scope: case PluginInstallationScope.OFFICIAL_ONLY: if ( plugin_verification is None @@ -627,10 +694,10 @@ class PluginService: raise PluginInstallationForbiddenError( "Plugin installation is restricted to official and specific partners" ) - case PluginInstallationScope.NONE: - raise PluginInstallationForbiddenError("Installing plugins is not allowed") case PluginInstallationScope.ALL: pass + case _: + raise PluginInstallationForbiddenError("Plugin installation policy is invalid") @staticmethod def get_debugging_key(tenant_id: str) -> str: @@ -656,6 +723,26 @@ class PluginService: plugins = manager.list_plugins(tenant_id) return plugins + @staticmethod + def list_installed_plugin_ids(tenant_id: str, category: PluginCategory) -> Sequence[str]: + """List all currently installed plugin IDs in one category through the daemon's lightweight query.""" + manager = PluginInstaller() + return manager.list_installed_plugin_ids(tenant_id, category) + + @staticmethod + def list_model_provider_bindings( + tenant_id: str, *, client: PluginModelClient | None = None + ) -> Sequence[PluginModelProviderBinding]: + """Return fresh model bindings and reconcile remote-debug provider metadata before it is read.""" + model_client = client or PluginModelClient() + bindings = model_client.fetch_model_provider_bindings(tenant_id) + if PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins(tenant_id, bindings): + PluginService.invalidate_plugin_model_providers_cache(tenant_id) + + marker = PluginService._get_remote_model_plugin_cache_marker(bindings) + PluginService._store_cached_remote_model_plugin_marker(tenant_id, marker) + return bindings + @staticmethod def list_with_total(tenant_id: str, user_id: str, page: int, page_size: int) -> PluginListResponse: """List tenant plugins with endpoint counts reconciled from live records. @@ -672,17 +759,33 @@ class PluginService: @staticmethod def list_by_category( - tenant_id: str, category: PluginCategory, page: int, page_size: int + tenant_id: str, + category: PluginCategory, + page: int, + page_size: int, + *, + query: str = "", + tags: Sequence[str] = (), + language: str = "en_US", ) -> PluginListWithoutTotalResponse: """ List plugins in one category with a has-more cursor signal and without calculating total. - The daemon scans tenant installations in the existing list order and stops once it finds one extra match. - This keeps pagination usable before category is persisted on installation rows. + The daemon applies category, search, and tag filters before pagination, then stops once it finds one extra + match. Only a complete, unfiltered first page may reconcile the model-provider cache; the unpaginated model + binding read is the authoritative marker source for larger result sets. """ manager = PluginInstaller() - plugins = manager.list_plugins_by_category(tenant_id, category, page, page_size) - if category == PluginCategory.Model: + plugins = manager.list_plugins_by_category( + tenant_id, + category, + page, + page_size, + query=query, + tags=tags, + language=language, + ) + if category == PluginCategory.Model and page == 1 and not plugins.has_more and not query and not tags: should_invalidate_model_provider_cache = ( PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins( tenant_id, @@ -900,23 +1003,26 @@ class PluginService: # check if plugin pkg is already downloaded manager = PluginInstaller() - features = FeatureService.get_system_features() + permission = PluginService._get_plugin_installation_permission() try: manager.fetch_plugin_manifest(tenant_id, new_plugin_unique_identifier) - # already downloaded, skip, and record install event - marketplace.record_install_plugin_event(new_plugin_unique_identifier) except Exception: - # plugin not installed, download and upload pkg + # plugin not downloaded yet, download and upload pkg pkg = download_plugin_pkg(new_plugin_unique_identifier) response = manager.upload_pkg( tenant_id, pkg, - verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only, + verify_signature=permission.restrict_to_marketplace_only, ) # check if the plugin is available to install PluginService._check_plugin_installation_scope(response.verification) + else: + # already downloaded, the cached pkg still has to satisfy the installation scope + decode_response = manager.decode_plugin_from_identifier(tenant_id, new_plugin_unique_identifier) + PluginService._check_plugin_installation_scope(decode_response.verification) + marketplace.record_install_plugin_event(new_plugin_unique_identifier) result = manager.upgrade_plugin( tenant_id, @@ -967,11 +1073,11 @@ class PluginService: """ PluginService._check_marketplace_only_permission() manager = PluginInstaller() - features = FeatureService.get_system_features() + permission = PluginService._get_plugin_installation_permission() response = manager.upload_pkg( tenant_id, pkg, - verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only, + verify_signature=permission.restrict_to_marketplace_only, ) PluginService._check_plugin_installation_scope(response.verification) @@ -989,13 +1095,13 @@ class PluginService: pkg = download_with_size_limit( f"https://github.com/{repo}/releases/download/{version}/{package}", dify_config.PLUGIN_MAX_PACKAGE_SIZE ) - features = FeatureService.get_system_features() + permission = PluginService._get_plugin_installation_permission() manager = PluginInstaller() response = manager.upload_pkg( tenant_id, pkg, - verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only, + verify_signature=permission.restrict_to_marketplace_only, ) PluginService._check_plugin_installation_scope(response.verification) @@ -1069,7 +1175,7 @@ class PluginService: if not dify_config.MARKETPLACE_ENABLED: raise ValueError("marketplace is not enabled") - features = FeatureService.get_system_features() + permission = PluginService._get_plugin_installation_permission() manager = PluginInstaller() try: @@ -1079,7 +1185,7 @@ class PluginService: response = manager.upload_pkg( tenant_id, pkg, - verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only, + verify_signature=permission.restrict_to_marketplace_only, ) # check if the plugin is available to install PluginService._check_plugin_installation_scope(response.verification) @@ -1101,31 +1207,31 @@ class PluginService: # collect actual plugin_unique_identifiers actual_plugin_unique_identifiers = [] metas = [] - features = FeatureService.get_system_features() + permission = PluginService._get_plugin_installation_permission() # check if already downloaded for plugin_unique_identifier in plugin_unique_identifiers: try: manager.fetch_plugin_manifest(tenant_id, plugin_unique_identifier) - plugin_decode_response = manager.decode_plugin_from_identifier(tenant_id, plugin_unique_identifier) - # check if the plugin is available to install - PluginService._check_plugin_installation_scope(plugin_decode_response.verification) - # already downloaded, skip - actual_plugin_unique_identifiers.append(plugin_unique_identifier) - metas.append({"plugin_unique_identifier": plugin_unique_identifier}) except Exception: - # plugin not installed, download and upload pkg + # plugin not downloaded yet, download and upload pkg pkg = download_plugin_pkg(plugin_unique_identifier) response = manager.upload_pkg( tenant_id, pkg, - verify_signature=features.plugin_installation_permission.restrict_to_marketplace_only, + verify_signature=permission.restrict_to_marketplace_only, ) # check if the plugin is available to install PluginService._check_plugin_installation_scope(response.verification) # use response plugin_unique_identifier actual_plugin_unique_identifiers.append(response.unique_identifier) metas.append({"plugin_unique_identifier": response.unique_identifier}) + else: + # already downloaded, the cached pkg still has to satisfy the installation scope + plugin_decode_response = manager.decode_plugin_from_identifier(tenant_id, plugin_unique_identifier) + PluginService._check_plugin_installation_scope(plugin_decode_response.verification) + actual_plugin_unique_identifiers.append(plugin_unique_identifier) + metas.append({"plugin_unique_identifier": plugin_unique_identifier}) result = manager.install_from_identifiers( tenant_id, @@ -1137,10 +1243,11 @@ class PluginService: return result @staticmethod - def uninstall(tenant_id: str, plugin_installation_id: str) -> bool: + def uninstall(tenant_id: str, plugin_installation_id: str, *, preserve_credentials: bool = False) -> bool: + """Uninstall a plugin and optionally retain model-provider credentials for replacement.""" manager = PluginInstaller() - # Get plugin info before uninstalling to delete associated credentials + # Resolve the trusted plugin ID before the daemon removes the installation record. plugins = manager.list_plugins(tenant_id) plugin = next((p for p in plugins if p.installation_id == plugin_installation_id), None) @@ -1150,66 +1257,81 @@ class PluginService: PluginService.invalidate_plugin_model_providers_cache(tenant_id) return result - if dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE: PluginManagerService.try_pre_uninstall_plugin( PreUninstallPluginRequest( tenant_id=tenant_id, plugin_unique_identifier=plugin.plugin_unique_identifier, ) ) - with Session(db.engine) as session, session.begin(): + + result = manager.uninstall(tenant_id, plugin_installation_id) + if not result: + return False + + if preserve_credentials: + logger.info("Preserving credentials while replacing plugin: %s", plugin.plugin_id) + else: plugin_id = plugin.plugin_id - logger.info("Deleting credentials for plugin: %s", plugin_id) + logger.info("Deleting credentials after uninstalling plugin: %s", plugin_id) + provider_ids: Sequence[str] = [] - session.execute( - delete(TenantPreferredModelProvider).where( - TenantPreferredModelProvider.tenant_id == tenant_id, - TenantPreferredModelProvider.provider_name.like(f"{plugin_id}/%"), + with Session(db.engine) as session, session.begin(): + session.execute( + delete(TenantPreferredModelProvider).where( + TenantPreferredModelProvider.tenant_id == tenant_id, + TenantPreferredModelProvider.provider_name.like(f"{plugin_id}/%"), + ) ) - ) - # Delete provider credentials that match this plugin - credential_ids = session.scalars( - select(ProviderCredential.id).where( - ProviderCredential.tenant_id == tenant_id, - ProviderCredential.provider_name.like(f"{plugin_id}/%"), - ) - ).all() - - if not credential_ids: - logger.info("No credentials found for plugin: %s", plugin_id) - else: - provider_ids = session.scalars( - select(Provider.id).where( - Provider.tenant_id == tenant_id, - Provider.provider_name.like(f"{plugin_id}/%"), - Provider.credential_id.in_(credential_ids), + credential_ids = session.scalars( + select(ProviderCredential.id).where( + ProviderCredential.tenant_id == tenant_id, + ProviderCredential.provider_name.like(f"{plugin_id}/%"), ) ).all() - session.execute(update(Provider).where(Provider.id.in_(provider_ids)).values(credential_id=None)) + if not credential_ids: + logger.info("No credentials found for plugin: %s", plugin_id) + else: + provider_ids = session.scalars( + select(Provider.id).where( + Provider.tenant_id == tenant_id, + Provider.provider_name.like(f"{plugin_id}/%"), + Provider.credential_id.in_(credential_ids), + ) + ).all() - for provider_id in provider_ids: - ProviderCredentialsCache( - tenant_id=tenant_id, - identity_id=provider_id, - cache_type=ProviderCredentialsCacheType.PROVIDER, - ).delete() - - session.execute( - delete(ProviderCredential).where( - ProviderCredential.id.in_(credential_ids), + session.execute(update(Provider).where(Provider.id.in_(provider_ids)).values(credential_id=None)) + session.execute( + delete(ProviderCredential).where( + ProviderCredential.id.in_(credential_ids), + ) ) - ) - logger.info( - "Completed deleting credentials and cleaning provider associations for plugin: %s", - plugin_id, - ) + logger.info( + "Completed deleting credentials and cleaning provider associations for plugin: %s", + plugin_id, + ) - result = manager.uninstall(tenant_id, plugin_installation_id) - if result: - PluginService.invalidate_plugin_model_providers_cache(tenant_id) + for provider_id in provider_ids: + ProviderCredentialsCache( + tenant_id=tenant_id, + identity_id=provider_id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ).delete() + + from core.provider_manager import ProviderConfigurationCacheSource, ProviderManager + + ProviderManager.invalidate_configurations_cache( + tenant_id, + sources=( + ProviderConfigurationCacheSource.PREFERRED_MODEL_PROVIDERS, + ProviderConfigurationCacheSource.PROVIDER_CREDENTIALS, + ), + ) + + PluginService.invalidate_plugin_model_providers_cache(tenant_id) return result @staticmethod diff --git a/api/core/prompt/advanced_prompt_transform.py b/api/core/prompt/advanced_prompt_transform.py index e7c88811fe5..5f9f2384568 100644 --- a/api/core/prompt/advanced_prompt_transform.py +++ b/api/core/prompt/advanced_prompt_transform.py @@ -19,6 +19,7 @@ from graphon.model_runtime.entities import ( ) from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent, PromptMessageContentUnionTypes from graphon.runtime import VariablePool +from graphon.variables.template_resolution import convert_template class AdvancedPromptTransform(PromptTransform): @@ -171,7 +172,7 @@ class AdvancedPromptTransform(PromptTransform): if k.startswith("#"): vp.add(k[1:-1].split("."), v) raw_prompt = raw_prompt.replace("{{#context#}}", context or "") - prompt = vp.convert_template(raw_prompt).text + prompt = convert_template(vp, raw_prompt).text else: parser = PromptTemplateParser(template=raw_prompt, with_variable_tmpl=self.with_variable_tmpl) prompt_inputs: Mapping[str, str] = {k: inputs[k] for k in parser.variable_keys if k in inputs} diff --git a/api/core/provider_manager.py b/api/core/provider_manager.py index 0064252b64a..4baeaed95bf 100644 --- a/api/core/provider_manager.py +++ b/api/core/provider_manager.py @@ -10,7 +10,7 @@ from enum import StrEnum from json import JSONDecodeError from typing import TYPE_CHECKING, Any, Protocol, Self -from pydantic import TypeAdapter +from pydantic import TypeAdapter, ValidationError from sqlalchemy import select from sqlalchemy.exc import IntegrityError @@ -34,7 +34,9 @@ from core.entities.provider_entities import ( from core.helper import encrypter from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType from core.helper.position_helper import is_filtered -from enums.deployment_edition import DeploymentEdition +from core.plugin.entities.plugin import PluginInstallationSource +from core.plugin.entities.plugin_daemon import PluginModelProviderDeclaration +from enums import DeploymentEdition from extensions import ext_hosting_provider from extensions.ext_database import db from extensions.ext_redis import redis_client @@ -572,17 +574,24 @@ class ProviderManager: instance scope. """ - decoding_rsa_key: Any | None - decoding_cipher_rsa: Any | None + # Keyed by tenant_id -- a single ProviderManager instance may be asked to decrypt + # credentials belonging to different tenants (e.g. load balancing configs each carry + # their own tenant_id), so this cache must not collapse to a single shared value. + _decoding_contexts: dict[str, Any] _model_runtime: ModelRuntime _configurations_cache: dict[str, ProviderConfigurations] def __init__(self, model_runtime: ModelRuntime): - self.decoding_rsa_key = None - self.decoding_cipher_rsa = None + self._decoding_contexts = {} self._model_runtime = model_runtime self._configurations_cache = {} + def _get_decoding_context(self, tenant_id: str) -> Any: + """Return this manager's cached decoding context for `tenant_id`, fetching it once if absent.""" + if tenant_id not in self._decoding_contexts: + self._decoding_contexts[tenant_id] = encrypter.get_decrypt_decoding(tenant_id) + return self._decoding_contexts[tenant_id] + def clear_configurations_cache(self, tenant_id: str | None = None) -> None: """Drop assembled provider configurations cached on this manager instance.""" if tenant_id is None: @@ -757,7 +766,11 @@ class ProviderManager: has_valid_quota = any(quota_conf.is_valid for quota_conf in system_configuration.quota_configurations) if preferred_provider_type == ProviderType.SYSTEM: - if not system_configuration.enabled or not has_valid_quota: + if not system_configuration.enabled or not system_configuration.quota_configurations: + using_provider_type = ProviderType.CUSTOM + elif not has_valid_quota and (custom_configuration.provider or custom_configuration.models): + # Only configured alternatives can serve as fallbacks; otherwise downstream checks must surface + # system quota exhaustion instead of reporting missing custom credentials. using_provider_type = ProviderType.CUSTOM else: @@ -1495,16 +1508,14 @@ class ProviderManager: return {} # Decrypt secret variables - if self.decoding_rsa_key is None or self.decoding_cipher_rsa is None: - self.decoding_rsa_key, self.decoding_cipher_rsa = encrypter.get_decrypt_decoding(tenant_id) + decoding_context = self._get_decoding_context(tenant_id) for variable in secret_variables: if variable in credentials: with contextlib.suppress(ValueError): credentials[variable] = encrypter.decrypt_token_with_decoding( credentials.get(variable) or "", - self.decoding_rsa_key, - self.decoding_cipher_rsa, + decoding_context, ) # Cache the decrypted credentials @@ -1529,6 +1540,19 @@ class ProviderManager: if provider_hosting_configuration is None or not provider_hosting_configuration.enabled: return SystemConfiguration(enabled=False) + try: + plugin_provider_entity = PluginModelProviderDeclaration.model_validate(provider_entity) + except ValidationError: + return SystemConfiguration(enabled=False) + + if plugin_provider_entity.installation_source != PluginInstallationSource.Marketplace: + return SystemConfiguration(enabled=False) + + from core.plugin.plugin_service import PluginService + + if not PluginService.is_plugin_verified(tenant_id, plugin_provider_entity.plugin_unique_identifier): + return SystemConfiguration(enabled=False) + # Convert provider_records to dict quota_type_to_provider_records_dict: dict[ProviderQuotaType, Provider] = {} for provider_record in provider_records: @@ -1542,16 +1566,17 @@ class ProviderManager: if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: from services.credit_pool_service import CreditPoolService - trail_pool = CreditPoolService.get_pool( - tenant_id=tenant_id, - pool_type=ProviderQuotaType.TRIAL, - session=db.session(), - ) - paid_pool = CreditPoolService.get_pool( - tenant_id=tenant_id, - pool_type=ProviderQuotaType.PAID, - session=db.session(), - ) + with session_factory.create_session() as session: + trail_pool = CreditPoolService.get_pool( + tenant_id=tenant_id, + pool_type=ProviderQuotaType.TRIAL, + session=session, + ) + paid_pool = CreditPoolService.get_pool( + tenant_id=tenant_id, + pool_type=ProviderQuotaType.PAID, + session=session, + ) else: trail_pool = None paid_pool = None @@ -1642,17 +1667,15 @@ class ProviderManager: else [] ) - # Get decoding rsa key and cipher for decrypting credentials - if self.decoding_rsa_key is None or self.decoding_cipher_rsa is None: - self.decoding_rsa_key, self.decoding_cipher_rsa = encrypter.get_decrypt_decoding(tenant_id) + # Get decoding context for decrypting credentials + decoding_context = self._get_decoding_context(tenant_id) for variable in provider_credential_secret_variables: if variable in provider_credentials: try: provider_credentials[variable] = encrypter.decrypt_token_with_decoding( provider_credentials.get(variable, ""), - self.decoding_rsa_key, - self.decoding_cipher_rsa, + decoding_context, ) except ValueError: pass @@ -1786,19 +1809,15 @@ class ProviderManager: except (ValueError, JSONDecodeError): continue - # Get decoding rsa key and cipher for decrypting credentials - if self.decoding_rsa_key is None or self.decoding_cipher_rsa is None: - self.decoding_rsa_key, self.decoding_cipher_rsa = encrypter.get_decrypt_decoding( - load_balancing_model_config.tenant_id - ) + # Get decoding context for decrypting credentials + decoding_context = self._get_decoding_context(load_balancing_model_config.tenant_id) for variable in model_credential_secret_variables: if variable in provider_model_credentials: try: provider_model_credentials[variable] = encrypter.decrypt_token_with_decoding( provider_model_credentials.get(variable) or "", - self.decoding_rsa_key, - self.decoding_cipher_rsa, + decoding_context, ) except ValueError: pass diff --git a/api/core/rag/cleaner/clean_processor.py b/api/core/rag/cleaner/clean_processor.py index 790253053de..452251584e6 100644 --- a/api/core/rag/cleaner/clean_processor.py +++ b/api/core/rag/cleaner/clean_processor.py @@ -9,7 +9,7 @@ class CleanProcessor: # remove invalid symbol text = re.sub(r"<\|", "<", text) text = re.sub(r"\|>", ">", text) - text = re.sub(r"[\x00-\x08\x0B\x0C\x0E-\x1F\x7F\xEF\xBF\xBE]", "", text) + text = re.sub(r"[\x00-\x08\x0B\x0C\x0E-\x1F\x7F]", "", text) # Unicode U+FFFE text = re.sub("\ufffe", "", text) diff --git a/api/core/rag/datasource/vdb/vector_factory.py b/api/core/rag/datasource/vdb/vector_factory.py index 1bd0d8cafbb..a355a90d0f4 100644 --- a/api/core/rag/datasource/vdb/vector_factory.py +++ b/api/core/rag/datasource/vdb/vector_factory.py @@ -128,15 +128,16 @@ class Vector: self._session = session self._vector_processor = self._init_vector(session=session) - def _init_vector(self, *, session: Session) -> BaseVector: + @staticmethod + def resolve_vector_type(dataset: Dataset, *, session: Session) -> str: vector_type = dify_config.VECTOR_STORE - if self._dataset.index_struct_dict: - vector_type = self._dataset.index_struct_dict["type"] + if dataset.index_struct_dict: + vector_type = dataset.index_struct_dict["type"] else: if dify_config.VECTOR_STORE_WHITELIST_ENABLE: stmt = select(Whitelist).where( - Whitelist.tenant_id == self._dataset.tenant_id, Whitelist.category == "vector_db" + Whitelist.tenant_id == dataset.tenant_id, Whitelist.category == "vector_db" ) whitelist = session.scalars(stmt).one_or_none() if whitelist: @@ -145,6 +146,10 @@ class Vector: if not vector_type: raise ValueError("Vector store must be specified.") + return vector_type + + def _init_vector(self, *, session: Session) -> BaseVector: + vector_type = self.resolve_vector_type(self._dataset, session=session) vector_factory_cls = self.get_vector_factory(vector_type) return vector_factory_cls().init_vector(self._dataset, self._attributes, self._embeddings) diff --git a/api/core/rag/extractor/firecrawl/firecrawl_app.py b/api/core/rag/extractor/firecrawl/firecrawl_app.py index 556158cf00a..77edffa77cd 100644 --- a/api/core/rag/extractor/firecrawl/firecrawl_app.py +++ b/api/core/rag/extractor/firecrawl/firecrawl_app.py @@ -6,6 +6,10 @@ import httpx from extensions.ext_storage import storage +# Bounded connect/read timeout so a slow or hanging Firecrawl endpoint cannot +# block extraction indefinitely (mirrors the WaterCrawl extractor client). +_REQUEST_TIMEOUT = httpx.Timeout(30.0, connect=5.0) + class FirecrawlDocumentData(TypedDict): title: str | None @@ -176,7 +180,7 @@ class FirecrawlApp: def _post_request(self, url, data, headers, retries=3, backoff_factor=0.5) -> httpx.Response: response: httpx.Response | None = None for attempt in range(retries): - response = httpx.post(url, headers=headers, json=data) + response = httpx.post(url, headers=headers, json=data, timeout=_REQUEST_TIMEOUT) if response.status_code == 502: time.sleep(backoff_factor * (2**attempt)) else: @@ -187,7 +191,7 @@ class FirecrawlApp: def _get_request(self, url, headers, retries=3, backoff_factor=0.5) -> httpx.Response: response: httpx.Response | None = None for attempt in range(retries): - response = httpx.get(url, headers=headers) + response = httpx.get(url, headers=headers, timeout=_REQUEST_TIMEOUT) if response.status_code == 502: time.sleep(backoff_factor * (2**attempt)) else: diff --git a/api/core/rag/extractor/notion_extractor.py b/api/core/rag/extractor/notion_extractor.py index 568ccb1912a..0fc57c830fb 100644 --- a/api/core/rag/extractor/notion_extractor.py +++ b/api/core/rag/extractor/notion_extractor.py @@ -21,6 +21,11 @@ SEARCH_URL = "https://api.notion.com/v1/search" RETRIEVE_PAGE_URL_TMPL = "https://api.notion.com/v1/pages/{page_id}" RETRIEVE_DATABASE_URL_TMPL = "https://api.notion.com/v1/databases/{database_id}" + +# Bounded connect/read timeout so a slow or hanging Notion API cannot block +# dataset extraction indefinitely. +_REQUEST_TIMEOUT = httpx.Timeout(30.0, connect=5.0) + # if user want split by headings, use the corresponding splitter HEADING_SPLITTER = { "heading_1": "# ", @@ -110,6 +115,7 @@ class NotionExtractor(BaseExtractor): "Notion-Version": "2022-06-28", }, json=current_query, + timeout=_REQUEST_TIMEOUT, ) response_data = res.json() @@ -179,6 +185,7 @@ class NotionExtractor(BaseExtractor): "Notion-Version": "2022-06-28", }, params=query_dict, + timeout=_REQUEST_TIMEOUT, ) if res.status_code != 200: raise ValueError(f"Error fetching Notion block data: {res.text}") @@ -241,6 +248,7 @@ class NotionExtractor(BaseExtractor): "Notion-Version": "2022-06-28", }, params=query_dict, + timeout=_REQUEST_TIMEOUT, ) data = res.json() if "results" not in data or data["results"] is None: @@ -301,6 +309,7 @@ class NotionExtractor(BaseExtractor): "Notion-Version": "2022-06-28", }, params=query_dict, + timeout=_REQUEST_TIMEOUT, ) data = res.json() # get table headers text @@ -375,6 +384,7 @@ class NotionExtractor(BaseExtractor): "Notion-Version": "2022-06-28", }, json=query_dict, + timeout=_REQUEST_TIMEOUT, ) data = res.json() diff --git a/api/core/rag/index_processor/index_processor.py b/api/core/rag/index_processor/index_processor.py index bf6eb0a1262..dc951c554fe 100644 --- a/api/core/rag/index_processor/index_processor.py +++ b/api/core/rag/index_processor/index_processor.py @@ -15,6 +15,7 @@ from core.rag.index_processor.index_processor_base import SummaryIndexSettingDic from core.workflow.nodes.knowledge_index.exc import KnowledgeIndexNodeError from core.workflow.nodes.knowledge_index.protocols import IndexingResultDict, Preview, PreviewItem, QaPreview from models.dataset import Dataset, Document, DocumentSegment +from services.vector_space_admission_service import VectorSpaceAdmissionService from .index_processor_factory import IndexProcessorFactory from .processor.paragraph_index_processor import ParagraphIndexProcessor @@ -103,7 +104,18 @@ class IndexProcessor: indexing_start_at = time.perf_counter() # The metadata reads above must not keep a transaction open across vector I/O. session.commit() - # delete from vector index + + # V1 guards only first-time indexing. + if not original_document_id: + VectorSpaceAdmissionService().ensure_pipeline_can_be_indexed( + dataset=dataset, + document_id=document.id, + chunk_structure=dataset.chunk_structure, + chunks=chunks, + include_summaries=bool(summary_index_setting and summary_index_setting.get("enable")), + session=session, + ) + if index_node_ids: index_processor.clean( dataset, index_node_ids, with_keywords=True, delete_child_chunks=True, session=session diff --git a/api/core/rag/index_processor/processor/qa_index_processor.py b/api/core/rag/index_processor/processor/qa_index_processor.py index 4b55d0d75a9..e1971da27fa 100644 --- a/api/core/rag/index_processor/processor/qa_index_processor.py +++ b/api/core/rag/index_processor/processor/qa_index_processor.py @@ -1,4 +1,4 @@ -"""Paragraph index processor.""" +"""Question-and-answer index processor.""" import logging import re diff --git a/api/core/rag/retrieval/dataset_retrieval.py b/api/core/rag/retrieval/dataset_retrieval.py index b89931f57ff..197ad138360 100644 --- a/api/core/rag/retrieval/dataset_retrieval.py +++ b/api/core/rag/retrieval/dataset_retrieval.py @@ -9,6 +9,7 @@ from collections.abc import Generator, Mapping from typing import Any, Union, cast from flask import Flask, current_app +from opentelemetry.trace import get_current_span from sqlalchemy import and_, func, literal, or_, select, update from sqlalchemy.orm import Session, sessionmaker @@ -650,16 +651,19 @@ class DatasetRetrieval: self._record_usage(router_usage) timer = None if dataset_id: - # get retrieval model config - dataset_stmt = select(Dataset).where(Dataset.id == dataset_id) - selected_dataset = session.scalar(dataset_stmt) + allowed_dataset = next((dataset for dataset in available_datasets if dataset.id == dataset_id), None) + selected_dataset = ( + session.scalar(select(Dataset).where(Dataset.id == allowed_dataset.id, Dataset.tenant_id == tenant_id)) + if allowed_dataset + else None + ) if selected_dataset: results = [] if selected_dataset.provider == "external": external_documents = ExternalDatasetService.fetch_external_knowledge_retrieval( session=session, tenant_id=selected_dataset.tenant_id, - dataset_id=dataset_id, + dataset_id=selected_dataset.id, query=query, external_retrieval_parameters=selected_dataset.retrieval_model, metadata_condition=metadata_condition, @@ -673,7 +677,7 @@ class DatasetRetrieval: if document.metadata is not None: document.metadata["score"] = external_document.get("score") document.metadata["title"] = external_document.get("title") - document.metadata["dataset_id"] = dataset_id + document.metadata["dataset_id"] = selected_dataset.id document.metadata["dataset_name"] = selected_dataset.name results.append(document) else: @@ -723,7 +727,7 @@ class DatasetRetrieval: weights=retrieval_model_config.get("weights", None), document_ids_filter=document_ids_filter, ) - self._on_query(query, None, [dataset_id], app_id, user_from, user_id) + self._on_query(query, None, [selected_dataset.id], app_id, user_from, user_id) if results: thread = threading.Thread( @@ -1204,8 +1208,9 @@ class DatasetRetrieval: attachment_ids: list[str] | None, cancel_event: threading.Event | None, thread_exceptions: list[Exception] | None, + skip_on_error: bool = False, ) -> None: - """Collect errors only after they pass through the traced retrieval method.""" + """Collect errors after tracing, or skip dataset-level failures when requested.""" try: self._run_retriever_thread( flask_app=flask_app, @@ -1218,6 +1223,25 @@ class DatasetRetrieval: attachment_ids=attachment_ids, ) except Exception as exc: + if skip_on_error: + logger.warning( + "Skipping dataset retrieval because retriever failed, dataset_id=%s, error_type=%s, error=%s", + dataset_id, + type(exc).__name__, + str(exc), + ) + span = get_current_span() + if span and span.is_recording(): + span.add_event( + "dataset_retrieval.skipped", + attributes={ + "dataset_id": dataset_id, + "error.type": type(exc).__name__, + "error.message": str(exc), + }, + ) + return + if cancel_event: cancel_event.set() if thread_exceptions is not None: @@ -1878,6 +1902,7 @@ class DatasetRetrieval: "attachment_ids": [attachment_id] if attachment_id else None, "cancel_event": cancel_event, "thread_exceptions": retrieval_thread_exceptions, + "skip_on_error": True, }, ) threads.append(retrieval_thread) @@ -2003,8 +2028,11 @@ class DatasetRetrieval: results = session.scalars( select(Dataset) .outerjoin(subquery, Dataset.id == subquery.c.dataset_id) - .where(Dataset.tenant_id == tenant_id, Dataset.id.in_(dataset_ids)) - .where((subquery.c.available_document_count > 0) | (Dataset.provider == "external")) + .where( + Dataset.tenant_id == tenant_id, + Dataset.id.in_(dataset_ids), + (subquery.c.available_document_count > 0) | (Dataset.provider == "external"), + ) ).all() available_datasets = [] @@ -2023,6 +2051,8 @@ class DatasetRetrieval: redis_client.zremrangebyscore(key, 0, current_time - 60000) request_count = redis_client.zcard(key) if request_count > knowledge_rate_limit.limit: + # The rate-limit exception is raised after this block, so commit the audit row + # explicitly instead of relying on the Session context, which only closes it. with session_factory.create_session() as session: rate_limit_log = RateLimitLog( tenant_id=tenant_id, @@ -2030,6 +2060,7 @@ class DatasetRetrieval: operation="knowledge", ) session.add(rate_limit_log) + session.commit() raise exc.RateLimitExceededError( "you have reached the knowledge base request rate limit of your subscription." ) diff --git a/api/core/rag/retrieval/router/multi_dataset_react_route.py b/api/core/rag/retrieval/router/multi_dataset_react_route.py index 21a9d04f7f2..95ffd4c84ca 100644 --- a/api/core/rag/retrieval/router/multi_dataset_react_route.py +++ b/api/core/rag/retrieval/router/multi_dataset_react_route.py @@ -2,7 +2,6 @@ from collections.abc import Generator, Sequence from typing import Any, Union from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity -from core.app.llm import deduct_llm_quota from core.model_manager import ModelInstance, ModelManager from core.prompt.advanced_prompt_transform import AdvancedPromptTransform from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate @@ -168,9 +167,6 @@ class ReactMultiDatasetRouter: # handle invoke result text, usage = self._handle_invoke_result(invoke_result=invoke_result) - # deduct quota - deduct_llm_quota(tenant_id=tenant_id, model_instance=bound_model_instance, usage=usage) - return text, usage def _handle_invoke_result(self, invoke_result: Generator) -> tuple[str, LLMUsage]: diff --git a/api/core/rag/splitter/fixed_text_splitter.py b/api/core/rag/splitter/fixed_text_splitter.py index 98ba4d7fcc7..687f544e874 100644 --- a/api/core/rag/splitter/fixed_text_splitter.py +++ b/api/core/rag/splitter/fixed_text_splitter.py @@ -91,8 +91,8 @@ class FixedRecursiveCharacterTextSplitter(EnhanceRecursiveCharacterTextSplitter) splits = re.split(r" +", text) else: splits = text.split(separator) - if self._keep_separator: - splits = [s + separator for s in splits[:-1]] + splits[-1:] + if self._keep_separator: + splits = [s + separator for s in splits[:-1]] + splits[-1:] else: splits = list(text) if separator == "\n": diff --git a/api/core/repositories/celery_workflow_node_execution_repository.py b/api/core/repositories/celery_workflow_node_execution_repository.py index 3b152a0e232..c393f6afd51 100644 --- a/api/core/repositories/celery_workflow_node_execution_repository.py +++ b/api/core/repositories/celery_workflow_node_execution_repository.py @@ -16,6 +16,9 @@ from core.repositories.factory import ( OrderConfig, WorkflowNodeExecutionRepository, ) +from core.repositories.sqlalchemy_workflow_node_execution_repository import ( + SQLAlchemyWorkflowNodeExecutionRepository, +) from graphon.entities import WorkflowNodeExecution from models import Account, CreatorUserRole, EndUser from models.workflow import WorkflowNodeExecutionTriggeredFrom @@ -36,7 +39,7 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository): Key features: - Asynchronous save operations using Celery tasks - - In-memory cache for immediate reads + - In-memory cache for immediate reads with database backfill across Celery tasks - Support for multi-tenancy through tenant/app filtering - Automatic retry and error handling through Celery """ @@ -49,6 +52,8 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository): _creator_user_role: CreatorUserRole _execution_cache: dict[str, WorkflowNodeExecution] _workflow_execution_mapping: dict[str, list[str]] + _database_loaded_workflow_executions: set[str] + _sql_repository: SQLAlchemyWorkflowNodeExecutionRepository def __init__( self, @@ -98,6 +103,14 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository): # Cache for mapping workflow_execution_ids to execution IDs for efficient retrieval self._workflow_execution_mapping = {} + self._database_loaded_workflow_executions = set() + self._sql_repository = SQLAlchemyWorkflowNodeExecutionRepository( + session_factory=self._session_factory, + tenant_id=tenant_id, + user=user, + app_id=app_id, + triggered_from=triggered_from, + ) logger.info( "Initialized CeleryWorkflowNodeExecutionRepository for tenant %s, app %s, triggered_from %s", @@ -149,6 +162,17 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository): # For now, we'll re-raise the exception raise + @override + def save_synchronously(self, execution: WorkflowNodeExecution) -> None: + """Create the Agent v2 caller row before runtime participant allocation.""" + + self._sql_repository.save_synchronously(execution) + self._execution_cache[execution.id] = execution + if execution.workflow_execution_id: + execution_ids = self._workflow_execution_mapping.setdefault(execution.workflow_execution_id, []) + if execution.id not in execution_ids: + execution_ids.append(execution.id) + @override def get_by_workflow_execution( self, @@ -156,7 +180,7 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository): order_config: OrderConfig | None = None, ) -> Sequence[WorkflowNodeExecution]: """ - Retrieve all workflow node executions for a workflow execution from cache. + Retrieve workflow node executions from cache after loading persisted history once. Args: workflow_execution_id: The workflow execution identifier @@ -166,6 +190,25 @@ class CeleryWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository): A sequence of WorkflowNodeExecution instances """ try: + if workflow_execution_id not in self._database_loaded_workflow_executions: + try: + persisted_executions = self._sql_repository.get_by_workflow_execution( + workflow_execution_id, + order_config, + ) + except Exception: + logger.exception( + "Failed to load persisted workflow node executions for execution %s", + workflow_execution_id, + ) + else: + execution_ids = self._workflow_execution_mapping.setdefault(workflow_execution_id, []) + for execution in persisted_executions: + self._execution_cache.setdefault(execution.id, execution) + if execution.id not in execution_ids: + execution_ids.append(execution.id) + self._database_loaded_workflow_executions.add(workflow_execution_id) + # Get execution IDs for this workflow execution from cache execution_ids = self._workflow_execution_mapping.get(workflow_execution_id, []) diff --git a/api/core/repositories/factory.py b/api/core/repositories/factory.py index 8a37ba933fc..3e8b5277402 100644 --- a/api/core/repositories/factory.py +++ b/api/core/repositories/factory.py @@ -35,6 +35,8 @@ class WorkflowExecutionRepository(Protocol): class WorkflowNodeExecutionRepository(Protocol): def save(self, execution: WorkflowNodeExecution): ... + def save_synchronously(self, execution: WorkflowNodeExecution) -> None: ... + def save_execution_data(self, execution: WorkflowNodeExecution): ... def get_by_workflow_execution( diff --git a/api/core/repositories/human_input_repository.py b/api/core/repositories/human_input_repository.py index ac3aa46fc21..b144d1f88d4 100644 --- a/api/core/repositories/human_input_repository.py +++ b/api/core/repositories/human_input_repository.py @@ -67,6 +67,7 @@ class FormCreateParams: # workflow_execution_id for chatflow runs; set alone (workflow_execution_id None) # for Agent v2 chat ask_human forms, which have no workflow run. conversation_id: str | None = None + form_id: str | None = None class HumanInputFormRecipientEntity(Protocol): @@ -110,7 +111,7 @@ class HumanInputFormEntity(Protocol): class HumanInputFormRepository(Protocol): - def get_form(self, node_id: str) -> HumanInputFormEntity | None: ... + def get_form(self, node_id: str, *, form_id: str | None = None) -> HumanInputFormEntity | None: ... def create_form(self, params: FormCreateParams) -> HumanInputFormEntity: ... @@ -460,8 +461,7 @@ class HumanInputFormRepositoryImpl: raise ValueError("a runtime human input form requires a workflow_execution_id or conversation_id") with session_factory.create_session() as session, session.begin(): - # Generate unique form ID - form_id = str(uuidv7()) + form_id = params.form_id or str(uuidv7()) start_time = naive_utc_now() node_expiration = form_config.expiration_time(start_time) form_definition = FormDefinition( @@ -546,7 +546,7 @@ class HumanInputFormRepositoryImpl: return _HumanInputFormEntityImpl(form_model=form_model, recipient_models=recipient_models) - def get_form(self, node_id: str) -> HumanInputFormEntity | None: + def get_form(self, node_id: str, *, form_id: str | None = None) -> HumanInputFormEntity | None: if self._workflow_execution_id is None: raise ValueError("workflow_execution_id is required to load runtime human input forms") @@ -555,6 +555,8 @@ class HumanInputFormRepositoryImpl: HumanInputForm.node_id == node_id, HumanInputForm.tenant_id == self._tenant_id, ) + if form_id is not None: + form_query = form_query.where(HumanInputForm.id == form_id) with session_factory.create_session() as session: form_model: HumanInputForm | None = session.scalars(form_query).first() if form_model is None: diff --git a/api/core/repositories/sqlalchemy_workflow_node_execution_repository.py b/api/core/repositories/sqlalchemy_workflow_node_execution_repository.py index 15dca1ff8cf..2fe5f427e5d 100644 --- a/api/core/repositories/sqlalchemy_workflow_node_execution_repository.py +++ b/api/core/repositories/sqlalchemy_workflow_node_execution_repository.py @@ -18,6 +18,7 @@ from tenacity import before_sleep_log, retry, retry_if_exception, stop_after_att from configs import dify_config from core.repositories.factory import OrderConfig, WorkflowNodeExecutionRepository +from core.workflow.node_execution_process_data import preserve_workflow_agent_binding_id from extensions.ext_storage import storage from graphon.entities import WorkflowNodeExecution from graphon.enums import WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus @@ -372,6 +373,12 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository) logger.exception("Failed to save workflow node execution after all retries") raise + @override + def save_synchronously(self, execution: WorkflowNodeExecution) -> None: + """Persist a caller row before an Agent v2 participant is materialized.""" + + self.save(execution) + def _persist_to_database(self, db_model: WorkflowNodeExecutionModel): """ Persist the database model to the database. @@ -386,6 +393,13 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository) existing = session.get(WorkflowNodeExecutionModel, db_model.id) if existing: + merged_process_data = preserve_workflow_agent_binding_id( + existing.process_data_dict, + db_model.process_data_dict, + ) + db_model.process_data = ( + _deterministic_json_dump(merged_process_data) if merged_process_data is not None else None + ) # Update existing record by copying all non-private attributes for key, value in db_model.__dict__.items(): if not key.startswith("_"): @@ -442,18 +456,25 @@ class SQLAlchemyWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository) else: db_model.outputs = self._json_encode(domain_model.outputs) - if domain_model.process_data is not None: + process_data = preserve_workflow_agent_binding_id(db_model.process_data_dict, domain_model.process_data) + if process_data is not None: result = self._truncate_and_upload( - domain_model.process_data, + process_data, domain_model.id, ExecutionOffLoadType.PROCESS_DATA, ) if result is not None: - db_model.process_data = self._json_encode(result.truncated_value) - domain_model.set_truncated_process_data(result.truncated_value) + truncated_process_data = preserve_workflow_agent_binding_id( + process_data, + result.truncated_value, + ) + if truncated_process_data is None: + raise ValueError("truncated process data is unavailable") + db_model.process_data = self._json_encode(truncated_process_data) + domain_model.set_truncated_process_data(truncated_process_data) offload_data = _replace_or_append_offload(offload_data, result.offload) else: - db_model.process_data = self._json_encode(domain_model.process_data) + db_model.process_data = self._json_encode(process_data) db_model.offload_data = offload_data with self._session_factory() as session, session.begin(): diff --git a/api/core/tools/builtin_tool/providers/audio/audio.yaml b/api/core/tools/builtin_tool/providers/audio/audio.yaml index 07db268dacc..31e72cd8433 100644 --- a/api/core/tools/builtin_tool/providers/audio/audio.yaml +++ b/api/core/tools/builtin_tool/providers/audio/audio.yaml @@ -3,6 +3,7 @@ identity: name: audio label: en_US: Audio + zh_Hans: 音频 description: en_US: A tool for tts and asr. zh_Hans: 一个用于文本转语音和语音转文本的工具。 diff --git a/api/core/tools/builtin_tool/providers/audio/tools/asr.yaml b/api/core/tools/builtin_tool/providers/audio/tools/asr.yaml index b2c82f80863..3b384c39f13 100644 --- a/api/core/tools/builtin_tool/providers/audio/tools/asr.yaml +++ b/api/core/tools/builtin_tool/providers/audio/tools/asr.yaml @@ -3,6 +3,7 @@ identity: author: hjlarry label: en_US: Speech To Text + zh_Hans: 语音转文本 description: human: en_US: Convert audio file to text. diff --git a/api/core/tools/builtin_tool/providers/audio/tools/tts.yaml b/api/core/tools/builtin_tool/providers/audio/tools/tts.yaml index 36f42bd689f..5f78b79266d 100644 --- a/api/core/tools/builtin_tool/providers/audio/tools/tts.yaml +++ b/api/core/tools/builtin_tool/providers/audio/tools/tts.yaml @@ -3,6 +3,7 @@ identity: author: hjlarry label: en_US: Text To Speech + zh_Hans: 文本转语音 description: human: en_US: Convert text to audio file. diff --git a/api/core/tools/entities/tool_entities.py b/api/core/tools/entities/tool_entities.py index 786910d91d4..cc8123f0adb 100644 --- a/api/core/tools/entities/tool_entities.py +++ b/api/core/tools/entities/tool_entities.py @@ -311,6 +311,8 @@ class ToolParameter(PluginParameter): MODEL_SELECTOR = PluginParameterType.MODEL_SELECTOR ANY = PluginParameterType.ANY DYNAMIC_SELECT = PluginParameterType.DYNAMIC_SELECT + DATE = PluginParameterType.DATE + DATE_RANGE = PluginParameterType.DATE_RANGE # MCP object and array type parameters ARRAY = MCPServerParameterType.ARRAY diff --git a/api/core/tools/mcp_tool/tool.py b/api/core/tools/mcp_tool/tool.py index 07c9ff63b30..899ffe56e6b 100644 --- a/api/core/tools/mcp_tool/tool.py +++ b/api/core/tools/mcp_tool/tool.py @@ -25,6 +25,7 @@ from core.tools.__base.tool import Tool from core.tools.__base.tool_runtime import ToolRuntime from core.tools.entities.tool_entities import ToolEntity, ToolInvokeMessage, ToolProviderType from core.tools.errors import ToolInvokeError +from enums import DeploymentEdition from graphon.model_runtime.entities.llm_entities import LLMUsage, LLMUsageMetadata logger = logging.getLogger(__name__) @@ -107,6 +108,8 @@ class MCPTool(Tool): if self.entity.output_schema and result.structuredContent: for k, v in result.structuredContent.items(): yield self.create_variable_message(k, v) + elif result.structuredContent: + yield self.create_json_message(result.structuredContent) def _process_text_content(self, content: TextContent) -> Generator[ToolInvokeMessage, None, None]: """Process text content and yield appropriate messages.""" @@ -270,7 +273,7 @@ class MCPTool(Tool): the deployment actually has the enterprise side that can mint tokens. Non-enterprise installs treat the DB value as a no-op — a stale row won't trigger a 5xx against a missing inner-API endpoint.""" - return self.identity_mode != IdentityMode.OFF and dify_config.ENTERPRISE_ENABLED + return self.identity_mode != IdentityMode.OFF and dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE def invoke_remote_mcp_tool( self, diff --git a/api/core/tools/signature.py b/api/core/tools/signature.py index 77d81bafbc1..725160aaf8a 100644 --- a/api/core/tools/signature.py +++ b/api/core/tools/signature.py @@ -4,29 +4,53 @@ import hmac import os import time import urllib.parse +from typing import Literal +from urllib.parse import urlsplit from configs import dify_config +def bind_file_uri(uri: str, base_url: str) -> str: + """Bind a Dify-owned file URI to one caller-selected origin. + + Explicit remote HTTP(S) URLs are already complete and pass through. Other + values must be origin-free ``/files/...`` URIs. + """ + + parsed = urlsplit(uri) + if parsed.scheme in {"http", "https"} and parsed.netloc: + return uri + if ( + parsed.scheme + or parsed.netloc + or parsed.fragment + or uri.startswith("//") + or not parsed.path.startswith("/files/") + ): + raise ValueError("file URI must be an absolute HTTP(S) URL or a /files/ URI") + return f"{base_url}{uri}" + + def _secret_key() -> bytes: return dify_config.SECRET_KEY.encode() -def sign_tool_file(tool_file_id: str, extension: str, for_external: bool = True) -> str: - """ - sign file to get a temporary url for plugin access - """ - # Use internal URL for plugin/tool file access in Docker environments, unless for_external is True - base_url = dify_config.FILES_URL if for_external else (dify_config.INTERNAL_FILES_URL or dify_config.FILES_URL) - file_preview_url = f"{base_url}/files/tools/{tool_file_id}{extension}" - +def sign_tool_file_uri(tool_file_id: str, extension: str) -> str: + """Sign a ToolFile path without selecting a network origin.""" timestamp = str(int(time.time())) nonce = os.urandom(16).hex() data_to_sign = f"file-preview|{tool_file_id}|{timestamp}|{nonce}" sign = hmac.new(_secret_key(), data_to_sign.encode(), hashlib.sha256).digest() encoded_sign = base64.urlsafe_b64encode(sign).decode() - return f"{file_preview_url}?timestamp={timestamp}&nonce={nonce}&sign={encoded_sign}" + return f"/files/tools/{tool_file_id}{extension}?timestamp={timestamp}&nonce={nonce}&sign={encoded_sign}" + + +def sign_tool_file(tool_file_id: str, extension: str, for_external: bool = True) -> str: + """Sign a ToolFile URL for the browser or an internal Dify service.""" + + base_url = dify_config.FILES_URL if for_external else (dify_config.INTERNAL_FILES_URL or dify_config.FILES_URL) + return bind_file_uri(sign_tool_file_uri(tool_file_id, extension), base_url) def sign_upload_file_preview_url(upload_file_id: str, extension: str) -> str: @@ -64,16 +88,28 @@ def verify_tool_file_signature(file_id: str, timestamp: str, nonce: str, sign: s return current_time - int(timestamp) <= dify_config.FILES_ACCESS_TIMEOUT -def get_signed_file_url_for_plugin( - filename: str, mimetype: str, tenant_id: str, user_id: str, conversation_id: str | None = None +def get_signed_file_uri_for_plugin( + filename: str, + mimetype: str, + tenant_id: str, + user_id: str, + conversation_id: str | None = None, + user_from: Literal["account", "end-user"] | None = None, ) -> str: - """Build the signed upload URL used by the plugin-facing file upload endpoint.""" + """Build a signed plugin-upload URI without selecting a network origin.""" - base_url = dify_config.INTERNAL_FILES_URL or dify_config.FILES_URL - upload_url = f"{base_url}/files/upload/for-plugin" timestamp = str(int(time.time())) nonce = os.urandom(16).hex() - data_to_sign = f"upload|{filename}|{mimetype}|{tenant_id}|{user_id}|{conversation_id or ''}|{timestamp}|{nonce}" + data_to_sign = _plugin_upload_signature_payload( + filename=filename, + mimetype=mimetype, + tenant_id=tenant_id, + user_id=user_id, + conversation_id=conversation_id, + timestamp=timestamp, + nonce=nonce, + user_from=user_from, + ) sign = hmac.new(_secret_key(), data_to_sign.encode(), hashlib.sha256).digest() encoded_sign = base64.urlsafe_b64encode(sign).decode() query_params = { @@ -85,8 +121,10 @@ def get_signed_file_url_for_plugin( } if conversation_id: query_params["conversation_id"] = conversation_id + if user_from is not None: + query_params["user_from"] = user_from query = urllib.parse.urlencode(query_params) - return f"{upload_url}?{query}" + return f"/files/upload/for-plugin?{query}" def verify_plugin_file_signature( @@ -96,13 +134,23 @@ def verify_plugin_file_signature( tenant_id: str, user_id: str, conversation_id: str | None = None, + user_from: Literal["account", "end-user"] | None = None, timestamp: str, nonce: str, sign: str, ) -> bool: """Verify the signature used by the plugin-facing file upload endpoint.""" - data_to_sign = f"upload|{filename}|{mimetype}|{tenant_id}|{user_id}|{conversation_id or ''}|{timestamp}|{nonce}" + data_to_sign = _plugin_upload_signature_payload( + filename=filename, + mimetype=mimetype, + tenant_id=tenant_id, + user_id=user_id, + conversation_id=conversation_id, + timestamp=timestamp, + nonce=nonce, + user_from=user_from, + ) recalculated_sign = hmac.new(_secret_key(), data_to_sign.encode(), hashlib.sha256).digest() recalculated_encoded_sign = base64.urlsafe_b64encode(recalculated_sign).decode() @@ -111,3 +159,27 @@ def verify_plugin_file_signature( current_time = int(time.time()) return current_time - int(timestamp) <= dify_config.FILES_ACCESS_TIMEOUT + + +def _plugin_upload_signature_payload( + *, + filename: str, + mimetype: str, + tenant_id: str, + user_id: str, + conversation_id: str | None, + timestamp: str, + nonce: str, + user_from: Literal["account", "end-user"] | None, +) -> str: + """Build the compatible upload signature payload with optional identity ownership. + + Omitting ``user_from`` preserves the legacy payload. When present, the + identity kind is appended and HMAC-protected so account/end-user ownership + cannot be altered. + """ + + payload = f"upload|{filename}|{mimetype}|{tenant_id}|{user_id}|{conversation_id or ''}|{timestamp}|{nonce}" + if user_from is not None: + payload = f"{payload}|{user_from}" + return payload diff --git a/api/core/tools/tool_manager.py b/api/core/tools/tool_manager.py index fc85f20bdd4..0f33b104914 100644 --- a/api/core/tools/tool_manager.py +++ b/api/core/tools/tool_manager.py @@ -55,6 +55,7 @@ from core.tools.workflow_as_tool.provider import WorkflowToolProviderController from core.tools.workflow_as_tool.tool import WorkflowTool from extensions.ext_database import db from graphon.runtime import VariablePool +from graphon.variables.template_resolution import convert_template from models.provider_ids import ToolProviderID from models.tools import ApiToolProvider, BuiltinToolProvider, WorkflowToolProvider from services.tools.mcp_tools_manage_service import MCPToolManageService @@ -1113,7 +1114,7 @@ class ToolManager: elif tool_input.type == "constant": parameter_value = tool_input.value elif tool_input.type == "mixed": - segment_group = variable_pool.convert_template(str(tool_input.value)) + segment_group = convert_template(variable_pool, str(tool_input.value)) parameter_value = segment_group.text else: raise ToolParameterError(f"Unknown tool input type '{tool_input.type}'") diff --git a/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py b/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py index 8f07310f3c7..489170053a8 100644 --- a/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py +++ b/api/core/tools/utils/dataset_retriever/dataset_retriever_tool.py @@ -123,7 +123,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): for hit_callback in self.hit_callbacks: hit_callback.return_retriever_resource_info(context_list) - return str("\n".join([item.page_content for item in results])) + return "\n".join([item.page_content for item in results]) else: if metadata_condition and not document_ids_filter: return "" @@ -139,7 +139,7 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): top_k=self.top_k, document_ids_filter=document_ids_filter, ) - return str("\n".join([document.page_content for document in documents])) + return "\n".join([document.page_content for document in documents]) else: if self.top_k > 0: # retrieval source @@ -241,5 +241,5 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool): hit_callback.return_retriever_resource_info(retrieval_resource_list) if document_context_list: document_context_list = sorted(document_context_list, key=lambda x: x.score or 0.0, reverse=True) - return str("\n".join([document_context.content for document_context in document_context_list])) + return "\n".join([document_context.content for document_context in document_context_list]) return "" diff --git a/api/core/tools/utils/web_reader_tool.py b/api/core/tools/utils/web_reader_tool.py index c156fd888f7..2e630b803eb 100644 --- a/api/core/tools/utils/web_reader_tool.py +++ b/api/core/tools/utils/web_reader_tool.py @@ -2,7 +2,7 @@ import mimetypes import re from collections.abc import Sequence from dataclasses import dataclass -from typing import Any, cast +from typing import Any from urllib.parse import unquote import charset_normalizer @@ -58,7 +58,7 @@ def get_url(url: str, user_agent: str | None = None) -> str: return f"Unsupported content-type [{main_content_type}] of URL." if main_content_type in extract_processor.SUPPORT_URL_CONTENT_TYPES: - return cast(str, ExtractProcessor.load_from_url(url, return_text=True)) + return ExtractProcessor.load_from_url(url, return_text=True) response = remote_fetcher.make_request("GET", url, headers=headers, follow_redirects=True, timeout=(120, 300)) elif response.status_code == 403: diff --git a/api/core/workflow/human_input_adapter.py b/api/core/workflow/human_input_adapter.py index 0865365ea68..52021f9f49b 100644 --- a/api/core/workflow/human_input_adapter.py +++ b/api/core/workflow/human_input_adapter.py @@ -21,6 +21,7 @@ from graphon.enums import BuiltinNodeTypes from graphon.nodes.base.variable_template_parser import VariableTemplateParser from graphon.runtime import VariablePool from graphon.variables.consts import SELECTORS_LENGTH +from graphon.variables.template_resolution import convert_template class DeliveryMethodType(enum.StrEnum): @@ -116,7 +117,7 @@ class EmailDeliveryConfig(BaseModel): templated_body = cls.replace_url_placeholder(body, url) if variable_pool is None: return templated_body - return variable_pool.convert_template(templated_body).text + return convert_template(variable_pool, templated_body).text @classmethod def render_markdown_body(cls, body: str) -> str: diff --git a/api/core/workflow/llm_environment_variable.py b/api/core/workflow/llm_environment_variable.py new file mode 100644 index 00000000000..aacca73a6d0 --- /dev/null +++ b/api/core/workflow/llm_environment_variable.py @@ -0,0 +1,166 @@ +"""LLM environment-variable types and model-reference resolution helpers. + +Graphon does not have an LLM segment type, so the runtime representation stays +an object variable. Workflow persistence and API boundaries use the semantic +``llm`` value type through the serialization helpers in this module. +""" + +from collections.abc import Mapping, Sequence +from typing import Any + +from pydantic import BaseModel, ConfigDict, Field, field_validator + +from core.workflow.variable_prefixes import ENVIRONMENT_VARIABLE_NODE_ID +from graphon.nodes.llm.entities import LLMMode, ModelConfig +from graphon.variables.variables import ObjectVariable, VariableBase + +LLM_ENVIRONMENT_VARIABLE_VALUE_TYPE = "llm" + + +class LLMModelSelection(BaseModel): + """The shared model configuration stored by an LLM environment variable.""" + + model_config = ConfigDict(extra="forbid", frozen=True, str_strip_whitespace=True) + + provider: str = Field(min_length=1) + name: str = Field(min_length=1) + mode: LLMMode + completion_params: dict[str, Any] | None = None + + +class LLMEnvironmentVariable(ObjectVariable): + """An object variable with a strictly validated LLM model-selection value.""" + + @field_validator("value", mode="before") + @classmethod + def validate_model_selection(cls, value: object) -> dict[str, Any]: + return LLMModelSelection.model_validate(value).model_dump(mode="json", exclude_none=True) + + +def dump_environment_variable(variable: VariableBase, *, mode: str = "python") -> dict[str, Any]: + """Serialize a variable while preserving its workflow-level semantic type.""" + + result = variable.model_dump(mode=mode) + if isinstance(variable, LLMEnvironmentVariable): + result["value_type"] = LLM_ENVIRONMENT_VARIABLE_VALUE_TYPE + return result + + +def environment_variable_value_type(variable: VariableBase) -> str: + """Return the public value type used by workflow environment-variable APIs.""" + + if isinstance(variable, LLMEnvironmentVariable): + return LLM_ENVIRONMENT_VARIABLE_VALUE_TYPE + return str(variable.value_type.exposed_type()) + + +def parse_llm_model_selector(selector: object) -> tuple[str, str]: + """Validate and normalize an LLM node's environment-variable selector.""" + + if ( + not isinstance(selector, Sequence) + or isinstance(selector, str | bytes) + or len(selector) != 2 + or selector[0] != ENVIRONMENT_VARIABLE_NODE_ID + or not isinstance(selector[1], str) + or not selector[1] + ): + raise ValueError("LLM model selector must have the form ['env', '']") + return ENVIRONMENT_VARIABLE_NODE_ID, selector[1] + + +def should_resolve_llm_model_selector(selector: object) -> bool: + """Return whether a selector uses the LLM environment-variable contract. + + Older snippet graphs can contain selectors such as ``["start", "MODEL"]``. + They predate this contract and must keep using the node's static model. + Malformed selectors and selectors beginning with ``env`` are handled by the + strict parser so new invalid references still fail fast. + """ + + if selector is None: + return False + if ( + isinstance(selector, Sequence) + and not isinstance(selector, str | bytes) + and len(selector) > 0 + and selector[0] != ENVIRONMENT_VARIABLE_NODE_ID + ): + return False + return True + + +def resolve_llm_model_config( + *, + node_model: ModelConfig, + variable_name: str, + variable_value: Mapping[str, Any], +) -> ModelConfig: + """Apply a shared model configuration, preserving legacy node-local parameters.""" + + try: + selection = LLMModelSelection.model_validate(variable_value) + except ValueError as exc: + raise ValueError(f"LLM environment variable '{variable_name}' has an invalid model selection: {exc}") from exc + + if selection.mode != node_model.mode: + raise ValueError( + f"LLM environment variable '{variable_name}' uses mode '{selection.mode.value}', " + f"but the referencing node uses mode '{node_model.mode.value}'" + ) + + update: dict[str, Any] = { + "provider": selection.provider, + "name": selection.name, + } + if selection.completion_params is not None: + update["completion_params"] = selection.completion_params + return node_model.model_copy(update=update) + + +def resolve_llm_model_config_from_environment( + *, + node_model: ModelConfig, + selector: object, + environment_variables: Mapping[str, VariableBase], +) -> ModelConfig: + """Resolve a persisted LLM environment reference against workflow variables.""" + + _, variable_name = parse_llm_model_selector(selector) + variable = environment_variables.get(variable_name) + if not isinstance(variable, LLMEnvironmentVariable): + raise ValueError(f"LLM environment variable '{variable_name}' was not found or is not an LLM variable") + return resolve_llm_model_config( + node_model=node_model, + variable_name=variable_name, + variable_value=variable.value, + ) + + +def validate_llm_environment_model_references( + *, + graph: Mapping[str, Any], + environment_variables: Sequence[VariableBase], +) -> None: + """Validate every LLM environment-model reference in a workflow graph.""" + + variables_by_name = {variable.name: variable for variable in environment_variables} + nodes = graph.get("nodes", []) + if not isinstance(nodes, Sequence) or isinstance(nodes, str | bytes): + return + + for node in nodes: + if not isinstance(node, Mapping): + continue + node_data = node.get("data") + if not isinstance(node_data, Mapping) or node_data.get("type") != "llm": + continue + selector = node_data.get("model_selector") + if not should_resolve_llm_model_selector(selector): + continue + model = ModelConfig.model_validate(node_data.get("model", {})) + resolve_llm_model_config_from_environment( + node_model=model, + selector=selector, + environment_variables=variables_by_name, + ) diff --git a/api/core/workflow/llm_node.py b/api/core/workflow/llm_node.py new file mode 100644 index 00000000000..1f1c863ab5a --- /dev/null +++ b/api/core/workflow/llm_node.py @@ -0,0 +1,44 @@ +from collections.abc import Callable, Generator, Sequence +from typing import Any, override + +from graphon.model_runtime.entities.llm_entities import LLMStructuredOutput +from graphon.model_runtime.entities.message_entities import PromptMessage +from graphon.node_events.base import NodeEventBase +from graphon.nodes.llm.node import LLMNode +from graphon.nodes.llm.runtime_protocols import LLMPollingCapableProtocol + + +# TODO: Remove this Dify-specific node once graphon exposes a polling finalization hook. +class DifyLLMNode(LLMNode): + """Dify-owned LLM node lifecycle extensions.""" + + @classmethod + @override + def version(cls) -> str: + return "1" + + def __init__( + self, + *args: Any, + polling_finalizer: Callable[[], None], + **kwargs: Any, + ) -> None: + super().__init__(*args, **kwargs) + self._polling_finalizer = polling_finalizer + + @override + def _invoke_llm_with_polling( + self, + *, + polling_model: LLMPollingCapableProtocol, + prompt_messages: Sequence[PromptMessage], + stop: Sequence[str] | None, + ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + try: + yield from super()._invoke_llm_with_polling( + polling_model=polling_model, + prompt_messages=prompt_messages, + stop=stop, + ) + finally: + self._polling_finalizer() diff --git a/api/core/workflow/node_execution_process_data.py b/api/core/workflow/node_execution_process_data.py new file mode 100644 index 00000000000..4133feb1b21 --- /dev/null +++ b/api/core/workflow/node_execution_process_data.py @@ -0,0 +1,27 @@ +from collections.abc import Mapping +from typing import Any + +WORKFLOW_AGENT_BINDING_ID_KEY = "workflow_agent_binding_id" + + +def preserve_workflow_agent_binding_id( + identity_source: Mapping[str, Any] | None, + process_data: Mapping[str, Any] | None, +) -> dict[str, Any] | None: + source_id = (identity_source or {}).get(WORKFLOW_AGENT_BINDING_ID_KEY) + target_id = (process_data or {}).get(WORKFLOW_AGENT_BINDING_ID_KEY) + for value in (source_id, target_id): + if value is not None and not isinstance(value, str): + raise ValueError("workflow_agent_binding_id must be a string") + if source_id is not None and target_id is not None and source_id != target_id: + raise ValueError("workflow_agent_binding_id does not match") + + if process_data is None and source_id is None: + return None + merged = dict(process_data or {}) + if source_id is not None: + merged[WORKFLOW_AGENT_BINDING_ID_KEY] = source_id + return merged + + +__all__ = ["WORKFLOW_AGENT_BINDING_ID_KEY", "preserve_workflow_agent_binding_id"] diff --git a/api/core/workflow/node_factory.py b/api/core/workflow/node_factory.py index 3b47e32adf1..f73dd47e7bc 100644 --- a/api/core/workflow/node_factory.py +++ b/api/core/workflow/node_factory.py @@ -22,6 +22,12 @@ from core.model_manager import ModelInstance from core.prompt.entities.advanced_prompt_entities import MemoryConfig from core.trigger.constants import TRIGGER_NODE_TYPES from core.workflow.human_input_adapter import adapt_node_config_for_graph +from core.workflow.llm_environment_variable import ( + parse_llm_model_selector, + resolve_llm_model_config, + should_resolve_llm_model_selector, +) +from core.workflow.llm_node import DifyLLMNode from core.workflow.node_runtime import ( DifyFileReferenceFactory, DifyHumanInputNodeRuntime, @@ -63,7 +69,7 @@ from graphon.nodes.http_request import build_http_request_config from graphon.nodes.llm.entities import LLMNodeData from graphon.nodes.parameter_extractor.entities import ParameterExtractorNodeData from graphon.nodes.question_classifier.entities import QuestionClassifierNodeData -from graphon.variables.segments import ArrayObjectSegment +from graphon.variables.segments import ArrayObjectSegment, ObjectSegment from models.model import Conversation if TYPE_CHECKING: @@ -361,6 +367,12 @@ class DifyNodeFactory(NodeFactory): self._agent_runtime_support = AgentRuntimeSupport() self._agent_message_transformer = AgentMessageTransformer() + def with_runtime_state(self, graph_runtime_state: "GraphRuntimeState") -> "DifyNodeFactory": + return DifyNodeFactory( + graph_init_params=self.graph_init_params, + graph_runtime_state=graph_runtime_state, + ) + @staticmethod def _resolve_dify_context(run_context: Mapping[str, Any]) -> DifyRunContext: raw_ctx = run_context.get(DIFY_RUN_CONTEXT_KEY) @@ -394,6 +406,9 @@ class DifyNodeFactory(NodeFactory): # stay explicit and constructors receive the concrete typed payload. resolved_node_data = self._validate_resolved_node_data(node_class, node_data) node_type = node_data.type + if node_type == BuiltinNodeTypes.LLM: + resolved_node_data = self._resolve_llm_model_reference(cast(LLMNodeData, resolved_node_data)) + node: Node | None = None node_init_kwargs_factories: Mapping[NodeType, Callable[[], dict[str, object]]] = { BuiltinNodeTypes.CODE: lambda: { "code_executor": self._code_executor, @@ -412,7 +427,8 @@ class DifyNodeFactory(NodeFactory): }, BuiltinNodeTypes.HUMAN_INPUT: lambda: { "hitl_callback": self._build_human_input_callback( - node_data=DifyHumanInputNodeData.model_validate(adapted_node_config["data"]) + node_data=DifyHumanInputNodeData.model_validate(adapted_node_config["data"]), + execution_id_getter=lambda: node.execution_id if node is not None else None, ), }, BuiltinNodeTypes.LLM: lambda: self._build_llm_compatible_node_init_kwargs( @@ -457,13 +473,14 @@ class DifyNodeFactory(NodeFactory): } node_init_kwargs = node_init_kwargs_factories.get(node_type, lambda: {})() constructor_node_data = resolved_node_data.model_dump(mode="python", by_alias=True) - return node_class( + node = node_class( node_id=node_id, data=constructor_node_data, graph_init_params=self.graph_init_params, graph_runtime_state=self.graph_runtime_state, **node_init_kwargs, ) + return node @staticmethod def _validate_resolved_node_data(node_class: type[Node], node_data: BaseNodeData) -> BaseNodeData: @@ -477,8 +494,29 @@ class DifyNodeFactory(NodeFactory): @staticmethod def _resolve_node_class(*, node_type: NodeType, node_version: str) -> type[Node]: + if node_type == BuiltinNodeTypes.LLM: + return DifyLLMNode return resolve_workflow_node_class(node_type=node_type, node_version=node_version) + def _resolve_llm_model_reference(self, node_data: LLMNodeData) -> LLMNodeData: + """Resolve an optional shared model selector from the workflow variable pool.""" + + model_selector = (node_data.model_extra or {}).get("model_selector") + if not should_resolve_llm_model_selector(model_selector): + return node_data + + selector = parse_llm_model_selector(model_selector) + variable = self.graph_runtime_state.variable_pool.get(selector) + if not isinstance(variable, ObjectSegment): + raise ValueError(f"LLM environment variable '{selector[1]}' was not found or is not an LLM variable") + + resolved_model = resolve_llm_model_config( + node_model=node_data.model, + variable_name=selector[1], + variable_value=variable.value, + ) + return node_data.model_copy(update={"model": resolved_model}) + def _build_agent_node_init_kwargs(self, *, node_class: type[Node]) -> dict[str, object]: if issubclass(node_class, DifyAgentNode): from clients.agent_backend import AgentBackendRunEventAdapter, AgentBackendRunRequestBuilder @@ -487,7 +525,7 @@ class DifyNodeFactory(NodeFactory): from core.workflow.nodes.agent_v2.output_failure_orchestrator import OutputFailureOrchestrator from core.workflow.nodes.agent_v2.output_file_rebacker import reback_tool_file_output from core.workflow.nodes.agent_v2.output_type_checker import PerOutputTypeChecker - from core.workflow.nodes.agent_v2.session_store import WorkflowAgentRuntimeSessionStore + from core.workflow.nodes.agent_v2.session_store import WorkflowAgentWorkspaceStore return { "binding_resolver": WorkflowAgentBindingResolver(), @@ -497,6 +535,7 @@ class DifyNodeFactory(NodeFactory): ), "agent_backend_client": create_agent_backend_run_client( base_url=dify_config.AGENT_BACKEND_BASE_URL, + api_token=dify_config.AGENT_BACKEND_API_TOKEN, use_fake=dify_config.AGENT_BACKEND_USE_FAKE, fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO, stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS, @@ -511,7 +550,7 @@ class DifyNodeFactory(NodeFactory): # tenant validator resolves ToolFile (canonical) + UploadFile refs. "type_checker": PerOutputTypeChecker(file_validator=AgentOutputFileTenantValidator()), "failure_orchestrator": OutputFailureOrchestrator(), - "session_store": WorkflowAgentRuntimeSessionStore(), + "session_store": WorkflowAgentWorkspaceStore(), } return { "strategy_resolver": self._agent_strategy_resolver, @@ -524,6 +563,7 @@ class DifyNodeFactory(NodeFactory): self, *, node_data: DifyHumanInputNodeData, + execution_id_getter: Callable[[], str | None], ) -> DifyHITLCallback: return DifyHITLCallback( form_repository=self._human_input_runtime.build_form_repository(), @@ -532,6 +572,7 @@ class DifyNodeFactory(NodeFactory): delivery_methods=self._human_input_runtime._resolve_delivery_methods(node_data=node_data), display_in_ui=self._human_input_runtime._display_in_ui(node_data=node_data), file_reference_factory=self._file_reference_factory, + execution_id_getter=execution_id_getter, ) def _build_llm_compatible_node_init_kwargs( @@ -548,18 +589,19 @@ class DifyNodeFactory(NodeFactory): ) -> dict[str, object]: validated_node_data = cast(LLMCompatibleNodeData, node_data) model_instance = self._build_model_instance_for_llm_node(validated_node_data) + node_model_instance = ( + self._wrap_model_instance_for_node( + node_data=validated_node_data, + model_instance=model_instance, + request_metadata={"app_id": self._dify_context.app_id}, + ) + if wrap_model_instance + else model_instance + ) node_init_kwargs: dict[str, object] = { "credentials_provider": self._llm_credentials_provider, "model_factory": self._llm_model_factory, - "model_instance": ( - self._wrap_model_instance_for_node( - node_data=validated_node_data, - model_instance=model_instance, - request_metadata={"app_id": self._dify_context.app_id}, - ) - if wrap_model_instance - else model_instance - ), + "model_instance": node_model_instance, "memory": self._build_memory_for_llm_node( node_data=validated_node_data, model_instance=model_instance, @@ -581,6 +623,7 @@ class DifyNodeFactory(NodeFactory): node_init_kwargs["jinja2_template_renderer"] = self._jinja2_template_renderer if validated_node_data.type == BuiltinNodeTypes.LLM: node_init_kwargs["default_query_selector"] = system_variable_selector(SystemVariableKey.QUERY) + node_init_kwargs["polling_finalizer"] = cast(DifyPreparedLLM, node_model_instance).finalize_llm_polling return node_init_kwargs @staticmethod diff --git a/api/core/workflow/node_runtime.py b/api/core/workflow/node_runtime.py index d0391b69c73..f113a88ccb9 100644 --- a/api/core/workflow/node_runtime.py +++ b/api/core/workflow/node_runtime.py @@ -3,7 +3,7 @@ from __future__ import annotations from collections.abc import Callable, Generator, Mapping, Sequence from dataclasses import dataclass from enum import Enum -from typing import TYPE_CHECKING, Any, Literal, cast, overload, override +from typing import TYPE_CHECKING, Any, Literal, Protocol, cast, overload, override from pydantic import JsonValue from sqlalchemy import select @@ -20,7 +20,7 @@ from core.db.session_factory import session_factory from core.helper.trace_id_helper import ParentTraceContext from core.llm_generator.output_parser.errors import OutputParserError from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output -from core.model_manager import ModelInstance +from core.model_manager import ModelInstance, QuotaManagedModelInstance from core.plugin.impl.exc import PluginDaemonClientSideError, PluginInvokeError from core.plugin.impl.plugin import PluginInstaller from core.prompt.utils.prompt_message_util import PromptMessageUtil @@ -49,6 +49,7 @@ from graphon.file import File, FileTransferMethod, FileType from graphon.model_runtime.entities import LLMMode from graphon.model_runtime.entities.llm_entities import ( LLMPollingResult, + LLMPollingStatus, LLMResult, LLMResultChunk, LLMResultChunkWithStructuredOutput, @@ -87,6 +88,33 @@ from .human_input_adapter import ( ) from .system_variables import SystemVariableKey, get_system_text + +class PollingLLMRuntimeProtocol(Protocol): + """Runtime capability required by the workflow polling adapter.""" + + def start_llm_polling( + self, + *, + provider: str, + model: str, + credentials: dict[str, Any], + model_parameters: dict[str, Any], + prompt_messages: Sequence[PromptMessage], + tools: Sequence[PromptMessageTool] | None, + stop: Sequence[str] | None, + json_schema: dict[str, Any] | None, + ) -> LLMPollingResult: ... + + def check_llm_polling( + self, + *, + provider: str, + model: str, + credentials: dict[str, Any], + plugin_state: dict[str, JsonValue], + ) -> LLMPollingResult: ... + + if TYPE_CHECKING: from core.tools.__base.tool import Tool from core.tools.entities.tool_entities import ToolInvokeMessage as CoreToolInvokeMessage @@ -281,23 +309,45 @@ class DifyPreparedLLM(LLMProtocol): def is_structured_output_parse_error(self, error: Exception) -> bool: return isinstance(error, OutputParserError) + def finalize_llm_polling(self) -> None: + """Finalize resources held by a polling invocation, if any.""" + class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol): """Prepared workflow LLM adapter that exposes Graphon's polling protocol.""" def __init__(self, model_instance: ModelInstance, request_metadata: Mapping[str, object] | None = None) -> None: - from core.plugin.impl.model_runtime import PluginModelRuntime - super().__init__(model_instance, request_metadata=request_metadata) - model_type_instance = model_instance.model_type_instance - if not isinstance(model_type_instance, LargeLanguageModel): - raise TypeError("Polling wrapper requires a large-language-model instance.") + model_type_instance = cast(LargeLanguageModel, model_instance.model_type_instance) + self._polling_runtime = cast(PollingLLMRuntimeProtocol, model_type_instance.model_runtime) + self._polling_quota_reservation = None - plugin_model_runtime = model_type_instance.model_runtime - if not isinstance(plugin_model_runtime, PluginModelRuntime): - raise TypeError("Polling wrapper requires a plugin-backed model runtime.") + @override + def finalize_llm_polling(self) -> None: + reservation = self._polling_quota_reservation + self._polling_quota_reservation = None + if reservation is not None: + QuotaManagedModelInstance.release_quota_safely(reservation) - self._plugin_model_runtime = plugin_model_runtime + def _settle_polling_quota(self, polling_result: LLMPollingResult) -> LLMPollingResult: + reservation = self._polling_quota_reservation + if reservation is None or polling_result.status == LLMPollingStatus.RUNNING: + return polling_result + + try: + if polling_result.status == LLMPollingStatus.SUCCEEDED: + if polling_result.result is None: + raise ValueError("A successful LLM polling result must include a model result.") + reservation.commit(polling_result.result.usage) + else: + reservation.release() + except Exception: + QuotaManagedModelInstance.release_quota_safely(reservation) + raise + finally: + self._polling_quota_reservation = None + + return polling_result @override def start_llm_polling( @@ -309,16 +359,26 @@ class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol): stop: Sequence[str] | None, json_schema: Mapping[str, Any] | None, ) -> LLMPollingResult: - return self._plugin_model_runtime.start_llm_polling( - provider=self.provider, - model=self.model_name, - credentials=self._model_instance.credentials, - prompt_messages=prompt_messages, - model_parameters=dict(model_parameters), - tools=tools, - stop=stop, - json_schema=dict(json_schema) if json_schema is not None else None, - ) + self.finalize_llm_polling() + + if isinstance(self._model_instance, QuotaManagedModelInstance): + self._polling_quota_reservation = self._model_instance.reserve_quota() + + try: + polling_result = self._polling_runtime.start_llm_polling( + provider=self.provider, + model=self.model_name, + credentials=self._model_instance.credentials, + prompt_messages=prompt_messages, + model_parameters=dict(model_parameters), + tools=tools, + stop=stop, + json_schema=dict(json_schema) if json_schema is not None else None, + ) + return self._settle_polling_quota(polling_result) + except Exception: + self.finalize_llm_polling() + raise @override def check_llm_polling( @@ -326,12 +386,17 @@ class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol): *, plugin_state: Mapping[str, JsonValue], ) -> LLMPollingResult: - return self._plugin_model_runtime.check_llm_polling( - provider=self.provider, - model=self.model_name, - credentials=self._model_instance.credentials, - plugin_state=dict(plugin_state), - ) + try: + polling_result = self._polling_runtime.check_llm_polling( + provider=self.provider, + model=self.model_name, + credentials=self._model_instance.credentials, + plugin_state=dict(plugin_state), + ) + return self._settle_polling_quota(polling_result) + except Exception: + self.finalize_llm_polling() + raise class DifyPromptMessageSerializer(PromptMessageSerializerProtocol): diff --git a/api/core/workflow/nodes/agent/agent_node.py b/api/core/workflow/nodes/agent/agent_node.py index 2b6745d46a9..90536f7ff82 100644 --- a/api/core/workflow/nodes/agent/agent_node.py +++ b/api/core/workflow/nodes/agent/agent_node.py @@ -1,16 +1,19 @@ from __future__ import annotations from collections.abc import Generator, Mapping, Sequence +from functools import singledispatchmethod from typing import TYPE_CHECKING, Any, override from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext from core.workflow.system_variables import SystemVariableKey, get_system_text from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus +from graphon.graph_events import GraphNodeEventBase from graphon.node_events import NodeEventBase, NodeRunResult, StreamCompletedEvent from graphon.nodes.base.node import Node from graphon.nodes.base.variable_template_parser import VariableTemplateParser from .entities import AgentNodeData +from .events import AgentLogEvent, NodeRunAgentLogEvent from .exceptions import ( AgentInvocationError, AgentMessageTransformError, @@ -71,6 +74,29 @@ class AgentNode(Node[AgentNodeData]): ), } + @override + @singledispatchmethod + def _dispatch( # pyrefly: ignore[missing-override-decorator] + self, event: NodeEventBase + ) -> GraphNodeEventBase: + return super()._dispatch(event) + + @_dispatch.register + def _dispatch_agent_log(self, event: AgentLogEvent) -> NodeRunAgentLogEvent: + return NodeRunAgentLogEvent( + id=self.execution_id, + node_id=self._node_id, + node_type=self.node_type, + message_id=event.message_id, + label=event.label, + node_execution_id=event.node_execution_id, + parent_id=event.parent_id, + error=event.error, + status=event.status, + data=event.data, + metadata=event.metadata, + ) + @override def _run(self) -> Generator[NodeEventBase, None, None]: from core.plugin.impl.exc import PluginDaemonClientSideError diff --git a/api/core/workflow/nodes/agent/entities.py b/api/core/workflow/nodes/agent/entities.py index 51452c29a3f..16e026939df 100644 --- a/api/core/workflow/nodes/agent/entities.py +++ b/api/core/workflow/nodes/agent/entities.py @@ -1,7 +1,7 @@ from enum import IntEnum, StrEnum, auto from typing import Any, Literal, Union -from pydantic import BaseModel +from pydantic import BaseModel, Field from core.prompt.entities.advanced_prompt_entities import MemoryConfig from core.tools.entities.tool_entities import ToolSelector @@ -11,9 +11,9 @@ from graphon.enums import BuiltinNodeTypes, NodeType class AgentNodeData(BaseNodeData): type: NodeType = BuiltinNodeTypes.AGENT - agent_strategy_provider_name: str - agent_strategy_name: str - agent_strategy_label: str + agent_strategy_provider_name: str = "" + agent_strategy_name: str = "" + agent_strategy_label: str = "" memory: MemoryConfig | None = None # The version of the tool parameter. # If this value is None, it indicates this is a previous version @@ -24,7 +24,7 @@ class AgentNodeData(BaseNodeData): value: Union[list[str], list[ToolSelector], Any] type: Literal["mixed", "variable", "constant"] - agent_parameters: dict[str, AgentInput] + agent_parameters: dict[str, AgentInput] = Field(default_factory=dict) class ParamsAutoGenerated(IntEnum): diff --git a/api/core/workflow/nodes/agent/events.py b/api/core/workflow/nodes/agent/events.py new file mode 100644 index 00000000000..824decb5199 --- /dev/null +++ b/api/core/workflow/nodes/agent/events.py @@ -0,0 +1,34 @@ +from collections.abc import Mapping +from typing import Any + +from pydantic import Field + +from graphon.graph_events import GraphNodeEventBase +from graphon.node_events import NodeEventBase + + +class AgentLogEvent(NodeEventBase): + message_id: str = Field(..., description="id") + label: str = Field(..., description="label") + node_execution_id: str = Field(..., description="node execution id") + parent_id: str | None = Field(..., description="parent id") + error: str | None = Field(..., description="error") + status: str = Field(..., description="status") + data: Mapping[str, Any] = Field(..., description="data") + metadata: Mapping[str, Any] = Field(default_factory=dict, description="metadata") + node_id: str = Field(..., description="node id") + + +class GraphAgentNodeEventBase(GraphNodeEventBase): + pass + + +class NodeRunAgentLogEvent(GraphAgentNodeEventBase): + message_id: str = Field(..., description="message id") + label: str = Field(..., description="label") + node_execution_id: str = Field(..., description="node execution id") + parent_id: str | None = Field(..., description="parent id") + error: str | None = Field(..., description="error") + status: str = Field(..., description="status") + data: Mapping[str, object] = Field(..., description="data") + metadata: Mapping[str, object] = Field(default_factory=dict) diff --git a/api/core/workflow/nodes/agent/message_transformer.py b/api/core/workflow/nodes/agent/message_transformer.py index f44681377dc..d3147a3581c 100644 --- a/api/core/workflow/nodes/agent/message_transformer.py +++ b/api/core/workflow/nodes/agent/message_transformer.py @@ -16,7 +16,6 @@ from graphon.file import File, FileTransferMethod, get_file_type_by_mime_type from graphon.model_runtime.entities.llm_entities import LLMUsage, LLMUsageMetadata from graphon.model_runtime.utils.encoders import jsonable_encoder from graphon.node_events import ( - AgentLogEvent, NodeEventBase, NodeRunResult, StreamChunkEvent, @@ -26,6 +25,7 @@ from graphon.variables.segments import ArrayFileSegment from models import ToolFile from services.tools.builtin_tools_manage_service import BuiltinToolManageService +from .events import AgentLogEvent from .exceptions import AgentNodeError, AgentVariableTypeError, ToolFileNotFoundError _file_access_controller = DatabaseFileAccessController() diff --git a/api/core/workflow/nodes/agent/runtime_support.py b/api/core/workflow/nodes/agent/runtime_support.py index a872774c98c..6ffccd70777 100644 --- a/api/core/workflow/nodes/agent/runtime_support.py +++ b/api/core/workflow/nodes/agent/runtime_support.py @@ -21,6 +21,7 @@ from core.workflow.system_variables import SystemVariableKey, get_system_text from extensions.ext_database import db from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelType from graphon.runtime import VariablePool +from graphon.variables.template_resolution import convert_template from models.model import Conversation from .entities import AgentNodeData, AgentOldVersionModelFeatures, ParamsAutoGenerated @@ -67,7 +68,7 @@ class AgentRuntimeSupport: except TypeError: parameter_value = str(agent_input.value) - segment_group = variable_pool.convert_template(parameter_value) + segment_group = convert_template(variable_pool, parameter_value) parameter_value = segment_group.log if for_log else segment_group.text try: if not isinstance(agent_input.value, str): @@ -198,8 +199,23 @@ class AgentRuntimeSupport: if model_schema: model_schema = self._remove_unsupported_model_features_for_old_version(model_schema) value["entity"] = model_schema.model_dump(mode="json") + # The model selector value from the workflow frontend only + # carries provider/model/mode — it does NOT include + # completion_params. AgentStrategy plugins (cot_agent, + # function_calling) read completion_params to build the + # LLMModelConfig that is backwards-invoked, and some model + # providers raise KeyError('required') when + # completion_params is empty because their parameter_rules + # declare required fields with no default. Populate + # completion_params with the defaults declared in the model + # schema so the plugin daemon always receives a valid set + # of model parameters. + if "completion_params" not in value: + value["completion_params"] = self._extract_default_completion_params(model_schema) else: value["entity"] = None + if "completion_params" not in value: + value["completion_params"] = {} result[parameter_name] = value return result @@ -275,6 +291,24 @@ class AgentRuntimeSupport: model_schema.features.remove(feature) return model_schema + @staticmethod + def _extract_default_completion_params(model_schema: AIModelEntity) -> dict[str, Any]: + """Build a completion_params dict from the model schema's parameter_rules. + + The workflow Agent node's model-selector parameter only stores + provider/model/mode — it never carries completion_params. When the + value is forwarded to the plugin daemon, AgentModelConfig defaults + completion_params to ``{}``, which causes some model providers to fail + because their parameter_rules declare required fields. This helper + collects the ``default`` value of every parameter_rule that has one so + the plugin daemon receives a valid, non-empty set of model parameters. + """ + completion_params: dict[str, Any] = {} + for rule in model_schema.parameter_rules: + if rule.default is not None: + completion_params[rule.name] = rule.default + return completion_params + @staticmethod def _filter_mcp_type_tool( strategy: ResolvedAgentStrategy, diff --git a/api/core/workflow/nodes/agent_v2/agent_node.py b/api/core/workflow/nodes/agent_v2/agent_node.py index e100b7414b3..9911ef57ad7 100644 --- a/api/core/workflow/nodes/agent_v2/agent_node.py +++ b/api/core/workflow/nodes/agent_v2/agent_node.py @@ -18,13 +18,10 @@ from clients.agent_backend import ( AgentBackendRunEventAdapter, AgentBackendRunFailedInternalEvent, AgentBackendRunSucceededInternalEvent, - AgentBackendSessionCleanupPayload, AgentBackendStreamError, AgentBackendStreamInternalEvent, AgentBackendTransportError, AgentBackendValidationError, - RuntimeLayerSpec, - extract_runtime_layer_specs, ) from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext from core.repositories.human_input_repository import HumanInputFormRepository, HumanInputFormRepositoryImpl @@ -33,11 +30,12 @@ from core.workflow.nodes.human_input.session_binding import default_session_bind from core.workflow.system_variables import SystemVariableKey, get_system_text from graphon.entities.pause_reason import HitlRequired, SchedulingPause from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus -from graphon.node_events import NodeEventBase, NodeRunResult, PauseRequestedEvent, StreamCompletedEvent +from graphon.graph_events import NodeRunPauseRequestedEvent +from graphon.node_events import NodeEventBase, NodeRunResult, StreamCompletedEvent from graphon.nodes.base.node import Node from models.agent_config_entities import AgentSoulConfig, WorkflowNodeJobConfig from services.agent.prompt_mentions import extract_workflow_node_output_selectors -from tasks.agent_backend_session_cleanup_task import cleanup_workflow_agent_runtime_session +from services.agent.workspace_service import AgentWorkspaceNotFoundError from .ask_human_hitl import AskHumanFormBuildError, build_ask_human_pause_reason from .ask_human_resume import build_deferred_tool_results, resolve_ask_human_form @@ -56,7 +54,7 @@ from .runtime_request_builder import ( WorkflowAgentRuntimeRequestBuilder, WorkflowAgentRuntimeRequestBuildError, ) -from .session_store import WorkflowAgentRuntimeSessionStore, WorkflowAgentSessionScope +from .session_store import WorkflowAgentSessionScope, WorkflowAgentWorkspaceStore if TYPE_CHECKING: from graphon.entities import GraphInitParams @@ -68,7 +66,7 @@ logger = logging.getLogger(__name__) # Stage 4 §5+§7: the terminal events that `_consume_event_stream` may return. # Stream + started events are filtered out before we yield; transport errors # are surfaced as a separate StreamCompletedEvent in the second tuple slot. -_TerminalAgentBackendEvent = ( +type _TerminalAgentBackendEvent = ( AgentBackendRunSucceededInternalEvent | AgentBackendRunFailedInternalEvent | AgentBackendRunCancelledInternalEvent @@ -93,7 +91,7 @@ class DifyAgentNode(Node[DifyAgentNodeData]): output_adapter: WorkflowAgentOutputAdapter, type_checker: PerOutputTypeChecker, failure_orchestrator: OutputFailureOrchestrator, - session_store: WorkflowAgentRuntimeSessionStore | None = None, + session_store: WorkflowAgentWorkspaceStore, ) -> None: super().__init__( node_id=node_id, @@ -130,7 +128,34 @@ class DifyAgentNode(Node[DifyAgentNodeData]): return reason @override - def _run(self) -> Generator[NodeEventBase, None, None]: + def _run(self) -> Generator[NodeEventBase | NodeRunPauseRequestedEvent, None, None]: + inputs: dict[str, Any] = {} + process_data: dict[str, Any] = {} + metadata: dict[str, Any] = { + "agent_backend": { + "status": "not_started", + } + } + try: + yield from self._run_inner(inputs=inputs, process_data=process_data, metadata=metadata) + except Exception as error: + if not process_data: + raise + yield self._failure_event( + inputs=inputs, + process_data=process_data, + metadata=metadata, + error=str(error), + error_type="agent_workflow_node_runtime_error", + ) + + def _run_inner( + self, + *, + inputs: dict[str, Any], + process_data: dict[str, Any], + metadata: dict[str, Any], + ) -> Generator[NodeEventBase | NodeRunPauseRequestedEvent, None, None]: dify_ctx = DifyRunContext.model_validate(self.require_run_context_value(DIFY_RUN_CONTEXT_KEY)) workflow_id = self.graph_init_params.workflow_id workflow_run_id = get_system_text( @@ -143,21 +168,24 @@ class DifyAgentNode(Node[DifyAgentNodeData]): self.graph_runtime_state.variable_pool, SystemVariableKey.CONVERSATION_ID, ) - inputs: dict[str, Any] = {} - process_data: dict[str, Any] = {} - metadata: dict[str, Any] = { - "agent_backend": { - "status": "not_started", - } - } # ──── Setup: resolve binding once + extract declared outputs for stage 4 checks ──── try: + existing_scope = self._session_store.load_existing_node_execution_scope( + tenant_id=dify_ctx.tenant_id, + app_id=dify_ctx.app_id, + workflow_id=workflow_id, + workflow_run_id=workflow_run_id, + node_id=self._node_id, + node_execution_id=self.execution_id, + ) bundle = self._binding_resolver.resolve( tenant_id=dify_ctx.tenant_id, app_id=dify_ctx.app_id, workflow_id=workflow_id, node_id=self._node_id, + binding_id=existing_scope.workflow_agent_binding_id if existing_scope is not None else None, + snapshot_id=existing_scope.agent_config_snapshot_id if existing_scope is not None else None, ) except WorkflowAgentBindingError as error: yield self._failure_event( @@ -168,20 +196,31 @@ class DifyAgentNode(Node[DifyAgentNodeData]): error_type=error.error_code, ) return + except AgentWorkspaceNotFoundError as error: + yield self._failure_event( + inputs=inputs, + process_data=process_data, + metadata=metadata, + error=str(error), + error_type="agent_workflow_node_runtime_error", + ) + return - process_data = { - "agent_id": bundle.agent.id, - "agent_config_snapshot_id": bundle.snapshot.id, - "binding_id": bundle.binding.id, - } - session_scope = WorkflowAgentSessionScope( + process_data.update( + { + "agent_id": bundle.agent.id, + "agent_config_snapshot_id": bundle.snapshot.id, + "workflow_agent_binding_id": bundle.binding.id, + } + ) + session_scope = existing_scope or WorkflowAgentSessionScope( tenant_id=dify_ctx.tenant_id, app_id=dify_ctx.app_id, workflow_id=workflow_id, workflow_run_id=workflow_run_id, node_id=self._node_id, - node_execution_id=self.id, - binding_id=bundle.binding.id, + node_execution_id=self.execution_id, + workflow_agent_binding_id=bundle.binding.id, agent_id=bundle.agent.id, agent_config_snapshot_id=bundle.snapshot.id, ) @@ -200,47 +239,53 @@ class DifyAgentNode(Node[DifyAgentNodeData]): # the second Agent run as deferred_tool_results; if it is somehow still # waiting, re-emit the same pause defensively. deferred_tool_results = None - if self._session_store is not None: - stored_session = self._session_store.load_active_session(session_scope) - if stored_session is not None and stored_session.pending_form_id is not None: - resume_outcome = resolve_ask_human_form( - form_id=stored_session.pending_form_id, - tenant_id=dify_ctx.tenant_id, - node_id=self._node_id, + stored_session = self._session_store.load_or_create_node_execution_session( + session_scope, + home_snapshot_id=bundle.snapshot.home_snapshot_id, + ) + if stored_session.pending_form_id is not None: + resume_outcome = resolve_ask_human_form( + form_id=stored_session.pending_form_id, + tenant_id=dify_ctx.tenant_id, + node_id=self._node_id, + ) + if resume_outcome is not None and resume_outcome.repause is not None: + yield self._pause_event( + reason=resume_outcome.repause, + inputs=inputs, + process_data=process_data, + metadata=metadata, + ) + return + if ( + resume_outcome is not None + and resume_outcome.deferred_result is not None + and stored_session.pending_tool_call_id is not None + ): + deferred_tool_results = build_deferred_tool_results( + tool_call_id=stored_session.pending_tool_call_id, + result=resume_outcome.deferred_result, ) - if resume_outcome is not None and resume_outcome.repause is not None: - yield PauseRequestedEvent(reason=self._to_graph_pause_reason(resume_outcome.repause)) - return - if ( - resume_outcome is not None - and resume_outcome.deferred_result is not None - and stored_session.pending_tool_call_id is not None - ): - deferred_tool_results = build_deferred_tool_results( - tool_call_id=stored_session.pending_tool_call_id, - result=resume_outcome.deferred_result, - ) # ──── Retry loop (Stage 4 §7) ──── attempt = 0 while True: try: - session_snapshot = None - if self._session_store is not None: - session_snapshot = self._session_store.load_active_snapshot(session_scope) runtime_request = self._runtime_request_builder.build( WorkflowAgentRuntimeBuildContext( dify_context=dify_ctx, workflow_id=workflow_id, workflow_run_id=workflow_run_id, node_id=self._node_id, - node_execution_id=self.id, + node_execution_id=self.execution_id, variable_pool=self.graph_runtime_state.variable_pool, binding=bundle.binding, agent=bundle.agent, snapshot=bundle.snapshot, + binding_id=stored_session.binding_id, + backend_binding_ref=stored_session.backend_binding_ref, attempt=attempt, - session_snapshot=session_snapshot, + session_snapshot=stored_session.session_snapshot, deferred_tool_results=deferred_tool_results, ) ) @@ -266,8 +311,9 @@ class DifyAgentNode(Node[DifyAgentNodeData]): # Capture inputs only from the first attempt so retry doesn't churn the # node's "inputs" payload that ends up in the workflow detail view. if attempt == 0: - inputs = {"agent_backend_request": runtime_request.redacted_request} - metadata = dict(runtime_request.metadata) + inputs["agent_backend_request"] = runtime_request.redacted_request + metadata.clear() + metadata.update(runtime_request.metadata) metadata["attempt"] = attempt try: @@ -288,7 +334,12 @@ class DifyAgentNode(Node[DifyAgentNodeData]): "status": create_response.status, } - terminal_event, exhausted = self._consume_event_stream(create_response.run_id, metadata) + terminal_event, exhausted = self._consume_event_stream( + create_response.run_id, + inputs=inputs, + process_data=process_data, + metadata=metadata, + ) if exhausted is not None: # Streaming error / unexpected end — surface immediately without # retrying because the failure is transport-level. @@ -349,29 +400,23 @@ class DifyAgentNode(Node[DifyAgentNodeData]): ) self._save_session_snapshot( session_scope=session_scope, - backend_run_id=terminal_event.run_id, + binding_id=stored_session.binding_id, snapshot=terminal_event.session_snapshot, - runtime_layer_specs=extract_runtime_layer_specs(runtime_request.request.composition), metadata=metadata, pending_form_id=pending_form_id, pending_tool_call_id=pending_tool_call_id, ) - yield PauseRequestedEvent(reason=self._to_graph_pause_reason(pause_reason)) - return - - # Non-success terminal (failed / cancelled) skips per-output - # post-processing — the backend itself already failed. We also retire - # the local ACTIVE session row so a workflow loop back into the same - # Agent node cannot resume from a stale snapshot. The failed agent - # backend layers (suspended per ``on_exit``) are left for agent - # backend's own GC; this row will no longer be picked up by the - # workflow-terminal cleanup layer. - if not isinstance(terminal_event, AgentBackendRunSucceededInternalEvent): - self._mark_session_cleaned_on_failure( - session_scope=session_scope, - backend_run_id=terminal_event.run_id, + yield self._pause_event( + reason=pause_reason, + inputs=inputs, + process_data=process_data, metadata=metadata, ) + return + + # A failed attempt does not retire the product-owned Binding. The + # Workflow Run terminal lifecycle event owns that transition. + if not isinstance(terminal_event, AgentBackendRunSucceededInternalEvent): yield StreamCompletedEvent( node_run_result=self._output_adapter.build_failure_result( event=terminal_event, @@ -384,9 +429,8 @@ class DifyAgentNode(Node[DifyAgentNodeData]): self._save_session_snapshot( session_scope=session_scope, - backend_run_id=terminal_event.run_id, + binding_id=stored_session.binding_id, snapshot=terminal_event.session_snapshot, - runtime_layer_specs=extract_runtime_layer_specs(runtime_request.request.composition), metadata=metadata, ) @@ -458,6 +502,9 @@ class DifyAgentNode(Node[DifyAgentNodeData]): def _consume_event_stream( self, run_id: str, + *, + inputs: dict[str, Any], + process_data: dict[str, Any], metadata: dict[str, Any], ) -> tuple[ _TerminalAgentBackendEvent | None, @@ -507,8 +554,8 @@ class DifyAgentNode(Node[DifyAgentNodeData]): return internal_event, None self._cancel_backend_run(run_id, reason="unexpected_event") return None, self._failure_event( - inputs={}, - process_data={}, + inputs=inputs, + process_data=process_data, metadata=metadata, error=f"Unexpected internal event type {internal_event.type!r}", error_type="agent_backend_stream_error", @@ -516,8 +563,8 @@ class DifyAgentNode(Node[DifyAgentNodeData]): except AgentBackendError as error: self._cancel_backend_run(run_id, reason=self._stream_stop_reason()) return None, self._failure_event( - inputs={}, - process_data={}, + inputs=inputs, + process_data=process_data, metadata=metadata, error=str(error), error_type=self._agent_backend_error_type(error), @@ -525,8 +572,8 @@ class DifyAgentNode(Node[DifyAgentNodeData]): except Exception as error: self._cancel_backend_run(run_id, reason=self._stream_stop_reason()) return None, self._failure_event( - inputs={}, - process_data={}, + inputs=inputs, + process_data=process_data, metadata=metadata, error=str(error), error_type="agent_backend_stream_error", @@ -596,21 +643,17 @@ class DifyAgentNode(Node[DifyAgentNodeData]): self, *, session_scope: WorkflowAgentSessionScope, - backend_run_id: str, + binding_id: str, snapshot: CompositorSessionSnapshot | None, - runtime_layer_specs: list[RuntimeLayerSpec], metadata: dict[str, Any], pending_form_id: str | None = None, pending_tool_call_id: str | None = None, ) -> None: - if self._session_store is None: - return try: self._session_store.save_active_snapshot( scope=session_scope, - backend_run_id=backend_run_id, + binding_id=binding_id, snapshot=snapshot, - runtime_layer_specs=runtime_layer_specs, pending_form_id=pending_form_id, pending_tool_call_id=pending_tool_call_id, ) @@ -619,88 +662,18 @@ class DifyAgentNode(Node[DifyAgentNodeData]): metadata["agent_backend"] = agent_backend except Exception: logger.warning( - "Failed to persist workflow Agent runtime session snapshot: " - "tenant_id=%s workflow_run_id=%s node_id=%s binding_id=%s agent_id=%s backend_run_id=%s", + "Failed to persist workflow Agent Binding session snapshot: " + "tenant_id=%s workflow_run_id=%s node_id=%s binding_id=%s agent_id=%s", session_scope.tenant_id, session_scope.workflow_run_id, session_scope.node_id, - session_scope.binding_id, + session_scope.workflow_agent_binding_id, session_scope.agent_id, - backend_run_id, exc_info=True, ) agent_backend = dict(metadata.get("agent_backend") or {}) agent_backend["session_snapshot_persisted"] = False - agent_backend["session_snapshot_persist_error"] = "workflow_agent_runtime_session_store_error" - metadata["agent_backend"] = agent_backend - - def _mark_session_cleaned_on_failure( - self, - *, - session_scope: WorkflowAgentSessionScope, - backend_run_id: str, - metadata: dict[str, Any], - ) -> None: - if self._session_store is None: - return - stored_session = self._session_store.load_active_session(session_scope) - try: - if stored_session is not None and stored_session.runtime_layer_specs: - payload = AgentBackendSessionCleanupPayload( - session_snapshot=stored_session.session_snapshot, - runtime_layer_specs=stored_session.runtime_layer_specs, - idempotency_key=( - f"{session_scope.tenant_id}:{session_scope.workflow_run_id}:{session_scope.node_id}:" - f"{session_scope.binding_id}:workflow-agent-failure-cleanup:" - f"{stored_session.backend_run_id or 'no-stored-run'}:{backend_run_id}" - ), - metadata={ - "tenant_id": session_scope.tenant_id, - "app_id": session_scope.app_id, - "workflow_id": session_scope.workflow_id, - "workflow_run_id": session_scope.workflow_run_id, - "node_id": session_scope.node_id, - "node_execution_id": session_scope.node_execution_id, - "binding_id": session_scope.binding_id, - "agent_id": session_scope.agent_id, - "agent_config_snapshot_id": session_scope.agent_config_snapshot_id, - "previous_agent_backend_run_id": stored_session.backend_run_id, - "failed_agent_backend_run_id": backend_run_id, - }, - ) - cleanup_workflow_agent_runtime_session.delay(payload.model_dump(mode="json")) - except Exception: - logger.warning( - "Failed to enqueue workflow Agent backend cleanup on agent run failure: " - "tenant_id=%s workflow_run_id=%s node_id=%s binding_id=%s agent_id=%s backend_run_id=%s", - session_scope.tenant_id, - session_scope.workflow_run_id, - session_scope.node_id, - session_scope.binding_id, - session_scope.agent_id, - backend_run_id, - exc_info=True, - ) - try: - self._session_store.mark_cleaned(scope=session_scope, backend_run_id=backend_run_id) - agent_backend = dict(metadata.get("agent_backend") or {}) - agent_backend["session_snapshot_cleaned_on_failure"] = True - metadata["agent_backend"] = agent_backend - except Exception: - logger.warning( - "Failed to mark workflow Agent runtime session cleaned on agent run failure: " - "tenant_id=%s workflow_run_id=%s node_id=%s binding_id=%s agent_id=%s backend_run_id=%s", - session_scope.tenant_id, - session_scope.workflow_run_id, - session_scope.node_id, - session_scope.binding_id, - session_scope.agent_id, - backend_run_id, - exc_info=True, - ) - agent_backend = dict(metadata.get("agent_backend") or {}) - agent_backend["session_snapshot_cleaned_on_failure"] = False - agent_backend["session_snapshot_cleanup_error"] = "workflow_agent_runtime_session_store_error" + agent_backend["session_snapshot_persist_error"] = "workflow_agent_workspace_store_error" metadata["agent_backend"] = agent_backend @staticmethod @@ -742,6 +715,27 @@ class DifyAgentNode(Node[DifyAgentNodeData]): ) ) + def _pause_event( + self, + *, + reason: HumanInputRequired | SchedulingPause, + inputs: dict[str, Any], + process_data: dict[str, Any], + metadata: dict[str, Any], + ) -> NodeRunPauseRequestedEvent: + return NodeRunPauseRequestedEvent( + id=self.execution_id, + node_id=self._node_id, + node_type=self.node_type, + node_run_result=NodeRunResult( + status=WorkflowNodeExecutionStatus.PAUSED, + inputs=inputs, + process_data=process_data, + metadata={WorkflowNodeExecutionMetadataKey.AGENT_LOG: metadata}, + ), + reason=self._to_graph_pause_reason(reason), + ) + @staticmethod def _agent_backend_error_type(error: AgentBackendError) -> str: if isinstance(error, AgentBackendValidationError): diff --git a/api/core/workflow/nodes/agent_v2/binding_resolver.py b/api/core/workflow/nodes/agent_v2/binding_resolver.py index 34812f670f7..3710af100b0 100644 --- a/api/core/workflow/nodes/agent_v2/binding_resolver.py +++ b/api/core/workflow/nodes/agent_v2/binding_resolver.py @@ -41,18 +41,27 @@ class WorkflowAgentBindingResolver: app_id: str, workflow_id: str, node_id: str, + binding_id: str | None = None, + snapshot_id: str | None = None, ) -> WorkflowAgentBindingBundle: - with session_factory.create_session() as session: - binding = session.scalar( - select(WorkflowAgentNodeBinding) - .where( - WorkflowAgentNodeBinding.tenant_id == tenant_id, - WorkflowAgentNodeBinding.app_id == app_id, - WorkflowAgentNodeBinding.workflow_id == workflow_id, - WorkflowAgentNodeBinding.node_id == node_id, - ) - .limit(1) + """Resolve the current binding, optionally at a generation pinned by an existing execution.""" + + if (binding_id is None) != (snapshot_id is None): + raise WorkflowAgentBindingError( + "agent_binding_generation_invalid", + "Workflow Agent binding and config snapshot must be pinned together.", ) + + with session_factory.create_session() as session: + binding_stmt = select(WorkflowAgentNodeBinding).where( + WorkflowAgentNodeBinding.tenant_id == tenant_id, + WorkflowAgentNodeBinding.app_id == app_id, + WorkflowAgentNodeBinding.workflow_id == workflow_id, + WorkflowAgentNodeBinding.node_id == node_id, + ) + if binding_id is not None: + binding_stmt = binding_stmt.where(WorkflowAgentNodeBinding.id == binding_id) + binding = session.scalar(binding_stmt.limit(1)) if binding is None: raise WorkflowAgentBindingError( "agent_binding_not_found", @@ -77,12 +86,16 @@ class WorkflowAgentBindingResolver: f"Agent {binding.agent_id} is not available or has not been published.", ) - snapshot_id = ( - agent.active_config_snapshot_id - if binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT - else binding.current_snapshot_id + effective_snapshot_id = ( + ( + agent.active_config_snapshot_id + if binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT + else binding.current_snapshot_id + ) + if snapshot_id is None + else snapshot_id ) - if snapshot_id is None: + if effective_snapshot_id is None: raise WorkflowAgentBindingError( "agent_config_snapshot_not_found", "Workflow Agent binding has no current config snapshot.", @@ -93,14 +106,14 @@ class WorkflowAgentBindingResolver: .where( AgentConfigSnapshot.tenant_id == tenant_id, AgentConfigSnapshot.agent_id == agent.id, - AgentConfigSnapshot.id == snapshot_id, + AgentConfigSnapshot.id == effective_snapshot_id, ) .limit(1) ) if snapshot is None: raise WorkflowAgentBindingError( "agent_config_snapshot_not_found", - f"Agent config snapshot {snapshot_id} not found.", + f"Agent config snapshot {effective_snapshot_id} not found.", ) session.expunge(binding) diff --git a/api/core/workflow/nodes/agent_v2/output_adapter.py b/api/core/workflow/nodes/agent_v2/output_adapter.py index b5ac93df6ae..f2d408a7999 100644 --- a/api/core/workflow/nodes/agent_v2/output_adapter.py +++ b/api/core/workflow/nodes/agent_v2/output_adapter.py @@ -98,7 +98,7 @@ class WorkflowAgentOutputAdapter: match event: case AgentBackendRunFailedInternalEvent(): error = event.error - error_type = event.reason or "agent_backend_run_failed" + error_type = event.error_type or event.reason or "agent_backend_run_failed" terminal_status = "failed" case AgentBackendRunCancelledInternalEvent(): error = event.message or "Agent backend run was cancelled." diff --git a/api/core/workflow/nodes/agent_v2/runtime_request_builder.py b/api/core/workflow/nodes/agent_v2/runtime_request_builder.py index 22b99412053..367dd2e6d42 100644 --- a/api/core/workflow/nodes/agent_v2/runtime_request_builder.py +++ b/api/core/workflow/nodes/agent_v2/runtime_request_builder.py @@ -33,7 +33,6 @@ from dify_agent.layers.shell import ( DifyShellCliToolConfig, DifyShellEnvVarConfig, DifyShellLayerConfig, - DifyShellSandboxConfig, DifyShellSecretRefConfig, ) from dify_agent.protocol import CreateRunRequest, DeferredToolResultsPayload @@ -136,6 +135,8 @@ class WorkflowAgentRuntimeBuildContext: binding: WorkflowAgentNodeBinding agent: Agent snapshot: AgentConfigSnapshot + binding_id: str + backend_binding_ref: str # Stage 4 §7 / D-4: 0 for the first run, then incremented per retry. Drives the # idempotency key so the backend treats each retry as a fresh request. attempt: int = 0 @@ -251,6 +252,7 @@ class WorkflowAgentRuntimeRequestBuilder: agent_mode=self._agent_backend_agent_mode(context.dify_context.invoke_from), invoke_from=cast(DifyExecutionContextInvokeFrom, context.dify_context.invoke_from.value), ), + backend_binding_ref=context.backend_binding_ref, agent_soul_prompt=soul_prompt or None, workflow_node_job_prompt=workflow_job_prompt, user_prompt=user_prompt, @@ -699,7 +701,7 @@ class WorkflowAgentRuntimeRequestBuilder: "only the accepted file-mapping shape and the returned `reference`; never invent the `reference` " "value.", "If you are replying to the user in natural language and want them to open or download the produced " - "file, include the returned `download_url` in that reply instead of copying it into structured " + "file, include the returned `public_download_url` in that reply instead of copying it into structured " "`final_output` unless the schema explicitly asks for it.", *file_output_lines, ] @@ -739,7 +741,6 @@ class WorkflowAgentRuntimeRequestBuilder: def build_shell_layer_config(agent_soul: AgentSoulConfig) -> DifyShellLayerConfig: """Map Agent Soul shell-adjacent fields into the Agent backend shell config.""" - sandbox_config = _plain_mapping(agent_soul.sandbox.config) return DifyShellLayerConfig( cli_tools=[ tool @@ -750,12 +751,6 @@ def build_shell_layer_config(agent_soul: AgentSoulConfig) -> DifyShellLayerConfi secret_refs=[ secret for secret in (_shell_secret_ref(item) for item in agent_soul.env.secret_refs) if secret is not None ], - sandbox=DifyShellSandboxConfig( - provider=agent_soul.sandbox.provider, - config=sandbox_config, - ) - if agent_soul.sandbox.provider or sandbox_config - else None, ) @@ -841,7 +836,12 @@ def _knowledge_metadata_filtering_config( return DifyKnowledgeMetadataFilteringConfig( mode=metadata_filtering.mode, model_config=_knowledge_model_config(metadata_filtering.metadata_model_config), - conditions=cast(Any, metadata_filtering.conditions.model_dump(mode="json")) + conditions=cast( + Any, + metadata_filtering.conditions.model_dump( + mode="json", exclude={"conditions": {"__all__": {"id", "metadata_id"}}} + ), + ) if metadata_filtering.conditions is not None else None, ) diff --git a/api/core/workflow/nodes/agent_v2/session_cleanup_layer.py b/api/core/workflow/nodes/agent_v2/session_cleanup_layer.py deleted file mode 100644 index c3a7c23a56c..00000000000 --- a/api/core/workflow/nodes/agent_v2/session_cleanup_layer.py +++ /dev/null @@ -1,126 +0,0 @@ -"""Workflow terminal layer that retires Agent backend sessions asynchronously.""" - -from __future__ import annotations - -import logging -from typing import override - -from clients.agent_backend import AgentBackendSessionCleanupPayload -from core.workflow.system_variables import SystemVariableKey, get_system_text -from graphon.graph_engine.layers import GraphEngineLayer -from graphon.graph_events import ( - GraphEngineEvent, - GraphRunAbortedEvent, - GraphRunFailedEvent, - GraphRunPartialSucceededEvent, - GraphRunSucceededEvent, -) -from tasks.agent_backend_session_cleanup_task import cleanup_workflow_agent_runtime_session - -from .session_store import StoredWorkflowAgentSession, WorkflowAgentRuntimeSessionStore - -logger = logging.getLogger(__name__) - - -class WorkflowAgentSessionCleanupLayer(GraphEngineLayer): - """Retire workflow-owned Agent runtime sessions when the workflow ends. - - Workflow termination is a product-lifecycle boundary: once the run reaches a - terminal graph event, the local session row must no longer be resumable. The - actual Agent backend cleanup is therefore dispatched asynchronously with the - persisted snapshot/specs payload, while the local row is marked CLEANED - immediately afterwards regardless of enqueue outcome. - """ - - _TERMINAL_EVENTS = ( - GraphRunSucceededEvent, - GraphRunPartialSucceededEvent, - GraphRunFailedEvent, - GraphRunAbortedEvent, - ) - - def __init__(self, *, session_store: WorkflowAgentRuntimeSessionStore) -> None: - super().__init__() - self._session_store = session_store - - @override - def on_graph_start(self) -> None: - return - - @override - def on_event(self, event: GraphEngineEvent) -> None: - if not isinstance(event, self._TERMINAL_EVENTS): - return - workflow_run_id = get_system_text( - self.graph_runtime_state.variable_pool, - SystemVariableKey.WORKFLOW_EXECUTION_ID, - ) - if not workflow_run_id: - logger.warning("Skipping workflow Agent session cleanup: workflow_run_id is missing.") - return - - for stored_session in self._session_store.list_active_sessions(workflow_run_id=workflow_run_id): - self._cleanup_session(stored_session) - - @override - def on_graph_end(self, error: Exception | None) -> None: - return - - def _cleanup_session(self, stored_session: StoredWorkflowAgentSession) -> None: - scope = stored_session.scope - try: - if stored_session.runtime_layer_specs: - payload = AgentBackendSessionCleanupPayload( - session_snapshot=stored_session.session_snapshot, - runtime_layer_specs=stored_session.runtime_layer_specs, - idempotency_key=f"{scope.workflow_run_id}:{scope.node_id}:{scope.binding_id}:agent-session-cleanup", - metadata={ - "tenant_id": scope.tenant_id, - "app_id": scope.app_id, - "workflow_id": scope.workflow_id, - "workflow_run_id": scope.workflow_run_id, - "node_id": scope.node_id, - "node_execution_id": scope.node_execution_id, - "binding_id": scope.binding_id, - "agent_id": scope.agent_id, - "agent_config_snapshot_id": scope.agent_config_snapshot_id, - "previous_agent_backend_run_id": stored_session.backend_run_id, - }, - ) - cleanup_workflow_agent_runtime_session.delay(payload.model_dump(mode="json")) - else: - logger.warning( - "Skipping workflow Agent backend cleanup enqueue: no runtime_layer_specs persisted. " - "workflow_run_id=%s node_id=%s agent_id=%s", - scope.workflow_run_id, - scope.node_id, - scope.agent_id, - ) - except Exception: - logger.warning( - "Failed to enqueue workflow Agent backend cleanup: " - "workflow_run_id=%s node_id=%s agent_id=%s previous_run_id=%s", - scope.workflow_run_id, - scope.node_id, - scope.agent_id, - stored_session.backend_run_id, - exc_info=True, - ) - finally: - try: - self._session_store.mark_cleaned(scope=scope, backend_run_id=stored_session.backend_run_id) - except Exception: - logger.warning( - "Failed to retire workflow Agent runtime session after cleanup enqueue: " - "workflow_run_id=%s node_id=%s agent_id=%s previous_run_id=%s", - scope.workflow_run_id, - scope.node_id, - scope.agent_id, - stored_session.backend_run_id, - exc_info=True, - ) - - -def build_workflow_agent_session_cleanup_layer() -> WorkflowAgentSessionCleanupLayer: - """Wire the cleanup layer with the standard workflow-owned session store.""" - return WorkflowAgentSessionCleanupLayer(session_store=WorkflowAgentRuntimeSessionStore()) diff --git a/api/core/workflow/nodes/agent_v2/session_store.py b/api/core/workflow/nodes/agent_v2/session_store.py index 215fa67b5eb..134ccc0a62a 100644 --- a/api/core/workflow/nodes/agent_v2/session_store.py +++ b/api/core/workflow/nodes/agent_v2/session_store.py @@ -1,31 +1,32 @@ +"""Workflow Agent participant persistence keyed by node execution.""" + from __future__ import annotations -from dataclasses import dataclass, field +import json +import time +from dataclasses import dataclass from agenton.compositor import CompositorSessionSnapshot -from dify_agent.protocol import RuntimeLayerSpec -from pydantic import TypeAdapter from sqlalchemy import select +from sqlalchemy.orm import Session from core.db.session_factory import session_factory -from libs.datetime_utils import naive_utc_now from models.agent import ( - AgentRuntimeSessionOwnerType, - WorkflowAgentRuntimeSession, - WorkflowAgentRuntimeSessionStatus, + AgentConfigVersionKind, + AgentWorkingResourceStatus, + AgentWorkspace, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, +) +from models.workflow import WorkflowNodeExecutionModel +from services.agent.workspace_service import ( + AgentWorkspaceNotFoundError, + AgentWorkspaceService, + WorkspaceOwnerScope, ) -_SPECS_ADAPTER: TypeAdapter[list[RuntimeLayerSpec]] = TypeAdapter(list[RuntimeLayerSpec]) - - -def _serialize_specs(specs: list[RuntimeLayerSpec]) -> str: - return _SPECS_ADAPTER.dump_json(specs).decode() - - -def _deserialize_specs(value: str | None) -> list[RuntimeLayerSpec]: - if not value: - return [] - return _SPECS_ADAPTER.validate_json(value) +_CALLER_VISIBILITY_ATTEMPTS = 60 +_CALLER_VISIBILITY_INTERVAL_SECONDS = 0.05 @dataclass(frozen=True, slots=True) @@ -36,171 +37,251 @@ class WorkflowAgentSessionScope: workflow_run_id: str | None node_id: str node_execution_id: str - binding_id: str + workflow_agent_binding_id: str agent_id: str agent_config_snapshot_id: str + @property + def workspace_owner(self) -> WorkspaceOwnerScope: + return WorkspaceOwnerScope( + tenant_id=self.tenant_id, + app_id=self.app_id, + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id=self.workflow_run_id or self.node_execution_id, + owner_scope_key=f"{self.node_id}:{self.workflow_agent_binding_id}", + ) + @dataclass(frozen=True, slots=True) class StoredWorkflowAgentSession: scope: WorkflowAgentSessionScope - session_snapshot: CompositorSessionSnapshot - backend_run_id: str | None - runtime_layer_specs: list[RuntimeLayerSpec] = field(default_factory=list) - # ENG-637: set while the session is paused on a dify.ask_human deferred call. + binding_id: str + workspace_id: str + backend_binding_ref: str + session_snapshot: CompositorSessionSnapshot | None pending_form_id: str | None = None pending_tool_call_id: str | None = None -class WorkflowAgentRuntimeSessionStore: - """Stores Agent backend session snapshots for workflow Agent node re-entry.""" +class WorkflowAgentWorkspaceStore: + """Load or create the participant named by a node execution caller row.""" - def load_active_snapshot(self, scope: WorkflowAgentSessionScope) -> CompositorSessionSnapshot | None: - stored = self.load_active_session(scope) - return stored.session_snapshot if stored is not None else None - - def load_active_session(self, scope: WorkflowAgentSessionScope) -> StoredWorkflowAgentSession | None: - """Load the active session row including any pending ask_human correlation.""" - if scope.workflow_run_id is None: - return None + def load_existing_node_execution_scope( + self, + *, + tenant_id: str, + app_id: str, + workflow_id: str, + workflow_run_id: str | None, + node_id: str, + node_execution_id: str, + ) -> WorkflowAgentSessionScope | None: + """Return the generation pinned by an existing node execution participant.""" with session_factory.create_session() as session: - row = session.scalar( - select(WorkflowAgentRuntimeSession).where( - WorkflowAgentRuntimeSession.tenant_id == scope.tenant_id, - WorkflowAgentRuntimeSession.workflow_run_id == scope.workflow_run_id, - WorkflowAgentRuntimeSession.node_id == scope.node_id, - WorkflowAgentRuntimeSession.binding_id == scope.binding_id, - WorkflowAgentRuntimeSession.agent_id == scope.agent_id, - WorkflowAgentRuntimeSession.status == WorkflowAgentRuntimeSessionStatus.ACTIVE, - ) + execution = self._load_execution_by_identity( + session=session, + tenant_id=tenant_id, + app_id=app_id, + workflow_id=workflow_id, + workflow_run_id=workflow_run_id, + node_id=node_id, + node_execution_id=node_execution_id, ) - if row is None: + binding_id = execution.agent_workspace_binding_id + if binding_id is None: return None - return StoredWorkflowAgentSession( - scope=scope, - session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot), - backend_run_id=row.backend_run_id, - runtime_layer_specs=_deserialize_specs(row.composition_layer_specs), - pending_form_id=row.pending_form_id, - pending_tool_call_id=row.pending_tool_call_id, + process_data = execution.process_data_dict + if not isinstance(process_data, dict): + raise AgentWorkspaceNotFoundError("Workflow node execution caller identity is invalid") + workflow_agent_binding_id = process_data.get("workflow_agent_binding_id") + if not isinstance(workflow_agent_binding_id, str): + raise AgentWorkspaceNotFoundError("Workflow node execution caller identity is missing") + owner_scope = WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=app_id, + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id=workflow_run_id or node_execution_id, + owner_scope_key=f"{node_id}:{workflow_agent_binding_id}", + ) + binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=tenant_id, + binding_id=binding_id, + expected_owner_scope=owner_scope, + ) + if binding is None or binding.agent_config_version_kind != AgentConfigVersionKind.SNAPSHOT: + raise AgentWorkspaceNotFoundError("Workflow node participant Binding is unavailable") + return WorkflowAgentSessionScope( + tenant_id=tenant_id, + app_id=app_id, + workflow_id=workflow_id, + workflow_run_id=workflow_run_id, + node_id=node_id, + node_execution_id=node_execution_id, + workflow_agent_binding_id=workflow_agent_binding_id, + agent_id=binding.agent_id, + agent_config_snapshot_id=binding.agent_config_version_id, ) - def list_active_sessions(self, *, workflow_run_id: str) -> list[StoredWorkflowAgentSession]: + def load_or_create_node_execution_session( + self, scope: WorkflowAgentSessionScope, *, home_snapshot_id: str | None + ) -> StoredWorkflowAgentSession: with session_factory.create_session() as session: - rows = session.scalars( - select(WorkflowAgentRuntimeSession).where( - WorkflowAgentRuntimeSession.workflow_run_id == workflow_run_id, - WorkflowAgentRuntimeSession.status == WorkflowAgentRuntimeSessionStatus.ACTIVE, + execution = self._load_execution(session=session, scope=scope) + process_data = execution.process_data_dict + if process_data is None: + process_data = {} + if not isinstance(process_data, dict): + raise AgentWorkspaceNotFoundError("Workflow node execution caller identity is invalid") + stored_workflow_binding_id = process_data.get("workflow_agent_binding_id") + if stored_workflow_binding_id is not None and stored_workflow_binding_id != scope.workflow_agent_binding_id: + raise AgentWorkspaceNotFoundError("Workflow node execution caller identity does not match") + + binding_id = execution.agent_workspace_binding_id + if binding_id is None: + binding = AgentWorkspaceService.create_binding( + session=session, + scope=scope.workspace_owner, + agent_id=scope.agent_id, + base_home_snapshot_id=home_snapshot_id, + agent_config_version_id=scope.agent_config_snapshot_id, + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, ) - ).all() - return [ - StoredWorkflowAgentSession( - scope=WorkflowAgentSessionScope( - tenant_id=row.tenant_id, - app_id=row.app_id, - # These columns are nullable on the unified runtime-session - # table (workflow_run ⊕ conversation owner), but are always - # populated for a workflow-owned row; coerce for the typed scope. - workflow_id=row.workflow_id or "", - workflow_run_id=row.workflow_run_id, - node_id=row.node_id or "", - node_execution_id=row.node_execution_id or "", - binding_id=row.binding_id or "", - agent_id=row.agent_id, - agent_config_snapshot_id=row.agent_config_snapshot_id or "", - ), - session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot), - backend_run_id=row.backend_run_id, - runtime_layer_specs=_deserialize_specs(row.composition_layer_specs), + execution.agent_workspace_binding_id = binding.id + execution.process_data = json.dumps( + { + **process_data, + "workflow_agent_binding_id": scope.workflow_agent_binding_id, + }, + ensure_ascii=False, ) - for row in rows - ] + session.commit() + else: + if stored_workflow_binding_id is None: + raise AgentWorkspaceNotFoundError("Workflow node execution caller identity is missing") + resolved_binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=scope.tenant_id, + binding_id=binding_id, + expected_owner_scope=scope.workspace_owner, + ) + if resolved_binding is None or resolved_binding.agent_id != scope.agent_id: + raise AgentWorkspaceNotFoundError("Workflow node participant Binding is unavailable") + binding = resolved_binding + AgentWorkspaceService.validate_binding_generation( + binding, + base_home_snapshot_id=home_snapshot_id, + agent_config_version_id=scope.agent_config_snapshot_id, + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + return self._stored(scope, binding) def save_active_snapshot( self, *, scope: WorkflowAgentSessionScope, - backend_run_id: str, + binding_id: str, snapshot: CompositorSessionSnapshot | None, - runtime_layer_specs: list[RuntimeLayerSpec], pending_form_id: str | None = None, pending_tool_call_id: str | None = None, ) -> None: - if scope.workflow_run_id is None or snapshot is None: + if snapshot is None: return + AgentWorkspaceService.save_binding_session_snapshot( + tenant_id=scope.tenant_id, + binding_id=binding_id, + session_snapshot=snapshot.model_dump_json(), + pending_form_id=pending_form_id, + pending_tool_call_id=pending_tool_call_id, + ) - snapshot_json = snapshot.model_dump_json() - specs_json = _serialize_specs(runtime_layer_specs) + def retire_workflow_run(self, *, tenant_id: str, app_id: str, workflow_run_id: str) -> list[str]: + """Retire active Workspaces, commit, and return active or already-retired IDs for collection.""" + + retired: list[str] = [] with session_factory.create_session() as session: - row = session.scalar( - select(WorkflowAgentRuntimeSession).where( - WorkflowAgentRuntimeSession.tenant_id == scope.tenant_id, - WorkflowAgentRuntimeSession.workflow_run_id == scope.workflow_run_id, - WorkflowAgentRuntimeSession.node_id == scope.node_id, - WorkflowAgentRuntimeSession.binding_id == scope.binding_id, - WorkflowAgentRuntimeSession.agent_id == scope.agent_id, + workspaces = session.scalars( + select(AgentWorkspace).where( + AgentWorkspace.tenant_id == tenant_id, + AgentWorkspace.app_id == app_id, + AgentWorkspace.owner_type == AgentWorkspaceOwnerType.WORKFLOW_RUN, + AgentWorkspace.owner_id == workflow_run_id, + AgentWorkspace.status.in_((AgentWorkingResourceStatus.ACTIVE, AgentWorkingResourceStatus.RETIRED)), ) - ) - if row is None: - row = WorkflowAgentRuntimeSession( - tenant_id=scope.tenant_id, - app_id=scope.app_id, - owner_type=AgentRuntimeSessionOwnerType.WORKFLOW_RUN, - workflow_id=scope.workflow_id, - workflow_run_id=scope.workflow_run_id, - node_id=scope.node_id, - node_execution_id=scope.node_execution_id, - binding_id=scope.binding_id, - agent_id=scope.agent_id, - agent_config_snapshot_id=scope.agent_config_snapshot_id, - backend_run_id=backend_run_id, - session_snapshot=snapshot_json, - composition_layer_specs=specs_json, - status=WorkflowAgentRuntimeSessionStatus.ACTIVE, - pending_form_id=pending_form_id, - pending_tool_call_id=pending_tool_call_id, + ).all() + for workspace in workspaces: + if workspace.status == AgentWorkingResourceStatus.RETIRED: + retired.append(workspace.id) + continue + workspace_id = AgentWorkspaceService.retire_workspace( + session=session, + tenant_id=tenant_id, + workspace_id=workspace.id, ) - session.add(row) - else: - row.node_execution_id = scope.node_execution_id - row.agent_config_snapshot_id = scope.agent_config_snapshot_id - row.backend_run_id = backend_run_id - row.session_snapshot = snapshot_json - row.composition_layer_specs = specs_json - row.status = WorkflowAgentRuntimeSessionStatus.ACTIVE - row.cleaned_at = None - # Set (or clear, when omitted) the ask_human pause correlation. - row.pending_form_id = pending_form_id - row.pending_tool_call_id = pending_tool_call_id + if workspace_id is not None: + retired.append(workspace_id) session.commit() + return retired - def mark_cleaned(self, *, scope: WorkflowAgentSessionScope, backend_run_id: str | None = None) -> None: - if scope.workflow_run_id is None: - return + @staticmethod + def _load_execution(*, session: Session, scope: WorkflowAgentSessionScope) -> WorkflowNodeExecutionModel: + return WorkflowAgentWorkspaceStore._load_execution_by_identity( + session=session, + tenant_id=scope.tenant_id, + app_id=scope.app_id, + workflow_id=scope.workflow_id, + workflow_run_id=scope.workflow_run_id, + node_id=scope.node_id, + node_execution_id=scope.node_execution_id, + ) - with session_factory.create_session() as session: - row = session.scalar( - select(WorkflowAgentRuntimeSession).where( - WorkflowAgentRuntimeSession.tenant_id == scope.tenant_id, - WorkflowAgentRuntimeSession.workflow_run_id == scope.workflow_run_id, - WorkflowAgentRuntimeSession.node_id == scope.node_id, - WorkflowAgentRuntimeSession.binding_id == scope.binding_id, - WorkflowAgentRuntimeSession.agent_id == scope.agent_id, - WorkflowAgentRuntimeSession.status == WorkflowAgentRuntimeSessionStatus.ACTIVE, - ) - ) - if row is None: - return - if backend_run_id is not None: - row.backend_run_id = backend_run_id - row.status = WorkflowAgentRuntimeSessionStatus.CLEANED - row.cleaned_at = naive_utc_now() - session.commit() + @staticmethod + def _load_execution_by_identity( + *, + session: Session, + tenant_id: str, + app_id: str, + workflow_id: str, + workflow_run_id: str | None, + node_id: str, + node_execution_id: str, + ) -> WorkflowNodeExecutionModel: + """Wait briefly for the already-emitted node-start event to persist its caller row.""" + + stmt = select(WorkflowNodeExecutionModel).where( + WorkflowNodeExecutionModel.id == node_execution_id, + WorkflowNodeExecutionModel.tenant_id == tenant_id, + WorkflowNodeExecutionModel.app_id == app_id, + WorkflowNodeExecutionModel.workflow_id == workflow_id, + WorkflowNodeExecutionModel.node_id == node_id, + WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id, + ) + for attempt in range(_CALLER_VISIBILITY_ATTEMPTS): + execution = session.scalar(stmt) + if execution is not None: + return execution + if attempt < _CALLER_VISIBILITY_ATTEMPTS - 1: + time.sleep(_CALLER_VISIBILITY_INTERVAL_SECONDS) + + raise AgentWorkspaceNotFoundError("Workflow node execution caller is unavailable") + + @staticmethod + def _stored(scope: WorkflowAgentSessionScope, binding: AgentWorkspaceBinding) -> StoredWorkflowAgentSession: + snapshot = ( + CompositorSessionSnapshot.model_validate_json(binding.session_snapshot) + if binding.session_snapshot + else None + ) + return StoredWorkflowAgentSession( + scope=scope, + binding_id=binding.id, + workspace_id=binding.workspace_id, + backend_binding_ref=binding.backend_binding_ref, + session_snapshot=snapshot, + pending_form_id=binding.pending_form_id, + pending_tool_call_id=binding.pending_tool_call_id, + ) -__all__ = [ - "StoredWorkflowAgentSession", - "WorkflowAgentRuntimeSessionStore", - "WorkflowAgentSessionScope", -] +__all__ = ["StoredWorkflowAgentSession", "WorkflowAgentSessionScope", "WorkflowAgentWorkspaceStore"] diff --git a/api/core/workflow/nodes/agent_v2/workspace_retirement_layer.py b/api/core/workflow/nodes/agent_v2/workspace_retirement_layer.py new file mode 100644 index 00000000000..7c44897fcb7 --- /dev/null +++ b/api/core/workflow/nodes/agent_v2/workspace_retirement_layer.py @@ -0,0 +1,89 @@ +"""Retire Workflow Agent Workspaces when the Workflow Run terminates.""" + +from __future__ import annotations + +import logging +from typing import override + +from core.app.entities.app_invoke_entities import DifyRunContext +from core.workflow.nodes.agent_v2.session_store import WorkflowAgentWorkspaceStore +from core.workflow.system_variables import SystemVariableKey, get_system_text +from graphon.graph_engine.layers import GraphEngineLayer +from graphon.graph_events import ( + GraphEngineEvent, + GraphRunAbortedEvent, + GraphRunFailedEvent, + GraphRunPartialSucceededEvent, + GraphRunSucceededEvent, +) +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection + +logger = logging.getLogger(__name__) + + +class WorkflowAgentWorkspaceRetirementLayer(GraphEngineLayer): + """Synchronously retire run Workspaces, then enqueue physical collection.""" + + _TERMINAL_EVENTS = ( + GraphRunSucceededEvent, + GraphRunPartialSucceededEvent, + GraphRunFailedEvent, + GraphRunAbortedEvent, + ) + + def __init__( + self, + *, + dify_run_context: DifyRunContext, + ) -> None: + super().__init__() + self._dify_run_context = dify_run_context + + @override + def on_graph_start(self) -> None: + return + + @override + def on_event(self, event: GraphEngineEvent) -> None: + if not isinstance(event, self._TERMINAL_EVENTS): + return + workflow_run_id = get_system_text( + self.graph_runtime_state.variable_pool, + SystemVariableKey.WORKFLOW_EXECUTION_ID, + ) + if not workflow_run_id: + logger.warning("Skipping Workflow Agent Workspace retirement: workflow_run_id is missing") + return + try: + workspace_ids = WorkflowAgentWorkspaceStore().retire_workflow_run( + tenant_id=self._dify_run_context.tenant_id, + app_id=self._dify_run_context.app_id, + workflow_run_id=workflow_run_id, + ) + except Exception: + logger.exception( + "Failed to retire Workflow Agent Workspaces", + extra={ + "tenant_id": self._dify_run_context.tenant_id, + "app_id": self._dify_run_context.app_id, + "workflow_run_id": workflow_run_id, + }, + ) + return + enqueue_agent_resource_collection( + tenant_id=self._dify_run_context.tenant_id, + workspace_ids=workspace_ids, + ) + + @override + def on_graph_end(self, error: Exception | None) -> None: + return + + +def build_workflow_agent_workspace_retirement_layer( + *, dify_run_context: DifyRunContext +) -> WorkflowAgentWorkspaceRetirementLayer: + return WorkflowAgentWorkspaceRetirementLayer(dify_run_context=dify_run_context) + + +__all__ = ["WorkflowAgentWorkspaceRetirementLayer", "build_workflow_agent_workspace_retirement_layer"] diff --git a/api/core/workflow/nodes/human_input/callback.py b/api/core/workflow/nodes/human_input/callback.py index 46c6f3ceeee..8095078b455 100644 --- a/api/core/workflow/nodes/human_input/callback.py +++ b/api/core/workflow/nodes/human_input/callback.py @@ -2,7 +2,7 @@ from __future__ import annotations import json import logging -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from datetime import datetime, timedelta from typing import Any @@ -11,10 +11,10 @@ from core.repositories.human_input_repository import FormCreateParams, HumanInpu from core.workflow.human_input_adapter import DeliveryChannelConfig from core.workflow.node_runtime import DifyFileReferenceFactory from graphon.nodes.human_input.entities import Completed, Expired, HITLContext, HITLDecision, PauseRequested -from graphon.runtime import VariablePool from graphon.runtime.graph_runtime_state_protocol import ReadOnlyVariablePool from graphon.variables.factory import build_segment from graphon.variables.segments import Segment +from graphon.variables.template_resolution import convert_template from libs.datetime_utils import ensure_naive_utc, naive_utc_now from .entities import ( @@ -31,24 +31,13 @@ from .session_binding import default_session_binding logger = logging.getLogger(__name__) -def _require_template_variable_pool(pool: ReadOnlyVariablePool) -> VariablePool: - """Return the concrete graphon pool required for template expansion.""" - if isinstance(pool, VariablePool): - return pool - - msg = "human input rendering requires graphon.runtime.VariablePool for template expansion" - raise TypeError(msg) - - def render_form_content_before_submission( node_data: HumanInputNodeData, *, variable_pool: ReadOnlyVariablePool, ) -> str: """Process form content by substituting runtime variables before pause.""" - # NOTE(QuantumGhost): This is not ideal, we should expose - # VariablePool method in Graphon. - rendered_form_content = _require_template_variable_pool(variable_pool).convert_template(node_data.form_content) + rendered_form_content = convert_template(variable_pool, node_data.form_content) return rendered_form_content.markdown @@ -91,6 +80,7 @@ class DifyHITLCallback: delivery_methods: Sequence[DeliveryChannelConfig] = (), display_in_ui: bool = False, file_reference_factory: DifyFileReferenceFactory | None = None, + execution_id_getter: Callable[[], str | None] | None = None, ) -> None: self._form_repository = form_repository self._session_binding = default_session_binding @@ -100,11 +90,17 @@ class DifyHITLCallback: self._delivery_methods = tuple(delivery_methods) self._display_in_ui = display_in_ui self._file_reference_factory = file_reference_factory + self._execution_id_getter = execution_id_getter def __call__(self, ctx: HITLContext) -> HITLDecision: - form = self._form_repository.get_form(ctx.node_id) + form_id = self._execution_id_getter() if self._execution_id_getter is not None else None + form = ( + self._form_repository.get_form(ctx.node_id, form_id=form_id) + if form_id is not None + else self._form_repository.get_form(ctx.node_id) + ) if form is None: - created = self._create_form(ctx) + created = self._create_form(ctx, form_id=form_id) return PauseRequested(session_id=self._session_binding.issue_session_id_for_form(form_id=created.id)) status = self._normalize_status(form.status) @@ -163,7 +159,7 @@ class DifyHITLCallback: outputs=outputs, ) - def _create_form(self, ctx: HITLContext) -> HumanInputFormEntity: + def _create_form(self, ctx: HITLContext, *, form_id: str | None = None) -> HumanInputFormEntity: params = FormCreateParams( workflow_execution_id=self._workflow_execution_id or ctx.workflow_execution_id, conversation_id=self._conversation_id, @@ -181,6 +177,7 @@ class DifyHITLCallback: variable_pool=ctx.variable_pool, ) ), + form_id=form_id, ) return self._form_repository.create_form(params) diff --git a/api/core/workflow/nodes/human_input/enums.py b/api/core/workflow/nodes/human_input/enums.py index 53a3bd0b964..da8406d9cf2 100644 --- a/api/core/workflow/nodes/human_input/enums.py +++ b/api/core/workflow/nodes/human_input/enums.py @@ -67,10 +67,10 @@ class FormInputType(enum.StrEnum): class ValueSourceType(enum.StrEnum): """ValueSourceType records whether the value comes from a static setting - in form definiton, or a variable while the workflow is running. + in form definition, or a variable while the workflow is running. """ # `VARIABLE` means that the value comes from a variable in workflow execution VARIABLE = enum.auto() - # `CONSTANT` measn that the value comes from a static setting in form definition. + # `CONSTANT` means that the value comes from a static setting in form definition. CONSTANT = enum.auto() diff --git a/api/core/workflow/nodes/human_input/pause_reason.py b/api/core/workflow/nodes/human_input/pause_reason.py index 934219392f0..26c28691508 100644 --- a/api/core/workflow/nodes/human_input/pause_reason.py +++ b/api/core/workflow/nodes/human_input/pause_reason.py @@ -15,8 +15,8 @@ class DifyHITLEventType(StrEnum): """ - # Ideally this should be a string constaint. However, we cannot put - # string constant into Literal type cosntructor. We have to warp it as a + # Ideally this should be a string constraint. However, we cannot put + # string constant into Literal type constructor. We have to wrap it as a # string enumeration. HUMAN_INPUT_REQUIRED = PauseReasonType.LEGACY_HUMAN_INPUT_REQUIRED.value diff --git a/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py b/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py index 11082d53fa2..b4e975dcfa6 100644 --- a/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py +++ b/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py @@ -25,7 +25,6 @@ from graphon.enums import ( from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.model_runtime.utils.encoders import jsonable_encoder from graphon.node_events import NodeRunResult -from graphon.nodes.base import LLMUsageTrackingMixin from graphon.nodes.base.node import Node from graphon.variables import ( ArrayFileSegment, @@ -33,6 +32,7 @@ from graphon.variables import ( StringSegment, ) from graphon.variables.segments import ArrayObjectSegment +from graphon.variables.template_resolution import convert_template from .entities import ( Condition, @@ -64,7 +64,7 @@ def _normalize_metadata_filter_sequence_item(value: object) -> str: return value if isinstance(value, str) else str(value) -class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeData]): +class KnowledgeRetrievalNode(Node[KnowledgeRetrievalNodeData]): node_type = BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL # Instance attributes specific to LLMNode. @@ -309,7 +309,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD resolved_value: str | Sequence[str] | int | float | None match value: case str(): - segment_group = variable_pool.convert_template(value) + segment_group = convert_template(variable_pool, value) if len(segment_group.value) == 1: resolved_value = _normalize_metadata_filter_scalar(segment_group.value[0].to_object()) else: @@ -317,7 +317,7 @@ class KnowledgeRetrievalNode(LLMUsageTrackingMixin, Node[KnowledgeRetrievalNodeD case _ if isinstance(value, Sequence) and all(isinstance(v, str) for v in value): resolved_values: list[str] = [] for v in value: - segment_group = variable_pool.convert_template(v) + segment_group = convert_template(variable_pool, v) if len(segment_group.value) == 1: resolved_values.append( _normalize_metadata_filter_sequence_item(segment_group.value[0].to_object()) diff --git a/api/core/workflow/workflow_entry.py b/api/core/workflow/workflow_entry.py index fb12922ed7f..866bc73fcf6 100644 --- a/api/core/workflow/workflow_entry.py +++ b/api/core/workflow/workflow_entry.py @@ -2,13 +2,13 @@ import logging import time from collections.abc import Generator, Mapping, Sequence from typing import Any, TypedDict +from uuid import uuid4 from configs import dify_config from context import capture_current_context from core.app.apps.exc import GenerateTaskStoppedError from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom, build_dify_run_context from core.app.file_access import DatabaseFileAccessController -from core.app.workflow.layers.llm_quota import LLMQuotaLayer from core.app.workflow.layers.observability import ObservabilityLayer from core.workflow.node_factory import ( DifyGraphInitContext, @@ -26,7 +26,6 @@ from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add from core.workflow.variable_prefixes import ENVIRONMENT_VARIABLE_NODE_ID from extensions.otel.runtime import is_instrument_flag_enabled from factories import file_factory -from graphon.entities import GraphInitParams from graphon.entities.graph_config import NodeConfigDictAdapter from graphon.errors import WorkflowNodeRunFailedError from graphon.file import File @@ -34,11 +33,12 @@ from graphon.filters import GraphEventFilterContext, ResponseStreamFilter, filte from graphon.graph import Graph from graphon.graph_engine import GraphEngine, GraphEngineConfig from graphon.graph_engine.command_channels import CommandChannel, InMemoryChannel -from graphon.graph_engine.layers import DebugLoggingLayer, ExecutionLimitsLayer -from graphon.graph_events import GraphEngineEvent, GraphNodeEventBase, GraphRunFailedEvent +from graphon.graph_engine.layers import DebugLoggingLayer, ExecutionLimitsLayer, GraphEngineLayer +from graphon.graph_events import GraphEngineEvent, GraphNodeEventBase, GraphRunFailedEvent, is_node_result_event from graphon.nodes import BuiltinNodeTypes from graphon.nodes.base.node import Node -from graphon.runtime import ChildGraphNotFoundError, GraphRuntimeState, VariablePool +from graphon.nodes.container_effects import ContainerAwaitRequest +from graphon.runtime import GraphRuntimeState, ReadOnlyGraphRuntimeStateWrapper, VariablePool from graphon.variable_loader import DUMMY_VARIABLE_LOADER, VariableLoader, load_into_variable_pool from models.workflow import Workflow @@ -69,77 +69,6 @@ def iter_dify_graph_engine_events( ) -class _WorkflowChildEngineBuilder: - tenant_id: str - - def __init__(self, *, tenant_id: str) -> None: - self.tenant_id = tenant_id - - @staticmethod - def _has_node_id(graph_config: Mapping[str, Any], node_id: str) -> bool | None: - """ - Return whether `graph_config["nodes"]` contains the given node id. - - Returns `None` when the nodes payload shape is unexpected, so graph-level - validation can surface the original configuration error. - """ - nodes = graph_config.get("nodes") - if not isinstance(nodes, list): - return None - - for node in nodes: - if not isinstance(node, Mapping): - return None - current_id = node.get("id") - if isinstance(current_id, str) and current_id == node_id: - return True - return False - - def build_child_engine( - self, - *, - workflow_id: str, - graph_init_params: GraphInitParams, - parent_graph_runtime_state: GraphRuntimeState, - root_node_id: str, - variable_pool: VariablePool | None = None, - ) -> GraphEngine: - """Build a child engine with a fresh runtime state and only child-safe layers.""" - child_graph_runtime_state = GraphRuntimeState( - variable_pool=variable_pool if variable_pool is not None else parent_graph_runtime_state.variable_pool, - start_at=time.perf_counter(), - execution_context=parent_graph_runtime_state.execution_context, - ) - node_factory = DifyNodeFactory( - graph_init_params=graph_init_params, - graph_runtime_state=child_graph_runtime_state, - ) - - graph_config = graph_init_params.graph_config - has_root_node = self._has_node_id(graph_config=graph_config, node_id=root_node_id) - if has_root_node is False: - raise ChildGraphNotFoundError(f"child graph root node '{root_node_id}' not found") - - child_graph = Graph.init( - graph_config=graph_config, - node_factory=node_factory, - root_node_id=root_node_id, - ) - - command_channel = InMemoryChannel() - config = GraphEngineConfig() - child_engine = GraphEngine( - workflow_id=workflow_id, - graph=child_graph, - graph_runtime_state=child_graph_runtime_state, - command_channel=command_channel, - config=config, - child_engine_builder=self, - ) - child_engine.layer(LLMQuotaLayer(tenant_id=self.tenant_id)) - return child_engine - - class _NodeConfigDict(TypedDict): id: str width: int @@ -208,8 +137,8 @@ class WorkflowEntry: self.command_channel = command_channel self._response_stream_filter = response_stream_filter or ResponseStreamFilter() execution_context = capture_current_context() - graph_runtime_state.execution_context = execution_context - self._child_engine_builder = _WorkflowChildEngineBuilder(tenant_id=tenant_id) + # ponytail: Graphon snapshots omit process-local context; use a public rebind API when Graphon exposes one. + graph_runtime_state._execution_context = execution_context self.graph_engine = GraphEngine( workflow_id=workflow_id, graph=graph, @@ -221,7 +150,6 @@ class WorkflowEntry: scale_up_threshold=dify_config.GRAPH_ENGINE_SCALE_UP_THRESHOLD, scale_down_idle_time=dify_config.GRAPH_ENGINE_SCALE_DOWN_IDLE_TIME, ), - child_engine_builder=self._child_engine_builder, ) # Add debug logging layer when in debug mode @@ -241,7 +169,6 @@ class WorkflowEntry: max_steps=dify_config.WORKFLOW_MAX_EXECUTION_STEPS, max_time=dify_config.WORKFLOW_MAX_EXECUTION_TIME ) self.graph_engine.layer(limits_layer) - self.graph_engine.layer(LLMQuotaLayer(tenant_id=tenant_id)) # Add observability layer when OTel is enabled if dify_config.ENABLE_OTEL or is_instrument_flag_enabled(): @@ -271,7 +198,7 @@ class WorkflowEntry: user_inputs: Mapping[str, Any], variable_pool: VariablePool, variable_loader: VariableLoader = DUMMY_VARIABLE_LOADER, - ) -> tuple[Node, Generator[GraphNodeEventBase, None, None]]: + ) -> tuple[Node, Generator[GraphNodeEventBase | ContainerAwaitRequest, None, None]]: """ Single step run workflow node :param workflow: Workflow instance @@ -285,6 +212,8 @@ class WorkflowEntry: # Get node type node_type = node_config_data.type + if node_type in {BuiltinNodeTypes.LOOP, BuiltinNodeTypes.ITERATION}: + raise ValueError("Loop and Iteration nodes must use their engine-backed debug endpoints") node_version = str(node_config_data.version) node_cls = resolve_workflow_node_class(node_type=node_type, node_version=node_version) @@ -358,7 +287,7 @@ class WorkflowEntry: node = node_factory.create_node(node_config) try: - generator = cls._traced_node_run(node) + generator = cls._run_node_with_layers(node, tenant_id=workflow.tenant_id) except Exception as e: logger.exception( "error while running node, workflow_id=%s, node_id=%s, node_type=%s, node_version=%s", @@ -419,7 +348,7 @@ class WorkflowEntry: @classmethod def run_free_node( cls, node_data: dict[str, Any], node_id: str, tenant_id: str, user_id: str, user_inputs: dict[str, Any] - ) -> tuple[Node, Generator[GraphNodeEventBase, None, None]]: + ) -> tuple[Node, Generator[GraphNodeEventBase | ContainerAwaitRequest, None, None]]: """ Run free node @@ -497,7 +426,7 @@ class WorkflowEntry: tenant_id=tenant_id, ) - generator = cls._traced_node_run(node) + generator = cls._run_node_with_layers(node, tenant_id=tenant_id) return node, generator except Exception as e: @@ -613,24 +542,48 @@ class WorkflowEntry: variable_pool.add([variable_node_id] + variable_key_list, input_value) @staticmethod - def _traced_node_run(node: Node) -> Generator[GraphNodeEventBase, None, None]: + def _run_node_with_layers( + node: Node, *, tenant_id: str + ) -> Generator[GraphNodeEventBase | ContainerAwaitRequest, None, None]: """ - Wraps a node's run method with OpenTelemetry tracing and returns a generator. + Run a standalone node with the same quota and observability hooks as GraphEngine. """ - # Wrap node.run() with ObservabilityLayer hooks to produce node-level spans - layer = ObservabilityLayer() - layer.on_graph_start() - node.ensure_execution_id() + layers: Sequence[GraphEngineLayer] = (ObservabilityLayer(),) + command_channel = InMemoryChannel() + runtime_state = ReadOnlyGraphRuntimeStateWrapper(node.graph_runtime_state) + for layer in layers: + layer.initialize(runtime_state, command_channel) + layer.on_graph_start() + + node.bind_execution_id(str(uuid4())) def _gen(): error: Exception | None = None - layer.on_node_run_start(node) + result_event: GraphNodeEventBase | None = None + layers_finished = False + + def finish_layers() -> None: + nonlocal layers_finished + if layers_finished: + return + layers_finished = True + for layer in layers: + layer.on_node_run_end(node, error, result_event) + for layer in layers: + layer.on_graph_end(error) + try: - yield from node.run() + for layer in layers: + layer.on_node_run_start(node) + for event in node.run(): + if isinstance(event, GraphNodeEventBase) and is_node_result_event(event): + result_event = event + finish_layers() + yield event except Exception as exc: error = exc raise finally: - layer.on_node_run_end(node, error) + finish_layers() return _gen() diff --git a/api/dev/lint_response_contracts.py b/api/dev/lint_response_contracts.py index 6cfc1c6b446..5914f1db469 100644 --- a/api/dev/lint_response_contracts.py +++ b/api/dev/lint_response_contracts.py @@ -668,6 +668,7 @@ def checks_for_file(file_path: Path, repo_root: Path) -> list[ContractCheck]: def as_jsonable(check: ContractCheck) -> dict[str, Any]: data = asdict(check) + # pyrefly: ignore [bad-assignment] data["documented"] = {str(status): model for status, model in check.documented.items()} return data diff --git a/api/docker/entrypoint.sh b/api/docker/entrypoint.sh index 67832da826a..53c598a5174 100755 --- a/api/docker/entrypoint.sh +++ b/api/docker/entrypoint.sh @@ -31,13 +31,13 @@ if [[ "${MODE}" == "worker" ]]; then CONCURRENCY_OPTION="-c ${CELERY_WORKER_AMOUNT:-1}" fi - # Configure queues based on edition if not explicitly set + # Configure queues based on product edition if not explicitly set if [[ -z "${CELERY_QUEUES}" ]]; then - if [[ "${EDITION}" == "CLOUD" ]]; then + if [[ "${DEPLOYMENT_EDITION:-COMMUNITY}" == "CLOUD" ]]; then # Cloud edition: separate queues for dataset and trigger tasks DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,app_rbac,plugin,workflow_storage,conversation,workflow_professional,workflow_team,workflow_sandbox,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution" else - # Community edition (SELF_HOSTED): dataset, pipeline and workflow have separate queues + # Self-hosted editions: dataset, pipeline and workflow have separate queues DEFAULT_QUEUES="api_token,dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,app_rbac,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_publisher,trigger_refresh_executor,retention,workflow_based_app_execution" fi else diff --git a/api/enterprise/telemetry/README.md b/api/enterprise/telemetry/README.md index e43c0b1ea29..4d0065a5416 100644 --- a/api/enterprise/telemetry/README.md +++ b/api/enterprise/telemetry/README.md @@ -32,7 +32,7 @@ The Enterprise OTEL exporter is configured via environment variables. | Variable | Description | Default | |----------|-------------|---------| -| `ENTERPRISE_ENABLED` | Master switch for all enterprise features. | `false` | +| `DEPLOYMENT_EDITION` | Product edition; enterprise telemetry is only available in `ENTERPRISE`. | `COMMUNITY` | | `ENTERPRISE_TELEMETRY_ENABLED` | Master switch for enterprise telemetry. | `false` | | `ENTERPRISE_OTLP_ENDPOINT` | OTLP collector endpoint (e.g., `http://otel-collector:4318`). | - | | `ENTERPRISE_OTLP_HEADERS` | Custom headers for OTLP requests (e.g., `x-scope-orgid=tenant1`). | - | diff --git a/api/enterprise/telemetry/exporter.py b/api/enterprise/telemetry/exporter.py index 80959514f28..a177b4f12a4 100644 --- a/api/enterprise/telemetry/exporter.py +++ b/api/enterprise/telemetry/exporter.py @@ -42,12 +42,15 @@ from enterprise.telemetry.id_generator import ( set_correlation_id, set_span_id_source, ) +from enums import DeploymentEdition logger = logging.getLogger(__name__) def is_enterprise_telemetry_enabled() -> bool: - return bool(dify_config.ENTERPRISE_ENABLED and dify_config.ENTERPRISE_TELEMETRY_ENABLED) + return bool( + dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE and dify_config.ENTERPRISE_TELEMETRY_ENABLED + ) def _parse_otlp_headers(raw: str) -> dict[str, str]: diff --git a/api/enums/__init__.py b/api/enums/__init__.py index e69de29bb2d..dd5d029f56f 100644 --- a/api/enums/__init__.py +++ b/api/enums/__init__.py @@ -0,0 +1,59 @@ +from enum import StrEnum, auto + + +class CloudPlan(StrEnum): + """ + Enum representing user plan types in the cloud platform. + + SANDBOX: Free/default plan with limited features + PROFESSIONAL: Professional paid plan + TEAM: Team collaboration paid plan + """ + + SANDBOX = auto() + PROFESSIONAL = auto() + TEAM = auto() + + +class DeploymentEdition(StrEnum): + """Enum representing the deployment edition of the platform.""" + + COMMUNITY = "COMMUNITY" + ENTERPRISE = "ENTERPRISE" + CLOUD = "CLOUD" + + +class HostedTrialProvider(StrEnum): + """Enum representing hosted model provider names for trial access.""" + + OPENAI = "langgenius/openai/openai" + ANTHROPIC = "langgenius/anthropic/anthropic" + GEMINI = "langgenius/gemini/google" + X = "langgenius/x/x" + DEEPSEEK = "langgenius/deepseek/deepseek" + TONGYI = "langgenius/tongyi/tongyi" + + @property + def config_key(self) -> str: + """Return the config key used in dify_config (e.g., HOSTED_{config_key}_PAID_ENABLED).""" + if self == HostedTrialProvider.X: + return "XAI" + return self.name + + +class QuotaType(StrEnum): + """Supported quota types for tenant feature usage.""" + + TRIGGER = auto() + WORKFLOW = auto() + UNLIMITED = auto() + + @property + def billing_key(self) -> str: + match self: + case QuotaType.TRIGGER: + return "trigger_event" + case QuotaType.WORKFLOW: + return "api_rate_limit" + case _: + raise ValueError(f"Invalid quota type: {self}") diff --git a/api/enums/cloud_plan.py b/api/enums/cloud_plan.py deleted file mode 100644 index 927cff5471a..00000000000 --- a/api/enums/cloud_plan.py +++ /dev/null @@ -1,15 +0,0 @@ -from enum import StrEnum, auto - - -class CloudPlan(StrEnum): - """ - Enum representing user plan types in the cloud platform. - - SANDBOX: Free/default plan with limited features - PROFESSIONAL: Professional paid plan - TEAM: Team collaboration paid plan - """ - - SANDBOX = auto() - PROFESSIONAL = auto() - TEAM = auto() diff --git a/api/enums/deployment_edition.py b/api/enums/deployment_edition.py deleted file mode 100644 index 5541651b576..00000000000 --- a/api/enums/deployment_edition.py +++ /dev/null @@ -1,11 +0,0 @@ -from enum import StrEnum - - -class DeploymentEdition(StrEnum): - """ - Enum representing the deployment edition of the platform. - """ - - COMMUNITY = "COMMUNITY" - ENTERPRISE = "ENTERPRISE" - CLOUD = "CLOUD" diff --git a/api/enums/hosted_provider.py b/api/enums/hosted_provider.py deleted file mode 100644 index c6d3715dc17..00000000000 --- a/api/enums/hosted_provider.py +++ /dev/null @@ -1,21 +0,0 @@ -from enum import StrEnum - - -class HostedTrialProvider(StrEnum): - """ - Enum representing hosted model provider names for trial access. - """ - - OPENAI = "langgenius/openai/openai" - ANTHROPIC = "langgenius/anthropic/anthropic" - GEMINI = "langgenius/gemini/google" - X = "langgenius/x/x" - DEEPSEEK = "langgenius/deepseek/deepseek" - TONGYI = "langgenius/tongyi/tongyi" - - @property - def config_key(self) -> str: - """Return the config key used in dify_config (e.g., HOSTED_{config_key}_PAID_ENABLED).""" - if self == HostedTrialProvider.X: - return "XAI" - return self.name diff --git a/api/enums/quota_type.py b/api/enums/quota_type.py deleted file mode 100644 index a10ac21f69e..00000000000 --- a/api/enums/quota_type.py +++ /dev/null @@ -1,21 +0,0 @@ -from enum import StrEnum, auto - - -class QuotaType(StrEnum): - """ - Supported quota types for tenant feature usage. - """ - - TRIGGER = auto() - WORKFLOW = auto() - UNLIMITED = auto() - - @property - def billing_key(self) -> str: - match self: - case QuotaType.TRIGGER: - return "trigger_event" - case QuotaType.WORKFLOW: - return "api_rate_limit" - case _: - raise ValueError(f"Invalid quota type: {self}") diff --git a/api/events/event_handlers/create_document_index.py b/api/events/event_handlers/create_document_index.py index 8bc4240251f..aba387dd187 100644 --- a/api/events/event_handlers/create_document_index.py +++ b/api/events/event_handlers/create_document_index.py @@ -21,7 +21,7 @@ def handle(sender, **kwargs): document_ids = kwargs.get("document_ids", []) start_at = time.perf_counter() try: - indexing_runner = IndexingRunner() + indexing_runner = IndexingRunner(enforce_vector_space_admission=True) with session_factory.create_session() as session: documents = [] for document_id in document_ids: diff --git a/api/events/event_handlers/queue_credential_sync_when_tenant_created.py b/api/events/event_handlers/queue_credential_sync_when_tenant_created.py index 6566c214b05..a9f5e815018 100644 --- a/api/events/event_handlers/queue_credential_sync_when_tenant_created.py +++ b/api/events/event_handlers/queue_credential_sync_when_tenant_created.py @@ -1,4 +1,5 @@ from configs import dify_config +from enums import DeploymentEdition from events.tenant_event import tenant_was_created from services.enterprise.workspace_sync import WorkspaceSyncService @@ -7,7 +8,7 @@ from services.enterprise.workspace_sync import WorkspaceSyncService def handle(sender, **kwargs): """Queue credential sync when a tenant/workspace is created.""" # Only queue sync tasks if plugin manager (enterprise feature) is enabled - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return tenant = sender diff --git a/api/extensions/ext_application_services.py b/api/extensions/ext_application_services.py new file mode 100644 index 00000000000..db647eafbe2 --- /dev/null +++ b/api/extensions/ext_application_services.py @@ -0,0 +1,102 @@ +"""Composition root for application services used by transport adapters.""" + +from dataclasses import dataclass +from typing import cast + +from flask import Flask, current_app +from sqlalchemy.orm import Session, sessionmaker + +from configs import dify_config +from constants.dsl_version import CURRENT_APP_DSL_VERSION +from core.db.session_factory import get_session_maker +from core.schemas.schema_manager import SchemaManager +from enums import DeploymentEdition +from extensions.ext_redis import RedisClientWrapper, redis_client +from repositories.explore_banner_query_repository import ExploreBannerQueryRepository +from repositories.installation_state_repository import InstallationStateRepository +from repositories.workspace_member_query_repository import WorkspaceMemberQueryRepository +from repositories.workspace_query_repository import WorkspaceQueryRepository +from services.explore_banner_query_service import ExploreBannerQueryService +from services.feature_query_service import FeatureQueryService +from services.feature_service import FeatureService +from services.feature_service_gateway import FeatureServiceGateway +from services.init_validation_service import InitValidationService +from services.schema_definition_service import SchemaDefinitionService +from services.setup_adapters import RedisSetupLock, RegisterServiceAccountProvisioner +from services.setup_service import SetupService +from services.workspace_member_query_service import WorkspaceMemberQueryService +from services.workspace_member_role_resolver import DeploymentWorkspaceMemberRoleResolver +from services.workspace_plan_gateway import DeploymentWorkspacePlanGateway +from services.workspace_query_service import WorkspaceQueryService + +_EXTENSION_KEY = "application_services" + + +@dataclass(frozen=True, slots=True) +class ApplicationServices: + explore_banner_queries: ExploreBannerQueryService + schema_definitions: SchemaDefinitionService + setup: SetupService + feature_queries: FeatureQueryService + init_validation: InitValidationService + workspace_queries: WorkspaceQueryService + workspace_member_queries: WorkspaceMemberQueryService + + +def build_application_services( + *, + database_client: sessionmaker[Session], + deployment_edition: DeploymentEdition, + initialization_password: str, + redis: RedisClientWrapper, +) -> ApplicationServices: + installation_state = InstallationStateRepository(client=database_client) + return ApplicationServices( + explore_banner_queries=ExploreBannerQueryService( + banners=ExploreBannerQueryRepository(client=database_client), + is_enabled=FeatureService.is_explore_banner_enabled, + ), + schema_definitions=SchemaDefinitionService(source_factory=SchemaManager), + setup=SetupService( + state=installation_state, + accounts=RegisterServiceAccountProvisioner(client=database_client), + lock=RedisSetupLock(client=redis), + setup_required=deployment_edition != DeploymentEdition.CLOUD, + ), + feature_queries=FeatureQueryService( + features=FeatureServiceGateway(), + trial_models=FeatureService.get_trial_models(), + app_dsl_version=CURRENT_APP_DSL_VERSION, + ), + init_validation=InitValidationService( + state=installation_state, + validation_required=(deployment_edition != DeploymentEdition.CLOUD and bool(initialization_password)), + expected_password=initialization_password, + ), + workspace_queries=WorkspaceQueryService( + workspaces=WorkspaceQueryRepository( + client=database_client, + ), + plans=DeploymentWorkspacePlanGateway(), + ), + workspace_member_queries=WorkspaceMemberQueryService( + members=WorkspaceMemberQueryRepository( + session_factory=database_client, + ), + roles=DeploymentWorkspaceMemberRoleResolver(), + ), + ) + + +def init_app(app: Flask) -> None: + app.extensions[_EXTENSION_KEY] = build_application_services( + database_client=get_session_maker(), + deployment_edition=dify_config.DEPLOYMENT_EDITION, + initialization_password=dify_config.INIT_PASSWORD, + redis=redis_client, + ) + + +def application_services() -> ApplicationServices: + """Return the application services bound to the current Flask app.""" + return cast(ApplicationServices, current_app.extensions[_EXTENSION_KEY]) diff --git a/api/extensions/ext_celery.py b/api/extensions/ext_celery.py index 2cf3505e918..a69971d7848 100644 --- a/api/extensions/ext_celery.py +++ b/api/extensions/ext_celery.py @@ -5,10 +5,12 @@ from typing import Any import pytz # type: ignore[import-untyped] from celery import Celery, Task from celery.schedules import crontab +from celery.signals import beat_init from typing_extensions import TypedDict from configs import dify_config from dify_app import DifyApp +from enums import DeploymentEdition from extensions.redis_names import normalize_redis_key_prefix from extensions.workflow_warm_shutdown import setup_workflow_warm_shutdown_handler @@ -36,6 +38,19 @@ class CeleryBeatScheduleEntry(TypedDict): schedule: crontab | timedelta +def _enqueue_initial_community_telemetry_heartbeat(sender: Any, **_: Any) -> None: + task_name = "community_telemetry.send_heartbeat" + if "community_telemetry_heartbeat" not in sender.app.conf.beat_schedule: + return + + task = sender.app.tasks.get(task_name) + if task is not None: + task.apply_async() + + +beat_init.connect(_enqueue_initial_community_telemetry_heartbeat, weak=False) + + def get_celery_ssl_options() -> CelerySSLOptionsDict | None: """Get SSL configuration for Celery broker/backend connections.""" # Only apply SSL if we're using Redis as broker/backend @@ -152,6 +167,7 @@ def init_app(app: DifyApp) -> Celery: imports = [ "tasks.async_workflow_tasks", # trigger workers + "tasks.collect_agent_resources_task", # retired Agent resource collection "tasks.trigger_processing_tasks", # async trigger processing "tasks.generate_summary_index_task", # summary index generation "tasks.regenerate_summary_index_task", # summary index regeneration @@ -260,7 +276,19 @@ def init_app(app: DifyApp) -> Celery: "schedule": timedelta(minutes=dify_config.API_TOKEN_LAST_USED_UPDATE_INTERVAL), } - if dify_config.ENTERPRISE_ENABLED and dify_config.ENTERPRISE_TELEMETRY_ENABLED: + if ( + dify_config.DEPLOYMENT_EDITION == DeploymentEdition.COMMUNITY + and not dify_config.DISABLE_TELEMETRY + and not dify_config.DO_NOT_TRACK + and not dify_config.CI + ): + imports.append("tasks.community_telemetry_task") + beat_schedule["community_telemetry_heartbeat"] = { + "task": "community_telemetry.send_heartbeat", + "schedule": timedelta(minutes=dify_config.TELEMETRY_HEARTBEAT_INTERVAL_MINUTES), + } + + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE and dify_config.ENTERPRISE_TELEMETRY_ENABLED: imports.append("tasks.enterprise_telemetry_task") celery_app.conf.update(beat_schedule=beat_schedule, imports=imports) diff --git a/api/extensions/ext_enterprise_telemetry.py b/api/extensions/ext_enterprise_telemetry.py index b3cfa01aee6..2d9bcbad6d1 100644 --- a/api/extensions/ext_enterprise_telemetry.py +++ b/api/extensions/ext_enterprise_telemetry.py @@ -4,7 +4,7 @@ Initializes the EnterpriseExporter singleton during ``create_app()`` (single-threaded), registers blinker event handlers, and hooks atexit for graceful shutdown. -Skipped entirely when either ``ENTERPRISE_ENABLED`` or ``ENTERPRISE_TELEMETRY_ENABLED`` +Skipped entirely outside the Enterprise edition or when ``ENTERPRISE_TELEMETRY_ENABLED`` is false (``is_enabled()`` gate). """ @@ -15,6 +15,7 @@ import logging from typing import TYPE_CHECKING from configs import dify_config +from enums import DeploymentEdition if TYPE_CHECKING: from dify_app import DifyApp @@ -26,7 +27,9 @@ _exporter: EnterpriseExporter | None = None def is_enabled() -> bool: - return bool(dify_config.ENTERPRISE_ENABLED and dify_config.ENTERPRISE_TELEMETRY_ENABLED) + return bool( + dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE and dify_config.ENTERPRISE_TELEMETRY_ENABLED + ) def init_app(app: DifyApp) -> None: diff --git a/api/extensions/ext_fastopenapi.py b/api/extensions/ext_fastopenapi.py index d51f2781fc5..236b43ff0f6 100644 --- a/api/extensions/ext_fastopenapi.py +++ b/api/extensions/ext_fastopenapi.py @@ -34,11 +34,11 @@ def init_app(app: DifyApp) -> None: # Ensure route decorators are evaluated. import controllers.console.init_validate as init_validate_module - import controllers.console.ping as ping_module + import controllers.console.system as system_module from controllers.console import remote_files, setup _ = init_validate_module - _ = ping_module + _ = system_module _ = remote_files _ = setup diff --git a/api/extensions/ext_key_provider.py b/api/extensions/ext_key_provider.py new file mode 100644 index 00000000000..add12ce57f3 --- /dev/null +++ b/api/extensions/ext_key_provider.py @@ -0,0 +1,50 @@ +import logging +from collections.abc import Callable + +from flask import Flask + +from configs import dify_config +from dify_app import DifyApp +from libs.key_providers.base import BaseKeyProvider +from libs.key_providers.key_provider_type import KeyProviderType + +logger = logging.getLogger(__name__) + + +class KeyProviderManager: + _provider: BaseKeyProvider | None = None + + def init_app(self, app: Flask): + with app.app_context(): + self._provider = self._build_provider() + + @property + def provider(self) -> BaseKeyProvider: + if self._provider is None: + self._provider = self._build_provider() + return self._provider + + def _build_provider(self) -> BaseKeyProvider: + provider_factory = self.get_provider_factory(dify_config.KEY_PROVIDER_TYPE) + return provider_factory() + + @staticmethod + def get_provider_factory(provider_type: str) -> Callable[[], BaseKeyProvider]: + match provider_type: + case KeyProviderType.LOCAL: + from libs.key_providers.rsa_key_provider import RSAKeyProvider + + return RSAKeyProvider + case KeyProviderType.AZURE_KEYVAULT: + from libs.key_providers.azure_keyvault_key_provider import AzureKeyVaultKeyProvider + + return AzureKeyVaultKeyProvider + case _: + raise ValueError(f"unsupported key provider type {provider_type}") + + +key_provider_manager = KeyProviderManager() + + +def init_app(app: DifyApp): + key_provider_manager.init_app(app) diff --git a/api/extensions/ext_login.py b/api/extensions/ext_login.py index fddefb14f52..4b91a2320a3 100644 --- a/api/extensions/ext_login.py +++ b/api/extensions/ext_login.py @@ -15,7 +15,12 @@ from core.db.session_factory import session_factory from core.logging.context import set_identity_context from dify_app import DifyApp from libs.passport import PassportService -from libs.token import extract_access_token, extract_console_cookie_token, extract_webapp_passport +from libs.token import ( + extract_access_token, + extract_console_cookie_token, + extract_webapp_passport, + is_admin_api_key_request, +) from models import Account, Tenant, TenantAccountJoin from models.enums import EndUserType from models.model import AppMCPServer, EndUser @@ -66,23 +71,22 @@ def _load_user_from_request(request_from_flask_login: Request, session: Session) auth_token = extract_access_token(request) # Check for admin API key authentication first - if dify_config.ADMIN_API_KEY_ENABLE and auth_token: - admin_api_key = dify_config.ADMIN_API_KEY - if admin_api_key and admin_api_key == auth_token: - workspace_id = request.headers.get("X-WORKSPACE-ID") - if workspace_id: - tenant_account_join = session.execute( - select(Tenant, TenantAccountJoin) - .where(Tenant.id == workspace_id) - .where(TenantAccountJoin.tenant_id == Tenant.id) - .where(TenantAccountJoin.role == "owner") - ).one_or_none() - if tenant_account_join: - tenant, ta = tenant_account_join - account = session.scalar(select(Account).where(Account.id == ta.account_id)) - if account: - account.set_current_tenant_with_session(tenant, session=session) - return account + if is_admin_api_key_request(request): + workspace_id = request.headers.get("X-WORKSPACE-ID") + if workspace_id: + tenant_account_join = session.execute( + select(Tenant, TenantAccountJoin).where( + Tenant.id == workspace_id, + TenantAccountJoin.tenant_id == Tenant.id, + TenantAccountJoin.role == "owner", + ) + ).one_or_none() + if tenant_account_join: + tenant, ta = tenant_account_join + account = session.scalar(select(Account).where(Account.id == ta.account_id)) + if account: + account.set_current_tenant_with_session(tenant, session=session) + return account if request.blueprint in {"console", "inner_api"}: if not auth_token: diff --git a/api/extensions/ext_otel.py b/api/extensions/ext_otel.py index 63edbe93e79..ed607847b47 100644 --- a/api/extensions/ext_otel.py +++ b/api/extensions/ext_otel.py @@ -59,7 +59,7 @@ def init_app(app: DifyApp): SERVICE_NAME: dify_config.APPLICATION_NAME, SERVICE_VERSION: f"dify-{dify_config.project.version}-{dify_config.COMMIT_SHA}", PROCESS_PID: os.getpid(), - DEPLOYMENT_ENVIRONMENT_NAME: f"{dify_config.DEPLOY_ENV}-{dify_config.EDITION}", + DEPLOYMENT_ENVIRONMENT_NAME: f"{dify_config.DEPLOY_ENV}-{dify_config.DEPLOYMENT_EDITION.value}", HOST_NAME: socket.gethostname(), HOST_ARCH: platform.machine(), "custom.deployment.git_commit": dify_config.COMMIT_SHA, diff --git a/api/extensions/ext_sentry.py b/api/extensions/ext_sentry.py index 69d1f1ab07e..19a153c6f78 100644 --- a/api/extensions/ext_sentry.py +++ b/api/extensions/ext_sentry.py @@ -1,5 +1,6 @@ from configs import dify_config from dify_app import DifyApp +from enums import DeploymentEdition def init_app(app: DifyApp): @@ -7,6 +8,7 @@ def init_app(app: DifyApp): import sentry_sdk from sentry_sdk.integrations.celery import CeleryIntegration from sentry_sdk.integrations.flask import FlaskIntegration + from sentry_sdk.integrations.logging import ignore_logger from werkzeug.exceptions import HTTPException from graphon.model_runtime.errors.invoke import InvokeRateLimitError @@ -45,3 +47,12 @@ def init_app(app: DifyApp): release=f"dify-{dify_config.project.version}-{dify_config.COMMIT_SHA}", before_send=before_send, ) + + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: + # Cloud only. `opentelemetry.context.detach()` catches its own failures and reports + # them through `logger.exception`, so a double-detach surfaces as an error event + # rather than as a raised exception. Under gevent this can fire about once per + # workflow run, which is enough volume to crowd real errors out of the issue stream. + # Self-hosted deployments rarely see enough of it to be worth silencing, and the + # records still reach the normal logging handlers either way. + ignore_logger("opentelemetry.context") diff --git a/api/extensions/ext_socketio.py b/api/extensions/ext_socketio.py index 887734c8b8d..78ff7576287 100644 --- a/api/extensions/ext_socketio.py +++ b/api/extensions/ext_socketio.py @@ -20,9 +20,14 @@ def _get_ssl_cert_reqs() -> ssl.VerifyMode: def _build_redis_options(redis_url: str) -> dict[str, Any]: - """Build Redis options for Socket.IO's cross-process pub/sub manager.""" + """Build Redis options for Socket.IO's cross-process pub/sub manager. + + Note: ``socket_timeout`` is intentionally omitted. The RedisManager runs a + blocking ``pubsub.listen()`` loop that idles indefinitely between messages; + applying a read timeout there causes a reconnect storm (issue #39423). + ``socket_connect_timeout`` still guards connection establishment. + """ options: dict[str, Any] = { - "socket_timeout": dify_config.REDIS_SOCKET_TIMEOUT, "socket_connect_timeout": dify_config.REDIS_SOCKET_CONNECT_TIMEOUT, "health_check_interval": dify_config.REDIS_HEALTH_CHECK_INTERVAL, "protocol": dify_config.REDIS_SERIALIZATION_PROTOCOL, diff --git a/api/extensions/logstore/repositories/logstore_workflow_node_execution_repository.py b/api/extensions/logstore/repositories/logstore_workflow_node_execution_repository.py index 1c62df3746d..fa11e210c05 100644 --- a/api/extensions/logstore/repositories/logstore_workflow_node_execution_repository.py +++ b/api/extensions/logstore/repositories/logstore_workflow_node_execution_repository.py @@ -277,6 +277,12 @@ class LogstoreWorkflowNodeExecutionRepository(WorkflowNodeExecutionRepository): logger.exception("Failed to dual-write node execution to SQL database: id=%s", execution.id) # Don't raise - LogStore write succeeded, SQL is just a backup + @override + def save_synchronously(self, execution: WorkflowNodeExecution) -> None: + """Create the SQL caller row required by Agent v2 participant ownership.""" + + self.sql_repository.save_synchronously(execution) + @override def save_execution_data(self, execution: WorkflowNodeExecution) -> None: """ diff --git a/api/extensions/otel/instrumentation.py b/api/extensions/otel/instrumentation.py index 1718f9eb65e..1180e658ca9 100644 --- a/api/extensions/otel/instrumentation.py +++ b/api/extensions/otel/instrumentation.py @@ -37,10 +37,7 @@ class SupportsFlaskInstrumentor(Protocol): # pyrefly infers `NoneType`. Narrow the instances to just the methods we use # while leaving runtime behavior unchanged. def _new_celery_instrumentor() -> SupportsInstrument: - return cast( - SupportsInstrument, - CeleryInstrumentor(tracer_provider=get_tracer_provider(), meter_provider=get_meter_provider()), - ) + return cast(SupportsInstrument, CeleryInstrumentor()) def _new_httpx_instrumentor() -> SupportsInstrument: @@ -155,7 +152,9 @@ def init_httpx_instrumentor() -> None: def init_instruments(app: DifyApp) -> None: if not is_celery_worker(): init_flask_instrumentor(app) - _new_celery_instrumentor().instrument() + _new_celery_instrumentor().instrument( + tracer_provider=get_tracer_provider(), meter_provider=get_meter_provider() + ) instrument_exception_logging() init_sqlalchemy_instrumentor(app) diff --git a/api/extensions/otel/parser/base.py b/api/extensions/otel/parser/base.py index 621541d1ccb..fc10f183c99 100644 --- a/api/extensions/otel/parser/base.py +++ b/api/extensions/otel/parser/base.py @@ -3,7 +3,7 @@ Base parser interface and utilities for OpenTelemetry node parsers. Content gating: ``should_include_content()`` controls whether content-bearing span attributes (inputs, outputs, prompts, completions, documents) are written. -Gate is only active in EE (``ENTERPRISE_ENABLED=True``) when +Gate is only active in the Enterprise edition when ``ENTERPRISE_INCLUDE_CONTENT=False``; CE behaviour is unchanged. """ @@ -15,6 +15,7 @@ from opentelemetry.trace.status import Status, StatusCode from pydantic import BaseModel from configs import dify_config +from enums import DeploymentEdition from extensions.otel.semconv.gen_ai import ChainAttributes, GenAIAttributes from graphon.enums import BuiltinNodeTypes from graphon.file import File @@ -26,9 +27,9 @@ from graphon.variables import Segment def should_include_content() -> bool: """Return True if content should be written to spans. - CE (ENTERPRISE_ENABLED=False): always True — no behaviour change. + Community and Cloud editions: always True — no behaviour change. """ - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return True return dify_config.ENTERPRISE_INCLUDE_CONTENT diff --git a/api/extensions/otel/runtime.py b/api/extensions/otel/runtime.py index 4f749a49e5d..9f80fef355a 100644 --- a/api/extensions/otel/runtime.py +++ b/api/extensions/otel/runtime.py @@ -98,7 +98,7 @@ def init_celery_worker(*args, **kwargs): metric_provider = get_meter_provider() if dify_config.DEBUG: logger.info("Initializing OpenTelemetry for Celery worker") - CeleryInstrumentor(tracer_provider=tracer_provider, meter_provider=metric_provider).instrument() + CeleryInstrumentor().instrument(tracer_provider=tracer_provider, meter_provider=metric_provider) setup_celery_sqlcommenter() diff --git a/api/factories/variable_factory.py b/api/factories/variable_factory.py index 586430c54bc..f02f4272608 100644 --- a/api/factories/variable_factory.py +++ b/api/factories/variable_factory.py @@ -9,6 +9,10 @@ from collections.abc import Mapping, Sequence from typing import Any, cast from configs import dify_config +from core.workflow.llm_environment_variable import ( + LLM_ENVIRONMENT_VARIABLE_VALUE_TYPE, + LLMEnvironmentVariable, +) from core.workflow.variable_prefixes import ( CONVERSATION_VARIABLE_NODE_ID, ENVIRONMENT_VARIABLE_NODE_ID, @@ -66,6 +70,17 @@ def build_conversation_variable_from_mapping(mapping: Mapping[str, Any], /) -> V def build_environment_variable_from_mapping(mapping: Mapping[str, Any], /) -> VariableBase: if not mapping.get("name"): raise VariableError("missing name") + if mapping.get("value_type") == LLM_ENVIRONMENT_VARIABLE_VALUE_TYPE: + llm_mapping = dict(mapping) + llm_mapping["value_type"] = SegmentType.OBJECT + llm_mapping["selector"] = [ENVIRONMENT_VARIABLE_NODE_ID, mapping["name"]] + try: + result = LLMEnvironmentVariable.model_validate(llm_mapping) + except ValueError as exc: + raise VariableError(f"invalid LLM environment variable: {exc}") from exc + if result.size > dify_config.MAX_VARIABLE_SIZE: + raise VariableError(f"variable size {result.size} exceeds limit {dify_config.MAX_VARIABLE_SIZE}") + return result return _build_variable_from_mapping(mapping=mapping, selector=[ENVIRONMENT_VARIABLE_NODE_ID, mapping["name"]]) diff --git a/api/fields/agent_fields.py b/api/fields/agent_fields.py index 50105273e1c..0f1542bc256 100644 --- a/api/fields/agent_fields.py +++ b/api/fields/agent_fields.py @@ -171,6 +171,12 @@ class AgentLogConversationItemResponse(ResponseModel): return to_timestamp(value) +class AgentLogFeedbackResponse(ResponseModel): + rating: Literal["like", "dislike"] + content: str | None = None + from_source: Literal["user", "admin"] + + class AgentLogMessageItemResponse(ResponseModel): id: str message_id: str @@ -181,6 +187,8 @@ class AgentLogMessageItemResponse(ResponseModel): error: str | None = None from_end_user_id: str | None = None from_account_id: str | None = None + feedback_enabled: bool = False + feedbacks: list[AgentLogFeedbackResponse] = Field(default_factory=list) message_tokens: int answer_tokens: int total_tokens: int @@ -383,6 +391,7 @@ class AgentAppComposerResponse(ResponseModel): variant: Literal[ComposerVariant.AGENT_APP] agent: AgentComposerAgentResponse active_config_snapshot: AgentConfigSnapshotSummaryResponse | None = None + active_config_is_published: bool draft: AgentConfigDraftSummaryResponse | None = None agent_soul: AgentSoulConfig save_options: list[ComposerSaveStrategy] diff --git a/api/fields/conversation_fields.py b/api/fields/conversation_fields.py index 073305d2dd9..7612d09f58d 100644 --- a/api/fields/conversation_fields.py +++ b/api/fields/conversation_fields.py @@ -34,7 +34,7 @@ class _SessionResponseSource[SourceT]: self._session = session def __getattr__(self, name: str) -> object: - return getattr(self._source, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self._source, name) # guard-ignore: no-new-getattr -- delegates model fields class _FeedbackResponseSource(_SessionResponseSource[MessageFeedback]): diff --git a/api/fields/dataset_fields.py b/api/fields/dataset_fields.py index 4846aa9689c..c81fb79df3f 100644 --- a/api/fields/dataset_fields.py +++ b/api/fields/dataset_fields.py @@ -227,7 +227,7 @@ class DatasetDetailResponseSource: return self.dataset.get_total_available_documents(session=self.session) def __getattr__(self, name: str) -> Any: - return getattr(self.dataset, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self.dataset, name) # guard-ignore: no-new-getattr -- delegates model fields def dataset_detail_response_source(dataset: Any, *, session: Session) -> DatasetDetailResponseSource: diff --git a/api/fields/document_fields.py b/api/fields/document_fields.py index aa3b4135ec6..64de3f00bb4 100644 --- a/api/fields/document_fields.py +++ b/api/fields/document_fields.py @@ -90,7 +90,7 @@ class DocumentWithSession: return self.document.get_doc_metadata_details(session=self.session) def __getattr__(self, name: str) -> Any: - return getattr(self.document, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self.document, name) # guard-ignore: no-new-getattr -- delegates model fields def document_response(document: Document, *, session: Session) -> DocumentResponse: @@ -121,6 +121,9 @@ class DocumentStatusResponse(ResponseModel): completed_at: int | None paused_at: int | None error: str | None + error_code: str | None = None + estimated_vector_space_mb: int | None = None + vector_space_limit_mb: int | None = None stopped_at: int | None completed_segments: int | None = None total_segments: int | None = None diff --git a/api/fields/file_fields.py b/api/fields/file_fields.py index 681dc3db2d2..094f6895bec 100644 --- a/api/fields/file_fields.py +++ b/api/fields/file_fields.py @@ -10,11 +10,13 @@ from libs.helper import to_timestamp class UploadConfig(ResponseModel): file_size_limit: int + knowledge_file_size_limit: int batch_count_limit: int file_upload_limit: int image_file_size_limit: int video_file_size_limit: int audio_file_size_limit: int + skill_file_size_limit: int workflow_file_upload_limit: int image_file_batch_limit: int single_chunk_attachment_limit: int diff --git a/api/libs/broadcast_channel/redis/_subscription.py b/api/libs/broadcast_channel/redis/_subscription.py index 01a9e668bcc..c155b9e1973 100644 --- a/api/libs/broadcast_channel/redis/_subscription.py +++ b/api/libs/broadcast_channel/redis/_subscription.py @@ -117,9 +117,12 @@ class RedisSubscriptionBase(Subscription): ) continue - self._enqueue_message(payload_bytes) if payload_bytes == SIG_CLOSE: - break + # Close signals are broadcast to every subscriber on the topic. + # The closing subscription is already handled by the _closed check above. + continue + + self._enqueue_message(payload_bytes) _logger.debug("%s listener thread stopped for channel %s", self._get_subscription_type().title(), self._topic) try: diff --git a/api/libs/broadcast_channel/redis/streams_channel.py b/api/libs/broadcast_channel/redis/streams_channel.py index b3385b05388..c86d77ada6b 100644 --- a/api/libs/broadcast_channel/redis/streams_channel.py +++ b/api/libs/broadcast_channel/redis/streams_channel.py @@ -128,11 +128,17 @@ class _StreamsSubscription(Subscription): data_bytes = data.encode() case bytes() | bytearray(): data_bytes = bytes(data) - if data_bytes is not None: - if data_bytes == SIG_CLOSE: - break - self._queue.put_nowait(data_bytes) last_id = entry_id + if data_bytes is None: + continue + if data_bytes == SIG_CLOSE: + # Close signals share the stream with normal events. Ignore signals + # emitted by another subscription while this one is still open. + with self._lock: + if self._closed: + break + continue + self._queue.put_nowait(data_bytes) finally: self._queue.put_nowait(self._SENTINEL) with self._lock: diff --git a/api/libs/datetime_utils.py b/api/libs/datetime_utils.py index e0a6ec2cacd..d962d81e78a 100644 --- a/api/libs/datetime_utils.py +++ b/api/libs/datetime_utils.py @@ -35,6 +35,15 @@ def ensure_naive_utc(dt: datetime.datetime) -> datetime.datetime: return dt.astimezone(datetime.UTC).replace(tzinfo=None) +def to_utc_timestamp(dt: datetime.datetime) -> int: + """Convert a datetime to Unix epoch seconds, assuming naive values are UTC. + + Persisted datetimes may be returned without timezone information. Treat + those values as UTC instead of interpreting them in the host timezone. + """ + return int(ensure_naive_utc(dt).replace(tzinfo=datetime.UTC).timestamp()) + + def parse_time_range( start: str | None, end: str | None, tzname: str ) -> tuple[datetime.datetime | None, datetime.datetime | None]: diff --git a/api/libs/device_flow_security.py b/api/libs/device_flow_security.py index 9f4c1f56f66..c10d3daaab6 100644 --- a/api/libs/device_flow_security.py +++ b/api/libs/device_flow_security.py @@ -17,7 +17,8 @@ from werkzeug.exceptions import NotFound from libs import jws from libs.token import is_secure -from services.feature_service import FeatureService, LicenseStatus +from services.entities.feature_entities import LicenseStatus +from services.feature_service import FeatureService logger = logging.getLogger(__name__) diff --git a/api/libs/email_i18n.py b/api/libs/email_i18n.py index 1519f07bb1b..606dd9cfde0 100644 --- a/api/libs/email_i18n.py +++ b/api/libs/email_i18n.py @@ -16,7 +16,8 @@ from flask import render_template from pydantic import BaseModel, Field from extensions.ext_mail import mail -from services.feature_service import BrandingModel, FeatureService +from services.entities.feature_entities import BrandingModel +from services.feature_service import FeatureService class EmailType(StrEnum): diff --git a/api/libs/helper.py b/api/libs/helper.py index 752342bfad6..7610b834d31 100644 --- a/api/libs/helper.py +++ b/api/libs/helper.py @@ -16,7 +16,7 @@ from zoneinfo import available_timezones from flask import Request, Response, stream_with_context from flask_restx import fields -from pydantic import BaseModel, ConfigDict, TypeAdapter, with_config +from pydantic import BaseModel, ConfigDict, TypeAdapter, WithJsonSchema, with_config from pydantic.functional_validators import AfterValidator from typing_extensions import TypedDict @@ -284,12 +284,20 @@ def _strict_uuid(value: str | UUID) -> str: raise ValueError("must be a valid UUID") from exc -UUIDStr = Annotated[str, AfterValidator(_strict_uuid)] +UUIDStr = Annotated[ + str, + AfterValidator(_strict_uuid), + WithJsonSchema({"format": "uuid", "type": "string"}), +] def alphanumeric(value: str): # check if the value is alphanumeric and underlined - if re.match(r"^[a-zA-Z0-9_]+$", value): + # 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 "tool_name\n". re.fullmatch requires the entire + # string to match. Regression for #39666 (sibling of #39234 / #39548). + if re.fullmatch(r"^[a-zA-Z0-9_]+$", value): return value raise ValueError(f"{value} is not a valid alphanumeric value") @@ -391,7 +399,7 @@ def extract_remote_ip(request: Request) -> str: if request.headers.get("CF-Connecting-IP"): return cast(str, request.headers.get("CF-Connecting-IP")) elif request.headers.getlist("X-Forwarded-For"): - return cast(str, request.headers.getlist("X-Forwarded-For")[0]) + return request.headers.getlist("X-Forwarded-For")[0] else: return cast(str, request.remote_addr) diff --git a/api/libs/key_providers/__init__.py b/api/libs/key_providers/__init__.py new file mode 100644 index 00000000000..6c6ed5d8091 --- /dev/null +++ b/api/libs/key_providers/__init__.py @@ -0,0 +1,14 @@ +from libs.key_providers.base import BaseKeyProvider + +__all__ = ["BaseKeyProvider", "generate_key_pair"] + + +def generate_key_pair(tenant_id: str) -> str: + """ + Provision the tenant credential encryption key using the configured KEY_PROVIDER_TYPE. + + Returns the opaque reference to be stored in Tenant.encrypt_public_key. + """ + from extensions.ext_key_provider import key_provider_manager + + return key_provider_manager.provider.generate_key_pair(tenant_id) diff --git a/api/libs/key_providers/azure_keyvault_key_provider.py b/api/libs/key_providers/azure_keyvault_key_provider.py new file mode 100644 index 00000000000..4d5979f301d --- /dev/null +++ b/api/libs/key_providers/azure_keyvault_key_provider.py @@ -0,0 +1,203 @@ +import json +from datetime import UTC, datetime +from typing import Any, override + +from azure.core.exceptions import AzureError +from azure.identity import DefaultAzureCredential +from azure.keyvault.keys import ( + KeyClient, + KeyRotationLifetimeAction, + KeyRotationPolicy, + KeyRotationPolicyAction, +) +from azure.keyvault.keys.crypto import CryptographyClient, KeyWrapAlgorithm +from Crypto.Cipher import AES +from Crypto.Random import get_random_bytes + +from configs import dify_config +from libs.key_providers.base import BaseKeyProvider + +# Marker kept identical to libs/rsa.py so ciphertext produced by either provider is +# self-describing, even though the two providers never decode each other's payloads. +_PREFIX = b"HYBRID:" + +# Bump this if the *binary layout* below ever changes. Adding new metadata fields does NOT +# require a bump (metadata is a JSON object -- old readers just ignore unknown keys via .get()). +_ENVELOPE_VERSION = 1 + +_DEFAULT_WRAP_ALGORITHM = KeyWrapAlgorithm.rsa_oaep_256 + + +class AzureKeyVaultKeyProvider(BaseKeyProvider): + """ + Envelope-encryption key provider backed by Azure Key Vault. + + Ciphertext envelope (self-describing, forward-compatible): + PREFIX (7 bytes) + + envelope_version (1 byte) + + metadata_len (2 bytes, big-endian) + metadata (UTF-8 JSON object) + + wrapped_key_len (2 bytes, big-endian) + wrapped_key + + nonce (16 bytes) + tag (16 bytes) + ciphertext + + `metadata` currently carries {"key_version": ..., "wrap_alg": ...}. It's a JSON object + rather than fixed-width fields so new attributes can be added later without touching the + binary layout or breaking old ciphertext; `envelope_version` exists separately to guard + the binary layout itself, in case that ever needs to change. + + Recording the key_version that wrapped each token (rather than always resolving "the + current version" at decrypt time) is what makes Key Vault's native automatic key rotation + safe to use here: old tokens keep decrypting against the version that encrypted them, + while new tokens pick up whatever version is current. This only holds as long as old + versions are never allowed to *expire* -- see generate_key_pair(). + """ + + def __init__(self): + vault_url = dify_config.AZURE_KEYVAULT_VAULT_URL + if not vault_url: + raise ValueError("AZURE_KEYVAULT_VAULT_URL must be configured when KEY_PROVIDER_TYPE=azure-keyvault") + + self._vault_url = vault_url + self._credential = DefaultAzureCredential() + self._key_client = KeyClient(vault_url=vault_url, credential=self._credential) + + @staticmethod + def _key_name(tenant_id: str) -> str: + return f"dify-tenant-{tenant_id}" + + def _get_crypto_client(self, tenant_id: str, version: str | None = None) -> tuple[CryptographyClient, str]: + """ + Return a CryptographyClient bound to `version`, along with the resolved version string + that was actually used. + """ + key_name = self._key_name(tenant_id) + if version is not None: + resolved_version = version + else: + versions = list(self._key_client.list_properties_of_key_versions(key_name)) + if not versions: + raise ValueError(f"No key versions found for key {key_name}") + current = max( + versions, + key=lambda properties: properties.created_on or datetime.min.replace(tzinfo=UTC), + ) + resolved_version = current.version or "" + return self._key_client.get_cryptography_client(key_name, key_version=resolved_version), resolved_version + + @override + def generate_key_pair(self, tenant_id: str) -> str: + key_name = self._key_name(tenant_id) + self._key_client.create_rsa_key(key_name, size=dify_config.AZURE_KEYVAULT_KEY_SIZE) + + rotation_interval_days = dify_config.AZURE_KEYVAULT_ROTATION_INTERVAL_DAYS + if rotation_interval_days: + self._key_client.update_key_rotation_policy( + key_name, + policy=KeyRotationPolicy( + lifetime_actions=[ + KeyRotationLifetimeAction( + KeyRotationPolicyAction.rotate, + time_after_create=f"P{rotation_interval_days}D", + ) + ], + # Deliberately no `expires_in` here: this provider pins each ciphertext to + # the key_version that encrypted it and relies on old versions staying + # usable forever. If versions were also given an expiry (time_before_expiry + # trigger / expires_in), old ciphertext would become permanently + # undecryptable once its version expired, unless a separate re-wrap/ + # migration job proactively moves it to the new version first. + ), + ) + return key_name + + @override + def encrypt(self, tenant_id: str, text: str) -> bytes: + aes_key = get_random_bytes(16) + cipher_aes = AES.new(aes_key, AES.MODE_EAX) + ciphertext, tag = cipher_aes.encrypt_and_digest(text.encode()) + + crypto_client, key_version = self._get_crypto_client(tenant_id) + wrapped_key = crypto_client.wrap_key(_DEFAULT_WRAP_ALGORITHM, aes_key).encrypted_key + + metadata = json.dumps({"key_version": key_version, "wrap_alg": _DEFAULT_WRAP_ALGORITHM.value}).encode() + + return ( + _PREFIX + + _ENVELOPE_VERSION.to_bytes(1, "big") + + len(metadata).to_bytes(2, "big") + + metadata + + len(wrapped_key).to_bytes(2, "big") + + wrapped_key + + cipher_aes.nonce + + tag + + ciphertext + ) + + @override + def get_decrypt_decoding(self, tenant_id: str) -> str: + return tenant_id + + @override + def decrypt_with_decoding(self, encrypted_text: bytes, decoding: str) -> str: + tenant_id = decoding + if not encrypted_text.startswith(_PREFIX): + raise ValueError("Unsupported ciphertext format for Azure Key Vault key provider") + + # Bytes slicing never raises on out-of-range indices in Python (it just returns a + # shorter/empty slice), so a truncated envelope wouldn't otherwise surface as an error + # until (maybe) AES decryption fails much later, or not at all. Validate lengths + # explicitly and turn any parsing failure into ValueError, matching what callers + # (e.g. core/provider_manager.py) already expect and suppress for malformed credentials. + try: + body = encrypted_text[len(_PREFIX) :] + if len(body) < 1: + raise ValueError("Malformed Azure Key Vault envelope: missing envelope version") + envelope_version = body[0] + if envelope_version != _ENVELOPE_VERSION: + raise ValueError(f"Unsupported Azure Key Vault envelope version: {envelope_version}") + offset = 1 + + if len(body) < offset + 2: + raise ValueError("Malformed Azure Key Vault envelope: truncated metadata length") + metadata_len = int.from_bytes(body[offset : offset + 2], "big") + offset += 2 + if len(body) < offset + metadata_len: + raise ValueError("Malformed Azure Key Vault envelope: truncated metadata") + metadata: Any = json.loads(body[offset : offset + metadata_len]) + if not isinstance(metadata, dict): + raise ValueError("Malformed Azure Key Vault envelope: metadata is not a JSON object") + offset += metadata_len + + if len(body) < offset + 2: + raise ValueError("Malformed Azure Key Vault envelope: truncated wrapped key length") + key_len = int.from_bytes(body[offset : offset + 2], "big") + offset += 2 + if len(body) < offset + key_len + 16 + 16: + raise ValueError("Malformed Azure Key Vault envelope: truncated wrapped key/nonce/tag") + wrapped_key = body[offset : offset + key_len] + offset += key_len + nonce = body[offset : offset + 16] + offset += 16 + tag = body[offset : offset + 16] + offset += 16 + ciphertext = body[offset:] + except (IndexError, TypeError) as exc: + raise ValueError("Malformed Azure Key Vault envelope") from exc + + wrap_alg = KeyWrapAlgorithm(metadata["wrap_alg"]) if metadata.get("wrap_alg") else _DEFAULT_WRAP_ALGORITHM + try: + # A specific key_version can legitimately become unusable after this ciphertext was + # created -- disabled, deleted, or (if a rotation policy with an expiry was + # misconfigured despite generate_key_pair()'s warning against it) expired. Every + # caller of decrypt_token_with_decoding (core/provider_manager.py, + # services/model_load_balancing_service.py) already only expects/suppresses + # ValueError for "this particular credential can't be decrypted right now", so Azure + # SDK errors must be translated here rather than left to escape as a different type + # and crash the whole call chain (e.g. building a tenant's full provider + # configuration just to create an unrelated new credential). + crypto_client, _ = self._get_crypto_client(tenant_id, version=metadata.get("key_version")) + aes_key = crypto_client.unwrap_key(wrap_alg, wrapped_key).key + except AzureError as exc: + raise ValueError(f"Failed to unwrap credential via Azure Key Vault: {exc}") from exc + + cipher_aes = AES.new(aes_key, AES.MODE_EAX, nonce=nonce) + return cipher_aes.decrypt_and_verify(ciphertext, tag).decode() diff --git a/api/libs/key_providers/base.py b/api/libs/key_providers/base.py new file mode 100644 index 00000000000..1cb9ce3625e --- /dev/null +++ b/api/libs/key_providers/base.py @@ -0,0 +1,34 @@ +"""Abstract interface for tenant credential encryption key providers.""" + +from abc import ABC, abstractmethod +from typing import Any + + +class BaseKeyProvider(ABC): + """Interface for providers that manage the keys used to encrypt/decrypt tenant credentials.""" + + @abstractmethod + def generate_key_pair(self, tenant_id: str) -> str: + """ + Provision the encryption key for a tenant. + + Returns an opaque reference to be stored in Tenant.encrypt_public_key + (e.g. a PEM public key, or a key vault key name/identifier). + """ + raise NotImplementedError + + @abstractmethod + def encrypt(self, tenant_id: str, text: str) -> bytes: + raise NotImplementedError + + @abstractmethod + def get_decrypt_decoding(self, tenant_id: str) -> Any: + """Return a reusable decoding context, so batch decryption can avoid repeated key lookups.""" + raise NotImplementedError + + @abstractmethod + def decrypt_with_decoding(self, encrypted_text: bytes, decoding: Any) -> str: + raise NotImplementedError + + def decrypt(self, tenant_id: str, encrypted_text: bytes) -> str: + return self.decrypt_with_decoding(encrypted_text, self.get_decrypt_decoding(tenant_id)) diff --git a/api/libs/key_providers/key_provider_type.py b/api/libs/key_providers/key_provider_type.py new file mode 100644 index 00000000000..8ecdd047600 --- /dev/null +++ b/api/libs/key_providers/key_provider_type.py @@ -0,0 +1,6 @@ +from enum import StrEnum + + +class KeyProviderType(StrEnum): + LOCAL = "local" + AZURE_KEYVAULT = "azure-keyvault" diff --git a/api/libs/key_providers/rsa_key_provider.py b/api/libs/key_providers/rsa_key_provider.py new file mode 100644 index 00000000000..e0582dfb397 --- /dev/null +++ b/api/libs/key_providers/rsa_key_provider.py @@ -0,0 +1,47 @@ +from typing import override + +from Crypto.PublicKey import RSA + +from libs import rsa +from libs.key_providers.base import BaseKeyProvider + + +class RSAKeyProvider(BaseKeyProvider): + """ + Default key provider: per-tenant RSA key pair. + + The private key is kept in the configured STORAGE_TYPE backend (see libs/rsa.py). + This provider only composes the existing libs.rsa implementation; the underlying + crypto logic is intentionally left untouched. + """ + + @override + def generate_key_pair(self, tenant_id: str) -> str: + return rsa.generate_key_pair(tenant_id) + + @override + def encrypt(self, tenant_id: str, text: str) -> bytes: + from models.account import Tenant + from models.engine import db + + if not (tenant := db.session.get(Tenant, tenant_id)): + raise ValueError(f"Tenant with id {tenant_id} not found") + if tenant.encrypt_public_key is None: + raise ValueError(f"Tenant with id {tenant_id} has no encrypt_public_key") + return rsa.encrypt(text, tenant.encrypt_public_key) + + @override + def get_decrypt_decoding(self, tenant_id: str) -> tuple[RSA.RsaKey, object]: + return rsa.get_decrypt_decoding(tenant_id) + + @override + def decrypt_with_decoding(self, encrypted_text: bytes, decoding: tuple[RSA.RsaKey, object]) -> str: + rsa_key, cipher_rsa = decoding + return rsa.decrypt_token_with_decoding(encrypted_text, rsa_key, cipher_rsa) + + @override + def decrypt(self, tenant_id: str, encrypted_text: bytes) -> str: + # Overrides BaseKeyProvider's generic get_decrypt_decoding()+decrypt_with_decoding() + # composition to call libs.rsa.decrypt() directly (a single-shot equivalent), matching + # this provider's one supported decrypt path in libs/rsa.py. + return rsa.decrypt(encrypted_text, tenant_id) diff --git a/api/libs/login.py b/api/libs/login.py index bbb8ba1611c..bf89ab75123 100644 --- a/api/libs/login.py +++ b/api/libs/login.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Callable from functools import wraps -from typing import TYPE_CHECKING, Any, Concatenate, cast, overload +from typing import TYPE_CHECKING, Any, Concatenate, NamedTuple, cast, overload from flask import Response, current_app, g, has_request_context, request from flask_login.config import EXEMPT_METHODS @@ -19,6 +19,13 @@ if TYPE_CHECKING: from models.model import EndUser +class AccountWithTenant(NamedTuple): + """Authenticated account and its active tenant.""" + + account: Account + tenant_id: str + + def _resolve_current_user() -> EndUser | Account | None: """ Resolve the current user proxy to its underlying user object. @@ -36,7 +43,7 @@ def _get_login_manager() -> DifyLoginManager: return app.login_manager -def current_account_with_tenant() -> tuple[Account, str]: +def current_account_with_tenant() -> AccountWithTenant: """ Resolve the underlying account for the current user proxy and ensure tenant context exists. Allows tests to supply plain Account mocks without the LocalProxy helper. @@ -46,7 +53,7 @@ def current_account_with_tenant() -> tuple[Account, str]: if not isinstance(user, Account): raise ValueError("current_user must be an Account instance") assert user.current_tenant_id is not None, "The tenant information should be loaded." - return user, user.current_tenant_id + return AccountWithTenant(account=user, tenant_id=user.current_tenant_id) def current_account_with_tenant_optional() -> tuple[Account | None, str | None]: @@ -67,7 +74,7 @@ def resolve_account_fallback( current_tenant_id: str | None = None, *, fallback_tenant_id: str | None = None, -) -> tuple[Account, str]: +) -> AccountWithTenant: """ If the provided current user and tenant ID is None, fallback to current_account_with_tenant. This is useful for those service layers whose controllers are not migrated to use DI for @@ -79,7 +86,7 @@ def resolve_account_fallback( tenant_id = current_tenant_id or fallback_tenant_id if tenant_id is None: raise ValueError("current_tenant_id is required when current_user is provided.") - return current_user, tenant_id + return AccountWithTenant(account=current_user, tenant_id=tenant_id) return current_account_with_tenant() @@ -92,8 +99,7 @@ def resolve_tenant_id_fallback(current_tenant_id: str | None = None) -> str: """ if current_tenant_id is not None: return current_tenant_id - _, tenant_id = current_account_with_tenant() - return tenant_id + return current_account_with_tenant().tenant_id @overload diff --git a/api/libs/oauth_bearer.py b/api/libs/oauth_bearer.py index 36de4b85ae0..ed17503bb3b 100644 --- a/api/libs/oauth_bearer.py +++ b/api/libs/oauth_bearer.py @@ -25,6 +25,7 @@ from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, ServiceUnavailable, Unauthorized from configs import dify_config +from enums import DeploymentEdition from extensions.ext_database import db from extensions.ext_redis import redis_client from libs.rate_limit import enforce_bearer_rate_limit @@ -531,7 +532,7 @@ def require_workspace_member(ctx: AuthContext, tenant_id: str) -> None: No-op on EE (gateway RBAC owns tenant isolation) and for SSO subjects (no `tenant_account_joins` row by definition). """ - if dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE: return if ctx.subject_type != SubjectType.ACCOUNT or ctx.account_id is None: return diff --git a/api/libs/rsa.py b/api/libs/rsa.py index c72032701f0..c401129b3a4 100644 --- a/api/libs/rsa.py +++ b/api/libs/rsa.py @@ -1,3 +1,20 @@ +""" +Low-level implementation of the default ("local") tenant credential encryption key provider. + +Do NOT import this module directly to encrypt/decrypt tenant credentials. It only implements +one specific key provider (per-tenant RSA key pair, private key kept in STORAGE_TYPE). Other +KEY_PROVIDER_TYPE options (e.g. 'azure-keyvault') are not implemented here. + +Instead use: + - core.helper.encrypter (encrypt_token / decrypt_token / batch_decrypt_token / ...) for + application code that needs to encrypt or decrypt tenant credentials. + - libs.key_providers (generate_key_pair) when provisioning the key for a new tenant. + +This module is only meant to be imported by libs.key_providers.rsa_key_provider.RSAKeyProvider, +which the rest of the codebase should reach through extensions.ext_key_provider.key_provider_manager. +This is enforced by the "no-direct-rsa-imports" contract in .importlinter (run via `make lint`). +""" + import hashlib from typing import Union diff --git a/api/libs/token.py b/api/libs/token.py index ce868dc3d8a..1b28b60a538 100644 --- a/api/libs/token.py +++ b/api/libs/token.py @@ -1,3 +1,4 @@ +import hmac import logging import re from datetime import UTC, datetime, timedelta @@ -84,6 +85,19 @@ def extract_access_token(request: Request) -> str | None: return extract_console_cookie_token(request) or _try_extract_from_header(request) +def is_admin_api_key_request(request: Request) -> bool: + """Return whether the request carries the configured admin API key as a bearer token. + + Admin API key authentication is header-only so an unrelated console session cookie + cannot shadow the bearer token used by server-to-server clients. + """ + admin_api_key = dify_config.ADMIN_API_KEY + bearer_token = _try_extract_from_header(request) + if not dify_config.ADMIN_API_KEY_ENABLE or not admin_api_key or not bearer_token: + return False + return hmac.compare_digest(bearer_token, admin_api_key) + + def extract_webapp_access_token(request: Request) -> str | None: return request.cookies.get(_real_cookie_name(COOKIE_NAME_WEBAPP_ACCESS_TOKEN)) or _try_extract_from_header(request) @@ -183,10 +197,8 @@ def build_force_logout_cookie_headers() -> list[str]: def check_csrf_token(request: Request, user_id: str): # some apis are sent by beacon, so we need to bypass csrf token check # since these APIs are post, they are already protected by SameSite: Lax, so csrf is not required. - if dify_config.ADMIN_API_KEY_ENABLE: - auth_token = extract_access_token(request) - if auth_token and auth_token == dify_config.ADMIN_API_KEY: - return + if is_admin_api_key_request(request): + return def _unauthorized(): raise Unauthorized("CSRF token is missing or invalid.") diff --git a/api/libs/workspace_permission.py b/api/libs/workspace_permission.py index 435b07dd6ea..969a81f5f9d 100644 --- a/api/libs/workspace_permission.py +++ b/api/libs/workspace_permission.py @@ -12,6 +12,7 @@ import logging from werkzeug.exceptions import Forbidden from configs import dify_config +from enums import DeploymentEdition from services.enterprise.enterprise_service import EnterpriseService from services.feature_service import FeatureService @@ -32,8 +33,8 @@ def check_workspace_member_invite_permission(workspace_id: str) -> None: Raises: Forbidden: If either billing plan or workspace policy prohibits member invitations """ - # Check enterprise workspace policy level (only if enterprise enabled) - if dify_config.ENTERPRISE_ENABLED: + # Check the enterprise workspace policy only in the Enterprise edition. + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE: try: permission = EnterpriseService.WorkspacePermissionService.get_permission(workspace_id) if not permission.allow_member_invite: @@ -62,8 +63,8 @@ def check_workspace_owner_transfer_permission(workspace_id: str) -> None: if not features.is_allow_transfer_workspace: raise Forbidden("Your current plan does not allow workspace ownership transfer") - # Check enterprise workspace policy level (only if enterprise enabled) - if dify_config.ENTERPRISE_ENABLED: + # Check the enterprise workspace policy only in the Enterprise edition. + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE: try: permission = EnterpriseService.WorkspacePermissionService.get_permission(workspace_id) if not permission.allow_owner_transfer: diff --git a/api/machinery/__init__.py b/api/machinery/__init__.py new file mode 100644 index 00000000000..6499b96c91c --- /dev/null +++ b/api/machinery/__init__.py @@ -0,0 +1 @@ +"""Framework-neutral API machinery.""" diff --git a/api/machinery/context.py b/api/machinery/context.py new file mode 100644 index 00000000000..bc462afaf68 --- /dev/null +++ b/api/machinery/context.py @@ -0,0 +1,10 @@ +"""Stable values passed from API admission into application services.""" + +from typing import NamedTuple + + +class RequestContext(NamedTuple): + request_id: str + trace_id: str | None + account_id: str + active_workspace_id: str | None diff --git a/api/machinery/errors.py b/api/machinery/errors.py new file mode 100644 index 00000000000..517d61fcb42 --- /dev/null +++ b/api/machinery/errors.py @@ -0,0 +1,19 @@ +"""Framework-neutral errors raised by API machinery. + +They deliberately do not carry Flask responses, Werkzeug exceptions, HTTP +status codes, or surface-specific wire models. +""" + +from typing import NamedTuple + + +class ErrorDetail(NamedTuple): + type: str + location: tuple[str | int, ...] + message: str + + +class MachineryError(Exception): + """Base class for failures owned by API machinery.""" + + code: str = "machinery_error" diff --git a/api/migrations/versions/2026_07_21_2251-2f39536b3feb_add_agent_home_snapshot_ledger.py b/api/migrations/versions/2026_07_21_2251-2f39536b3feb_add_agent_home_snapshot_ledger.py new file mode 100644 index 00000000000..7bb38f657c0 --- /dev/null +++ b/api/migrations/versions/2026_07_21_2251-2f39536b3feb_add_agent_home_snapshot_ledger.py @@ -0,0 +1,64 @@ +"""add agent home snapshot ledger + +Revision ID: 2f39536b3feb +Revises: 6f5a9c2d8e1b +Create Date: 2026-07-21 22:51:07.268658 + +""" +from alembic import op +import models as models +import sqlalchemy as sa + + +# revision identifiers, used by Alembic. +revision = '2f39536b3feb' +down_revision = '6f5a9c2d8e1b' +branch_labels = None +depends_on = None + + +def upgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.create_table('agent_home_snapshots', + sa.Column('id', models.types.StringUUID(), nullable=False), + sa.Column('tenant_id', models.types.StringUUID(), nullable=False), + sa.Column('agent_id', models.types.StringUUID(), nullable=False), + sa.Column('snapshot_ref', sa.String(length=255), nullable=False), + sa.Column('created_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False), + sa.PrimaryKeyConstraint('id', name='agent_home_snapshot_pkey') + ) + with op.batch_alter_table('agent_home_snapshots', schema=None) as batch_op: + batch_op.create_index('agent_home_snapshot_tenant_agent_idx', ['tenant_id', 'agent_id'], unique=False) + + with op.batch_alter_table('agent_config_drafts', schema=None) as batch_op: + batch_op.add_column(sa.Column('home_snapshot_id', models.types.StringUUID(), nullable=True)) + + with op.batch_alter_table('agent_config_snapshots', schema=None) as batch_op: + batch_op.add_column(sa.Column('home_snapshot_id', models.types.StringUUID(), nullable=True)) + + with op.batch_alter_table('agent_runtime_sessions', schema=None) as batch_op: + batch_op.add_column(sa.Column('home_snapshot_id', models.types.StringUUID(), nullable=True)) + batch_op.drop_index(batch_op.f('agent_runtime_session_conversation_scope_unique'), postgresql_where='(conversation_id IS NOT NULL)') + batch_op.create_index('agent_runtime_session_conversation_scope_unique', ['tenant_id', 'conversation_id', 'agent_id', 'agent_config_snapshot_id', 'home_snapshot_id'], unique=True, postgresql_where=sa.text('conversation_id IS NOT NULL')) + + # ### end Alembic commands ### + + +def downgrade(): + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table('agent_runtime_sessions', schema=None) as batch_op: + batch_op.drop_index('agent_runtime_session_conversation_scope_unique', postgresql_where=sa.text('conversation_id IS NOT NULL')) + batch_op.create_index(batch_op.f('agent_runtime_session_conversation_scope_unique'), ['tenant_id', 'conversation_id', 'agent_id', 'agent_config_snapshot_id'], unique=True, postgresql_where='(conversation_id IS NOT NULL)') + batch_op.drop_column('home_snapshot_id') + + with op.batch_alter_table('agent_config_snapshots', schema=None) as batch_op: + batch_op.drop_column('home_snapshot_id') + + with op.batch_alter_table('agent_config_drafts', schema=None) as batch_op: + batch_op.drop_column('home_snapshot_id') + + with op.batch_alter_table('agent_home_snapshots', schema=None) as batch_op: + batch_op.drop_index('agent_home_snapshot_tenant_agent_idx') + + op.drop_table('agent_home_snapshots') + # ### end Alembic commands ### diff --git a/api/migrations/versions/2026_07_23_0203-f6e4c5686857_replace_agent_runtime_sessions_with_.py b/api/migrations/versions/2026_07_23_0203-f6e4c5686857_replace_agent_runtime_sessions_with_.py new file mode 100644 index 00000000000..aed7bc93215 --- /dev/null +++ b/api/migrations/versions/2026_07_23_0203-f6e4c5686857_replace_agent_runtime_sessions_with_.py @@ -0,0 +1,153 @@ +"""replace agent runtime sessions with workspaces and bindings + +Revision ID: f6e4c5686857 +Revises: 2f39536b3feb +Create Date: 2026-07-23 02:03:05.641638 + +""" +from alembic import op +import models as models +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + +# revision identifiers, used by Alembic. +revision = 'f6e4c5686857' +down_revision = '2f39536b3feb' +branch_labels = None +depends_on = None + + +def upgrade(): + # ### commands auto generated by Alembic - please adjust! ### + op.create_table('agent_workspace_bindings', + sa.Column('tenant_id', models.types.StringUUID(), nullable=False), + sa.Column('app_id', models.types.StringUUID(), nullable=False), + sa.Column('workspace_id', models.types.StringUUID(), nullable=False), + sa.Column('agent_id', models.types.StringUUID(), nullable=False), + sa.Column('base_home_snapshot_id', models.types.StringUUID(), nullable=True), + sa.Column('agent_config_version_id', models.types.StringUUID(), nullable=False), + sa.Column('agent_config_version_kind', sa.String(length=32), nullable=False), + sa.Column('backend_binding_ref', sa.String(length=255), nullable=False), + sa.Column('session_snapshot', models.types.LongText(), nullable=True), + sa.Column('status', sa.String(length=32), server_default='active', nullable=False), + sa.Column('retired_at', sa.DateTime(), nullable=True), + sa.Column('pending_form_id', models.types.StringUUID(), nullable=True), + sa.Column('pending_tool_call_id', sa.String(length=255), nullable=True), + sa.Column('id', models.types.StringUUID(), nullable=False), + sa.Column('created_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False), + sa.Column('updated_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False), + sa.PrimaryKeyConstraint('id', name='agent_workspace_binding_pkey') + ) + with op.batch_alter_table('agent_workspace_bindings', schema=None) as batch_op: + batch_op.create_index('agent_workspace_binding_agent_status_idx', ['tenant_id', 'agent_id', 'status'], unique=False) + batch_op.create_index('agent_workspace_binding_status_retired_idx', ['status', 'retired_at'], unique=False) + batch_op.create_index('agent_workspace_binding_workspace_status_idx', ['tenant_id', 'workspace_id', 'status'], unique=False) + + op.create_table('agent_workspaces', + sa.Column('tenant_id', models.types.StringUUID(), nullable=False), + sa.Column('app_id', models.types.StringUUID(), nullable=False), + sa.Column('owner_type', sa.String(length=32), nullable=False), + sa.Column('owner_id', models.types.StringUUID(), nullable=False), + sa.Column('owner_scope_key', sa.String(length=255), nullable=False), + sa.Column('backend_workspace_ref', sa.String(length=255), nullable=False), + sa.Column('status', sa.String(length=32), server_default='active', nullable=False), + sa.Column('active_guard', sa.SmallInteger(), server_default='1', nullable=True), + sa.Column('retired_at', sa.DateTime(), nullable=True), + sa.Column('id', models.types.StringUUID(), nullable=False), + sa.Column('created_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False), + sa.Column('updated_at', sa.DateTime(), server_default=sa.text('CURRENT_TIMESTAMP'), nullable=False), + sa.PrimaryKeyConstraint('id', name='agent_workspace_pkey') + ) + with op.batch_alter_table('agent_workspaces', schema=None) as batch_op: + batch_op.create_index('agent_workspace_owner_active_unique', ['tenant_id', 'owner_type', 'owner_id', 'owner_scope_key', 'active_guard'], unique=True) + batch_op.create_index('agent_workspace_status_retired_idx', ['status', 'retired_at'], unique=False) + batch_op.create_index('agent_workspace_tenant_app_status_idx', ['tenant_id', 'app_id', 'status'], unique=False) + batch_op.create_index('agent_workspace_tenant_status_idx', ['tenant_id', 'status'], unique=False) + + with op.batch_alter_table('agent_runtime_sessions', schema=None) as batch_op: + batch_op.drop_index(batch_op.f('agent_runtime_session_backend_run_idx')) + batch_op.drop_index(batch_op.f('agent_runtime_session_conversation_lookup_idx')) + batch_op.drop_index(batch_op.f('agent_runtime_session_conversation_scope_unique'), postgresql_where='(conversation_id IS NOT NULL)') + batch_op.drop_index(batch_op.f('agent_runtime_session_workflow_lookup_idx')) + batch_op.drop_index(batch_op.f('agent_runtime_session_workflow_scope_unique'), postgresql_where='(workflow_run_id IS NOT NULL)') + + op.drop_table('agent_runtime_sessions') + with op.batch_alter_table('agent_home_snapshots', schema=None) as batch_op: + batch_op.add_column(sa.Column('status', sa.String(length=32), server_default='active', nullable=False)) + batch_op.add_column(sa.Column('retired_at', sa.DateTime(), nullable=True)) + batch_op.create_index('agent_home_snapshot_status_retired_idx', ['status', 'retired_at'], unique=False) + + with op.batch_alter_table('conversations', schema=None) as batch_op: + batch_op.add_column(sa.Column('agent_workspace_binding_id', models.types.StringUUID(), nullable=True)) + + with op.batch_alter_table('agent_config_drafts', schema=None) as batch_op: + batch_op.add_column(sa.Column('agent_workspace_binding_id', models.types.StringUUID(), nullable=True)) + + with op.batch_alter_table('workflow_node_executions', schema=None) as batch_op: + batch_op.add_column(sa.Column('agent_workspace_binding_id', models.types.StringUUID(), nullable=True)) + + # ### end Alembic commands ### + + +def downgrade(): + # ### commands auto generated by Alembic - please adjust! ### + with op.batch_alter_table('workflow_node_executions', schema=None) as batch_op: + batch_op.drop_column('agent_workspace_binding_id') + + with op.batch_alter_table('agent_config_drafts', schema=None) as batch_op: + batch_op.drop_column('agent_workspace_binding_id') + + with op.batch_alter_table('conversations', schema=None) as batch_op: + batch_op.drop_column('agent_workspace_binding_id') + + with op.batch_alter_table('agent_home_snapshots', schema=None) as batch_op: + batch_op.drop_index('agent_home_snapshot_status_retired_idx') + batch_op.drop_column('retired_at') + batch_op.drop_column('status') + + op.create_table('agent_runtime_sessions', + sa.Column('id', sa.UUID(), server_default=sa.text('uuidv7()'), autoincrement=False, nullable=False), + sa.Column('tenant_id', sa.UUID(), autoincrement=False, nullable=False), + sa.Column('app_id', sa.UUID(), autoincrement=False, nullable=False), + sa.Column('owner_type', sa.VARCHAR(length=32), autoincrement=False, nullable=False), + sa.Column('agent_id', sa.UUID(), autoincrement=False, nullable=False), + sa.Column('backend_run_id', sa.VARCHAR(length=255), autoincrement=False, nullable=True), + sa.Column('session_snapshot', sa.TEXT(), autoincrement=False, nullable=False), + sa.Column('workflow_id', sa.UUID(), autoincrement=False, nullable=True), + sa.Column('workflow_run_id', sa.UUID(), autoincrement=False, nullable=True), + sa.Column('node_id', sa.VARCHAR(length=255), autoincrement=False, nullable=True), + sa.Column('node_execution_id', sa.VARCHAR(length=255), autoincrement=False, nullable=True), + sa.Column('binding_id', sa.UUID(), autoincrement=False, nullable=True), + sa.Column('agent_config_snapshot_id', sa.UUID(), autoincrement=False, nullable=True), + sa.Column('composition_layer_specs', sa.TEXT(), autoincrement=False, nullable=False), + sa.Column('conversation_id', sa.UUID(), autoincrement=False, nullable=True), + sa.Column('status', sa.VARCHAR(length=32), server_default=sa.text("'active'::character varying"), autoincrement=False, nullable=False), + sa.Column('cleaned_at', postgresql.TIMESTAMP(), autoincrement=False, nullable=True), + sa.Column('created_at', postgresql.TIMESTAMP(), server_default=sa.text('CURRENT_TIMESTAMP'), autoincrement=False, nullable=False), + sa.Column('updated_at', postgresql.TIMESTAMP(), server_default=sa.text('CURRENT_TIMESTAMP'), autoincrement=False, nullable=False), + sa.Column('pending_form_id', sa.UUID(), autoincrement=False, nullable=True), + sa.Column('pending_tool_call_id', sa.VARCHAR(length=255), autoincrement=False, nullable=True), + sa.Column('home_snapshot_id', sa.UUID(), autoincrement=False, nullable=True), + sa.PrimaryKeyConstraint('id', name=op.f('agent_runtime_session_pkey')) + ) + with op.batch_alter_table('agent_runtime_sessions', schema=None) as batch_op: + batch_op.create_index(batch_op.f('agent_runtime_session_workflow_scope_unique'), ['tenant_id', 'workflow_run_id', 'node_id', 'binding_id', 'agent_id'], unique=True, postgresql_where='(workflow_run_id IS NOT NULL)') + batch_op.create_index(batch_op.f('agent_runtime_session_workflow_lookup_idx'), ['tenant_id', 'workflow_run_id', 'node_id', 'status'], unique=False) + batch_op.create_index(batch_op.f('agent_runtime_session_conversation_scope_unique'), ['tenant_id', 'conversation_id', 'agent_id', 'agent_config_snapshot_id', 'home_snapshot_id'], unique=True, postgresql_where='(conversation_id IS NOT NULL)') + batch_op.create_index(batch_op.f('agent_runtime_session_conversation_lookup_idx'), ['tenant_id', 'conversation_id', 'status'], unique=False) + batch_op.create_index(batch_op.f('agent_runtime_session_backend_run_idx'), ['backend_run_id'], unique=False) + + with op.batch_alter_table('agent_workspaces', schema=None) as batch_op: + batch_op.drop_index('agent_workspace_tenant_status_idx') + batch_op.drop_index('agent_workspace_tenant_app_status_idx') + batch_op.drop_index('agent_workspace_status_retired_idx') + batch_op.drop_index('agent_workspace_owner_active_unique') + + op.drop_table('agent_workspaces') + with op.batch_alter_table('agent_workspace_bindings', schema=None) as batch_op: + batch_op.drop_index('agent_workspace_binding_workspace_status_idx') + batch_op.drop_index('agent_workspace_binding_status_retired_idx') + batch_op.drop_index('agent_workspace_binding_agent_status_idx') + + op.drop_table('agent_workspace_bindings') + # ### end Alembic commands ### diff --git a/api/migrations/versions/2026_07_23_1200-6f5a9c2d8e1b_add_telemetry_fields_to_dify_setups.py b/api/migrations/versions/2026_07_23_1200-6f5a9c2d8e1b_add_telemetry_fields_to_dify_setups.py new file mode 100644 index 00000000000..ca5ad4a1608 --- /dev/null +++ b/api/migrations/versions/2026_07_23_1200-6f5a9c2d8e1b_add_telemetry_fields_to_dify_setups.py @@ -0,0 +1,30 @@ +"""add telemetry fields to dify_setups + +Revision ID: 6f5a9c2d8e1b +Revises: d2825e7b9c10 +Create Date: 2026-07-23 12:00:00.000000 + +""" + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "6f5a9c2d8e1b" +down_revision = "d2825e7b9c10" +branch_labels = None +depends_on = None + + +def upgrade(): + with op.batch_alter_table("dify_setups", schema=None) as batch_op: + batch_op.add_column(sa.Column("instance_id", sa.String(length=255), nullable=True)) + batch_op.add_column(sa.Column("install_reported_at", sa.DateTime(), nullable=True)) + batch_op.add_column(sa.Column("last_heartbeat_at", sa.DateTime(), nullable=True)) + + +def downgrade(): + with op.batch_alter_table("dify_setups", schema=None) as batch_op: + batch_op.drop_column("last_heartbeat_at") + batch_op.drop_column("install_reported_at") + batch_op.drop_column("instance_id") diff --git a/api/migrations/versions/2026_07_28_2331-e4708db55c1d_make_home_snapshot_references_nullable.py b/api/migrations/versions/2026_07_28_2331-e4708db55c1d_make_home_snapshot_references_nullable.py new file mode 100644 index 00000000000..5eb4b722c2b --- /dev/null +++ b/api/migrations/versions/2026_07_28_2331-e4708db55c1d_make_home_snapshot_references_nullable.py @@ -0,0 +1,49 @@ +"""make home snapshot references nullable + +Revision ID: e4708db55c1d +Revises: f6e4c5686857 +Create Date: 2026-07-28 23:31:26.284147 + +The corrected preceding revisions make fresh upgrades nullable. This linear +convergence revision intentionally repeats those alters so databases that +already ran the old NOT NULL revisions are repaired too. + +""" +from alembic import op +import models as models + + +# revision identifiers, used by Alembic. +revision = 'e4708db55c1d' +down_revision = 'f6e4c5686857' +branch_labels = None +depends_on = None + + +def upgrade(): + with op.batch_alter_table("agent_config_drafts", schema=None) as batch_op: + batch_op.alter_column( + "home_snapshot_id", + existing_type=models.types.StringUUID(), + nullable=True, + ) + + with op.batch_alter_table("agent_config_snapshots", schema=None) as batch_op: + batch_op.alter_column( + "home_snapshot_id", + existing_type=models.types.StringUUID(), + nullable=True, + ) + + with op.batch_alter_table("agent_workspace_bindings", schema=None) as batch_op: + batch_op.alter_column( + "base_home_snapshot_id", + existing_type=models.types.StringUUID(), + nullable=True, + ) + + +def downgrade(): + # The corrected down revision already defines nullable columns, and rows + # created under this revision may contain NULL values. + pass diff --git a/api/migrations/versions/2026_08_05_1030-a1c7f4e9b3d2_add_workflow_version_number.py b/api/migrations/versions/2026_08_05_1030-a1c7f4e9b3d2_add_workflow_version_number.py new file mode 100644 index 00000000000..3e092e01d8a --- /dev/null +++ b/api/migrations/versions/2026_08_05_1030-a1c7f4e9b3d2_add_workflow_version_number.py @@ -0,0 +1,60 @@ +"""add workflow version number + +Revision ID: a1c7f4e9b3d2 +Revises: e4708db55c1d +Create Date: 2026-08-05 10:30:00.000000 + +Introduces user-facing workflow version numbers (`#N`), unique and monotonically +increasing per app. `workflow_version_counters` holds one row per app with the +highest number handed out so far, so numbers are never reused when a published +version is deleted. + +DDL only. Versions published before this revision keep `version_number` NULL and +continue to render as "Untitled Version"; numbering starts at #1 on the first +publish after the upgrade. + +""" + +import sqlalchemy as sa +from alembic import op + +import models as models + +# revision identifiers, used by Alembic. +revision = "a1c7f4e9b3d2" +down_revision = "e4708db55c1d" +branch_labels = None +depends_on = None + + +def upgrade(): + op.create_table( + "workflow_version_counters", + sa.Column("app_id", models.types.StringUUID(), nullable=False), + sa.Column("last_version_number", sa.Integer(), nullable=False), + sa.PrimaryKeyConstraint("app_id", name="workflow_version_counter_pkey"), + ) + + with op.batch_alter_table("workflows", schema=None) as batch_op: + batch_op.add_column(sa.Column("version_number", sa.Integer(), nullable=True)) + + # Excluding NULLs keeps the index off every pre-existing version row. The + # partial-WHERE clause is PG-only (SQLAlchemy drops the kwarg on MySQL → + # plain unique index); both dialects treat NULLs as distinct, so unnumbered + # rows stay unconstrained either way. + op.create_index( + "workflow_app_version_number_idx", + "workflows", + ["app_id", "version_number"], + unique=True, + postgresql_where=sa.text("version_number IS NOT NULL"), + ) + + +def downgrade(): + op.drop_index("workflow_app_version_number_idx", table_name="workflows") + + with op.batch_alter_table("workflows", schema=None) as batch_op: + batch_op.drop_column("version_number") + + op.drop_table("workflow_version_counters") diff --git a/api/models/__init__.py b/api/models/__init__.py index b4c1362b414..b0c3058f0b7 100644 --- a/api/models/__init__.py +++ b/api/models/__init__.py @@ -15,21 +15,22 @@ from .agent import ( AgentConfigRevision, AgentConfigRevisionOperation, AgentConfigSnapshot, + AgentConfigVersionKind, AgentDebugConversation, AgentDriveFile, AgentDriveFileKind, + AgentHomeSnapshot, AgentIconType, AgentKind, - AgentRuntimeSession, - AgentRuntimeSessionOwnerType, - AgentRuntimeSessionStatus, AgentScope, AgentSource, AgentStatus, + AgentWorkingResourceStatus, + AgentWorkspace, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, WorkflowAgentBindingType, WorkflowAgentNodeBinding, - WorkflowAgentRuntimeSession, - WorkflowAgentRuntimeSessionStatus, ) from .api_based_extension import APIBasedExtension, APIBasedExtensionPoint from .comment import ( @@ -147,6 +148,7 @@ from .workflow import ( WorkflowRun, WorkflowRunArchiveBundle, WorkflowType, + WorkflowVersionCounter, resolve_workflow_kind, ) @@ -164,17 +166,20 @@ __all__ = [ "AgentConfigRevision", "AgentConfigRevisionOperation", "AgentConfigSnapshot", + "AgentConfigVersionKind", "AgentDebugConversation", "AgentDriveFile", "AgentDriveFileKind", + "AgentHomeSnapshot", "AgentIconType", "AgentKind", - "AgentRuntimeSession", - "AgentRuntimeSessionOwnerType", - "AgentRuntimeSessionStatus", "AgentScope", "AgentSource", "AgentStatus", + "AgentWorkingResourceStatus", + "AgentWorkspace", + "AgentWorkspaceBinding", + "AgentWorkspaceOwnerType", "ApiRequest", "ApiToken", "ApiToolProvider", @@ -271,8 +276,6 @@ __all__ = [ "Workflow", "WorkflowAgentBindingType", "WorkflowAgentNodeBinding", - "WorkflowAgentRuntimeSession", - "WorkflowAgentRuntimeSessionStatus", "WorkflowAppLog", "WorkflowAppLogCreatedFrom", "WorkflowArchiveLog", @@ -291,5 +294,6 @@ __all__ = [ "WorkflowToolProvider", "WorkflowTriggerStatus", "WorkflowType", + "WorkflowVersionCounter", "resolve_workflow_kind", ] diff --git a/api/models/account.py b/api/models/account.py index 919ee7da820..5bf4e3e641e 100644 --- a/api/models/account.py +++ b/api/models/account.py @@ -165,11 +165,8 @@ class Account(UserMixin, TypeBase): def set_tenant_id_with_session(self, tenant_id: str, *, session: Session) -> None: """Set the current tenant by id using the caller-owned session.""" - query = ( - select(Tenant, TenantAccountJoin) - .where(Tenant.id == tenant_id) - .where(TenantAccountJoin.tenant_id == Tenant.id) - .where(TenantAccountJoin.account_id == self.id) + query = select(Tenant, TenantAccountJoin).where( + Tenant.id == tenant_id, TenantAccountJoin.tenant_id == Tenant.id, TenantAccountJoin.account_id == self.id ) tenant_account_join = session.execute(query).first() if not tenant_account_join: diff --git a/api/models/agent.py b/api/models/agent.py index cd3d371481d..a5863dd1259 100644 --- a/api/models/agent.py +++ b/api/models/agent.py @@ -116,35 +116,29 @@ class WorkflowAgentBindingType(StrEnum): INLINE_AGENT = "inline_agent" -class AgentRuntimeSessionStatus(StrEnum): - """Lifecycle state of an Agent backend session snapshot. +class AgentWorkingResourceStatus(StrEnum): + """Product lifecycle state for a persistent working-environment resource.""" - Owner-agnostic: applies both to workflow Agent Node runs (owner = - workflow_run) and to Agent App conversations (owner = conversation). - """ - - # Snapshot can be reused by a later Agent run in the same session. ACTIVE = "active" - # Snapshot has been retired and must not be submitted to Agent backend again. - CLEANED = "cleaned" + RETIRED = "retired" -class AgentRuntimeSessionOwnerType(StrEnum): - """Which product surface owns an Agent runtime session row.""" +class AgentWorkspaceOwnerType(StrEnum): + """Product scope that owns a Workspace.""" - # Owned by one workflow Agent Node execution scope. WORKFLOW_RUN = "workflow_run" - # Owned by one Agent App conversation (multi-turn chat). CONVERSATION = "conversation" + BUILD_DRAFT = "build_draft" -# Back-compat alias: the workflow lifecycle code (shipped in PR #36724) imports -# the old name. Kept so unifying the table does not churn that path. -WorkflowAgentRuntimeSessionStatus = AgentRuntimeSessionStatus +class AgentConfigVersionKind(StrEnum): + SNAPSHOT = "snapshot" + DRAFT = "draft" + BUILD_DRAFT = "build_draft" class Agent(DefaultFieldsMixin, Base): - """Workspace-scoped Agent identity used by Agent Roster and workflow-only agents.""" + """Agent Soul and source lineage; ``AgentWorkspaceBinding.id`` identifies each materialized participant.""" __tablename__ = "agents" __table_args__ = ( @@ -221,14 +215,42 @@ class Agent(DefaultFieldsMixin, Base): archived_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) -class AgentDebugConversation(DefaultFieldsMixin, Base): - """Per-account, per-draft console debug conversation for an Agent App. +class AgentHomeSnapshot(Base): + """Append-only mapping from one Agent-owned Home identity to its backend ref. - Agent App preview state must be isolated by editor account. The Agent row is - shared by everyone in the workspace, so this table owns the user-specific - conversation pointers used by console debug chat. ``draft`` is the Preview - conversation and ``debug_build`` is the Build conversation; they must never - share persisted messages or runtime sessions. + Product tables reference ``id``. ``snapshot_ref`` remains an opaque + deployment-specific handle and is only consumed at Dify Agent boundaries. + Snapshot bytes and ``snapshot_ref`` are immutable. Lifecycle metadata can + transition ACTIVE -> RETIRED; successful physical collection deletes row. + """ + + __tablename__ = "agent_home_snapshots" + __table_args__ = ( + sa.PrimaryKeyConstraint("id", name="agent_home_snapshot_pkey"), + Index("agent_home_snapshot_tenant_agent_idx", "tenant_id", "agent_id"), + Index("agent_home_snapshot_status_retired_idx", "status", "retired_at"), + ) + + id: Mapped[str] = mapped_column(StringUUID, default=lambda: str(uuidv7())) + tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False) + agent_id: Mapped[str] = mapped_column(StringUUID, nullable=False) + snapshot_ref: Mapped[str] = mapped_column(String(255), nullable=False) + status: Mapped[AgentWorkingResourceStatus] = mapped_column( + EnumText(AgentWorkingResourceStatus, length=32), + nullable=False, + default=AgentWorkingResourceStatus.ACTIVE, + server_default=AgentWorkingResourceStatus.ACTIVE.value, + ) + retired_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + created_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, server_default=func.current_timestamp()) + + +class AgentDebugConversation(DefaultFieldsMixin, Base): + """Current console Conversation pointer for one account and draft surface. + + This row owns no Binding or runtime. A Preview Conversation holds its + CONVERSATION Binding pointer, while a DEBUG_BUILD AgentConfigDraft holds its + BUILD_DRAFT Binding pointer. """ __tablename__ = "agent_debug_conversations" @@ -259,7 +281,11 @@ class AgentDebugConversation(DefaultFieldsMixin, Base): class AgentConfigDraft(DefaultFieldsMixin, Base): - """Editable Agent Soul draft separated from immutable published snapshots.""" + """Editable Agent Soul draft separated from immutable published snapshots. + + A DEBUG_BUILD draft owns its materialized participant through + ``agent_workspace_binding_id``. Normal drafts leave that pointer unset. + """ __tablename__ = "agent_config_drafts" __table_args__ = ( @@ -281,6 +307,8 @@ class AgentConfigDraft(DefaultFieldsMixin, Base): account_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) draft_owner_key: Mapped[str] = mapped_column(String(255), nullable=False, default="") base_snapshot_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) + home_snapshot_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) + agent_workspace_binding_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) config_snapshot: Mapped[Any] = mapped_column(JSONModelColumn(AgentSoulConfig), nullable=False) created_by: Mapped[str | None] = mapped_column(StringUUID, nullable=True) updated_by: Mapped[str | None] = mapped_column(StringUUID, nullable=True) @@ -316,6 +344,7 @@ class AgentConfigSnapshot(DefaultFieldsMixin, Base): agent_id: Mapped[str] = mapped_column(StringUUID, nullable=False) version: Mapped[int] = mapped_column(sa.Integer, nullable=False) config_snapshot: Mapped[Any] = mapped_column(JSONModelColumn(AgentSoulConfig), nullable=False) + home_snapshot_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) summary: Mapped[str | None] = mapped_column(LongText, nullable=True) version_note: Mapped[str | None] = mapped_column(LongText, nullable=True) created_by: Mapped[str | None] = mapped_column(StringUUID, nullable=True) @@ -432,102 +461,83 @@ class WorkflowAgentNodeBinding(DefaultFieldsMixin, Base): return dict(self.node_job_config) -class AgentRuntimeSession(DefaultFieldsMixin, Base): - """Persisted Agent backend session snapshot, owner-agnostic. +class AgentWorkspace(DefaultFieldsMixin, Base): + """Mutable Workspace owned by one product scope, independent of Agents.""" - One unified table serves both owners (decision Q2): - - workflow Agent Node runs: ``owner_type = workflow_run``; the - ``workflow_id / workflow_run_id / node_id / binding_id / - agent_config_snapshot_id / composition_layer_specs`` columns are set. - - Agent App conversations: ``owner_type = conversation``; the - ``conversation_id`` column is set and the workflow columns stay NULL. - Runtime state is scoped by ``agent_config_snapshot_id``. For published - web/API runs this points to an immutable AgentConfigSnapshot; for console - debugger/build runs it points to the editable AgentConfigDraft row. - - The snapshot is runtime state returned by Agent backend, kept separate from - Agent Soul snapshots and workflow node-job config. - """ - - __tablename__ = "agent_runtime_sessions" + __tablename__ = "agent_workspaces" __table_args__ = ( - sa.PrimaryKeyConstraint("id", name="agent_runtime_session_pkey"), - # Workflow owner uniqueness (partial: only rows with a workflow_run_id). + sa.PrimaryKeyConstraint("id", name="agent_workspace_pkey"), Index( - "agent_runtime_session_workflow_scope_unique", + "agent_workspace_owner_active_unique", "tenant_id", - "workflow_run_id", - "node_id", - "binding_id", - "agent_id", + "owner_type", + "owner_id", + "owner_scope_key", + "active_guard", unique=True, - postgresql_where=sa.text("workflow_run_id IS NOT NULL"), ), - # Conversation owner uniqueness (partial: only rows with a conversation_id). - Index( - "agent_runtime_session_conversation_scope_unique", - "tenant_id", - "conversation_id", - "agent_id", - "agent_config_snapshot_id", - unique=True, - postgresql_where=sa.text("conversation_id IS NOT NULL"), - ), - Index( - "agent_runtime_session_workflow_lookup_idx", - "tenant_id", - "workflow_run_id", - "node_id", - "status", - ), - Index( - "agent_runtime_session_conversation_lookup_idx", - "tenant_id", - "conversation_id", - "status", - ), - Index("agent_runtime_session_backend_run_idx", "backend_run_id"), + Index("agent_workspace_tenant_status_idx", "tenant_id", "status"), + Index("agent_workspace_tenant_app_status_idx", "tenant_id", "app_id", "status"), + Index("agent_workspace_status_retired_idx", "status", "retired_at"), ) tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False) app_id: Mapped[str] = mapped_column(StringUUID, nullable=False) - owner_type: Mapped[AgentRuntimeSessionOwnerType] = mapped_column( - EnumText(AgentRuntimeSessionOwnerType, length=32), nullable=False + owner_type: Mapped[AgentWorkspaceOwnerType] = mapped_column( + EnumText(AgentWorkspaceOwnerType, length=32), nullable=False ) - agent_id: Mapped[str] = mapped_column(StringUUID, nullable=False) - backend_run_id: Mapped[str | None] = mapped_column(String(255), nullable=True) - session_snapshot: Mapped[str] = mapped_column(LongText, nullable=False) - # Workflow-owner columns (NULL for conversation owner). - workflow_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) - workflow_run_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) - node_id: Mapped[str | None] = mapped_column(String(255), nullable=True) - node_execution_id: Mapped[str | None] = mapped_column(String(255), nullable=True) - binding_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) - agent_config_snapshot_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) - # JSON-encoded list of non-sensitive runtime layer specs ({name, type, deps, - # config}). The persisted schema keeps its original name because the sandbox - # refactor intentionally avoids a storage migration. - composition_layer_specs: Mapped[str] = mapped_column(LongText, nullable=False, server_default="[]") - # Conversation-owner column (NULL for workflow owner). - conversation_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) - status: Mapped[AgentRuntimeSessionStatus] = mapped_column( - EnumText(AgentRuntimeSessionStatus, length=32), + owner_id: Mapped[str] = mapped_column(StringUUID, nullable=False) + owner_scope_key: Mapped[str] = mapped_column(String(255), nullable=False) + backend_workspace_ref: Mapped[str] = mapped_column(String(255), nullable=False) + status: Mapped[AgentWorkingResourceStatus] = mapped_column( + EnumText(AgentWorkingResourceStatus, length=32), nullable=False, - default=AgentRuntimeSessionStatus.ACTIVE, + default=AgentWorkingResourceStatus.ACTIVE, + server_default=AgentWorkingResourceStatus.ACTIVE.value, ) - cleaned_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) - # ENG-637: when a run pauses for a dify.ask_human deferred call, these link - # the session to the awaiting HITL form and the deferred tool_call_id, so a - # resumed node can map the submitted form back into deferred_tool_results. - # Both NULL whenever the session is not paused on human input. + active_guard: Mapped[int | None] = mapped_column(sa.SmallInteger, nullable=True, default=1, server_default="1") + retired_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) + + +class AgentWorkspaceBinding(DefaultFieldsMixin, Base): + """One materialized Agent participant and session attached to a Workspace. + + All resource IDs are logical associations rather than database foreign + keys, so RETIRED rows can outlive their Workspace or base Home Snapshot. + ``agent_id`` identifies the source Agent Soul; this row's ``id`` identifies + the participant and its private Materialized Home. + """ + + __tablename__ = "agent_workspace_bindings" + __table_args__ = ( + sa.PrimaryKeyConstraint("id", name="agent_workspace_binding_pkey"), + Index("agent_workspace_binding_workspace_status_idx", "tenant_id", "workspace_id", "status"), + Index("agent_workspace_binding_agent_status_idx", "tenant_id", "agent_id", "status"), + Index("agent_workspace_binding_status_retired_idx", "status", "retired_at"), + ) + + tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False) + app_id: Mapped[str] = mapped_column(StringUUID, nullable=False) + workspace_id: Mapped[str] = mapped_column(StringUUID, nullable=False) + agent_id: Mapped[str] = mapped_column(StringUUID, nullable=False) + base_home_snapshot_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) + agent_config_version_id: Mapped[str] = mapped_column(StringUUID, nullable=False) + agent_config_version_kind: Mapped[AgentConfigVersionKind] = mapped_column( + EnumText(AgentConfigVersionKind, length=32), nullable=False + ) + backend_binding_ref: Mapped[str] = mapped_column(String(255), nullable=False) + session_snapshot: Mapped[str | None] = mapped_column(LongText, nullable=True) + status: Mapped[AgentWorkingResourceStatus] = mapped_column( + EnumText(AgentWorkingResourceStatus, length=32), + nullable=False, + default=AgentWorkingResourceStatus.ACTIVE, + server_default=AgentWorkingResourceStatus.ACTIVE.value, + ) + retired_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) pending_form_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) pending_tool_call_id: Mapped[str | None] = mapped_column(String(255), nullable=True) -# Back-compat alias for the shipped workflow lifecycle code (PR #36724). -WorkflowAgentRuntimeSession = AgentRuntimeSession - - class AgentDriveFileKind(StrEnum): """Kind of existing file record an agent-drive KV entry points at.""" diff --git a/api/models/agent_config_entities.py b/api/models/agent_config_entities.py index a6a0cd544a2..fc845144e68 100644 --- a/api/models/agent_config_entities.py +++ b/api/models/agent_config_entities.py @@ -420,8 +420,19 @@ class AgentKnowledgeRetrievalConfig(BaseModel): class AgentKnowledgeMetadataCondition(BaseModel): + """One manual metadata filter clause. + + ``id`` and ``metadata_id`` are UI-only bookkeeping the composer sends on + every save (a stable row key and a reference to the selected metadata + field). They are persisted here for round-tripping the composer's draft + state but are stripped before building the Agent runtime request, whose + DTO only accepts ``name``/``comparison_operator``/``value``. + """ + model_config = ConfigDict(extra="forbid") + id: str | None = None + metadata_id: str | None = None name: str = Field(min_length=1, max_length=255) comparison_operator: SupportedComparisonOperator value: ConditionValue = None @@ -527,8 +538,14 @@ class AgentModelResponseFormatConfig(AgentFlexibleConfig): type: str | None = Field(default=None, max_length=64) -class AgentSoulModelSettings(BaseModel): - model_config = ConfigDict(extra="ignore") +class AgentSoulModelSettings(AgentFlexibleConfig): + """Model parameters for the Agent Soul model. + + Model plugins can declare arbitrary parameters via ``parameter_rules`` + (e.g. Qwen/Tongyi's ``enable_thinking``) beyond the common OpenAI-style + fields typed below, so extra keys must round-trip through persistence + rather than being dropped. + """ temperature: float | None = None top_p: float | None = None diff --git a/api/models/dataset.py b/api/models/dataset.py index 1432d819486..b82681e636c 100644 --- a/api/models/dataset.py +++ b/api/models/dataset.py @@ -347,7 +347,11 @@ class Dataset(Base): def get_doc_form(self, *, session: Session) -> str | None: if self.chunk_structure: return self.chunk_structure - return session.scalar(select(Document.doc_form).where(Document.dataset_id == self.id).limit(1)) + return session.scalar( + select(Document.doc_form) + .where(Document.dataset_id == self.id, Document.tenant_id == self.tenant_id) + .limit(1) + ) @property def retrieval_model_dict(self): @@ -736,7 +740,11 @@ class Document(Base): select(DatasetMetadata) .join(DatasetMetadataBinding, DatasetMetadataBinding.metadata_id == DatasetMetadata.id) .where( - DatasetMetadataBinding.dataset_id == self.dataset_id, DatasetMetadataBinding.document_id == self.id + DatasetMetadata.tenant_id == self.tenant_id, + DatasetMetadata.dataset_id == self.dataset_id, + DatasetMetadataBinding.tenant_id == self.tenant_id, + DatasetMetadataBinding.dataset_id == self.dataset_id, + DatasetMetadataBinding.document_id == self.id, ) ).all() metadata_list: list[DocMetadataDetailItem] = [] @@ -919,10 +927,10 @@ class DocumentSegment(TypeBase): tenant_id: Mapped[str] = mapped_column(StringUUID, nullable=False) dataset_id: Mapped[str] = mapped_column(StringUUID, nullable=False) document_id: Mapped[str] = mapped_column(StringUUID, nullable=False) - position: Mapped[int] + position: Mapped[int] = mapped_column(sa.Integer, nullable=False) content: Mapped[str] = mapped_column(LongText, nullable=False) - word_count: Mapped[int] - tokens: Mapped[int] + word_count: Mapped[int] = mapped_column(sa.Integer, nullable=False) + tokens: Mapped[int] = mapped_column(sa.Integer, nullable=False) created_by: Mapped[str] = mapped_column(StringUUID, nullable=False) # basic fields @@ -1740,7 +1748,9 @@ class Pipeline(TypeBase): ) def retrieve_dataset(self, session: Session | scoped_session): - return session.scalar(select(Dataset).where(Dataset.pipeline_id == self.id)) + return session.scalar( + select(Dataset).where(Dataset.pipeline_id == self.id, Dataset.tenant_id == self.tenant_id) + ) class DocumentPipelineExecutionLog(TypeBase): diff --git a/api/models/model.py b/api/models/model.py index c9f27a78b91..0a09f50b844 100644 --- a/api/models/model.py +++ b/api/models/model.py @@ -362,6 +362,9 @@ class DifySetup(TypeBase): __table_args__ = (sa.PrimaryKeyConstraint("version", name="dify_setup_pkey"),) version: Mapped[str] = mapped_column(String(255), nullable=False) + instance_id: Mapped[str | None] = mapped_column(String(255), nullable=True, default=None) + install_reported_at: Mapped[datetime | None] = mapped_column(sa.DateTime, nullable=True, default=None) + last_heartbeat_at: Mapped[datetime | None] = mapped_column(sa.DateTime, nullable=True, default=None) setup_at: Mapped[datetime] = mapped_column( sa.DateTime, nullable=False, server_default=func.current_timestamp(), init=False ) @@ -1114,14 +1117,14 @@ class ExporleBanner(TypeBase): status: Mapped[BannerStatus] = mapped_column( EnumText(BannerStatus, length=255), nullable=False, - server_default=sa.text("'enabled'::character varying"), + server_default=sa.text("'enabled'"), default=BannerStatus.ENABLED, ) created_at: Mapped[datetime] = mapped_column( sa.DateTime, nullable=False, server_default=func.current_timestamp(), init=False ) language: Mapped[str] = mapped_column( - String(255), nullable=False, server_default=sa.text("'en-US'::character varying"), default="en-US" + String(255), nullable=False, server_default=sa.text("'en-US'"), default="en-US" ) @@ -1157,6 +1160,13 @@ class OAuthProviderApp(TypeBase): class Conversation(Base): + """Conversation state, including the exact Agent participant when applicable. + + ``agent_workspace_binding_id`` is a logical pointer rather than a foreign + key because retired Binding ledger rows may be collected before the + conversation history is deleted. + """ + __tablename__ = "conversations" __table_args__ = ( sa.PrimaryKeyConstraint("id", name="conversation_pkey"), @@ -1178,6 +1188,7 @@ class Conversation(Base): id: Mapped[str] = mapped_column(StringUUID, default=lambda: str(uuid4())) app_id = mapped_column(StringUUID, nullable=False) app_model_config_id = mapped_column(StringUUID, nullable=True) + agent_workspace_binding_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) model_provider = mapped_column(String(255), nullable=True) override_model_configs = mapped_column(LongText) model_id = mapped_column(String(255), nullable=True) @@ -1719,11 +1730,10 @@ class Message(Base): select(MessageFeedback).where(MessageFeedback.message_id == self.id, MessageFeedback.from_source == "user") ) - @property - def admin_feedback(self) -> MessageFeedback | None: - return self.admin_feedback_with_session(session=db.session()) + def admin_feedback(self, session: Session) -> MessageFeedback | None: + return self.admin_feedback_with_session(session=session) - def admin_feedback_with_session(self, *, session: Session) -> MessageFeedback | None: + def admin_feedback_with_session(self, session: Session) -> MessageFeedback | None: return session.scalar( select(MessageFeedback).where(MessageFeedback.message_id == self.id, MessageFeedback.from_source == "admin") ) diff --git a/api/models/workflow.py b/api/models/workflow.py index 58135e79334..fa44dee7ec0 100644 --- a/api/models/workflow.py +++ b/api/models/workflow.py @@ -25,6 +25,7 @@ from typing_extensions import deprecated from core.trigger.constants import TRIGGER_PLUGIN_NODE_TYPE from core.workflow.human_input_adapter import adapt_node_config_for_graph +from core.workflow.llm_environment_variable import LLMEnvironmentVariable, dump_environment_variable from core.workflow.nodes.human_input.pause_reason import ( HumanInputRequired, ) @@ -210,6 +211,13 @@ class Workflow(Base): # bug __table_args__ = ( sa.PrimaryKeyConstraint("id", name="workflow_pkey"), sa.Index("workflow_version_idx", "tenant_id", "app_id", "version"), + sa.Index( + "workflow_app_version_number_idx", + "app_id", + "version_number", + unique=True, + postgresql_where=sa.text("version_number IS NOT NULL"), + ), ) id: Mapped[str] = mapped_column(StringUUID, default=lambda: str(uuid4())) @@ -223,6 +231,9 @@ class Workflow(Base): # bug server_default=sa.text("'standard'"), ) version: Mapped[str] = mapped_column(String(255), nullable=False) + # User-facing version number, unique and monotonically increasing within an app, displayed as `#N`. + # NULL for draft workflows and for versions published before numbering was introduced. + version_number: Mapped[int | None] = mapped_column(sa.Integer, nullable=True, default=None) marked_name: Mapped[str] = mapped_column(String(255), default="", server_default="") marked_comment: Mapped[str] = mapped_column(String(255), default="", server_default="") graph: Mapped[str] = mapped_column(LongText) @@ -264,6 +275,7 @@ class Workflow(Base): # bug marked_name: str = "", marked_comment: str = "", kind: str | None = WorkflowKind.STANDARD.value, + version_number: int | None = None, ) -> "Workflow": workflow = Workflow() workflow.id = str(uuid4()) @@ -272,6 +284,7 @@ class Workflow(Base): # bug workflow.type = WorkflowType(type) workflow.kind = resolve_workflow_kind(kind) workflow.version = version + workflow.version_number = version_number workflow.graph = graph workflow.features = features workflow.created_by = created_by @@ -581,7 +594,7 @@ class Workflow(Base): # bug @property def environment_variables( self, - ) -> Sequence[StringVariable | IntegerVariable | FloatVariable | SecretVariable]: + ) -> Sequence[StringVariable | IntegerVariable | FloatVariable | SecretVariable | LLMEnvironmentVariable]: # Use workflow.tenant_id to avoid relying on request user in background threads tenant_id = self.tenant_id @@ -596,21 +609,21 @@ class Workflow(Base): # bug # decrypt secret variables value def decrypt_func( var: VariableBase, - ) -> StringVariable | IntegerVariable | FloatVariable | SecretVariable: + ) -> StringVariable | IntegerVariable | FloatVariable | SecretVariable | LLMEnvironmentVariable: match var: case SecretVariable(): return var.model_copy( update={"value": encrypter.decrypt_token(tenant_id=tenant_id, token=var.value)} ) - case StringVariable() | IntegerVariable() | FloatVariable(): + case StringVariable() | IntegerVariable() | FloatVariable() | LLMEnvironmentVariable(): return var case _: # Other variable types are not supported for environment variables raise AssertionError(f"Unexpected variable type for environment variable: {type(var)}") - decrypted_results: list[SecretVariable | StringVariable | IntegerVariable | FloatVariable] = [ - decrypt_func(var) for var in results - ] + decrypted_results: list[ + SecretVariable | StringVariable | IntegerVariable | FloatVariable | LLMEnvironmentVariable + ] = [decrypt_func(var) for var in results] return decrypted_results @environment_variables.setter @@ -646,7 +659,7 @@ class Workflow(Base): # bug encrypted_vars = list(map(encrypt_func, value)) environment_variables_json = json.dumps( - {var.name: var.model_dump() for var in encrypted_vars}, + {var.name: dump_environment_variable(var) for var in encrypted_vars}, ensure_ascii=False, ) self._environment_variables = environment_variables_json @@ -686,7 +699,7 @@ class Workflow(Base): # bug result: WorkflowContentDict = { "graph": self.graph_dict, "features": self.features_dict, - "environment_variables": [var.model_dump(mode="json") for var in environment_variables], + "environment_variables": [dump_environment_variable(var, mode="json") for var in environment_variables], "conversation_variables": [var.model_dump(mode="json") for var in self.conversation_variables], "rag_pipeline_variables": self.rag_pipeline_variables, } @@ -734,6 +747,24 @@ class Workflow(Base): # bug return str(d) +class WorkflowVersionCounter(Base): + """Monotonic per-app allocator for `Workflow.version_number`. + + One row per app, holding the highest number handed out so far. Numbers are never + reused, so deleting a published version does not free its number. + + `app_id` mirrors `Workflow.app_id`, which is polymorphic: it holds an app id, a + pipeline id or a snippet id depending on the workflow kind. UUID uniqueness across + those tables is why no owner-type column is needed here. + """ + + __tablename__ = "workflow_version_counters" + __table_args__ = (sa.PrimaryKeyConstraint("app_id", name="workflow_version_counter_pkey"),) + + app_id: Mapped[str] = mapped_column(StringUUID) + last_version_number: Mapped[int] = mapped_column(sa.Integer, nullable=False, default=0) + + class WorkflowRunDict(TypedDict): id: str tenant_id: str @@ -1031,6 +1062,7 @@ class WorkflowNodeExecutionModel(Base): # This model is expected to have `offlo node_id: Mapped[str] = mapped_column(String(255)) node_type: Mapped[str] = mapped_column(String(255)) title: Mapped[str] = mapped_column(String(255)) + agent_workspace_binding_id: Mapped[str | None] = mapped_column(StringUUID, nullable=True) inputs: Mapped[str | None] = mapped_column(LongText) process_data: Mapped[str | None] = mapped_column(LongText) outputs: Mapped[str | None] = mapped_column(LongText) diff --git a/api/openapi/markdown/console-openapi.md b/api/openapi/markdown/console-openapi.md index a40ee97685d..c2d159ec4c5 100644 --- a/api/openapi/markdown/console-openapi.md +++ b/api/openapi/markdown/console-openapi.md @@ -964,12 +964,6 @@ Stop a running Agent App chat message generation | ---- | ---------- | ----------- | -------- | ------ | | agent_id | path | | Yes | string (uuid) | -#### Request Body - -| Required | Schema | -| -------- | ------ | -| No | **application/json**: [AgentDebugConversationRefreshPayload](#agentdebugconversationrefreshpayload)
| - #### Responses | Code | Description | Schema | @@ -1261,7 +1255,8 @@ Get basic information for an Agent App conversation sandbox | Name | Located in | Description | Required | Schema | | ---- | ---------- | ----------- | -------- | ------ | | agent_id | path | Agent ID | Yes | string (uuid) | -| conversation_id | query | Agent App conversation ID | Yes | string | +| caller_id | query | Agent App caller ID | Yes | string | +| caller_type | query | | Yes | string,
**Available values:** "build_draft", "conversation" | #### Responses @@ -1277,8 +1272,9 @@ List a directory in an Agent App conversation sandbox | Name | Located in | Description | Required | Schema | | ---- | ---------- | ----------- | -------- | ------ | | agent_id | path | Agent ID | Yes | string (uuid) | -| conversation_id | query | Agent App conversation ID | Yes | string | -| path | query | Directory path relative to the sandbox workspace | No | string,
**Default:** . | +| caller_id | query | Agent App caller ID | Yes | string | +| caller_type | query | | Yes | string,
**Available values:** "build_draft", "conversation" | +| path | query | Binding path: relative paths start in Workspace; exact `~` and paths beginning with `~/` start in Home; `~user` is an ordinary relative path from Workspace; absolute paths remain absolute; `..` and paths outside Workspace are governed by backend isolation, not a Workspace-root restriction | No | string,
**Default:** . | #### Responses @@ -1286,25 +1282,8 @@ List a directory in an Agent App conversation sandbox | ---- | ----------- | ------ | | 200 | Listing returned | **application/json**: [SandboxListResponse](#sandboxlistresponse)
| -### [GET] /agent/{agent_id}/sandbox/files/read -Read a text/binary preview file in an Agent App conversation sandbox - -#### Parameters - -| Name | Located in | Description | Required | Schema | -| ---- | ---------- | ----------- | -------- | ------ | -| agent_id | path | Agent ID | Yes | string (uuid) | -| conversation_id | query | Agent App conversation ID | Yes | string | -| path | query | File path relative to the sandbox workspace | Yes | string | - -#### Responses - -| Code | Description | Schema | -| ---- | ----------- | ------ | -| 200 | Preview returned | **application/json**: [SandboxReadResponse](#sandboxreadresponse)
| - -### [POST] /agent/{agent_id}/sandbox/files/upload -Upload one Agent App sandbox file and return a signed download URL +### [POST] /agent/{agent_id}/sandbox/files/download +Create a ToolFile from one Agent App Binding file and return its download URL #### Parameters @@ -1316,13 +1295,31 @@ Upload one Agent App sandbox file and return a signed download URL | Required | Schema | | -------- | ------ | -| Yes | **application/json**: [AgentSandboxUploadPayload](#agentsandboxuploadpayload)
| +| Yes | **application/json**: [AgentSandboxDownloadPayload](#agentsandboxdownloadpayload)
| #### Responses | Code | Description | Schema | | ---- | ----------- | ------ | -| 200 | Uploaded | **application/json**: [SandboxUploadResponse](#sandboxuploadresponse)
| +| 200 | Download URL returned | **application/json**: [SandboxDownloadResponse](#sandboxdownloadresponse)
| + +### [GET] /agent/{agent_id}/sandbox/files/read +Read a text/binary preview file in an Agent App conversation sandbox + +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| agent_id | path | Agent ID | Yes | string (uuid) | +| caller_id | query | Agent App caller ID | Yes | string | +| caller_type | query | | Yes | string,
**Available values:** "build_draft", "conversation" | +| path | query | Binding path: relative paths start in Workspace; exact `~` and paths beginning with `~/` start in Home; `~user` is an ordinary relative path from Workspace; absolute paths remain absolute; `..` and paths outside Workspace are governed by backend isolation, not a Workspace-root restriction | Yes | string | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Preview returned | **application/json**: [SandboxReadResponse](#sandboxreadresponse)
| ### [POST] /agent/{agent_id}/skills/upload Upload + standardize a Skill into an Agent App drive @@ -1672,6 +1669,23 @@ Create a new application | 200 | Import confirmed | **application/json**: [Import](#import)
| | 400 | Import failed | **application/json**: [Import](#import)
| +### [GET] /apps/recent +**Return the lightweight app cards needed by the Explore home page** + +Get recently modified apps for the home Continue Work section + +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| limit | query | Number of recently modified apps to return (1-8) | No | integer,
**Default:** 8 | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Success | **application/json**: [RecentAppListResponse](#recentapplistresponse)
| + ### [GET] /apps/starred Get applications starred by the current account @@ -3190,6 +3204,23 @@ Update MCP server configuration for an application | 403 | Insufficient permissions | | | 404 | Server not found | | +### [POST] /apps/{app_id}/server/refresh +Refresh MCP server configuration and regenerate server code + +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| app_id | path | App ID | Yes | string (uuid) | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | MCP server refreshed successfully | **application/json**: [AppMCPServerResponse](#appmcpserverresponse)
| +| 403 | Insufficient permissions | | +| 404 | Server not found | | + ### [POST] /apps/{app_id}/site Update application site configuration @@ -3790,8 +3821,8 @@ List a directory in a workflow Agent node sandbox | app_id | path | Application ID | Yes | string (uuid) | | node_id | path | Workflow Agent node ID | Yes | string | | workflow_run_id | path | Workflow run ID | Yes | string (uuid) | -| node_execution_id | query | Optional workflow node execution ID. When omitted, the latest active session for the node is used. | No | string | -| path | query | Directory path relative to the sandbox workspace | No | string,
**Default:** . | +| node_execution_id | query | Workflow node execution ID | Yes | string | +| path | query | Binding path: relative paths start in Workspace; exact `~` and paths beginning with `~/` start in Home; `~user` is an ordinary relative path from Workspace; absolute paths remain absolute; `..` and paths outside Workspace are governed by backend isolation, not a Workspace-root restriction | No | string,
**Default:** . | #### Responses @@ -3799,27 +3830,8 @@ List a directory in a workflow Agent node sandbox | ---- | ----------- | ------ | | 200 | Listing returned | **application/json**: [SandboxListResponse](#sandboxlistresponse)
| -### [GET] /apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/read -Read a text/binary preview file in a workflow Agent node sandbox - -#### Parameters - -| Name | Located in | Description | Required | Schema | -| ---- | ---------- | ----------- | -------- | ------ | -| app_id | path | Application ID | Yes | string (uuid) | -| node_id | path | Workflow Agent node ID | Yes | string | -| workflow_run_id | path | Workflow run ID | Yes | string (uuid) | -| node_execution_id | query | Optional workflow node execution ID. When omitted, the latest active session for the node is used. | No | string | -| path | query | File path relative to the sandbox workspace | Yes | string | - -#### Responses - -| Code | Description | Schema | -| ---- | ----------- | ------ | -| 200 | Preview returned | **application/json**: [SandboxReadResponse](#sandboxreadresponse)
| - -### [POST] /apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/upload -Upload one workflow Agent sandbox file and return a signed download URL +### [POST] /apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/download +Create a ToolFile from one workflow Agent Binding file and return its download URL #### Parameters @@ -3833,13 +3845,32 @@ Upload one workflow Agent sandbox file and return a signed download URL | Required | Schema | | -------- | ------ | -| Yes | **application/json**: [WorkflowAgentSandboxUploadPayload](#workflowagentsandboxuploadpayload)
| +| Yes | **application/json**: [WorkflowAgentSandboxDownloadPayload](#workflowagentsandboxdownloadpayload)
| #### Responses | Code | Description | Schema | | ---- | ----------- | ------ | -| 200 | Uploaded | **application/json**: [SandboxUploadResponse](#sandboxuploadresponse)
| +| 200 | Download URL returned | **application/json**: [SandboxDownloadResponse](#sandboxdownloadresponse)
| + +### [GET] /apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/read +Read a text/binary preview file in a workflow Agent node sandbox + +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| app_id | path | Application ID | Yes | string (uuid) | +| node_id | path | Workflow Agent node ID | Yes | string | +| workflow_run_id | path | Workflow run ID | Yes | string (uuid) | +| node_execution_id | query | Workflow node execution ID | Yes | string | +| path | query | Binding path: relative paths start in Workspace; exact `~` and paths beginning with `~/` start in Home; `~user` is an ordinary relative path from Workspace; absolute paths remain absolute; `..` and paths outside Workspace are governed by backend isolation, not a Workspace-root restriction | Yes | string | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Preview returned | **application/json**: [SandboxReadResponse](#sandboxreadresponse)
| ### [GET] /apps/{app_id}/workflow/comments **Get all comments for a workflow** @@ -5126,23 +5157,6 @@ Restore a published workflow version into the draft workflow | ---- | ----------- | | 204 | API key deleted successfully | -### [GET] /apps/{server_id}/server/refresh -Refresh MCP server configuration and regenerate server code - -#### Parameters - -| Name | Located in | Description | Required | Schema | -| ---- | ---------- | ----------- | -------- | ------ | -| server_id | path | Server ID | Yes | string (uuid) | - -#### Responses - -| Code | Description | Schema | -| ---- | ----------- | ------ | -| 200 | MCP server refreshed successfully | **application/json**: [AppMCPServerResponse](#appmcpserverresponse)
| -| 403 | Insufficient permissions | | -| 404 | Server not found | | - ### [GET] /auth/plugin/datasource/default-list #### Responses @@ -5338,7 +5352,7 @@ Sync partner tenants bindings | Code | Description | Schema | | ---- | ----------- | ------ | -| 200 | Success | **application/json**: [BillingResponse](#billingresponse)
| +| 200 | Success | **application/json**: [BillingSubscriptionResponse](#billingsubscriptionresponse)
| ### [GET] /code-based-extension Get code-based extension data by module name @@ -5961,6 +5975,7 @@ then asynchronously generates summary indexes for the provided documents. | Code | Description | | ---- | ----------- | | 204 | Documents metadata updated successfully | +| 404 | Dataset, document, or metadata not found | ### [PATCH] /datasets/{dataset_id}/documents/status/{action}/batch #### Parameters @@ -6770,7 +6785,7 @@ Check if dataset is in use | Required | Schema | | -------- | ------ | -| Yes | **application/json**: [EmailPayload](#emailpayload)
| +| Yes | **application/json**: [EmailCodeSendPayload](#emailcodesendpayload)
| #### Responses @@ -6885,7 +6900,7 @@ Check if dataset is in use | Code | Description | Schema | | ---- | ----------- | ------ | -| 200 | Success | **application/json**: [LimitationModel](#limitationmodel)
| +| 200 | Success | **application/json**: [VectorSpaceLimitationModel](#vectorspacelimitationmodel)
| ### [GET] /files/support-type #### Responses @@ -7023,19 +7038,15 @@ Request body: | ---- | ----------- | ------ | | 200 | Success | **application/json**: [ConsoleHumanInputFormSubmitResponse](#consolehumaninputformsubmitresponse)
| -### [POST] /info -#### Responses - -| Code | Description | Schema | -| ---- | ----------- | ------ | -| 200 | Success | **application/json**: [TenantInfoResponse](#tenantinforesponse)
| - ### [GET] /installed-apps #### Parameters | Name | Located in | Description | Required | Schema | | ---- | ---------- | ----------- | -------- | ------ | | app_id | query | App ID to filter by | No | string | +| cursor | query | Opaque cursor returned by the previous page | No | string | +| limit | query | Number of installed apps to return | No | integer,
**Default:** 20 | +| name | query | App name to search for | No | string | #### Responses @@ -7069,6 +7080,19 @@ Request body: | ---- | ----------- | | 204 | App uninstalled successfully | +### [GET] /installed-apps/{installed_app_id} +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| installed_app_id | path | | Yes | string (uuid) | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Success | **application/json**: [InstalledAppResponse](#installedappresponse)
| + ### [PATCH] /installed-apps/{installed_app_id} #### Parameters @@ -7867,6 +7891,7 @@ Update account-level Step-by-step Tour state | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Success | **application/json**: [SimpleDataResponse](#simpledataresponse)
| +| 404 | Customized pipeline template not found | | ### [POST] /rag/pipeline/dataset #### Request Body @@ -7915,6 +7940,7 @@ Update account-level Step-by-step Tour state | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Pipeline template | **application/json**: [PipelineTemplateDetailResponse](#pipelinetemplatedetailresponse)
| +| 404 | Pipeline template not found | | ### [GET] /rag/pipelines/datasource-plugins #### Responses @@ -7990,6 +8016,7 @@ Update account-level Step-by-step Tour state | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Success | **application/json**: [RagPipelineOpaqueResponse](#ragpipelineopaqueresponse)
| +| 404 | Dataset or pipeline not found | | ### [POST] /rag/pipelines/{pipeline_id}/customized/publish #### Parameters @@ -8009,6 +8036,7 @@ Update account-level Step-by-step Tour state | Code | Description | | ---- | ----------- | | 204 | Pipeline template published | +| 404 | Pipeline, workflow, or dataset not found | ### [GET] /rag/pipelines/{pipeline_id}/exports #### Parameters @@ -10039,13 +10067,6 @@ Returns information about why and where the workflow is paused. | ---- | ----------- | ------ | | 200 | Success | **application/json**: [TenantListResponse](#tenantlistresponse)
| -### [POST] /workspaces/current -#### Responses - -| Code | Description | Schema | -| ---- | ----------- | ------ | -| 200 | Success | **application/json**: [TenantInfoResponse](#tenantinforesponse)
| - ### [GET] /workspaces/current/agent-provider/{provider_name} Get specific agent provider details @@ -10567,6 +10588,20 @@ Update a plugin endpoint | ---- | ----------- | ------ | | 200 | Model providers retrieved successfully | **application/json**: [ModelProviderListResponse](#modelproviderlistresponse)
| +### [GET] /workspaces/current/model-providers/credits +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Model provider credits retrieved successfully | **application/json**: [ModelProviderCreditsResponse](#modelprovidercreditsresponse)
| + +### [GET] /workspaces/current/model-providers/summary +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Model provider summaries retrieved successfully | **application/json**: [ModelProviderSummaryListResponse](#modelprovidersummarylistresponse)
| + ### [GET] /workspaces/current/model-providers/{provider}/checkout-url #### Parameters @@ -11112,6 +11147,19 @@ Returns permission flags that control workspace features like member invitations | ---- | ----------- | ------ | | 200 | Success | **application/json**: [PluginInstallTaskStartResponse](#plugininstalltaskstartresponse)
| +### [GET] /workspaces/current/plugin/installed-ids +#### Parameters + +| Name | Located in | Description | Required | Schema | +| ---- | ---------- | ----------- | -------- | ------ | +| category | query | Plugin category to include | Yes | string,
**Available values:** "agent-strategy", "datasource", "extension", "model", "tool", "trigger" | + +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Success | **application/json**: [PluginInstalledIdsResponse](#plugininstalledidsresponse)
| + ### [GET] /workspaces/current/plugin/list #### Parameters @@ -11370,8 +11418,11 @@ Returns permission flags that control workspace features like member invitations | Name | Located in | Description | Required | Schema | | ---- | ---------- | ----------- | -------- | ------ | +| language | query | Language used for localized label and description search | No | string,
**Available values:** "en_US", "ja_JP", "pt_BR", "zh_Hans",
**Default:** en_US | | page | query | Page number | No | integer,
**Default:** 1 | | page_size | query | Page size (1-256) | No | integer,
**Default:** 256 | +| query | query | Case-insensitive search query | No | string | +| tags | query | Match any plugin tag | No | [ string ] | | category | path | | Yes | string | #### Responses @@ -11971,6 +12022,14 @@ Returns permission flags that control workspace features like member invitations | ---- | ----------- | ------ | | 200 | Success | **application/json**: [WorkspaceAccessMatrix](#workspaceaccessmatrix)
| +### [GET] /workspaces/current/summary +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Success | **application/json**: [CurrentWorkspaceSummaryResponse](#currentworkspacesummaryresponse)
| +| 409 | Current workspace is archived | | + ### [GET] /workspaces/current/tool-labels #### Responses @@ -12778,6 +12837,13 @@ Returns permission flags that control workspace features like member invitations | ---- | ----------- | ------ | | 200 | Trigger providers retrieved successfully | **application/json**: [TriggerProviderListResponse](#triggerproviderlistresponse)
| +### [GET] /workspaces/custom-config +#### Responses + +| Code | Description | Schema | +| ---- | ----------- | ------ | +| 200 | Success | **application/json**: [WorkspaceCustomConfigResponse](#workspacecustomconfigresponse)
| + ### [POST] /workspaces/custom-config #### Request Body @@ -13218,6 +13284,7 @@ Model class for AI model. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | +| access_ready | boolean | | Yes | | api_key_count | integer | | Yes | | api_rph | integer | | Yes | | api_rpm | integer | | Yes | @@ -13243,6 +13310,7 @@ Model class for AI model. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | +| active_config_is_published | boolean | | Yes | | active_config_snapshot | [AgentConfigSnapshotSummaryResponse](#agentconfigsnapshotsummaryresponse) | | No | | agent | [AgentComposerAgentResponse](#agentcomposeragentresponse) | | Yes | | agent_soul | [AgentSoulConfig](#agentsoulconfig) | | Yes | @@ -13282,7 +13350,7 @@ Model class for AI model. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | | access_mode | string | | No | -| active_config_is_published | boolean | | No | +| access_ready | boolean | | No | | api_base_url | string | | No | | app_id | string | | No | | backing_app_id | string | | No | @@ -13926,12 +13994,6 @@ Stable Agent Soul reference to one normalized skill archive. | date | string | | Yes | | message_count | integer | | Yes | -#### AgentDebugConversationRefreshPayload - -| Name | Type | Description | Required | -| ---- | ---- | ----------- | -------- | -| draft_type | [AgentConfigDraftType](#agentconfigdrafttype) | Agent draft surface whose conversation should be refreshed | No | - #### AgentDebugConversationRefreshResponse | Name | Type | Description | Required | @@ -14245,9 +14307,19 @@ the current roster/workflow APIs scoped to Dify Agent. #### AgentKnowledgeMetadataCondition +One manual metadata filter clause. + +``id`` and ``metadata_id`` are UI-only bookkeeping the composer sends on +every save (a stable row key and a reference to the selected metadata +field). They are persisted here for round-tripping the composer's draft +state but are stripped before building the Agent runtime request, whose +DTO only accepts ``name``/``comparison_operator``/``value``. + | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | | comparison_operator | string,
**Available values:** "<", "=", ">", "after", "before", "contains", "empty", "end with", "in", "is", "is not", "not contains", "not empty", "not in", "start with", "≠", "≤", "≥" | *Enum:* `"<"`, `"="`, `">"`, `"after"`, `"before"`, `"contains"`, `"empty"`, `"end with"`, `"in"`, `"is"`, `"is not"`, `"not contains"`, `"not empty"`, `"not in"`, `"start with"`, `"≠"`, `"≤"`, `"≥"` | Yes | +| id | string | | No | +| metadata_id | string | | No | | name | string | | Yes | | value | string
[ string ]
number | | No | @@ -14375,6 +14447,14 @@ section may be empty, which is how callers express "no knowledge layer". | updated_at | integer | | No | | user_rate | number | | No | +#### AgentLogFeedbackResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| content | string | | No | +| from_source | string,
**Available values:** "admin", "user" | *Enum:* `"admin"`, `"user"` | Yes | +| rating | string,
**Available values:** "dislike", "like" | *Enum:* `"dislike"`, `"like"` | Yes | + #### AgentLogListResponse | Name | Type | Description | Required | @@ -14395,6 +14475,8 @@ section may be empty, which is how callers express "no knowledge layer". | created_at | integer | | No | | currency | string | | Yes | | error | string | | No | +| feedback_enabled | boolean | | No | +| feedbacks | [ [AgentLogFeedbackResponse](#agentlogfeedbackresponse) ] | | No | | from_account_id | string | | No | | from_end_user_id | string | | No | | id | string | | Yes | @@ -14636,6 +14718,14 @@ section may be empty, which is how callers express "no knowledge layer". | workflow_id | string | | No | | workflow_node_id | string | | No | +#### AgentSandboxDownloadPayload + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| caller_id | string | Agent App caller ID | Yes | +| caller_type | string,
**Available values:** "build_draft", "conversation" | *Enum:* `"build_draft"`, `"conversation"` | Yes | +| path | string | Binding path: relative paths start in Workspace; exact `~` and paths beginning with `~/` start in Home; `~user` is an ordinary relative path from Workspace; absolute paths remain absolute; `..` and paths outside Workspace are governed by backend isolation, not a Workspace-root restriction | Yes | + #### AgentSandboxProviderConfig | Name | Type | Description | Required | @@ -14645,13 +14735,6 @@ section may be empty, which is how callers express "no knowledge layer". | image | string | | No | | working_dir | string | | No | -#### AgentSandboxUploadPayload - -| Name | Type | Description | Required | -| ---- | ---- | ----------- | -------- | -| conversation_id | string | Agent App conversation ID | Yes | -| path | string | File path relative to the sandbox workspace | Yes | - #### AgentScope Visibility and lifecycle scope of an Agent record. @@ -14855,6 +14938,13 @@ Reference to model credentials resolved only at runtime. #### AgentSoulModelSettings +Model parameters for the Agent Soul model. + +Model plugins can declare arbitrary parameters via ``parameter_rules`` +(e.g. Qwen/Tongyi's ``enable_thinking``) beyond the common OpenAI-style +fields typed below, so extra keys must round-trip through persistence +rather than being dropped. + | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | | frequency_penalty | number | | No | @@ -15552,6 +15642,12 @@ AppMCPServer Status Enum | ---- | ---- | ----------- | -------- | | AppMCPServerStatus | string | AppMCPServer Status Enum | | +#### AppMode + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| AppMode | string | | | + #### AppModelConfigResponse | Name | Type | Description | Required | @@ -15769,6 +15865,15 @@ AppMCPServer Status Enum | ---- | ---- | ----------- | -------- | | data | [ [AverageSessionInteractionStatisticItem](#averagesessioninteractionstatisticitem) ] | | Yes | +#### BannerContentResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| category | string | | Yes | +| description | string | | Yes | +| img-src | string | | Yes | +| title | string | | Yes | + #### BannerListResponse | Name | Type | Description | Required | @@ -15779,12 +15884,20 @@ AppMCPServer Status Enum | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| content | | | Yes | -| created_at | string | | No | +| content | [BannerContentResponse](#bannercontentresponse) | | Yes | +| created_at | string | | Yes | | id | string | | Yes | -| link | string | | No | +| link | string | | Yes | | sort | integer | | Yes | -| status | string | | Yes | +| status | [BannerStatus](#bannerstatus) | | Yes | + +#### BannerStatus + +ExporleBanner status + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| BannerStatus | string | ExporleBanner status | | #### BatchImportPayload @@ -15834,7 +15947,7 @@ Retrieval settings for Amazon Bedrock knowledge base queries. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| enabled | boolean | | Yes | +| enabled | boolean | Deprecated. Use system features deployment_edition to determine the product edition. | Yes | | subscription | [SubscriptionModel](#subscriptionmodel) | | Yes | #### BillingResponse @@ -15843,6 +15956,12 @@ Retrieval settings for Amazon Bedrock knowledge base queries. | ---- | ---- | ----------- | -------- | | BillingResponse | object | | | +#### BillingSubscriptionResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| url | string | | Yes | + #### BinaryFileResponse | Name | Type | Description | Required | @@ -16081,6 +16200,18 @@ Button styles for user actions. | install_commands | [ string ] | | No | | name | string | | Yes | +#### CloudPlan + +Enum representing user plan types in the cloud platform. + +SANDBOX: Free/default plan with limited features +PROFESSIONAL: Professional paid plan +TEAM: Team collaboration paid plan + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| CloudPlan | string | Enum representing user plan types in the cloud platform. SANDBOX: Free/default plan with limited features PROFESSIONAL: Professional paid plan TEAM: Team collaboration paid plan | | + #### CodeBasedExtensionQuery | Name | Type | Description | Required | @@ -16529,6 +16660,16 @@ Model class for credential form schema. | ---- | ---- | ----------- | -------- | | CredentialType | string | | | +#### CurrentWorkspaceSummaryResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| credits | integer | Remaining credits in the effective pool; -1 means unlimited. | Yes | +| id | string | | Yes | +| name | string | | Yes | +| plan | [CloudPlan](#cloudplan) | | Yes | +| role | [TenantAccountRole](#tenantaccountrole) | | Yes | + #### CustomConfigurationResponse Model class for provider custom configuration response. @@ -17389,7 +17530,7 @@ Request payload for bulk downloading documents as a zip archive. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| document_id | string | Document ID whose metadata should be updated. | Yes | +| document_id | string (uuid) | Document ID whose metadata should be updated. | Yes | | metadata_list | [ [MetadataDetail](#metadatadetail) ] | Metadata fields to update. | Yes | | partial_update | boolean | Whether to partially update metadata, keeping existing values for unspecified fields. | No | @@ -17645,6 +17786,14 @@ Portable DSL reference that could not be restored in the target workspace. | timezone | string | | No | | token | string | | Yes | +#### EmailCodeSendPayload + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| email | string | | Yes | +| language | string | | No | +| turnstile_token | string | Cloudflare Turnstile token. Required at runtime for Dify Cloud. | No | + #### EmailPayload | Name | Type | Description | Required | @@ -17885,7 +18034,9 @@ declaration of an endpoint group | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | +| deleted_environment_variable_ids | [ string ] | Environment variable IDs to delete when patch is true | No | | environment_variables | [ [EnvironmentVariableItemPayload](#environmentvariableitempayload) ] | Environment variables for the draft workflow | Yes | +| patch | boolean | Treat environment_variables as per-ID upserts instead of replacing the full collection | No | #### ErrorDocsResponse @@ -18021,6 +18172,13 @@ Built-in tool icons are URL strings; API-based tool icons are provider-defined p #### ExternalDatasetCreatePayload +Validated fields required to create an external dataset binding. + +The console controller owns HTTP concerns, but the service also needs this +contract when creating the tenant-scoped dataset and external knowledge +binding. Keep it outside controllers so service imports do not depend on +Flask blueprint initialization. + | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | | description | string | | No | @@ -18654,21 +18812,23 @@ Input field definition for snippet parameters. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| description | string | | No | -| icon | string | | No | -| icon_background | string | | No | -| icon_type | string | | No | +| description | string | | Yes | +| icon | string | | Yes | +| icon_background | string | | Yes | +| icon_type | [IconType](#icontype) | | Yes | | icon_url | string | | Yes | | id | string | | Yes | -| mode | string | | No | -| name | string | | No | -| use_icon_as_answer_icon | boolean | | No | +| mode | [AppMode](#appmode) | | Yes | +| name | string | | Yes | +| use_icon_as_answer_icon | boolean | | Yes | #### InstalledAppListResponse | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | +| has_more | boolean | | Yes | | installed_apps | [ [InstalledAppResponse](#installedappresponse) ] | | Yes | +| next_cursor | string | | Yes | #### InstalledAppResponse @@ -18679,7 +18839,7 @@ Input field definition for snippet parameters. | editable | boolean | | Yes | | id | string | | Yes | | is_pinned | boolean | | Yes | -| last_used_at | integer | | No | +| last_used_at | integer | | Yes | | uninstallable | boolean | | Yes | #### InstalledAppUpdatePayload @@ -18693,6 +18853,9 @@ Input field definition for snippet parameters. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | | app_id | string | App ID to filter by | No | +| cursor | string | Opaque cursor returned by the previous page | No | +| limit | integer,
**Default:** 20 | Number of installed apps to return | No | +| name | string | App name to search for | No | #### InstructionGeneratePayload @@ -18822,6 +18985,7 @@ Enum class for large language model mode. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | | expired_at | string | | Yes | +| license_expiry_notice_enabled | boolean | | Yes | | seats | [LicenseLimitationModel](#licenselimitationmodel) | | Yes | | status | [LicenseStatus](#licensestatus) | | Yes | | workspaces | [LicenseLimitationModel](#licenselimitationmodel) | | Yes | @@ -19171,7 +19335,7 @@ Enum class for large language model mode. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| id | string | Metadata field ID. | Yes | +| id | string (uuid) | Metadata field ID. | Yes | | name | string | Metadata field name. | Yes | | value | string
integer
number | Metadata value. Can be a string, number, or `null`. | No | @@ -19307,6 +19471,30 @@ Enum class for model property key. | ---- | ---- | ----------- | -------- | | ModelPropertyKey | string | Enum class for model property key. | | +#### ModelProviderCreditsResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| exhausted_at | integer | | Yes | +| is_exhausted | boolean | | Yes | +| is_unlimited | boolean | | Yes | +| next_credit_reset_date | integer | | Yes | +| pool_type | string | | Yes | +| quota_limit | integer | Credit limit for the effective pool; -1 means unlimited. | Yes | +| quota_used | integer | | Yes | +| remaining_credits | integer | Remaining credits; -1 means unlimited. | Yes | + +#### ModelProviderCustomConfigurationSummaryResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| available_credentials | [ [CredentialConfiguration](#credentialconfiguration) ] | | Yes | +| current_credential_id | string | | No | +| current_credential_name | string | | No | +| current_credential_usable | boolean | | Yes | +| has_custom_models | boolean | Whether custom model configuration exists, including saved model credentials. | Yes | +| status | [CustomConfigurationStatus](#customconfigurationstatus) | | Yes | + #### ModelProviderListResponse | Name | Type | Description | Required | @@ -19319,6 +19507,49 @@ Enum class for model property key. | ---- | ---- | ----------- | -------- | | payment_link | string | | Yes | +#### ModelProviderPluginSummaryResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| installation_id | string | | Yes | +| plugin_id | string | | Yes | +| plugin_unique_identifier | string | | Yes | +| runtime_type | string | | Yes | +| source | [PluginInstallationSource](#plugininstallationsource) | | Yes | +| version | string | | Yes | + +#### ModelProviderSummaryListResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| data | [ [ModelProviderSummaryResponse](#modelprovidersummaryresponse) ] | | Yes | +| plugins | object | | Yes | + +#### ModelProviderSummaryResponse + +Fields required to render the collapsed model-provider list. + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| configurate_methods | [ [ConfigurateMethod](#configuratemethod) ] | | Yes | +| custom_configuration | [ModelProviderCustomConfigurationSummaryResponse](#modelprovidercustomconfigurationsummaryresponse) | | Yes | +| description | [I18nObject](#i18nobject) | | No | +| icon_small | [I18nObject](#i18nobject) | | No | +| icon_small_dark | [I18nObject](#i18nobject) | | No | +| is_configured | boolean | | Yes | +| label | [I18nObject](#i18nobject) | | Yes | +| plugin_id | string | | Yes | +| preferred_provider_type | [ProviderType](#providertype) | | Yes | +| provider | string | | Yes | +| supported_model_types | [ [ModelType](#modeltype) ] | | Yes | +| system_configuration | [ModelProviderSystemConfigurationSummaryResponse](#modelprovidersystemconfigurationsummaryresponse) | | Yes | + +#### ModelProviderSystemConfigurationSummaryResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| enabled | boolean | | Yes | + #### ModelSelectorScope | Name | Type | Description | Required | @@ -20047,6 +20278,7 @@ Enum class for parameter type. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | | plugin_installation_id | string | | Yes | +| preserve_credentials | boolean | | No | #### ParserUpdateCredential @@ -20307,8 +20539,11 @@ Shared permission levels for resources (datasets, credentials, etc.) | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | +| language | string,
**Available values:** "en_US", "ja_JP", "pt_BR", "zh_Hans",
**Default:** en_US | Language used for localized label and description search
*Enum:* `"en_US"`, `"ja_JP"`, `"pt_BR"`, `"zh_Hans"` | No | | page | integer,
**Default:** 1 | Page number | No | | page_size | integer,
**Default:** 256 | Page size (1-256) | No | +| query | string | Case-insensitive search query | No | +| tags | [ string ] | Match any plugin tag | No | #### PluginCategoryListResponse @@ -20509,6 +20744,18 @@ Shared permission levels for resources (datasets, credentials, etc.) | ---- | ---- | ----------- | -------- | | plugins | [ [PluginInstallationItemResponse](#plugininstallationitemresponse) ] | | Yes | +#### PluginInstalledIdsQuery + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| category | [PluginCategory](#plugincategory) | Plugin category to include | Yes | + +#### PluginInstalledIdsResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| plugin_ids | [ string ] | | Yes | + #### PluginListResponse | Name | Type | Description | Required | @@ -21018,6 +21265,28 @@ Whitelist scopes accepted by RBAC app and dataset access config APIs. | result | string | | Yes | | updated_at | integer | | Yes | +#### RecentAppListResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| data | [ [RecentAppResponse](#recentappresponse) ] | | Yes | + +#### RecentAppResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| author_name | string | | No | +| icon | string | | No | +| icon_background | string | | No | +| icon_type | [IconType](#icontype) | | No | +| icon_url | string | | Yes | +| id | string | | Yes | +| maintainer | string | | No | +| mode | string,
**Available values:** "advanced-chat", "agent-chat", "chat", "completion", "workflow" | *Enum:* `"advanced-chat"`, `"agent-chat"`, `"chat"`, `"completion"`, `"workflow"` | Yes | +| name | string | | Yes | +| permission_keys | [ string ] | | No | +| updated_at | integer | | Yes | + #### RecommendedAppDetailNullableResponse | Name | Type | Description | Required | @@ -21028,7 +21297,7 @@ Whitelist scopes accepted by RBAC app and dataset access config APIs. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| can_trial | boolean | | No | +| can_trial | boolean | | Yes | | export_data | string | | Yes | | icon | string | | No | | icon_background | string | | No | @@ -21061,7 +21330,7 @@ Whitelist scopes accepted by RBAC app and dataset access config APIs. | ---- | ---- | ----------- | -------- | | app | [RecommendedAppInfoResponse](#recommendedappinforesponse) | | No | | app_id | string | | Yes | -| can_trial | boolean | | No | +| can_trial | boolean | | Yes | | categories | [ string ] | | No | | copyright | string | | No | | custom_disclaimer | string | | No | @@ -21295,6 +21564,18 @@ Whitelist scopes accepted by RBAC app and dataset access config APIs. | instruction | string | Structured output generation instruction | Yes | | model_config | [ModelConfig](#modelconfig) | Model configuration | Yes | +#### SSOProtocol + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| SSOProtocol | string | | | + +#### SandboxDownloadResponse + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| url | string | | Yes | + #### SandboxFileEntryResponse | Name | Type | Description | Required | @@ -21308,7 +21589,6 @@ Whitelist scopes accepted by RBAC app and dataset access config APIs. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| session_id | string | | Yes | | workspace_cwd | string | | Yes | #### SandboxListResponse @@ -21329,12 +21609,6 @@ Whitelist scopes accepted by RBAC app and dataset access config APIs. | text | string | | No | | truncated | boolean | | Yes | -#### SandboxUploadResponse - -| Name | Type | Description | Required | -| ---- | ---- | ----------- | -------- | -| url | string | | Yes | - #### SavedMessageCreatePayload | Name | Type | Description | Required | @@ -21853,6 +22127,7 @@ Query parameters for listing snippet published workflows. | updated_at | integer | | Yes | | updated_by | [SimpleAccountResponse](#simpleaccountresponse) | | No | | version | string | | Yes | +| version_number | integer | | No | #### StarredAppListQuery @@ -21954,7 +22229,7 @@ The subscription constructor of the trigger provider | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | | interval | string | | Yes | -| plan | string,
**Default:** sandbox | | Yes | +| plan | [CloudPlan](#cloudplan) | | Yes | #### SubscriptionQuery @@ -22016,7 +22291,7 @@ The subscription constructor of the trigger provider | ---- | ---- | ----------- | -------- | | _is_collaborative | boolean | | No | | conversation_variables | [ object ] | | No | -| environment_variables | [ object ] | | No | +| environment_variable_patch | [SyncEnvironmentVariablePatchPayload](#syncenvironmentvariablepatchpayload) | | No | | features | object | | Yes | | graph | object | | Yes | | hash | string | | No | @@ -22029,6 +22304,13 @@ The subscription constructor of the trigger provider | result | string | | Yes | | updated_at | integer | | Yes | +#### SyncEnvironmentVariablePatchPayload + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| deleted_environment_variable_ids | [ string ] | | No | +| environment_variables | [ object ] | | No | + #### SystemConfigurationResponse Model class for provider system configuration response. @@ -22058,7 +22340,6 @@ Non-sensitive bootstrap snapshot exposed before Console or Web authentication. | enable_marketplace | boolean | | Yes | | enable_social_oauth_login | boolean | | Yes | | enable_step_by_step_tour | boolean | | Yes | -| enable_trial_app | boolean | | Yes | | is_allow_register | boolean | | Yes | | is_email_setup | boolean | | Yes | | knowledge_fs_enabled | boolean | | Yes | @@ -22066,7 +22347,7 @@ Non-sensitive bootstrap snapshot exposed before Console or Web authentication. | plugin_installation_permission | [PluginInstallationPermissionModel](#plugininstallationpermissionmodel) | | Yes | | rbac_enabled | boolean | | Yes | | sso_enforced_for_signin | boolean | | Yes | -| sso_enforced_for_signin_protocol | string | | Yes | +| sso_enforced_for_signin_protocol | [SSOProtocol](#ssoprotocol) | | Yes | | webapp_auth | [WebAppAuthModel](#webappauthmodel) | | Yes | #### SystemParameters @@ -22162,7 +22443,7 @@ Tag type | in_trial | boolean | | No | | name | string | | No | | next_credit_reset_date | integer | | No | -| plan | string | | No | +| plan | [CloudPlan](#cloudplan) | | No | | role | string | | No | | status | string | | No | | trial_credits | integer | | No | @@ -22179,7 +22460,7 @@ Tag type | id | string | | Yes | | last_opened_at | integer | | No | | name | string | | No | -| plan | string | | No | +| plan | [CloudPlan](#cloudplan) | | No | | status | string | | No | #### TenantListResponse @@ -22931,7 +23212,9 @@ Payload for updating a snippet. | file_upload_limit | integer | | Yes | | image_file_batch_limit | integer | | Yes | | image_file_size_limit | integer | | Yes | +| knowledge_file_size_limit | integer | | Yes | | single_chunk_attachment_limit | integer | | Yes | +| skill_file_size_limit | integer | | Yes | | video_file_size_limit | integer | | Yes | | workflow_file_upload_limit | integer | | Yes | @@ -22993,11 +23276,19 @@ User action configuration. #### ValueSourceType ValueSourceType records whether the value comes from a static setting -in form definiton, or a variable while the workflow is running. +in form definition, or a variable while the workflow is running. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| ValueSourceType | string | ValueSourceType records whether the value comes from a static setting in form definiton, or a variable while the workflow is running. | | +| ValueSourceType | string | ValueSourceType records whether the value comes from a static setting in form definition, or a variable while the workflow is running. | | + +#### VectorSpaceLimitationModel + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| limit | integer | | Yes | +| size | integer | | Yes | +| usage_unknown | boolean | | No | #### VerificationTokenResponse @@ -23022,7 +23313,7 @@ in form definiton, or a variable while the workflow is running. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| protocol | string | | Yes | +| protocol | [SSOProtocol](#ssoprotocol) | | Yes | #### WebhookTriggerResponse @@ -23125,12 +23416,12 @@ How a workflow node is bound to an Agent. | variant | string | | Yes | | workflow_id | string | | No | -#### WorkflowAgentSandboxUploadPayload +#### WorkflowAgentSandboxDownloadPayload | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| node_execution_id | string | Optional workflow node execution ID. When omitted, the latest active session for the node is used. | No | -| path | string | File path relative to the sandbox workspace | Yes | +| node_execution_id | string | Workflow node execution ID | Yes | +| path | string | Binding path: relative paths start in Workspace; exact `~` and paths beginning with `~/` start in Home; `~user` is an ordinary relative path from Workspace; absolute paths remain absolute; `..` and paths outside Workspace are governed by backend isolation, not a Workspace-root restriction | Yes | #### WorkflowAppLogPaginationResponse @@ -23852,6 +24143,7 @@ tenant's default model. The underlying generator never raises — an empty | updated_at | integer | | Yes | | updated_by | [SimpleAccountResponse](#simpleaccountresponse) | | No | | version | string | | Yes | +| version_number | integer | | No | #### WorkflowRestoreResponse @@ -24450,7 +24742,7 @@ FastOpenAPI proof of concept for Dify API **Initialize system setup with admin account. NOTE: This endpoint is unauthenticated by design for first-time bootstrap. - Access is restricted by deployment mode (`SELF_HOSTED`), one-time setup guards, + Access is restricted to self-hosted editions (`COMMUNITY` and `ENTERPRISE`), one-time setup guards, and init-password validation rather than user session authentication. ** @@ -24536,19 +24828,9 @@ FastOpenAPI proof of concept for Dify API | setup_at | string | Setup completion time (ISO format) | No | | step | string,
**Available values:** "finished", "not_started" | Setup step status
*Enum:* `"finished"`, `"not_started"` | Yes | -###### VersionFeatures - -| Name | Type | Description | Required | -| ---- | ---- | ----------- | -------- | -| can_replace_logo | boolean | Whether logo replacement is supported | Yes | -| model_load_balancing_enabled | boolean | Whether model load balancing is enabled | Yes | - ###### VersionResponse | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| can_auto_update | boolean | Whether auto-update is supported | Yes | -| features | [VersionFeatures](#versionfeatures) | Feature flags and capabilities | Yes | -| release_date | string | Release date of latest version | Yes | | release_notes | string | Release notes for latest version | Yes | | version | string | Latest version number | Yes | diff --git a/api/openapi/markdown/openapi-openapi.md b/api/openapi/markdown/openapi-openapi.md index e67b1f87d10..75882587299 100644 --- a/api/openapi/markdown/openapi-openapi.md +++ b/api/openapi/markdown/openapi-openapi.md @@ -83,7 +83,7 @@ User-scoped operations | mode | query | App types the ``app`` usage face (``get app``) lists and filters. A curated subset of :class:`AppMode`: the real, user-facing app categories. Excludes runtime-only mode tags that are not standalone apps (``rag-pipeline`` is a knowledge ``Pipeline``; ``channel`` is unused) and the roster-owned ``agent`` type (surfaced through the roster, not this list). Members reference ``AppMode.*.value`` so the subset relationship is type-checked: dropping a member from ``AppMode`` breaks this at import. This is the single source for the listable set — params, filters, and the generated CLI whitelist all derive from it. | No | string,
**Available values:** "advanced-chat", "agent-chat", "chat", "completion", "workflow" | | name | query | | No | string | | page | query | | No | integer,
**Default:** 1 | -| workspace_id | query | | Yes | string | +| workspace_id | query | | Yes | string (uuid) | #### Responses @@ -129,7 +129,7 @@ User-scoped operations | Name | Located in | Description | Required | Schema | | ---- | ---------- | ----------- | -------- | ------ | | include_secret | query | Include encrypted secret values in the exported DSL | No | boolean | -| workflow_id | query | Export a specific workflow version instead of the current draft | No | string | +| workflow_id | query | Export a specific workflow version instead of the current draft | No | string (uuid) | | app_id | path | | Yes | string | #### Responses @@ -600,7 +600,7 @@ mode is a closed enum of listable app types. | mode | [SupportedAppType](#supportedapptype) | | No | | name | string | | No | | page | integer,
**Default:** 1 | | No | -| workspace_id | string | | Yes | +| workspace_id | string (uuid) | | Yes | #### AppListResponse @@ -648,6 +648,14 @@ mode is a closed enum of listable app types. | ---- | ---- | ----------- | -------- | | leaked_dependencies | [ [PluginDependency](#plugindependency) ] | | No | +#### DeploymentEdition + +Enum representing the deployment edition of the platform. + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| DeploymentEdition | string | Enum representing the deployment edition of the platform. | | + #### DeviceCodeRequest | Name | Type | Description | Required | @@ -974,7 +982,7 @@ Meta endpoint payload for `GET /openapi/v1/_version` — no auth required. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| edition | string,
**Available values:** "CLOUD", "SELF_HOSTED" | *Enum:* `"CLOUD"`, `"SELF_HOSTED"` | Yes | +| edition | [DeploymentEdition](#deploymentedition) | | Yes | | version | string | | Yes | #### SessionListQuery diff --git a/api/openapi/markdown/service-openapi.md b/api/openapi/markdown/service-openapi.md index 3cc8e4c3410..493fb19d316 100644 --- a/api/openapi/markdown/service-openapi.md +++ b/api/openapi/markdown/service-openapi.md @@ -1165,9 +1165,10 @@ Create a document by uploading a file. Supports common document formats (PDF, TX | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Document created successfully. | **application/json**: [DocumentAndBatchResponse](#documentandbatchresponse)
| -| 400 | - `no_file_uploaded` : Please upload your file. - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, file too large, unsupported file type, missing required fields, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | +| 400 | - `no_file_uploaded` : Please upload your file. - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, unsupported file type, missing required fields, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | | 401 | Unauthorized - invalid API token | | | 403 | Forbidden - dataset API access or workspace access denied | | +| 413 | `file_too_large` : File size exceeded. | | ### [POST] /datasets/{dataset_id}/document/create-by-text **Create Document by Text** @@ -1220,9 +1221,10 @@ Create a document by uploading a file. Supports common document formats (PDF, TX | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Document created successfully. | **application/json**: [DocumentAndBatchResponse](#documentandbatchresponse)
| -| 400 | - `no_file_uploaded` : Please upload your file. - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, file too large, unsupported file type, missing required fields, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | +| 400 | - `no_file_uploaded` : Please upload your file. - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, unsupported file type, missing required fields, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | | 401 | Unauthorized - invalid API token | | | 403 | Forbidden - dataset API access or workspace access denied | | +| 413 | `file_too_large` : File size exceeded. | | ### [GET] /datasets/{dataset_id}/documents **List Documents** @@ -1391,10 +1393,11 @@ Update an existing document by uploading a new file. Re-triggers indexing — us | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Document updated successfully. | **application/json**: [DocumentAndBatchResponse](#documentandbatchresponse)
| -| 400 | - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, file too large, unsupported file type, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | +| 400 | - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, unsupported file type, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | | 401 | Unauthorized - invalid API token | | | 403 | Forbidden - dataset API access or workspace access denied | | | 404 | Document not found | | +| 413 | `file_too_large` : File size exceeded. | | ### [GET] /datasets/{dataset_id}/documents/{document_id}/download **Download Document** @@ -1443,10 +1446,11 @@ Update an existing document by uploading a new file. Re-triggers indexing — us | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Document updated successfully. | **application/json**: [DocumentAndBatchResponse](#documentandbatchresponse)
| -| 400 | - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, file too large, unsupported file type, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | +| 400 | - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, unsupported file type, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | | 401 | Unauthorized - invalid API token | | | 403 | Forbidden - dataset API access or workspace access denied | | | 404 | Document not found | | +| 413 | `file_too_large` : File size exceeded. | | ### [POST] /datasets/{dataset_id}/documents/{document_id}/update-by-text **Update Document by Text** @@ -1502,10 +1506,11 @@ Update an existing document by uploading a new file. Re-triggers indexing — us | Code | Description | Schema | | ---- | ----------- | ------ | | 200 | Document updated successfully. | **application/json**: [DocumentAndBatchResponse](#documentandbatchresponse)
| -| 400 | - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, file too large, unsupported file type, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | +| 400 | - `too_many_files` : Only one file is allowed. - `filename_not_exists_error` : The specified filename does not exist. - `provider_not_initialize` : No valid model provider credentials found. Please go to Settings -> Model Provider to complete your provider credentials. - `invalid_param` : Knowledge base does not exist, external datasets not supported, unsupported file type, or invalid doc_form (must be `text_model`, `hierarchical_model`, or `qa_model`). | | | 401 | Unauthorized - invalid API token | | | 403 | Forbidden - dataset API access or workspace access denied | | | 404 | Document not found | | +| 413 | `file_too_large` : File size exceeded. | | --- ## default @@ -1534,7 +1539,7 @@ Update metadata values for multiple documents at once. Each document in the requ | 200 | Document metadata updated successfully. | **application/json**: [DatasetMetadataActionResponse](#datasetmetadataactionresponse)
| | 401 | Unauthorized - invalid API token | | | 403 | Forbidden - dataset API access or workspace access denied | | -| 404 | Dataset not found | | +| 404 | Dataset, document, or metadata not found | | ### [GET] /datasets/{dataset_id}/metadata **List Metadata Fields** @@ -2995,7 +3000,7 @@ Request payload for bulk downloading documents as a zip archive. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| document_id | string | Document ID whose metadata should be updated. | Yes | +| document_id | string (uuid) | Document ID whose metadata should be updated. | Yes | | metadata_list | [ [MetadataDetail](#metadatadetail) ] | Metadata fields to update. | Yes | | partial_update | boolean | Whether to partially update metadata, keeping existing values for unspecified fields. | No | @@ -3057,6 +3062,8 @@ Request payload for bulk downloading documents as a zip archive. | completed_at | integer | | Yes | | completed_segments | integer | | No | | error | string | | Yes | +| error_code | string | | No | +| estimated_vector_space_mb | integer | | No | | id | string | | Yes | | indexing_status | string | | Yes | | parsing_completed_at | integer | | Yes | @@ -3065,6 +3072,7 @@ Request payload for bulk downloading documents as a zip archive. | splitting_completed_at | integer | | Yes | | stopped_at | integer | | Yes | | total_segments | integer | | No | +| vector_space_limit_mb | integer | | No | #### DocumentTextCreatePayload @@ -3505,7 +3513,7 @@ Model class for i18n object. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| id | string | Metadata field ID. | Yes | +| id | string (uuid) | Metadata field ID. | Yes | | name | string | Metadata field name. | Yes | | value | string
integer
number | Metadata value. Can be a string, number, or `null`. | No | @@ -4052,11 +4060,11 @@ User action configuration. #### ValueSourceType ValueSourceType records whether the value comes from a static setting -in form definiton, or a variable while the workflow is running. +in form definition, or a variable while the workflow is running. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| ValueSourceType | string | ValueSourceType records whether the value comes from a static setting in form definiton, or a variable while the workflow is running. | | +| ValueSourceType | string | ValueSourceType records whether the value comes from a static setting in form definition, or a variable while the workflow is running. | | #### WeightKeywordSetting diff --git a/api/openapi/markdown/web-openapi.md b/api/openapi/markdown/web-openapi.md index 4b292991a63..36f2f7a084a 100644 --- a/api/openapi/markdown/web-openapi.md +++ b/api/openapi/markdown/web-openapi.md @@ -965,6 +965,12 @@ Returns Server-Sent Events stream. | ---- | ---- | ----------- | -------- | | tool_icons | object | Tool icon metadata keyed by tool name | No | +#### AppMode + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| AppMode | string | | | + #### AppPermissionQuery | Name | Type | Description | Required | @@ -1442,6 +1448,12 @@ Form input definition. | summary | string | | No | | word_count | integer | | No | +#### SSOProtocol + +| Name | Type | Description | Required | +| ---- | ---- | ----------- | -------- | +| SSOProtocol | string | | | + #### SavedMessageCreatePayload | Name | Type | Description | Required | @@ -1557,7 +1569,6 @@ Non-sensitive bootstrap snapshot exposed before Console or Web authentication. | enable_marketplace | boolean | | Yes | | enable_social_oauth_login | boolean | | Yes | | enable_step_by_step_tour | boolean | | Yes | -| enable_trial_app | boolean | | Yes | | is_allow_register | boolean | | Yes | | is_email_setup | boolean | | Yes | | knowledge_fs_enabled | boolean | | Yes | @@ -1565,7 +1576,7 @@ Non-sensitive bootstrap snapshot exposed before Console or Web authentication. | plugin_installation_permission | [PluginInstallationPermissionModel](#plugininstallationpermissionmodel) | | Yes | | rbac_enabled | boolean | | Yes | | sso_enforced_for_signin | boolean | | Yes | -| sso_enforced_for_signin_protocol | string | | Yes | +| sso_enforced_for_signin_protocol | [SSOProtocol](#ssoprotocol) | | Yes | | webapp_auth | [WebAppAuthModel](#webappauthmodel) | | Yes | #### SystemParameters @@ -1600,11 +1611,11 @@ User action configuration. #### ValueSourceType ValueSourceType records whether the value comes from a static setting -in form definiton, or a variable while the workflow is running. +in form definition, or a variable while the workflow is running. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| ValueSourceType | string | ValueSourceType records whether the value comes from a static setting in form definiton, or a variable while the workflow is running. | | +| ValueSourceType | string | ValueSourceType records whether the value comes from a static setting in form definition, or a variable while the workflow is running. | | #### VerificationTokenResponse @@ -1629,7 +1640,7 @@ in form definiton, or a variable while the workflow is running. | Name | Type | Description | Required | | ---- | ---- | ----------- | -------- | -| protocol | string | | Yes | +| protocol | [SSOProtocol](#ssoprotocol) | | Yes | #### WebAppCustomConfigResponse @@ -1647,6 +1658,7 @@ in form definiton, or a variable while the workflow is running. | custom_config | [WebAppCustomConfigResponse](#webappcustomconfigresponse) | | No | | enable_site | boolean | | Yes | | end_user_id | string | | No | +| mode | [AppMode](#appmode) | | Yes | | model_config | [WebModelConfigResponse](#webmodelconfigresponse) | | No | | plan | string | | Yes | | site | [WebSiteResponse](#websiteresponse) | | Yes | diff --git a/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/aliyun_trace.py b/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/aliyun_trace.py index b2de91860e0..8a519a04ad5 100644 --- a/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/aliyun_trace.py +++ b/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/aliyun_trace.py @@ -1,6 +1,6 @@ import logging -from collections.abc import Sequence -from typing import override +from collections.abc import Mapping, Sequence +from typing import Any, override from opentelemetry.trace import SpanKind from sqlalchemy.orm import sessionmaker @@ -29,33 +29,50 @@ from dify_trace_aliyun.data_exporter.traceclient import ( from dify_trace_aliyun.entities.aliyun_trace_entity import SpanData, TraceMetadata from dify_trace_aliyun.entities.semconv import ( DIFY_APP_ID, + GEN_AI_AGENT_NAME, GEN_AI_COMPLETION, GEN_AI_INPUT_MESSAGE, + GEN_AI_OPERATION_NAME, GEN_AI_OUTPUT_MESSAGE, GEN_AI_PROMPT, GEN_AI_PROVIDER_NAME, + GEN_AI_REACT_FINISH_REASON, + GEN_AI_REACT_ROUND, GEN_AI_REQUEST_MODEL, GEN_AI_RESPONSE_FINISH_REASON, + GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN, GEN_AI_USAGE_INPUT_TOKENS, GEN_AI_USAGE_OUTPUT_TOKENS, GEN_AI_USAGE_TOTAL_TOKENS, + OPERATION_NAME_CHAT, + OPERATION_NAME_INVOKE_AGENT, + OPERATION_NAME_REACT, RETRIEVAL_DOCUMENT, RETRIEVAL_QUERY, - TOOL_DESCRIPTION, - TOOL_NAME, - TOOL_PARAMETERS, GenAISpanKind, ) from dify_trace_aliyun.utils import ( + AgentLogEntry, + convert_seconds_to_nanoseconds, create_common_span_attributes, + create_gen_ai_tool_attributes, create_links_from_trace_id, + create_status_from_agent_log_entry, create_status_from_error, + extract_model_name_from_thought_label, + extract_react_round_number, extract_retrieval_documents, + extract_tool_description, + extract_tool_name_from_call_label, format_input_messages, format_output_messages, format_retrieval_documents, get_user_id_from_message_data, get_workflow_node_status, + is_llm_thought_entry, + is_tool_call_entry, + map_gen_ai_tool_type, + parse_agent_log_entries, serialize_json_data, ) from extensions.ext_database import db @@ -130,6 +147,9 @@ class AliyunDataTrace(BaseTraceInstance): for node_execution in workflow_node_executions: node_span = self.build_workflow_node_span(node_execution, trace_info, trace_metadata) self.trace_client.add_span(node_span) + if node_span is not None and node_execution.node_type == BuiltinNodeTypes.AGENT: + for react_span in self.build_agent_react_spans(node_execution, trace_metadata): + self.trace_client.add_span(react_span) def message_trace(self, trace_info: MessageTraceInfo): message_data = trace_info.message_data @@ -175,6 +195,28 @@ class AliyunDataTrace(BaseTraceInstance): ) self.trace_client.add_span(message_span) + llm_attributes: dict[str, Any] = { + **create_common_span_attributes( + session_id=trace_metadata.session_id, + user_id=trace_metadata.user_id, + span_kind=GenAISpanKind.LLM, + inputs=inputs_json, + outputs=outputs_str, + ), + GEN_AI_OPERATION_NAME: OPERATION_NAME_CHAT, + GEN_AI_REQUEST_MODEL: trace_info.metadata.get("ls_model_name") or "", + GEN_AI_PROVIDER_NAME: trace_info.metadata.get("ls_provider") or "", + GEN_AI_USAGE_INPUT_TOKENS: str(trace_info.message_tokens), + GEN_AI_USAGE_OUTPUT_TOKENS: str(trace_info.answer_tokens), + GEN_AI_USAGE_TOTAL_TOKENS: str(trace_info.total_tokens), + GEN_AI_PROMPT: inputs_json, + GEN_AI_COMPLETION: outputs_str, + } + if trace_info.gen_ai_server_time_to_first_token is not None: + llm_attributes[GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN] = convert_seconds_to_nanoseconds( + trace_info.gen_ai_server_time_to_first_token + ) + llm_span = SpanData( trace_id=trace_metadata.trace_id, parent_span_id=message_span_id, @@ -182,22 +224,7 @@ class AliyunDataTrace(BaseTraceInstance): name="llm", start_time=convert_datetime_to_nanoseconds(trace_info.start_time), end_time=convert_datetime_to_nanoseconds(trace_info.end_time), - attributes={ - **create_common_span_attributes( - session_id=trace_metadata.session_id, - user_id=trace_metadata.user_id, - span_kind=GenAISpanKind.LLM, - inputs=inputs_json, - outputs=outputs_str, - ), - GEN_AI_REQUEST_MODEL: trace_info.metadata.get("ls_model_name") or "", - GEN_AI_PROVIDER_NAME: trace_info.metadata.get("ls_provider") or "", - GEN_AI_USAGE_INPUT_TOKENS: str(trace_info.message_tokens), - GEN_AI_USAGE_OUTPUT_TOKENS: str(trace_info.answer_tokens), - GEN_AI_USAGE_TOTAL_TOKENS: str(trace_info.total_tokens), - GEN_AI_PROMPT: inputs_json, - GEN_AI_COMPLETION: outputs_str, - }, + attributes=llm_attributes, status=status, links=trace_metadata.links, ) @@ -258,9 +285,11 @@ class AliyunDataTrace(BaseTraceInstance): links=create_links_from_trace_id(trace_info.trace_id), ) - tool_config_json = serialize_json_data(trace_info.tool_config) + tool_config = trace_info.tool_config if isinstance(trace_info.tool_config, Mapping) else {} tool_inputs_json = serialize_json_data(trace_info.tool_inputs) + tool_result = str(trace_info.tool_outputs) inputs_json = serialize_json_data(trace_info.inputs) + provider_type = tool_config.get("tool_provider_type") or tool_config.get("provider_type") tool_span = SpanData( trace_id=trace_metadata.trace_id, @@ -275,11 +304,16 @@ class AliyunDataTrace(BaseTraceInstance): user_id=trace_metadata.user_id, span_kind=GenAISpanKind.TOOL, inputs=inputs_json, - outputs=str(trace_info.tool_outputs), + outputs=tool_result, + ), + **create_gen_ai_tool_attributes( + tool_name=trace_info.tool_name, + tool_type=map_gen_ai_tool_type(str(provider_type) if provider_type else None), + tool_description=extract_tool_description(tool_config), + tool_call_id=str(trace_info.metadata.get("node_execution_id") or ""), + tool_call_arguments=tool_inputs_json, + tool_call_result=tool_result, ), - TOOL_NAME: trace_info.tool_name, - TOOL_DESCRIPTION: tool_config_json, - TOOL_PARAMETERS: tool_inputs_json, }, status=status, links=trace_metadata.links, @@ -316,6 +350,8 @@ class AliyunDataTrace(BaseTraceInstance): node_span = self.build_workflow_retrieval_span(trace_info, node_execution, trace_metadata) elif node_execution.node_type == BuiltinNodeTypes.TOOL: node_span = self.build_workflow_tool_span(trace_info, node_execution, trace_metadata) + elif node_execution.node_type == BuiltinNodeTypes.AGENT: + node_span = self.build_workflow_agent_span(trace_info, node_execution, trace_metadata) else: node_span = self.build_workflow_task_span(trace_info, node_execution, trace_metadata) return node_span @@ -349,12 +385,15 @@ class AliyunDataTrace(BaseTraceInstance): def build_workflow_tool_span( self, trace_info: WorkflowTraceInfo, node_execution: WorkflowNodeExecution, trace_metadata: TraceMetadata ) -> SpanData: - tool_des = {} - if node_execution.metadata: - tool_des = node_execution.metadata.get(WorkflowNodeExecutionMetadataKey.TOOL_INFO, {}) + tool_info: Mapping[str, Any] = {} + if isinstance(node_execution.metadata, Mapping): + raw_tool_info = node_execution.metadata.get(WorkflowNodeExecutionMetadataKey.TOOL_INFO, {}) + if isinstance(raw_tool_info, Mapping): + tool_info = raw_tool_info inputs_json = serialize_json_data(node_execution.inputs or {}) outputs_json = serialize_json_data(node_execution.outputs) + provider_type = tool_info.get("provider_type") return SpanData( trace_id=trace_metadata.trace_id, @@ -371,9 +410,14 @@ class AliyunDataTrace(BaseTraceInstance): inputs=inputs_json, outputs=outputs_json, ), - TOOL_NAME: node_execution.title, - TOOL_DESCRIPTION: serialize_json_data(tool_des), - TOOL_PARAMETERS: inputs_json, + **create_gen_ai_tool_attributes( + tool_name=node_execution.title, + tool_type=map_gen_ai_tool_type(str(provider_type) if provider_type else None), + tool_description=extract_tool_description(tool_info), + tool_call_id=str(node_execution.id or ""), + tool_call_arguments=inputs_json, + tool_call_result=outputs_json, + ), }, status=get_workflow_node_status(node_execution), links=trace_metadata.links, @@ -414,16 +458,58 @@ class AliyunDataTrace(BaseTraceInstance): def build_workflow_llm_span( self, trace_info: WorkflowTraceInfo, node_execution: WorkflowNodeExecution, trace_metadata: TraceMetadata ) -> SpanData: - process_data = node_execution.process_data or {} - outputs = node_execution.outputs or {} + process_data = node_execution.process_data if isinstance(node_execution.process_data, Mapping) else {} + inputs = node_execution.inputs if isinstance(node_execution.inputs, Mapping) else {} + outputs = node_execution.outputs if isinstance(node_execution.outputs, Mapping) else {} usage_data = process_data.get("usage", {}) if "usage" in process_data else outputs.get("usage", {}) + if not isinstance(usage_data, Mapping): + usage_data = {} - prompts_json = serialize_json_data(process_data.get("prompts", [])) - text_output = str(outputs.get("text", "")) + # On invoke failure graphon leaves process_data empty, but prep already wrote + # model identity (and template variables / context) into node inputs. + prompts = process_data.get("prompts") or [] + if prompts: + prompts_json = serialize_json_data(prompts) + elif inputs: + prompts_json = serialize_json_data(inputs) + else: + prompts_json = serialize_json_data([]) + + text_output = str(outputs.get("text") or "") + if not text_output: + text_output = str(outputs.get("error_message") or node_execution.error or "") + + finish_reason = outputs.get("finish_reason") or outputs.get("error_type") or "" + model_name = process_data.get("model_name") or inputs.get("model_name") or "" + model_provider = process_data.get("model_provider") or inputs.get("model_provider") or "" gen_ai_input_message = format_input_messages(process_data) gen_ai_output_message = format_output_messages(outputs) + attributes: dict[str, Any] = { + **create_common_span_attributes( + session_id=trace_metadata.session_id, + user_id=trace_metadata.user_id, + span_kind=GenAISpanKind.LLM, + inputs=prompts_json, + outputs=text_output, + ), + GEN_AI_OPERATION_NAME: OPERATION_NAME_CHAT, + GEN_AI_REQUEST_MODEL: str(model_name), + GEN_AI_PROVIDER_NAME: str(model_provider), + GEN_AI_USAGE_INPUT_TOKENS: str(usage_data.get("prompt_tokens", 0)), + GEN_AI_USAGE_OUTPUT_TOKENS: str(usage_data.get("completion_tokens", 0)), + GEN_AI_USAGE_TOTAL_TOKENS: str(usage_data.get("total_tokens", 0)), + GEN_AI_PROMPT: prompts_json, + GEN_AI_COMPLETION: text_output, + GEN_AI_RESPONSE_FINISH_REASON: str(finish_reason), + GEN_AI_INPUT_MESSAGE: gen_ai_input_message, + GEN_AI_OUTPUT_MESSAGE: gen_ai_output_message, + } + time_to_first_token = usage_data.get("time_to_first_token") + if isinstance(time_to_first_token, (int, float)): + attributes[GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN] = convert_seconds_to_nanoseconds(float(time_to_first_token)) + return SpanData( trace_id=trace_metadata.trace_id, parent_span_id=trace_metadata.workflow_span_id, @@ -431,26 +517,225 @@ class AliyunDataTrace(BaseTraceInstance): name=node_execution.title, start_time=convert_datetime_to_nanoseconds(node_execution.created_at), end_time=convert_datetime_to_nanoseconds(node_execution.finished_at), + attributes=attributes, + status=get_workflow_node_status(node_execution), + links=trace_metadata.links, + ) + + def build_workflow_agent_span( + self, trace_info: WorkflowTraceInfo, node_execution: WorkflowNodeExecution, trace_metadata: TraceMetadata + ) -> SpanData: + """Build an AGENT-kind span for an agent-strategy node (instead of a generic TASK span).""" + inputs_json = serialize_json_data(node_execution.inputs) + outputs = node_execution.outputs if isinstance(node_execution.outputs, Mapping) else {} + usage_data = outputs.get("usage", {}) + if not isinstance(usage_data, Mapping): + usage_data = {} + text_output = str(outputs.get("text", "")) + + attributes: dict[str, Any] = { + **create_common_span_attributes( + session_id=trace_metadata.session_id, + user_id=trace_metadata.user_id, + span_kind=GenAISpanKind.AGENT, + inputs=inputs_json, + outputs=text_output, + ), + GEN_AI_OPERATION_NAME: OPERATION_NAME_INVOKE_AGENT, + GEN_AI_AGENT_NAME: node_execution.title, + GEN_AI_USAGE_INPUT_TOKENS: str(usage_data.get("prompt_tokens", 0)), + GEN_AI_USAGE_OUTPUT_TOKENS: str(usage_data.get("completion_tokens", 0)), + GEN_AI_USAGE_TOTAL_TOKENS: str(usage_data.get("total_tokens", 0)), + } + time_to_first_token = usage_data.get("time_to_first_token") + if isinstance(time_to_first_token, (int, float)): + attributes[GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN] = convert_seconds_to_nanoseconds(float(time_to_first_token)) + + return SpanData( + trace_id=trace_metadata.trace_id, + parent_span_id=trace_metadata.workflow_span_id, + span_id=convert_to_span_id(node_execution.id, "node"), + name=node_execution.title, + start_time=convert_datetime_to_nanoseconds(node_execution.created_at), + end_time=convert_datetime_to_nanoseconds(node_execution.finished_at), + attributes=attributes, + status=get_workflow_node_status(node_execution), + links=trace_metadata.links, + ) + + def build_agent_react_spans( + self, node_execution: WorkflowNodeExecution, trace_metadata: TraceMetadata + ) -> list[SpanData]: + """Build ReAct STEP spans and child LLM / TOOL spans from the agent execution log. + + The agent log lives in ``outputs["json"]``; ``started_at``/``finished_at`` there are + monotonic-clock seconds, so they are mapped onto wall-clock time by anchoring the + earliest ``started_at`` to the node's start time. Entries without timing fall back + to the node's start/end times. Returns an empty list when no log is available. + """ + try: + outputs = node_execution.outputs or {} + round_entries = parse_agent_log_entries(outputs) + if not round_entries: + return [] + + agent_span_id = convert_to_span_id(node_execution.id, "node") + node_start_ns = convert_datetime_to_nanoseconds(node_execution.created_at) + node_end_ns = convert_datetime_to_nanoseconds(node_execution.finished_at) + + monotonic_starts = [ + entry.metadata["started_at"] + for round_entry in round_entries + for entry in [round_entry, *round_entry.children] + if isinstance(entry.metadata.get("started_at"), (int, float)) + ] + base_monotonic = min(monotonic_starts) if monotonic_starts else None + + def to_wall_clock_ns(monotonic_seconds: Any, fallback: int | None) -> int | None: + if ( + isinstance(monotonic_seconds, (int, float)) + and base_monotonic is not None + and node_start_ns is not None + ): + return node_start_ns + convert_seconds_to_nanoseconds(float(monotonic_seconds) - base_monotonic) + return fallback + + spans: list[SpanData] = [] + for index, round_entry in enumerate(round_entries, start=1): + round_number = extract_react_round_number(round_entry.label, index) + step_span_id = generate_span_id() + step_attributes: dict[str, Any] = { + **create_common_span_attributes( + session_id=trace_metadata.session_id, + user_id=trace_metadata.user_id, + span_kind=GenAISpanKind.STEP, + inputs="", + outputs=serialize_json_data(round_entry.data), + ), + GEN_AI_OPERATION_NAME: OPERATION_NAME_REACT, + GEN_AI_REACT_ROUND: round_number, + } + if round_entry.error: + step_attributes[GEN_AI_REACT_FINISH_REASON] = "error" + spans.append( + SpanData( + trace_id=trace_metadata.trace_id, + parent_span_id=agent_span_id, + span_id=step_span_id, + name=round_entry.label or f"react step {round_number}", + start_time=to_wall_clock_ns(round_entry.metadata.get("started_at"), node_start_ns), + end_time=to_wall_clock_ns(round_entry.metadata.get("finished_at"), node_end_ns), + attributes=step_attributes, + status=create_status_from_agent_log_entry(round_entry), + links=trace_metadata.links, + ) + ) + + for child in round_entry.children: + child_start = to_wall_clock_ns(child.metadata.get("started_at"), node_start_ns) + child_end = to_wall_clock_ns(child.metadata.get("finished_at"), node_end_ns) + if is_tool_call_entry(child): + spans.append( + self._build_agent_tool_call_span( + entry=child, + step_span_id=step_span_id, + trace_metadata=trace_metadata, + start_time=child_start, + end_time=child_end, + ) + ) + elif is_llm_thought_entry(child): + spans.append( + self._build_agent_llm_call_span( + entry=child, + step_span_id=step_span_id, + trace_metadata=trace_metadata, + start_time=child_start, + end_time=child_end, + ) + ) + return spans + except Exception as e: + logger.warning("Error occurred in build_agent_react_spans: %s", e, exc_info=True) + return [] + + def _build_agent_llm_call_span( + self, + entry: AgentLogEntry, + step_span_id: int, + trace_metadata: TraceMetadata, + start_time: int | None, + end_time: int | None, + ) -> SpanData: + completion = str(entry.data.get("thought") or entry.data.get("action") or "") + return SpanData( + trace_id=trace_metadata.trace_id, + parent_span_id=step_span_id, + span_id=generate_span_id(), + name=entry.label or "llm", + start_time=start_time, + end_time=end_time, attributes={ **create_common_span_attributes( session_id=trace_metadata.session_id, user_id=trace_metadata.user_id, span_kind=GenAISpanKind.LLM, - inputs=prompts_json, - outputs=text_output, + inputs="", + outputs=serialize_json_data(entry.data), ), - GEN_AI_REQUEST_MODEL: process_data.get("model_name") or "", - GEN_AI_PROVIDER_NAME: process_data.get("model_provider") or "", - GEN_AI_USAGE_INPUT_TOKENS: str(usage_data.get("prompt_tokens", 0)), - GEN_AI_USAGE_OUTPUT_TOKENS: str(usage_data.get("completion_tokens", 0)), - GEN_AI_USAGE_TOTAL_TOKENS: str(usage_data.get("total_tokens", 0)), - GEN_AI_PROMPT: prompts_json, - GEN_AI_COMPLETION: text_output, - GEN_AI_RESPONSE_FINISH_REASON: outputs.get("finish_reason") or "", - GEN_AI_INPUT_MESSAGE: gen_ai_input_message, - GEN_AI_OUTPUT_MESSAGE: gen_ai_output_message, + GEN_AI_OPERATION_NAME: OPERATION_NAME_CHAT, + GEN_AI_REQUEST_MODEL: extract_model_name_from_thought_label(entry.label), + GEN_AI_PROVIDER_NAME: str(entry.metadata.get("provider") or ""), + GEN_AI_USAGE_TOTAL_TOKENS: str(entry.metadata.get("total_tokens", 0)), + GEN_AI_COMPLETION: completion, }, - status=get_workflow_node_status(node_execution), + status=create_status_from_agent_log_entry(entry), + links=trace_metadata.links, + ) + + def _build_agent_tool_call_span( + self, + entry: AgentLogEntry, + step_span_id: int, + trace_metadata: TraceMetadata, + start_time: int | None, + end_time: int | None, + ) -> SpanData: + tool_name = str(entry.data.get("tool_name") or extract_tool_name_from_call_label(entry.label) or "tool") + tool_parameters = entry.data.get("tool_call_args") + if tool_parameters is None: + tool_parameters = entry.data.get("tool_call_input") + if tool_parameters is None: + tool_parameters = entry.data + tool_arguments_json = serialize_json_data(tool_parameters) + tool_result = entry.data.get("output", entry.data) + tool_result_json = tool_result if isinstance(tool_result, str) else serialize_json_data(tool_result) + provider_type = entry.metadata.get("provider_type") or entry.data.get("provider_type") + return SpanData( + trace_id=trace_metadata.trace_id, + parent_span_id=step_span_id, + span_id=generate_span_id(), + name=entry.label or f"CALL {tool_name}", + start_time=start_time, + end_time=end_time, + attributes={ + **create_common_span_attributes( + session_id=trace_metadata.session_id, + user_id=trace_metadata.user_id, + span_kind=GenAISpanKind.TOOL, + inputs=tool_arguments_json, + outputs=tool_result_json, + ), + **create_gen_ai_tool_attributes( + tool_name=tool_name, + tool_type=map_gen_ai_tool_type(str(provider_type) if provider_type else None), + tool_description=extract_tool_description(entry.data) or extract_tool_description(entry.metadata), + tool_call_id=entry.id, + tool_call_arguments=tool_arguments_json, + tool_call_result=tool_result_json, + ), + }, + status=create_status_from_agent_log_entry(entry), links=trace_metadata.links, ) diff --git a/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/data_exporter/traceclient.py b/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/data_exporter/traceclient.py index 00aab6bf891..374d20ade9d 100644 --- a/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/data_exporter/traceclient.py +++ b/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/data_exporter/traceclient.py @@ -53,7 +53,7 @@ class TraceClient: attributes={ service_attributes.SERVICE_NAME: service_name, service_attributes.SERVICE_VERSION: f"dify-{dify_config.project.version}-{dify_config.COMMIT_SHA}", - DEPLOYMENT_ENVIRONMENT: f"{dify_config.DEPLOY_ENV}-{dify_config.EDITION}", + DEPLOYMENT_ENVIRONMENT: f"{dify_config.DEPLOY_ENV}-{dify_config.DEPLOYMENT_EDITION.value}", HOST_NAME: socket.gethostname(), ACS_ARMS_SERVICE_FEATURE: "genai_app", } diff --git a/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/entities/semconv.py b/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/entities/semconv.py index b6e46c5262a..19925410a64 100644 --- a/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/entities/semconv.py +++ b/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/entities/semconv.py @@ -11,6 +11,7 @@ GEN_AI_SESSION_ID: Final[str] = "gen_ai.session.id" GEN_AI_USER_ID: Final[str] = "gen_ai.user.id" GEN_AI_USER_NAME: Final[str] = "gen_ai.user.name" GEN_AI_SPAN_KIND: Final[str] = "gen_ai.span.kind" +GEN_AI_OPERATION_NAME: Final[str] = "gen_ai.operation.name" GEN_AI_FRAMEWORK: Final[str] = "gen_ai.framework" # Chain attributes @@ -30,14 +31,43 @@ GEN_AI_USAGE_TOTAL_TOKENS: Final[str] = "gen_ai.usage.total_tokens" GEN_AI_PROMPT: Final[str] = "gen_ai.prompt" GEN_AI_COMPLETION: Final[str] = "gen_ai.completion" GEN_AI_RESPONSE_FINISH_REASON: Final[str] = "gen_ai.response.finish_reason" +# Time to first token of the model response in a streaming scenario, in nanoseconds. +GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN: Final[str] = "gen_ai.response.time_to_first_token" GEN_AI_INPUT_MESSAGE: Final[str] = "gen_ai.input.messages" GEN_AI_OUTPUT_MESSAGE: Final[str] = "gen_ai.output.messages" -# Tool attributes -TOOL_NAME: Final[str] = "tool.name" -TOOL_DESCRIPTION: Final[str] = "tool.description" -TOOL_PARAMETERS: Final[str] = "tool.parameters" +# Tool attributes (GenAI semantic conventions) +GEN_AI_TOOL_CALL_ID: Final[str] = "gen_ai.tool.call.id" +GEN_AI_TOOL_DESCRIPTION: Final[str] = "gen_ai.tool.description" +GEN_AI_TOOL_NAME: Final[str] = "gen_ai.tool.name" +GEN_AI_TOOL_TYPE: Final[str] = "gen_ai.tool.type" +GEN_AI_TOOL_CALL_ARGUMENTS: Final[str] = "gen_ai.tool.call.arguments" +GEN_AI_TOOL_CALL_RESULT: Final[str] = "gen_ai.tool.call.result" + +# Skill attributes (conditionally required when loading a Skill) +GEN_AI_SKILL_ID: Final[str] = "gen_ai.skill.id" +GEN_AI_SKILL_NAME: Final[str] = "gen_ai.skill.name" +GEN_AI_SKILL_DESCRIPTION: Final[str] = "gen_ai.skill.description" +GEN_AI_SKILL_VERSION: Final[str] = "gen_ai.skill.version" + +# Agent attributes +GEN_AI_AGENT_NAME: Final[str] = "gen_ai.agent.name" + +# ReAct step attributes +GEN_AI_REACT_ROUND: Final[str] = "gen_ai.react.round" +GEN_AI_REACT_FINISH_REASON: Final[str] = "gen_ai.react.finish_reason" + +# gen_ai.operation.name values (see Aliyun LLM Trace field definitions) +OPERATION_NAME_CHAT: Final[str] = "chat" +OPERATION_NAME_EXECUTE_TOOL: Final[str] = "execute_tool" +OPERATION_NAME_INVOKE_AGENT: Final[str] = "invoke_agent" +OPERATION_NAME_REACT: Final[str] = "react" + +# gen_ai.tool.type values +TOOL_TYPE_FUNCTION: Final[str] = "function" +TOOL_TYPE_EXTENSION: Final[str] = "extension" +TOOL_TYPE_DATASTORE: Final[str] = "datastore" class GenAISpanKind(StrEnum): @@ -49,3 +79,5 @@ class GenAISpanKind(StrEnum): TOOL = "TOOL" AGENT = "AGENT" TASK = "TASK" + # Marks one Reasoning-Acting iteration of an agent (ReAct step). + STEP = "STEP" diff --git a/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/utils.py b/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/utils.py index 5678c66adbf..c9145a95b54 100644 --- a/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/utils.py +++ b/api/providers/trace/trace-aliyun/src/dify_trace_aliyun/utils.py @@ -1,5 +1,7 @@ import json +import re from collections.abc import Mapping +from dataclasses import dataclass, field from typing import Any, TypedDict from opentelemetry.trace import Link, Status, StatusCode @@ -7,11 +9,22 @@ from opentelemetry.trace import Link, Status, StatusCode from core.rag.models.document import Document from dify_trace_aliyun.entities.semconv import ( GEN_AI_FRAMEWORK, + GEN_AI_OPERATION_NAME, GEN_AI_SESSION_ID, GEN_AI_SPAN_KIND, + GEN_AI_TOOL_CALL_ARGUMENTS, + GEN_AI_TOOL_CALL_ID, + GEN_AI_TOOL_CALL_RESULT, + GEN_AI_TOOL_DESCRIPTION, + GEN_AI_TOOL_NAME, + GEN_AI_TOOL_TYPE, GEN_AI_USER_ID, INPUT_VALUE, + OPERATION_NAME_EXECUTE_TOOL, OUTPUT_VALUE, + TOOL_TYPE_DATASTORE, + TOOL_TYPE_EXTENSION, + TOOL_TYPE_FUNCTION, GenAISpanKind, ) from extensions.ext_database import db @@ -106,6 +119,47 @@ def create_common_span_attributes( } +def map_gen_ai_tool_type(provider_type: str | None) -> str: + """Map Dify tool provider type to GenAI ``gen_ai.tool.type`` values.""" + normalized = (provider_type or "").strip().lower() + if normalized in {"dataset-retrieval", "datastore"}: + return TOOL_TYPE_DATASTORE + if normalized == "extension": + return TOOL_TYPE_EXTENSION + return TOOL_TYPE_FUNCTION + + +def extract_tool_description(tool_meta: Mapping[str, Any] | None) -> str: + if not isinstance(tool_meta, Mapping): + return "" + for key in ("description", "tool_description"): + value = tool_meta.get(key) + if value: + return str(value) + return "" + + +def create_gen_ai_tool_attributes( + *, + tool_name: str, + tool_type: str = TOOL_TYPE_FUNCTION, + tool_description: str = "", + tool_call_id: str = "", + tool_call_arguments: str = "", + tool_call_result: str = "", +) -> dict[str, str]: + """Build GenAI-compliant attributes for a TOOL span.""" + return { + GEN_AI_OPERATION_NAME: OPERATION_NAME_EXECUTE_TOOL, + GEN_AI_TOOL_NAME: tool_name, + GEN_AI_TOOL_TYPE: tool_type or TOOL_TYPE_FUNCTION, + GEN_AI_TOOL_DESCRIPTION: tool_description, + GEN_AI_TOOL_CALL_ID: tool_call_id, + GEN_AI_TOOL_CALL_ARGUMENTS: tool_call_arguments, + GEN_AI_TOOL_CALL_RESULT: tool_call_result, + } + + def format_retrieval_documents(retrieval_documents: list) -> list: try: if not isinstance(retrieval_documents, list): @@ -174,6 +228,121 @@ def format_input_messages(process_data: Mapping[str, Any]) -> str: return serialize_json_data([]) +def convert_seconds_to_nanoseconds(seconds: float) -> int: + return int(seconds * 1e9) + + +_REACT_ROUND_LABEL_PATTERN = re.compile(r"ROUND\s+(\d+)", re.IGNORECASE) +_LLM_THOUGHT_LABEL_SUFFIX = " Thought" +_TOOL_CALL_LABEL_PREFIX = "CALL " + + +@dataclass +class AgentLogEntry: + """One entry of the agent-strategy execution log (``outputs["json"]`` of an agent node). + + Entries form a tree via ``parent_id``: top-level entries are ReAct rounds + (label like ``ROUND 1``) and their children are LLM thoughts / tool calls. + ``started_at``/``finished_at`` in ``metadata`` are monotonic-clock seconds + (``time.perf_counter``), not epoch timestamps. + """ + + id: str + parent_id: str | None + label: str + status: str + error: str | None + data: dict[str, Any] + metadata: dict[str, Any] + children: list["AgentLogEntry"] = field(default_factory=list) + + +def parse_agent_log_entries(outputs: Mapping[str, Any]) -> list[AgentLogEntry]: + """Parse agent node outputs into a tree of log entries, returning top-level rounds in order. + + Entries without an ``id`` (e.g. the trailing ``{"data": []}`` element) are skipped. + Children whose parent is missing are dropped. + """ + raw_entries = outputs.get("json") + if not isinstance(raw_entries, list): + return [] + + entries: list[AgentLogEntry] = [] + entries_by_id: dict[str, AgentLogEntry] = {} + for raw_entry in raw_entries: + if not isinstance(raw_entry, dict): + continue + entry_id = raw_entry.get("id") + if not entry_id: + continue + data = raw_entry.get("data") + metadata = raw_entry.get("metadata") + entry = AgentLogEntry( + id=str(entry_id), + parent_id=raw_entry.get("parent_id"), + label=str(raw_entry.get("label") or ""), + status=str(raw_entry.get("status") or ""), + error=raw_entry.get("error"), + data=data if isinstance(data, dict) else {}, + metadata=metadata if isinstance(metadata, dict) else {}, + ) + entries.append(entry) + entries_by_id[entry.id] = entry + + roots: list[AgentLogEntry] = [] + for entry in entries: + if entry.parent_id is None: + roots.append(entry) + else: + parent = entries_by_id.get(entry.parent_id) + if parent is not None: + parent.children.append(entry) + return roots + + +def extract_react_round_number(label: str, fallback: int) -> int: + match = _REACT_ROUND_LABEL_PATTERN.search(label) + if match: + return int(match.group(1)) + return fallback + + +def is_tool_call_entry(entry: AgentLogEntry) -> bool: + """Tool invocations use labels like ``CALL {tool_name}`` (also carry a provider in metadata).""" + return entry.label.startswith(_TOOL_CALL_LABEL_PREFIX) + + +def is_llm_thought_entry(entry: AgentLogEntry) -> bool: + """LLM thought entries use labels like ``{model} Thought``. + + Do not key off ``metadata.provider`` alone: tool CALL entries also set provider + (to the tool provider), which previously misclassified them as LLM spans. + """ + if is_tool_call_entry(entry): + return False + return entry.label.endswith(_LLM_THOUGHT_LABEL_SUFFIX) or bool(entry.metadata.get("provider")) + + +def extract_model_name_from_thought_label(label: str) -> str: + if label.endswith(_LLM_THOUGHT_LABEL_SUFFIX): + return label.removesuffix(_LLM_THOUGHT_LABEL_SUFFIX) + return "" + + +def extract_tool_name_from_call_label(label: str) -> str: + if label.startswith(_TOOL_CALL_LABEL_PREFIX): + return label.removeprefix(_TOOL_CALL_LABEL_PREFIX).strip() + return "" + + +def create_status_from_agent_log_entry(entry: AgentLogEntry) -> Status: + if entry.error: + return Status(StatusCode.ERROR, str(entry.error)) + if entry.status == "success": + return Status(StatusCode.OK) + return Status(StatusCode.UNSET) + + def format_output_messages(outputs: Mapping[str, Any]) -> str: try: if not isinstance(outputs, dict): diff --git a/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/data_exporter/test_traceclient.py b/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/data_exporter/test_traceclient.py index 797134b3619..0048c2ea6c0 100644 --- a/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/data_exporter/test_traceclient.py +++ b/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/data_exporter/test_traceclient.py @@ -56,6 +56,7 @@ class TestTraceClient: assert client.worker_thread.is_alive() client.shutdown() + # pyrefly: ignore [unnecessary-comparison] assert client.done is True @patch("dify_trace_aliyun.data_exporter.traceclient.OTLPSpanExporter") diff --git a/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/entities/test_semconv.py b/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/entities/test_semconv.py index 9cab40748f6..e07aa718f29 100644 --- a/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/entities/test_semconv.py +++ b/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/entities/test_semconv.py @@ -1,27 +1,43 @@ from dify_trace_aliyun.entities.semconv import ( ACS_ARMS_SERVICE_FEATURE, + GEN_AI_AGENT_NAME, GEN_AI_COMPLETION, GEN_AI_FRAMEWORK, GEN_AI_INPUT_MESSAGE, + GEN_AI_OPERATION_NAME, GEN_AI_OUTPUT_MESSAGE, GEN_AI_PROMPT, GEN_AI_PROVIDER_NAME, + GEN_AI_REACT_FINISH_REASON, + GEN_AI_REACT_ROUND, GEN_AI_REQUEST_MODEL, GEN_AI_RESPONSE_FINISH_REASON, + GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN, GEN_AI_SESSION_ID, + GEN_AI_SKILL_DESCRIPTION, + GEN_AI_SKILL_ID, + GEN_AI_SKILL_NAME, + GEN_AI_SKILL_VERSION, GEN_AI_SPAN_KIND, + GEN_AI_TOOL_CALL_ARGUMENTS, + GEN_AI_TOOL_CALL_ID, + GEN_AI_TOOL_CALL_RESULT, + GEN_AI_TOOL_DESCRIPTION, + GEN_AI_TOOL_NAME, + GEN_AI_TOOL_TYPE, GEN_AI_USAGE_INPUT_TOKENS, GEN_AI_USAGE_OUTPUT_TOKENS, GEN_AI_USAGE_TOTAL_TOKENS, GEN_AI_USER_ID, GEN_AI_USER_NAME, INPUT_VALUE, + OPERATION_NAME_EXECUTE_TOOL, OUTPUT_VALUE, RETRIEVAL_DOCUMENT, RETRIEVAL_QUERY, - TOOL_DESCRIPTION, - TOOL_NAME, - TOOL_PARAMETERS, + TOOL_TYPE_DATASTORE, + TOOL_TYPE_EXTENSION, + TOOL_TYPE_FUNCTION, GenAISpanKind, ) @@ -47,9 +63,25 @@ def test_constants(): assert GEN_AI_RESPONSE_FINISH_REASON == "gen_ai.response.finish_reason" assert GEN_AI_INPUT_MESSAGE == "gen_ai.input.messages" assert GEN_AI_OUTPUT_MESSAGE == "gen_ai.output.messages" - assert TOOL_NAME == "tool.name" - assert TOOL_DESCRIPTION == "tool.description" - assert TOOL_PARAMETERS == "tool.parameters" + assert GEN_AI_TOOL_CALL_ID == "gen_ai.tool.call.id" + assert GEN_AI_TOOL_DESCRIPTION == "gen_ai.tool.description" + assert GEN_AI_TOOL_NAME == "gen_ai.tool.name" + assert GEN_AI_TOOL_TYPE == "gen_ai.tool.type" + assert GEN_AI_TOOL_CALL_ARGUMENTS == "gen_ai.tool.call.arguments" + assert GEN_AI_TOOL_CALL_RESULT == "gen_ai.tool.call.result" + assert GEN_AI_SKILL_ID == "gen_ai.skill.id" + assert GEN_AI_SKILL_NAME == "gen_ai.skill.name" + assert GEN_AI_SKILL_DESCRIPTION == "gen_ai.skill.description" + assert GEN_AI_SKILL_VERSION == "gen_ai.skill.version" + assert GEN_AI_OPERATION_NAME == "gen_ai.operation.name" + assert OPERATION_NAME_EXECUTE_TOOL == "execute_tool" + assert TOOL_TYPE_FUNCTION == "function" + assert TOOL_TYPE_EXTENSION == "extension" + assert TOOL_TYPE_DATASTORE == "datastore" + assert GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN == "gen_ai.response.time_to_first_token" + assert GEN_AI_AGENT_NAME == "gen_ai.agent.name" + assert GEN_AI_REACT_ROUND == "gen_ai.react.round" + assert GEN_AI_REACT_FINISH_REASON == "gen_ai.react.finish_reason" def test_gen_ai_span_kind_enum(): @@ -61,8 +93,9 @@ def test_gen_ai_span_kind_enum(): assert GenAISpanKind.TOOL == "TOOL" assert GenAISpanKind.AGENT == "AGENT" assert GenAISpanKind.TASK == "TASK" + assert GenAISpanKind.STEP == "STEP" # Verify iteration works (covers the class definition) kinds = list(GenAISpanKind) - assert len(kinds) == 8 + assert len(kinds) == 9 assert "LLM" in kinds diff --git a/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/test_aliyun_trace.py b/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/test_aliyun_trace.py index 06bdb61b4d1..cf9f5c41978 100644 --- a/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/test_aliyun_trace.py +++ b/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/test_aliyun_trace.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json import logging from datetime import UTC, datetime from types import SimpleNamespace @@ -12,18 +13,34 @@ from dify_trace_aliyun.aliyun_trace import AliyunDataTrace from dify_trace_aliyun.config import AliyunConfig from dify_trace_aliyun.entities.aliyun_trace_entity import SpanData, TraceMetadata from dify_trace_aliyun.entities.semconv import ( + GEN_AI_AGENT_NAME, GEN_AI_COMPLETION, GEN_AI_INPUT_MESSAGE, + GEN_AI_OPERATION_NAME, GEN_AI_OUTPUT_MESSAGE, GEN_AI_PROMPT, + GEN_AI_PROVIDER_NAME, + GEN_AI_REACT_ROUND, GEN_AI_REQUEST_MODEL, GEN_AI_RESPONSE_FINISH_REASON, + GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN, + GEN_AI_TOOL_CALL_ARGUMENTS, + GEN_AI_TOOL_CALL_ID, + GEN_AI_TOOL_CALL_RESULT, + GEN_AI_TOOL_DESCRIPTION, + GEN_AI_TOOL_NAME, + GEN_AI_TOOL_TYPE, GEN_AI_USAGE_TOTAL_TOKENS, + INPUT_VALUE, + OPERATION_NAME_CHAT, + OPERATION_NAME_EXECUTE_TOOL, + OPERATION_NAME_INVOKE_AGENT, + OPERATION_NAME_REACT, + OUTPUT_VALUE, RETRIEVAL_DOCUMENT, RETRIEVAL_QUERY, - TOOL_DESCRIPTION, - TOOL_NAME, - TOOL_PARAMETERS, + TOOL_TYPE_DATASTORE, + TOOL_TYPE_FUNCTION, GenAISpanKind, ) from opentelemetry.trace import Link, SpanContext, SpanKind, Status, StatusCode, TraceFlags @@ -397,8 +414,10 @@ def test_tool_trace_creates_span(trace_instance: AliyunDataTrace, monkeypatch: p _make_tool_trace_info( tool_name="my-tool", tool_inputs={"a": 1}, - tool_config={"description": "x"}, + tool_outputs="tool-out", + tool_config={"description": "x", "tool_provider_type": "builtin"}, inputs={"i": 1}, + metadata={"conversation_id": "conv", "user_id": "u", "node_execution_id": "exec-1"}, ) ) @@ -407,8 +426,13 @@ def test_tool_trace_creates_span(trace_instance: AliyunDataTrace, monkeypatch: p span = spans[0] assert span.name == "my-tool" assert span.status == status - assert span.attributes[TOOL_NAME] == "my-tool" - assert span.attributes[TOOL_DESCRIPTION] == '{"description": "x"}' + assert span.attributes[GEN_AI_OPERATION_NAME] == OPERATION_NAME_EXECUTE_TOOL + assert span.attributes[GEN_AI_TOOL_NAME] == "my-tool" + assert span.attributes[GEN_AI_TOOL_DESCRIPTION] == "x" + assert span.attributes[GEN_AI_TOOL_TYPE] == TOOL_TYPE_FUNCTION + assert span.attributes[GEN_AI_TOOL_CALL_ID] == "exec-1" + assert span.attributes[GEN_AI_TOOL_CALL_ARGUMENTS] == '{"a": 1}' + assert span.attributes[GEN_AI_TOOL_CALL_RESULT] == "tool-out" def test_get_workflow_node_executions_requires_app_id(trace_instance: AliyunDataTrace): @@ -534,20 +558,30 @@ def test_build_workflow_tool_span(trace_instance: AliyunDataTrace, monkeypatch: node_execution.outputs = {"b": 2} node_execution.created_at = _dt() node_execution.finished_at = _dt() - node_execution.metadata = {WorkflowNodeExecutionMetadataKey.TOOL_INFO: {"k": "v"}} + node_execution.metadata = { + WorkflowNodeExecutionMetadataKey.TOOL_INFO: { + "provider_type": "dataset-retrieval", + "description": "search docs", + } + } span = trace_instance.build_workflow_tool_span(_make_workflow_trace_info(), node_execution, trace_metadata) - assert span.attributes[TOOL_NAME] == "my-tool" - assert span.attributes[TOOL_DESCRIPTION] == '{"k": "v"}' - assert span.attributes[TOOL_PARAMETERS] == '{"a": 1}' + assert span.attributes[GEN_AI_OPERATION_NAME] == OPERATION_NAME_EXECUTE_TOOL + assert span.attributes[GEN_AI_TOOL_NAME] == "my-tool" + assert span.attributes[GEN_AI_TOOL_DESCRIPTION] == "search docs" + assert span.attributes[GEN_AI_TOOL_TYPE] == TOOL_TYPE_DATASTORE + assert span.attributes[GEN_AI_TOOL_CALL_ID] == "node-id" + assert span.attributes[GEN_AI_TOOL_CALL_ARGUMENTS] == '{"a": 1}' + assert span.attributes[GEN_AI_TOOL_CALL_RESULT] == '{"b": 2}' assert span.status.status_code == StatusCode.OK # Cover metadata is None and inputs is None node_execution.metadata = None node_execution.inputs = None span2 = trace_instance.build_workflow_tool_span(_make_workflow_trace_info(), node_execution, trace_metadata) - assert span2.attributes[TOOL_DESCRIPTION] == "{}" - assert span2.attributes[TOOL_PARAMETERS] == "{}" + assert span2.attributes[GEN_AI_TOOL_DESCRIPTION] == "" + assert span2.attributes[GEN_AI_TOOL_TYPE] == TOOL_TYPE_FUNCTION + assert span2.attributes[GEN_AI_TOOL_CALL_ARGUMENTS] == "{}" def test_build_workflow_retrieval_span(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch): @@ -592,6 +626,8 @@ def test_build_workflow_llm_span(trace_instance: AliyunDataTrace, monkeypatch: p node_execution = MagicMock(spec=WorkflowNodeExecution) node_execution.id = "node-id" node_execution.title = "llm" + node_execution.inputs = {} + node_execution.error = None node_execution.process_data = { "usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3}, "prompts": ["p"], @@ -618,6 +654,288 @@ def test_build_workflow_llm_span(trace_instance: AliyunDataTrace, monkeypatch: p assert span2.attributes[GEN_AI_USAGE_TOTAL_TOKENS] == "10" +def test_build_workflow_llm_span_falls_back_to_inputs_on_invoke_failure( + trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch +): + """Model invoke failures leave process_data empty; prep still populates inputs.""" + monkeypatch.setattr(aliyun_trace_module, "convert_to_span_id", lambda _, __: 9) + monkeypatch.setattr(aliyun_trace_module, "convert_datetime_to_nanoseconds", lambda _: 123) + monkeypatch.setattr(aliyun_trace_module, "get_workflow_node_status", lambda _: Status(StatusCode.ERROR, "boom")) + + trace_metadata = _make_trace_metadata() + node_execution = MagicMock(spec=WorkflowNodeExecution) + node_execution.id = "node-id" + node_execution.title = "llm" + node_execution.process_data = {} + node_execution.inputs = { + "model_name": "qwen-plus", + "model_provider": "langgenius/tongyi/tongyi", + "#context#": "ctx", + "query": "hello", + } + node_execution.outputs = {"error_message": "API rate limit", "error_type": "InvokeError"} + node_execution.error = "API rate limit" + node_execution.created_at = _dt() + node_execution.finished_at = _dt() + + span = trace_instance.build_workflow_llm_span(_make_workflow_trace_info(), node_execution, trace_metadata) + + expected_prompt = json.dumps(node_execution.inputs, ensure_ascii=False) + assert span.attributes[GEN_AI_REQUEST_MODEL] == "qwen-plus" + assert span.attributes[GEN_AI_PROVIDER_NAME] == "langgenius/tongyi/tongyi" + assert span.attributes[GEN_AI_PROMPT] == expected_prompt + assert span.attributes[INPUT_VALUE] == expected_prompt + assert span.attributes[GEN_AI_COMPLETION] == "API rate limit" + assert span.attributes[OUTPUT_VALUE] == "API rate limit" + assert span.attributes[GEN_AI_RESPONSE_FINISH_REASON] == "InvokeError" + # prompts were never persisted, so structured input.messages stays empty + assert span.attributes[GEN_AI_INPUT_MESSAGE] == "[]" + + +def _make_agent_outputs() -> dict: + return { + "text": " pong!", + "usage": { + "prompt_tokens": 100, + "completion_tokens": 225, + "total_tokens": 325, + "time_to_first_token": 0.5, + }, + "files": [], + "json": [ + { + "id": "round-1", + "parent_id": None, + "error": None, + "status": "success", + "data": {"action_input": "", "action_name": "", "observation": "", "thought": ""}, + "label": "ROUND 1", + "metadata": {"started_at": 6055.211092814, "finished_at": 6056.990671049, "total_tokens": 325}, + "node_id": "n1", + }, + { + "id": "thought-1", + "parent_id": "round-1", + "error": None, + "status": "success", + "data": {"action": " pong!", "thought": ""}, + "label": "deepseek-v4-flash Thought", + "metadata": { + "started_at": 6055.211450345, + "finished_at": 6056.990473571, + "provider": "langgenius/deepseek/deepseek", + "total_tokens": 325, + }, + "node_id": "n1", + }, + { + "id": "tool-1", + "parent_id": "round-1", + "error": None, + "status": "success", + "data": { + "tool_name": "current_time", + "tool_call_args": {"timezone": "Asia/Shanghai"}, + "output": "2026-07-28 18:00:00", + }, + "label": "CALL current_time", + # Tool CALL logs also set provider (tool provider); must not become LLM spans. + "metadata": { + "started_at": 6056.990500000, + "finished_at": 6056.990600000, + "provider": "langgenius/time/time", + }, + "node_id": "n1", + }, + {"data": []}, + ], + } + + +def test_build_workflow_node_span_routes_agent_type(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch): + node_execution = MagicMock(spec=WorkflowNodeExecution) + trace_info = _make_workflow_trace_info() + trace_metadata = _make_trace_metadata() + + monkeypatch.setattr(trace_instance, "build_workflow_agent_span", MagicMock(return_value="agent")) + + node_execution.node_type = BuiltinNodeTypes.AGENT + assert trace_instance.build_workflow_node_span(node_execution, trace_info, trace_metadata) == "agent" + + +def test_build_workflow_agent_span(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(aliyun_trace_module, "convert_to_span_id", lambda _, __: 9) + monkeypatch.setattr(aliyun_trace_module, "convert_datetime_to_nanoseconds", lambda _: 123) + status = Status(StatusCode.OK) + monkeypatch.setattr(aliyun_trace_module, "get_workflow_node_status", lambda _: status) + + trace_metadata = _make_trace_metadata() + node_execution = MagicMock(spec=WorkflowNodeExecution) + node_execution.id = "node-id" + node_execution.title = "my-agent" + node_execution.inputs = {"query": "ping"} + node_execution.outputs = _make_agent_outputs() + node_execution.created_at = _dt() + node_execution.finished_at = _dt() + + span = trace_instance.build_workflow_agent_span(_make_workflow_trace_info(), node_execution, trace_metadata) + assert span.attributes["gen_ai.span.kind"] == GenAISpanKind.AGENT + assert span.attributes[GEN_AI_OPERATION_NAME] == OPERATION_NAME_INVOKE_AGENT + assert span.attributes[GEN_AI_AGENT_NAME] == "my-agent" + assert span.attributes[GEN_AI_USAGE_TOTAL_TOKENS] == "325" + assert span.attributes[GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN] == 500_000_000 + assert span.attributes["output.value"] == " pong!" + + # TTFT attribute must be absent when usage does not carry it (e.g. blocking mode) + node_execution.outputs = {"text": "t", "usage": {"total_tokens": 1, "time_to_first_token": None}} + span2 = trace_instance.build_workflow_agent_span(_make_workflow_trace_info(), node_execution, trace_metadata) + assert GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN not in span2.attributes + + # Malformed usage payloads must not break agent span building + node_execution.outputs = {"text": "t", "usage": "not-a-mapping"} + span3 = trace_instance.build_workflow_agent_span(_make_workflow_trace_info(), node_execution, trace_metadata) + assert span3.attributes[GEN_AI_USAGE_TOTAL_TOKENS] == "0" + assert GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN not in span3.attributes + + +def test_build_agent_react_spans(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch): + node_start_ns = 1_000_000_000 + monkeypatch.setattr(aliyun_trace_module, "convert_to_span_id", lambda _, __: 9) + monkeypatch.setattr(aliyun_trace_module, "convert_datetime_to_nanoseconds", lambda _: node_start_ns) + span_ids = iter([100, 200, 300]) + monkeypatch.setattr(aliyun_trace_module, "generate_span_id", lambda: next(span_ids)) + + trace_metadata = _make_trace_metadata() + node_execution = MagicMock(spec=WorkflowNodeExecution) + node_execution.id = "node-id" + node_execution.outputs = _make_agent_outputs() + node_execution.created_at = _dt() + node_execution.finished_at = _dt() + + spans = trace_instance.build_agent_react_spans(node_execution, trace_metadata) + assert len(spans) == 3 + step_span, llm_span, tool_span = spans + + assert step_span.parent_span_id == 9 + assert step_span.span_id == 100 + assert step_span.name == "ROUND 1" + assert step_span.attributes["gen_ai.span.kind"] == GenAISpanKind.STEP + assert step_span.attributes[GEN_AI_OPERATION_NAME] == OPERATION_NAME_REACT + assert step_span.attributes[GEN_AI_REACT_ROUND] == 1 + # Round starts at the earliest monotonic timestamp, anchored to node start time + assert step_span.start_time == node_start_ns + assert step_span.end_time == node_start_ns + int((6056.990671049 - 6055.211092814) * 1e9) + assert step_span.status.status_code == StatusCode.OK + + assert llm_span.parent_span_id == 100 + assert llm_span.span_id == 200 + assert llm_span.name == "deepseek-v4-flash Thought" + assert llm_span.attributes["gen_ai.span.kind"] == GenAISpanKind.LLM + assert llm_span.attributes[GEN_AI_OPERATION_NAME] == OPERATION_NAME_CHAT + assert llm_span.attributes[GEN_AI_REQUEST_MODEL] == "deepseek-v4-flash" + assert llm_span.attributes[GEN_AI_PROVIDER_NAME] == "langgenius/deepseek/deepseek" + assert llm_span.attributes[GEN_AI_USAGE_TOTAL_TOKENS] == "325" + assert llm_span.attributes[GEN_AI_COMPLETION] == " pong!" + assert llm_span.start_time == node_start_ns + int((6055.211450345 - 6055.211092814) * 1e9) + + assert tool_span.parent_span_id == 100 + assert tool_span.span_id == 300 + assert tool_span.name == "CALL current_time" + assert tool_span.attributes["gen_ai.span.kind"] == GenAISpanKind.TOOL + assert tool_span.attributes[GEN_AI_OPERATION_NAME] == OPERATION_NAME_EXECUTE_TOOL + assert tool_span.attributes[GEN_AI_TOOL_NAME] == "current_time" + assert tool_span.attributes[GEN_AI_TOOL_TYPE] == TOOL_TYPE_FUNCTION + assert tool_span.attributes[GEN_AI_TOOL_CALL_ID] == "tool-1" + assert tool_span.attributes[GEN_AI_TOOL_CALL_ARGUMENTS] == '{"timezone": "Asia/Shanghai"}' + assert tool_span.attributes[GEN_AI_TOOL_CALL_RESULT] == "2026-07-28 18:00:00" + assert tool_span.start_time == node_start_ns + int((6056.990500000 - 6055.211092814) * 1e9) + + +def test_build_agent_react_spans_returns_empty_without_log(trace_instance: AliyunDataTrace): + node_execution = MagicMock(spec=WorkflowNodeExecution) + node_execution.id = "node-id" + node_execution.outputs = {"text": "t"} + node_execution.created_at = _dt() + node_execution.finished_at = _dt() + + assert trace_instance.build_agent_react_spans(node_execution, _make_trace_metadata()) == [] + + +def test_workflow_trace_adds_react_spans_for_agent_nodes( + trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(aliyun_trace_module, "convert_to_trace_id", lambda _: 111) + monkeypatch.setattr(aliyun_trace_module, "convert_to_span_id", lambda _, __: 222) + monkeypatch.setattr(aliyun_trace_module, "create_links_from_trace_id", lambda _: []) + + agent_node = MagicMock(spec=WorkflowNodeExecution) + agent_node.node_type = BuiltinNodeTypes.AGENT + code_node = MagicMock(spec=WorkflowNodeExecution) + code_node.node_type = BuiltinNodeTypes.CODE + + monkeypatch.setattr(trace_instance, "add_workflow_span", MagicMock()) + monkeypatch.setattr(trace_instance, "get_workflow_node_executions", MagicMock(return_value=[agent_node, code_node])) + monkeypatch.setattr(trace_instance, "build_workflow_node_span", MagicMock(side_effect=["agent-span", "task-span"])) + build_agent_react_spans = MagicMock(return_value=["step-span", "llm-span"]) + monkeypatch.setattr(trace_instance, "build_agent_react_spans", build_agent_react_spans) + + trace_instance.workflow_trace(_make_workflow_trace_info()) + + build_agent_react_spans.assert_called_once() + assert _recording_trace_client(trace_instance).added_spans == [ + "agent-span", + "step-span", + "llm-span", + "task-span", + ] + + +def test_build_workflow_llm_span_records_time_to_first_token( + trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(aliyun_trace_module, "convert_to_span_id", lambda _, __: 9) + monkeypatch.setattr(aliyun_trace_module, "convert_datetime_to_nanoseconds", lambda _: 123) + monkeypatch.setattr(aliyun_trace_module, "get_workflow_node_status", lambda _: Status(StatusCode.OK)) + + trace_metadata = _make_trace_metadata() + node_execution = MagicMock(spec=WorkflowNodeExecution) + node_execution.id = "node-id" + node_execution.title = "llm" + node_execution.inputs = {} + node_execution.error = None + node_execution.process_data = {"prompts": []} + node_execution.outputs = {"text": "t", "usage": {"total_tokens": 1, "time_to_first_token": 0.123}} + node_execution.created_at = _dt() + node_execution.finished_at = _dt() + + span = trace_instance.build_workflow_llm_span(_make_workflow_trace_info(), node_execution, trace_metadata) + assert span.attributes[GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN] == 123_000_000 + assert span.attributes[GEN_AI_OPERATION_NAME] == OPERATION_NAME_CHAT + + # Absent when usage does not carry TTFT (blocking mode) + node_execution.outputs = {"text": "t", "usage": {"total_tokens": 1, "time_to_first_token": None}} + span2 = trace_instance.build_workflow_llm_span(_make_workflow_trace_info(), node_execution, trace_metadata) + assert GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN not in span2.attributes + + # Non-numeric TTFT must not raise and must be omitted, so the span is still built + node_execution.outputs = {"text": "t", "usage": {"total_tokens": 1, "time_to_first_token": "n/a"}} + span3 = trace_instance.build_workflow_llm_span(_make_workflow_trace_info(), node_execution, trace_metadata) + assert GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN not in span3.attributes + + +def test_message_trace_records_time_to_first_token(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(aliyun_trace_module, "convert_to_trace_id", lambda _: 10) + monkeypatch.setattr(aliyun_trace_module, "convert_to_span_id", lambda _, span_type: 0) + monkeypatch.setattr(aliyun_trace_module, "convert_datetime_to_nanoseconds", lambda _: 123) + monkeypatch.setattr(aliyun_trace_module, "get_user_id_from_message_data", lambda _: "user") + monkeypatch.setattr(aliyun_trace_module, "create_links_from_trace_id", lambda _: []) + + trace_instance.message_trace(_make_message_trace_info(gen_ai_server_time_to_first_token=0.25)) + + llm_span = _recorded_span_data(trace_instance)[1] + assert llm_span.attributes[GEN_AI_RESPONSE_TIME_TO_FIRST_TOKEN] == 250_000_000 + + def test_add_workflow_span(trace_instance: AliyunDataTrace, monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr( aliyun_trace_module, "convert_to_span_id", lambda _, span_type: {"message": 20}.get(span_type, 0) diff --git a/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/test_aliyun_trace_utils.py b/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/test_aliyun_trace_utils.py index 0900dfda97f..1895f5a08aa 100644 --- a/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/test_aliyun_trace_utils.py +++ b/api/providers/trace/trace-aliyun/tests/unit_tests/aliyun_trace/test_aliyun_trace_utils.py @@ -1,3 +1,5 @@ +"""Unit tests for Aliyun trace utility transformations and database lookups.""" + import json from collections.abc import Mapping from typing import Any, cast @@ -6,30 +8,43 @@ from unittest.mock import MagicMock import pytest from dify_trace_aliyun.entities.semconv import ( GEN_AI_FRAMEWORK, + GEN_AI_OPERATION_NAME, GEN_AI_SESSION_ID, GEN_AI_SPAN_KIND, + GEN_AI_TOOL_CALL_ARGUMENTS, + GEN_AI_TOOL_NAME, + GEN_AI_TOOL_TYPE, GEN_AI_USER_ID, INPUT_VALUE, + OPERATION_NAME_EXECUTE_TOOL, OUTPUT_VALUE, + TOOL_TYPE_DATASTORE, + TOOL_TYPE_EXTENSION, + TOOL_TYPE_FUNCTION, ) from dify_trace_aliyun.utils import ( create_common_span_attributes, + create_gen_ai_tool_attributes, create_links_from_trace_id, create_status_from_error, extract_retrieval_documents, + extract_tool_description, format_input_messages, format_output_messages, format_retrieval_documents, get_user_id_from_message_data, get_workflow_node_status, + map_gen_ai_tool_type, serialize_json_data, ) from opentelemetry.trace import Link, StatusCode +from sqlalchemy.orm import Session from core.rag.models.document import Document from graphon.entities import WorkflowNodeExecution from graphon.enums import WorkflowNodeExecutionStatus from models import EndUser +from models.enums import EndUserType def test_get_user_id_from_message_data_no_end_user(monkeypatch: pytest.MonkeyPatch): @@ -40,35 +55,40 @@ def test_get_user_id_from_message_data_no_end_user(monkeypatch: pytest.MonkeyPat assert get_user_id_from_message_data(message_data) == "account_id" -def test_get_user_id_from_message_data_with_end_user(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [(EndUser,)], indirect=True) +def test_get_user_id_from_message_data_with_end_user(monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None: message_data = MagicMock() message_data.from_account_id = "account_id" message_data.from_end_user_id = "end_user_id" - end_user_data = MagicMock(spec=EndUser) - end_user_data.session_id = "session_id" - - mock_session = MagicMock() - mock_session.get.return_value = end_user_data + end_user_data = EndUser( + id="end_user_id", + tenant_id="tenant_id", + app_id="app_id", + type=EndUserType.BROWSER, + session_id="session_id", + ) + sqlite3_session.add(end_user_data) + sqlite3_session.commit() from dify_trace_aliyun.utils import db - monkeypatch.setattr(db, "session", mock_session) + monkeypatch.setattr(db, "session", sqlite3_session) assert get_user_id_from_message_data(message_data) == "session_id" -def test_get_user_id_from_message_data_end_user_not_found(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [(EndUser,)], indirect=True) +def test_get_user_id_from_message_data_end_user_not_found( + monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: message_data = MagicMock() message_data.from_account_id = "account_id" message_data.from_end_user_id = "end_user_id" - mock_session = MagicMock() - mock_session.get.return_value = None - from dify_trace_aliyun.utils import db - monkeypatch.setattr(db, "session", mock_session) + monkeypatch.setattr(db, "session", sqlite3_session) assert get_user_id_from_message_data(message_data) == "account_id" @@ -270,3 +290,29 @@ def test_format_output_messages(): # Exception path # Trigger exception in serialize_json_data by passing non-serializable assert format_output_messages({"text": MagicMock()}) == serialize_json_data([]) + + +def test_map_gen_ai_tool_type(): + assert map_gen_ai_tool_type("dataset-retrieval") == TOOL_TYPE_DATASTORE + assert map_gen_ai_tool_type("extension") == TOOL_TYPE_EXTENSION + assert map_gen_ai_tool_type("builtin") == TOOL_TYPE_FUNCTION + assert map_gen_ai_tool_type(None) == TOOL_TYPE_FUNCTION + + +def test_extract_tool_description(): + assert extract_tool_description({"description": "d"}) == "d" + assert extract_tool_description({"tool_description": "td"}) == "td" + assert extract_tool_description({"other": 1}) == "" + assert extract_tool_description(None) == "" + + +def test_create_gen_ai_tool_attributes(): + attrs = create_gen_ai_tool_attributes( + tool_name="search", + tool_type=TOOL_TYPE_DATASTORE, + tool_call_arguments='{"q": 1}', + ) + assert attrs[GEN_AI_OPERATION_NAME] == OPERATION_NAME_EXECUTE_TOOL + assert attrs[GEN_AI_TOOL_NAME] == "search" + assert attrs[GEN_AI_TOOL_TYPE] == TOOL_TYPE_DATASTORE + assert attrs[GEN_AI_TOOL_CALL_ARGUMENTS] == '{"q": 1}' diff --git a/api/providers/trace/trace-arize-phoenix/tests/unit_tests/arize_phoenix_trace/test_arize_phoenix_trace.py b/api/providers/trace/trace-arize-phoenix/tests/unit_tests/arize_phoenix_trace/test_arize_phoenix_trace.py index 9e3cac2255b..f3f67e8ddbe 100644 --- a/api/providers/trace/trace-arize-phoenix/tests/unit_tests/arize_phoenix_trace/test_arize_phoenix_trace.py +++ b/api/providers/trace/trace-arize-phoenix/tests/unit_tests/arize_phoenix_trace/test_arize_phoenix_trace.py @@ -1,3 +1,5 @@ +"""Unit tests for Arize/Phoenix tracing with a real SQLite session factory.""" + import json from collections.abc import Sequence from datetime import UTC, datetime, timedelta @@ -36,9 +38,12 @@ from opentelemetry.context import Context from opentelemetry.sdk import trace as trace_sdk from opentelemetry.sdk.trace import ReadableSpan, Tracer from opentelemetry.sdk.trace.export import SimpleSpanProcessor, SpanExporter, SpanExportResult + +# pyrefly: ignore [deprecated] from opentelemetry.semconv.trace import SpanAttributes as OTELSpanAttributes from opentelemetry.trace import NonRecordingSpan, SpanContext, StatusCode, TraceFlags, TraceState, use_span from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator +from sqlalchemy.orm import Session from core.ops.entities.trace_entity import ( DatasetRetrievalTraceInfo, @@ -53,10 +58,24 @@ from core.ops.entities.trace_entity import ( ) from core.ops.exceptions import PendingTraceParentContextError from graphon.enums import BUILT_IN_NODE_TYPES, BuiltinNodeTypes, WorkflowNodeExecutionStatus +from models import WorkflowNodeExecutionModel + +pytestmark = pytest.mark.parametrize("sqlite3_session", [(WorkflowNodeExecutionModel,)], indirect=True) # --- Helpers --- +@pytest.fixture(autouse=True) +def _sqlite_trace_database(monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None: + """Route trace-created session factories to an isolated SQLite engine.""" + + monkeypatch.setattr( + arize_phoenix_trace_module, + "db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) + + def _dt(): return datetime(2024, 1, 1, 0, 0, 0, tzinfo=UTC) @@ -316,7 +335,7 @@ def test_normalize_wrapper_index_rejects_unstable_values(value): assert _normalize_wrapper_index(value) is None -def test_parent_workflow_can_publish_span_context_keeps_unknown_parent_retryable(monkeypatch): +def test_parent_workflow_can_publish_span_context_keeps_unknown_parent_retryable(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr( "dify_trace_arize_phoenix.arize_phoenix_trace.db.session.query", lambda model: _FakeQuery(None), @@ -325,7 +344,7 @@ def test_parent_workflow_can_publish_span_context_keeps_unknown_parent_retryable assert _parent_workflow_can_publish_span_context("missing-run") is True -def test_parent_workflow_can_publish_span_context_checks_parent_app_tracing(monkeypatch): +def test_parent_workflow_can_publish_span_context_checks_parent_app_tracing(monkeypatch: pytest.MonkeyPatch): parent_run = SimpleNamespace(app_id="parent-app") parent_app = SimpleNamespace(tracing=json.dumps({"enabled": True, "tracing_provider": "phoenix"})) @@ -771,11 +790,8 @@ def test_trace_exception(trace_instance): trace_instance.trace(_make_workflow_info()) -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") -def test_workflow_trace_full(mock_db, mock_repo_factory, mock_sessionmaker, trace_instance): - mock_db.engine = MagicMock() +def test_workflow_trace_full(mock_repo_factory, trace_instance): info = _make_workflow_info() repo = MagicMock() mock_repo_factory.create_workflow_node_execution_repository.return_value = repo @@ -805,22 +821,39 @@ def test_workflow_trace_full(mock_db, mock_repo_factory, mock_sessionmaker, trac assert trace_instance.tracer.start_span.call_count >= 2 -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") -def test_workflow_trace_no_app_id(mock_db, trace_instance): - mock_db.engine = MagicMock() +def test_workflow_trace_no_app_id(trace_instance): info = _make_workflow_info() info.metadata = {} with pytest.raises(ValueError, match="No app_id found in trace_info metadata"): trace_instance.workflow_trace(info) -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") +def test_workflow_trace_queries_real_repository_with_sqlite_session_factory( + trace_instance, sqlite3_session: Session +) -> None: + """Exercise the repository query through the session factory created by the trace.""" + + info = _make_workflow_info( + tenant_id="00000000-0000-0000-0000-000000000001", + workflow_id="00000000-0000-0000-0000-000000000002", + workflow_run_id="00000000-0000-0000-0000-000000000003", + metadata={"app_id": "00000000-0000-0000-0000-000000000004"}, + ) + service_account = MagicMock() + service_account.id = "00000000-0000-0000-0000-000000000005" + + with patch.object(trace_instance, "get_service_account_with_tenant", return_value=service_account): + trace_instance.workflow_trace(info) + + workflow_span_call = _get_start_span_call( + trace_instance.tracer.start_span, + span_name="workflow_00000000-0000-0000-0000-000000000003", + ) + assert workflow_span_call.kwargs["attributes"][SpanAttributes.SESSION_ID] == info.workflow_run_id + + @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") -def test_workflow_trace_uses_canonical_root_context_for_top_level_workflow( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance -): - mock_db.engine = MagicMock() +def test_workflow_trace_uses_canonical_root_context_for_top_level_workflow(mock_repo_factory, trace_instance): info = _make_workflow_info( message_id="message-1", workflow_run_id="workflow-run-1", @@ -856,16 +889,11 @@ def test_workflow_trace_uses_canonical_root_context_for_top_level_workflow( assert workflow_span_call.kwargs["context"] is root_context -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_uses_workflow_run_id_for_root_span_and_populates_root_inputs_outputs( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, ): - mock_db.engine = MagicMock() info = _make_workflow_info( workflow_run_inputs={"prompt": "hello"}, workflow_run_outputs={"result": "world"}, @@ -891,16 +919,11 @@ def test_workflow_trace_uses_workflow_run_id_for_root_span_and_populates_root_in assert root_span_call.kwargs["attributes"][SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_propagates_workflow_error_to_root_span( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, ): - mock_db.engine = MagicMock() info = _make_workflow_info( workflow_run_status="failed", error="Traceback (most recent call last): RuntimeError: workflow failed", @@ -919,16 +942,11 @@ def test_workflow_trace_propagates_workflow_error_to_root_span( assert mock_ensure_root_span.call_args.kwargs["root_span_error"] == info.error -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_falls_back_to_dify_name_when_workflow_run_id_is_blank( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, ): - mock_db.engine = MagicMock() info = _make_workflow_info( metadata={ "app_id": "app1", @@ -947,13 +965,10 @@ def test_workflow_trace_falls_back_to_dify_name_when_workflow_run_id_is_blank( assert root_span_call.kwargs["attributes"]["dify_trace_id"] == "" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_reuses_upstream_parent_workflow_context_when_no_parent_node_execution_id_is_available( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance + mock_repo_factory, trace_instance ): - mock_db.engine = MagicMock() info = _make_workflow_info( message_id="message-1", workflow_run_id="workflow-run-1", @@ -994,16 +1009,11 @@ def test_workflow_trace_reuses_upstream_parent_workflow_context_when_no_parent_n assert workflow_span_call.kwargs["context"] is parent_context -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_uses_published_parent_node_context_for_nested_workflow( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, ): - mock_db.engine = MagicMock() info = _make_workflow_info( message_id="message-1", workflow_run_id="workflow-run-1", @@ -1040,16 +1050,11 @@ def test_workflow_trace_uses_published_parent_node_context_for_nested_workflow( assert workflow_span_call.kwargs["context"] is parent_context -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_raises_pending_parent_error_when_parent_node_context_is_missing( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, ): - mock_db.engine = MagicMock() info = _make_workflow_info( message_id="message-1", workflow_run_id="workflow-run-1", @@ -1084,16 +1089,11 @@ def test_workflow_trace_raises_pending_parent_error_when_parent_node_context_is_ mock_ensure_root_span.assert_not_called() -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_falls_back_when_parent_app_tracing_cannot_publish_parent_context( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, ): - mock_db.engine = MagicMock() info = _make_workflow_info( message_id="message-1", workflow_run_id="workflow-run-1", @@ -1140,16 +1140,11 @@ def test_workflow_trace_falls_back_when_parent_app_tracing_cannot_publish_parent assert workflow_span_call.kwargs["context"] is parent_context -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_still_retries_when_parent_app_can_publish_parent_context( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, ): - mock_db.engine = MagicMock() info = _make_workflow_info( message_id="message-1", workflow_run_id="workflow-run-1", @@ -1181,13 +1176,10 @@ def test_workflow_trace_still_retries_when_parent_app_can_publish_parent_context mock_ensure_root_span.assert_not_called() -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_uses_parent_workflow_run_id_for_workflow_and_nodes_when_nested_context_is_present( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance + mock_repo_factory, trace_instance ): - mock_db.engine = MagicMock() info = _make_workflow_info( conversation_id=None, metadata={ @@ -1223,13 +1215,8 @@ def test_workflow_trace_uses_parent_workflow_run_id_for_workflow_and_nodes_when_ assert node_span_call.kwargs["attributes"][SpanAttributes.SESSION_ID] == "outer-workflow-run-1" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") -def test_workflow_trace_falls_back_to_node_type_when_node_title_is_blank( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance -): - mock_db.engine = MagicMock() +def test_workflow_trace_falls_back_to_node_type_when_node_title_is_blank(mock_repo_factory, trace_instance): info = _make_workflow_info() repo = MagicMock() node_execution = _make_node_execution( @@ -1249,13 +1236,8 @@ def test_workflow_trace_falls_back_to_node_type_when_node_title_is_blank( assert node_span_call.kwargs["attributes"][SpanAttributes.SESSION_ID] == "r1" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") -def test_workflow_trace_prefers_workflow_graph_node_title_over_execution_title( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance -): - mock_db.engine = MagicMock() +def test_workflow_trace_prefers_workflow_graph_node_title_over_execution_title(mock_repo_factory, trace_instance): info = _make_workflow_info( workflow_data={ "graph": { @@ -1289,13 +1271,10 @@ def test_workflow_trace_prefers_workflow_graph_node_title_over_execution_title( assert node_span_call.kwargs["attributes"][SpanAttributes.SESSION_ID] == "r1" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_keeps_nested_conversation_session_while_reusing_parent_root_context( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance + mock_repo_factory, trace_instance ): - mock_db.engine = MagicMock() info = _make_workflow_info( conversation_id="conversation-1", message_id="message-1", @@ -1346,16 +1325,11 @@ def test_workflow_trace_keeps_nested_conversation_session_while_reusing_parent_r assert node_span_call.kwargs["attributes"][SpanAttributes.SESSION_ID] == "conversation-1" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_publishes_tool_node_parent_span_context_to_redis( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, ): - mock_db.engine = MagicMock() info = _make_workflow_info() repo = MagicMock() node_execution = _make_node_execution( @@ -1403,18 +1377,13 @@ def test_workflow_trace_publishes_tool_node_parent_span_context_to_redis( ("publish", "publish failed"), ], ) -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_cleans_up_tool_span_when_parent_context_publish_fails( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, failing_step, expected_message, ): - mock_db.engine = MagicMock() info = _make_workflow_info() repo = MagicMock() node_execution = _make_node_execution( @@ -1458,13 +1427,8 @@ def test_workflow_trace_cleans_up_tool_span_when_parent_context_publish_fails( workflow_span.end.assert_called_once() -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") -def test_workflow_trace_parents_serial_nodes_to_resolved_predecessor_span( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance -): - mock_db.engine = MagicMock() +def test_workflow_trace_parents_serial_nodes_to_resolved_predecessor_span(mock_repo_factory, trace_instance): info = _make_workflow_info() repo = MagicMock() second_node = _make_node_execution( @@ -1520,18 +1484,13 @@ def test_workflow_trace_parents_serial_nodes_to_resolved_predecessor_span( ("iteration", "iteration_id"), ], ) -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_parents_structured_start_nodes_to_enclosing_structure_span( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, enclosing_node_type, structured_field, ): - mock_db.engine = MagicMock() info = _make_workflow_info() repo = MagicMock() enclosing_node = _make_node_execution( @@ -1581,18 +1540,13 @@ def test_workflow_trace_parents_structured_start_nodes_to_enclosing_structure_sp ("iteration", "iteration_id"), ], ) -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_keeps_duplicate_body_node_children_under_enclosing_structure( - mock_sessionmaker, mock_repo_factory, - mock_db, trace_instance, enclosing_node_type, structured_field, ): - mock_db.engine = MagicMock() info = _make_workflow_info() repo = MagicMock() enclosing_node = _make_node_execution( @@ -1670,13 +1624,8 @@ def test_workflow_trace_keeps_duplicate_body_node_children_under_enclosing_struc assert child_node_call.kwargs["context"] == f"context:{enclosing_node_type}" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") -def test_workflow_trace_records_exception_node_event_without_failing_root_span( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance -): - mock_db.engine = MagicMock() +def test_workflow_trace_records_exception_node_event_without_failing_root_span(mock_repo_factory, trace_instance): info = _make_workflow_info(workflow_run_status="succeeded", error=None) repo = MagicMock() handled_error_node = _make_node_execution( @@ -1716,13 +1665,8 @@ def test_workflow_trace_records_exception_node_event_without_failing_root_span( ) -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") -def test_workflow_trace_groups_loop_iteration_children_under_wrapper_spans( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance -): - mock_db.engine = MagicMock() +def test_workflow_trace_groups_loop_iteration_children_under_wrapper_spans(mock_repo_factory, trace_instance): info = _make_workflow_info(conversation_id="conversation-1") repo = MagicMock() loop_node = _make_node_execution( @@ -1799,13 +1743,10 @@ def test_workflow_trace_groups_loop_iteration_children_under_wrapper_spans( assert first_body_call.kwargs["attributes"]["dify.node.loop_index"] == 0 -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_finalizes_loop_wrapper_with_child_time_bounds_and_error_status( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance + mock_repo_factory, trace_instance ): - mock_db.engine = MagicMock() info = _make_workflow_info() repo = MagicMock() loop_node = _make_node_execution( @@ -1877,13 +1818,10 @@ def test_workflow_trace_finalizes_loop_wrapper_with_child_time_bounds_and_error_ assert wrapper_span.set_status.call_args.args[0].status_code == StatusCode.ERROR -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") @patch("dify_trace_arize_phoenix.arize_phoenix_trace.DifyCoreRepositoryFactory") -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.sessionmaker") def test_workflow_trace_falls_back_to_workflow_span_for_parallel_like_ambiguous_predecessors( - mock_sessionmaker, mock_repo_factory, mock_db, trace_instance + mock_repo_factory, trace_instance ): - mock_db.engine = MagicMock() info = _make_workflow_info() repo = MagicMock() child_node = _make_node_execution( @@ -1945,9 +1883,7 @@ def test_workflow_trace_falls_back_to_workflow_span_for_parallel_like_ambiguous_ assert child_node_call.kwargs["context"] == "context:workflow" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") -def test_message_trace_keeps_conversation_id_as_session(mock_db, trace_instance): - mock_db.engine = MagicMock() +def test_message_trace_keeps_conversation_id_as_session(trace_instance): info = _make_message_info() info.message_data = MagicMock() info.message_data.conversation_id = "conversation-2" @@ -1975,9 +1911,7 @@ def test_message_trace_keeps_conversation_id_as_session(mock_db, trace_instance) assert message_span_call.kwargs["attributes"][SpanAttributes.SESSION_ID] == "conversation-2" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") -def test_message_trace_uses_trace_session_id_metadata_as_session(mock_db, trace_instance): - mock_db.engine = MagicMock() +def test_message_trace_uses_trace_session_id_metadata_as_session(trace_instance): info = _make_message_info(metadata={"app_id": "app-1", "trace_session_id": "session-1"}) info.message_data = MagicMock() info.message_data.conversation_id = "conversation-2" @@ -2009,9 +1943,7 @@ def test_message_trace_uses_trace_session_id_metadata_as_session(mock_db, trace_ assert llm_span_call.kwargs["attributes"][SpanAttributes.SESSION_ID] == "session-1" -@patch("dify_trace_arize_phoenix.arize_phoenix_trace.db") -def test_message_trace_with_error(mock_db, trace_instance): - mock_db.engine = MagicMock() +def test_message_trace_with_error(trace_instance): info = _make_message_info() info.message_data = MagicMock() info.message_data.from_account_id = "acc1" diff --git a/api/providers/trace/trace-langfuse/tests/unit_tests/langfuse_trace/test_langfuse_trace.py b/api/providers/trace/trace-langfuse/tests/unit_tests/langfuse_trace/test_langfuse_trace.py index 8da2dc9fa0d..93f48073363 100644 --- a/api/providers/trace/trace-langfuse/tests/unit_tests/langfuse_trace/test_langfuse_trace.py +++ b/api/providers/trace/trace-langfuse/tests/unit_tests/langfuse_trace/test_langfuse_trace.py @@ -1,3 +1,5 @@ +"""Unit tests for Langfuse trace translation with real SQLite-backed lookups.""" + import collections import logging from datetime import UTC, datetime, timedelta @@ -15,6 +17,7 @@ from dify_trace_langfuse.entities.langfuse_trace_entity import ( UnitEnum, ) from dify_trace_langfuse.langfuse_trace import LangFuseDataTrace +from sqlalchemy.orm import Session from core.ops.entities.trace_entity import ( DatasetRetrievalTraceInfo, @@ -28,7 +31,7 @@ from core.ops.entities.trace_entity import ( ) from graphon.enums import BuiltinNodeTypes from models import EndUser -from models.enums import MessageStatus +from models.enums import EndUserType, MessageStatus def _dt() -> datetime: @@ -187,7 +190,10 @@ def test_trace_dispatch(trace_instance, monkeypatch: pytest.MonkeyPatch): mocks["generate_name_trace"].assert_called_once_with(info) -def test_workflow_trace_with_message_id(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_with_message_id( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: # Setup trace info trace_info = WorkflowTraceInfo( workflow_id="wf-1", @@ -211,10 +217,10 @@ def test_workflow_trace_with_message_id(trace_instance, monkeypatch: pytest.Monk error="", ) - # Mock DB and Repositories - mock_session = MagicMock() - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.sessionmaker", lambda bind: lambda: mock_session) - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_langfuse.langfuse_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) # Mock node executions node_llm = MagicMock() @@ -291,7 +297,10 @@ def test_workflow_trace_with_message_id(trace_instance, monkeypatch: pytest.Monk assert other_span.level == LevelEnum.ERROR -def test_workflow_trace_no_message_id(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_no_message_id( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: trace_info = WorkflowTraceInfo( workflow_id="wf-1", tenant_id="tenant-1", @@ -314,8 +323,10 @@ def test_workflow_trace_no_message_id(trace_instance, monkeypatch: pytest.Monkey error="", ) - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.sessionmaker", lambda bind: lambda: MagicMock()) - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_langfuse.langfuse_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) repo = MagicMock() repo.get_by_workflow_execution.return_value = [] mock_factory = MagicMock() @@ -332,7 +343,10 @@ def test_workflow_trace_no_message_id(trace_instance, monkeypatch: pytest.Monkey assert trace_data.name == TraceTaskName.WORKFLOW_TRACE -def test_workflow_trace_missing_app_id(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_missing_app_id( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: trace_info = WorkflowTraceInfo( workflow_id="wf-1", tenant_id="tenant-1", @@ -353,8 +367,10 @@ def test_workflow_trace_missing_app_id(trace_instance, monkeypatch: pytest.Monke workflow_app_log_id="log-1", error="", ) - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.sessionmaker", lambda bind: lambda: MagicMock()) - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_langfuse.langfuse_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) with pytest.raises(ValueError, match="No app_id found in trace_info metadata"): trace_instance.workflow_trace(trace_info) @@ -404,7 +420,8 @@ def test_message_trace_basic(trace_instance, monkeypatch: pytest.MonkeyPatch): assert gen_data.usage.total == 30 -def test_message_trace_with_end_user(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [(EndUser,)], indirect=True) +def test_message_trace_with_end_user(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None: message_data = MagicMock() message_data.id = "msg-1" message_data.from_account_id = "acc-1" @@ -434,20 +451,32 @@ def test_message_trace_with_end_user(trace_instance, monkeypatch: pytest.MonkeyP error=None, ) - # Mock DB session for EndUser lookup - mock_end_user = MagicMock(spec=EndUser) - mock_end_user.session_id = "session-id-123" + end_user = EndUser( + id="end-user-1", + tenant_id="tenant-1", + app_id="app-1", + type=EndUserType.BROWSER, + session_id="session-id-123", + ) + engine = sqlite3_session.get_bind() + with Session(engine) as write_session: + write_session.add(end_user) + write_session.commit() - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.db.session.get", lambda model, pk: mock_end_user) + with Session(engine) as read_session: + monkeypatch.setattr( + "dify_trace_langfuse.langfuse_trace.db", + SimpleNamespace(engine=engine, session=read_session), + ) - trace_instance.add_trace = MagicMock() - trace_instance.add_generation = MagicMock() + trace_instance.add_trace = MagicMock() + trace_instance.add_generation = MagicMock() - trace_instance.message_trace(trace_info) + trace_instance.message_trace(trace_info) - trace_data = trace_instance.add_trace.call_args[1]["langfuse_trace_data"] - assert trace_data.user_id == "session-id-123" - assert trace_data.metadata["user_id"] == "session-id-123" + trace_data = trace_instance.add_trace.call_args[1]["langfuse_trace_data"] + assert trace_data.user_id == "session-id-123" + assert trace_data.metadata["user_id"] == "session-id-123" def test_message_trace_none_data(trace_instance): @@ -709,8 +738,12 @@ def test_langfuse_trace_entity_with_list_dict_input(): assert data.input[0]["content"] == "hello" +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) def test_workflow_trace_handles_usage_extraction_error( - trace_instance, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture + trace_instance, + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + sqlite3_session: Session, ): # Setup trace info to trigger LLM node usage extraction trace_info = WorkflowTraceInfo( @@ -758,8 +791,10 @@ def test_workflow_trace_handles_usage_extraction_error( mock_factory = MagicMock() mock_factory.create_workflow_node_execution_repository.return_value = repo monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.DifyCoreRepositoryFactory", mock_factory) - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.sessionmaker", lambda bind: lambda: MagicMock()) - monkeypatch.setattr("dify_trace_langfuse.langfuse_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_langfuse.langfuse_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) monkeypatch.setattr(trace_instance, "get_service_account_with_tenant", lambda app_id: MagicMock()) trace_instance.add_trace = MagicMock() diff --git a/api/providers/trace/trace-langsmith/tests/unit_tests/langsmith_trace/test_langsmith_trace.py b/api/providers/trace/trace-langsmith/tests/unit_tests/langsmith_trace/test_langsmith_trace.py index 76d4c99caf7..f9406e13048 100644 --- a/api/providers/trace/trace-langsmith/tests/unit_tests/langsmith_trace/test_langsmith_trace.py +++ b/api/providers/trace/trace-langsmith/tests/unit_tests/langsmith_trace/test_langsmith_trace.py @@ -1,5 +1,8 @@ +"""Unit tests for LangSmith trace translation with SQLite-backed lookups.""" + import collections from datetime import datetime, timedelta +from types import SimpleNamespace from typing import override from unittest.mock import MagicMock @@ -11,6 +14,7 @@ from dify_trace_langsmith.entities.langsmith_trace_entity import ( LangSmithRunUpdateModel, ) from dify_trace_langsmith.langsmith_trace import LangSmithDataTrace +from sqlalchemy.orm import Session from core.ops.entities.trace_entity import ( DatasetRetrievalTraceInfo, @@ -24,6 +28,7 @@ from core.ops.entities.trace_entity import ( ) from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey from models import EndUser +from models.enums import EndUserType def _dt() -> datetime: @@ -108,7 +113,8 @@ def test_trace_dispatch(trace_instance, monkeypatch: pytest.MonkeyPatch): mocks["generate_name_trace"].assert_called_once_with(info) -def test_workflow_trace(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None: # Setup trace info workflow_data = MagicMock() workflow_data.created_at = _dt() @@ -137,10 +143,10 @@ def test_workflow_trace(trace_instance, monkeypatch: pytest.MonkeyPatch): workflow_data=workflow_data, ) - # Mock dependencies - mock_session = MagicMock() - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: mock_session) - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_langsmith.langsmith_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) # Mock node executions node_llm = MagicMock() @@ -228,7 +234,10 @@ def test_workflow_trace(trace_instance, monkeypatch: pytest.MonkeyPatch): assert call_args[4].run_type == LangSmithRunType.retriever -def test_workflow_trace_no_start_time(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_no_start_time( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: workflow_data = MagicMock() workflow_data.created_at = _dt() workflow_data.finished_at = _dt() + timedelta(seconds=1) @@ -256,9 +265,10 @@ def test_workflow_trace_no_start_time(trace_instance, monkeypatch: pytest.Monkey workflow_data=workflow_data, ) - mock_session = MagicMock() - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: mock_session) - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_langsmith.langsmith_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) repo = MagicMock() repo.get_by_workflow_execution.return_value = [] mock_factory = MagicMock() @@ -271,7 +281,10 @@ def test_workflow_trace_no_start_time(trace_instance, monkeypatch: pytest.Monkey assert trace_instance.add_run.called -def test_workflow_trace_missing_app_id(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_missing_app_id( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: trace_info = MagicMock(spec=WorkflowTraceInfo) trace_info.trace_id = "trace-1" trace_info.message_id = None @@ -287,15 +300,17 @@ def test_workflow_trace_missing_app_id(trace_instance, monkeypatch: pytest.Monke trace_info.workflow_run_outputs = {} trace_info.error = "" - mock_session = MagicMock() - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: mock_session) - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_langsmith.langsmith_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) with pytest.raises(ValueError, match="No app_id found in trace_info metadata"): trace_instance.workflow_trace(trace_info) -def test_message_trace(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [(EndUser,)], indirect=True) +def test_message_trace(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None: message_data = MagicMock() message_data.id = "msg-1" message_data.from_account_id = "acc-1" @@ -321,10 +336,19 @@ def test_message_trace(trace_instance, monkeypatch: pytest.MonkeyPatch): message_file_data=MagicMock(url="file-url"), ) - # Mock EndUser lookup - mock_end_user = MagicMock(spec=EndUser) - mock_end_user.session_id = "session-id-123" - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db.session.get", lambda model, pk: mock_end_user) + end_user = EndUser( + id="end-user-1", + tenant_id="tenant-1", + app_id="app-1", + type=EndUserType.BROWSER, + session_id="session-id-123", + ) + sqlite3_session.add(end_user) + sqlite3_session.commit() + monkeypatch.setattr( + "dify_trace_langsmith.langsmith_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) trace_instance.add_run = MagicMock() @@ -521,9 +545,13 @@ def test_update_run_error(trace_instance): trace_instance.update_run(update_data) +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) def test_workflow_trace_usage_extraction_error( - trace_instance, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture -): + trace_instance, + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + sqlite3_session: Session, +) -> None: workflow_data = MagicMock() workflow_data.created_at = _dt() workflow_data.finished_at = _dt() + timedelta(seconds=1) @@ -576,8 +604,10 @@ def test_workflow_trace_usage_extraction_error( mock_factory = MagicMock() mock_factory.create_workflow_node_execution_repository.return_value = repo monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.DifyCoreRepositoryFactory", mock_factory) - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: MagicMock()) - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_langsmith.langsmith_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) monkeypatch.setattr(trace_instance, "get_service_account_with_tenant", lambda app_id: MagicMock()) trace_instance.add_run = MagicMock() @@ -644,9 +674,11 @@ def _make_workflow_trace_info( ) -def _patch_workflow_trace_deps(monkeypatch, trace_instance): - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.sessionmaker", lambda bind: lambda: MagicMock()) - monkeypatch.setattr("dify_trace_langsmith.langsmith_trace.db", MagicMock(engine="engine")) +def _patch_workflow_trace_deps(monkeypatch, trace_instance, sqlite3_session: Session) -> None: + monkeypatch.setattr( + "dify_trace_langsmith.langsmith_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) repo = MagicMock() repo.get_by_workflow_execution.return_value = [] factory = MagicMock() @@ -656,14 +688,17 @@ def _patch_workflow_trace_deps(monkeypatch, trace_instance): trace_instance.add_run = MagicMock() -def test_workflow_trace_id_uses_message_id_not_external(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_id_uses_message_id_not_external( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: """Chatflow with external trace_id: LangSmith trace_id must be message_id, not external.""" trace_info = _make_workflow_trace_info( message_id="msg-abc", workflow_run_id="run-xyz", trace_id="external-999", ) - _patch_workflow_trace_deps(monkeypatch, trace_instance) + _patch_workflow_trace_deps(monkeypatch, trace_instance, sqlite3_session) trace_instance.workflow_trace(trace_info) @@ -677,14 +712,17 @@ def test_workflow_trace_id_uses_message_id_not_external(trace_instance, monkeypa assert trace_info.metadata.get("external_trace_id") == "external-999" -def test_workflow_trace_id_pure_workflow_uses_run_id(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_id_pure_workflow_uses_run_id( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: """Pure workflow (no message_id) with external trace_id: trace_id must be workflow_run_id.""" trace_info = _make_workflow_trace_info( message_id=None, workflow_run_id="run-xyz", trace_id="external-999", ) - _patch_workflow_trace_deps(monkeypatch, trace_instance) + _patch_workflow_trace_deps(monkeypatch, trace_instance, sqlite3_session) trace_instance.workflow_trace(trace_info) diff --git a/api/providers/trace/trace-opik/tests/unit_tests/opik_trace/test_opik_trace.py b/api/providers/trace/trace-opik/tests/unit_tests/opik_trace/test_opik_trace.py index ab20a783748..3761a011312 100644 --- a/api/providers/trace/trace-opik/tests/unit_tests/opik_trace/test_opik_trace.py +++ b/api/providers/trace/trace-opik/tests/unit_tests/opik_trace/test_opik_trace.py @@ -1,3 +1,5 @@ +"""Unit tests for Opik trace translation with SQLite-backed lookups.""" + import collections import logging from datetime import UTC, datetime, timedelta @@ -8,6 +10,7 @@ from unittest.mock import MagicMock import pytest from dify_trace_opik.config import OpikConfig from dify_trace_opik.opik_trace import OpikDataTrace, prepare_opik_uuid, wrap_dict, wrap_metadata +from sqlalchemy.orm import Session from core.ops.entities.trace_entity import ( DatasetRetrievalTraceInfo, @@ -21,7 +24,7 @@ from core.ops.entities.trace_entity import ( ) from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey from models import EndUser -from models.enums import MessageStatus +from models.enums import EndUserType, MessageStatus def _dt() -> datetime: @@ -133,7 +136,10 @@ def test_trace_dispatch(trace_instance, monkeypatch: pytest.MonkeyPatch): mocks["generate_name_trace"].assert_called_once_with(info) -def test_workflow_trace_with_message_id(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_with_message_id( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: # Define constants for better readability WORKFLOW_ID = "fb05c7cd-6cec-4add-8a84-df03a408b4ce" WORKFLOW_RUN_ID = "33c67568-7a8a-450e-8916-a5f135baeaef" @@ -166,9 +172,10 @@ def test_workflow_trace_with_message_id(trace_instance, monkeypatch: pytest.Monk error="", ) - mock_session = MagicMock() - monkeypatch.setattr("dify_trace_opik.opik_trace.sessionmaker", lambda bind: lambda: mock_session) - monkeypatch.setattr("dify_trace_opik.opik_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_opik.opik_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) node_llm = MagicMock() node_llm.id = LLM_NODE_ID @@ -222,7 +229,10 @@ def test_workflow_trace_with_message_id(trace_instance, monkeypatch: pytest.Monk assert trace_instance.add_span.call_count >= 1 -def test_workflow_trace_no_message_id(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_no_message_id( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: # Define constants for better readability WORKFLOW_ID = "f0708b36-b1d7-42b3-a876-1d01b7d8f1a3" WORKFLOW_RUN_ID = "d42ec285-c2fd-4248-8866-5c9386b101ac" @@ -251,8 +261,10 @@ def test_workflow_trace_no_message_id(trace_instance, monkeypatch: pytest.Monkey error="", ) - monkeypatch.setattr("dify_trace_opik.opik_trace.sessionmaker", lambda bind: lambda: MagicMock()) - monkeypatch.setattr("dify_trace_opik.opik_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_opik.opik_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) repo = MagicMock() repo.get_by_workflow_execution.return_value = [] mock_factory = MagicMock() @@ -266,7 +278,10 @@ def test_workflow_trace_no_message_id(trace_instance, monkeypatch: pytest.Monkey trace_instance.add_trace.assert_called_once() -def test_workflow_trace_missing_app_id(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) +def test_workflow_trace_missing_app_id( + trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session +) -> None: trace_info = WorkflowTraceInfo( workflow_id="5745f1b8-f8e6-4859-8110-996acb6c8d6a", tenant_id="tenant-1", @@ -287,8 +302,10 @@ def test_workflow_trace_missing_app_id(trace_instance, monkeypatch: pytest.Monke workflow_app_log_id="339760b2-4b94-4532-8c81-133a97e4680e", error="", ) - monkeypatch.setattr("dify_trace_opik.opik_trace.sessionmaker", lambda bind: lambda: MagicMock()) - monkeypatch.setattr("dify_trace_opik.opik_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_opik.opik_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) with pytest.raises(ValueError, match="No app_id found in trace_info metadata"): trace_instance.workflow_trace(trace_info) @@ -341,7 +358,9 @@ def test_message_trace_basic(trace_instance, monkeypatch: pytest.MonkeyPatch): trace_instance.add_span.assert_called_once() -def test_message_trace_with_end_user(trace_instance, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite3_session", [(EndUser,)], indirect=True) +def test_message_trace_with_end_user(trace_instance, monkeypatch: pytest.MonkeyPatch, sqlite3_session: Session) -> None: + write_session = sqlite3_session message_data = MagicMock() message_data.id = "85411059-79fb-4deb-a76c-c2e215f1b97e" message_data.from_account_id = "acc-1" @@ -371,15 +390,26 @@ def test_message_trace_with_end_user(trace_instance, monkeypatch: pytest.MonkeyP error=None, ) - mock_end_user = MagicMock(spec=EndUser) - mock_end_user.session_id = "session-id-123" + end_user = EndUser( + id="end-user-1", + tenant_id="tenant-1", + app_id="app-1", + type=EndUserType.BROWSER, + session_id="session-id-123", + ) + write_session.add(end_user) + write_session.commit() - monkeypatch.setattr("dify_trace_opik.opik_trace.db.session.get", lambda model, pk: mock_end_user) + with Session(write_session.get_bind()) as read_session: + monkeypatch.setattr( + "dify_trace_opik.opik_trace.db", + SimpleNamespace(engine=read_session.get_bind(), session=read_session), + ) - trace_instance.add_trace = MagicMock(return_value=MagicMock(id="trace_id_2")) - trace_instance.add_span = MagicMock() + trace_instance.add_trace = MagicMock(return_value=MagicMock(id="trace_id_2")) + trace_instance.add_span = MagicMock() - trace_instance.message_trace(trace_info) + trace_instance.message_trace(trace_info) trace_data = trace_instance.add_trace.call_args[0][0] assert trace_data["metadata"]["user_id"] == "acc-1" @@ -615,9 +645,13 @@ def test_get_project_url_error(trace_instance): trace_instance.get_project_url() +@pytest.mark.parametrize("sqlite3_session", [()], indirect=True) def test_workflow_trace_usage_extraction_error_fixed( - trace_instance, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture -): + trace_instance, + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + sqlite3_session: Session, +) -> None: trace_info = WorkflowTraceInfo( workflow_id="86a52565-4a6b-4a1b-9bfd-98e4595e70de", tenant_id="66e8e918-472e-4b69-8051-12502c34fc07", @@ -663,8 +697,10 @@ def test_workflow_trace_usage_extraction_error_fixed( mock_factory = MagicMock() mock_factory.create_workflow_node_execution_repository.return_value = repo monkeypatch.setattr("dify_trace_opik.opik_trace.DifyCoreRepositoryFactory", mock_factory) - monkeypatch.setattr("dify_trace_opik.opik_trace.sessionmaker", lambda bind: lambda: MagicMock()) - monkeypatch.setattr("dify_trace_opik.opik_trace.db", MagicMock(engine="engine")) + monkeypatch.setattr( + "dify_trace_opik.opik_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind(), session=sqlite3_session), + ) monkeypatch.setattr(trace_instance, "get_service_account_with_tenant", lambda app_id: MagicMock()) trace_instance.add_trace = MagicMock() diff --git a/api/providers/trace/trace-tencent/src/dify_trace_tencent/client.py b/api/providers/trace/trace-tencent/src/dify_trace_tencent/client.py index be06ab4a36a..c616de5724b 100644 --- a/api/providers/trace/trace-tencent/src/dify_trace_tencent/client.py +++ b/api/providers/trace/trace-tencent/src/dify_trace_tencent/client.py @@ -81,7 +81,7 @@ class TencentTraceClient: attributes={ service_attributes.SERVICE_NAME: service_name, service_attributes.SERVICE_VERSION: f"dify-{dify_config.project.version}-{dify_config.COMMIT_SHA}", - DEPLOYMENT_ENVIRONMENT: f"{dify_config.DEPLOY_ENV}-{dify_config.EDITION}", + DEPLOYMENT_ENVIRONMENT: f"{dify_config.DEPLOY_ENV}-{dify_config.DEPLOYMENT_EDITION.value}", HOST_NAME: socket.gethostname(), "telemetry.sdk.language": "python", "telemetry.sdk.name": "opentelemetry", diff --git a/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_client.py b/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_client.py index 98f829d1b7d..f0c2752ffec 100644 --- a/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_client.py +++ b/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_client.py @@ -16,6 +16,8 @@ from dify_trace_tencent.entities.tencent_trace_entity import SpanData from opentelemetry.sdk.trace import Event from opentelemetry.trace import SpanContext, Status, StatusCode, TraceFlags +from enums import DeploymentEdition + metric_reader_instances: list[DummyMetricReader] = [] meter_provider_instances: list[DummyMeterProvider] = [] @@ -158,7 +160,7 @@ def patch_core_components(monkeypatch: pytest.MonkeyPatch) -> PatchedCoreCompone project=SimpleNamespace(version="test"), COMMIT_SHA="sha", DEPLOY_ENV="dev", - EDITION="cloud", + DEPLOYMENT_EDITION=DeploymentEdition.CLOUD, ) monkeypatch.setattr(client_module, "dify_config", fake_config) @@ -331,6 +333,7 @@ def test_create_and_export_span_exception_logs_error( client = _build_client() span = patch_core_components["span"] span.get_span_context.return_value = _make_span_context(span_id=2) + # pyrefly: ignore [missing-attribute] client.tracer.start_span.side_effect = RuntimeError("boom") caplog.set_level(logging.DEBUG, logger=client_module.logger.name) @@ -430,6 +433,7 @@ def test_shutdown_logs_when_meter_provider_fails(caplog: pytest.LogCaptureFixtur meter_provider = meter_provider_instances[-1] meter_provider.shutdown.side_effect = RuntimeError("boom") assert client.metric_reader is not None + # pyrefly: ignore [missing-attribute] client.metric_reader.shutdown.side_effect = RuntimeError("boom") caplog.set_level(logging.DEBUG, logger=client_module.logger.name) diff --git a/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_tencent_trace.py b/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_tencent_trace.py index 71bcabd1416..feea2560ec5 100644 --- a/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_tencent_trace.py +++ b/api/providers/trace/trace-tencent/tests/unit_tests/tencent_trace/test_tencent_trace.py @@ -1,11 +1,15 @@ +"""Unit tests for Tencent tracing, including SQLite-backed account resolution.""" + import gc import logging import warnings +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest from dify_trace_tencent.config import TencentConfig from dify_trace_tencent.tencent_trace import TencentDataTrace +from sqlalchemy.orm import Session from core.ops.entities.trace_entity import ( DatasetRetrievalTraceInfo, @@ -18,7 +22,9 @@ from core.ops.entities.trace_entity import ( ) from graphon.entities import WorkflowNodeExecution from graphon.enums import BuiltinNodeTypes -from models import Account, App +from models import Account, App, Tenant, TenantAccountJoin +from models.account import TenantAccountRole +from models.model import AppMode, IconType logger = logging.getLogger(__name__) @@ -412,59 +418,100 @@ class TestTencentDataTrace: assert result is None assert len([r for r in caplog.records if r.levelno == logging.DEBUG]) >= 1 - def test_get_workflow_node_executions(self, tencent_data_trace): + @pytest.mark.parametrize( + "sqlite3_session", + [(Account, App, Tenant, TenantAccountJoin)], + indirect=True, + ) + def test_get_workflow_node_executions( + self, + tencent_data_trace, + monkeypatch: pytest.MonkeyPatch, + sqlite3_session: Session, + ) -> None: + account = Account(name="Trace User", email="trace-user@example.com") + tenant = Tenant(name="Trace Tenant") + sqlite3_session.add_all([account, tenant]) + sqlite3_session.flush() + app = App( + id="app-1", + tenant_id=tenant.id, + name="Trace App", + description="", + mode=AppMode.WORKFLOW, + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + created_by=account.id, + max_active_requests=0, + ) + tenant_join = TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + current=True, + role=TenantAccountRole.OWNER, + ) + sqlite3_session.add_all([app, tenant_join]) + sqlite3_session.commit() + trace_info = MagicMock(spec=WorkflowTraceInfo) - trace_info.metadata = {"app_id": "app-1"} + trace_info.metadata = {"app_id": app.id} trace_info.workflow_run_id = "run-1" + database = SimpleNamespace(engine=sqlite3_session.get_bind()) + monkeypatch.setattr("dify_trace_tencent.tencent_trace.db", database) + monkeypatch.setattr("models.account.db", database) - app = MagicMock(spec=App) - app.id = "app-1" - app.created_by = "user-1" - app.tenant_id = "tenant-1" + with patch("dify_trace_tencent.tencent_trace.SQLAlchemyWorkflowNodeExecutionRepository") as mock_repo: + mock_repo.return_value.get_by_workflow_execution.return_value = [] + results = tencent_data_trace._get_workflow_node_executions(trace_info) - account = MagicMock(spec=Account) - account.id = "user-1" + assert results == [] + service_account = mock_repo.call_args.kwargs["user"] + assert isinstance(service_account, Account) + assert service_account.id == account.id + assert mock_repo.call_args.kwargs["tenant_id"] == tenant.id - mock_executions = [MagicMock()] - - with patch("dify_trace_tencent.tencent_trace.db") as mock_db: - mock_db.engine = "engine" - with patch("dify_trace_tencent.tencent_trace.Session") as mock_session_ctx: - session = mock_session_ctx.return_value.__enter__.return_value - session.scalar.side_effect = [app, account] - - with patch("dify_trace_tencent.tencent_trace.SQLAlchemyWorkflowNodeExecutionRepository") as mock_repo: - mock_repo.return_value.get_by_workflow_execution.return_value = mock_executions - - results = tencent_data_trace._get_workflow_node_executions(trace_info) - - assert results == mock_executions - assert mock_repo.call_args.kwargs["tenant_id"] == "tenant-1" - - def test_get_workflow_node_executions_no_app_id(self, tencent_data_trace, caplog: pytest.LogCaptureFixture): + @pytest.mark.parametrize("sqlite3_session", [()], indirect=True) + def test_get_workflow_node_executions_no_app_id( + self, + tencent_data_trace, + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, + sqlite3_session: Session, + ) -> None: trace_info = MagicMock(spec=WorkflowTraceInfo) trace_info.metadata = {} + monkeypatch.setattr( + "dify_trace_tencent.tencent_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind()), + ) with caplog.at_level(logging.ERROR): results = tencent_data_trace._get_workflow_node_executions(trace_info) assert results == [] assert len([r for r in caplog.records if r.levelno == logging.ERROR]) >= 1 - def test_get_workflow_node_executions_app_not_found(self, tencent_data_trace, caplog: pytest.LogCaptureFixture): + @pytest.mark.parametrize("sqlite3_session", [(App,)], indirect=True) + def test_get_workflow_node_executions_app_not_found( + self, + tencent_data_trace, + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, + sqlite3_session: Session, + ) -> None: trace_info = MagicMock(spec=WorkflowTraceInfo) trace_info.metadata = {"app_id": "app-1"} + monkeypatch.setattr( + "dify_trace_tencent.tencent_trace.db", + SimpleNamespace(engine=sqlite3_session.get_bind()), + ) - with patch("dify_trace_tencent.tencent_trace.db") as mock_db: - mock_db.init_app = MagicMock() # Ensure init_app is mocked - mock_db.engine = "engine" - with patch("dify_trace_tencent.tencent_trace.Session") as mock_session_ctx: - session = mock_session_ctx.return_value.__enter__.return_value - session.scalar.return_value = None - - with caplog.at_level(logging.ERROR): - results = tencent_data_trace._get_workflow_node_executions(trace_info) - assert results == [] - assert len([r for r in caplog.records if r.levelno == logging.ERROR]) >= 1 + with caplog.at_level(logging.ERROR): + results = tencent_data_trace._get_workflow_node_executions(trace_info) + assert results == [] + assert len([r for r in caplog.records if r.levelno == logging.ERROR]) >= 1 def test_get_user_id_workflow(self, tencent_data_trace): trace_info = MagicMock(spec=WorkflowTraceInfo) diff --git a/api/providers/vdb/README.md b/api/providers/vdb/README.md index b5b4197f63c..e398cb938b9 100644 --- a/api/providers/vdb/README.md +++ b/api/providers/vdb/README.md @@ -29,7 +29,7 @@ In `pyproject.toml`: pgvector = "dify_vdb_pgvector.pgvector:PGVectorFactory" ``` -The value is **`module:attribute`**: a importable module path and the class implementing `AbstractVectorFactory`. +The value is **`module:attribute`**: an importable module path and the class implementing `AbstractVectorFactory`. ### How registration works diff --git a/api/providers/vdb/vdb-alibabacloud-mysql/pyproject.toml b/api/providers/vdb/vdb-alibabacloud-mysql/pyproject.toml index 9103f3e4f17..d54627d1f72 100644 --- a/api/providers/vdb/vdb-alibabacloud-mysql/pyproject.toml +++ b/api/providers/vdb/vdb-alibabacloud-mysql/pyproject.toml @@ -2,7 +2,7 @@ name = "dify-vdb-alibabacloud-mysql" version = "0.0.1" dependencies = [ - "mysql-connector-python>=9.3.0,<10.0.0", + "mysql-connector-python>=26.7.0,<27.0.0", ] description = "Dify vector store backend (dify-vdb-alibabacloud-mysql)." diff --git a/api/providers/vdb/vdb-analyticdb/pyproject.toml b/api/providers/vdb/vdb-analyticdb/pyproject.toml index f22e3e8e122..e8e120db0ff 100644 --- a/api/providers/vdb/vdb-analyticdb/pyproject.toml +++ b/api/providers/vdb/vdb-analyticdb/pyproject.toml @@ -2,7 +2,7 @@ name = "dify-vdb-analyticdb" version = "0.0.1" dependencies = [ - "alibabacloud_gpdb20160503~=5.2.0", + "alibabacloud_gpdb20160503~=5.9.0", "alibabacloud_tea_openapi==0.4.4", "clickhouse-connect==0.15.1", ] diff --git a/api/providers/vdb/vdb-baidu/pyproject.toml b/api/providers/vdb/vdb-baidu/pyproject.toml index bacff08793a..42368307993 100644 --- a/api/providers/vdb/vdb-baidu/pyproject.toml +++ b/api/providers/vdb/vdb-baidu/pyproject.toml @@ -2,7 +2,7 @@ name = "dify-vdb-baidu" version = "0.0.1" dependencies = [ - "pymochow==2.4.0", + "pymochow==2.4.1", ] description = "Dify vector store backend (dify-vdb-baidu)." diff --git a/api/providers/vdb/vdb-couchbase/pyproject.toml b/api/providers/vdb/vdb-couchbase/pyproject.toml index 6bc348b2eb5..f7281d10fd2 100644 --- a/api/providers/vdb/vdb-couchbase/pyproject.toml +++ b/api/providers/vdb/vdb-couchbase/pyproject.toml @@ -3,7 +3,7 @@ name = "dify-vdb-couchbase" version = "0.0.1" dependencies = [ - "couchbase~=4.6.0", + "couchbase~=4.6.2", ] description = "Dify vector store backend (dify-vdb-couchbase)." diff --git a/api/providers/vdb/vdb-iris/pyproject.toml b/api/providers/vdb/vdb-iris/pyproject.toml index c4da985032e..b79b665f32f 100644 --- a/api/providers/vdb/vdb-iris/pyproject.toml +++ b/api/providers/vdb/vdb-iris/pyproject.toml @@ -3,7 +3,7 @@ name = "dify-vdb-iris" version = "0.0.1" dependencies = [ - "intersystems-irispython>=5.1.0,<6.0.0", + "intersystems-irispython>=5.4.0,<6.0.0", ] description = "Dify vector store backend (dify-vdb-iris)." diff --git a/api/providers/vdb/vdb-matrixone/src/dify_vdb_matrixone/matrixone_vector.py b/api/providers/vdb/vdb-matrixone/src/dify_vdb_matrixone/matrixone_vector.py index f9fc7f45b7a..7a51776ad26 100644 --- a/api/providers/vdb/vdb-matrixone/src/dify_vdb_matrixone/matrixone_vector.py +++ b/api/providers/vdb/vdb-matrixone/src/dify_vdb_matrixone/matrixone_vector.py @@ -119,9 +119,8 @@ class MatrixoneVector(BaseVector): assert self.client is not None ids = [] for doc in documents: - if doc.metadata is not None: - doc_id = doc.metadata.get("doc_id", str(uuid.uuid4())) - ids.append(doc_id) + doc_id = doc.metadata.get("doc_id") if doc.metadata else None + ids.append(str(doc_id or uuid.uuid4())) self.client.insert( texts=[doc.page_content for doc in documents], embeddings=embeddings, @@ -167,6 +166,7 @@ class MatrixoneVector(BaseVector): filter = None if document_ids_filter: filter = {"document_id": {"$in": document_ids_filter}} + score_threshold = float(kwargs.get("score_threshold") or 0.0) results = self.client.query( query_vector=query_vector, @@ -175,15 +175,17 @@ class MatrixoneVector(BaseVector): ) docs = [] - # TODO: add the score threshold to the query for result in results: - metadata = result.metadata - docs.append( - Document( - page_content=result.document, - metadata=metadata, + metadata = parse_metadata_json(result.metadata) + score = 1.0 / (1.0 + float(result.distance)) + if score >= score_threshold: + metadata["score"] = score + docs.append( + Document( + page_content=result.document, + metadata=metadata, + ) ) - ) return docs @ensure_client diff --git a/api/providers/vdb/vdb-matrixone/tests/unit_tests/test_matrixone_vector.py b/api/providers/vdb/vdb-matrixone/tests/unit_tests/test_matrixone_vector.py index 762ec330b29..c14c9529fdf 100644 --- a/api/providers/vdb/vdb-matrixone/tests/unit_tests/test_matrixone_vector.py +++ b/api/providers/vdb/vdb-matrixone/tests/unit_tests/test_matrixone_vector.py @@ -168,7 +168,8 @@ def test_get_client_handles_full_text_index_creation_error(matrixone_module, mon def test_add_texts_generates_ids_and_inserts(matrixone_module, monkeypatch: pytest.MonkeyPatch): vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module)) vector.client = MagicMock() - monkeypatch.setattr(matrixone_module.uuid, "uuid4", lambda: "generated-uuid") + generated_ids = iter(["generated-id-b", "generated-id-c"]) + monkeypatch.setattr(matrixone_module.uuid, "uuid4", lambda: next(generated_ids)) docs = [ Document(page_content="a", metadata={"doc_id": "doc-a", "document_id": "d-1"}), Document(page_content="b", metadata={"document_id": "d-2"}), @@ -177,18 +178,15 @@ def test_add_texts_generates_ids_and_inserts(matrixone_module, monkeypatch: pyte ids = vector.add_texts(docs, [[0.1], [0.2], [0.3]]) - # For current prod code, only docs with metadata get ids, so only two ids - assert ids == ["doc-a", "generated-uuid"] + assert ids == ["doc-a", "generated-id-b", "generated-id-c"] vector.client.insert.assert_called_once() insert_kwargs = vector.client.insert.call_args.kwargs - # All lists passed to insert should be the same length texts = insert_kwargs["texts"] embeddings = insert_kwargs["embeddings"] metadatas = insert_kwargs["metadatas"] ids_insert = insert_kwargs["ids"] - assert len(texts) == len(embeddings) == len(metadatas) == len(docs) - # ids may be shorter than docs for current prod code, but should match number of docs with metadata - assert ids_insert == ["doc-a", "generated-uuid"] + assert len(ids_insert) == len(texts) == len(embeddings) == len(metadatas) == len(docs) + assert ids_insert == ids def test_delete_and_metadata_methods(matrixone_module): @@ -208,19 +206,25 @@ def test_delete_and_metadata_methods(matrixone_module): assert vector.client.delete.call_count == 3 -def test_search_by_vector_builds_documents(matrixone_module): +def test_search_by_vector_applies_score_threshold(matrixone_module): vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module)) vector.client = MagicMock() vector.client.query.return_value = [ - SimpleNamespace(document="doc-a", metadata={"doc_id": "1"}), - SimpleNamespace(document="doc-b", metadata={"doc_id": "2"}), + SimpleNamespace(document="doc-a", metadata={"doc_id": "1"}, distance=0.25), + SimpleNamespace(document="doc-b", metadata={"doc_id": "2"}, distance=2.0), ] - docs = vector.search_by_vector([0.1, 0.2], top_k=2, document_ids_filter=["d-1"]) + docs = vector.search_by_vector( + [0.1, 0.2], + top_k=2, + score_threshold=0.5, + document_ids_filter=["d-1"], + ) - assert len(docs) == 2 + assert len(docs) == 1 assert docs[0].page_content == "doc-a" - assert docs[1].metadata["doc_id"] == "2" + assert docs[0].metadata["doc_id"] == "1" + assert docs[0].metadata["score"] == pytest.approx(0.8) assert vector.client.query.call_args.kwargs["filter"] == {"document_id": {"$in": ["d-1"]}} diff --git a/api/providers/vdb/vdb-oceanbase/pyproject.toml b/api/providers/vdb/vdb-oceanbase/pyproject.toml index 7888c89724e..a417ecdd72f 100644 --- a/api/providers/vdb/vdb-oceanbase/pyproject.toml +++ b/api/providers/vdb/vdb-oceanbase/pyproject.toml @@ -4,7 +4,7 @@ version = "0.0.1" dependencies = [ "pyobvector==0.2.25", - "mysql-connector-python>=9.3.0,<10.0.0", + "mysql-connector-python>=26.7.0,<27.0.0", ] description = "Dify vector store backend (dify-vdb-oceanbase)." diff --git a/api/providers/vdb/vdb-oceanbase/src/dify_vdb_oceanbase/oceanbase_vector.py b/api/providers/vdb/vdb-oceanbase/src/dify_vdb_oceanbase/oceanbase_vector.py index 93be92a62f2..0cd445f857e 100644 --- a/api/providers/vdb/vdb-oceanbase/src/dify_vdb_oceanbase/oceanbase_vector.py +++ b/api/providers/vdb/vdb-oceanbase/src/dify_vdb_oceanbase/oceanbase_vector.py @@ -321,8 +321,13 @@ class OceanBaseVector(BaseVector): from sqlalchemy import text - # Validate key to prevent injection in JSON path - if not re.match(r"^[a-zA-Z0-9_.]+$", key): + # Validate key to prevent injection in JSON path. + # Use re.fullmatch instead of re.match to reject trailing newlines. + # Python's '$' matches at end-of-string OR just before a trailing + # newline, so re.match accepts "user_id\n". re.fullmatch requires + # the whole string to match. + # Regression for #39884 (sibling of #39234 / #39548 / #39666 / #39730 / #39880). + if not re.fullmatch(r"[a-zA-Z0-9_.]+", key): raise ValueError(f"Invalid characters in metadata key: {key}") # Use parameterized query to prevent SQL injection @@ -454,11 +459,14 @@ class OceanBaseVector(BaseVector): _where_clause = None if document_ids_filter: # Validate document IDs to prevent SQL injection - # Document IDs should be alphanumeric with hyphens and underscores + # Document IDs should be alphanumeric with hyphens and underscores. + # Use re.fullmatch instead of re.match to reject trailing newlines. + # See the metadata-key validator above for the rationale. + # Regression for #39884 (sibling of #39234 / #39548 / #39666 / #39730 / #39880). import re for doc_id in document_ids_filter: - if not isinstance(doc_id, str) or not re.match(r"^[a-zA-Z0-9_-]+$", doc_id): + if not isinstance(doc_id, str) or not re.fullmatch(r"[a-zA-Z0-9_-]+", doc_id): raise ValueError(f"Invalid document ID format: {doc_id}") # Safe to use in query after validation diff --git a/api/providers/vdb/vdb-oceanbase/tests/integration_tests/bench_oceanbase.py b/api/providers/vdb/vdb-oceanbase/tests/integration_tests/bench_oceanbase.py index 50f67369425..a6bf8f01c7f 100644 --- a/api/providers/vdb/vdb-oceanbase/tests/integration_tests/bench_oceanbase.py +++ b/api/providers/vdb/vdb-oceanbase/tests/integration_tests/bench_oceanbase.py @@ -204,7 +204,9 @@ def main(): # Re-query doc_id from one of the rows we inserted with client.engine.connect() as conn: res = conn.execute(text(f"SELECT metadata->>'$.document_id' FROM `{tbl_meta}` LIMIT 1")) - doc_id_1000 = res.fetchone()[0] + res1 = res.fetchone() + assert res1 is not None + doc_id_1000 = res1[0] logger.info("\n[Metadata filter query — 1000 rows, by document_id]") times_no_idx = bench_metadata_query(client, tbl_meta, doc_id_1000, with_index=False) diff --git a/api/providers/vdb/vdb-oceanbase/tests/unit_tests/test_oceanbase_vector.py b/api/providers/vdb/vdb-oceanbase/tests/unit_tests/test_oceanbase_vector.py index 36393cc486d..c5ca9d0a873 100644 --- a/api/providers/vdb/vdb-oceanbase/tests/unit_tests/test_oceanbase_vector.py +++ b/api/providers/vdb/vdb-oceanbase/tests/unit_tests/test_oceanbase_vector.py @@ -551,3 +551,54 @@ def test_oceanbase_factory_uses_existing_or_generated_collection(oceanbase_modul assert vector_cls.call_args_list[0].args[0] == "existing_collection" assert vector_cls.call_args_list[1].args[0] == "auto_collection" assert dataset_without_index.index_struct is not None + + +@pytest.mark.parametrize( + "bad_doc_id", + [ + "doc-123\n", # trailing LF -- the bug the #39884 fix closes + "doc-123\r", # trailing CR + "doc-123\r\n", # trailing CRLF + ], +) +def test_search_by_vector_rejects_document_id_with_trailing_newline(oceanbase_module, bad_doc_id): + """Regression for #39884 (sibling of #39234 / #39548 / #39666 / #39730 / #39880). + + The old `re.match(r"^[a-zA-Z0-9_-]+$", doc_id)` accepted a doc_id ending in + `\n` because Python's `$` matches just before a trailing newline. The fix + uses `re.fullmatch`, which requires the whole string to match. The + doc_ids are joined into a SQL `IN` clause so a trailing newline would + land in the SQL fragment. + """ + vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector) + vector._collection_name = "collection_1" + vector._hnsw_ef_search = -1 + vector._config = SimpleNamespace(metric_type="cosine") + vector._client = MagicMock() + + with pytest.raises(ValueError, match="Invalid document ID format"): + vector.search_by_vector([0.1, 0.2], document_ids_filter=[bad_doc_id]) + + +@pytest.mark.parametrize( + "bad_key", + [ + "user_id\n", # trailing LF -- the bug the #39884 fix closes + "user_id\r", # trailing CR + "user_id\r\n", # trailing CRLF + ], +) +def test_get_ids_by_metadata_field_rejects_key_with_trailing_newline(oceanbase_module, bad_key): + """Regression for #39884 (sibling of #39234 / #39548 / #39666 / #39730 / #39880). + + The old `re.match(r"^[a-zA-Z0-9_.]+$", key)` accepted a key ending in `\n`. + The key is interpolated into a SQL JSON-path expression + (`WHERE metadata->>'$.{key}' = :value`), so a trailing newline would + land in the SQL fragment. + """ + vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector) + vector._collection_name = "collection_1" + vector._client = MagicMock() + + with pytest.raises(Exception, match="Failed to query documents by metadata field"): + vector.get_ids_by_metadata_field(bad_key, "doc-1") diff --git a/api/providers/vdb/vdb-oracle/pyproject.toml b/api/providers/vdb/vdb-oracle/pyproject.toml index 6747485041e..d60b90c1813 100644 --- a/api/providers/vdb/vdb-oracle/pyproject.toml +++ b/api/providers/vdb/vdb-oracle/pyproject.toml @@ -3,7 +3,7 @@ name = "dify-vdb-oracle" version = "0.0.1" dependencies = [ - "oracledb==3.4.2", + "oracledb==4.0.2", ] description = "Dify vector store backend (dify-vdb-oracle)." diff --git a/api/providers/vdb/vdb-oracle/src/dify_vdb_oracle/oraclevector.py b/api/providers/vdb/vdb-oracle/src/dify_vdb_oracle/oraclevector.py index b8639dae619..831ccd32009 100644 --- a/api/providers/vdb/vdb-oracle/src/dify_vdb_oracle/oraclevector.py +++ b/api/providers/vdb/vdb-oracle/src/dify_vdb_oracle/oraclevector.py @@ -316,10 +316,10 @@ class OracleVector(BaseVector): entities.append(current_entity) else: try: - nltk.data.find("tokenizers/punkt") + nltk.data.find("tokenizers/punkt_tab") nltk.data.find("corpora/stopwords") except LookupError: - raise LookupError("Unable to find the required NLTK data package: punkt and stopwords") + raise LookupError("Unable to find the required NLTK data package: punkt_tab and stopwords") e_str = re.sub(r"[^\w ]", "", query) all_tokens = nltk.word_tokenize(e_str) stop_words = stopwords.words("english") diff --git a/api/providers/vdb/vdb-pgvector/pyproject.toml b/api/providers/vdb/vdb-pgvector/pyproject.toml index 2a972aa2772..b9dd73892cc 100644 --- a/api/providers/vdb/vdb-pgvector/pyproject.toml +++ b/api/providers/vdb/vdb-pgvector/pyproject.toml @@ -3,7 +3,7 @@ name = "dify-vdb-pgvector" version = "0.0.1" dependencies = [ - "pgvector==0.4.2", + "pgvector==0.5.0", ] description = "Dify vector store backend (dify-vdb-pgvector)." diff --git a/api/providers/vdb/vdb-tablestore/pyproject.toml b/api/providers/vdb/vdb-tablestore/pyproject.toml index fd1a2d54e0c..04633b43024 100644 --- a/api/providers/vdb/vdb-tablestore/pyproject.toml +++ b/api/providers/vdb/vdb-tablestore/pyproject.toml @@ -3,7 +3,7 @@ name = "dify-vdb-tablestore" version = "0.0.1" dependencies = [ - "tablestore==6.4.4", + "tablestore==6.4.8", ] description = "Dify vector store backend (dify-vdb-tablestore)." diff --git a/api/providers/vdb/vdb-tidb-on-qdrant/src/dify_vdb_tidb_on_qdrant/tidb_on_qdrant_vector.py b/api/providers/vdb/vdb-tidb-on-qdrant/src/dify_vdb_tidb_on_qdrant/tidb_on_qdrant_vector.py index b352243b92a..be6df15cae1 100644 --- a/api/providers/vdb/vdb-tidb-on-qdrant/src/dify_vdb_tidb_on_qdrant/tidb_on_qdrant_vector.py +++ b/api/providers/vdb/vdb-tidb-on-qdrant/src/dify_vdb_tidb_on_qdrant/tidb_on_qdrant_vector.py @@ -46,6 +46,11 @@ if TYPE_CHECKING: type MetadataFilter = DictFilter | common_types.Filter +# Bounded connect/read timeout so a slow or hanging TiDB Cloud API call +# cannot block a cluster provisioning or password rotation forever. +_TIDB_CLOUD_REQUEST_TIMEOUT = httpx.Timeout(30.0, connect=5.0) + + class TidbOnQdrantConfig(BaseModel): endpoint: str api_key: str | None = None @@ -356,7 +361,7 @@ class TidbOnQdrantVector(BaseVector): query_filter=filter, limit=kwargs.get("top_k", 4), with_payload=True, - with_vectors=True, + with_vectors=False, score_threshold=kwargs.get("score_threshold", 0.0), ) docs = [] @@ -551,6 +556,7 @@ class TidbOnQdrantVectorFactory(AbstractVectorFactory): f"{tidb_config.api_url}/clusters", json=cluster_data, auth=DigestAuth(tidb_config.public_key, tidb_config.private_key), + timeout=_TIDB_CLOUD_REQUEST_TIMEOUT, ) if response.status_code == 200: @@ -574,6 +580,7 @@ class TidbOnQdrantVectorFactory(AbstractVectorFactory): f"{tidb_config.api_url}/clusters/{cluster_id}/password", json=body, auth=DigestAuth(tidb_config.public_key, tidb_config.private_key), + timeout=_TIDB_CLOUD_REQUEST_TIMEOUT, ) if response.status_code == 200: diff --git a/api/providers/vdb/vdb-tidb-on-qdrant/tests/unit_tests/test_tidb_on_qdrant_vector.py b/api/providers/vdb/vdb-tidb-on-qdrant/tests/unit_tests/test_tidb_on_qdrant_vector.py index 76802de62ef..f94c6496897 100644 --- a/api/providers/vdb/vdb-tidb-on-qdrant/tests/unit_tests/test_tidb_on_qdrant_vector.py +++ b/api/providers/vdb/vdb-tidb-on-qdrant/tests/unit_tests/test_tidb_on_qdrant_vector.py @@ -226,3 +226,40 @@ class TestInitVectorEndpointSelection: qdrant_url = binding_endpoint or global_url or "" assert qdrant_url == "https://qdrant-global.tidb.com" + + +class TestTidbCloudRequestTimeouts: + """The TiDB Cloud API calls must be bounded: a hanging endpoint must not + block cluster provisioning or password rotation forever.""" + + @patch("dify_vdb_tidb_on_qdrant.tidb_on_qdrant_vector.httpx.post") + def test_create_cluster_passes_bounded_timeout(self, mock_post): + from dify_vdb_tidb_on_qdrant.tidb_on_qdrant_vector import ( + TidbConfig, + TidbOnQdrantVectorFactory, + ) + + mock_post.return_value.status_code = 200 + mock_post.return_value.json.return_value = {} + factory = TidbOnQdrantVectorFactory() + config = TidbConfig(api_url="https://api.tidbcloud.test", public_key="pub", private_key="priv") + + factory.create_tidb_serverless_cluster(config, "display", "us-east-1") + + assert mock_post.call_args.kwargs["timeout"] is not None + + @patch("dify_vdb_tidb_on_qdrant.tidb_on_qdrant_vector.httpx.put") + def test_change_password_passes_bounded_timeout(self, mock_put): + from dify_vdb_tidb_on_qdrant.tidb_on_qdrant_vector import ( + TidbConfig, + TidbOnQdrantVectorFactory, + ) + + mock_put.return_value.status_code = 200 + mock_put.return_value.json.return_value = {} + factory = TidbOnQdrantVectorFactory() + config = TidbConfig(api_url="https://api.tidbcloud.test", public_key="pub", private_key="priv") + + factory.change_tidb_serverless_root_password(config, "c-1", "new-password") + + assert mock_put.call_args.kwargs["timeout"] is not None diff --git a/api/providers/vdb/vdb-tidb-vector/src/dify_vdb_tidb_vector/tidb_vector.py b/api/providers/vdb/vdb-tidb-vector/src/dify_vdb_tidb_vector/tidb_vector.py index 9f80ae5a76f..1fadb5aef58 100644 --- a/api/providers/vdb/vdb-tidb-vector/src/dify_vdb_tidb_vector/tidb_vector.py +++ b/api/providers/vdb/vdb-tidb-vector/src/dify_vdb_tidb_vector/tidb_vector.py @@ -20,6 +20,8 @@ from models.dataset import Dataset logger = logging.getLogger(__name__) +FULLTEXT_INDEX_NAME = "idx_text" + class TiDBVectorConfig(BaseModel): host: str @@ -28,6 +30,7 @@ class TiDBVectorConfig(BaseModel): password: str database: str program_name: str + enable_fulltext_search: bool = False @model_validator(mode="before") @classmethod @@ -95,10 +98,15 @@ class TiDBVector(BaseVector): logger.info("_create_collection, collection_name %s", self._collection_name) lock_name = f"vector_indexing_lock_{self._collection_name}" with redis_client.lock(lock_name, timeout=20): - collection_exist_cache_key = f"vector_indexing_{self._collection_name}" + collection_exist_cache_key = self._collection_exist_cache_key() if redis_client.get(collection_exist_cache_key): return tidb_dist_func = self._get_distance_func() + fulltext_index_statement = ( + f",\n FULLTEXT INDEX {FULLTEXT_INDEX_NAME} (text) WITH PARSER MULTILINGUAL" + if self._client_config.enable_fulltext_search + else "" + ) with sessionmaker(bind=self._engine).begin() as session: create_statement = sql_text(f""" CREATE TABLE IF NOT EXISTS {self._collection_name} ( @@ -113,11 +121,52 @@ class TiDBVector(BaseVector): KEY (doc_id), KEY (document_id), VECTOR INDEX idx_vector (({tidb_dist_func}(vector))) USING HNSW + {fulltext_index_statement} ); """) session.execute(create_statement) + if self._client_config.enable_fulltext_search: + self._ensure_fulltext_index(session) redis_client.set(collection_exist_cache_key, 1, ex=3600) + def _collection_exist_cache_key(self) -> str: + search_mode = "fulltext" if self._client_config.enable_fulltext_search else "semantic" + return f"vector_indexing_{self._collection_name}_{search_mode}" + + def _ensure_fulltext_index(self, session) -> None: + index_check_statement = sql_text(""" + SELECT COUNT(1) + FROM INFORMATION_SCHEMA.STATISTICS + WHERE TABLE_SCHEMA = DATABASE() + AND TABLE_NAME = :table_name + AND INDEX_NAME = :index_name + """) + result = session.execute( + index_check_statement, + params={ + "index_name": FULLTEXT_INDEX_NAME, + "table_name": self._collection_name, + }, + ) + if result.scalar(): + return + + session.execute( + sql_text( + f"ALTER TABLE {self._collection_name} " + f"ADD FULLTEXT INDEX {FULLTEXT_INDEX_NAME} (text) WITH PARSER MULTILINGUAL;" + ) + ) + + @staticmethod + def _document_ids_filter_condition(document_ids_filter: list[str] | None) -> tuple[str, dict[str, str]]: + if not document_ids_filter: + return "", {} + + filter_params = {f"document_id_{index}": document_id for index, document_id in enumerate(document_ids_filter)} + placeholders = ", ".join(f":{param_name}" for param_name in filter_params) + return f"document_id IN ({placeholders})", filter_params + @override def add_texts(self, documents: list[Document], embeddings: list[list[float]], **kwargs): table = self._table(len(embeddings[0])) @@ -127,8 +176,8 @@ class TiDBVector(BaseVector): chunks_table_data = [] with self._engine.connect() as conn, conn.begin(): - for id, text, meta, embedding in zip(ids, texts, metas, embeddings): - chunks_table_data.append({"id": id, "vector": embedding, "text": text, "meta": meta}) + for doc_id, text, meta, embedding in zip(ids, texts, metas, embeddings): + chunks_table_data.append({"id": doc_id, "vector": embedding, "text": text, "meta": meta}) # Execute the batch insert when the batch size is reached if len(chunks_table_data) == 500: @@ -205,9 +254,10 @@ class TiDBVector(BaseVector): tidb_dist_func = self._get_distance_func() document_ids_filter = kwargs.get("document_ids_filter") where_clause = "" + filter_params: dict[str, str] = {} if document_ids_filter: - document_ids = ", ".join(f"'{id}'" for id in document_ids_filter) - where_clause = f" WHERE meta->>'$.document_id' in ({document_ids}) " + document_ids_filter_condition, filter_params = self._document_ids_filter_condition(document_ids_filter) + where_clause = f" WHERE {document_ids_filter_condition} " with Session(self._engine) as session: select_statement = sql_text(f""" @@ -230,6 +280,7 @@ class TiDBVector(BaseVector): "query_vector_str": query_vector_str, "distance": distance, "top_k": top_k, + **filter_params, }, ) results = [(row[0], row[1], row[2]) for row in res] @@ -241,8 +292,51 @@ class TiDBVector(BaseVector): @override def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]: - # tidb doesn't support bm25 search - return [] + if not self._client_config.enable_fulltext_search or not query: + return [] + + top_k = kwargs.get("top_k", 4) + score_threshold = float(kwargs.get("score_threshold") or 0.0) + document_ids_filter = kwargs.get("document_ids_filter") + + where_conditions = ["FTS_MATCH_WORD(text, :query)"] + filter_params: dict[str, str] = {} + if document_ids_filter: + document_ids_filter_condition, filter_params = self._document_ids_filter_condition(document_ids_filter) + where_conditions.append(document_ids_filter_condition) + where_clause = " AND ".join(where_conditions) + + docs = [] + with Session(self._engine) as session: + select_statement = sql_text(f""" + SELECT meta, text, score + FROM ( + SELECT + meta, + text, + FTS_MATCH_WORD(text, :query) AS score + FROM {self._collection_name} + WHERE {where_clause} + ORDER BY score DESC + LIMIT :top_k + ) t + WHERE score >= :score_threshold + """) + res = session.execute( + select_statement, + params={ + "query": query, + "score_threshold": score_threshold, + "top_k": top_k, + **filter_params, + }, + ) + results = [(row[0], row[1], row[2]) for row in res] + for meta, text, score in results: + metadata = parse_metadata_json(meta) + metadata["score"] = score + docs.append(Document(page_content=text, metadata=metadata)) + return docs @override def delete(self): @@ -280,5 +374,6 @@ class TiDBVectorFactory(AbstractVectorFactory): password=dify_config.TIDB_VECTOR_PASSWORD or "", database=dify_config.TIDB_VECTOR_DATABASE or "", program_name=dify_config.APPLICATION_NAME, + enable_fulltext_search=dify_config.TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH, ), ) diff --git a/api/providers/vdb/vdb-tidb-vector/tests/unit_tests/test_tidb_vector.py b/api/providers/vdb/vdb-tidb-vector/tests/unit_tests/test_tidb_vector.py index ed03cbee88d..47d63f0641d 100644 --- a/api/providers/vdb/vdb-tidb-vector/tests/unit_tests/test_tidb_vector.py +++ b/api/providers/vdb/vdb-tidb-vector/tests/unit_tests/test_tidb_vector.py @@ -25,6 +25,7 @@ def _config(tidb_module): password="secret", database="dify", program_name="dify-app", + enable_fulltext_search=False, ) @@ -118,11 +119,13 @@ def test_create_collection_skips_when_cache_hit(tidb_module, monkeypatch: pytest vector = tidb_module.TiDBVector.__new__(tidb_module.TiDBVector) vector._collection_name = "collection_1" vector._engine = MagicMock() + vector._client_config = _config(tidb_module) tidb_module.Session = MagicMock() vector._create_collection(3) + tidb_module.redis_client.get.assert_called_once_with("vector_indexing_collection_1_semantic") tidb_module.Session.assert_not_called() tidb_module.redis_client.set.assert_not_called() @@ -151,15 +154,90 @@ def test_create_collection_executes_create_sql_and_sets_cache(tidb_module, monke vector._collection_name = "collection_1" vector._engine = MagicMock() vector._distance_func = "l2" + vector._client_config = _config(tidb_module) vector._create_collection(3) sql = str(session.execute.call_args.args[0]) assert "VECTOR(3)" in sql assert "VEC_L2_DISTANCE" in sql + assert "FULLTEXT INDEX" not in sql tidb_module.redis_client.set.assert_called_once() +def test_create_collection_adds_fulltext_index_when_enabled(tidb_module, monkeypatch: pytest.MonkeyPatch): + lock = MagicMock() + lock.__enter__.return_value = None + lock.__exit__.return_value = None + monkeypatch.setattr(tidb_module.redis_client, "lock", MagicMock(return_value=lock)) + monkeypatch.setattr(tidb_module.redis_client, "get", MagicMock(return_value=None)) + monkeypatch.setattr(tidb_module.redis_client, "set", MagicMock()) + + session = MagicMock() + + class _BeginCtx: + def __enter__(self): + return session + + def __exit__(self, exc_type, exc, tb): + return False + + mock_sm = MagicMock(begin=MagicMock(return_value=_BeginCtx())) + monkeypatch.setattr(tidb_module, "sessionmaker", lambda **kwargs: mock_sm) + + vector = tidb_module.TiDBVector.__new__(tidb_module.TiDBVector) + vector._collection_name = "collection_1" + vector._engine = MagicMock() + vector._distance_func = "cosine" + vector._client_config = _config(tidb_module).model_copy(update={"enable_fulltext_search": True}) + + vector._create_collection(3) + + sql = str(session.execute.call_args_list[0].args[0]) + assert "FULLTEXT INDEX idx_text (text) WITH PARSER MULTILINGUAL" in sql + + +def test_create_collection_ensures_fulltext_index_when_enabled(tidb_module, monkeypatch: pytest.MonkeyPatch): + lock = MagicMock() + lock.__enter__.return_value = None + lock.__exit__.return_value = None + monkeypatch.setattr(tidb_module.redis_client, "lock", MagicMock(return_value=lock)) + monkeypatch.setattr(tidb_module.redis_client, "get", MagicMock(return_value=None)) + monkeypatch.setattr(tidb_module.redis_client, "set", MagicMock()) + + session = MagicMock() + index_check_result = MagicMock() + index_check_result.scalar.return_value = 0 + session.execute.side_effect = [None, index_check_result, None] + + class _BeginCtx: + def __enter__(self): + return session + + def __exit__(self, exc_type, exc, tb): + return False + + mock_sm = MagicMock(begin=MagicMock(return_value=_BeginCtx())) + monkeypatch.setattr(tidb_module, "sessionmaker", lambda **kwargs: mock_sm) + + vector = tidb_module.TiDBVector.__new__(tidb_module.TiDBVector) + vector._collection_name = "collection_1" + vector._engine = MagicMock() + vector._distance_func = "cosine" + vector._client_config = _config(tidb_module).model_copy(update={"enable_fulltext_search": True}) + + vector._create_collection(3) + + executed_sql = [str(call.args[0]) for call in session.execute.call_args_list] + assert "INFORMATION_SCHEMA.STATISTICS" in executed_sql[1] + assert "ALTER TABLE collection_1 ADD FULLTEXT INDEX idx_text (text) WITH PARSER MULTILINGUAL" in executed_sql[2] + assert session.execute.call_args_list[1].kwargs["params"] == { + "index_name": "idx_text", + "table_name": "collection_1", + } + tidb_module.redis_client.get.assert_called_once_with("vector_indexing_collection_1_fulltext") + + def test_add_texts_batches_inserts_and_returns_ids(tidb_module, monkeypatch: pytest.MonkeyPatch): class _InsertStmt: def __init__(self, table): @@ -215,10 +293,45 @@ def tidb_vector_with_session(tidb_module, monkeypatch: pytest.MonkeyPatch): return vector, session, tidb_module -# 1. search_by_full_text returns empty -def test_search_by_full_text_returns_empty(tidb_vector_with_session): - vector, _, _ = tidb_vector_with_session +# 1. search_by_full_text returns empty when disabled +def test_search_by_full_text_returns_empty_when_disabled(tidb_vector_with_session): + vector, session, tidb_module = tidb_vector_with_session + vector._client_config = _config(tidb_module) assert vector.search_by_full_text("query") == [] + session.execute.assert_not_called() + + +def test_search_by_full_text_queries_tidb_fts_and_scores(tidb_vector_with_session): + vector, session, tidb_module = tidb_vector_with_session + vector._client_config = _config(tidb_module).model_copy(update={"enable_fulltext_search": True}) + session.execute.return_value = [ + ('{"doc_id":"id-1","document_id":"d-1"}', "text-1", 0.8), + ('{"doc_id":"id-2","document_id":"d-2"}', "text-2", 0.6), + ] + + docs = vector.search_by_full_text( + "search query", + top_k=2, + score_threshold=0.5, + document_ids_filter=["d-1", "d'2"], + ) + + assert len(docs) == 2 + assert docs[0].page_content == "text-1" + assert docs[0].metadata["score"] == pytest.approx(0.8) + assert docs[1].metadata["score"] == pytest.approx(0.6) + sql = str(session.execute.call_args.args[0]) + params = session.execute.call_args.kwargs["params"] + assert "FTS_MATCH_WORD(text, :query)" in sql + assert "document_id IN (:document_id_0, :document_id_1)" in sql + assert "d'2" not in sql + assert params == { + "document_id_0": "d-1", + "document_id_1": "d'2", + "query": "search query", + "score_threshold": 0.5, + "top_k": 2, + } # 2. text_exists returns True when ids found @@ -378,15 +491,18 @@ def test_search_by_vector_filters_and_scores(tidb_module, monkeypatch: pytest.Mo [0.1, 0.2], top_k=2, score_threshold=0.5, - document_ids_filter=["d-1", "d-2"], + document_ids_filter=["d-1", "d'2"], ) assert len(docs) == 2 assert docs[0].metadata["score"] == pytest.approx(0.8) assert docs[1].metadata["score"] == pytest.approx(0.6) sql = str(session.execute.call_args.args[0]) params = session.execute.call_args.kwargs["params"] - assert "meta->>'$.document_id' in ('d-1', 'd-2')" in sql + assert "document_id IN (:document_id_0, :document_id_1)" in sql + assert "d'2" not in sql assert params["distance"] == pytest.approx(0.5) + assert params["document_id_0"] == "d-1" + assert params["document_id_1"] == "d'2" assert params["top_k"] == 2 session.commit.assert_not_called() @@ -428,6 +544,7 @@ def test_tidb_factory_uses_existing_or_generated_collection(tidb_module, monkeyp monkeypatch.setattr(tidb_module.dify_config, "TIDB_VECTOR_USER", "root") monkeypatch.setattr(tidb_module.dify_config, "TIDB_VECTOR_PASSWORD", "secret") monkeypatch.setattr(tidb_module.dify_config, "TIDB_VECTOR_DATABASE", "dify") + monkeypatch.setattr(tidb_module.dify_config, "TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH", True) monkeypatch.setattr(tidb_module.dify_config, "APPLICATION_NAME", "dify-app") with patch.object(tidb_module, "TiDBVector", return_value="vector") as vector_cls: @@ -438,4 +555,5 @@ def test_tidb_factory_uses_existing_or_generated_collection(tidb_module, monkeyp assert result_2 == "vector" assert vector_cls.call_args_list[0].kwargs["collection_name"] == "existing_collection" assert vector_cls.call_args_list[1].kwargs["collection_name"] == "auto_collection" + assert vector_cls.call_args_list[0].kwargs["config"].enable_fulltext_search is True assert dataset_without_index.index_struct is not None diff --git a/api/providers/vdb/vdb-weaviate/pyproject.toml b/api/providers/vdb/vdb-weaviate/pyproject.toml index 035fbd396d1..f2e75f16073 100644 --- a/api/providers/vdb/vdb-weaviate/pyproject.toml +++ b/api/providers/vdb/vdb-weaviate/pyproject.toml @@ -3,7 +3,7 @@ name = "dify-vdb-weaviate" version = "0.0.1" dependencies = [ - "weaviate-client==4.20.5", + "weaviate-client==4.22.0", ] description = "Dify vector store backend (dify-vdb-weaviate)." diff --git a/api/pyproject.toml b/api/pyproject.toml index 53ed8927292..fc84c35fa68 100644 --- a/api/pyproject.toml +++ b/api/pyproject.toml @@ -1,12 +1,12 @@ [project] name = "dify-api" -version = "1.16.0" +version = "1.16.1" requires-python = "~=3.12.0" dependencies = [ # Legacy: mature and widely deployed "bleach>=6.4.0,<7.0.0", - "boto3>=1.43.46,<2.0.0", + "boto3>=1.43.56,<2.0.0", "celery>=5.6.3,<6.0.0", "croniter>=6.2.2,<7.0.0", "dify-agent", @@ -33,19 +33,19 @@ dependencies = [ "flask-restx>=1.3.2,<2.0.0", "google-cloud-aiplatform>=1.160.0,<2.0.0", "httpx[socks]==0.28.1", - "opentelemetry-distro==0.62b1", - "opentelemetry-instrumentation-celery==0.62b1", - "opentelemetry-instrumentation-flask==0.62b1", - "opentelemetry-instrumentation-httpx==0.62b1", - "opentelemetry-instrumentation-redis==0.62b1", - "opentelemetry-instrumentation-sqlalchemy==0.62b1", + "opentelemetry-distro==0.65b0", + "opentelemetry-instrumentation-celery==0.65b0", + "opentelemetry-instrumentation-flask==0.65b0", + "opentelemetry-instrumentation-httpx==0.65b0", + "opentelemetry-instrumentation-redis==0.65b0", + "opentelemetry-instrumentation-sqlalchemy==0.65b0", "opentelemetry-propagator-b3>=1.41.1,<2.0.0", "readabilipy==0.3.0", "resend>=2.27.0,<3.0.0", "zstandard==0.25.0", # Emerging: newer and fast-moving, use compatible pins "fastopenapi[flask]==0.7.0", - "graphon==0.6.0", + "graphon==0.7.0", "httpx-sse==0.4.3", "json-repair==0.60.1", ] @@ -107,12 +107,12 @@ dify-trace-tencent = { workspace = true } dify-trace-weave = { workspace = true } [tool.uv] -default-groups = ["storage", "tools", "vdb-all", "trace-all"] +default-groups = ["kms", "storage", "tools", "vdb-all", "trace-all"] package = false override-dependencies = [ "litellm>=1.83.10,<2.0.0", "pyarrow>=23.0.1,<24.0.0", - "cryptography>=49.0.0,<50.0.0", + "cryptography>=50.0.0,<51.0.0", "setuptools>=80.10.2,<81", ] @@ -183,8 +183,19 @@ dev = [ # "locust>=2.40.4", # Temporarily removed due to compatibility issues. Uncomment when resolved. "pytest-timeout>=2.4.0", "pytest-xdist>=3.8.0", - "pyrefly>=1.0.0", + "pyrefly>=1.2.0", "xinference-client>=2.7.0", + "types-pytz>=2026.3.1.20260727", + "types-bleach>=6.4.0.20260728", + "types-croniter>=6.2.4.20260711", +] + +############################################################ +# [ KMS ] dependency group +# Required for key management / credential encryption key providers +############################################################ +kms = [ + "azure-keyvault-keys>=4.10.0,<5.0.0", ] ############################################################ @@ -193,10 +204,10 @@ dev = [ ############################################################ storage = [ "azure-storage-blob>=12.30.0,<13.0.0", - "bce-python-sdk==0.9.72", + "bce-python-sdk==0.9.76", "cos-python-sdk-v5>=1.9.44,<2.0.0", "esdk-obs-python>=3.26.6,<4.0.0", - "google-cloud-storage>=3.12.1,<4.0.0", + "google-cloud-storage>=3.13.0,<4.0.0", "opendal==0.46.0", "oss2>=2.19.1,<3.0.0", "supabase>=2.31.0,<3.0.0", @@ -206,7 +217,7 @@ storage = [ ############################################################ # [ Tools ] dependency group ############################################################ -tools = ["cloudscraper>=1.2.71,<2.0.0", "nltk>=3.9.1,<4.0.0"] +tools = ["cloudscraper>=1.2.71,<2.0.0", "nltk>=3.10.0,<4.0.0"] ############################################################ # [ VDB ] workspace plugins — hollow packages under providers/vdb/* @@ -304,4 +315,4 @@ python-platform = "linux" python-version = "3.12.0" infer-with-first-use = true min-severity = "warn" -errors = { missing-override-decorator = "error" } +errors = { missing-override-decorator = "error", unnecessary-type-conversion = "info" } diff --git a/api/repositories/explore_banner_query_repository.py b/api/repositories/explore_banner_query_repository.py new file mode 100644 index 00000000000..011828b7977 --- /dev/null +++ b/api/repositories/explore_banner_query_repository.py @@ -0,0 +1,47 @@ +"""Database repository for the Explore banner read model.""" + +from typing import override + +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker + +from models.enums import BannerStatus +from models.model import ExporleBanner +from services.explore_banner_query_service import ExploreBannerQuery, ExploreBannerRecord + + +class ExploreBannerQueryRepository(ExploreBannerQuery): + def __init__(self, client: sessionmaker[Session]) -> None: + self._client = client + + @override + def list_enabled(self, language: str) -> tuple[ExploreBannerRecord, ...]: + stmt = ( + select( + ExporleBanner.id, + ExporleBanner.content, + ExporleBanner.link, + ExporleBanner.sort, + ExporleBanner.status, + ExporleBanner.created_at, + ) + .where( + ExporleBanner.status == BannerStatus.ENABLED, + ExporleBanner.language == language, + ) + .order_by(ExporleBanner.sort) + ) + + with self._client() as session: + rows = session.execute(stmt).all() + return tuple( + ExploreBannerRecord( + id=banner_id, + content=content, + link=link, + sort=sort, + status=status.value, + created_at=created_at, + ) + for banner_id, content, link, sort, status, created_at in rows + ) diff --git a/api/repositories/installation_state_repository.py b/api/repositories/installation_state_repository.py new file mode 100644 index 00000000000..755a3134933 --- /dev/null +++ b/api/repositories/installation_state_repository.py @@ -0,0 +1,27 @@ +"""Persistence adapter for installation setup and tenant-existence state.""" + +from datetime import datetime + +from sqlalchemy import exists, select +from sqlalchemy.orm import Session, sessionmaker + +from models.account import Tenant +from models.model import DifySetup + + +class InstallationStateRepository: + """Read persistent state shared by installation bootstrap use cases.""" + + def __init__(self, client: sessionmaker[Session]) -> None: + self._client = client + + def get_setup_at(self) -> datetime | None: + with self._client() as session: + return session.scalar(select(DifySetup.setup_at).limit(1)) + + def is_setup(self) -> bool: + return self.get_setup_at() is not None + + def has_tenants(self) -> bool: + with self._client() as session: + return session.scalar(select(exists().select_from(Tenant))) is True diff --git a/api/repositories/workspace_member_query_repository.py b/api/repositories/workspace_member_query_repository.py new file mode 100644 index 00000000000..1c67ae13e7f --- /dev/null +++ b/api/repositories/workspace_member_query_repository.py @@ -0,0 +1,62 @@ +"""Database repository for the workspace-member read model.""" + +from typing import override + +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker + +from models.account import Account, TenantAccountJoin +from services.workspace_member_query_service import WorkspaceMemberQuery, WorkspaceMemberRecord + + +class WorkspaceMemberQueryRepository(WorkspaceMemberQuery): + def __init__(self, session_factory: sessionmaker[Session]) -> None: + self._session_factory = session_factory + + @override + def list_for_workspace(self, workspace_id: str) -> tuple[WorkspaceMemberRecord, ...]: + stmt = ( + select( + Account.id, + Account.name, + Account.email, + Account.avatar, + Account.last_login_at, + Account.last_active_at, + Account.created_at, + Account.status, + TenantAccountJoin.role, + ) + .select_from(Account) + .join(TenantAccountJoin, TenantAccountJoin.account_id == Account.id) + .where(TenantAccountJoin.tenant_id == workspace_id) + ) + + with self._session_factory() as session: + rows = session.execute(stmt).all() + records = tuple( + WorkspaceMemberRecord( + id=account_id, + name=name, + email=email, + avatar=avatar, + last_login_at=last_login_at, + last_active_at=last_active_at, + created_at=created_at, + status=status.value, + legacy_role=legacy_role.value, + ) + for ( + account_id, + name, + email, + avatar, + last_login_at, + last_active_at, + created_at, + status, + legacy_role, + ) in rows + ) + + return records diff --git a/api/repositories/workspace_query_repository.py b/api/repositories/workspace_query_repository.py new file mode 100644 index 00000000000..d221601f9ea --- /dev/null +++ b/api/repositories/workspace_query_repository.py @@ -0,0 +1,45 @@ +"""Database repository for the workspace-list read model.""" + +from typing import override + +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker + +from models.account import Tenant, TenantAccountJoin, TenantStatus +from services.workspace_query_service import WorkspaceQuery, WorkspaceRecord + + +class WorkspaceQueryRepository(WorkspaceQuery): + def __init__(self, client: sessionmaker[Session]) -> None: + self._client = client + + @override + def list_for_account(self, account_id: str) -> tuple[WorkspaceRecord, ...]: + stmt = ( + select( + Tenant.id, + Tenant.name, + Tenant.status, + Tenant.created_at, + TenantAccountJoin.last_opened_at, + ) + .join(TenantAccountJoin, TenantAccountJoin.tenant_id == Tenant.id) + .where( + TenantAccountJoin.account_id == account_id, + Tenant.status == TenantStatus.NORMAL, + ) + .order_by(Tenant.created_at.asc()) + ) + + with self._client() as session: + rows = session.execute(stmt).all() + return tuple( + WorkspaceRecord( + id=workspace_id, + name=name, + status=status.value, + created_at=created_at, + last_opened_at=last_opened_at, + ) + for workspace_id, name, status, created_at, last_opened_at in rows + ) diff --git a/api/schedule/clean_messages.py b/api/schedule/clean_messages.py index be5f483b959..6cad69f34e9 100644 --- a/api/schedule/clean_messages.py +++ b/api/schedule/clean_messages.py @@ -19,15 +19,14 @@ def clean_messages(): Clean expired messages based on clean policy. This task uses MessagesCleanService to efficiently clean messages in batches. - The behavior depends on BILLING_ENABLED configuration: - - BILLING_ENABLED=True: only delete messages from sandbox tenants (with whitelist/grace period) - - BILLING_ENABLED=False: delete all messages within the time range + Cloud only deletes messages from sandbox tenants (with whitelist/grace period). + Self-hosted editions delete all messages within the configured time range. """ click.echo(click.style("clean_messages: start clean messages.", fg="green")) start_at = time.perf_counter() try: - # Create policy based on billing configuration + # Create policy based on deployment edition. policy = create_message_clean_policy( graceful_period_days=dify_config.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD, ) diff --git a/api/schedule/clean_unused_datasets_task.py b/api/schedule/clean_unused_datasets_task.py index 9bdb074647c..02770c83232 100644 --- a/api/schedule/clean_unused_datasets_task.py +++ b/api/schedule/clean_unused_datasets_task.py @@ -10,7 +10,7 @@ import app from configs import dify_config from core.db.session_factory import session_factory from core.rag.index_processor.index_processor_factory import IndexProcessorFactory -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from extensions.ext_redis import redis_client from libs.pagination import paginate_query from models.dataset import Dataset, DatasetAutoDisableLog, DatasetQuery, Document diff --git a/api/schedule/mail_clean_document_notify_task.py b/api/schedule/mail_clean_document_notify_task.py index 1a76a4aa306..549cf9cef0d 100644 --- a/api/schedule/mail_clean_document_notify_task.py +++ b/api/schedule/mail_clean_document_notify_task.py @@ -8,7 +8,7 @@ from sqlalchemy import select import app from configs import dify_config from core.db.session_factory import session_factory -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from extensions.ext_mail import mail from libs.email_i18n import EmailType, get_email_i18n_service from models import Account, Tenant, TenantAccountJoin diff --git a/api/services/account_service.py b/api/services/account_service.py index cc2d983c4ef..3f460275e20 100644 --- a/api/services/account_service.py +++ b/api/services/account_service.py @@ -22,15 +22,16 @@ from werkzeug.exceptions import Unauthorized from configs import dify_config from constants.languages import get_valid_language, language_timezone_mapping +from enums import DeploymentEdition from events.tenant_event import tenant_was_created from extensions.ext_database import db from extensions.ext_redis import redis_client, redis_fallback from libs.datetime_utils import naive_utc_now from libs.helper import RateLimiter, TokenManager from libs.helper import timezone as validate_timezone +from libs.key_providers import generate_key_pair from libs.passport import PassportService from libs.password import compare_password, hash_password, valid_password -from libs.rsa import generate_key_pair from libs.token import generate_csrf_token from models.account import ( Account, @@ -75,6 +76,7 @@ from services.errors.account import ( from services.errors.workspace import WorkSpaceNotAllowedCreateError, WorkspacesLimitExceededError from services.feature_service import FeatureService from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService +from services.telemetry_service import CommunityTelemetryService from tasks.delete_account_task import delete_account_task from tasks.mail_account_deletion_task import send_account_deletion_verification_code from tasks.mail_change_mail_task import ( @@ -118,7 +120,7 @@ class InvitationDetailDict(TypedDict): def _try_join_enterprise_default_workspace(account_id: str) -> None: """Best-effort join to enterprise default workspace.""" - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return from services.enterprise.enterprise_service import try_join_default_workspace @@ -324,26 +326,41 @@ class AccountService: if account.status == AccountStatus.BANNED: raise Unauthorized("Account is banned.") - current_tenant = session.scalar( + current_tenant_join = session.scalar( select(TenantAccountJoin) .where(TenantAccountJoin.account_id == account.id, TenantAccountJoin.current == True) .limit(1) ) - if current_tenant: - account.set_tenant_id_with_session(current_tenant.tenant_id, session=session) - else: - available_ta = session.scalar( + if current_tenant_join is not None: + account.set_tenant_id_with_session(current_tenant_join.tenant_id, session=session) + + has_valid_current_tenant = ( + current_tenant_join is not None + and account.current_tenant is not None + and account.current_tenant.status == TenantStatus.NORMAL + ) + if not has_valid_current_tenant: + if current_tenant_join is not None: + current_tenant_join.current = False + + available_tenant_join = session.scalar( select(TenantAccountJoin) - .where(TenantAccountJoin.account_id == account.id) + .join(Tenant, TenantAccountJoin.tenant_id == Tenant.id) + .where( + TenantAccountJoin.account_id == account.id, + Tenant.status == TenantStatus.NORMAL, + ) .order_by(TenantAccountJoin.id.asc()) .limit(1) ) - if not available_ta: + if available_tenant_join is None: + if current_tenant_join is not None: + session.commit() return None - account.set_tenant_id_with_session(available_ta.tenant_id, session=session) - available_ta.current = True - available_ta.last_opened_at = naive_utc_now() + account.set_tenant_id_with_session(available_tenant_join.tenant_id, session=session) + available_tenant_join.current = True + available_tenant_join.last_opened_at = naive_utc_now() session.commit() AccountService._refresh_account_last_active(account, session) @@ -362,7 +379,7 @@ class AccountService: payload = { "user_id": account.id, "exp": exp, - "iss": dify_config.EDITION, + "iss": dify_config.DEPLOYMENT_EDITION.value, "sub": "Console API Passport", } @@ -447,7 +464,7 @@ class AccountService: if not FeatureService.get_license().seats.is_available(): raise SeatsLimitExceededError("licensed seats limit exceeded") - if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(email): + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and BillingService.is_email_in_freeze(email): raise AccountRegisterError( description=( "This email account has been deleted within the past " @@ -1037,7 +1054,7 @@ class AccountService: @classmethod def get_user_through_email(cls, email: str, *, session: Session): - if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(email): + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and BillingService.is_email_in_freeze(email): raise AccountRegisterError( description=( "This email account has been deleted within the past " @@ -1056,7 +1073,7 @@ class AccountService: @classmethod def is_account_in_freeze(cls, email: str) -> bool: - if dify_config.BILLING_ENABLED and BillingService.is_email_in_freeze(email): + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and BillingService.is_email_in_freeze(email): return True return False @@ -1379,7 +1396,7 @@ class TenantService: session.add(ta) session.commit() - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: BillingService.clean_billing_info_cache(tenant.id) return ta @@ -1623,8 +1640,7 @@ class TenantService: select(Account, TenantAccountJoin.role) .select_from(Account) .join(TenantAccountJoin, Account.id == TenantAccountJoin.account_id) - .where(TenantAccountJoin.tenant_id == tenant.id) - .where(TenantAccountJoin.role == "dataset_operator") + .where(TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.role == "dataset_operator") ) # Initialize an empty list to store the updated accounts @@ -1813,7 +1829,7 @@ class TenantService: account_email, ) - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: BillingService.clean_billing_info_cache(tenant.id) # Queue account deletion sync task for enterprise backend to reassign resources (enterprise only) @@ -1953,7 +1969,7 @@ class RegisterService: TenantService.create_owner_tenant_if_not_exist(account=account, is_setup=True, session=session) - dify_setup = DifySetup(version=dify_config.project.version) + dify_setup = DifySetup(version=dify_config.project.version, instance_id=str(uuid.uuid4())) session.add(dify_setup) session.commit() except Exception as e: @@ -1966,6 +1982,11 @@ class RegisterService: logger.exception("Setup account failed, email: %s, name: %s", email, name) raise ValueError(f"Setup failed: {e}") + try: + CommunityTelemetryService.report_install(session=session) + except Exception: + logger.debug("Failed to report install telemetry", exc_info=True) + @classmethod def register( cls, diff --git a/api/services/agent/composer_service.py b/api/services/agent/composer_service.py index 96f0b8dde23..13924aae4b1 100644 --- a/api/services/agent/composer_service.py +++ b/api/services/agent/composer_service.py @@ -9,7 +9,7 @@ from sqlalchemy.sql.elements import ColumnElement from core.agent.publish_visibility import agent_has_workflow_callable_active_snapshot from libs.helper import to_timestamp -from models import Account +from models import Account, App, Conversation from models.agent import ( APP_BACKED_AGENT_SOURCES, Agent, @@ -18,12 +18,15 @@ from models.agent import ( AgentConfigRevision, AgentConfigRevisionOperation, AgentConfigSnapshot, + AgentConfigVersionKind, + AgentDebugConversation, AgentDriveFile, AgentIconType, AgentKind, AgentScope, AgentSource, AgentStatus, + AgentWorkspaceOwnerType, WorkflowAgentBindingType, WorkflowAgentNodeBinding, ) @@ -35,6 +38,7 @@ from models.workflow import Workflow from services.agent.agent_soul_state import agent_soul_has_model from services.agent.composer_validator import ComposerConfigValidator from services.agent.errors import ( + AgentBuildSandboxNotFoundError, AgentModelNotConfiguredError, AgentNameConflictError, AgentNotFoundError, @@ -42,11 +46,17 @@ from services.agent.errors import ( AgentVersionNotFoundError, InvalidComposerConfigError, ) +from services.agent.home_snapshot_service import ( + AgentHomeSnapshotService, + validate_home_snapshot_binding, +) from services.agent.knowledge_datasets import ( get_tenant_knowledge_dataset_rows, list_missing_tenant_knowledge_dataset_ids, ) +from services.agent.retirement_service import WorkflowAgentRetirementService from services.agent.roster_service import AgentRosterService +from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService, WorkspaceOwnerScope from services.app_service import AppService, CreateAppParams from services.entities.agent_entities import ( AgentSoulConfig, @@ -56,6 +66,7 @@ from services.entities.agent_entities import ( ComposerVariant, WorkflowNodeJobConfig, ) +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection # WorkflowAgentNodeBinding.workflow_version tag for the draft workflow row. # Mirrors Workflow.version when it is "draft" (see models/workflow.py). @@ -198,6 +209,18 @@ class AgentComposerService: binding = cls._get_workflow_binding( session=session, tenant_id=tenant_id, workflow_id=workflow.id, node_id=node_id ) + retirement_candidates = ( + {binding.agent_id} + if binding is not None + and binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT + and binding.agent_id + and payload.save_strategy + in { + ComposerSaveStrategy.SAVE_AS_NEW_AGENT, + ComposerSaveStrategy.SAVE_TO_ROSTER, + } + else set() + ) match payload.save_strategy: case ComposerSaveStrategy.NODE_JOB_ONLY: @@ -232,7 +255,11 @@ class AgentComposerService: ) case ComposerSaveStrategy.SAVE_TO_ROSTER: binding = cls._save_to_roster( - session=session, tenant_id=tenant_id, account_id=account_id, binding=binding, payload=payload + session=session, + tenant_id=tenant_id, + account_id=account_id, + binding=binding, + payload=payload, ) session.flush() @@ -257,6 +284,17 @@ class AgentComposerService: payload=payload, agent_id=binding.agent_id, ) + session.commit() + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=tenant_id, + agent_ids=retirement_candidates, + account_id=account_id, + ) + enqueue_agent_resource_collection( + tenant_id=tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) return state @classmethod @@ -390,12 +428,10 @@ class AgentComposerService: @classmethod def _load_agent_composer_for_agent(cls, *, session: Session, tenant_id: str, agent: Agent) -> dict[str, Any]: - draft = cls._get_or_create_agent_draft( + draft = cls.get_or_create_normal_agent_draft( session=session, tenant_id=tenant_id, agent=agent, - draft_type=AgentConfigDraftType.DRAFT, - account_id=None, created_by=agent.updated_by or agent.created_by, ) version = cls._get_version_if_present( @@ -405,6 +441,7 @@ class AgentComposerService: "variant": ComposerVariant.AGENT_APP.value, "agent": cls._serialize_agent(agent), "active_config_snapshot": cls._serialize_version(version), + "active_config_is_published": bool(agent.active_config_snapshot_id and agent.active_config_is_published), "draft": cls._serialize_draft(draft), "agent_soul": draft.config_snapshot_dict, "save_options": [ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION.value], @@ -416,7 +453,34 @@ class AgentComposerService: @classmethod def save_agent_app_composer( - cls, *, session: Session, tenant_id: str, app_id: str, account_id: str, payload: ComposerSavePayload + cls, + *, + session: Session, + tenant_id: str, + app_id: str, + account_id: str, + payload: ComposerSavePayload, + ) -> dict[str, Any]: + try: + return cls._save_agent_app_composer_impl( + session=session, + tenant_id=tenant_id, + app_id=app_id, + account_id=account_id, + payload=payload, + ) + except IntegrityError as exc: + raise AgentNameConflictError() from exc + + @classmethod + def _save_agent_app_composer_impl( + cls, + *, + session: Session, + tenant_id: str, + app_id: str, + account_id: str, + payload: ComposerSavePayload, ) -> dict[str, Any]: if payload.variant != ComposerVariant.AGENT_APP: raise ValueError("Agent App composer endpoint only accepts agent_app variant") @@ -445,11 +509,20 @@ class AgentComposerService: updated_by=account_id, ) session.add(agent) - try: - session.flush() - except IntegrityError as exc: - session.rollback() - raise AgentNameConflictError() from exc + session.flush() + initial_version = cls._create_config_version( + session=session, + tenant_id=tenant_id, + agent_id=agent.id, + account_id=account_id, + agent_soul=AgentSoulConfig(), + operation=AgentConfigRevisionOperation.CREATE_VERSION, + version_note=None, + home_snapshot_id=None, + ) + agent.active_config_snapshot_id = initial_version.id + agent.active_config_has_model = False + agent.active_config_is_published = False return cls._save_agent_composer_for_agent( session=session, tenant_id=tenant_id, @@ -487,7 +560,7 @@ class AgentComposerService: ) -> dict[str, Any]: if payload.agent_soul is None: raise ValueError("agent_soul is required") - cls._save_agent_draft( + draft = cls._save_agent_draft( session=session, tenant_id=tenant_id, agent=agent, @@ -502,6 +575,7 @@ class AgentComposerService: tenant_id=tenant_id, agent=agent, agent_soul=payload.agent_soul, + home_snapshot_id=draft.home_snapshot_id, ) session.flush() @@ -522,6 +596,7 @@ class AgentComposerService: tenant_id: str, agent: Agent, agent_soul: AgentSoulConfig, + home_snapshot_id: str | None, ) -> bool: if not agent.active_config_snapshot_id: return False @@ -537,7 +612,9 @@ class AgentComposerService: if not agent_has_workflow_callable_active_snapshot(session=session, agent=agent): return False - return _agent_soul_config_json(agent_soul) == _agent_soul_config_json(active_version.config_snapshot_dict) + return home_snapshot_id == active_version.home_snapshot_id and _agent_soul_config_json( + agent_soul + ) == _agent_soul_config_json(active_version.config_snapshot_dict) @classmethod def publish_agent_app_draft( @@ -546,6 +623,7 @@ class AgentComposerService: agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) if agent.scope != AgentScope.ROSTER or agent.source not in APP_BACKED_AGENT_SOURCES: raise AgentNotFoundError() + access_was_ready = agent_has_workflow_callable_active_snapshot(session=session, agent=agent) draft = cls._get_or_create_agent_draft( session=session, tenant_id=tenant_id, @@ -566,6 +644,11 @@ class AgentComposerService: if not agent_soul_has_model(agent_soul): raise AgentModelNotConfiguredError() cls.validate_knowledge_datasets(session=session, tenant_id=tenant_id, agent_soul=agent_soul) + validate_home_snapshot_binding( + session=session, + agent=agent, + home_snapshot_id=draft.home_snapshot_id, + ) version = cls._create_config_version( session=session, tenant_id=tenant_id, @@ -575,6 +658,7 @@ class AgentComposerService: operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, version_note=version_note, previous_snapshot_id=agent.active_config_snapshot_id, + home_snapshot_id=draft.home_snapshot_id, ) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(agent_soul) @@ -582,6 +666,22 @@ class AgentComposerService: agent.updated_by = account_id draft.base_snapshot_id = version.id draft.updated_by = account_id + if not access_was_ready: + if not agent.app_id: + raise AgentNotFoundError() + app = session.scalar( + select(App) + .where( + App.tenant_id == tenant_id, + App.id == agent.app_id, + ) + .limit(1) + ) + if app is None: + raise AgentNotFoundError() + app.enable_site = True + app.enable_api = True + app.updated_by = account_id session.flush() return { "result": "success", @@ -594,6 +694,29 @@ class AgentComposerService: def checkout_agent_app_build_draft( cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, force: bool = False ) -> dict[str, Any]: + try: + result, retired_binding_id = cls._checkout_agent_app_build_draft_in_transaction( + session=session, + tenant_id=tenant_id, + agent_id=agent_id, + account_id=account_id, + force=force, + ) + session.commit() + except Exception: + session.rollback() + raise + if retired_binding_id is not None: + enqueue_agent_resource_collection( + tenant_id=tenant_id, + binding_ids=(retired_binding_id,), + ) + return result + + @classmethod + def _checkout_agent_app_build_draft_in_transaction( + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str, force: bool + ) -> tuple[dict[str, Any], str | None]: agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) normal_draft = cls._get_or_create_agent_draft( session=session, @@ -611,7 +734,23 @@ class AgentComposerService: account_id=account_id, ) if build_draft is not None and not force: - return cls._serialize_build_draft_state(build_draft) + return cls._serialize_build_draft_state(build_draft), None + retired_binding_id: str | None = None + if build_draft is not None and build_draft.agent_workspace_binding_id is not None: + cls._validate_active_build_draft_binding( + session=session, + tenant_id=tenant_id, + agent=agent, + build_draft=build_draft, + ) + retired_binding_id = AgentWorkspaceService.retire_binding( + session=session, + tenant_id=tenant_id, + binding_id=build_draft.agent_workspace_binding_id, + ) + if retired_binding_id is None: + raise AgentBuildSandboxNotFoundError() + build_draft.agent_workspace_binding_id = None if build_draft is None: build_draft = AgentConfigDraft( tenant_id=tenant_id, @@ -623,10 +762,44 @@ class AgentComposerService: ) session.add(build_draft) build_draft.base_snapshot_id = normal_draft.base_snapshot_id + build_draft.home_snapshot_id = normal_draft.home_snapshot_id build_draft.config_snapshot = AgentSoulConfig.model_validate(normal_draft.config_snapshot_dict) build_draft.updated_by = account_id session.flush() - return cls._serialize_build_draft_state(build_draft) + return cls._serialize_build_draft_state(build_draft), retired_binding_id + + @classmethod + def _validate_active_build_draft_binding( + cls, + *, + session: Session, + tenant_id: str, + agent: Agent, + build_draft: AgentConfigDraft, + ) -> None: + binding_id = build_draft.agent_workspace_binding_id + runtime_app_id = AgentRosterService.runtime_backing_app_id(agent) + if binding_id is None or runtime_app_id is None: + raise AgentBuildSandboxNotFoundError() + binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=tenant_id, + binding_id=binding_id, + expected_owner_scope=WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=runtime_app_id, + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id=build_draft.id, + ), + ) + if binding is None or binding.agent_id != agent.id: + raise AgentBuildSandboxNotFoundError() + AgentWorkspaceService.validate_binding_generation( + binding, + base_home_snapshot_id=build_draft.home_snapshot_id, + agent_config_version_id=build_draft.id, + agent_config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) @classmethod def load_agent_app_build_draft( @@ -668,6 +841,29 @@ class AgentComposerService: def apply_agent_app_build_draft( cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str ) -> dict[str, Any]: + try: + result, retired_binding_ids = cls._apply_agent_app_build_draft_in_transaction( + session=session, + tenant_id=tenant_id, + agent_id=agent_id, + account_id=account_id, + ) + session.commit() + except Exception: + session.rollback() + raise + enqueue_agent_resource_collection(tenant_id=tenant_id, binding_ids=retired_binding_ids) + return result + + @classmethod + def _apply_agent_app_build_draft_in_transaction( + cls, + *, + session: Session, + tenant_id: str, + agent_id: str, + account_id: str, + ) -> tuple[dict[str, Any], list[str]]: agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) build_draft = cls._get_agent_draft( session=session, @@ -679,6 +875,21 @@ class AgentComposerService: if build_draft is None: raise AgentVersionNotFoundError() applied_agent_soul = AgentSoulConfig.model_validate(build_draft.config_snapshot_dict) + ComposerConfigValidator.validate_publish_payload( + ComposerSavePayload( + variant=ComposerVariant.AGENT_APP, + agent_soul=applied_agent_soul, + save_strategy=ComposerSaveStrategy.SAVE_AS_NEW_VERSION, + ) + ) + cls.validate_knowledge_datasets(session=session, tenant_id=tenant_id, agent_soul=applied_agent_soul) + source_binding_id = build_draft.agent_workspace_binding_id + if source_binding_id is None: + raise AgentBuildSandboxNotFoundError() + home_snapshot = AgentHomeSnapshotService.create_for_build_apply( + session=session, + build_draft=build_draft, + ) normal_draft = cls._save_agent_draft( session=session, tenant_id=tenant_id, @@ -689,32 +900,145 @@ class AgentComposerService: account_id_for_audit=account_id, base_snapshot_id=build_draft.base_snapshot_id, ) + retired_binding_ids = cls._retire_normal_preview_bindings( + session=session, + tenant_id=tenant_id, + agent=agent, + normal_draft=normal_draft, + ) + normal_draft.home_snapshot_id = home_snapshot.id agent.active_config_is_published = cls._agent_soul_matches_active_config( session=session, tenant_id=tenant_id, agent=agent, agent_soul=applied_agent_soul, + home_snapshot_id=home_snapshot.id, ) agent.updated_by = account_id + retired_binding_id = AgentWorkspaceService.retire_binding( + session=session, + tenant_id=tenant_id, + binding_id=source_binding_id, + ) + if retired_binding_id is None: + raise AgentBuildSandboxNotFoundError() + retired_binding_ids.append(source_binding_id) session.delete(build_draft) - session.flush() - return {"result": "success", "draft": cls._serialize_draft(normal_draft)} + return {"result": "success", "draft": cls._serialize_draft(normal_draft)}, retired_binding_ids + + @classmethod + def _retire_normal_preview_bindings( + cls, + *, + session: Session, + tenant_id: str, + agent: Agent, + normal_draft: AgentConfigDraft, + ) -> list[str]: + """Retire Preview participants before Build Apply replaces the shared Draft Home.""" + + mappings = session.scalars( + select(AgentDebugConversation).where( + AgentDebugConversation.tenant_id == tenant_id, + AgentDebugConversation.agent_id == agent.id, + AgentDebugConversation.draft_type == AgentConfigDraftType.DRAFT, + ) + ).all() + retired_binding_ids: list[str] = [] + for mapping in mappings: + conversation = session.scalar( + select(Conversation).where( + Conversation.id == mapping.conversation_id, + Conversation.app_id == mapping.app_id, + Conversation.is_deleted.is_(False), + ) + ) + if conversation is None or conversation.agent_workspace_binding_id is None: + continue + binding_id = conversation.agent_workspace_binding_id + binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=tenant_id, + binding_id=binding_id, + expected_owner_scope=WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=mapping.app_id, + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id=conversation.id, + ), + ) + if binding is None or binding.agent_id != agent.id: + raise AgentWorkspaceNotFoundError("Agent Preview participant Binding is unavailable") + AgentWorkspaceService.validate_binding_generation( + binding, + base_home_snapshot_id=normal_draft.home_snapshot_id, + agent_config_version_id=normal_draft.id, + agent_config_version_kind=AgentConfigVersionKind.DRAFT, + ) + retired_binding_id = AgentWorkspaceService.retire_binding( + session=session, + tenant_id=tenant_id, + binding_id=binding_id, + ) + if retired_binding_id is None: + raise AgentWorkspaceNotFoundError("Agent Preview participant Binding is unavailable") + conversation.agent_workspace_binding_id = None + retired_binding_ids.append(binding_id) + return retired_binding_ids @classmethod def discard_agent_app_build_draft( cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str ) -> dict[str, Any]: + try: + result, retired_binding_id = cls._discard_agent_app_build_draft_in_transaction( + session=session, + tenant_id=tenant_id, + agent_id=agent_id, + account_id=account_id, + ) + session.commit() + except Exception: + session.rollback() + raise + if retired_binding_id is not None: + enqueue_agent_resource_collection( + tenant_id=tenant_id, + binding_ids=(retired_binding_id,), + ) + return result + + @classmethod + def _discard_agent_app_build_draft_in_transaction( + cls, *, session: Session, tenant_id: str, agent_id: str, account_id: str + ) -> tuple[dict[str, Any], str | None]: + agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=agent_id) build_draft = cls._get_agent_draft( session=session, tenant_id=tenant_id, - agent_id=agent_id, + agent_id=agent.id, draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id=account_id, ) - if build_draft is not None: - session.delete(build_draft) - session.flush() - return {"result": "success"} + if build_draft is None: + return {"result": "success"}, None + retired_binding_id: str | None = None + if build_draft.agent_workspace_binding_id is not None: + cls._validate_active_build_draft_binding( + session=session, + tenant_id=tenant_id, + agent=agent, + build_draft=build_draft, + ) + retired_binding_id = AgentWorkspaceService.retire_binding( + session=session, + tenant_id=tenant_id, + binding_id=build_draft.agent_workspace_binding_id, + ) + if retired_binding_id is None: + raise AgentBuildSandboxNotFoundError() + session.delete(build_draft) + return {"result": "success"}, retired_binding_id @classmethod def collect_validation_findings( @@ -1278,6 +1602,12 @@ class AgentComposerService: binding = cls._require_binding(binding) if not binding.agent_id or payload.agent_soul is None: raise ValueError("agent_id and agent_soul are required") + current_snapshot = cls._require_version( + session=session, + tenant_id=tenant_id, + agent_id=binding.agent_id, + version_id=binding.current_snapshot_id, + ) version = cls._create_config_version( session=session, tenant_id=tenant_id, @@ -1286,6 +1616,7 @@ class AgentComposerService: agent_soul=payload.agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, version_note=payload.version_note, + home_snapshot_id=current_snapshot.home_snapshot_id, ) agent = cls._require_agent(session=session, tenant_id=tenant_id, agent_id=binding.agent_id) agent.active_config_snapshot_id = version.id @@ -1456,6 +1787,7 @@ class AgentComposerService: agent_soul=agent_soul, operation=AgentConfigRevisionOperation.CREATE_VERSION, version_note=None, + home_snapshot_id=None, ) agent.active_config_snapshot_id = version.id agent.active_config_has_model = agent_soul_has_model(agent_soul) @@ -1589,7 +1921,6 @@ class AgentComposerService: session=session, ) except IntegrityError as exc: - session.rollback() raise AgentNameConflictError() from exc agent = AgentRosterService(session).get_app_backing_agent(tenant_id=tenant_id, app_id=app.id) @@ -1627,6 +1958,7 @@ class AgentComposerService: agent_soul: AgentSoulConfig, operation: AgentConfigRevisionOperation, version_note: str | None, + home_snapshot_id: str | None, previous_snapshot_id: str | None = None, ) -> AgentConfigSnapshot: next_version = ( @@ -1643,6 +1975,7 @@ class AgentComposerService: agent_id=agent_id, version=next_version, config_snapshot=agent_soul, + home_snapshot_id=home_snapshot_id, version_note=version_note, created_by=account_id, ) @@ -1682,6 +2015,7 @@ class AgentComposerService: operation=operation, version_note=version_note, previous_snapshot_id=current_snapshot.id, + home_snapshot_id=current_snapshot.home_snapshot_id, ) @classmethod @@ -1748,7 +2082,10 @@ class AgentComposerService: agent: Agent, created_by: str | None, ) -> AgentConfigDraft: - """Resolve the shared Preview draft, rebasing inline agents when needed.""" + """Resolve the normal Draft, rebasing only stale WORKFLOW_ONLY DRAFT rows whose account_id is None. + + Roster and DEBUG_BUILD Drafts are never rebased. + """ return cls._get_or_create_agent_draft( session=session, tenant_id=tenant_id, @@ -1766,6 +2103,8 @@ class AgentComposerService: snapshot: AgentConfigSnapshot, updated_by: str | None, ) -> bool: + """Sync a stale normal Draft's base_snapshot_id, home_snapshot_id, config_snapshot, and updated_by.""" + if ( agent.scope != AgentScope.WORKFLOW_ONLY or draft.draft_type != AgentConfigDraftType.DRAFT @@ -1776,6 +2115,7 @@ class AgentComposerService: ): return False draft.base_snapshot_id = snapshot.id + draft.home_snapshot_id = snapshot.home_snapshot_id draft.config_snapshot = AgentSoulConfig.model_validate(snapshot.config_snapshot_dict) draft.updated_by = updated_by return True @@ -1812,7 +2152,9 @@ class AgentComposerService: agent_id=agent.id, version_id=agent.active_config_snapshot_id, ) - if active_snapshot is not None and cls._rebase_workflow_only_normal_draft( + if active_snapshot is None: + raise AgentVersionNotFoundError() + if cls._rebase_workflow_only_normal_draft( agent=agent, draft=draft, snapshot=active_snapshot, @@ -1826,18 +2168,17 @@ class AgentComposerService: agent_id=agent.id, version_id=agent.active_config_snapshot_id, ) - agent_soul = ( - AgentSoulConfig.model_validate(base_snapshot.config_snapshot_dict) - if base_snapshot is not None - else AgentSoulConfig() - ) + if base_snapshot is None: + raise AgentVersionNotFoundError() + agent_soul = AgentSoulConfig.model_validate(base_snapshot.config_snapshot_dict) draft = AgentConfigDraft( tenant_id=tenant_id, agent_id=agent.id, draft_type=draft_type, account_id=account_id if draft_type == AgentConfigDraftType.DEBUG_BUILD else None, draft_owner_key=account_id if draft_type == AgentConfigDraftType.DEBUG_BUILD and account_id else "", - base_snapshot_id=base_snapshot.id if base_snapshot else None, + base_snapshot_id=base_snapshot.id, + home_snapshot_id=base_snapshot.home_snapshot_id, config_snapshot=agent_soul, created_by=created_by, updated_by=created_by, @@ -2139,7 +2480,7 @@ class AgentComposerService: from services.agent.roster_service import AgentRosterService - return AgentRosterService(session).get_or_create_agent_app_debug_conversation_id( + return AgentRosterService(session).get_or_create_build_conversation( tenant_id=tenant_id, agent_id=agent.id, account_id=account_id, diff --git a/api/services/agent/dsl_service.py b/api/services/agent/dsl_service.py index b2ffae45dad..e1f9afd0142 100644 --- a/api/services/agent/dsl_service.py +++ b/api/services/agent/dsl_service.py @@ -224,6 +224,7 @@ class AgentDslService: account_id=None, draft_owner_key="", base_snapshot_id=snapshot.id, + home_snapshot_id=snapshot.home_snapshot_id, config_snapshot=soul, created_by=account.id, updated_by=account.id, @@ -243,7 +244,7 @@ class AgentDslService: portable_graph: Mapping[str, Any], raw_packages: Mapping[str, Any], account: Account, - ) -> tuple[dict[str, Any], list[DslImportWarning]]: + ) -> tuple[dict[str, Any], list[DslImportWarning], set[str]]: """Materialize every packaged Agent as a node-owned inline Agent.""" graph = copy.deepcopy(dict(portable_graph)) @@ -256,6 +257,11 @@ class AgentDslService: WorkflowAgentNodeBinding.workflow_version == Workflow.VERSION_DRAFT, ) ).all() + retirement_candidates = { + binding.agent_id + for binding in previous_bindings + if binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT and binding.agent_id + } for binding in previous_bindings: self.session.delete(binding) self.session.flush() @@ -312,7 +318,7 @@ class AgentDslService: workflow.graph = json.dumps(graph) self.session.flush() - return graph, warnings + return graph, warnings, retirement_candidates def clone_inline_binding_for_node( self, @@ -567,6 +573,7 @@ class AgentDslService: agent_id=agent.id, version=next_version, config_snapshot=soul, + home_snapshot_id=None, created_by=account_id, ) self.session.add(snapshot) diff --git a/api/services/agent/errors.py b/api/services/agent/errors.py index 163687815d8..45f10031491 100644 --- a/api/services/agent/errors.py +++ b/api/services/agent/errors.py @@ -29,6 +29,18 @@ class AgentModelNotConfiguredError(BaseHTTPException): code = 400 +class AgentAccessNotReadyError(BaseHTTPException): + error_code = "agent_not_published" + description = "Publish the Agent before enabling Web App or API access." + code = 409 + + +class AgentBuildSandboxNotFoundError(BaseHTTPException): + error_code = "agent_build_sandbox_not_found" + description = "The retained Build Sandbox is no longer available." + code = 404 + + class AgentSoulLockedError(BadRequest): description = "Agent Soul is locked for this workflow node." diff --git a/api/services/agent/home_snapshot_service.py b/api/services/agent/home_snapshot_service.py new file mode 100644 index 00000000000..937c2506441 --- /dev/null +++ b/api/services/agent/home_snapshot_service.py @@ -0,0 +1,212 @@ +"""Own immutable Agent Home Snapshot ledger rows and physical collection.""" + +from __future__ import annotations + +import logging + +from dify_agent.client import Client, DifyAgentNotFoundError +from dify_agent.protocol import CreateHomeSnapshotFromBindingRequest +from sqlalchemy import select +from sqlalchemy.orm import Session + +from configs import dify_config +from core.db.session_factory import session_factory +from libs.datetime_utils import naive_utc_now +from libs.uuid_utils import uuidv7 +from models.agent import ( + Agent, + AgentConfigDraft, + AgentConfigSnapshot, + AgentConfigVersionKind, + AgentHomeSnapshot, + AgentStatus, + AgentWorkingResourceStatus, + AgentWorkspaceOwnerType, +) +from services.agent.errors import AgentBuildSandboxNotFoundError +from services.agent.workspace_service import AgentWorkspaceService, WorkspaceOwnerScope + +logger = logging.getLogger(__name__) + + +class AgentHomeSnapshotUnavailableError(RuntimeError): + """The requested owner-scoped Home Snapshot cannot be used.""" + + +class AgentHomeSnapshotService: + """Create, retire, and collect Agent-owned immutable Home Snapshots.""" + + @classmethod + def create_for_build_apply( + cls, + *, + session: Session, + build_draft: AgentConfigDraft, + ) -> AgentHomeSnapshot: + """Checkpoint the exact participant owned by ``build_draft``.""" + + source_binding_id = build_draft.agent_workspace_binding_id + if source_binding_id is None: + raise AgentBuildSandboxNotFoundError() + agent = session.scalar( + select(Agent).where( + Agent.id == build_draft.agent_id, + Agent.tenant_id == build_draft.tenant_id, + ) + ) + if agent is None: + raise AgentBuildSandboxNotFoundError() + from services.agent.roster_service import AgentRosterService + + runtime_app_id = AgentRosterService.runtime_backing_app_id(agent) + if runtime_app_id is None: + raise AgentBuildSandboxNotFoundError() + binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=build_draft.tenant_id, + binding_id=source_binding_id, + expected_owner_scope=WorkspaceOwnerScope( + tenant_id=build_draft.tenant_id, + app_id=runtime_app_id, + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id=build_draft.id, + ), + ) + if binding is None or binding.agent_id != build_draft.agent_id: + raise AgentBuildSandboxNotFoundError() + AgentWorkspaceService.validate_binding_generation( + binding, + base_home_snapshot_id=build_draft.home_snapshot_id, + agent_config_version_id=build_draft.id, + agent_config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) + + home_snapshot_id = str(uuidv7()) + try: + with cls._client() as client: + response = client.create_home_snapshot_from_binding_sync( + CreateHomeSnapshotFromBindingRequest( + tenant_id=build_draft.tenant_id, + agent_id=build_draft.agent_id, + home_snapshot_id=home_snapshot_id, + backend_binding_ref=binding.backend_binding_ref, + ) + ) + except DifyAgentNotFoundError as exc: + raise AgentBuildSandboxNotFoundError() from exc + + home_snapshot = AgentHomeSnapshot( + id=home_snapshot_id, + tenant_id=build_draft.tenant_id, + agent_id=build_draft.agent_id, + snapshot_ref=response.snapshot_ref, + status=AgentWorkingResourceStatus.ACTIVE, + ) + session.add(home_snapshot) + return home_snapshot + + @classmethod + def retire_all_for_agent(cls, *, session: Session, tenant_id: str, agent_id: str) -> list[str]: + rows = session.scalars( + select(AgentHomeSnapshot).where( + AgentHomeSnapshot.tenant_id == tenant_id, + AgentHomeSnapshot.agent_id == agent_id, + AgentHomeSnapshot.status == AgentWorkingResourceStatus.ACTIVE, + ) + ).all() + now = naive_utc_now() + for row in rows: + row.status = AgentWorkingResourceStatus.RETIRED + row.retired_at = now + return [row.id for row in rows] + + @classmethod + def collect_retired_home_snapshot(cls, *, tenant_id: str, home_snapshot_id: str) -> None: + try: + cls._collect_retired_home_snapshot(tenant_id=tenant_id, home_snapshot_id=home_snapshot_id) + except Exception: + logger.exception( + "Failed to collect retired Agent Home Snapshot", + extra={"tenant_id": tenant_id, "home_snapshot_id": home_snapshot_id}, + ) + + @classmethod + def _collect_retired_home_snapshot(cls, *, tenant_id: str, home_snapshot_id: str) -> None: + with session_factory.create_session() as session: + snapshot = session.scalar( + select(AgentHomeSnapshot).where( + AgentHomeSnapshot.id == home_snapshot_id, + AgentHomeSnapshot.tenant_id == tenant_id, + AgentHomeSnapshot.status == AgentWorkingResourceStatus.RETIRED, + ) + ) + if snapshot is None: + return + referenced = session.scalar( + select(AgentConfigDraft.id).where(AgentConfigDraft.home_snapshot_id == home_snapshot_id).limit(1) + ) or session.scalar( + select(AgentConfigSnapshot.id).where(AgentConfigSnapshot.home_snapshot_id == home_snapshot_id).limit(1) + ) + if referenced is not None: + return + snapshot_ref = snapshot.snapshot_ref + try: + cls.delete(snapshot_ref=snapshot_ref) + except Exception: + logger.exception( + "Failed to collect retired Agent Home Snapshot", + extra={"tenant_id": tenant_id, "home_snapshot_id": home_snapshot_id}, + ) + return + with session_factory.create_session() as session: + snapshot = session.scalar( + select(AgentHomeSnapshot).where( + AgentHomeSnapshot.id == home_snapshot_id, + AgentHomeSnapshot.tenant_id == tenant_id, + AgentHomeSnapshot.status == AgentWorkingResourceStatus.RETIRED, + ) + ) + if snapshot is not None: + session.delete(snapshot) + session.commit() + + @classmethod + def delete(cls, *, snapshot_ref: str) -> None: + with cls._client() as client: + client.delete_home_snapshot_sync(snapshot_ref) + + @staticmethod + def _client() -> Client: + base_url = dify_config.AGENT_BACKEND_BASE_URL + if not base_url: + raise AgentHomeSnapshotUnavailableError("Dify Agent backend is required for Home Snapshot operations") + return Client(base_url=base_url) + + +def validate_home_snapshot_binding(*, session: Session, agent: Agent, home_snapshot_id: str | None) -> None: + if home_snapshot_id is None: + return + _require_owned_home_snapshot(session=session, agent=agent, home_snapshot_id=home_snapshot_id) + + +def _require_owned_home_snapshot(*, session: Session, agent: Agent, home_snapshot_id: str) -> AgentHomeSnapshot: + if agent.status != AgentStatus.ACTIVE: + raise AgentHomeSnapshotUnavailableError(f"Agent {agent.id} is not active") + home_snapshot = session.scalar( + select(AgentHomeSnapshot).where( + AgentHomeSnapshot.id == home_snapshot_id, + AgentHomeSnapshot.tenant_id == agent.tenant_id, + AgentHomeSnapshot.agent_id == agent.id, + AgentHomeSnapshot.status == AgentWorkingResourceStatus.ACTIVE, + ) + ) + if home_snapshot is None: + raise AgentHomeSnapshotUnavailableError(f"Home Snapshot {home_snapshot_id} is unavailable for Agent {agent.id}") + return home_snapshot + + +__all__ = [ + "AgentHomeSnapshotService", + "AgentHomeSnapshotUnavailableError", + "validate_home_snapshot_binding", +] diff --git a/api/services/agent/observability_service.py b/api/services/agent/observability_service.py index cabb4c135e5..a8551e4ee27 100644 --- a/api/services/agent/observability_service.py +++ b/api/services/agent/observability_service.py @@ -1,7 +1,7 @@ from __future__ import annotations import json -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import datetime from decimal import Decimal @@ -17,8 +17,8 @@ from core.app.entities.app_invoke_entities import InvokeFrom from graphon.enums import WorkflowNodeExecutionStatus from libs.helper import convert_datetime_to_date, escape_like_pattern, to_timestamp from models.agent import WorkflowAgentNodeBinding -from models.enums import CreatorUserRole, MessageStatus -from models.model import App, Conversation, Message +from models.enums import CreatorUserRole, FeedbackFromSource, FeedbackRating, MessageStatus +from models.model import App, Conversation, Message, MessageFeedback from models.workflow import WorkflowNodeExecutionModel, WorkflowRun, WorkflowType @@ -138,7 +138,12 @@ class AgentObservabilityService: return int(message.message_tokens or 0) + int(message.answer_tokens or 0) @classmethod - def serialize_log_message(cls, message: Message, conversation: Conversation | None = None) -> dict[str, Any]: + def serialize_log_message( + cls, + message: Message, + conversation: Conversation | None = None, + feedbacks: Sequence[MessageFeedback] = (), + ) -> dict[str, Any]: invoke_from = message.invoke_from.value if message.invoke_from else None return { "id": message.id, @@ -153,6 +158,8 @@ class AgentObservabilityService: "from_source": message.from_source.value if message.from_source else None, "from_end_user_id": message.from_end_user_id, "from_account_id": message.from_account_id, + "feedback_enabled": True, + "feedbacks": [cls._serialize_message_feedback(feedback) for feedback in feedbacks], "message_tokens": int(message.message_tokens or 0), "answer_tokens": int(message.answer_tokens or 0), "total_tokens": cls._total_tokens(message), @@ -201,14 +208,19 @@ class AgentObservabilityService: rows: list[dict[str, Any]] = [] for source_filter in source_filters: if source_filter.kind in {"all", "webapp"}: + messages = self._list_webapp_messages( + app=app, + conversation_id=conversation_id, + params=params, + source_filter=source_filter, + ) + feedbacks_by_message = self._list_message_feedbacks(app=app, messages=messages) rows.extend( - self.serialize_log_message(message) - for message in self._list_webapp_messages( - app=app, - conversation_id=conversation_id, - params=params, - source_filter=source_filter, + self.serialize_log_message( + message, + feedbacks=feedbacks_by_message.get(message.id, ()), ) + for message in messages ) if source_filter.kind in {"all", "workflow"}: rows.extend( @@ -270,6 +282,10 @@ class AgentObservabilityService: ) stmt = self._apply_observability_filters(stmt, params=params, source_filter=source_filter) rows = list(self._session.execute(stmt).all()) + feedback_rates = self._list_conversation_feedback_rates( + app=app, + conversation_ids=[row[0].id for row in rows], + ) return [ self._serialize_conversation_log( conversation=row[0], @@ -279,6 +295,8 @@ class AgentObservabilityService: source=self._serialize_webapp_source(app), created_at=row.created_at, updated_at=row.updated_at, + user_rate=feedback_rates.get(row[0].id, {}).get("user_rate"), + operation_rate=feedback_rates.get(row[0].id, {}).get("operation_rate"), ) for row in rows ] @@ -350,6 +368,62 @@ class AgentObservabilityService: stmt = self._apply_message_filters(stmt, params=params, source_filter=source_filter) return list(self._session.scalars(stmt.order_by(Message.created_at.desc(), Message.id.desc())).all()) + def _list_message_feedbacks(self, *, app: App, messages: Sequence[Message]) -> dict[str, list[MessageFeedback]]: + message_ids = [message.id for message in messages] + if not message_ids: + return {} + + stmt = ( + select(MessageFeedback) + .where( + MessageFeedback.app_id == app.id, + MessageFeedback.message_id.in_(message_ids), + MessageFeedback.from_source.in_((FeedbackFromSource.USER, FeedbackFromSource.ADMIN)), + ) + .order_by(MessageFeedback.created_at.asc(), MessageFeedback.id.asc()) + ) + feedbacks_by_message: dict[str, list[MessageFeedback]] = {} + for feedback in self._session.scalars(stmt).all(): + feedbacks_by_message.setdefault(feedback.message_id, []).append(feedback) + return feedbacks_by_message + + def _list_conversation_feedback_rates( + self, *, app: App, conversation_ids: Sequence[str] + ) -> dict[str, dict[str, float]]: + if not conversation_ids: + return {} + + stmt = ( + select( + MessageFeedback.conversation_id, + MessageFeedback.from_source, + func.sum(sa.case((MessageFeedback.rating == FeedbackRating.LIKE, 1), else_=0)).label("like_count"), + func.count(MessageFeedback.id).label("total_count"), + ) + .where( + MessageFeedback.app_id == app.id, + MessageFeedback.conversation_id.in_(conversation_ids), + MessageFeedback.from_source.in_((FeedbackFromSource.USER, FeedbackFromSource.ADMIN)), + ) + .group_by(MessageFeedback.conversation_id, MessageFeedback.from_source) + ) + rates: dict[str, dict[str, float]] = {} + for row in self._session.execute(stmt).all(): + rate = self._positive_feedback_rate(like_count=row.like_count, total_count=row.total_count) + if rate is None: + continue + source = self._enum_value(row.from_source) + rate_key = "user_rate" if source == FeedbackFromSource.USER.value else "operation_rate" + rates.setdefault(row.conversation_id, {})[rate_key] = rate + return rates + + @staticmethod + def _positive_feedback_rate(*, like_count: int | None, total_count: int | None) -> float | None: + total = int(total_count or 0) + if total == 0: + return None + return int(like_count or 0) / total + def _list_workflow_messages( self, *, @@ -547,6 +621,8 @@ class AgentObservabilityService: source: dict[str, Any], created_at: datetime | None, updated_at: datetime | None, + user_rate: float | None, + operation_rate: float | None, ) -> dict[str, Any]: return { "id": conversation.id, @@ -554,8 +630,8 @@ class AgentObservabilityService: "title": conversation.name, "end_user_id": conversation.from_end_user_id, "message_count": int(message_count or 0), - "user_rate": None, - "operation_rate": None, + "user_rate": user_rate, + "operation_rate": operation_rate, "unread": conversation.read_at is None, "source": source, "status": cls._conversation_status(paused_count=paused_count, failed_count=failed_count), @@ -622,6 +698,8 @@ class AgentObservabilityService: "from_account_id": ( node_execution.created_by if created_by_role == CreatorUserRole.ACCOUNT.value else None ), + "feedback_enabled": False, + "feedbacks": [], "message_tokens": prompt_tokens, "answer_tokens": completion_tokens, "total_tokens": total_tokens, @@ -632,6 +710,14 @@ class AgentObservabilityService: "updated_at": to_timestamp(node_execution.finished_at or node_execution.created_at), } + @classmethod + def _serialize_message_feedback(cls, feedback: MessageFeedback) -> dict[str, Any]: + return { + "rating": cls._enum_value(feedback.rating), + "content": feedback.content, + "from_source": cls._enum_value(feedback.from_source), + } + @staticmethod def _json_mapping(value: object) -> Mapping[str, Any]: if isinstance(value, Mapping): diff --git a/api/services/agent/prompt_mentions.py b/api/services/agent/prompt_mentions.py index 5f42bffe3ff..a3690f4093a 100644 --- a/api/services/agent/prompt_mentions.py +++ b/api/services/agent/prompt_mentions.py @@ -320,7 +320,7 @@ def _format_output_mention(output: DeclaredOutputConfig) -> str: f"{output.name} (file output; create the file locally, run " f"`dify-agent file upload `, then set final_output.{output.name} to a `tool_file` mapping " f"using the returned `reference`; if replying to the user in natural language, use the returned " - f"`download_url`; do not call final_output before upload succeeds, and do not use the local path, " + f"`public_download_url`; do not call final_output before upload succeeds, and do not use the local path, " "filename, URL, or a synthesized dify-file-ref as the reference)" ) if ( @@ -332,7 +332,7 @@ def _format_output_mention(output: DeclaredOutputConfig) -> str: f"{output.name} (array[file] output; upload each produced file with " f"`dify-agent file upload `, then set final_output.{output.name} to `tool_file` mappings " f"using the returned `reference` values; if replying to the user in natural language, use the returned " - f"`download_url`; do not call final_output before all uploads succeed, and do not use local paths, " + f"`public_download_url`; do not call final_output before all uploads succeed, and do not use local paths, " "filenames, URLs, or synthesized dify-file-ref values as references)" ) return f"{output.name} ({output.type.value})" diff --git a/api/services/agent/retirement_service.py b/api/services/agent/retirement_service.py new file mode 100644 index 00000000000..25b133e14f9 --- /dev/null +++ b/api/services/agent/retirement_service.py @@ -0,0 +1,166 @@ +"""Workflow-only Agent ownership retirement after product transactions commit.""" + +from __future__ import annotations + +import logging +from collections.abc import Iterable + +from sqlalchemy import or_, select +from sqlalchemy.orm import Session + +from core.db.session_factory import session_factory +from libs.datetime_utils import naive_utc_now +from models.agent import ( + Agent, + AgentScope, + AgentStatus, + AgentWorkingResourceStatus, + AgentWorkspaceBinding, + WorkflowAgentNodeBinding, +) +from models.enums import AppStatus +from models.model import App +from models.workflow import Workflow +from services.agent.home_snapshot_service import AgentHomeSnapshotService +from services.agent.workspace_service import AgentWorkspaceService + +logger = logging.getLogger(__name__) + + +class WorkflowAgentRetirementService: + """Archive workflow-only Agents once no effective binding owns them.""" + + @classmethod + def retire_unowned( + cls, + *, + tenant_id: str, + agent_ids: Iterable[str], + account_id: str | None, + ) -> tuple[list[str], list[str]]: + """Re-check ownership, archive orphans, and commit their resource retirement.""" + + candidates = tuple(sorted({agent_id for agent_id in agent_ids if agent_id})) + if not candidates: + return [], [] + retired_bindings: list[str] = [] + retired_snapshots: list[str] = [] + try: + with session_factory.create_session() as session: + retired_agent_ids = cls.archive_unowned( + session=session, + tenant_id=tenant_id, + agent_ids=candidates, + account_id=account_id, + ) + for agent_id in retired_agent_ids: + bindings = session.scalars( + select(AgentWorkspaceBinding).where( + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.agent_id == agent_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + ) + ).all() + for binding in bindings: + binding_id = AgentWorkspaceService.retire_binding( + session=session, + tenant_id=tenant_id, + binding_id=binding.id, + ) + if binding_id is not None: + retired_bindings.append(binding_id) + retired_snapshots.extend( + AgentHomeSnapshotService.retire_all_for_agent( + session=session, + tenant_id=tenant_id, + agent_id=agent_id, + ) + ) + session.commit() + except Exception: + logger.exception( + "Failed to retire unowned Workflow Agents", + extra={ + "tenant_id": tenant_id, + "agent_ids": candidates, + }, + ) + return [], [] + return retired_bindings, retired_snapshots + + @classmethod + def archive_unowned( + cls, + *, + session: Session, + tenant_id: str, + agent_ids: Iterable[str], + account_id: str | None, + ) -> list[str]: + """Archive active orphans and return every orphan eligible for Home cleanup.""" + candidates = tuple(sorted({agent_id for agent_id in agent_ids if agent_id})) + if not candidates: + return [] + agents = session.scalars( + select(Agent).where( + Agent.tenant_id == tenant_id, + Agent.id.in_(candidates), + Agent.scope == AgentScope.WORKFLOW_ONLY, + Agent.status.in_((AgentStatus.ACTIVE, AgentStatus.ARCHIVED)), + ) + ).all() + effective_agent_ids = cls._effective_agent_ids( + session=session, + tenant_id=tenant_id, + agent_ids=[agent.id for agent in agents], + ) + now = naive_utc_now() + cleanup_candidates: list[str] = [] + for agent in agents: + if agent.id in effective_agent_ids: + continue + if agent.status == AgentStatus.ACTIVE: + agent.status = AgentStatus.ARCHIVED + agent.archived_by = account_id + agent.archived_at = now + agent.updated_by = account_id or agent.updated_by + agent.updated_at = now + cleanup_candidates.append(agent.id) + session.flush() + return cleanup_candidates + + @staticmethod + def _effective_agent_ids( + *, + session: Session, + tenant_id: str, + agent_ids: list[str], + ) -> set[str]: + if not agent_ids: + return set() + values = session.scalars( + select(WorkflowAgentNodeBinding.agent_id) + .join( + Workflow, + Workflow.id == WorkflowAgentNodeBinding.workflow_id, + ) + .join(App, App.id == WorkflowAgentNodeBinding.app_id) + .where( + WorkflowAgentNodeBinding.tenant_id == tenant_id, + WorkflowAgentNodeBinding.agent_id.in_(agent_ids), + Workflow.tenant_id == tenant_id, + Workflow.app_id == WorkflowAgentNodeBinding.app_id, + Workflow.version == WorkflowAgentNodeBinding.workflow_version, + App.tenant_id == tenant_id, + App.status == AppStatus.NORMAL, + or_( + Workflow.version == Workflow.VERSION_DRAFT, + App.workflow_id == Workflow.id, + ), + ) + .distinct() + ).all() + return {agent_id for agent_id in values if agent_id} + + +__all__ = ["WorkflowAgentRetirementService"] diff --git a/api/services/agent/roster_service.py b/api/services/agent/roster_service.py index c366f2af47b..af7bed03879 100644 --- a/api/services/agent/roster_service.py +++ b/api/services/agent/roster_service.py @@ -4,10 +4,8 @@ from typing import Any, TypedDict from sqlalchemy import and_, func, or_, select from sqlalchemy.exc import IntegrityError -from clients.agent_backend.session_cleanup import AgentBackendSessionCleanupPayload from constants.model_template import default_app_templates from core.agent.publish_visibility import workflow_callable_active_snapshot_filter -from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore from core.app.entities.app_invoke_entities import InvokeFrom from libs.datetime_utils import naive_utc_now from libs.helper import to_timestamp @@ -25,6 +23,9 @@ from models.agent import ( AgentScope, AgentSource, AgentStatus, + AgentWorkingResourceStatus, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, WorkflowAgentBindingType, WorkflowAgentNodeBinding, ) @@ -36,15 +37,18 @@ from services.agent.agent_soul_state import agent_soul_has_model from services.agent.composer_validator import ComposerConfigValidator from services.agent.errors import ( AgentArchivedError, + AgentBuildSandboxNotFoundError, AgentNameConflictError, AgentNotFoundError, AgentVersionNotFoundError, ) +from services.agent.home_snapshot_service import AgentHomeSnapshotService +from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService, WorkspaceOwnerScope from services.app_service import AppService, CreateAppParams from services.enterprise.enterprise_service import EnterpriseService from services.entities.agent_entities import RosterAgentCreatePayload, RosterAgentUpdatePayload from services.feature_service import FeatureService -from tasks.agent_backend_session_cleanup_task import cleanup_conversation_agent_runtime_session +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection logger = logging.getLogger(__name__) @@ -304,6 +308,27 @@ class AgentRosterService: account_id: str, payload: RosterAgentCreatePayload, source: AgentSource = AgentSource.ROSTER, + ) -> Agent: + try: + agent = self._create_roster_agent_in_transaction( + tenant_id=tenant_id, + account_id=account_id, + payload=payload, + source=source, + ) + self._session.commit() + return agent + except IntegrityError as exc: + self._session.rollback() + raise AgentNameConflictError() from exc + + def _create_roster_agent_in_transaction( + self, + *, + tenant_id: str, + account_id: str, + payload: RosterAgentCreatePayload, + source: AgentSource, ) -> Agent: ComposerConfigValidator.validate_agent_soul(payload.agent_soul) @@ -323,17 +348,14 @@ class AgentRosterService: updated_by=account_id, ) self._session.add(agent) - try: - self._session.flush() - except IntegrityError as exc: - self._session.rollback() - raise AgentNameConflictError() from exc + self._session.flush() version = AgentConfigSnapshot( tenant_id=tenant_id, agent_id=agent.id, version=1, config_snapshot=payload.agent_soul, + home_snapshot_id=None, version_note=payload.version_note, created_by=account_id, ) @@ -354,11 +376,6 @@ class AgentRosterService: agent.active_config_has_model = agent_soul_has_model(payload.agent_soul) agent.active_config_is_published = True - try: - self._session.commit() - except IntegrityError as exc: - self._session.rollback() - raise AgentNameConflictError() from exc return agent def create_backing_agent_for_app( @@ -405,17 +422,14 @@ class AgentRosterService: updated_by=account_id, ) self._session.add(agent) - try: - self._session.flush() - except IntegrityError as exc: - self._session.rollback() - raise AgentNameConflictError() from exc + self._session.flush() version = AgentConfigSnapshot( tenant_id=tenant_id, agent_id=agent.id, version=1, config_snapshot=soul, + home_snapshot_id=None, created_by=account_id, ) self._session.add(version) @@ -590,16 +604,15 @@ class AgentRosterService: self._session.flush() return conversation_id - def get_or_create_agent_app_debug_conversation_id( + def get_or_create_build_conversation( 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 Build or Preview conversation for an Agent App.""" + """Return the current editor's stable Build conversation.""" agent = self._session.scalar( select(Agent).where( @@ -614,21 +627,20 @@ class AgentRosterService: conversation_id = self._get_or_create_agent_app_debug_conversation( agent=agent, account_id=account_id, - draft_type=draft_type, + draft_type=AgentConfigDraftType.DEBUG_BUILD, ) if commit: self._session.commit() return conversation_id - def load_agent_app_debug_conversation_id( + def get_current_preview_conversation( 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 the editor's current Preview conversation without creating one.""" return self._session.scalar( select(Conversation.id) @@ -637,7 +649,7 @@ class AgentRosterService: AgentDebugConversation.tenant_id == tenant_id, AgentDebugConversation.agent_id == agent_id, AgentDebugConversation.account_id == account_id, - AgentDebugConversation.draft_type == draft_type, + AgentDebugConversation.draft_type == AgentConfigDraftType.DRAFT, AgentDebugConversation.app_id == Conversation.app_id, Conversation.from_source == ConversationFromSource.CONSOLE, Conversation.from_account_id == account_id, @@ -657,25 +669,12 @@ class AgentRosterService: or 0 ) - def refresh_agent_app_debug_conversation_id( - self, - *, - tenant_id: str, - agent_id: str, - account_id: str, - draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD, - ) -> str: - """Start a new scoped console conversation for the current Agent App editor. + def rotate_preview_conversation(self, *, tenant_id: str, agent_id: str, account_id: str) -> str: + """Rotate Preview and retire its exact Conversation-owned Binding. - If this account already has a mapping for the requested draft surface, the previous - conversation is abandoned after the replacement mapping is committed: 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. This - order prevents a failed database commit from retiring the still-current runtime session. - The other draft surface is left untouched. - - A user and draft surface own one current mapping. If new-conversation requests overlap, - the last committed rotation becomes current and earlier response IDs cannot be continued. + The mapping update and exact CONVERSATION Binding retirement commit in + one transaction. Validation failures fail fast; collection is enqueued + only after commit. """ agent = self._session.scalar( @@ -694,142 +693,194 @@ class AgentRosterService: if not backing_app_id: raise AgentNotFoundError() - conversation_id = self._create_agent_app_debug_conversation( - app_id=backing_app_id, - account_id=account_id, - ) - previous_conversation: tuple[str, str] | None = None - mapping = self._session.scalar( - select(AgentDebugConversation).where( - AgentDebugConversation.tenant_id == tenant_id, - AgentDebugConversation.agent_id == agent_id, - AgentDebugConversation.account_id == account_id, - AgentDebugConversation.draft_type == draft_type, + retired_binding_id: str | None = None + try: + conversation_id = self._create_agent_app_debug_conversation( + app_id=backing_app_id, + account_id=account_id, ) - ) - if mapping is None: - self._session.add( - AgentDebugConversation( - tenant_id=tenant_id, - agent_id=agent_id, - app_id=backing_app_id, - account_id=account_id, - draft_type=draft_type, - conversation_id=conversation_id, + mapping = self._session.scalar( + select(AgentDebugConversation).where( + AgentDebugConversation.tenant_id == tenant_id, + AgentDebugConversation.agent_id == agent_id, + AgentDebugConversation.account_id == account_id, + AgentDebugConversation.draft_type == AgentConfigDraftType.DRAFT, ) ) - else: - previous_app_id = mapping.app_id - previous_conversation_id = mapping.conversation_id - if previous_conversation_id: - previous_conversation = (previous_app_id or backing_app_id, previous_conversation_id) - mapping.app_id = backing_app_id - mapping.conversation_id = conversation_id - self._session.flush() - self._session.commit() - - if previous_conversation: - previous_app_id, previous_conversation_id = previous_conversation - self._cleanup_debug_conversation_runtime_sessions( + if mapping is None: + self._session.add( + AgentDebugConversation( + tenant_id=tenant_id, + agent_id=agent_id, + app_id=backing_app_id, + account_id=account_id, + draft_type=AgentConfigDraftType.DRAFT, + conversation_id=conversation_id, + ) + ) + else: + previous_app_id = mapping.app_id or backing_app_id + previous_conversation_id = mapping.conversation_id + if previous_conversation_id: + previous_conversation = self._session.scalar( + select(Conversation).where( + Conversation.id == previous_conversation_id, + Conversation.app_id == previous_app_id, + Conversation.from_source == ConversationFromSource.CONSOLE, + Conversation.from_account_id == account_id, + Conversation.is_deleted.is_(False), + ) + ) + if ( + previous_conversation is not None + and previous_conversation.agent_workspace_binding_id is not None + ): + binding_id = previous_conversation.agent_workspace_binding_id + binding = AgentWorkspaceService.get_active_binding( + session=self._session, + tenant_id=tenant_id, + binding_id=binding_id, + expected_owner_scope=WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=previous_app_id, + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id=previous_conversation.id, + ), + ) + if binding is None or binding.agent_id != agent_id: + raise AgentWorkspaceNotFoundError( + "Agent debug Conversation participant Binding is unavailable" + ) + retired_binding_id = AgentWorkspaceService.retire_binding( + session=self._session, + tenant_id=tenant_id, + binding_id=binding_id, + ) + if retired_binding_id is None: + raise AgentWorkspaceNotFoundError( + "Agent debug Conversation participant Binding is unavailable" + ) + mapping.app_id = backing_app_id + mapping.conversation_id = conversation_id + self._session.flush() + self._session.commit() + except Exception: + self._session.rollback() + raise + if retired_binding_id is not None: + enqueue_agent_resource_collection( tenant_id=tenant_id, - agent_id=agent_id, - account_id=account_id, - draft_type=draft_type, - app_id=previous_app_id, - conversation_id=previous_conversation_id, + binding_ids=(retired_binding_id,), ) return conversation_id - def _cleanup_debug_conversation_runtime_sessions( - self, - *, - tenant_id: str, - agent_id: str, - account_id: str, - draft_type: AgentConfigDraftType, - app_id: str, - conversation_id: str, - ) -> None: + def reset_build_conversation(self, *, tenant_id: str, agent_id: str, account_id: str) -> str: + """Reset Build and retire its exact DEBUG_BUILD Draft-owned Binding. + + The mapping update, exact BUILD_DRAFT Binding retirement, and Draft + pointer clear commit in one transaction. Validation failures fail fast; + collection is enqueued only after commit. + """ + + agent = self._session.scalar( + select(Agent).where( + Agent.tenant_id == tenant_id, + Agent.id == agent_id, + Agent.status == AgentStatus.ACTIVE, + ) + ) + if agent is None: + raise AgentNotFoundError() + backing_app_id = self._ensure_workflow_agent_backing_app( + agent=agent, + account_id=agent.updated_by or agent.created_by, + ) + if not backing_app_id: + raise AgentNotFoundError() + + retired_binding_id: str | None = None try: - session_store = AgentAppRuntimeSessionStore() - stored_sessions = session_store.list_active_sessions_for_conversation( - tenant_id=tenant_id, - app_id=app_id, - conversation_id=conversation_id, + conversation_id = self._create_agent_app_debug_conversation( + app_id=backing_app_id, + account_id=account_id, ) - except Exception: - logger.warning( - "Failed to load Agent App runtime sessions for debug conversation refresh: " - "tenant_id=%s app_id=%s conversation_id=%s", - tenant_id, - app_id, - conversation_id, - exc_info=True, - ) - return - - for stored_session in stored_sessions: - try: - if stored_session.runtime_layer_specs: - payload = AgentBackendSessionCleanupPayload( - session_snapshot=stored_session.session_snapshot, - runtime_layer_specs=stored_session.runtime_layer_specs, - idempotency_key=( - f"{tenant_id}:{agent_id}:{account_id}:{draft_type.value}:{conversation_id}:" - "debug-session-cleanup:" - f"{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'}" - ), - metadata={ - "tenant_id": stored_session.scope.tenant_id, - "app_id": stored_session.scope.app_id, - "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, - }, - ) - cleanup_conversation_agent_runtime_session.delay(payload.model_dump(mode="json")) - except Exception: - logger.warning( - "Failed to enqueue Agent backend cleanup for debug conversation refresh: " - "tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s", - stored_session.scope.tenant_id, - stored_session.scope.app_id, - stored_session.scope.conversation_id, - stored_session.scope.agent_id, - stored_session.backend_run_id, - exc_info=True, + mapping = self._session.scalar( + select(AgentDebugConversation).where( + AgentDebugConversation.tenant_id == tenant_id, + AgentDebugConversation.agent_id == agent_id, + AgentDebugConversation.account_id == account_id, + AgentDebugConversation.draft_type == AgentConfigDraftType.DEBUG_BUILD, ) - finally: - try: - session_store.mark_cleaned( - scope=stored_session.scope, - backend_run_id=stored_session.backend_run_id, - ) - except Exception: - logger.warning( - "Failed to retire Agent App runtime session for debug conversation refresh: " - "tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s", - stored_session.scope.tenant_id, - stored_session.scope.app_id, - stored_session.scope.conversation_id, - stored_session.scope.agent_id, - stored_session.backend_run_id, - exc_info=True, - ) + ) + build_draft = self._session.scalar( + select(AgentConfigDraft) + .where( + AgentConfigDraft.tenant_id == tenant_id, + AgentConfigDraft.agent_id == agent_id, + AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD, + AgentConfigDraft.account_id == account_id, + ) + .order_by(AgentConfigDraft.updated_at.desc()) + .limit(1) + ) + if build_draft is not None and build_draft.agent_workspace_binding_id is not None: + binding_id = build_draft.agent_workspace_binding_id + binding = AgentWorkspaceService.get_active_binding( + session=self._session, + tenant_id=tenant_id, + binding_id=binding_id, + expected_owner_scope=WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=backing_app_id, + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id=build_draft.id, + ), + ) + if binding is None or binding.agent_id != agent_id: + raise AgentBuildSandboxNotFoundError() + retired_binding_id = AgentWorkspaceService.retire_binding( + session=self._session, + tenant_id=tenant_id, + binding_id=binding_id, + ) + if retired_binding_id is None: + raise AgentBuildSandboxNotFoundError() + build_draft.agent_workspace_binding_id = None - def load_or_create_agent_app_debug_conversation_ids_by_agent_id( + if mapping is None: + self._session.add( + AgentDebugConversation( + tenant_id=tenant_id, + agent_id=agent_id, + app_id=backing_app_id, + account_id=account_id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + conversation_id=conversation_id, + ) + ) + else: + mapping.app_id = backing_app_id + mapping.conversation_id = conversation_id + self._session.flush() + self._session.commit() + except Exception: + self._session.rollback() + raise + if retired_binding_id is not None: + enqueue_agent_resource_collection( + tenant_id=tenant_id, + binding_ids=(retired_binding_id,), + ) + return conversation_id + + def load_or_create_build_conversation_ids_by_agent_id( self, *, tenant_id: str, agents: list[Agent], account_id: str, - draft_type: AgentConfigDraftType = AgentConfigDraftType.DEBUG_BUILD, ) -> dict[str, str]: - """Return per-account scoped conversations for a page of Agent Apps.""" + """Return per-account Build conversations for a page of Agent Apps.""" conversation_ids_by_agent_id: dict[str, str] = {} changed = False @@ -839,7 +890,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, + draft_type=AgentConfigDraftType.DEBUG_BUILD, ) changed = True if changed: @@ -1044,8 +1095,10 @@ class AgentRosterService: session=self._session, ) - target_app.enable_site = source_app.enable_site - target_app.enable_api = source_app.enable_api + # A copy owns a new publication history. It remains private until its + # first successful publish even when the source Agent is public. + target_app.enable_site = False + target_app.enable_api = False target_app.use_icon_as_answer_icon = source_app.use_icon_as_answer_icon target_app.tracing = source_app.tracing @@ -1057,7 +1110,6 @@ class AgentRosterService: account_id=account.id, ) self._session.commit() - if FeatureService.get_system_features().webapp_auth.enabled: try: original_settings = EnterpriseService.WebAppAuth.get_app_access_mode_by_id(source_app.id) @@ -1114,7 +1166,7 @@ class AgentRosterService: target_version.version_note = source_version.version_note target_version.created_by = account_id target_agent.active_config_has_model = agent_soul_has_model(target_version.config_snapshot) - target_agent.active_config_is_published = source_agent.active_config_is_published + target_agent.active_config_is_published = False target_agent.updated_by = account_id def _next_duplicate_agent_name(self, *, tenant_id: str, base_name: str) -> str: @@ -1211,13 +1263,38 @@ class AgentRosterService: def archive_roster_agent(self, *, tenant_id: str, agent_id: str, account_id: str) -> None: agent = self._get_agent(tenant_id=tenant_id, agent_id=agent_id, roster_only=True) - if agent.status == AgentStatus.ARCHIVED: - return - agent.status = AgentStatus.ARCHIVED - agent.archived_by = account_id - agent.archived_at = naive_utc_now() - agent.updated_by = account_id + retired_binding_ids: list[str] = [] + if agent.status != AgentStatus.ARCHIVED: + agent.status = AgentStatus.ARCHIVED + agent.archived_by = account_id + agent.archived_at = naive_utc_now() + agent.updated_by = account_id + bindings = self._session.scalars( + select(AgentWorkspaceBinding).where( + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.agent_id == agent_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + ) + ).all() + for binding in bindings: + retired_id = AgentWorkspaceService.retire_binding( + session=self._session, + tenant_id=tenant_id, + binding_id=binding.id, + ) + if retired_id is not None: + retired_binding_ids.append(retired_id) + retired_snapshot_ids = AgentHomeSnapshotService.retire_all_for_agent( + session=self._session, + tenant_id=tenant_id, + agent_id=agent_id, + ) self._session.commit() + enqueue_agent_resource_collection( + tenant_id=tenant_id, + binding_ids=retired_binding_ids, + home_snapshot_ids=retired_snapshot_ids, + ) @staticmethod def _visible_version_operations(agent: Agent) -> set[AgentConfigRevisionOperation]: @@ -1382,9 +1459,11 @@ class AgentRosterService: account_id=None, draft_owner_key="", created_by=account_id, + home_snapshot_id=version.home_snapshot_id, ) self._session.add(draft) draft.base_snapshot_id = version.id + draft.home_snapshot_id = version.home_snapshot_id draft.config_snapshot = AgentSoulConfig.model_validate(version.config_snapshot_dict) draft.updated_by = account_id agent.active_config_is_published = version.id == agent.active_config_snapshot_id diff --git a/api/services/agent/skill_package_service.py b/api/services/agent/skill_package_service.py index d5601987abc..fbfd2ababfc 100644 --- a/api/services/agent/skill_package_service.py +++ b/api/services/agent/skill_package_service.py @@ -26,8 +26,9 @@ import zlib import yaml from pydantic import BaseModel +from configs import dify_config + # Bounds — generous but finite so a hostile upload can't exhaust memory/disk. -_MAX_ARCHIVE_BYTES = 50 * 1024 * 1024 _MAX_UNCOMPRESSED_BYTES = 200 * 1024 * 1024 _MAX_SKILL_MD_BYTES = 1 * 1024 * 1024 _MAX_ENTRIES = 5000 @@ -127,7 +128,8 @@ class SkillPackageService: self._check_extension(filename) if not content: raise SkillPackageError("empty_archive", "skill archive is empty", status_code=400) - if len(content) > _MAX_ARCHIVE_BYTES: + max_archive_bytes = dify_config.UPLOAD_SKILL_FILE_SIZE_LIMIT * 1024 * 1024 + if len(content) > max_archive_bytes: raise SkillPackageError("archive_too_large", "skill archive exceeds size limit", status_code=400) try: diff --git a/api/services/agent/workflow_publish_service.py b/api/services/agent/workflow_publish_service.py index e74a93a899b..b43f3091f88 100644 --- a/api/services/agent/workflow_publish_service.py +++ b/api/services/agent/workflow_publish_service.py @@ -24,6 +24,7 @@ from models.agent_config_entities import ( WorkflowNodeJobConfig, WorkflowPreviousNodeOutputRef, ) +from models.model import App from models.workflow import Workflow from services.agent.composer_validator import ComposerConfigValidator from services.agent.prompt_mentions import ( @@ -224,7 +225,7 @@ class WorkflowAgentPublishService: session: Session, draft_workflow: Workflow, account_id: str, - ) -> None: + ) -> set[str]: agent_nodes = dict(WorkflowAgentNodeValidator.iter_agent_v2_nodes(draft_workflow.graph_dict)) existing_bindings = list( session.scalars( @@ -237,9 +238,12 @@ class WorkflowAgentPublishService: ).all() ) existing_by_node_id = {binding.node_id: binding for binding in existing_bindings} + retirement_candidates: set[str] = set() for binding in existing_bindings: if binding.node_id not in agent_nodes: + if binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT and binding.agent_id: + retirement_candidates.add(binding.agent_id) session.delete(binding) for node_id, node_data in agent_nodes.items(): @@ -252,16 +256,34 @@ class WorkflowAgentPublishService: not binding_payload.get("agent_id") or not binding_payload.get("current_snapshot_id") ): continue + existing_binding = existing_by_node_id.get(node_id) + replaced_inline_agent_id = ( + existing_binding.agent_id + if existing_binding is not None + and existing_binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT + and existing_binding.agent_id + else None + ) cls._sync_agent_binding_for_node( session=session, draft_workflow=draft_workflow, node_id=node_id, node_data=node_data, node_binding=binding_payload, - existing_binding=existing_by_node_id.get(node_id), + existing_binding=existing_binding, account_id=account_id, ) + if ( + replaced_inline_agent_id + and existing_binding is not None + and ( + existing_binding.binding_type != WorkflowAgentBindingType.INLINE_AGENT + or existing_binding.agent_id != replaced_inline_agent_id + ) + ): + retirement_candidates.add(replaced_inline_agent_id) session.flush() + return retirement_candidates @classmethod def sync_roster_agent_bindings_for_draft( @@ -270,8 +292,8 @@ class WorkflowAgentPublishService: session: Session, draft_workflow: Workflow, account_id: str, - ) -> None: - cls.sync_agent_bindings_for_draft( + ) -> set[str]: + return cls.sync_agent_bindings_for_draft( session=session, draft_workflow=draft_workflow, account_id=account_id, @@ -561,12 +583,32 @@ class WorkflowAgentPublishService: session: Session, draft_workflow: Workflow, published_workflow: Workflow, - ) -> None: + ) -> set[str]: + current_workflow_id = session.scalar( + select(App.workflow_id).where( + App.tenant_id == draft_workflow.tenant_id, + App.id == draft_workflow.app_id, + ) + ) + retirement_candidates: set[str] = set() + if current_workflow_id: + retirement_candidates = { + agent_id + for agent_id in session.scalars( + select(WorkflowAgentNodeBinding.agent_id).where( + WorkflowAgentNodeBinding.tenant_id == draft_workflow.tenant_id, + WorkflowAgentNodeBinding.app_id == draft_workflow.app_id, + WorkflowAgentNodeBinding.workflow_id == current_workflow_id, + WorkflowAgentNodeBinding.binding_type == WorkflowAgentBindingType.INLINE_AGENT, + ) + ).all() + if agent_id + } node_ids = { node_id for node_id, _node_data in WorkflowAgentNodeValidator.iter_agent_v2_nodes(draft_workflow.graph_dict) } if not node_ids: - return + return retirement_candidates bindings = session.scalars( select(WorkflowAgentNodeBinding).where( @@ -578,7 +620,7 @@ class WorkflowAgentPublishService: ) ).all() if not bindings: - return + return retirement_candidates agents_by_id = { agent.id: agent @@ -611,6 +653,7 @@ class WorkflowAgentPublishService: updated_by=binding.updated_by, ) session.add(copied) + return retirement_candidates @classmethod def restore_agent_node_bindings_to_draft( @@ -620,7 +663,7 @@ class WorkflowAgentPublishService: source_workflow: Workflow, draft_workflow: Workflow, account_id: str, - ) -> None: + ) -> set[str]: """Replace draft bindings with the frozen bindings of a published workflow.""" existing = session.scalars( @@ -631,6 +674,11 @@ class WorkflowAgentPublishService: WorkflowAgentNodeBinding.workflow_version == cls._DRAFT_WORKFLOW_VERSION, ) ).all() + retirement_candidates = { + binding.agent_id + for binding in existing + if binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT and binding.agent_id + } for binding in existing: session.delete(binding) @@ -681,3 +729,4 @@ class WorkflowAgentPublishService: ) ) session.flush() + return retirement_candidates diff --git a/api/services/agent/workspace_service.py b/api/services/agent/workspace_service.py new file mode 100644 index 00000000000..574054238c0 --- /dev/null +++ b/api/services/agent/workspace_service.py @@ -0,0 +1,472 @@ +"""Own Workspace and AgentWorkspaceBinding product lifecycle. + +Dify API is the lifecycle ledger. Dify Agent only executes physical create, +acquire, and destroy operations selected by this service. Retire methods only +mutate the caller's transaction; collection performs network I/O after commit +and deletes ledger rows only after idempotent physical cleanup succeeds. +""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass + +from dify_agent.client import Client +from dify_agent.protocol import CreateExecutionBindingRequest, DestroyExecutionBindingRequest +from sqlalchemy import select +from sqlalchemy.orm import Session + +from configs import dify_config +from core.db.session_factory import session_factory +from libs.datetime_utils import naive_utc_now +from libs.uuid_utils import uuidv7 +from models.agent import ( + AgentConfigVersionKind, + AgentHomeSnapshot, + AgentWorkingResourceStatus, + AgentWorkspace, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, +) + +logger = logging.getLogger(__name__) + + +class AgentWorkspaceError(RuntimeError): + pass + + +class AgentWorkspaceNotFoundError(AgentWorkspaceError): + pass + + +class AgentWorkspaceBindingGenerationMismatchError(AgentWorkspaceError): + pass + + +@dataclass(frozen=True, slots=True) +class WorkspaceOwnerScope: + tenant_id: str + app_id: str + owner_type: AgentWorkspaceOwnerType + owner_id: str + owner_scope_key: str = "root" + + +class AgentWorkspaceService: + """Allocate and manage working-environment resources. + + A Binding ID is the participant identity. Product callers persist that ID + and use :meth:`get_active_binding`; Agent and Workspace attributes are not + participant lookup keys. + """ + + @classmethod + def resolve_active_workspace(cls, *, session: Session, scope: WorkspaceOwnerScope) -> AgentWorkspace | None: + return session.scalar( + select(AgentWorkspace).where( + AgentWorkspace.tenant_id == scope.tenant_id, + AgentWorkspace.app_id == scope.app_id, + AgentWorkspace.owner_type == scope.owner_type, + AgentWorkspace.owner_id == scope.owner_id, + AgentWorkspace.owner_scope_key == scope.owner_scope_key, + AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE, + ) + ) + + @classmethod + def get_active_binding( + cls, + *, + session: Session, + tenant_id: str, + binding_id: str, + expected_owner_scope: WorkspaceOwnerScope, + ) -> AgentWorkspaceBinding | None: + return session.scalar( + select(AgentWorkspaceBinding) + .join( + AgentWorkspace, + (AgentWorkspace.tenant_id == AgentWorkspaceBinding.tenant_id) + & (AgentWorkspace.id == AgentWorkspaceBinding.workspace_id), + ) + .where( + AgentWorkspaceBinding.id == binding_id, + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + AgentWorkspace.tenant_id == expected_owner_scope.tenant_id, + AgentWorkspace.app_id == expected_owner_scope.app_id, + AgentWorkspace.owner_type == expected_owner_scope.owner_type, + AgentWorkspace.owner_id == expected_owner_scope.owner_id, + AgentWorkspace.owner_scope_key == expected_owner_scope.owner_scope_key, + AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE, + ) + ) + + @classmethod + def create_binding( + cls, + *, + session: Session, + scope: WorkspaceOwnerScope, + agent_id: str, + base_home_snapshot_id: str | None, + agent_config_version_id: str, + agent_config_version_kind: AgentConfigVersionKind, + ) -> AgentWorkspaceBinding: + """Allocate one new participant in the caller-owned transaction. + + After backend creation returns successfully, any later Python, flush, + or commit failure may leave an orphan. Dify API does not perform + cross-system compensation; a future global reconciler is responsible + for those orphans. Backend-local cleanup applies only when creation + fails before the backend returns success. + """ + + home_snapshot_ref: str | None = None + if base_home_snapshot_id is not None: + home_snapshot = session.scalar( + select(AgentHomeSnapshot).where( + AgentHomeSnapshot.id == base_home_snapshot_id, + AgentHomeSnapshot.tenant_id == scope.tenant_id, + AgentHomeSnapshot.agent_id == agent_id, + AgentHomeSnapshot.status == AgentWorkingResourceStatus.ACTIVE, + ) + ) + if home_snapshot is None: + raise AgentWorkspaceNotFoundError("base Home Snapshot is unavailable") + home_snapshot_ref = home_snapshot.snapshot_ref + workspace = cls.resolve_active_workspace(session=session, scope=scope) + workspace_id = workspace.id if workspace is not None else str(uuidv7()) + binding_id = str(uuidv7()) + with cls._client() as client: + allocation = client.create_execution_binding_sync( + CreateExecutionBindingRequest( + tenant_id=scope.tenant_id, + agent_id=agent_id, + binding_id=binding_id, + workspace_id=workspace_id, + existing_workspace_ref=workspace.backend_workspace_ref if workspace is not None else None, + home_snapshot_ref=home_snapshot_ref, + ) + ) + if workspace is not None and allocation.workspace_ref != workspace.backend_workspace_ref: + raise AgentWorkspaceError("backend changed the existing Workspace ref") + if workspace is None: + workspace = AgentWorkspace( + id=workspace_id, + tenant_id=scope.tenant_id, + app_id=scope.app_id, + owner_type=scope.owner_type, + owner_id=scope.owner_id, + owner_scope_key=scope.owner_scope_key, + backend_workspace_ref=allocation.workspace_ref, + status=AgentWorkingResourceStatus.ACTIVE, + active_guard=1, + ) + session.add(workspace) + binding = AgentWorkspaceBinding( + id=binding_id, + tenant_id=scope.tenant_id, + app_id=scope.app_id, + workspace_id=workspace_id, + agent_id=agent_id, + base_home_snapshot_id=base_home_snapshot_id, + agent_config_version_id=agent_config_version_id, + agent_config_version_kind=agent_config_version_kind, + backend_binding_ref=allocation.binding_ref, + status=AgentWorkingResourceStatus.ACTIVE, + ) + session.add(binding) + return binding + + @classmethod + def save_binding_session_snapshot( + cls, + *, + tenant_id: str, + binding_id: str, + session_snapshot: str, + pending_form_id: str | None = None, + pending_tool_call_id: str | None = None, + ) -> None: + with session_factory.create_session() as session: + binding = session.scalar( + select(AgentWorkspaceBinding).where( + AgentWorkspaceBinding.id == binding_id, + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + ) + ) + if binding is None: + raise AgentWorkspaceNotFoundError("ACTIVE Binding is unavailable") + binding.session_snapshot = session_snapshot + binding.pending_form_id = pending_form_id + binding.pending_tool_call_id = pending_tool_call_id + session.commit() + + @classmethod + def retire_binding(cls, *, session: Session, tenant_id: str, binding_id: str) -> str | None: + binding = session.scalar( + select(AgentWorkspaceBinding) + .where( + AgentWorkspaceBinding.id == binding_id, + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + ) + .with_for_update() + ) + if binding is None: + return None + workspace = session.scalar( + select(AgentWorkspace) + .where( + AgentWorkspace.id == binding.workspace_id, + AgentWorkspace.tenant_id == tenant_id, + AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE, + ) + .with_for_update() + ) + now = naive_utc_now() + binding.status = AgentWorkingResourceStatus.RETIRED + binding.retired_at = now + if workspace is not None: + other_binding = session.scalar( + select(AgentWorkspaceBinding.id).where( + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.workspace_id == workspace.id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + AgentWorkspaceBinding.id != binding.id, + ) + ) + if other_binding is None: + workspace.status = AgentWorkingResourceStatus.RETIRED + workspace.active_guard = None + workspace.retired_at = now + return binding.id + + @classmethod + def retire_workspace(cls, *, session: Session, tenant_id: str, workspace_id: str) -> str | None: + workspace = session.scalar( + select(AgentWorkspace) + .where( + AgentWorkspace.id == workspace_id, + AgentWorkspace.tenant_id == tenant_id, + AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE, + ) + .with_for_update() + ) + if workspace is None: + return None + now = naive_utc_now() + workspace.status = AgentWorkingResourceStatus.RETIRED + workspace.active_guard = None + workspace.retired_at = now + bindings = session.scalars( + select(AgentWorkspaceBinding).where( + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.workspace_id == workspace.id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + ) + ).all() + for binding in bindings: + binding.status = AgentWorkingResourceStatus.RETIRED + binding.retired_at = now + return workspace.id + + @classmethod + def retire_all_for_app(cls, *, session: Session, tenant_id: str, app_id: str) -> list[str]: + """Retire all ACTIVE Workspaces owned by an App in the caller's transaction.""" + + workspaces = session.scalars( + select(AgentWorkspace).where( + AgentWorkspace.tenant_id == tenant_id, + AgentWorkspace.app_id == app_id, + AgentWorkspace.status == AgentWorkingResourceStatus.ACTIVE, + ) + ).all() + retired: list[str] = [] + for workspace in workspaces: + workspace_id = cls.retire_workspace( + session=session, + tenant_id=tenant_id, + workspace_id=workspace.id, + ) + if workspace_id is not None: + retired.append(workspace_id) + return retired + + @classmethod + def collect_retired_binding(cls, *, tenant_id: str, binding_id: str) -> None: + try: + cls._collect_retired_binding(tenant_id=tenant_id, binding_id=binding_id) + except Exception: + logger.exception( + "Failed to collect retired Agent Workspace Binding", + extra={"tenant_id": tenant_id, "binding_id": binding_id}, + ) + + @classmethod + def _collect_retired_binding(cls, *, tenant_id: str, binding_id: str) -> None: + with session_factory.create_session() as session: + binding = session.scalar( + select(AgentWorkspaceBinding).where( + AgentWorkspaceBinding.id == binding_id, + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.RETIRED, + ) + ) + if binding is None: + return + backend_binding_ref = binding.backend_binding_ref + workspace = session.scalar( + select(AgentWorkspace).where( + AgentWorkspace.id == binding.workspace_id, + AgentWorkspace.tenant_id == tenant_id, + ) + ) + if workspace is not None and workspace.status == AgentWorkingResourceStatus.RETIRED: + workspace_id = workspace.id + else: + workspace_id = None + if workspace_id is not None: + cls.collect_retired_workspace(tenant_id=tenant_id, workspace_id=workspace_id) + return + try: + with cls._client() as client: + client.destroy_execution_binding_sync( + DestroyExecutionBindingRequest( + binding_ref=backend_binding_ref, + destroy_workspace=False, + ) + ) + except Exception: + logger.exception( + "Failed to collect retired Agent Workspace Binding", + extra={"tenant_id": tenant_id, "binding_id": binding_id}, + ) + return + with session_factory.create_session() as session: + binding = session.scalar( + select(AgentWorkspaceBinding).where( + AgentWorkspaceBinding.id == binding_id, + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.RETIRED, + ) + ) + if binding is not None: + session.delete(binding) + session.commit() + + @classmethod + def collect_retired_workspace(cls, *, tenant_id: str, workspace_id: str) -> None: + try: + cls._collect_retired_workspace(tenant_id=tenant_id, workspace_id=workspace_id) + except Exception: + logger.exception( + "Failed to collect retired Agent Workspace", + extra={"tenant_id": tenant_id, "workspace_id": workspace_id}, + ) + + @classmethod + def _collect_retired_workspace(cls, *, tenant_id: str, workspace_id: str) -> None: + with session_factory.create_session() as session: + workspace = session.scalar( + select(AgentWorkspace).where( + AgentWorkspace.id == workspace_id, + AgentWorkspace.tenant_id == tenant_id, + AgentWorkspace.status == AgentWorkingResourceStatus.RETIRED, + ) + ) + if workspace is None: + return + bindings = session.scalars( + select(AgentWorkspaceBinding) + .where( + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.workspace_id == workspace_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.RETIRED, + ) + .order_by(AgentWorkspaceBinding.created_at) + ).all() + if not bindings: + logger.error( + "RETIRED Workspace has no Binding available for physical collection", + extra={"tenant_id": tenant_id, "workspace_id": workspace_id}, + ) + return + anchor = bindings[0] + remaining_ids = [binding.id for binding in bindings[1:]] + workspace_ref = workspace.backend_workspace_ref + binding_ref = anchor.backend_binding_ref + anchor_id = anchor.id + try: + with cls._client() as client: + client.destroy_execution_binding_sync( + DestroyExecutionBindingRequest( + binding_ref=binding_ref, + workspace_ref=workspace_ref, + destroy_workspace=True, + ) + ) + except Exception: + logger.exception( + "Failed to collect retired Agent Workspace", + extra={"tenant_id": tenant_id, "workspace_id": workspace_id, "binding_id": anchor_id}, + ) + return + with session_factory.create_session() as session: + stored_workspace = session.scalar( + select(AgentWorkspace).where( + AgentWorkspace.id == workspace_id, + AgentWorkspace.tenant_id == tenant_id, + AgentWorkspace.status == AgentWorkingResourceStatus.RETIRED, + ) + ) + stored_anchor = session.scalar( + select(AgentWorkspaceBinding).where( + AgentWorkspaceBinding.id == anchor_id, + AgentWorkspaceBinding.tenant_id == tenant_id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.RETIRED, + ) + ) + if stored_workspace is not None: + session.delete(stored_workspace) + if stored_anchor is not None: + session.delete(stored_anchor) + session.commit() + for remaining_id in remaining_ids: + cls.collect_retired_binding(tenant_id=tenant_id, binding_id=remaining_id) + + @staticmethod + def validate_binding_generation( + binding: AgentWorkspaceBinding, + *, + base_home_snapshot_id: str | None, + agent_config_version_id: str, + agent_config_version_kind: AgentConfigVersionKind, + ) -> None: + if ( + binding.base_home_snapshot_id != base_home_snapshot_id + or binding.agent_config_version_id != agent_config_version_id + or binding.agent_config_version_kind != agent_config_version_kind + ): + raise AgentWorkspaceBindingGenerationMismatchError( + "ACTIVE Binding belongs to a different Agent config/Home generation" + ) + + @staticmethod + def _client() -> Client: + base_url = dify_config.AGENT_BACKEND_BASE_URL + if not base_url: + raise AgentWorkspaceError("Dify Agent backend is required for Workspace operations") + return Client(base_url=base_url) + + +__all__ = [ + "AgentWorkspaceBindingGenerationMismatchError", + "AgentWorkspaceError", + "AgentWorkspaceNotFoundError", + "AgentWorkspaceService", + "WorkspaceOwnerScope", +] diff --git a/api/services/agent_app_sandbox_service.py b/api/services/agent_app_sandbox_service.py index 3f5a0bf41b2..91f16c2771c 100644 --- a/api/services/agent_app_sandbox_service.py +++ b/api/services/agent_app_sandbox_service.py @@ -1,39 +1,44 @@ -"""Resolve and proxy sandbox file access for Agent App and workflow Agent sessions. - -These services keep product-facing locators (conversation, workflow run, node) -on the API boundary and translate them into the agent backend's -``SandboxLocator`` using persisted non-sensitive runtime layer specs plus the -saved Agenton session snapshot. Upload responses stay console-facing here: the -agent backend still returns a canonical ToolFile mapping, while this API layer -re-resolves that mapping into a signed browser download URL. -""" +"""Resolve product locators to ACTIVE Workspace Bindings and proxy file access.""" from __future__ import annotations import urllib.parse from collections.abc import Callable -from typing import Any +from dataclasses import dataclass +from typing import Literal, cast -from agenton.compositor import CompositorSessionSnapshot from dify_agent.client import Client -from dify_agent.protocol import RuntimeLayerSpec, SandboxLocator, build_sandbox_locator_from_layer_specs -from pydantic import BaseModel, TypeAdapter +from dify_agent.layers.execution_context import ( + DifyExecutionContextAgentConfigVersionKind, + DifyExecutionContextLayerConfig, +) +from dify_agent.protocol import ( + BindingFileDownloadRequest, + BindingFileListResponse, + BindingFileReadResponse, +) +from pydantic import BaseModel from sqlalchemy import select from sqlalchemy.orm import Session from configs import dify_config -from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore -from core.app.file_access import DatabaseFileAccessController -from core.app.workflow.file_runtime import DifyWorkflowFileRuntime -from factories import file_factory -from models.agent import AgentRuntimeSessionOwnerType, WorkflowAgentRuntimeSession, WorkflowAgentRuntimeSessionStatus - -_RUNTIME_LAYER_SPECS_ADAPTER: TypeAdapter[list[RuntimeLayerSpec]] = TypeAdapter(list[RuntimeLayerSpec]) +from core.db.session_factory import session_factory +from core.tools.signature import bind_file_uri +from models.agent import ( + Agent, + AgentConfigDraft, + AgentConfigDraftType, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, +) +from models.model import App, AppMode, Conversation +from models.workflow import WorkflowNodeExecutionModel +from services.agent.roster_service import AgentRosterService +from services.agent.workspace_service import AgentWorkspaceService, WorkspaceOwnerScope +from services.file_request_service import FileRequestService class AgentSandboxInspectorError(Exception): - """A sandbox inspection failure mapped to an HTTP status by the controller.""" - code: str message: str status_code: int @@ -46,85 +51,239 @@ class AgentSandboxInspectorError(Exception): class AgentSandboxInfo(BaseModel): - """Basic Agent App sandbox metadata returned after a successful availability probe.""" - - session_id: str workspace_cwd: str -class AgentSandboxUploadDownload(BaseModel): - """Signed browser download URL for one sandbox upload result.""" - +class AgentSandboxDownload(BaseModel): url: str -class AgentAppSandboxService: - """Inspect and proxy file access for an Agent App conversation sandbox.""" +@dataclass(frozen=True, slots=True) +class _ResolvedBinding: + """Detached scalar Binding data safe to carry beyond its read transaction. - def __init__( - self, - *, - session_store: AgentAppRuntimeSessionStore | None = None, - client_factory: Callable[[], Client] | None = None, - ) -> None: - self._session_store = session_store or AgentAppRuntimeSessionStore() + Resolvers must end the transaction before Dify Agent network I/O; ORM and + session-bound objects never cross that boundary. + """ + + backend_binding_ref: str + agent_id: str + agent_config_version_id: str + agent_config_version_kind: str + + +class AgentAppSandboxService: + def __init__(self, *, client_factory: Callable[[], Client] | None = None) -> None: self._client_factory = client_factory or _default_client_factory - def get_info(self, *, tenant_id: str, app_id: str, conversation_id: str) -> AgentSandboxInfo: - locator = self._resolve_locator(tenant_id=tenant_id, app_id=app_id, conversation_id=conversation_id) - session_id, workspace_cwd = _extract_shell_workspace_or_raise( - snapshot=locator.session_snapshot, - not_found_message="this conversation's agent has no sandbox workspace", - ) + @staticmethod + def resolve_app_id(*, tenant_id: str, agent_id: str) -> str: + with session_factory.create_session() as session: + app = AgentRosterService(session).get_agent_runtime_app_model( + tenant_id=tenant_id, + agent_id=agent_id, + ) + return app.id - return AgentSandboxInfo( - session_id=session_id, - workspace_cwd=workspace_cwd, - ) - - def list_files(self, *, tenant_id: str, app_id: str, conversation_id: str, path: str): - locator = self._resolve_locator(tenant_id=tenant_id, app_id=app_id, conversation_id=conversation_id) - return self._client_factory().list_sandbox_files_sync(locator, path) - - def read_file(self, *, tenant_id: str, app_id: str, conversation_id: str, path: str): - locator = self._resolve_locator(tenant_id=tenant_id, app_id=app_id, conversation_id=conversation_id) - return self._client_factory().read_sandbox_file_sync(locator, path) - - def upload_file( - self, *, tenant_id: str, app_id: str, conversation_id: str, path: str - ) -> AgentSandboxUploadDownload: - locator = self._resolve_locator(tenant_id=tenant_id, app_id=app_id, conversation_id=conversation_id) - uploaded = self._client_factory().upload_sandbox_file_sync(locator, path) - return _upload_download_response( - tenant_id=tenant_id, - file_mapping=uploaded.file.model_dump(mode="python"), - ) - - def _resolve_locator(self, *, tenant_id: str, app_id: str, conversation_id: str) -> SandboxLocator: - stored = self._session_store.load_active_session_for_conversation( + def get_info( + self, + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: Literal["conversation", "build_draft"], + caller_id: str, + account_id: str, + ) -> AgentSandboxInfo: + self._resolve_binding( tenant_id=tenant_id, app_id=app_id, - conversation_id=conversation_id, + agent_id=agent_id, + caller_type=caller_type, + caller_id=caller_id, + account_id=account_id, ) - if stored is None: - raise AgentSandboxInspectorError( - "no_active_session", - "this conversation has no active sandbox session yet", - status_code=404, + return AgentSandboxInfo(workspace_cwd=".") + + def list_files( + self, + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: Literal["conversation", "build_draft"], + caller_id: str, + account_id: str, + path: str, + ) -> BindingFileListResponse: + binding = self._resolve_binding( + tenant_id=tenant_id, + app_id=app_id, + agent_id=agent_id, + caller_type=caller_type, + caller_id=caller_id, + account_id=account_id, + ) + with self._client_factory() as client: + return client.list_binding_files_sync(binding.backend_binding_ref, path) + + def read_file( + self, + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: Literal["conversation", "build_draft"], + caller_id: str, + account_id: str, + path: str, + ) -> BindingFileReadResponse: + binding = self._resolve_binding( + tenant_id=tenant_id, + app_id=app_id, + agent_id=agent_id, + caller_type=caller_type, + caller_id=caller_id, + account_id=account_id, + ) + with self._client_factory() as client: + return client.read_binding_file_sync(binding.backend_binding_ref, path) + + def download_file( + self, + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: Literal["conversation", "build_draft"], + caller_id: str, + account_id: str, + path: str, + ) -> AgentSandboxDownload: + binding = self._resolve_binding( + tenant_id=tenant_id, + app_id=app_id, + agent_id=agent_id, + caller_type=caller_type, + caller_id=caller_id, + account_id=account_id, + ) + with self._client_factory() as client: + downloaded = client.download_binding_file_sync( + BindingFileDownloadRequest( + backend_binding_ref=binding.backend_binding_ref, + path=path, + execution_context=DifyExecutionContextLayerConfig( + tenant_id=tenant_id, + user_id=account_id, + user_from="account", + app_id=app_id, + conversation_id=caller_id if caller_type == "conversation" else None, + agent_id=agent_id, + agent_config_version_id=binding.agent_config_version_id, + agent_config_version_kind=cast( + DifyExecutionContextAgentConfigVersionKind, + binding.agent_config_version_kind, + ), + agent_mode="agent_app", + invoke_from="debugger", + ), + ) ) - return _build_locator_or_raise( - snapshot=stored.session_snapshot, - runtime_layer_specs=stored.runtime_layer_specs, - not_found_message="this conversation's agent has no sandbox workspace", - ) + return _download_response(tenant_id=tenant_id, account_id=account_id, reference=downloaded.reference) + + @staticmethod + def _resolve_binding( + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: Literal["conversation", "build_draft"], + caller_id: str, + account_id: str, + ) -> _ResolvedBinding: + with session_factory.create_session() as session: + caller: AgentConfigDraft | Conversation | None + if caller_type == "build_draft": + agent = session.scalar( + select(Agent).where( + Agent.id == agent_id, + Agent.tenant_id == tenant_id, + ) + ) + if agent is None or AgentRosterService.runtime_backing_app_id(agent) != app_id: + caller = None + else: + caller = session.scalar( + select(AgentConfigDraft).where( + AgentConfigDraft.id == caller_id, + AgentConfigDraft.tenant_id == tenant_id, + AgentConfigDraft.agent_id == agent_id, + AgentConfigDraft.account_id == account_id, + AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD, + ) + ) + owner_scope = WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=app_id, + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id=caller_id, + ) + else: + caller = session.scalar( + select(Conversation) + .join(App, App.id == Conversation.app_id) + .where( + App.tenant_id == tenant_id, + Conversation.app_id == app_id, + Conversation.id == caller_id, + Conversation.from_account_id == account_id, + Conversation.is_deleted.is_(False), + ) + ) + owner_scope = WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=app_id, + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id=caller_id, + ) + if caller is None or caller.agent_workspace_binding_id is None: + raise AgentSandboxInspectorError( + "no_active_binding", + "this caller has no active Agent Workspace Binding", + status_code=404, + ) + binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=tenant_id, + binding_id=caller.agent_workspace_binding_id, + expected_owner_scope=owner_scope, + ) + if binding is None or binding.agent_id != agent_id: + raise AgentSandboxInspectorError( + "no_active_binding", + "this caller has no active Agent Workspace Binding", + status_code=404, + ) + return _binding_value(binding) class WorkflowAgentSandboxService: - """List/read/upload files in a workflow Agent node sandbox.""" - def __init__(self, *, client_factory: Callable[[], Client] | None = None) -> None: self._client_factory = client_factory or _default_client_factory + @staticmethod + def resolve_app_id(*, tenant_id: str, app_id: str) -> str | None: + with session_factory.create_session() as session: + return session.scalar( + select(App.id).where( + App.id == app_id, + App.tenant_id == tenant_id, + App.status == "normal", + App.mode.in_((AppMode.ADVANCED_CHAT.value, AppMode.WORKFLOW.value)), + ) + ) + def list_files( self, *, @@ -132,11 +291,11 @@ class WorkflowAgentSandboxService: app_id: str, workflow_run_id: str, node_id: str, - node_execution_id: str | None, + node_execution_id: str, path: str, session: Session, - ): - locator = self._resolve_locator( + ) -> BindingFileListResponse: + binding = self._resolve_binding( tenant_id=tenant_id, app_id=app_id, workflow_run_id=workflow_run_id, @@ -144,7 +303,8 @@ class WorkflowAgentSandboxService: node_execution_id=node_execution_id, session=session, ) - return self._client_factory().list_sandbox_files_sync(locator, path) + with self._client_factory() as client: + return client.list_binding_files_sync(binding.backend_binding_ref, path) def read_file( self, @@ -153,11 +313,11 @@ class WorkflowAgentSandboxService: app_id: str, workflow_run_id: str, node_id: str, - node_execution_id: str | None, + node_execution_id: str, path: str, session: Session, - ): - locator = self._resolve_locator( + ) -> BindingFileReadResponse: + binding = self._resolve_binding( tenant_id=tenant_id, app_id=app_id, workflow_run_id=workflow_run_id, @@ -165,142 +325,136 @@ class WorkflowAgentSandboxService: node_execution_id=node_execution_id, session=session, ) - return self._client_factory().read_sandbox_file_sync(locator, path) + with self._client_factory() as client: + return client.read_binding_file_sync(binding.backend_binding_ref, path) - def upload_file( + def download_file( self, *, tenant_id: str, app_id: str, workflow_run_id: str, node_id: str, - node_execution_id: str | None, + node_execution_id: str, + account_id: str, path: str, - session: Session, - ) -> AgentSandboxUploadDownload: - locator = self._resolve_locator( - tenant_id=tenant_id, - app_id=app_id, - workflow_run_id=workflow_run_id, - node_id=node_id, - node_execution_id=node_execution_id, - session=session, - ) - uploaded = self._client_factory().upload_sandbox_file_sync(locator, path) - return _upload_download_response( - tenant_id=tenant_id, - file_mapping=uploaded.file.model_dump(mode="python"), - ) + ) -> AgentSandboxDownload: + with session_factory.create_session() as session: + binding = self._resolve_binding( + tenant_id=tenant_id, + app_id=app_id, + workflow_run_id=workflow_run_id, + node_id=node_id, + node_execution_id=node_execution_id, + session=session, + ) + with self._client_factory() as client: + downloaded = client.download_binding_file_sync( + BindingFileDownloadRequest( + backend_binding_ref=binding.backend_binding_ref, + path=path, + execution_context=DifyExecutionContextLayerConfig( + tenant_id=tenant_id, + user_id=account_id, + user_from="account", + app_id=app_id, + workflow_run_id=workflow_run_id, + node_id=node_id, + node_execution_id=node_execution_id, + agent_id=binding.agent_id, + agent_config_version_id=binding.agent_config_version_id, + agent_config_version_kind=cast( + DifyExecutionContextAgentConfigVersionKind, + binding.agent_config_version_kind, + ), + agent_mode="workflow_run", + invoke_from="debugger", + ), + ) + ) + return _download_response(tenant_id=tenant_id, account_id=account_id, reference=downloaded.reference) - def _resolve_locator( - self, + @staticmethod + def _resolve_binding( *, tenant_id: str, app_id: str, workflow_run_id: str, node_id: str, - node_execution_id: str | None, + node_execution_id: str, session: Session, - ) -> SandboxLocator: - """Resolve one workflow Agent sandbox from product-facing identifiers. - - Callers may target either a specific node execution or the current node - as a whole. When ``node_execution_id`` is provided, lookup narrows to - that execution's ACTIVE runtime-session row. When it is omitted, the - service falls back to the most recently updated ACTIVE session for the - same ``workflow_run_id + node_id`` pair so console sandbox inspection can - still work from the broader workflow/node locator. - """ - stmt = select(WorkflowAgentRuntimeSession).where( - WorkflowAgentRuntimeSession.owner_type == AgentRuntimeSessionOwnerType.WORKFLOW_RUN, - WorkflowAgentRuntimeSession.tenant_id == tenant_id, - WorkflowAgentRuntimeSession.app_id == app_id, - WorkflowAgentRuntimeSession.workflow_run_id == workflow_run_id, - WorkflowAgentRuntimeSession.node_id == node_id, - WorkflowAgentRuntimeSession.status == WorkflowAgentRuntimeSessionStatus.ACTIVE, + ) -> _ResolvedBinding: + execution = session.scalar( + select(WorkflowNodeExecutionModel).where( + WorkflowNodeExecutionModel.id == node_execution_id, + WorkflowNodeExecutionModel.tenant_id == tenant_id, + WorkflowNodeExecutionModel.app_id == app_id, + WorkflowNodeExecutionModel.workflow_run_id == workflow_run_id, + WorkflowNodeExecutionModel.node_id == node_id, + ) ) - if node_execution_id: - stmt = stmt.where(WorkflowAgentRuntimeSession.node_execution_id == node_execution_id) - stmt = stmt.order_by(WorkflowAgentRuntimeSession.updated_at.desc()).limit(1) - - row = session.scalar(stmt) - - if row is None: + process_data = execution.process_data_dict if execution is not None else None + workflow_agent_binding_id = process_data.get("workflow_agent_binding_id") if process_data is not None else None + if ( + execution is None + or execution.agent_workspace_binding_id is None + or not isinstance(workflow_agent_binding_id, str) + ): raise AgentSandboxInspectorError( - "no_active_session", - "this workflow Agent node has no active sandbox session yet", + "no_active_binding", + "this Workflow Agent node execution has no active Workspace Binding", status_code=404, ) - return _build_locator_or_raise( - snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot), - runtime_layer_specs=_deserialize_runtime_layer_specs(row.composition_layer_specs), - not_found_message="this workflow Agent node has no sandbox workspace", - ) - - -def _build_locator_or_raise( - *, - snapshot: CompositorSessionSnapshot, - runtime_layer_specs: list[RuntimeLayerSpec], - not_found_message: str, -) -> SandboxLocator: - try: - return build_sandbox_locator_from_layer_specs( - layer_specs=runtime_layer_specs, - session_snapshot=snapshot, - ) - except ValueError as exc: - raise AgentSandboxInspectorError("no_sandbox", not_found_message, status_code=404) from exc - - -def _extract_shell_workspace_or_raise( - *, - snapshot: CompositorSessionSnapshot, - not_found_message: str, -) -> tuple[str, str]: - shell_layer = next((layer for layer in snapshot.layers if layer.name == "shell"), None) - if shell_layer is None: - raise AgentSandboxInspectorError("no_sandbox", not_found_message, status_code=404) - - session_id = shell_layer.runtime_state.get("session_id") - workspace_cwd = shell_layer.runtime_state.get("workspace_cwd") - if not isinstance(session_id, str) or not isinstance(workspace_cwd, str): - raise AgentSandboxInspectorError("no_sandbox", not_found_message, status_code=404) - return session_id, workspace_cwd - - -def _deserialize_runtime_layer_specs(value: str | None) -> list[RuntimeLayerSpec]: - if not value: - return [] - return _RUNTIME_LAYER_SPECS_ADAPTER.validate_json(value) - - -def _upload_download_response(*, tenant_id: str, file_mapping: dict[str, Any]) -> AgentSandboxUploadDownload: - """Resolve one uploaded ToolFile mapping into a signed external download URL.""" - - controller = DatabaseFileAccessController() - runtime = DifyWorkflowFileRuntime(file_access_controller=controller) - try: - file = file_factory.build_from_mapping( - mapping=file_mapping, + binding = AgentWorkspaceService.get_active_binding( + session=session, tenant_id=tenant_id, - access_controller=controller, + binding_id=execution.agent_workspace_binding_id, + expected_owner_scope=WorkspaceOwnerScope( + tenant_id=tenant_id, + app_id=app_id, + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id=workflow_run_id, + owner_scope_key=f"{node_id}:{workflow_agent_binding_id}", + ), ) - url = runtime.resolve_file_url(file=file, for_external=True) + if binding is None: + raise AgentSandboxInspectorError( + "no_active_binding", + "this Workflow Agent node execution has no active Workspace Binding", + status_code=404, + ) + resolved = _binding_value(binding) + # Deliberately end the read transaction before the caller performs Dify Agent I/O. + session.rollback() + return resolved + + +def _binding_value(binding: AgentWorkspaceBinding) -> _ResolvedBinding: + return _ResolvedBinding( + backend_binding_ref=binding.backend_binding_ref, + agent_id=binding.agent_id, + agent_config_version_id=binding.agent_config_version_id, + agent_config_version_kind=binding.agent_config_version_kind.value, + ) + + +def _download_response(*, tenant_id: str, account_id: str, reference: str) -> AgentSandboxDownload: + try: + result = FileRequestService().request_download( + tenant_id=tenant_id, + user_id=account_id, + user_from="account", + invoke_from="debugger", + file_mapping={"transfer_method": "tool_file", "reference": reference}, + ) + url = bind_file_uri(result.download_uri, dify_config.FILES_URL) except ValueError as exc: raise AgentSandboxInspectorError( - "sandbox_upload_download_unavailable", - "uploaded sandbox file could not be converted to a download URL", + "binding_file_download_unavailable", + "Binding file could not be converted to a download URL", status_code=502, ) from exc - - if not url: - raise AgentSandboxInspectorError( - "sandbox_upload_download_unavailable", - "uploaded sandbox file does not support download URL generation", - status_code=502, - ) - return AgentSandboxUploadDownload(url=_with_as_attachment(url)) + return AgentSandboxDownload(url=_with_as_attachment(url)) def _with_as_attachment(url: str) -> str: @@ -315,7 +469,7 @@ def _default_client_factory() -> Client: if not base_url: raise AgentSandboxInspectorError( "inspector_unavailable", - "the sandbox file inspector is not available (agent backend not configured)", + "the Binding file inspector is not available (Agent backend not configured)", status_code=503, ) return Client(base_url=base_url) @@ -323,8 +477,8 @@ def _default_client_factory() -> Client: __all__ = [ "AgentAppSandboxService", + "AgentSandboxDownload", "AgentSandboxInfo", "AgentSandboxInspectorError", - "AgentSandboxUploadDownload", "WorkflowAgentSandboxService", ] diff --git a/api/services/agent_config_service.py b/api/services/agent_config_service.py index ded7570eb0f..76edc3b199c 100644 --- a/api/services/agent_config_service.py +++ b/api/services/agent_config_service.py @@ -14,8 +14,11 @@ Soul reference only. from __future__ import annotations import io +import mimetypes +import os.path import urllib.parse import zipfile +from collections.abc import Callable from dataclasses import dataclass from enum import StrEnum from operator import itemgetter @@ -27,11 +30,13 @@ from sqlalchemy import select from sqlalchemy.exc import DataError, SQLAlchemyError from sqlalchemy.orm import Session +from configs import dify_config from core.app.file_access.controller import DatabaseFileAccessController -from core.db.session_factory import session_factory +from core.app.workflow.file_runtime import DifyWorkflowFileRuntime +from core.db.session_factory import session_factory as default_session_factory +from core.tools.signature import bind_file_uri, sign_tool_file_uri from core.tools.tool_file_manager import ToolFileManager from extensions.ext_storage import storage -from factories import file_factory from models.agent import Agent, AgentConfigDraft, AgentConfigDraftType, AgentConfigSnapshot from models.agent_config_entities import ( AgentConfigFileRefConfig, @@ -113,19 +118,40 @@ class ConfigDownload: payload: bytes +@dataclass(frozen=True, slots=True) +class ConfigDownloadRequest: + """Short-lived data-plane metadata for one authorized Config asset.""" + + filename: str + mime_type: str + size: int + download_uri: str + + class AgentConfigService: - """Read and update Agent Soul-backed config assets for one version target.""" + """Read and update Agent Soul-backed config assets for one version target. + + The service owns the lifecycle of its database sessions. Callers may inject + a session creator for an alternate engine; production defaults to the + application-wide session factory. + """ PREVIEW_MAX_BYTES = 64 * 1024 + _session_factory: Callable[[], Session] + def __init__( self, *, tool_file_manager: ToolFileManager | None = None, skill_normalize_service: ConfigSkillNormalizeService | None = None, + session_factory: Callable[[], Session] | None = None, ) -> None: + """Initialize external collaborators and the service-owned session creator.""" + self._tool_files = tool_file_manager or ToolFileManager() self._skill_normalizer = skill_normalize_service or ConfigSkillNormalizeService() + self._session_factory = session_factory or default_session_factory.create_session def resolve_target( self, @@ -136,7 +162,7 @@ class AgentConfigService: config_version_kind: AgentConfigVersionKind, user_id: str | None = None, ) -> AgentConfigTarget: - with session_factory.create_session() as session: + with self._session_factory() as session: target = self._resolve_target_in_session( session, tenant_id=tenant_id, @@ -216,16 +242,19 @@ class AgentConfigService: "items": [self._serialize_file_item(file_ref) for file_ref in target.agent_soul.config_files], } - def pull_skill( + def request_download( self, *, tenant_id: str, agent_id: str, config_version_id: str, config_version_kind: AgentConfigVersionKind, + kind: Literal["file", "skill"], name: str, user_id: str | None = None, - ) -> ConfigDownload: + ) -> ConfigDownloadRequest: + """Authorize one Config reference and return origin-free download metadata.""" + target = self.resolve_target( tenant_id=tenant_id, agent_id=agent_id, @@ -233,10 +262,28 @@ class AgentConfigService: config_version_kind=config_version_kind, user_id=user_id, ) - skill = self._require_skill(target.agent_soul, name=name) - file_id = self._available_skill_file_id(skill) - payload, mime_type = self._load_tool_file_bytes(tenant_id=tenant_id, file_id=file_id) - return ConfigDownload(filename=f"{skill.name}.zip", mime_type=mime_type or "application/zip", payload=payload) + if kind == "skill": + skill = self._require_skill(target.agent_soul, name=name) + return self._resolve_download_request( + tenant_id=tenant_id, + file_kind=skill.file_kind, + file_id=self._available_skill_file_id(skill), + filename=f"{skill.name}.zip", + default_mime_type="application/zip", + missing_code="config_skill_not_found", + missing_message="config skill payload is missing", + ) + + file_ref = self._require_file(target.agent_soul, name=name) + return self._resolve_download_request( + tenant_id=tenant_id, + file_kind=file_ref.file_kind, + file_id=self._available_file_id(file_ref), + filename=file_ref.name, + default_mime_type="application/octet-stream", + missing_code="config_file_not_found", + missing_message="config file payload is missing", + ) def download_skill_url( self, @@ -248,19 +295,16 @@ class AgentConfigService: name: str, user_id: str | None = None, ) -> str: - target = self.resolve_target( + result = self.request_download( tenant_id=tenant_id, agent_id=agent_id, config_version_id=config_version_id, config_version_kind=config_version_kind, + kind="skill", + name=name, user_id=user_id, ) - skill = self._require_skill(target.agent_soul, name=name) - file_id = self._available_skill_file_id(skill) - url = self._resolve_download_url(tenant_id=tenant_id, file_kind=skill.file_kind, file_id=file_id) - if url is None: - raise AgentConfigServiceError("config_skill_not_found", "config skill payload is missing", status_code=404) - return url + return bind_file_uri(result.download_uri, dify_config.FILES_URL) def inspect_skill( self, @@ -386,34 +430,6 @@ class AgentConfigService: ) return member_path - def pull_file( - self, - *, - tenant_id: str, - agent_id: str, - config_version_id: str, - config_version_kind: AgentConfigVersionKind, - name: str, - user_id: str | None = None, - ) -> ConfigDownload: - target = self.resolve_target( - tenant_id=tenant_id, - agent_id=agent_id, - config_version_id=config_version_id, - config_version_kind=config_version_kind, - user_id=user_id, - ) - file_ref = self._require_file(target.agent_soul, name=name) - file_id = self._available_file_id(file_ref) - payload, filename, mime_type = self._load_file_ref_bytes( - tenant_id=tenant_id, - file_kind=file_ref.file_kind, - file_id=file_id, - ) - return ConfigDownload( - filename=filename or file_ref.name, mime_type=mime_type or "application/octet-stream", payload=payload - ) - def download_file_url( self, *, @@ -424,19 +440,16 @@ class AgentConfigService: name: str, user_id: str | None = None, ) -> str: - target = self.resolve_target( + result = self.request_download( tenant_id=tenant_id, agent_id=agent_id, config_version_id=config_version_id, config_version_kind=config_version_kind, + kind="file", + name=name, user_id=user_id, ) - file_ref = self._require_file(target.agent_soul, name=name) - file_id = self._available_file_id(file_ref) - url = self._resolve_download_url(tenant_id=tenant_id, file_kind=file_ref.file_kind, file_id=file_id) - if url is None: - raise AgentConfigServiceError("config_file_not_found", "config file payload is missing", status_code=404) - return url + return bind_file_uri(result.download_uri, dify_config.FILES_URL) def upload_skill( self, @@ -494,7 +507,7 @@ class AgentConfigService: filename: str, surface: AgentConfigMutationSurface, ) -> dict[str, object]: - with session_factory.create_session() as session: + with self._session_factory() as session: target = self._resolve_target_in_session( session, tenant_id=tenant_id, @@ -579,7 +592,7 @@ class AgentConfigService: config_version_kind: AgentConfigVersionKind, upload_file_id: str, ) -> dict[str, object]: - with session_factory.create_session() as session: + with self._session_factory() as session: target = self._resolve_target_in_session( session, tenant_id=tenant_id, @@ -615,7 +628,7 @@ class AgentConfigService: payload: ConfigPushPayload, surface: AgentConfigMutationSurface, ) -> dict[str, object]: - with session_factory.create_session() as session: + with self._session_factory() as session: target = self._resolve_target_in_session( session, tenant_id=tenant_id, @@ -709,7 +722,7 @@ class AgentConfigService: env_text: str, surface: AgentConfigMutationSurface, ) -> dict[str, object]: - with session_factory.create_session() as session: + with self._session_factory() as session: target = self._resolve_target_in_session( session, tenant_id=tenant_id, @@ -756,7 +769,7 @@ class AgentConfigService: note: str, surface: AgentConfigMutationSurface, ) -> dict[str, object]: - with session_factory.create_session() as session: + with self._session_factory() as session: target = self._resolve_target_in_session( session, tenant_id=tenant_id, @@ -1339,7 +1352,7 @@ class AgentConfigService: return file_ref.file_id def _load_tool_file_bytes(self, *, tenant_id: str, file_id: str) -> tuple[bytes, str | None]: - with session_factory.create_session() as session: + with self._session_factory() as session: tool_file = session.scalar(select(ToolFile).where(ToolFile.id == file_id, ToolFile.tenant_id == tenant_id)) if tool_file is None: raise AgentConfigServiceError("config_skill_not_found", "config skill payload is missing", status_code=404) @@ -1352,7 +1365,7 @@ class AgentConfigService: file_kind: Literal["upload_file", "tool_file"], file_id: str, ) -> tuple[bytes, str | None, str | None]: - with session_factory.create_session() as session: + with self._session_factory() as session: if file_kind == "tool_file": tool_file = session.scalar( select(ToolFile).where(ToolFile.id == file_id, ToolFile.tenant_id == tenant_id) @@ -1369,35 +1382,56 @@ class AgentConfigService: raise AgentConfigServiceError("config_file_not_found", "config file payload is missing", status_code=404) return storage.load_once(upload_file.key), upload_file.name, upload_file.mime_type - @staticmethod - def _resolve_download_url( - *, tenant_id: str, file_kind: Literal["upload_file", "tool_file"], file_id: str - ) -> str | None: - controller = DatabaseFileAccessController() - from core.app.workflow.file_runtime import DifyWorkflowFileRuntime - - runtime = DifyWorkflowFileRuntime(file_access_controller=controller) - try: - if file_kind == "upload_file": - return runtime.resolve_upload_file_url( - upload_file_id=file_id, - for_external=True, - as_attachment=True, + def _resolve_download_request( + self, + *, + tenant_id: str, + file_kind: Literal["upload_file", "tool_file"], + file_id: str, + filename: str, + default_mime_type: str, + missing_code: str, + missing_message: str, + ) -> ConfigDownloadRequest: + with self._session_factory() as session: + if file_kind == "tool_file": + tool_file = session.scalar( + select(ToolFile).where(ToolFile.id == file_id, ToolFile.tenant_id == tenant_id) ) - file = file_factory.build_from_mapping( - mapping={"transfer_method": "tool_file", "tool_file_id": file_id}, - tenant_id=tenant_id, - access_controller=controller, - ) - url = runtime.resolve_file_url(file=file, for_external=True) - if not url: - return None - parsed = urllib.parse.urlsplit(url) + if tool_file is None: + raise AgentConfigServiceError(missing_code, missing_message, status_code=404) + extension = ( + os.path.splitext(tool_file.name)[1].lower() + or mimetypes.guess_extension(tool_file.mimetype) + or os.path.splitext(tool_file.file_key)[1].lower() + or ".bin" + ) + uri = sign_tool_file_uri(tool_file_id=file_id, extension=extension) + mime_type = tool_file.mimetype or default_mime_type + size = tool_file.size + else: + upload_file = session.scalar( + select(UploadFile).where(UploadFile.id == file_id, UploadFile.tenant_id == tenant_id) + ) + if upload_file is None: + raise AgentConfigServiceError(missing_code, missing_message, status_code=404) + runtime = DifyWorkflowFileRuntime(file_access_controller=DatabaseFileAccessController()) + uri = runtime.resolve_upload_file_uri(upload_file_id=file_id, as_attachment=True) + mime_type = upload_file.mime_type or default_mime_type + size = upload_file.size + + if file_kind == "tool_file": + parsed = urllib.parse.urlsplit(uri) query = urllib.parse.parse_qsl(parsed.query, keep_blank_values=True) query.append(("as_attachment", "true")) - return urllib.parse.urlunsplit(parsed._replace(query=urllib.parse.urlencode(query))) - except ValueError: - return None + uri = urllib.parse.urlunsplit(parsed._replace(query=urllib.parse.urlencode(query))) + + return ConfigDownloadRequest( + filename=filename, + mime_type=mime_type, + size=size, + download_uri=uri, + ) __all__ = [ @@ -1407,6 +1441,7 @@ __all__ = [ "AgentConfigTarget", "AgentConfigVersionKind", "ConfigDownload", + "ConfigDownloadRequest", "ConfigPushFileItem", "ConfigPushPayload", "ConfigPushSkillItem", diff --git a/api/services/agent_file_request_service.py b/api/services/agent_file_request_service.py deleted file mode 100644 index 091fc9fc59b..00000000000 --- a/api/services/agent_file_request_service.py +++ /dev/null @@ -1,93 +0,0 @@ -"""Resolve a download request for a workflow file ref to a signed URL (Agent Files §3.1.1/§4.5). - -The dify-agent server calls this on behalf of a sandbox that needs to pull a -``File`` / ``Array[File]`` workflow input. It binds the flattened file-access -context as a ``FileAccessScope``, rebuilds the graphon ``File`` from the mapping -(reusing tenant/user access checks), and returns an internal signed download URL -plus metadata — never the file bytes. The dify-agent server / sandbox then GETs -the URL directly from Dify API. -""" - -from __future__ import annotations - -from collections.abc import Mapping -from typing import Any - -from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom -from core.app.file_access.controller import DatabaseFileAccessController -from core.app.file_access.scope import FileAccessScope, bind_file_access_scope -from core.app.workflow.file_runtime import DifyWorkflowFileRuntime -from factories import file_factory - - -class FileDownloadRequestError(Exception): - """A download-request failure mapped to an HTTP status by the controller.""" - - code: str - message: str - status_code: int - - def __init__(self, code: str, message: str, *, status_code: int = 400) -> None: - super().__init__(message) - self.code = code - self.message = message - self.status_code = status_code - - -class AgentFileDownloadRequestService: - """Resolve a workflow file ref to a sandbox-accessible internal signed download URL.""" - - @classmethod - def resolve( - cls, - *, - tenant_id: str, - user_id: str, - user_from: str, - invoke_from: str, - file_mapping: Mapping[str, Any], - ) -> dict[str, Any]: - try: - scope_user_from = UserFrom(user_from) - scope_invoke_from = InvokeFrom(invoke_from) - except ValueError as exc: - raise FileDownloadRequestError("invalid_access_context", str(exc), status_code=400) from exc - - if not isinstance(file_mapping, Mapping) or not file_mapping.get("transfer_method"): - raise FileDownloadRequestError("invalid_file_mapping", "file.transfer_method is required", status_code=400) - - scope = FileAccessScope( - tenant_id=tenant_id, - user_id=user_id, - user_from=scope_user_from, - invoke_from=scope_invoke_from, - ) - controller = DatabaseFileAccessController() - runtime = DifyWorkflowFileRuntime(file_access_controller=controller) - try: - with bind_file_access_scope(scope): - file = file_factory.build_from_mapping( - mapping=file_mapping, - tenant_id=tenant_id, - access_controller=controller, - ) - # Internal URL (for_external=False): the consumer is the agent backend / - # sandbox, not a browser. Resolves against INTERNAL_FILES_URL, falling - # back to FILES_URL when not configured. - download_url = runtime.resolve_file_url(file=file, for_external=False) - except ValueError as exc: - raise FileDownloadRequestError("file_not_accessible", str(exc), status_code=404) from exc - - if not download_url: - raise FileDownloadRequestError( - "download_url_unavailable", "could not resolve a download URL for the file", status_code=502 - ) - return { - "filename": file.filename, - "mime_type": file.mime_type, - "size": file.size, - "download_url": download_url, - } - - -__all__ = ["AgentFileDownloadRequestService", "FileDownloadRequestError"] diff --git a/api/services/annotation_service.py b/api/services/annotation_service.py index 947c35fc0a0..087bbd9be2b 100644 --- a/api/services/annotation_service.py +++ b/api/services/annotation_service.py @@ -230,12 +230,12 @@ class AppAnnotationService: escaped_keyword = escape_like_pattern(keyword) stmt = ( select(MessageAnnotation) - .where(MessageAnnotation.app_id == app_id) .where( + MessageAnnotation.app_id == app_id, or_( MessageAnnotation.question.ilike(f"%{escaped_keyword}%", escape="\\"), MessageAnnotation.content.ilike(f"%{escaped_keyword}%", escape="\\"), - ) + ), ) .order_by(MessageAnnotation.created_at.desc(), MessageAnnotation.id.desc()) ) diff --git a/api/services/app_dsl_service.py b/api/services/app_dsl_service.py index 8f41ca109fd..98fe5f9d758 100644 --- a/api/services/app_dsl_service.py +++ b/api/services/app_dsl_service.py @@ -25,6 +25,12 @@ from core.trigger.constants import ( TRIGGER_SCHEDULE_NODE_TYPE, TRIGGER_WEBHOOK_NODE_TYPE, ) +from core.workflow.llm_environment_variable import ( + LLMEnvironmentVariable, + parse_llm_model_selector, + resolve_llm_model_config, + should_resolve_llm_model_selector, +) from core.workflow.nodes.knowledge_retrieval.entities import KnowledgeRetrievalNodeData from core.workflow.nodes.trigger_schedule.trigger_schedule_node import TriggerScheduleNode from events.app_event import app_model_config_was_updated, app_was_created @@ -32,7 +38,7 @@ from extensions.ext_redis import redis_client from factories import variable_factory from graphon.enums import BuiltinNodeTypes from graphon.model_runtime.utils.encoders import jsonable_encoder -from graphon.nodes.llm.entities import LLMNodeData +from graphon.nodes.llm.entities import LLMNodeData, ModelConfig from graphon.nodes.parameter_extractor.entities import ParameterExtractorNodeData from graphon.nodes.question_classifier.entities import QuestionClassifierNodeData from graphon.nodes.tool.entities import ToolNodeData @@ -41,16 +47,24 @@ from models import Account, App, AppMode from models.model import AppModelConfig, AppModelConfigDict, IconType, load_annotation_reply_config from models.workflow import Workflow from services.agent.dsl_service import AgentDslService, AgentPackage +from services.agent.retirement_service import WorkflowAgentRetirementService from services.agent.workflow_publish_service import WorkflowAgentPublishService from services.dsl_content import DSL_MAX_SIZE, dsl_content_size from services.dsl_version import check_version_compatibility from services.enterprise.rbac_service import RBACService -from services.entities.dsl_entities import CheckDependenciesResult, DslImportWarning, ImportMode, ImportStatus +from services.entities.dsl_entities import ( + CheckDependenciesResult, + DslImportWarning, + ImportMode, + ImportStatus, + PendingImportOwner, +) from services.errors.account import NoPermissionError from services.errors.app import WorkflowNotFoundError from services.plugin.dependencies_analysis import DependenciesAnalysisService from services.workflow_draft_variable_service import WorkflowDraftVariableService from services.workflow_service import WorkflowService +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection logger = logging.getLogger(__name__) @@ -72,7 +86,7 @@ class Import(BaseModel): warnings: list[DslImportWarning] = Field(default_factory=list) -class PendingData(BaseModel): +class PendingData(PendingImportOwner): import_mode: str yaml_content: str name: str | None = None @@ -233,6 +247,8 @@ class AppDslService: # If major version mismatch, store import info in Redis if status == ImportStatus.PENDING: pending_data = PendingData( + tenant_id=account.current_tenant_id, + account_id=account.id, import_mode=import_mode, yaml_content=content, name=name, @@ -338,6 +354,15 @@ class AppDslService: error="Invalid import information", ) pending_data = PendingData.model_validate_json(pending_data) + if not pending_data.is_accessible_by( + tenant_id=account.current_tenant_id, + account_id=account.id, + ): + return Import( + id=import_id, + status=ImportStatus.FAILED, + error="Import information expired or does not exist", + ) data = yaml.safe_load(pending_data.yaml_content) app = None @@ -474,8 +499,8 @@ class AppDslService: app.icon_type = resolved_icon_type app.icon = icon app.icon_background = icon_background or app_data.get("icon_background", "#FFFFFF") - app.enable_site = True - app.enable_api = True + app.enable_site = app_mode != AppMode.AGENT + app.enable_api = app_mode != AppMode.AGENT app.use_icon_as_answer_icon = app_data.get("use_icon_as_answer_icon", False) app.created_by = account.id app.maintainer = account.id @@ -546,7 +571,7 @@ class AppDslService: sync_agent_bindings=not raw_agent_packages, ) if raw_agent_packages: - _, warnings = AgentDslService(self._session).import_workflow_packages( + _, warnings, retirement_candidates = AgentDslService(self._session).import_workflow_packages( workflow=draft_workflow, portable_graph=graph, raw_packages=raw_agent_packages, @@ -557,6 +582,17 @@ class AppDslService: session=self._session, draft_workflow=draft_workflow, ) + self._session.commit() + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=app.tenant_id, + agent_ids=retirement_candidates, + account_id=account.id, + ) + enqueue_agent_resource_collection( + tenant_id=app.tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) case AppMode.CHAT | AppMode.AGENT_CHAT | AppMode.COMPLETION: # Initialize model config model_config = data.get("model_config") @@ -616,6 +652,9 @@ class AppDslService: :param app_model: App instance :param session: Database session used to load export data :param include_secret: Whether include secret variable + :param workflow_id: Optional published workflow version to export + :raises WorkflowNotFoundError: If the selected workflow version does not exist + :raises IsDraftWorkflowError: If the selected workflow is a draft :return: """ app_mode = AppMode.value_of(app_model.mode) @@ -675,10 +714,13 @@ class AppDslService: Append workflow export data :param export_data: export data :param app_model: App instance + :param workflow_id: Optional published workflow version to export """ workflow_service = WorkflowService() workflow = workflow_service.get_draft_workflow(app_model, workflow_id, session=session) if not workflow: + if workflow_id: + raise WorkflowNotFoundError(f"Workflow version not found. Workflow ID: {workflow_id}.") raise WorkflowNotFoundError("Missing draft workflow configuration, please check.") workflow_dict = workflow.to_dict(include_secret=include_secret) @@ -776,6 +818,34 @@ class AppDslService: :return: dependencies list format like ["langgenius/google"] """ graph = workflow.graph_dict + referenced_llm_nodes = [ + node.get("data", {}) + for node in graph.get("nodes", []) + if node.get("data", {}).get("type") == BuiltinNodeTypes.LLM + and should_resolve_llm_model_selector(node.get("data", {}).get("model_selector")) + ] + environment_variables = ( + {variable.name: variable for variable in workflow.environment_variables} if referenced_llm_nodes else {} + ) + for node_data in referenced_llm_nodes: + try: + selector = parse_llm_model_selector(node_data["model_selector"]) + variable = environment_variables.get(selector[1]) + if not isinstance(variable, LLMEnvironmentVariable): + raise ValueError( + f"LLM environment variable '{selector[1]}' was not found or is not an LLM variable" + ) + node_data["model"] = resolve_llm_model_config( + node_model=ModelConfig.model_validate(node_data.get("model", {})), + variable_name=selector[1], + variable_value=variable.value, + ).model_dump(mode="json") + except ValueError as exc: + logger.warning( + "Skipping unresolved LLM environment model while extracting dependencies for selector %r: %s", + node_data.get("model_selector"), + exc, + ) dependencies = cls._extract_dependencies_from_workflow_graph(graph) return dependencies diff --git a/api/services/app_generate_service.py b/api/services/app_generate_service.py index 80e6c574fd1..1fc9d7be2da 100644 --- a/api/services/app_generate_service.py +++ b/api/services/app_generate_service.py @@ -21,7 +21,7 @@ from core.app.features.rate_limiting import RateLimit from core.app.features.rate_limiting.rate_limit import rate_limit_context from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig from core.db import session_factory -from enums.quota_type import QuotaType +from enums import DeploymentEdition, QuotaType from extensions.otel import AppGenerateHandler, trace_span from models.model import Account, App, AppMode, EndUser from models.workflow import Workflow, WorkflowRun @@ -134,7 +134,7 @@ class AppGenerateService: action: Callable[[RateLimit, str], Any], ): quota_charge = unlimited() - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: try: quota_charge = QuotaService.reserve(QuotaType.WORKFLOW, app_model.tenant_id) except QuotaExceededError: diff --git a/api/services/app_service.py b/api/services/app_service.py index 3c32d1984e2..1b49983a219 100644 --- a/api/services/app_service.py +++ b/api/services/app_service.py @@ -1,6 +1,7 @@ import json import logging from collections.abc import Sequence +from dataclasses import dataclass from datetime import datetime from typing import Any, Literal, NotRequired, TypedDict, cast, override @@ -13,10 +14,12 @@ from sqlalchemy.orm import Session from configs import dify_config from constants.model_template import default_app_templates from core.agent.entities import AgentToolEntity +from core.agent.publish_visibility import agent_has_workflow_callable_active_snapshot from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.model_manager import ModelManager from core.tools.tool_manager import ToolManager from core.tools.utils.configuration import ToolParameterConfigurationManager +from enums import DeploymentEdition from events.app_event import app_was_created, app_was_deleted, app_was_updated from extensions.ext_database import db # noqa: F401 from graphon.model_runtime.entities.model_entities import ModelPropertyKey, ModelType @@ -25,22 +28,48 @@ from libs.datetime_utils import naive_utc_now from libs.login import current_user from libs.pagination import PaginatedResult, paginate_query from models import Account, AppStar -from models.agent import APP_BACKED_AGENT_SOURCES, Agent, AgentIconType, AgentScope, AgentStatus +from models.agent import ( + APP_BACKED_AGENT_SOURCES, + Agent, + AgentIconType, + AgentScope, + AgentStatus, + AgentWorkingResourceStatus, + AgentWorkspaceBinding, +) from models.model import App, AppMode, AppModelConfig, IconType, Site, load_annotation_reply_config from models.tools import ApiToolProvider from models.workflow import Workflow -from services.agent.errors import AgentNameConflictError +from services.agent.errors import AgentAccessNotReadyError, AgentNameConflictError +from services.agent.home_snapshot_service import AgentHomeSnapshotService +from services.agent.retirement_service import WorkflowAgentRetirementService +from services.agent.workspace_service import AgentWorkspaceService from services.billing_service import BillingService from services.enterprise import rbac_service as enterprise_rbac_service from services.enterprise.enterprise_service import EnterpriseService from services.feature_service import FeatureService from services.openapi.visibility import apply_openapi_gate, is_openapi_visible from services.tag_service import TagService +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection from tasks.remove_app_and_related_data_task import remove_app_and_related_data_task logger = logging.getLogger(__name__) AppListSortBy = Literal["last_modified", "recently_created", "earliest_created"] +RecentAppMode = Literal[ + AppMode.COMPLETION, + AppMode.WORKFLOW, + AppMode.CHAT, + AppMode.ADVANCED_CHAT, + AppMode.AGENT_CHAT, +] +RECENT_APP_MODES: tuple[RecentAppMode, ...] = ( + AppMode.COMPLETION, + AppMode.WORKFLOW, + AppMode.CHAT, + AppMode.ADVANCED_CHAT, + AppMode.AGENT_CHAT, +) class AppListBaseParams(BaseModel): @@ -65,6 +94,19 @@ class StarredAppListParams(AppListBaseParams): pass +@dataclass(frozen=True) +class RecentAppListItem: + id: str + name: str + icon_type: IconType | None + icon: str | None + icon_background: str | None + mode: RecentAppMode + author_name: str | None + updated_at: datetime + maintainer: str | None + + class CreateAppParams(BaseModel): name: str = Field(min_length=1) description: str | None = None @@ -86,7 +128,7 @@ class AppModelConfigResponseView: self._session = session def __getattr__(self, name: str) -> Any: - return getattr(self._app_model_config, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self._app_model_config, name) # guard-ignore: no-new-getattr -- delegates model fields @property def annotation_reply_dict(self) -> Any: @@ -101,7 +143,7 @@ class AppResponseView: self._session = session def __getattr__(self, name: str) -> Any: - return getattr(self._app, name) # noqa: no-new-getattr response adapter delegates model fields + return getattr(self._app, name) # guard-ignore: no-new-getattr -- delegates model fields @property def desc_or_prompt(self) -> str: @@ -323,6 +365,62 @@ class AppService: return app_models + def get_recent_apps( + self, + user_id: str, + tenant_id: str, + params: AppListParams, + session: Session, + ) -> list[RecentAppListItem]: + """Return recently modified apps as one lightweight, non-paginated projection.""" + filters = self._build_app_list_filters(user_id, tenant_id, params, session) + if not filters: + return [] + + stmt = ( + sa.select( + App.id, + App.name, + App.icon_type, + App.icon, + App.icon_background, + App.mode, + Account.name.label("author_name"), + App.updated_at, + App.maintainer, + ) + .outerjoin(Account, Account.id == App.created_by) + .where(*filters, App.mode.in_(RECENT_APP_MODES)) + .order_by(App.updated_at.desc()) + .limit(params.limit) + ) + rows = session.execute(stmt).all() + + return [ + RecentAppListItem( + id=str(app_id), + name=name, + icon_type=icon_type, + icon=icon, + icon_background=icon_background, + mode=cast(RecentAppMode, mode), + author_name=author_name, + updated_at=updated_at, + maintainer=maintainer, + ) + for ( + app_id, + name, + icon_type, + icon, + icon_background, + mode, + author_name, + updated_at, + maintainer, + ) in rows + ] + def get_paginate_starred_apps( self, user_id: str, @@ -394,7 +492,14 @@ class AppService: session.delete(existing_star) - def create_app(self, tenant_id: str, params: CreateAppParams, account: Account, *, session: Session) -> App: + def create_app( + self, + tenant_id: str, + params: CreateAppParams, + account: Account, + *, + session: Session, + ) -> App: """ Create app :param tenant_id: tenant id @@ -548,7 +653,7 @@ class AppService: # update web app setting as private EnterpriseService.WebAppAuth.update_app_access_mode(app.id, "private") - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: BillingService.clean_billing_info_cache(app.tenant_id) return app @@ -812,6 +917,30 @@ class AppService: return app + @staticmethod + def is_agent_app_access_ready(app: App, *, session: Session) -> bool: + """Return whether an Agent App has a publish-visible active snapshot.""" + + if app.mode != AppMode.AGENT: + return True + agent = session.scalar( + select(Agent) + .where( + Agent.tenant_id == app.tenant_id, + Agent.app_id == app.id, + Agent.scope == AgentScope.ROSTER, + Agent.source.in_(APP_BACKED_AGENT_SOURCES), + Agent.status == AgentStatus.ACTIVE, + ) + .limit(1) + ) + return bool(agent and agent_has_workflow_callable_active_snapshot(session=session, agent=agent)) + + @classmethod + def ensure_agent_app_access_ready(cls, app: App, *, session: Session) -> None: + if not cls.is_agent_app_access_ready(app, session=session): + raise AgentAccessNotReadyError() + def update_app_site_status(self, app: App, enable_site: bool, *, session: Session) -> App: """ Update app site status @@ -819,6 +948,8 @@ class AppService: :param enable_site: enable site status :return: App instance """ + if enable_site: + self.ensure_agent_app_access_ready(app, session=session) if enable_site == app.enable_site: return app assert current_user is not None @@ -838,6 +969,8 @@ class AppService: :param enable_api: enable api status :return: App instance """ + if enable_api: + self.ensure_agent_app_access_ready(app, session=session) if enable_api == app.enable_api: return app assert current_user is not None @@ -859,23 +992,72 @@ class AppService: app_was_deleted.send(app) backing_agent = self._get_backing_agent_for_update(app, session=session) + workflow_agent_ids = session.scalars( + select(Agent.id).where( + Agent.tenant_id == app.tenant_id, + Agent.app_id == app.id, + Agent.scope == AgentScope.WORKFLOW_ONLY, + Agent.status == AgentStatus.ACTIVE, + ) + ).all() + account_id = current_user.id if current_user else None if backing_agent is not None: now = naive_utc_now() - account_id = getattr(current_user, "id", None) backing_agent.status = AgentStatus.ARCHIVED backing_agent.archived_by = account_id backing_agent.archived_at = now backing_agent.updated_by = account_id backing_agent.updated_at = now + retired_binding_ids: list[str] = [] + retired_snapshot_ids: list[str] = [] + if backing_agent is not None: + bindings = session.scalars( + select(AgentWorkspaceBinding).where( + AgentWorkspaceBinding.tenant_id == app.tenant_id, + AgentWorkspaceBinding.agent_id == backing_agent.id, + AgentWorkspaceBinding.status == AgentWorkingResourceStatus.ACTIVE, + ) + ).all() + for binding in bindings: + binding_id = AgentWorkspaceService.retire_binding( + session=session, + tenant_id=app.tenant_id, + binding_id=binding.id, + ) + if binding_id is not None: + retired_binding_ids.append(binding_id) + retired_snapshot_ids = AgentHomeSnapshotService.retire_all_for_agent( + session=session, + tenant_id=app.tenant_id, + agent_id=backing_agent.id, + ) + + retired_workspace_ids = AgentWorkspaceService.retire_all_for_app( + session=session, + tenant_id=app.tenant_id, + app_id=app.id, + ) session.delete(app) session.commit() + workflow_binding_ids, workflow_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=app.tenant_id, + agent_ids=workflow_agent_ids, + account_id=account_id, + ) + enqueue_agent_resource_collection( + tenant_id=app.tenant_id, + workspace_ids=retired_workspace_ids, + binding_ids=[*retired_binding_ids, *workflow_binding_ids], + home_snapshot_ids=[*retired_snapshot_ids, *workflow_snapshot_ids], + ) + # clean up web app settings if FeatureService.get_system_features().webapp_auth.enabled: EnterpriseService.WebAppAuth.cleanup_webapp(app.id) - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: BillingService.clean_billing_info_cache(app.tenant_id) # Trigger asynchronous deletion of app and related data diff --git a/api/services/async_workflow_service.py b/api/services/async_workflow_service.py index 012e979aff2..48a113e8fcc 100644 --- a/api/services/async_workflow_service.py +++ b/api/services/async_workflow_service.py @@ -14,7 +14,7 @@ from celery.result import AsyncResult from sqlalchemy import select from sqlalchemy.orm import Session, sessionmaker -from enums.quota_type import QuotaType +from enums import QuotaType from extensions.ext_database import db from models.account import Account from models.enums import CreatorUserRole, WorkflowTriggerStatus diff --git a/api/services/auth/watercrawl/watercrawl.py b/api/services/auth/watercrawl/watercrawl.py index ac8ec24d893..2b2dd20641c 100644 --- a/api/services/auth/watercrawl/watercrawl.py +++ b/api/services/auth/watercrawl/watercrawl.py @@ -6,6 +6,10 @@ import httpx from services.auth.api_key_auth_base import ApiKeyAuthBase, AuthCredentials +# Explicit bounded timeout for credential-validation requests so a slow or +# hanging WaterCrawl endpoint cannot block the worker indefinitely. +_CREDENTIAL_TIMEOUT = httpx.Timeout(10.0) + class WatercrawlAuth(ApiKeyAuthBase): def __init__(self, credentials: AuthCredentials): @@ -33,7 +37,7 @@ class WatercrawlAuth(ApiKeyAuthBase): return {"Content-Type": "application/json", "X-API-KEY": self.api_key} def _get_request(self, url, headers): - return httpx.get(url, headers=headers) + return httpx.get(url, headers=headers, timeout=_CREDENTIAL_TIMEOUT) def _handle_error(self, response): if response.status_code in {402, 409, 500}: diff --git a/api/services/billing_service.py b/api/services/billing_service.py index aef5ae2f02b..0519f66f2c6 100644 --- a/api/services/billing_service.py +++ b/api/services/billing_service.py @@ -12,7 +12,7 @@ from tenacity import retry, retry_if_exception_type, stop_before_delay, wait_fix from werkzeug.exceptions import InternalServerError from core.helper.http_client_pooling import get_pooled_http_client -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from extensions.ext_redis import redis_client from libs.helper import RateLimiter from models import Account, TenantAccountJoin, TenantAccountRole @@ -21,7 +21,10 @@ logger = logging.getLogger(__name__) _http_client: httpx.Client = get_pooled_http_client( "billing:default", - lambda: httpx.Client(limits=httpx.Limits(max_keepalive_connections=50, max_connections=100)), + lambda: httpx.Client( + timeout=httpx.Timeout(30.0, connect=5.0), + limits=httpx.Limits(max_keepalive_connections=50, max_connections=100), + ), ) @@ -102,6 +105,7 @@ class _BillingQuota(TypedDict): class _VectorSpaceQuota(TypedDict): size: float limit: int + usage_unknown: NotRequired[bool] class _KnowledgeRateLimit(TypedDict): diff --git a/api/services/clear_free_plan_tenant_expired_logs.py b/api/services/clear_free_plan_tenant_expired_logs.py index 963938982d0..96885ebb225 100644 --- a/api/services/clear_free_plan_tenant_expired_logs.py +++ b/api/services/clear_free_plan_tenant_expired_logs.py @@ -10,7 +10,7 @@ from sqlalchemy import delete, func, select from sqlalchemy.orm import Session, sessionmaker from configs import dify_config -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from extensions.ext_database import db from extensions.ext_storage import storage from graphon.model_runtime.utils.encoders import jsonable_encoder @@ -371,7 +371,7 @@ class ClearFreePlanTenantExpiredLogs: def process_tenant(flask_app: Flask, tenant_id: str): try: if ( - not dify_config.BILLING_ENABLED + dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD or BillingService.get_info(tenant_id)["subscription"]["plan"] == CloudPlan.SANDBOX ): # only process sandbox tenant diff --git a/api/services/conversation_service.py b/api/services/conversation_service.py index 6ab9c5a8a2b..406af068194 100644 --- a/api/services/conversation_service.py +++ b/api/services/conversation_service.py @@ -6,9 +6,7 @@ from typing import Any from sqlalchemy import asc, desc, func, or_, select from sqlalchemy.orm import Session -from clients.agent_backend import AgentBackendSessionCleanupPayload from configs import dify_config -from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore from core.app.entities.app_invoke_entities import InvokeFrom from core.llm_generator.llm_generator import LLMGenerator from factories import variable_factory @@ -16,7 +14,9 @@ from graphon.variables.types import SegmentType from libs.datetime_utils import naive_utc_now from libs.infinite_scroll_pagination import InfiniteScrollPagination from models import Account, ConversationVariable +from models.agent import AgentWorkspaceOwnerType from models.model import App, Conversation, EndUser, Message +from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService, WorkspaceOwnerScope from services.errors.conversation import ( ConversationNotExistsError, ConversationVariableNotExistsError, @@ -24,7 +24,7 @@ from services.errors.conversation import ( LastConversationNotExistsError, ) from services.errors.message import MessageNotExistsError -from tasks.agent_backend_session_cleanup_task import cleanup_conversation_agent_runtime_session +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection from tasks.delete_conversation_task import delete_conversation_related_data logger = logging.getLogger(__name__) @@ -189,23 +189,30 @@ class ConversationService: """ Delete a conversation only if it belongs to the given user and app context. - Before removing the conversation row, this best-effort lifecycle path - enumerates any ACTIVE conversation-owned Agent backend runtime sessions, - enqueues asynchronous backend cleanup for rows with persisted runtime - layer specs, and then retires the local session rows even if enqueueing - fails. Conversation deletion and related-data cleanup scheduling still - proceed when that lifecycle bookkeeping only partially succeeds. + Conversation deletion is the product lifecycle boundary for its + Workspace. Physical collection happens only after the retire commit. Raises: ConversationNotExistsError: When the conversation is not visible to the current user. """ conversation = cls.get_conversation(app_model, conversation_id, user, session=session) - session_store = AgentAppRuntimeSessionStore() - stored_sessions = session_store.list_active_sessions_for_conversation( - tenant_id=app_model.tenant_id, - app_id=app_model.id, - conversation_id=conversation.id, - ) + binding_id = conversation.agent_workspace_binding_id + retired_binding_id: str | None = None + if binding_id is not None: + owner_scope = WorkspaceOwnerScope( + tenant_id=app_model.tenant_id, + app_id=app_model.id, + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id=conversation.id, + ) + binding = AgentWorkspaceService.get_active_binding( + session=session, + tenant_id=app_model.tenant_id, + binding_id=binding_id, + expected_owner_scope=owner_scope, + ) + if binding is None: + raise AgentWorkspaceNotFoundError("Conversation participant Binding is unavailable") try: logger.info( @@ -213,66 +220,25 @@ class ConversationService: app_model.name, conversation_id, ) - for stored_session in stored_sessions: - try: - if stored_session.runtime_layer_specs: - payload = AgentBackendSessionCleanupPayload( - session_snapshot=stored_session.session_snapshot, - runtime_layer_specs=stored_session.runtime_layer_specs, - idempotency_key=( - f"{stored_session.scope.tenant_id}:{stored_session.scope.app_id}:" - f"{stored_session.scope.conversation_id}:agent-runtime-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'}" - ), - metadata={ - "tenant_id": stored_session.scope.tenant_id, - "app_id": stored_session.scope.app_id, - "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, - "previous_agent_backend_run_id": stored_session.backend_run_id, - }, - ) - cleanup_conversation_agent_runtime_session.delay(payload.model_dump(mode="json")) - except Exception: - logger.warning( - "Failed to enqueue Agent backend cleanup for conversation deletion: " - "tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s", - stored_session.scope.tenant_id, - stored_session.scope.app_id, - stored_session.scope.conversation_id, - stored_session.scope.agent_id, - stored_session.backend_run_id, - exc_info=True, - ) - finally: - try: - session_store.mark_cleaned( - scope=stored_session.scope, - backend_run_id=stored_session.backend_run_id, - ) - except Exception: - logger.warning( - "Failed to retire Agent App runtime session for conversation deletion: " - "tenant_id=%s app_id=%s conversation_id=%s agent_id=%s backend_run_id=%s", - stored_session.scope.tenant_id, - stored_session.scope.app_id, - stored_session.scope.conversation_id, - stored_session.scope.agent_id, - stored_session.backend_run_id, - exc_info=True, - ) - + if binding_id is not None: + retired_binding_id = AgentWorkspaceService.retire_binding( + session=session, + tenant_id=app_model.tenant_id, + binding_id=binding_id, + ) + if retired_binding_id is None: + raise AgentWorkspaceNotFoundError("Conversation participant Binding is unavailable") session.delete(conversation) session.commit() - - delete_conversation_related_data.delay(conversation.id) - except Exception: session.rollback() raise + if retired_binding_id is not None: + enqueue_agent_resource_collection( + tenant_id=app_model.tenant_id, + binding_ids=(retired_binding_id,), + ) + delete_conversation_related_data.delay(conversation.id) @classmethod def get_conversational_variable( @@ -290,8 +256,7 @@ class ConversationService: stmt = ( select(ConversationVariable) - .where(ConversationVariable.app_id == app_model.id) - .where(ConversationVariable.conversation_id == conversation.id) + .where(ConversationVariable.app_id == app_model.id, ConversationVariable.conversation_id == conversation.id) .order_by(ConversationVariable.created_at) ) @@ -376,11 +341,10 @@ class ConversationService: conversation = cls.get_conversation(app_model, conversation_id, user, session=session) # Get the existing conversation variable - stmt = ( - select(ConversationVariable) - .where(ConversationVariable.app_id == app_model.id) - .where(ConversationVariable.conversation_id == conversation.id) - .where(ConversationVariable.id == variable_id) + stmt = select(ConversationVariable).where( + ConversationVariable.app_id == app_model.id, + ConversationVariable.conversation_id == conversation.id, + ConversationVariable.id == variable_id, ) existing_variable = session.scalar(stmt) diff --git a/api/services/credit_pool_service.py b/api/services/credit_pool_service.py index f398ed6e6e4..2a18e52a34c 100644 --- a/api/services/credit_pool_service.py +++ b/api/services/credit_pool_service.py @@ -7,7 +7,9 @@ from piling up database transactions while preserving cross-tenant concurrency. import logging from collections.abc import Callable -from dataclasses import dataclass +from dataclasses import dataclass, field +from enum import StrEnum, auto +from typing import Any from uuid import uuid4 from sqlalchemy import select @@ -15,6 +17,7 @@ from sqlalchemy.orm import Session from configs import dify_config from core.errors.error import QuotaExceededError +from enums import DeploymentEdition from extensions.ext_redis import redis_client from models import TenantCreditPool from models.enums import ProviderQuotaType @@ -44,6 +47,77 @@ class CreditPoolBalance: return self.quota_limit == -1 or self.remaining_credits >= required_credits +class CreditPoolReservationState(StrEnum): + RESERVED = auto() + COMMITTED = auto() + RELEASED = auto() + + +@dataclass +class CreditPoolReservation: + """A strict credit-pool reservation spanning one billable operation.""" + + tenant_id: str + pool_type: str + amount: int + request_id: str + reservation_id: str | None + meta: dict[str, Any] = field(default_factory=dict) + _session_factory: Callable[[], Session] | None = field(default=None, repr=False) + _state: CreditPoolReservationState = field(default=CreditPoolReservationState.RESERVED, init=False, repr=False) + + @property + def state(self) -> CreditPoolReservationState: + return self._state + + def commit(self) -> None: + if self._state == CreditPoolReservationState.COMMITTED: + return + if self._state == CreditPoolReservationState.RELEASED: + raise RuntimeError("Cannot commit a released credit reservation.") + + if self.reservation_id is not None: + from services.billing_service import BillingService + + BillingService.quota_commit( + tenant_id=self.tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=self.pool_type, + reservation_id=self.reservation_id, + actual_amount=self.amount, + meta={**self.meta, "request_id": self.request_id}, + ) + + # The database fallback reserves by deducting under the tenant lock, so + # commit only makes that already durable reservation final. + self._state = CreditPoolReservationState.COMMITTED + + def release(self) -> None: + if self._state in {CreditPoolReservationState.COMMITTED, CreditPoolReservationState.RELEASED}: + return + + if self.reservation_id is not None: + from services.billing_service import BillingService + + BillingService.quota_release( + tenant_id=self.tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=self.pool_type, + reservation_id=self.reservation_id, + ) + else: + if self._session_factory is None: + raise RuntimeError("Database credit reservation requires a session factory.") + CreditPoolService._release_database_reservation( + tenant_id=self.tenant_id, + pool_type=self.pool_type, + credits=self.amount, + session=self._session_factory(), + ) + + self._state = CreditPoolReservationState.RELEASED + + class CreditPoolService: @staticmethod def _normalize_pool_type(pool_type: str | ProviderQuotaType) -> str: @@ -51,7 +125,7 @@ class CreditPoolService: @staticmethod def _use_billing_quota() -> bool: - return bool(dify_config.BILLING_ENABLED) + return bool(dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD) @staticmethod def _require_session(session: Session | None) -> Session: @@ -162,6 +236,110 @@ class CreditPoolService: return False return pool.has_sufficient_credits(credits_required) + @classmethod + def reserve_credits( + cls, + tenant_id: str, + credits_required: int, + pool_type: str | ProviderQuotaType = "trial", + *, + request_id: str, + session_factory: Callable[[], Session] | None = None, + meta: dict[str, Any] | None = None, + ) -> CreditPoolReservation: + """Reserve the full amount or raise before the billable operation starts.""" + if credits_required <= 0: + raise ValueError("credits_required must be greater than 0") + if not request_id: + raise ValueError("request_id is required") + + normalized_pool_type = cls._normalize_pool_type(pool_type) + reservation_meta = {"source": "credit_pool.reservation", **(meta or {})} + if cls._use_billing_quota(): + from services.billing_service import BillingService + + result = BillingService.quota_reserve( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=normalized_pool_type, + request_id=request_id, + amount=credits_required, + meta=reservation_meta, + ) + reservation_id = result.get("reservation_id", "") + if not reservation_id: + raise QuotaExceededError("Insufficient credits remaining") + return CreditPoolReservation( + tenant_id=tenant_id, + pool_type=normalized_pool_type, + amount=credits_required, + request_id=request_id, + reservation_id=reservation_id, + meta=reservation_meta, + ) + + if session_factory is None: + raise ValueError("session_factory is required when billing quota is disabled") + + session = session_factory() + + def reserve() -> int: + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=normalized_pool_type) + if not pool: + raise QuotaExceededError("Credit pool not found") + if not pool.has_sufficient_credits(credits_required): + raise QuotaExceededError("Insufficient credits remaining") + + pool.quota_used += credits_required + session.commit() + return credits_required + + try: + cls._deduct_with_tenant_lock(tenant_id, reserve) + except QuotaExceededError: + session.rollback() + raise + except Exception: + session.rollback() + logger.exception("Failed to reserve credits for tenant %s", tenant_id) + raise QuotaExceededError("Failed to reserve credits") + + return CreditPoolReservation( + tenant_id=tenant_id, + pool_type=normalized_pool_type, + amount=credits_required, + request_id=request_id, + reservation_id=None, + meta=reservation_meta, + _session_factory=session_factory, + ) + + @classmethod + def _release_database_reservation( + cls, + *, + tenant_id: str, + pool_type: str, + credits: int, + session: Session, + ) -> None: + def release() -> int: + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) + if not pool: + raise QuotaExceededError("Credit pool not found") + if pool.quota_used < credits: + raise RuntimeError("Reserved credits exceed recorded usage.") + + pool.quota_used -= credits + session.commit() + return credits + + try: + cls._deduct_with_tenant_lock(tenant_id, release) + except Exception: + session.rollback() + raise + @classmethod def check_and_deduct_credits( cls, diff --git a/api/services/data_migration/import_service.py b/api/services/data_migration/import_service.py index 6b999119b2e..48ff50ac655 100644 --- a/api/services/data_migration/import_service.py +++ b/api/services/data_migration/import_service.py @@ -26,6 +26,7 @@ from models import Account, ApiToken, Tenant, TenantAccountJoin, TenantAccountRo from models.enums import ApiTokenType from models.model import App from models.tools import ApiToolProvider, MCPToolProvider, WorkflowToolProvider +from services.agent.retirement_service import WorkflowAgentRetirementService from services.app_dsl_service import AppDslService from services.data_migration.dependency_discovery_service import DependencyDiscoveryService from services.data_migration.entities import ( @@ -47,6 +48,7 @@ from services.tools.api_tools_manage_service import ApiToolManageService from services.tools.mcp_tools_manage_service import MCPToolManageService from services.tools.workflow_tools_manage_service import WorkflowToolManageService from services.workflow_service import WorkflowService +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection @dataclass(frozen=True) @@ -322,7 +324,7 @@ class MigrationImportService: options: ImportOptions, session: Session, ) -> str: - import_service = AppDslService(cast(Session, session)) + import_service = AppDslService(session) if existing_app is not None: import_result = import_service.import_app( account=account, @@ -711,7 +713,7 @@ class MigrationImportService: raise MigrationDataError(f"Referenced workflow app was not found in target tenant: {app_id}") if account_in_session is None: raise MigrationDataError(f"Operator account not found: {account.id}") - workflow = workflow_service.publish_workflow( + workflow, retirement_candidates = workflow_service.publish_workflow( session=session, app_model=app_in_session, account=account_in_session, @@ -721,6 +723,16 @@ class MigrationImportService: app_in_session.workflow_id = workflow.id app_in_session.updated_by = account.id app_in_session.updated_at = naive_utc_now() + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=target.tenant_id, + agent_ids=retirement_candidates, + account_id=account.id, + ) + enqueue_agent_resource_collection( + tenant_id=target.tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) def _import_mcp_tools( self, @@ -756,7 +768,7 @@ class MigrationImportService: report_items.append(ResourceReportItem(ResourceType.MCP_TOOL, existing.id, name, "skipped")) continue - service = MCPToolManageService(session=cast(Session, session)) + service = MCPToolManageService(session=session) configuration = MCPConfiguration.model_validate(mcp_data.get("configuration") or {}) authentication = ( MCPAuthentication.model_validate(mcp_data["authentication"]) if mcp_data.get("authentication") else None diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index 27650e68101..004e927b675 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -23,7 +23,7 @@ from core.model_manager import ModelManager from core.rag.index_processor.constant.built_in_field import BuiltInField from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType from core.rag.retrieval.retrieval_methods import RetrievalMethod -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from events.dataset_event import dataset_was_deleted from events.document_event import document_was_deleted from extensions.ext_redis import redis_client @@ -64,10 +64,11 @@ from models.model import UploadFile from models.provider_ids import ModelProviderID from models.source import DataSourceOauthBinding from models.workflow import Workflow -from services.dataset_ref_service import DatasetRef, SegmentRef +from services.dataset_ref_service import DatasetRef, DatasetRefService, SegmentRef from services.document_indexing_proxy.document_indexing_task_proxy import DocumentIndexingTaskProxy from services.document_indexing_proxy.duplicate_document_indexing_task_proxy import DuplicateDocumentIndexingTaskProxy from services.enterprise import rbac_service as enterprise_rbac_service +from services.entities.feature_entities import FeatureModel from services.entities.knowledge_entities.knowledge_entities import ( ChildChunkUpdateArgs, KnowledgeConfig, @@ -85,7 +86,7 @@ from services.errors.dataset import DatasetNameDuplicateError from services.errors.document import DocumentIndexingError from services.errors.file import FileNotExistsError from services.external_knowledge_service import ExternalDatasetService -from services.feature_service import FeatureModel, FeatureService +from services.feature_service import FeatureService from services.file_service import FileService from services.rag_pipeline.rag_pipeline import RagPipelineService from services.tag_service import TagService @@ -1364,8 +1365,8 @@ class DatasetService: return True @staticmethod - def dataset_use_check(dataset_id, session: Session) -> bool: - stmt = select(exists().where(AppDatasetJoin.dataset_id == dataset_id)) + def dataset_use_check(dataset_ref: DatasetRef, session: Session) -> bool: + stmt = select(exists().where(AppDatasetJoin.dataset_id == dataset_ref.dataset_id)) return session.execute(stmt).scalar_one() @staticmethod @@ -1383,7 +1384,11 @@ class DatasetService: if dataset.maintainer != user.id: user_permission = session.scalar( select(DatasetPermission) - .where(DatasetPermission.dataset_id == dataset.id, DatasetPermission.account_id == user.id) + .where( + DatasetPermission.dataset_id == dataset.id, + DatasetPermission.account_id == user.id, + DatasetPermission.tenant_id == dataset.tenant_id, + ) .limit(1) ) if not user_permission: @@ -1406,12 +1411,16 @@ class DatasetService: raise NoPermissionError("You do not have permission to access this dataset.") elif dataset.permission == DatasetPermissionEnum.PARTIAL_TEAM: - if not any( - dp.dataset_id == dataset.id - for dp in session.scalars( - select(DatasetPermission).where(DatasetPermission.account_id == user.id) - ).all() - ): + user_permission = session.scalar( + select(DatasetPermission.id) + .where( + DatasetPermission.dataset_id == dataset.id, + DatasetPermission.account_id == user.id, + DatasetPermission.tenant_id == dataset.tenant_id, + ) + .limit(1) + ) + if user_permission is None: raise NoPermissionError("You do not have permission to access this dataset.") @staticmethod @@ -1431,22 +1440,17 @@ class DatasetService: ).all() @staticmethod - def update_dataset_api_status(dataset_id: str, status: bool, session: Session): - dataset = DatasetService.get_dataset(dataset_id, session) - if dataset is None: - raise NotFound("Dataset not found.") - dataset.enable_api = status - if not current_user or not current_user.id: + def update_dataset_api_status(dataset: Dataset, status: bool, actor: Account, session: Session): + if not actor.id: raise ValueError("Current user or current user id not found") - dataset.updated_by = current_user.id + dataset.enable_api = status + dataset.updated_by = actor.id dataset.updated_at = naive_utc_now() session.flush() @staticmethod - def get_dataset_auto_disable_logs(dataset_id: str, session: Session) -> AutoDisableLogsDict: - assert isinstance(current_user, Account) - assert current_user.current_tenant_id is not None - features = FeatureService.get_features(current_user.current_tenant_id, exclude_vector_space=True) + def get_dataset_auto_disable_logs(dataset_ref: DatasetRef, session: Session) -> AutoDisableLogsDict: + features = FeatureService.get_features(dataset_ref.tenant_id, exclude_vector_space=True) if not features.billing.enabled or features.billing.subscription.plan == CloudPlan.SANDBOX: return { "document_ids": [], @@ -1456,7 +1460,8 @@ class DatasetService: start_date = datetime.datetime.now() - datetime.timedelta(days=30) dataset_auto_disable_logs = session.scalars( select(DatasetAutoDisableLog).where( - DatasetAutoDisableLog.dataset_id == dataset_id, + DatasetAutoDisableLog.tenant_id == dataset_ref.tenant_id, + DatasetAutoDisableLog.dataset_id == dataset_ref.dataset_id, DatasetAutoDisableLog.created_at >= start_date, ) ).all() @@ -1653,7 +1658,9 @@ class DocumentService: return None @staticmethod - def get_documents_by_ids(dataset_id: str, document_ids: Sequence[str], session: Session) -> Sequence[Document]: + def get_documents_by_ids( + dataset_ref: DatasetRef, document_ids: Sequence[str], session: Session + ) -> Sequence[Document]: """Fetch documents for a dataset in a single batch query.""" if not document_ids: return [] @@ -1661,7 +1668,8 @@ class DocumentService: # Fetch all requested documents in one query to avoid N+1 lookups. documents: Sequence[Document] = session.scalars( select(Document).where( - Document.dataset_id == dataset_id, + Document.tenant_id == dataset_ref.tenant_id, + Document.dataset_id == dataset_ref.dataset_id, Document.id.in_(document_id_list), ) ).all() @@ -1855,7 +1863,9 @@ class DocumentService: """ document_id_list: list[str] = [str(document_id) for document_id in document_ids] - documents = DocumentService.get_documents_by_ids(dataset_id, document_id_list, session) + documents = DocumentService.get_documents_by_ids( + DatasetRef(tenant_id=tenant_id, dataset_id=dataset_id), document_id_list, session + ) documents_by_id: dict[str, Document] = {str(document.id): document for document in documents} missing_document_ids: set[str] = set(document_id_list) - set(documents_by_id.keys()) @@ -1865,9 +1875,6 @@ class DocumentService: upload_file_ids: list[str] = [] upload_file_ids_by_document_id: dict[str, str] = {} for document_id, document in documents_by_id.items(): - if document.tenant_id != tenant_id: - raise Forbidden("No permission.") - upload_file_id = DocumentService._get_upload_file_id_for_upload_file_document( document, invalid_source_message="Only uploaded-file documents can be downloaded as ZIP.", @@ -1894,9 +1901,13 @@ class DocumentService: return document @staticmethod - def get_document_by_ids(document_ids: list[str], session: Session) -> Sequence[Document]: + def get_document_by_ids( + dataset_ref: DatasetRef, document_ids: Sequence[str], session: Session + ) -> Sequence[Document]: documents = session.scalars( select(Document).where( + Document.tenant_id == dataset_ref.tenant_id, + Document.dataset_id == dataset_ref.dataset_id, Document.id.in_(document_ids), Document.enabled == True, Document.indexing_status == IndexingStatus.COMPLETED, @@ -1930,10 +1941,11 @@ class DocumentService: return documents @staticmethod - def get_error_documents_by_dataset_id(dataset_id: str, session: Session) -> Sequence[Document]: + def get_error_documents_by_dataset_ref(dataset_ref: DatasetRef, session: Session) -> Sequence[Document]: documents = session.scalars( select(Document).where( - Document.dataset_id == dataset_id, + Document.tenant_id == dataset_ref.tenant_id, + Document.dataset_id == dataset_ref.dataset_id, Document.indexing_status.in_([IndexingStatus.ERROR, IndexingStatus.PAUSED]), ) ).all() @@ -2099,26 +2111,54 @@ class DocumentService: @staticmethod def retry_document(dataset_id: str, documents: list[Document], session: Session): - for document in documents: - # add retry flag - retry_indexing_cache_key = f"document_{document.id}_is_retried" - cache_result = redis_client.get(retry_indexing_cache_key) - if cache_result is not None: - raise ValueError("Document is being retried, please try again later") - # retry document indexing - document.indexing_status = IndexingStatus.WAITING - session.add(document) - session.commit() + """Reserve the whole retry batch before changing any document state. - redis_client.setex(retry_indexing_cache_key, 600, 1) - # trigger async task - document_ids = [document.id for document in documents] + Redis lock acquisition is intentionally coupled to this bounded status + transaction so a concurrent request cannot partially admit the batch. + """ if not current_user or not current_user.id: raise ValueError("Current user or current user id not found") + + unique_documents = list({document.id: document for document in documents}.values()) + retry_indexing_cache_keys = [f"document_{document.id}_is_retried" for document in unique_documents] + acquired_locks: list[Any] = [] + + def release_acquired_locks() -> None: + for retry_lock in acquired_locks: + try: + retry_lock.release() + except Exception: + logger.warning("Failed to release document retry lock", exc_info=True) + + try: + for retry_indexing_cache_key in retry_indexing_cache_keys: + retry_lock = redis_client.lock(retry_indexing_cache_key, timeout=600, thread_local=False) + if not retry_lock.acquire(blocking=False): + raise ValueError("Document is being retried, please try again later") + acquired_locks.append(retry_lock) + except Exception: + release_acquired_locks() + raise + + try: + for document in unique_documents: + document.indexing_status = IndexingStatus.WAITING + session.add(document) + session.commit() + except Exception: + session.rollback() + release_acquired_locks() + raise + + document_ids = [document.id for document in unique_documents] retry_document_indexing_task.delay(dataset_id, document_ids, current_user.id) @staticmethod - def sync_website_document(dataset_id: str, document: Document, session: Session): + def sync_website_document(dataset: Dataset, document: Document, session: Session): + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + if DatasetRefService.create_document_ref(dataset_ref, document) is None: + raise ValueError("Document not found.") + # add sync flag sync_indexing_cache_key = f"document_{document.id}_is_sync" cache_result = redis_client.get(sync_indexing_cache_key) @@ -2135,7 +2175,7 @@ class DocumentService: redis_client.setex(sync_indexing_cache_key, 600, 1) - sync_website_document_indexing_task.delay(dataset_id, document.id) + sync_website_document_indexing_task.delay(dataset.id, document.id) @staticmethod def get_documents_position(dataset_id, session: Session): @@ -2402,13 +2442,14 @@ class DocumentService: truncated_page_name, batch, ) + document.id = str(uuid.uuid4()) session.add(document) - session.flush() document_ids.append(document.id) documents.append(document) position += 1 else: exist_document.pop(page.page_id) + session.flush() # delete not selected documents if len(exist_document) > 0: clean_notion_document_task.delay(list(exist_document.values()), dataset.id) @@ -3313,6 +3354,7 @@ class DocumentService: # Only set async task and cache if document is currently enabled if document.enabled: + # pyrefly: ignore [bad-assignment] update_info["async_task"] = {"function": remove_document_from_index_task, "args": [document.id]} update_info["set_cache"] = True @@ -3333,6 +3375,7 @@ class DocumentService: # Only re-index if the document is currently enabled if document.enabled: + # pyrefly: ignore [bad-assignment] update_info["async_task"] = {"function": add_document_to_index_task, "args": [document.id]} update_info["set_cache"] = True @@ -3900,7 +3943,14 @@ class SegmentService: session.add(document) # Delete database records - session.execute(delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids))) + session.execute( + delete(DocumentSegment).where( + DocumentSegment.id.in_(segment_db_ids), + DocumentSegment.dataset_id == dataset.id, + DocumentSegment.document_id == document.id, + DocumentSegment.tenant_id == current_user.current_tenant_id, + ) + ) session.commit() @classmethod diff --git a/api/services/document_indexing_proxy/base.py b/api/services/document_indexing_proxy/base.py index 02df6752f30..91835c8abb2 100644 --- a/api/services/document_indexing_proxy/base.py +++ b/api/services/document_indexing_proxy/base.py @@ -4,7 +4,7 @@ from collections.abc import Callable from functools import cached_property from typing import Any, ClassVar -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from services.feature_service import FeatureService logger = logging.getLogger(__name__) diff --git a/api/services/enterprise/account_deletion_sync.py b/api/services/enterprise/account_deletion_sync.py index 89c4b80e670..57a306ed094 100644 --- a/api/services/enterprise/account_deletion_sync.py +++ b/api/services/enterprise/account_deletion_sync.py @@ -8,6 +8,7 @@ from sqlalchemy import select from sqlalchemy.orm import Session from configs import dify_config +from enums import DeploymentEdition from extensions.ext_redis import redis_client from models.account import TenantAccountJoin @@ -79,9 +80,9 @@ def sync_workspace_member_removal(workspace_id: str, member_id: str, *, source: source: Source of the sync request (e.g., "workspace_member_removed") Returns: - bool: True if task was queued (or skipped in community), False if queueing failed + bool: True if task was queued (or skipped outside the Enterprise edition), False if queueing failed """ - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return True return _queue_task(workspace_id=workspace_id, member_id=member_id, source=source) @@ -100,9 +101,9 @@ def sync_account_deletion(account_id: str, *, source: str, session: Session) -> session: SQLAlchemy session used to fetch workspace memberships Returns: - bool: True if all tasks were queued (or skipped in community), False if any queueing failed + bool: True if all tasks were queued (or skipped outside the Enterprise edition), False if any queueing failed """ - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return True # Fetch all workspaces the account belongs to diff --git a/api/services/enterprise/enterprise_service.py b/api/services/enterprise/enterprise_service.py index d7dbd12973d..e8c768f20be 100644 --- a/api/services/enterprise/enterprise_service.py +++ b/api/services/enterprise/enterprise_service.py @@ -4,12 +4,12 @@ import enum import logging import uuid from datetime import datetime -from typing import TYPE_CHECKING from cachetools.func import ttl_cache from pydantic import BaseModel, ConfigDict, Field, model_validator from configs import dify_config +from enums import DeploymentEdition from extensions.ext_redis import redis_client from services.enterprise.base import ( EnterpriseRequest, @@ -17,13 +17,11 @@ from services.enterprise.base import ( MCPNoRefreshTokenError, MCPTokenError, ) +from services.entities.feature_entities import LicenseStatus from services.errors.enterprise import ( EnterpriseServiceError, ) -if TYPE_CHECKING: - from services.feature_service import LicenseStatus - logger = logging.getLogger(__name__) DEFAULT_WORKSPACE_JOIN_TIMEOUT_SECONDS = 1.0 @@ -98,7 +96,7 @@ def try_join_default_workspace(account_id: str) -> None: This is a best-effort integration. Failures must not block user registration. """ - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return try: @@ -376,9 +374,9 @@ class EnterpriseService: caching, every request on an expired license would hit the enterprise API. Returns: - LicenseStatus enum value, or None if enterprise is disabled / unreachable. + LicenseStatus enum value, or None outside the Enterprise edition or when unreachable. """ - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return None cached = cls._read_cached_license_status() @@ -390,8 +388,6 @@ class EnterpriseService: @classmethod def _read_cached_license_status(cls) -> LicenseStatus | None: """Read license status from Redis cache, returning None on miss or failure.""" - from services.feature_service import LicenseStatus - try: raw = redis_client.get(LICENSE_STATUS_CACHE_KEY) if raw: @@ -404,8 +400,6 @@ class EnterpriseService: @classmethod def _fetch_and_cache_license_status(cls) -> LicenseStatus | None: """Fetch license status from enterprise API and cache the result.""" - from services.feature_service import LicenseStatus - try: info = cls.get_info() license_info = info.get("License") diff --git a/api/services/enterprise/rbac_service.py b/api/services/enterprise/rbac_service.py index 8a6842538ee..5bfefd8a916 100644 --- a/api/services/enterprise/rbac_service.py +++ b/api/services/enterprise/rbac_service.py @@ -830,10 +830,13 @@ class RBACService: tenant_id: str, account_id: str | None = None, include_owner: int | None = None, + biiling_enabled: bool | None = None, *, options: ListOption | None = None, ) -> Paginated[RBACRole]: - params = (options or ListOption()).to_params({"include_owner": include_owner}) + params = (options or ListOption()).to_params( + {"include_owner": include_owner, "biiling_enabled": biiling_enabled} + ) params["dataset_operator_enabled"] = dify_config.DATASET_OPERATOR_ENABLED data = _inner_call( "GET", @@ -869,13 +872,13 @@ class RBACService: ) @staticmethod - def get(tenant_id: str, account_id: str | None, role_id: str) -> RBACRole: + def get(tenant_id: str, account_id: str | None, role_id: str, billing_enabled: bool = True) -> RBACRole: data = _inner_call( "GET", f"{_INNER_PREFIX}/roles/item", tenant_id=tenant_id, account_id=account_id, - params={"id": role_id}, + params={"id": role_id, "billing_enabled": billing_enabled}, ) return RBACRole.model_validate(data or {}) diff --git a/api/services/entities/dsl_entities.py b/api/services/entities/dsl_entities.py index 0d5f49dc4cc..2bff3fcac24 100644 --- a/api/services/entities/dsl_entities.py +++ b/api/services/entities/dsl_entities.py @@ -18,6 +18,18 @@ class ImportStatus(StrEnum): FAILED = "failed" +class PendingImportOwner(BaseModel): + tenant_id: str | None = None + account_id: str | None = None + + def is_accessible_by(self, *, tenant_id: str | None, account_id: str) -> bool: + if tenant_id is None: + return False + owner = (self.tenant_id, self.account_id) + # Ownerless payloads come from older pods and expire after 10 minutes; #40106 removes this bridge. + return owner in ((None, None), (tenant_id, account_id)) + + class DslImportWarning(BaseModel): """Portable DSL reference that could not be restored in the target workspace.""" diff --git a/api/services/entities/external_knowledge_entities/external_knowledge_entities.py b/api/services/entities/external_knowledge_entities/external_knowledge_entities.py index 110dbe5a5e7..be4bd467731 100644 --- a/api/services/entities/external_knowledge_entities/external_knowledge_entities.py +++ b/api/services/entities/external_knowledge_entities/external_knowledge_entities.py @@ -1,6 +1,6 @@ from typing import Any, Literal, Union -from pydantic import BaseModel +from pydantic import BaseModel, Field class AuthorizationConfig(BaseModel): @@ -24,3 +24,19 @@ class ExternalKnowledgeApiSetting(BaseModel): request_method: str headers: dict[str, Any] | None = None params: dict[str, Any] | None = None + + +class ExternalDatasetCreatePayload(BaseModel): + """Validated fields required to create an external dataset binding. + + The console controller owns HTTP concerns, but the service also needs this + contract when creating the tenant-scoped dataset and external knowledge + binding. Keep it outside controllers so service imports do not depend on + Flask blueprint initialization. + """ + + external_knowledge_api_id: str + external_knowledge_id: str + name: str = Field(..., min_length=1, max_length=100) + description: str | None = Field(None, max_length=400) + external_retrieval_model: dict[str, object] | None = Field(default=None) diff --git a/api/services/entities/feature_entities.py b/api/services/entities/feature_entities.py new file mode 100644 index 00000000000..666943a5276 --- /dev/null +++ b/api/services/entities/feature_entities.py @@ -0,0 +1,205 @@ +"""Feature query results and policy values shared by their consumers.""" + +from enum import StrEnum + +from pydantic import BaseModel, ConfigDict, Field + +from enums import CloudPlan, DeploymentEdition + + +class FeatureResponseModel(BaseModel): + model_config = ConfigDict(json_schema_serialization_defaults_required=True, protected_namespaces=()) + + +class SubscriptionModel(FeatureResponseModel): + plan: CloudPlan = CloudPlan.SANDBOX + interval: str = "" + + +class BillingModel(FeatureResponseModel): + # Deprecated compatibility field. Deployment edition is the only source of truth for product edition. + # TODO: Remove after clients migrate to `SystemFeatureModel.deployment_edition`. + enabled: bool = Field( + default=False, + deprecated=True, + description="Deprecated. Use system features deployment_edition to determine the product edition.", + ) + subscription: SubscriptionModel = SubscriptionModel() + + +class EducationModel(FeatureResponseModel): + enabled: bool = False + activated: bool = False + + +class LimitationModel(FeatureResponseModel): + size: int = 0 + limit: int = 0 + + +class VectorSpaceLimitationModel(LimitationModel): + model_config = ConfigDict(json_schema_serialization_defaults_required=False, protected_namespaces=()) + + size: int + limit: int + usage_unknown: bool = Field(default=False, exclude_if=lambda value: not value) + + +class LicenseLimitationModel(FeatureResponseModel): + """ + - enabled: whether this limit is enforced + - size: current usage count + - limit: maximum allowed count; 0 means unlimited + """ + + enabled: bool = Field(False, description="Whether this limit is currently active") + size: int = Field(0, description="Number of resources already consumed") + limit: int = Field(0, description="Maximum number of resources allowed; 0 means no limit") + + def is_available(self, required: int = 1) -> bool: + """ + Determine whether the requested amount can be allocated. + + Returns True if: + - this limit is not active, or + - the limit is zero (unlimited), or + - there is enough remaining quota. + """ + if not self.enabled or self.limit == 0: + return True + + return (self.limit - self.size) >= required + + +class Quota(FeatureResponseModel): + usage: int = 0 + limit: int = 0 + reset_date: int = -1 + + +class LicenseStatus(StrEnum): + NONE = "none" + INACTIVE = "inactive" + ACTIVE = "active" + EXPIRING = "expiring" + EXPIRED = "expired" + LOST = "lost" + + +class LicenseStatusModel(FeatureResponseModel): + status: LicenseStatus = LicenseStatus.NONE + + +class LicenseModel(LicenseStatusModel): + expired_at: str = "" + workspaces: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0) + seats: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0) + license_expiry_notice_enabled: bool = False + + +class BrandingModel(FeatureResponseModel): + enabled: bool = False + application_title: str = "" + login_page_logo: str = "" + workspace_logo: str = "" + favicon: str = "" + + +class SSOProtocol(StrEnum): + SAML = "saml" + OIDC = "oidc" + OAUTH2 = "oauth2" + + +class WebAppAuthSSOModel(FeatureResponseModel): + protocol: SSOProtocol | None = None + + +class WebAppAuthModel(FeatureResponseModel): + enabled: bool = False + allow_sso: bool = False + sso_config: WebAppAuthSSOModel = Field(default_factory=WebAppAuthSSOModel) + allow_email_code_login: bool = False + allow_email_password_login: bool = False + allow_public_access: bool = True + + +class KnowledgePipeline(FeatureResponseModel): + publish_enabled: bool = False + + +class PluginInstallationScope(StrEnum): + NONE = "none" + OFFICIAL_ONLY = "official_only" + OFFICIAL_AND_SPECIFIC_PARTNERS = "official_and_specific_partners" + ALL = "all" + + +class PluginInstallationPermissionModel(FeatureResponseModel): + # Plugin installation scope – possible values: + # none: prohibit all plugin installations + # official_only: allow only Dify official plugins + # official_and_specific_partners: allow official and specific partner plugins + # all: allow installation of all plugins + plugin_installation_scope: PluginInstallationScope = PluginInstallationScope.ALL + + # If True, restrict plugin installation to the marketplace only + # Equivalent to ForceEnablePluginVerification + restrict_to_marketplace_only: bool = False + + +class FeatureModel(FeatureResponseModel): + billing: BillingModel = BillingModel() + education: EducationModel = EducationModel() + members: LimitationModel = LimitationModel(size=0, limit=1) + apps: LimitationModel = LimitationModel(size=0, limit=10) + vector_space: LimitationModel | None = LimitationModel(size=0, limit=5) + knowledge_rate_limit: int = 10 + annotation_quota_limit: LimitationModel = LimitationModel(size=0, limit=10) + documents_upload_quota: LimitationModel = LimitationModel(size=0, limit=50) + docs_processing: str = "standard" + can_replace_logo: bool = False + model_load_balancing_enabled: bool = False + dataset_operator_enabled: bool = False + webapp_copyright_enabled: bool = False + workspace_members: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0) + is_allow_transfer_workspace: bool = True + trigger_event: Quota = Quota(usage=0, limit=3000, reset_date=0) + api_rate_limit: Quota = Quota(usage=0, limit=5000, reset_date=0) + # Controls whether email delivery is allowed for HumanInput nodes. + human_input_email_delivery_enabled: bool = False + knowledge_pipeline: KnowledgePipeline = KnowledgePipeline() + next_credit_reset_date: int = 0 + + +class KnowledgeRateLimitModel(FeatureResponseModel): + enabled: bool = False + limit: int = 10 + subscription_plan: str = "" + + +class SystemFeatureModel(FeatureResponseModel): + """Non-sensitive bootstrap snapshot exposed before Console or Web authentication.""" + + deployment_edition: DeploymentEdition + enable_app_deploy: bool = False + sso_enforced_for_signin: bool = False + sso_enforced_for_signin_protocol: SSOProtocol | None = None + enable_marketplace: bool = False + enable_email_code_login: bool = False + enable_email_password_login: bool = True + enable_social_oauth_login: bool = False + enable_collaboration_mode: bool = True + is_allow_register: bool = False + is_email_setup: bool = False + license: LicenseStatusModel = LicenseStatusModel() + branding: BrandingModel = BrandingModel() + webapp_auth: WebAppAuthModel = Field(default_factory=WebAppAuthModel) + plugin_installation_permission: PluginInstallationPermissionModel = PluginInstallationPermissionModel() + enable_change_email: bool = True + enable_creators_platform: bool = False + enable_explore_banner: bool = False + enable_learn_app: bool = True + enable_step_by_step_tour: bool = False + rbac_enabled: bool = False + knowledge_fs_enabled: bool = False diff --git a/api/services/entities/knowledge_entities/knowledge_entities.py b/api/services/entities/knowledge_entities/knowledge_entities.py index 8cd99c4a98d..b5d74436fd6 100644 --- a/api/services/entities/knowledge_entities/knowledge_entities.py +++ b/api/services/entities/knowledge_entities/knowledge_entities.py @@ -6,6 +6,7 @@ from core.rag.entities import Rule from core.rag.entities.metadata_entities import MetadataFilteringCondition from core.rag.index_processor.constant.index_type import IndexStructureType from core.rag.retrieval.retrieval_methods import RetrievalMethod +from libs.helper import UUIDStr from models.enums import ProcessRuleMode DocForm = Annotated[ @@ -283,7 +284,7 @@ class MetadataUpdateArgs(BaseModel): class MetadataDetail(BaseModel): - id: str = Field(description="Metadata field ID.") + id: UUIDStr = Field(description="Metadata field ID.") name: str = Field(description="Metadata field name.") value: str | int | float | None = Field( default=None, @@ -292,7 +293,7 @@ class MetadataDetail(BaseModel): class DocumentMetadataOperation(BaseModel): - document_id: str = Field(description="Document ID whose metadata should be updated.") + document_id: UUIDStr = Field(description="Document ID whose metadata should be updated.") metadata_list: list[MetadataDetail] = Field(description="Metadata fields to update.") partial_update: bool = Field( default=False, diff --git a/api/services/entities/model_provider_entities.py b/api/services/entities/model_provider_entities.py index 020dc4a2ea9..8a9e8dd66b3 100644 --- a/api/services/entities/model_provider_entities.py +++ b/api/services/entities/model_provider_entities.py @@ -17,6 +17,7 @@ from core.entities.provider_entities import ( QuotaConfiguration, UnaddedModelConfiguration, ) +from core.plugin.entities.plugin import PluginInstallationSource from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import ( FetchFrom, @@ -69,6 +70,67 @@ class SystemConfigurationResponse(BaseModel): quota_configurations: list[QuotaConfiguration] = [] +class ModelProviderCustomConfigurationSummaryResponse(BaseModel): + status: CustomConfigurationStatus + has_custom_models: bool = Field( + description="Whether custom model configuration exists, including saved model credentials." + ) + available_credentials: list[CredentialConfiguration] + current_credential_id: str | None = None + current_credential_name: str | None = None + current_credential_usable: bool + + +class ModelProviderSystemConfigurationSummaryResponse(BaseModel): + enabled: bool + + +class ModelProviderPluginSummaryResponse(BaseModel): + installation_id: str + plugin_id: str + plugin_unique_identifier: str + runtime_type: str + source: PluginInstallationSource + version: str + + +class ModelProviderSummaryResponse(BaseModel): + """Fields required to render the collapsed model-provider list.""" + + tenant_id: str = Field(exclude=True) + provider: str + plugin_id: str + label: I18nObject + description: I18nObject | None = None + icon_small: I18nObject | None = None + icon_small_dark: I18nObject | None = None + supported_model_types: Sequence[ModelType] + configurate_methods: list[ConfigurateMethod] + preferred_provider_type: ProviderType + is_configured: bool + custom_configuration: ModelProviderCustomConfigurationSummaryResponse + system_configuration: ModelProviderSystemConfigurationSummaryResponse + + model_config = ConfigDict(protected_namespaces=()) + + @model_validator(mode="after") + def build_icon_urls(self): + url_prefix = ( + dify_config.CONSOLE_API_URL + f"/console/api/workspaces/{self.tenant_id}/model-providers/{self.provider}" + ) + if self.icon_small is not None: + self.icon_small = I18nObject( + en_US=f"{url_prefix}/icon_small/en_US", + zh_Hans=f"{url_prefix}/icon_small/zh_Hans", + ) + if self.icon_small_dark is not None: + self.icon_small_dark = I18nObject( + en_US=f"{url_prefix}/icon_small_dark/en_US", + zh_Hans=f"{url_prefix}/icon_small_dark/zh_Hans", + ) + return self + + class ProviderResponse(BaseModel): """ Model class for provider response. diff --git a/api/services/errors/metadata.py b/api/services/errors/metadata.py new file mode 100644 index 00000000000..97b864484b1 --- /dev/null +++ b/api/services/errors/metadata.py @@ -0,0 +1,2 @@ +class MetadataResourceNotFoundError(Exception): + pass diff --git a/api/services/errors/rag_pipeline.py b/api/services/errors/rag_pipeline.py new file mode 100644 index 00000000000..bce31a214a6 --- /dev/null +++ b/api/services/errors/rag_pipeline.py @@ -0,0 +1,2 @@ +class RagPipelineResourceNotFoundError(Exception): + pass diff --git a/api/services/explore_banner_query_service.py b/api/services/explore_banner_query_service.py new file mode 100644 index 00000000000..5a64245bb94 --- /dev/null +++ b/api/services/explore_banner_query_service.py @@ -0,0 +1,44 @@ +"""Application service for listing Explore banners. + +ExploreBanner is the legacy contract name shared by the API, feature flag, and database model. +""" + +from collections.abc import Callable, Sequence +from datetime import datetime +from typing import Any, NamedTuple, Protocol + +_DEFAULT_LANGUAGE = "en-US" + + +class ExploreBannerRecord(NamedTuple): + id: str + content: dict[str, Any] + link: str + sort: int + status: str + created_at: datetime + + +class ExploreBannerQuery(Protocol): + def list_enabled(self, language: str) -> Sequence[ExploreBannerRecord]: ... + + +class ExploreBannerQueryService: + def __init__( + self, + *, + banners: ExploreBannerQuery, + is_enabled: Callable[[], bool], + ) -> None: + self._banners = banners + self._is_enabled = is_enabled + + def list_for_language(self, language: str) -> tuple[ExploreBannerRecord, ...]: + if not self._is_enabled(): + return () + + banners = tuple(self._banners.list_enabled(language)) + if banners or language == _DEFAULT_LANGUAGE: + return banners + + return tuple(self._banners.list_enabled(_DEFAULT_LANGUAGE)) diff --git a/api/services/external_knowledge_service.py b/api/services/external_knowledge_service.py index 1a15fee1e97..03fe274fe89 100644 --- a/api/services/external_knowledge_service.py +++ b/api/services/external_knowledge_service.py @@ -19,8 +19,10 @@ from models.dataset import ( ExternalKnowledgeApis, ExternalKnowledgeBindings, ) +from services.enterprise import rbac_service as enterprise_rbac_service from services.entities.external_knowledge_entities.external_knowledge_entities import ( Authorization, + ExternalDatasetCreatePayload, ExternalKnowledgeApiSetting, ) from services.errors.dataset import DatasetNameDuplicateError @@ -114,7 +116,7 @@ class ExternalDatasetService: if response.status_code == 404: raise ValueError(f"Not Found: failed to connect to the endpoint: {endpoint}") if response.status_code == 403: - raise ValueError(f"Forbidden: Authorization failed with api_key: {api_key}") + raise ValueError("Forbidden: Authorization failed with the provided api_key") @staticmethod def get_external_knowledge_api( @@ -286,16 +288,17 @@ class ExternalDatasetService: return ExternalKnowledgeApiSetting.model_validate(settings) @staticmethod - def create_external_dataset(tenant_id: str, user_id: str, args: dict[str, Any], *, session: Session) -> Dataset: + def create_external_dataset( + tenant_id: str, user_id: str, args: ExternalDatasetCreatePayload, *, session: Session + ) -> Dataset: + """Create a tenant-scoped external dataset and binding in the caller's transaction.""" # check if dataset name already exists - if session.scalar( - select(Dataset).where(Dataset.name == args.get("name"), Dataset.tenant_id == tenant_id).limit(1) - ): - raise DatasetNameDuplicateError(f"Dataset with name {args.get('name')} already exists.") + if session.scalar(select(Dataset).where(Dataset.name == args.name, Dataset.tenant_id == tenant_id).limit(1)): + raise DatasetNameDuplicateError(f"Dataset with name {args.name} already exists.") external_knowledge_api = session.scalar( select(ExternalKnowledgeApis) .where( - ExternalKnowledgeApis.id == args.get("external_knowledge_api_id"), + ExternalKnowledgeApis.id == args.external_knowledge_api_id, ExternalKnowledgeApis.tenant_id == tenant_id, ) .limit(1) @@ -306,31 +309,33 @@ class ExternalDatasetService: dataset = Dataset( tenant_id=tenant_id, - name=args.get("name"), - description=args.get("description", ""), + name=args.name, + description=args.description or "", provider="external", - retrieval_model=args.get("external_retrieval_model"), + retrieval_model=args.external_retrieval_model, created_by=user_id, maintainer=user_id, ) session.add(dataset) session.flush() - if args.get("external_knowledge_id") is None: - raise ValueError("external_knowledge_id is required") - if args.get("external_knowledge_api_id") is None: - raise ValueError("external_knowledge_api_id is required") external_knowledge_binding = ExternalKnowledgeBindings( tenant_id=tenant_id, dataset_id=dataset.id, - external_knowledge_api_id=args.get("external_knowledge_api_id") or "", - external_knowledge_id=args.get("external_knowledge_id") or "", + external_knowledge_api_id=args.external_knowledge_api_id or "", + external_knowledge_id=args.external_knowledge_id or "", created_by=user_id, ) session.add(external_knowledge_binding) session.commit() + enterprise_rbac_service.try_sync_creator_access_policy_member_bindings( + tenant_id, + user_id, + enterprise_rbac_service.RBACResourceType.DATASET, + dataset.id, + ) return dataset diff --git a/api/services/feature_query_service.py b/api/services/feature_query_service.py new file mode 100644 index 00000000000..f24e387bb0f --- /dev/null +++ b/api/services/feature_query_service.py @@ -0,0 +1,62 @@ +"""Application service for feature queries exposed by API adapters.""" + +from collections.abc import Sequence +from typing import Protocol + +from machinery.context import RequestContext +from services.entities.feature_entities import ( + FeatureModel, + LicenseModel, + SystemFeatureModel, + VectorSpaceLimitationModel, +) + + +class FeatureQueryGateway(Protocol): + """Read dynamic feature resources without exposing their current implementation.""" + + def get_workspace_features(self, workspace_id: str) -> FeatureModel: ... + + def get_vector_space(self, workspace_id: str) -> VectorSpaceLimitationModel: ... + + def get_public_system_features(self) -> SystemFeatureModel: ... + + def get_license(self) -> LicenseModel: ... + + +class FeatureQueryService: + def __init__( + self, + *, + features: FeatureQueryGateway, + trial_models: Sequence[str], + app_dsl_version: str, + ) -> None: + self._features = features + self._trial_models = tuple(trial_models) + self._app_dsl_version = app_dsl_version + + def get_features(self, context: RequestContext) -> FeatureModel: + return self._features.get_workspace_features(self._require_active_workspace(context)) + + def get_vector_space(self, context: RequestContext) -> VectorSpaceLimitationModel: + return self._features.get_vector_space(self._require_active_workspace(context)) + + def get_trial_models(self) -> list[str]: + return list(self._trial_models) + + def get_app_dsl_version(self) -> str: + return self._app_dsl_version + + def get_system_features(self) -> SystemFeatureModel: + return self._features.get_public_system_features() + + def get_license(self) -> LicenseModel: + return self._features.get_license() + + @staticmethod + def _require_active_workspace(context: RequestContext) -> str: + workspace_id = context.active_workspace_id + if workspace_id is None: + raise RuntimeError("Console account admission did not resolve an active workspace") + return workspace_id diff --git a/api/services/feature_service.py b/api/services/feature_service.py index 6225562ec25..b32387ffea0 100644 --- a/api/services/feature_service.py +++ b/api/services/feature_service.py @@ -1,211 +1,41 @@ -from enum import StrEnum +import logging +from collections.abc import Mapping -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, ValidationError from configs import dify_config -from constants.dsl_version import CURRENT_APP_DSL_VERSION -from enums.cloud_plan import CloudPlan -from enums.deployment_edition import DeploymentEdition -from enums.hosted_provider import HostedTrialProvider +from enums import CloudPlan, DeploymentEdition, HostedTrialProvider from services.billing_service import BillingInfo, BillingService from services.enterprise.enterprise_service import EnterpriseService +from services.entities import feature_entities + +logger = logging.getLogger(__name__) -class FeatureResponseModel(BaseModel): - model_config = ConfigDict(json_schema_serialization_defaults_required=True, protected_namespaces=()) +class _EnterprisePluginInstallationPermission(BaseModel): + model_config = ConfigDict(extra="ignore") - -class SubscriptionModel(FeatureResponseModel): - plan: str = CloudPlan.SANDBOX - interval: str = "" - - -class BillingModel(FeatureResponseModel): - enabled: bool = False - subscription: SubscriptionModel = SubscriptionModel() - - -class EducationModel(FeatureResponseModel): - enabled: bool = False - activated: bool = False - - -class LimitationModel(FeatureResponseModel): - size: int = 0 - limit: int = 0 - - -class LicenseLimitationModel(FeatureResponseModel): - """ - - enabled: whether this limit is enforced - - size: current usage count - - limit: maximum allowed count; 0 means unlimited - """ - - enabled: bool = Field(False, description="Whether this limit is currently active") - size: int = Field(0, description="Number of resources already consumed") - limit: int = Field(0, description="Maximum number of resources allowed; 0 means no limit") - - def is_available(self, required: int = 1) -> bool: - """ - Determine whether the requested amount can be allocated. - - Returns True if: - - this limit is not active, or - - the limit is zero (unlimited), or - - there is enough remaining quota. - """ - if not self.enabled or self.limit == 0: - return True - - return (self.limit - self.size) >= required - - -class Quota(FeatureResponseModel): - usage: int = 0 - limit: int = 0 - reset_date: int = -1 - - -class LicenseStatus(StrEnum): - NONE = "none" - INACTIVE = "inactive" - ACTIVE = "active" - EXPIRING = "expiring" - EXPIRED = "expired" - LOST = "lost" - - -class LicenseStatusModel(FeatureResponseModel): - status: LicenseStatus = LicenseStatus.NONE - - -class LicenseModel(LicenseStatusModel): - expired_at: str = "" - workspaces: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0) - seats: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0) - - -class BrandingModel(FeatureResponseModel): - enabled: bool = False - application_title: str = "" - login_page_logo: str = "" - workspace_logo: str = "" - favicon: str = "" - - -class WebAppAuthSSOModel(FeatureResponseModel): - protocol: str = "" - - -class WebAppAuthModel(FeatureResponseModel): - enabled: bool = False - allow_sso: bool = False - sso_config: WebAppAuthSSOModel = WebAppAuthSSOModel() - allow_email_code_login: bool = False - allow_email_password_login: bool = False - allow_public_access: bool = True - - -class KnowledgePipeline(FeatureResponseModel): - publish_enabled: bool = False - - -class PluginInstallationScope(StrEnum): - NONE = "none" - OFFICIAL_ONLY = "official_only" - OFFICIAL_AND_SPECIFIC_PARTNERS = "official_and_specific_partners" - ALL = "all" - - -class PluginInstallationPermissionModel(FeatureResponseModel): - # Plugin installation scope – possible values: - # none: prohibit all plugin installations - # official_only: allow only Dify official plugins - # official_and_specific_partners: allow official and specific partner plugins - # all: allow installation of all plugins - plugin_installation_scope: PluginInstallationScope = PluginInstallationScope.ALL - - # If True, restrict plugin installation to the marketplace only - # Equivalent to ForceEnablePluginVerification - restrict_to_marketplace_only: bool = False - - -class FeatureModel(FeatureResponseModel): - billing: BillingModel = BillingModel() - education: EducationModel = EducationModel() - members: LimitationModel = LimitationModel(size=0, limit=1) - apps: LimitationModel = LimitationModel(size=0, limit=10) - vector_space: LimitationModel | None = LimitationModel(size=0, limit=5) - knowledge_rate_limit: int = 10 - annotation_quota_limit: LimitationModel = LimitationModel(size=0, limit=10) - documents_upload_quota: LimitationModel = LimitationModel(size=0, limit=50) - docs_processing: str = "standard" - can_replace_logo: bool = False - model_load_balancing_enabled: bool = False - dataset_operator_enabled: bool = False - webapp_copyright_enabled: bool = False - workspace_members: LicenseLimitationModel = LicenseLimitationModel(enabled=False, size=0, limit=0) - is_allow_transfer_workspace: bool = True - trigger_event: Quota = Quota(usage=0, limit=3000, reset_date=0) - api_rate_limit: Quota = Quota(usage=0, limit=5000, reset_date=0) - # Controls whether email delivery is allowed for HumanInput nodes. - human_input_email_delivery_enabled: bool = False - knowledge_pipeline: KnowledgePipeline = KnowledgePipeline() - next_credit_reset_date: int = 0 - - -class KnowledgeRateLimitModel(FeatureResponseModel): - enabled: bool = False - limit: int = 10 - subscription_plan: str = "" - - -class SystemFeatureModel(FeatureResponseModel): - """Non-sensitive bootstrap snapshot exposed before Console or Web authentication.""" - - deployment_edition: DeploymentEdition - enable_app_deploy: bool = False - sso_enforced_for_signin: bool = False - sso_enforced_for_signin_protocol: str = "" - enable_marketplace: bool = False - enable_email_code_login: bool = False - enable_email_password_login: bool = True - enable_social_oauth_login: bool = False - enable_collaboration_mode: bool = True - is_allow_register: bool = False - is_email_setup: bool = False - license: LicenseStatusModel = LicenseStatusModel() - branding: BrandingModel = BrandingModel() - webapp_auth: WebAppAuthModel = WebAppAuthModel() - plugin_installation_permission: PluginInstallationPermissionModel = PluginInstallationPermissionModel() - enable_change_email: bool = True - enable_creators_platform: bool = False - enable_trial_app: bool = False - enable_explore_banner: bool = False - enable_learn_app: bool = True - enable_step_by_step_tour: bool = False - rbac_enabled: bool = False - knowledge_fs_enabled: bool = False + plugin_installation_scope: feature_entities.PluginInstallationScope = Field(alias="pluginInstallationScope") + restrict_to_marketplace_only: bool = Field(alias="restrictToMarketplaceOnly", strict=True) class FeatureService: @classmethod - def get_features(cls, tenant_id: str, exclude_vector_space: bool = False) -> FeatureModel: - features = FeatureModel() + def get_features(cls, tenant_id: str, exclude_vector_space: bool = False) -> feature_entities.FeatureModel: + features = feature_entities.FeatureModel() if exclude_vector_space: features.vector_space = None cls._fulfill_params_from_env(features) - if dify_config.BILLING_ENABLED and tenant_id: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and tenant_id: cls._fulfill_params_from_billing_api( features, tenant_id, exclude_vector_space=exclude_vector_space, ) - if dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE: features.webapp_copyright_enabled = True features.knowledge_pipeline.publish_enabled = True cls._fulfill_params_from_workspace_info(features, tenant_id) @@ -218,21 +48,22 @@ class FeatureService: return features @classmethod - def get_vector_space(cls, tenant_id: str) -> LimitationModel: - vector_space = LimitationModel(size=0, limit=5) - if dify_config.BILLING_ENABLED and tenant_id: + def get_vector_space(cls, tenant_id: str) -> feature_entities.VectorSpaceLimitationModel: + vector_space = feature_entities.VectorSpaceLimitationModel(size=0, limit=5) + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and tenant_id: billing_vector_space = BillingService.get_vector_space(tenant_id) # NOTE: billing API returns vector_space.size as float (e.g. 0.0), # but feature API keeps LimitationModel.size as int for compatibility. vector_space.size = int(billing_vector_space["size"]) vector_space.limit = billing_vector_space["limit"] + vector_space.usage_unknown = billing_vector_space.get("usage_unknown", False) return vector_space @classmethod def get_knowledge_rate_limit(cls, tenant_id: str): - knowledge_rate_limit = KnowledgeRateLimitModel() - if dify_config.BILLING_ENABLED and tenant_id: + knowledge_rate_limit = feature_entities.KnowledgeRateLimitModel() + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and tenant_id: knowledge_rate_limit.enabled = True limit_info = BillingService.get_knowledge_rate_limit(tenant_id) knowledge_rate_limit.limit = limit_info.get("limit", 10) @@ -240,8 +71,25 @@ class FeatureService: return knowledge_rate_limit @classmethod - def _resolve_human_input_email_delivery_enabled(cls, *, features: FeatureModel, tenant_id: str | None) -> bool: - if dify_config.ENTERPRISE_ENABLED or not dify_config.BILLING_ENABLED: + def get_knowledge_file_size_limit(cls, tenant_id: str | None) -> int: + default_limit = dify_config.UPLOAD_FILE_SIZE_LIMIT + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD or not tenant_id: + return default_limit + + billing_info = BillingService.get_info(tenant_id, exclude_vector_space=True) + if billing_info["enabled"] and billing_info["subscription"]["plan"] in ( + CloudPlan.PROFESSIONAL, + CloudPlan.TEAM, + ): + return max(default_limit, dify_config.KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN) + + return default_limit + + @classmethod + def _resolve_human_input_email_delivery_enabled( + cls, *, features: feature_entities.FeatureModel, tenant_id: str | None + ) -> bool: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return True if not tenant_id: return False @@ -251,13 +99,13 @@ class FeatureService: ) @classmethod - def get_system_features(cls) -> SystemFeatureModel: - system_features = SystemFeatureModel(deployment_edition=dify_config.DEPLOYMENT_EDITION) + def get_system_features(cls) -> feature_entities.SystemFeatureModel: + system_features = feature_entities.SystemFeatureModel(deployment_edition=dify_config.DEPLOYMENT_EDITION) system_features.rbac_enabled = dify_config.RBAC_ENABLED cls._fulfill_system_params_from_env(system_features) - if dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE: system_features.branding.enabled = True system_features.webapp_auth.enabled = True system_features.enable_change_email = False @@ -275,7 +123,7 @@ class FeatureService: def is_workspace_creation_allowed(cls) -> bool: """Resolve the backend workspace-creation policy, including the Enterprise override.""" is_allowed = dify_config.ALLOW_CREATE_WORKSPACE - if not dify_config.ENTERPRISE_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: return is_allowed enterprise_info = EnterpriseService.get_info() @@ -284,25 +132,35 @@ class FeatureService: @classmethod def is_plugin_manager_enabled(cls) -> bool: """Return whether Enterprise plugin credential policies must be enforced.""" - return dify_config.ENTERPRISE_ENABLED + return dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE @classmethod - def get_license(cls) -> LicenseModel: + def get_plugin_installation_permission(cls) -> feature_entities.PluginInstallationPermissionModel: + """Resolve the validated deployment-wide plugin installation policy.""" + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: + return feature_entities.PluginInstallationPermissionModel() + + return cls._resolve_plugin_installation_permission(EnterpriseService.get_info()) + + @classmethod + def get_license(cls) -> feature_entities.LicenseModel: """Return full license detail. Enterprise-only; requires an authenticated caller. Non-enterprise deployments have no license, so an unconstrained default (unlimited seats/workspaces) is returned. """ - if not dify_config.ENTERPRISE_ENABLED: - return LicenseModel() - return cls._build_license(EnterpriseService.get_info()) + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.ENTERPRISE: + return feature_entities.LicenseModel() + license_model = cls._build_license(EnterpriseService.get_info()) + license_model.license_expiry_notice_enabled = dify_config.ENABLE_LICENSE_EXPIRY_NOTICE + return license_model + + @staticmethod + def is_explore_banner_enabled() -> bool: + return dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and dify_config.ENABLE_EXPLORE_BANNER @classmethod - def get_app_dsl_version(cls) -> str: - return CURRENT_APP_DSL_VERSION - - @classmethod - def _fulfill_system_params_from_env(cls, system_features: SystemFeatureModel): + def _fulfill_system_params_from_env(cls, system_features: feature_entities.SystemFeatureModel): system_features.enable_email_code_login = dify_config.ENABLE_EMAIL_CODE_LOGIN system_features.enable_email_password_login = dify_config.ENABLE_EMAIL_PASSWORD_LOGIN system_features.enable_social_oauth_login = dify_config.ENABLE_SOCIAL_OAUTH_LOGIN @@ -310,8 +168,7 @@ class FeatureService: system_features.is_allow_register = dify_config.ALLOW_REGISTER system_features.is_email_setup = dify_config.MAIL_TYPE is not None and dify_config.MAIL_TYPE != "" system_features.enable_change_email = dify_config.ENABLE_CHANGE_EMAIL - system_features.enable_trial_app = dify_config.ENABLE_TRIAL_APP - system_features.enable_explore_banner = dify_config.ENABLE_EXPLORE_BANNER + system_features.enable_explore_banner = cls.is_explore_banner_enabled() 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 @@ -334,14 +191,14 @@ class FeatureService: return cls._fulfill_trial_models_from_env() @classmethod - def _fulfill_params_from_env(cls, features: FeatureModel): + def _fulfill_params_from_env(cls, features: feature_entities.FeatureModel): features.can_replace_logo = dify_config.CAN_REPLACE_LOGO features.model_load_balancing_enabled = dify_config.MODEL_LB_ENABLED features.dataset_operator_enabled = dify_config.DATASET_OPERATOR_ENABLED features.education.enabled = dify_config.EDUCATION_ENABLED @classmethod - def _fulfill_params_from_workspace_info(cls, features: FeatureModel, tenant_id: str): + def _fulfill_params_from_workspace_info(cls, features: feature_entities.FeatureModel, tenant_id: str): workspace_info = EnterpriseService.get_workspace_info(tenant_id) if "WorkspaceMembers" in workspace_info: features.workspace_members.size = workspace_info["WorkspaceMembers"]["used"] @@ -351,7 +208,7 @@ class FeatureService: @classmethod def _fulfill_params_from_billing_api( cls, - features: FeatureModel, + features: feature_entities.FeatureModel, tenant_id: str, exclude_vector_space: bool = False, ): @@ -363,7 +220,7 @@ class FeatureService: features_usage_info = BillingService.get_quota_info(tenant_id) features.billing.enabled = billing_info["enabled"] - features.billing.subscription.plan = billing_info["subscription"]["plan"] + features.billing.subscription.plan = CloudPlan(billing_info["subscription"]["plan"]) features.billing.subscription.interval = billing_info["subscription"]["interval"] features.education.activated = billing_info["subscription"].get("education", False) @@ -425,7 +282,9 @@ class FeatureService: features.next_credit_reset_date = billing_info["next_credit_reset_date"] @classmethod - def _fulfill_vector_space_from_billing_info(cls, vector_space: LimitationModel, billing_info: BillingInfo): + def _fulfill_vector_space_from_billing_info( + cls, vector_space: feature_entities.LimitationModel, billing_info: BillingInfo + ): if "vector_space" not in billing_info: return @@ -435,19 +294,21 @@ class FeatureService: vector_space.limit = billing_info["vector_space"]["limit"] @classmethod - def _build_license(cls, enterprise_info: dict) -> LicenseModel: - license_model = LicenseModel() + def _build_license(cls, enterprise_info: dict) -> feature_entities.LicenseModel: + license_model = feature_entities.LicenseModel() if license_info := enterprise_info.get("License"): - license_model.status = LicenseStatus(license_info.get("status", LicenseStatus.INACTIVE)) + license_model.status = feature_entities.LicenseStatus( + license_info.get("status", feature_entities.LicenseStatus.INACTIVE) + ) license_model.expired_at = license_info.get("expiredAt", "") if workspaces_info := license_info.get("workspaces"): - license_model.workspaces = LicenseLimitationModel( + license_model.workspaces = feature_entities.LicenseLimitationModel( enabled=workspaces_info.get("enabled", False), limit=workspaces_info.get("limit", 0), size=workspaces_info.get("used", 0), ) if seats_info := license_info.get("licensedSeats"): - license_model.seats = LicenseLimitationModel( + license_model.seats = feature_entities.LicenseLimitationModel( enabled=seats_info.get("enabled", False), limit=seats_info.get("limit", 0), size=seats_info.get("used", 0), @@ -455,14 +316,60 @@ class FeatureService: return license_model @classmethod - def _fulfill_params_from_enterprise(cls, features: SystemFeatureModel): + def _resolve_plugin_installation_permission( + cls, enterprise_info: Mapping[str, object] + ) -> feature_entities.PluginInstallationPermissionModel: + if "PluginInstallationPermission" not in enterprise_info: + return feature_entities.PluginInstallationPermissionModel() + + try: + permission = _EnterprisePluginInstallationPermission.model_validate( + enterprise_info["PluginInstallationPermission"] + ) + except ValidationError as exc: + # Do not attach the exception because it may contain raw Enterprise configuration values. + logger.error( # noqa: TRY400 + "Invalid Enterprise plugin installation permission; denying all plugin installations: %s", + exc.errors(include_input=False), + ) + return feature_entities.PluginInstallationPermissionModel( + plugin_installation_scope=feature_entities.PluginInstallationScope.NONE, + restrict_to_marketplace_only=True, + ) + + return feature_entities.PluginInstallationPermissionModel( + plugin_installation_scope=permission.plugin_installation_scope, + restrict_to_marketplace_only=permission.restrict_to_marketplace_only, + ) + + @staticmethod + def _resolve_sso_protocol(value: object, *, field_name: str) -> feature_entities.SSOProtocol | None: + if value is None or (isinstance(value, str) and not value.strip()): + return None + + if not isinstance(value, str): + logger.error("Invalid Enterprise SSO protocol for %s; disabling the protocol", field_name) + return None + + try: + return feature_entities.SSOProtocol(value) + except ValueError: + logger.error( # noqa: TRY400 + "Invalid Enterprise SSO protocol for %s; disabling the protocol", field_name + ) + return None + + @classmethod + def _fulfill_params_from_enterprise(cls, features: feature_entities.SystemFeatureModel): enterprise_info = EnterpriseService.get_info() if "SSOEnforcedForSignin" in enterprise_info: features.sso_enforced_for_signin = enterprise_info["SSOEnforcedForSignin"] - if "SSOEnforcedForSigninProtocol" in enterprise_info: - features.sso_enforced_for_signin_protocol = enterprise_info["SSOEnforcedForSigninProtocol"] + features.sso_enforced_for_signin_protocol = cls._resolve_sso_protocol( + enterprise_info.get("SSOEnforcedForSigninProtocol"), + field_name="SSOEnforcedForSigninProtocol", + ) if "EnableEmailCodeLogin" in enterprise_info: features.enable_email_code_login = enterprise_info["EnableEmailCodeLogin"] @@ -490,22 +397,20 @@ class FeatureService: features.webapp_auth.allow_email_password_login = enterprise_info["WebAppAuth"].get( "allowEmailPasswordLogin", False ) - features.webapp_auth.sso_config.protocol = enterprise_info.get("SSOEnforcedForWebProtocol", "") + features.webapp_auth.sso_config.protocol = cls._resolve_sso_protocol( + enterprise_info.get("SSOEnforcedForWebProtocol"), + field_name="SSOEnforcedForWebProtocol", + ) # SECURITY NOTE: system-features is unauthenticated, so it exposes only license # *status* — enough for the login page to detect an expired/inactive license after # force-logout. Full license detail (expiry, workspace/seat usage) is served # separately by get_license() behind an authenticated endpoint. if license_info := enterprise_info.get("License"): - features.license = LicenseStatusModel( - status=LicenseStatus(license_info.get("status", LicenseStatus.INACTIVE)) + features.license = feature_entities.LicenseStatusModel( + status=feature_entities.LicenseStatus( + license_info.get("status", feature_entities.LicenseStatus.INACTIVE) + ) ) - if "PluginInstallationPermission" in enterprise_info: - plugin_installation_info = enterprise_info["PluginInstallationPermission"] - features.plugin_installation_permission.plugin_installation_scope = plugin_installation_info[ - "pluginInstallationScope" - ] - features.plugin_installation_permission.restrict_to_marketplace_only = plugin_installation_info[ - "restrictToMarketplaceOnly" - ] + features.plugin_installation_permission = cls._resolve_plugin_installation_permission(enterprise_info) diff --git a/api/services/feature_service_gateway.py b/api/services/feature_service_gateway.py new file mode 100644 index 00000000000..3afa5096442 --- /dev/null +++ b/api/services/feature_service_gateway.py @@ -0,0 +1,32 @@ +"""Feature-query gateway backed by the existing FeatureService.""" + +from typing import override + +from services.entities.feature_entities import ( + FeatureModel, + LicenseModel, + SystemFeatureModel, + VectorSpaceLimitationModel, +) +from services.feature_query_service import FeatureQueryGateway +from services.feature_service import FeatureService + + +class FeatureServiceGateway(FeatureQueryGateway): + """Read dynamic feature resources through FeatureService.""" + + @override + def get_workspace_features(self, workspace_id: str) -> FeatureModel: + return FeatureService.get_features(workspace_id, exclude_vector_space=True) + + @override + def get_vector_space(self, workspace_id: str) -> VectorSpaceLimitationModel: + return FeatureService.get_vector_space(workspace_id) + + @override + def get_public_system_features(self) -> SystemFeatureModel: + return FeatureService.get_system_features() + + @override + def get_license(self) -> LicenseModel: + return FeatureService.get_license() diff --git a/api/services/feedback_service.py b/api/services/feedback_service.py index 6e60a026ce2..af231394953 100644 --- a/api/services/feedback_service.py +++ b/api/services/feedback_service.py @@ -1,7 +1,7 @@ import csv import io import json -from datetime import datetime +from datetime import datetime, timedelta from flask import Response from sqlalchemy import or_, select @@ -75,8 +75,8 @@ class FeedbackService: if end_date: try: - end_dt = datetime.strptime(end_date, "%Y-%m-%d") - stmt = stmt.where(MessageFeedback.created_at <= end_dt) + end_dt = datetime.strptime(end_date, "%Y-%m-%d") + timedelta(days=1) + stmt = stmt.where(MessageFeedback.created_at < end_dt) except ValueError: raise ValueError(f"Invalid end_date format: {end_date}. Use YYYY-MM-DD") diff --git a/api/services/file_request_service.py b/api/services/file_request_service.py index a795595d095..66b56684073 100644 --- a/api/services/file_request_service.py +++ b/api/services/file_request_service.py @@ -1,9 +1,9 @@ """Service helpers for trusted file request control-plane endpoints. -These helpers are used by inner APIs that return signed upload/download URLs to -trusted external runtimes such as ``dify-agent``. They do not transfer file -bytes themselves; they only rebuild access-scoped ``graphon.file.File`` values -and resolve the signed URL that the caller should use directly. +These helpers are used by inner APIs that allocate file access for trusted +external runtimes. They rebuild access-scoped ``graphon.file.File`` values and +return origin-free signed URIs so each transport adapter can select its own +network origin without signing the file twice. """ from __future__ import annotations @@ -14,30 +14,32 @@ from typing import Any from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom from core.app.file_access import DatabaseFileAccessController, FileAccessScope, bind_file_access_scope +from core.app.workflow.file_runtime import DifyWorkflowFileRuntime from factories.file_factory.builders import build_from_mapping from graphon.file import File -from graphon.file import helpers as file_helpers @dataclass(frozen=True, slots=True) class DownloadFileRequestResult: - """Resolved metadata and signed URL returned to trusted download callers.""" + """Resolved metadata and signed URI returned to trusted download callers.""" filename: str mime_type: str | None size: int - download_url: str + download_uri: str class FileRequestService: - """Resolve signed download URLs for trusted external file consumers.""" + """Resolve signed download URIs for trusted external file consumers.""" _access_controller: DatabaseFileAccessController + _runtime: DifyWorkflowFileRuntime def __init__(self, access_controller: DatabaseFileAccessController | None = None) -> None: self._access_controller = access_controller or DatabaseFileAccessController() + self._runtime = DifyWorkflowFileRuntime(file_access_controller=self._access_controller) - def request_download_url( + def request_download( self, *, tenant_id: str, @@ -45,7 +47,6 @@ class FileRequestService: user_from: UserFrom | str, invoke_from: InvokeFrom | str, file_mapping: Mapping[str, Any], - for_external: bool = True, ) -> DownloadFileRequestResult: """Resolve one file mapping into signed download metadata. @@ -62,15 +63,15 @@ class FileRequestService: ) with bind_file_access_scope(scope): file = self._build_file(mapping=file_mapping, tenant_id=tenant_id) - download_url = file_helpers.resolve_file_url(file, for_external=for_external) + download_uri = self._runtime.resolve_file_uri(file=file) - if not download_url: + if not download_uri: raise ValueError("file does not support signed download") return DownloadFileRequestResult( filename=file.filename or "download.bin", mime_type=file.mime_type, size=file.size, - download_url=download_url, + download_uri=download_uri, ) def _build_file(self, *, mapping: Mapping[str, Any], tenant_id: str) -> File: diff --git a/api/services/file_service.py b/api/services/file_service.py index 9409d2ac64d..3ddc81ab41f 100644 --- a/api/services/file_service.py +++ b/api/services/file_service.py @@ -56,6 +56,7 @@ class FileService: tenant_id: str | None = None, source: Literal["datasets"] | None = None, source_url: str = "", + default_file_size_limit: int | None = None, ) -> UploadFile: # get file extension extension = os.path.splitext(filename)[1].lstrip(".").lower() @@ -79,7 +80,11 @@ class FileService: file_size = len(content) # check if the file size is exceeded - if not FileService.is_file_size_within_limit(extension=extension, file_size=file_size): + if not FileService.is_file_size_within_limit( + extension=extension, + file_size=file_size, + default_file_size_limit=default_file_size_limit, + ): raise FileTooLargeError # generate file key @@ -119,7 +124,12 @@ class FileService: return upload_file @staticmethod - def is_file_size_within_limit(*, extension: str, file_size: int) -> bool: + def is_file_size_within_limit( + *, + extension: str, + file_size: int, + default_file_size_limit: int | None = None, + ) -> bool: if extension in IMAGE_EXTENSIONS: file_size_limit = dify_config.UPLOAD_IMAGE_FILE_SIZE_LIMIT * 1024 * 1024 elif extension in VIDEO_EXTENSIONS: @@ -127,7 +137,12 @@ class FileService: elif extension in AUDIO_EXTENSIONS: file_size_limit = dify_config.UPLOAD_AUDIO_FILE_SIZE_LIMIT * 1024 * 1024 else: - file_size_limit = dify_config.UPLOAD_FILE_SIZE_LIMIT * 1024 * 1024 + # Context-specific uploads may override the default limit without changing media-specific limits. + file_size_limit = ( + (default_file_size_limit if default_file_size_limit is not None else dify_config.UPLOAD_FILE_SIZE_LIMIT) + * 1024 + * 1024 + ) return file_size <= file_size_limit @@ -170,9 +185,10 @@ class FileService: # user uuid as file name file_uuid = str(uuid.uuid4()) file_key = "upload_files/" + tenant_id + "/" + file_uuid + ".txt" + content = text.encode("utf-8") # save file to storage - storage.save(file_key, text.encode("utf-8")) + storage.save(file_key, content) # save file to db upload_file = UploadFile( @@ -180,7 +196,7 @@ class FileService: storage_type=StorageType(dify_config.STORAGE_TYPE), key=file_key, name=text_name, - size=len(text), + size=len(content), extension="txt", mime_type="text/plain", created_by=user_id, diff --git a/api/services/init_validation_service.py b/api/services/init_validation_service.py new file mode 100644 index 00000000000..ef305f52b31 --- /dev/null +++ b/api/services/init_validation_service.py @@ -0,0 +1,52 @@ +"""Application service for the self-hosted initialization gate.""" + +import hmac +from typing import Protocol + + +class InitValidationState(Protocol): + def has_tenants(self) -> bool: ... + + def is_setup(self) -> bool: ... + + +class AlreadyInitializedError(Exception): + """Raised when initialization has already created a tenant.""" + + +class InvalidInitializationPasswordError(Exception): + """Raised when the supplied initialization password does not match.""" + + +class InitValidationService: + def __init__( + self, + *, + state: InitValidationState, + validation_required: bool, + expected_password: str, + ) -> None: + self._state = state + self._validation_required = validation_required + self._expected_password = expected_password + + def is_validated(self, *, session_validated: bool) -> bool: + if not self._validation_required or session_validated: + return True + + return self._state.is_setup() + + def validate_password(self, password: str) -> None: + if self._state.has_tenants(): + raise AlreadyInitializedError + + expected_password = self._expected_password + if ( + not password + or not expected_password + or not hmac.compare_digest( + password.encode("utf-8"), + expected_password.encode("utf-8"), + ) + ): + raise InvalidInitializationPasswordError diff --git a/api/services/installed_app_service.py b/api/services/installed_app_service.py new file mode 100644 index 00000000000..e3ff637190e --- /dev/null +++ b/api/services/installed_app_service.py @@ -0,0 +1,177 @@ +from datetime import datetime + +from pydantic import BaseModel +from sqlalchemy import and_, exists, or_, select +from sqlalchemy.orm import Session, scoped_session + +from libs.helper import escape_like_pattern +from models import App, AppModelConfig, InstalledApp, Workflow +from models.model import AppMode +from services.enterprise.enterprise_service import EnterpriseService +from services.feature_service import FeatureService + + +class InstalledAppCursor(BaseModel): + is_pinned: bool + last_used_at: datetime | None + installed_app_id: str + + +def _published_app_filter(): + """Return the SQL predicate for installed-app web API availability. + + The installed-app parameters endpoint reads the published workflow for + workflow-style apps and the published app model config for easy UI apps. + Keep the list endpoint aligned in SQL so it does not return entries that + will immediately fail with app_unavailable when opened. + """ + workflow_app_modes = (AppMode.ADVANCED_CHAT, AppMode.WORKFLOW) + has_published_workflow = exists(select(Workflow.id).where(Workflow.id == App.workflow_id)) + has_published_model_config = exists(select(AppModelConfig.id).where(AppModelConfig.id == App.app_model_config_id)) + + return and_( + App.mode != AppMode.AGENT, + or_( + and_(App.mode.in_(workflow_app_modes), App.workflow_id.isnot(None), has_published_workflow), + and_(~App.mode.in_(workflow_app_modes), App.app_model_config_id.isnot(None), has_published_model_config), + ), + ) + + +def _installed_app_cursor_filter(cursor: InstalledAppCursor): + same_pin_group = InstalledApp.is_pinned == cursor.is_pinned + if cursor.last_used_at is None: + later_in_pin_group = and_( + InstalledApp.last_used_at.is_(None), + InstalledApp.id > cursor.installed_app_id, + ) + else: + later_in_pin_group = or_( + InstalledApp.last_used_at < cursor.last_used_at, + InstalledApp.last_used_at.is_(None), + and_( + InstalledApp.last_used_at == cursor.last_used_at, + InstalledApp.id > cursor.installed_app_id, + ), + ) + + if cursor.is_pinned: + return or_( + InstalledApp.is_pinned.is_(False), + and_(same_pin_group, later_in_pin_group), + ) + return and_(same_pin_group, later_in_pin_group) + + +def _installed_app_order_by(): + return ( + InstalledApp.is_pinned.desc(), + InstalledApp.last_used_at.desc().nulls_last(), + InstalledApp.id.asc(), + ) + + +def _filter_rows_by_webapp_auth( + rows: list[tuple[InstalledApp, App]], + *, + user_id: str, +) -> list[tuple[InstalledApp, App]]: + if not rows: + return [] + + app_ids = [app.id for _, app in rows] + webapp_settings = EnterpriseService.WebAppAuth.batch_get_app_access_mode_by_id(app_ids) + candidates = [ + (installed_app, app) + for installed_app, app in rows + if (setting := webapp_settings.get(app.id)) is not None and setting.access_mode != "sso_verified" + ] + permissions = EnterpriseService.WebAppAuth.batch_is_user_allowed_to_access_webapps( + user_id=user_id, + app_ids=[app.id for _, app in candidates], + ) + return [(installed_app, app) for installed_app, app in candidates if permissions.get(app.id)] + + +class InstalledAppService: + @classmethod + def get_visible_page( + cls, + *, + tenant_id: str, + user_id: str, + cursor: InstalledAppCursor | None, + limit: int, + app_id: str | None, + name: str | None, + session: Session | scoped_session, + ) -> tuple[list[tuple[InstalledApp, App]], bool, InstalledAppCursor | None]: + """Scan ordered candidates until one page of authorized apps is complete.""" + stmt = ( + select(InstalledApp, App) + .join(App, App.id == InstalledApp.app_id) + .where(InstalledApp.tenant_id == tenant_id, _published_app_filter()) + ) + if app_id: + stmt = stmt.where(InstalledApp.app_id == app_id) + if name and (normalized_name := name.strip()): + escaped_name = escape_like_pattern(normalized_name) + stmt = stmt.where(App.name.ilike(f"%{escaped_name}%", escape="\\")) + + webapp_auth_enabled = FeatureService.get_system_features().webapp_auth.enabled + scan_size = limit * 2 if webapp_auth_enabled else limit + 1 + visible_rows: list[tuple[InstalledApp, App]] = [] + scan_cursor = cursor + has_more = False + last_consumed_app: InstalledApp | None = None + + while True: + page_stmt = stmt + if scan_cursor is not None: + page_stmt = page_stmt.where(_installed_app_cursor_filter(scan_cursor)) + candidate_result = session.execute(page_stmt.order_by(*_installed_app_order_by()).limit(scan_size)).all() + candidate_rows = [(installed_app, app) for installed_app, app in candidate_result] + if not candidate_rows: + break + + authorized_rows = candidate_rows + if webapp_auth_enabled: + authorized_rows = _filter_rows_by_webapp_auth(candidate_rows, user_id=user_id) + + authorized_installed_app_ids = {installed_app.id for installed_app, _ in authorized_rows} + for row in candidate_rows: + installed_app = row[0] + if installed_app.id not in authorized_installed_app_ids: + last_consumed_app = installed_app + continue + if len(visible_rows) == limit: + has_more = True + break + visible_rows.append(row) + last_consumed_app = installed_app + if has_more: + break + + if len(candidate_rows) < scan_size: + break + last_scanned_app = candidate_rows[-1][0] + scan_cursor = InstalledAppCursor( + is_pinned=last_scanned_app.is_pinned, + last_used_at=last_scanned_app.last_used_at, + installed_app_id=last_scanned_app.id, + ) + + next_cursor = ( + InstalledAppCursor( + is_pinned=last_consumed_app.is_pinned, + last_used_at=last_consumed_app.last_used_at, + installed_app_id=last_consumed_app.id, + ) + if has_more and last_consumed_app + else None + ) + return visible_rows, has_more, next_cursor + + @staticmethod + def get_published_app(app_id: str, *, session: Session | scoped_session) -> App | None: + return session.scalar(select(App).where(App.id == app_id, _published_app_filter()).limit(1)) diff --git a/api/services/legacy_model_type_migration.py b/api/services/legacy_model_type_migration.py index 1465fc0912f..2034246991f 100644 --- a/api/services/legacy_model_type_migration.py +++ b/api/services/legacy_model_type_migration.py @@ -38,7 +38,7 @@ from typing import Protocol, cast, override import sqlalchemy as sa from sqlalchemy.exc import OperationalError -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, defer from sqlalchemy.sql import select from core.helper.model_provider_cache import ProviderCredentialsCache, ProviderCredentialsCacheType @@ -747,6 +747,7 @@ class Migration: with _session_factory(self._engine) as session: stmt = ( select(ProviderModel, raw_model_type) + .options(defer(ProviderModel.model_type)) .where( ProviderModel.tenant_id == self._tenant_id, sa.type_coerce(ProviderModel.model_type, sa.String()).in_(self._selected_legacy_values()), @@ -787,6 +788,7 @@ class Migration: raw_model_type = sa.type_coerce(ProviderModel.model_type, sa.String()).label(_RAW_MODEL_TYPE_COLUMN) stmt = ( select(ProviderModel, raw_model_type) + .options(defer(ProviderModel.model_type)) .where( ProviderModel.tenant_id == candidate.row.tenant_id, ProviderModel.provider_name == candidate.row.provider_name, @@ -983,6 +985,7 @@ class Migration: with _session_factory(self._engine) as session: stmt = ( select(TenantDefaultModel, raw_model_type) + .options(defer(TenantDefaultModel.model_type)) .where( TenantDefaultModel.tenant_id == self._tenant_id, sa.type_coerce(TenantDefaultModel.model_type, sa.String()).in_(self._selected_legacy_values()), @@ -1023,6 +1026,7 @@ class Migration: raw_model_type = sa.type_coerce(TenantDefaultModel.model_type, sa.String()).label(_RAW_MODEL_TYPE_COLUMN) stmt = ( select(TenantDefaultModel, raw_model_type) + .options(defer(TenantDefaultModel.model_type)) .where( TenantDefaultModel.tenant_id == candidate.row.tenant_id, sa.type_coerce(TenantDefaultModel.model_type, sa.String()).in_( @@ -1193,6 +1197,7 @@ class Migration: with _session_factory(self._engine) as session: stmt = ( select(ProviderModelSetting, raw_model_type) + .options(defer(ProviderModelSetting.model_type)) .where( ProviderModelSetting.tenant_id == self._tenant_id, sa.type_coerce(ProviderModelSetting.model_type, sa.String()).in_(self._selected_legacy_values()), @@ -1233,6 +1238,7 @@ class Migration: raw_model_type = sa.type_coerce(ProviderModelSetting.model_type, sa.String()).label(_RAW_MODEL_TYPE_COLUMN) stmt = ( select(ProviderModelSetting, raw_model_type) + .options(defer(ProviderModelSetting.model_type)) .where( ProviderModelSetting.tenant_id == candidate.row.tenant_id, ProviderModelSetting.provider_name == candidate.row.provider_name, @@ -1437,6 +1443,7 @@ class Migration: with _session_factory(self._engine) as session: stmt = ( select(LoadBalancingModelConfig, raw_model_type) + .options(defer(LoadBalancingModelConfig.model_type)) .where( LoadBalancingModelConfig.tenant_id == self._tenant_id, LoadBalancingModelConfig.name == "__inherit__", @@ -1485,6 +1492,7 @@ class Migration: raw_model_type = sa.type_coerce(LoadBalancingModelConfig.model_type, sa.String()).label(_RAW_MODEL_TYPE_COLUMN) stmt = ( select(LoadBalancingModelConfig, raw_model_type) + .options(defer(LoadBalancingModelConfig.model_type)) .where( LoadBalancingModelConfig.tenant_id == candidate.row.tenant_id, LoadBalancingModelConfig.provider_name == candidate.row.provider_name, @@ -1617,6 +1625,7 @@ class Migration: with _session_factory(self._engine) as session: stmt = ( select(LoadBalancingModelConfig, raw_model_type) + .options(defer(LoadBalancingModelConfig.model_type)) .where( LoadBalancingModelConfig.tenant_id == self._tenant_id, sa.type_coerce(LoadBalancingModelConfig.model_type, sa.String()).in_( @@ -1660,9 +1669,13 @@ class Migration: lock_rows: bool, ) -> _RowWithRawModelType[LoadBalancingModelConfig] | None: raw_model_type = sa.type_coerce(LoadBalancingModelConfig.model_type, sa.String()).label(_RAW_MODEL_TYPE_COLUMN) - stmt = select(LoadBalancingModelConfig, raw_model_type).where( - LoadBalancingModelConfig.id == candidate.row.id, - LoadBalancingModelConfig.tenant_id == self._tenant_id, + stmt = ( + select(LoadBalancingModelConfig, raw_model_type) + .options(defer(LoadBalancingModelConfig.model_type)) + .where( + LoadBalancingModelConfig.id == candidate.row.id, + LoadBalancingModelConfig.tenant_id == self._tenant_id, + ) ) if lock_rows: stmt = stmt.with_for_update() @@ -1818,6 +1831,7 @@ class Migration: with _session_factory(self._engine) as session: stmt = ( select(ProviderModelCredential, raw_model_type) + .options(defer(ProviderModelCredential.model_type)) .where( ProviderModelCredential.tenant_id == self._tenant_id, sa.type_coerce(ProviderModelCredential.model_type, sa.String()).in_(self._selected_legacy_values()), @@ -1858,6 +1872,7 @@ class Migration: raw_model_type = sa.type_coerce(ProviderModelCredential.model_type, sa.String()).label(_RAW_MODEL_TYPE_COLUMN) stmt = ( select(ProviderModelCredential, raw_model_type) + .options(defer(ProviderModelCredential.model_type)) .where( ProviderModelCredential.tenant_id == candidate.row.tenant_id, ProviderModelCredential.provider_name == candidate.row.provider_name, @@ -2147,6 +2162,7 @@ class Migration: stmt = ( select(ProviderModel) + .options(defer(ProviderModel.model_type)) .where( ProviderModel.tenant_id == self._tenant_id, ProviderModel.credential_id.in_(loser_ids), @@ -2182,6 +2198,7 @@ class Migration: stmt = ( select(LoadBalancingModelConfig) + .options(defer(LoadBalancingModelConfig.model_type)) .where( LoadBalancingModelConfig.tenant_id == self._tenant_id, LoadBalancingModelConfig.credential_id.in_(loser_ids), @@ -2295,9 +2312,14 @@ class Migration: def _row_to_dict(self, row: TypeBase, *, raw_model_type: str | None = None) -> dict[str, object]: mapper = sa.inspect(row).mapper - row_dict = {column.key: row.__dict__[column.key] for column in mapper.column_attrs} - if raw_model_type is not None and "model_type" in row_dict: - row_dict["model_type"] = raw_model_type + row_dict = { + column.key: ( + raw_model_type + if column.key == "model_type" and raw_model_type is not None + else row.__dict__[column.key] + ) + for column in mapper.column_attrs + } return _normalize_log_mapping(row_dict) def _log_row_deleted[T: TypeBase]( diff --git a/api/services/metadata_service.py b/api/services/metadata_service.py index 004a059acc6..317e371fc2a 100644 --- a/api/services/metadata_service.py +++ b/api/services/metadata_service.py @@ -9,13 +9,15 @@ from extensions.ext_redis import redis_client from libs.datetime_utils import naive_utc_now from libs.login import resolve_account_fallback from models import Account -from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding +from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding, Document from models.enums import DatasetMetadataType +from services.dataset_ref_service import DatasetRefService from services.dataset_service import DocumentService from services.entities.knowledge_entities.knowledge_entities import ( MetadataArgs, MetadataOperationData, ) +from services.errors.metadata import MetadataResourceNotFoundError logger = logging.getLogger(__name__) @@ -61,11 +63,10 @@ class MetadataService: @staticmethod def update_metadata_name( - dataset_id: str, + dataset: Dataset, metadata_id: str, name: str, - current_user: Account | None = None, - current_tenant_id: str | None = None, # TODO: the service_api is not migrated yet + current_user: Account, *, session: Session, ) -> DatasetMetadata | None: @@ -73,14 +74,13 @@ class MetadataService: if len(name) > 255: raise ValueError("Metadata name cannot exceed 255 characters.") - lock_key = f"dataset_metadata_lock_{dataset_id}" + lock_key = f"dataset_metadata_lock_{dataset.id}" # check if metadata name already exists - current_user, current_tenant_id = resolve_account_fallback(current_user, current_tenant_id) if session.scalar( select(DatasetMetadata) .where( - DatasetMetadata.tenant_id == current_tenant_id, - DatasetMetadata.dataset_id == dataset_id, + DatasetMetadata.tenant_id == dataset.tenant_id, + DatasetMetadata.dataset_id == dataset.id, DatasetMetadata.name == name, ) .limit(1) @@ -90,10 +90,14 @@ class MetadataService: if field.value == name: raise ValueError("Metadata name already exists in Built-in fields.") try: - MetadataService.knowledge_base_metadata_lock_check(dataset_id, None) + MetadataService.knowledge_base_metadata_lock_check(dataset.id, None) metadata = session.scalar( select(DatasetMetadata) - .where(DatasetMetadata.id == metadata_id, DatasetMetadata.dataset_id == dataset_id) + .where( + DatasetMetadata.id == metadata_id, + DatasetMetadata.tenant_id == dataset.tenant_id, + DatasetMetadata.dataset_id == dataset.id, + ) .limit(1) ) if metadata is None: @@ -105,11 +109,17 @@ class MetadataService: # update related documents dataset_metadata_bindings = session.scalars( - select(DatasetMetadataBinding).where(DatasetMetadataBinding.metadata_id == metadata_id) + select(DatasetMetadataBinding).where( + DatasetMetadataBinding.metadata_id == metadata_id, + DatasetMetadataBinding.tenant_id == dataset.tenant_id, + DatasetMetadataBinding.dataset_id == dataset.id, + ) ).all() if dataset_metadata_bindings: document_ids = [binding.document_id for binding in dataset_metadata_bindings] - documents = DocumentService.get_document_by_ids(document_ids, session) + documents = DocumentService.get_document_by_ids( + DatasetRefService.create_dataset_ref(dataset), document_ids, session + ) for document in documents: if not document.doc_metadata: doc_metadata = {} @@ -128,13 +138,17 @@ class MetadataService: redis_client.delete(lock_key) @staticmethod - def delete_metadata(dataset_id: str, metadata_id: str, session: Session): - lock_key = f"dataset_metadata_lock_{dataset_id}" + def delete_metadata(dataset: Dataset, metadata_id: str, session: Session): + lock_key = f"dataset_metadata_lock_{dataset.id}" try: - MetadataService.knowledge_base_metadata_lock_check(dataset_id, None) + MetadataService.knowledge_base_metadata_lock_check(dataset.id, None) metadata = session.scalar( select(DatasetMetadata) - .where(DatasetMetadata.id == metadata_id, DatasetMetadata.dataset_id == dataset_id) + .where( + DatasetMetadata.id == metadata_id, + DatasetMetadata.tenant_id == dataset.tenant_id, + DatasetMetadata.dataset_id == dataset.id, + ) .limit(1) ) if metadata is None: @@ -143,11 +157,17 @@ class MetadataService: # deal related documents dataset_metadata_bindings = session.scalars( - select(DatasetMetadataBinding).where(DatasetMetadataBinding.metadata_id == metadata_id) + select(DatasetMetadataBinding).where( + DatasetMetadataBinding.metadata_id == metadata_id, + DatasetMetadataBinding.tenant_id == dataset.tenant_id, + DatasetMetadataBinding.dataset_id == dataset.id, + ) ).all() if dataset_metadata_bindings: document_ids = [binding.document_id for binding in dataset_metadata_bindings] - documents = DocumentService.get_document_by_ids(document_ids, session) + documents = DocumentService.get_document_by_ids( + DatasetRefService.create_dataset_ref(dataset), document_ids, session + ) for document in documents: if not document.doc_metadata: doc_metadata = {} @@ -237,27 +257,61 @@ class MetadataService: def update_documents_metadata( dataset: Dataset, metadata_args: MetadataOperationData, - current_user: Account | None = None, # TODO: the service_api is not migrated yet - current_tenant_id: str | None = None, + current_user: Account, *, session: Session, ): - current_user, current_tenant_id = resolve_account_fallback( - current_user, current_tenant_id, fallback_tenant_id=dataset.tenant_id + metadata_ids = { + metadata_value.id + for operation in metadata_args.operation_data + for metadata_value in operation.metadata_list + } + metadatas = session.scalars( + select(DatasetMetadata).where( + DatasetMetadata.id.in_(metadata_ids), + DatasetMetadata.tenant_id == dataset.tenant_id, + DatasetMetadata.dataset_id == dataset.id, + ) + ).all() + metadata_by_id = {metadata.id: metadata for metadata in metadatas} + if metadata_ids != set(metadata_by_id): + raise MetadataResourceNotFoundError("Metadata not found.") + + document_ids = {operation.document_id for operation in metadata_args.operation_data} + owned_document_ids = set( + session.scalars( + select(Document.id).where( + Document.id.in_(document_ids), + Document.tenant_id == dataset.tenant_id, + Document.dataset_id == dataset.id, + ) + ).all() ) + if document_ids != owned_document_ids: + raise MetadataResourceNotFoundError("Document not found.") + for operation in metadata_args.operation_data: lock_key = f"document_metadata_lock_{operation.document_id}" try: MetadataService.knowledge_base_metadata_lock_check(None, operation.document_id) - document = DocumentService.get_document(dataset.id, operation.document_id, session=session) + document = session.scalar( + select(Document) + .where( + Document.id == operation.document_id, + Document.tenant_id == dataset.tenant_id, + Document.dataset_id == dataset.id, + ) + .with_for_update() + .execution_options(populate_existing=True) + ) if document is None: - raise ValueError("Document not found.") + raise MetadataResourceNotFoundError("Document not found.") if operation.partial_update: doc_metadata = copy.deepcopy(document.doc_metadata) if document.doc_metadata else {} else: doc_metadata = {} for metadata_value in operation.metadata_list: - doc_metadata[metadata_value.name] = metadata_value.value + doc_metadata[metadata_by_id[metadata_value.id].name] = metadata_value.value if dataset.built_in_field_enabled: doc_metadata[BuiltInField.document_name] = document.name doc_metadata[BuiltInField.uploader] = document.get_uploader(session=session) @@ -271,7 +325,9 @@ class MetadataService: if not operation.partial_update: session.execute( delete(DatasetMetadataBinding).where( - DatasetMetadataBinding.document_id == operation.document_id + DatasetMetadataBinding.tenant_id == dataset.tenant_id, + DatasetMetadataBinding.dataset_id == dataset.id, + DatasetMetadataBinding.document_id == document.id, ) ) @@ -281,7 +337,9 @@ class MetadataService: existing_binding = session.scalar( select(DatasetMetadataBinding) .where( - DatasetMetadataBinding.document_id == operation.document_id, + DatasetMetadataBinding.tenant_id == dataset.tenant_id, + DatasetMetadataBinding.dataset_id == dataset.id, + DatasetMetadataBinding.document_id == document.id, DatasetMetadataBinding.metadata_id == metadata_value.id, ) .limit(1) @@ -290,9 +348,9 @@ class MetadataService: continue dataset_metadata_binding = DatasetMetadataBinding( - tenant_id=current_tenant_id, + tenant_id=dataset.tenant_id, dataset_id=dataset.id, - document_id=operation.document_id, + document_id=document.id, metadata_id=metadata_value.id, created_by=current_user.id, ) diff --git a/api/services/model_load_balancing_service.py b/api/services/model_load_balancing_service.py index 6eab1ffbe3e..8dc4824fa66 100644 --- a/api/services/model_load_balancing_service.py +++ b/api/services/model_load_balancing_service.py @@ -178,8 +178,8 @@ class ModelLoadBalancingService: # Get credential form schemas from model credential schema or provider credential schema credential_schemas = self._get_credential_schema(provider_configuration) - # Get decoding rsa key and cipher for decrypting credentials - decoding_rsa_key, decoding_cipher_rsa = encrypter.get_decrypt_decoding(tenant_id) + # Get decoding context for decrypting credentials + decoding_context = encrypter.get_decrypt_decoding(tenant_id) # fetch status and ttl for each config datas: list[LoadBalancingConfigSummaryDict] = [] @@ -213,8 +213,7 @@ class ModelLoadBalancingService: if isinstance(token_value, str): credentials[variable] = encrypter.decrypt_token_with_decoding( token_value, - decoding_rsa_key, - decoding_cipher_rsa, + decoding_context, ) except ValueError: pass diff --git a/api/services/model_provider_service.py b/api/services/model_provider_service.py index 7c34afd42e1..11ba3723751 100644 --- a/api/services/model_provider_service.py +++ b/api/services/model_provider_service.py @@ -1,18 +1,43 @@ import logging +from collections import defaultdict +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any +from sqlalchemy import and_, select + if TYPE_CHECKING: from models.account import Account +from configs import dify_config +from core.db.session_factory import session_factory from core.entities.model_entities import ModelWithProviderEntity, ProviderModelWithStatusEntity +from core.entities.provider_entities import CredentialConfiguration +from core.helper.position_helper import is_filtered +from core.plugin.entities.plugin import PluginInstallationSource +from core.plugin.entities.plugin_daemon import PluginModelProviderBinding from core.plugin.impl.model_runtime_factory import create_plugin_model_provider_factory, create_plugin_provider_manager +from core.plugin.plugin_service import PluginService from core.provider_manager import ProviderManager +from enums import DeploymentEdition +from extensions import ext_hosting_provider from graphon.model_runtime.entities.model_entities import ModelType, ParameterRule -from models.provider import ProviderType +from models.provider import ( + Provider, + ProviderCredential, + ProviderModel, + ProviderModelCredential, + ProviderType, + TenantPreferredModelProvider, +) +from models.provider_ids import ModelProviderID from services.entities.model_provider_entities import ( CustomConfigurationResponse, CustomConfigurationStatus, DefaultModelResponse, + ModelProviderCustomConfigurationSummaryResponse, + ModelProviderPluginSummaryResponse, + ModelProviderSummaryResponse, + ModelProviderSystemConfigurationSummaryResponse, ModelWithProviderEntityResponse, ProviderResponse, ProviderWithModelsResponse, @@ -24,6 +49,17 @@ from services.errors.app_model_config import ProviderNotFoundError logger = logging.getLogger(__name__) +@dataclass(slots=True) +class _ProviderSummaryState: + has_custom_provider: bool = False + available_credentials: list[CredentialConfiguration] = field(default_factory=list) + has_custom_models: bool = False + current_credential_id: str | None = None + current_credential_name: str | None = None + current_credential_usable: bool = False + preferred_provider_type: ProviderType | None = None + + class ModelProviderService: """ Model Provider Service @@ -132,6 +168,249 @@ class ModelProviderService: return provider_responses + @staticmethod + def _load_provider_summary_states(tenant_id: str) -> dict[str, _ProviderSummaryState]: + """Load only the workspace columns required by the collapsed provider list.""" + with session_factory.create_session() as session: + custom_provider_rows = session.execute( + select( + Provider.provider_name, + Provider.credential_id, + ProviderCredential.provider_name.label("credential_provider_name"), + ProviderCredential.credential_name, + ) + .outerjoin( + ProviderCredential, + and_( + ProviderCredential.id == Provider.credential_id, + ProviderCredential.tenant_id == tenant_id, + ), + ) + .where( + Provider.tenant_id == tenant_id, + Provider.provider_type == ProviderType.CUSTOM, + Provider.is_valid.is_(True), + ) + ).all() + credential_rows = session.execute( + select( + ProviderCredential.id, + ProviderCredential.provider_name, + ProviderCredential.credential_name, + ) + .where(ProviderCredential.tenant_id == tenant_id) + .order_by( + ProviderCredential.created_at.desc(), + ProviderCredential.id.desc(), + ) + ).all() + custom_model_rows = session.execute( + select(ProviderModel.provider_name.label("provider_name")) + .where( + ProviderModel.tenant_id == tenant_id, + ProviderModel.is_valid.is_(True), + ) + .union( + select(ProviderModelCredential.provider_name.label("provider_name")).where( + ProviderModelCredential.tenant_id == tenant_id + ) + ) + ).all() + preferred_provider_rows = session.execute( + select( + TenantPreferredModelProvider.provider_name, + TenantPreferredModelProvider.preferred_provider_type, + ).where(TenantPreferredModelProvider.tenant_id == tenant_id) + ).all() + + states: defaultdict[str, _ProviderSummaryState] = defaultdict(_ProviderSummaryState) + for credential in credential_rows: + provider_name = str(ModelProviderID(credential.provider_name)) + states[provider_name].available_credentials.append( + CredentialConfiguration( + credential_id=credential.id, + credential_name=credential.credential_name, + ) + ) + + selected_provider_priorities: dict[str, bool] = {} + for provider in custom_provider_rows: + provider_name = str(ModelProviderID(provider.provider_name)) + state = states[provider_name] + state.has_custom_provider = True + + is_canonical_row = provider.provider_name == provider_name + if provider_name in selected_provider_priorities and not is_canonical_row: + continue + selected_provider_priorities[provider_name] = is_canonical_row + state.current_credential_id = provider.credential_id + if ( + provider.credential_provider_name is not None + and str(ModelProviderID(provider.credential_provider_name)) == provider_name + ): + state.current_credential_name = provider.credential_name + state.current_credential_usable = True + else: + state.current_credential_name = None + state.current_credential_usable = False + + for model in custom_model_rows: + states[str(ModelProviderID(model.provider_name))].has_custom_models = True + + preferred_provider_priorities: dict[str, bool] = {} + for preferred_provider in preferred_provider_rows: + provider_name = str(ModelProviderID(preferred_provider.provider_name)) + is_canonical_row = preferred_provider.provider_name == provider_name + if provider_name in preferred_provider_priorities and not is_canonical_row: + continue + preferred_provider_priorities[provider_name] = is_canonical_row + states[provider_name].preferred_provider_type = preferred_provider.preferred_provider_type + + return dict(states) + + @staticmethod + def _has_system_provider_hosting_configuration(provider: str) -> bool: + configuration = ext_hosting_provider.hosting_configuration.provider_map.get(provider) + return bool(configuration and configuration.enabled and configuration.quotas) + + @staticmethod + def _select_binding( + current_binding: PluginModelProviderBinding | None, + candidate_binding: PluginModelProviderBinding, + ) -> PluginModelProviderBinding: + """Prefer a remote-debug runtime when one shadows an installed plugin.""" + if current_binding is None: + return candidate_binding + if ( + candidate_binding.source == PluginInstallationSource.Remote + and current_binding.source != PluginInstallationSource.Remote + ): + return candidate_binding + return current_binding + + @staticmethod + def _get_preferred_provider_type( + state: _ProviderSummaryState, + *, + custom_present: bool, + system_enabled: bool, + ) -> ProviderType: + if state.preferred_provider_type is not None: + return state.preferred_provider_type + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and system_enabled: + return ProviderType.SYSTEM + if custom_present: + return ProviderType.CUSTOM + if system_enabled: + return ProviderType.SYSTEM + return ProviderType.CUSTOM + + def get_provider_summary_list( + self, tenant_id: str + ) -> tuple[list[ModelProviderSummaryResponse], dict[str, ModelProviderPluginSummaryResponse]]: + """Build the complete first-screen provider projection without assembling provider configurations.""" + # Read bindings first: remote-debug identity changes invalidate provider metadata + # before the provider cache is consulted. + bindings = PluginService.list_model_provider_bindings(tenant_id) + provider_entities = PluginService.fetch_plugin_model_providers(tenant_id=tenant_id) + states = self._load_provider_summary_states(tenant_id) + + bindings_by_provider: dict[str, PluginModelProviderBinding] = {} + for binding in bindings: + provider_name = ( + str(ModelProviderID(binding.provider)) + if binding.provider.count("/") == 2 + else str(ModelProviderID(f"{binding.plugin_id}/{binding.provider}")) + ) + bindings_by_provider[provider_name] = self._select_binding( + bindings_by_provider.get(provider_name), + binding, + ) + + provider_summaries: list[ModelProviderSummaryResponse] = [] + emitted_provider_names: set[str] = set() + for provider_entity in provider_entities: + if is_filtered( + include_set=dify_config.POSITION_PROVIDER_INCLUDES_SET, + exclude_set=dify_config.POSITION_PROVIDER_EXCLUDES_SET, + data=provider_entity, + name_func=lambda provider: provider.provider, + ): + continue + + provider_id = ModelProviderID(provider_entity.provider) + provider_name = str(provider_id) + if provider_name in emitted_provider_names: + continue + emitted_provider_names.add(provider_name) + + state = states.get(provider_name, _ProviderSummaryState()) + custom_configured = ( + state.has_custom_provider and bool(state.available_credentials) + ) or state.has_custom_models + custom_present = state.has_custom_provider or state.has_custom_models + provider_binding = bindings_by_provider.get(provider_name) + system_enabled = bool( + provider_binding + and self._has_system_provider_hosting_configuration(provider_name) + and provider_binding.source != PluginInstallationSource.Package + and provider_binding.verified + ) + preferred_provider_type = self._get_preferred_provider_type( + state, + custom_present=custom_present, + system_enabled=system_enabled, + ) + + provider_summaries.append( + ModelProviderSummaryResponse( + tenant_id=tenant_id, + provider=provider_name, + plugin_id=provider_id.plugin_id, + label=provider_entity.label, + description=provider_entity.description, + icon_small=provider_entity.icon_small, + icon_small_dark=provider_entity.icon_small_dark, + supported_model_types=provider_entity.supported_model_types, + configurate_methods=provider_entity.configurate_methods, + preferred_provider_type=preferred_provider_type, + is_configured=custom_configured or system_enabled, + custom_configuration=ModelProviderCustomConfigurationSummaryResponse( + status=CustomConfigurationStatus.ACTIVE + if custom_configured + else CustomConfigurationStatus.NO_CONFIGURE, + has_custom_models=state.has_custom_models, + available_credentials=state.available_credentials, + current_credential_id=state.current_credential_id, + current_credential_name=state.current_credential_name, + current_credential_usable=state.current_credential_usable, + ), + system_configuration=ModelProviderSystemConfigurationSummaryResponse( + enabled=system_enabled, + ), + ) + ) + + plugin_bindings: dict[str, PluginModelProviderBinding] = {} + for binding in bindings_by_provider.values(): + plugin_bindings[binding.plugin_id] = self._select_binding( + plugin_bindings.get(binding.plugin_id), + binding, + ) + + plugin_summaries = { + plugin_id: ModelProviderPluginSummaryResponse( + installation_id=binding.installation_id, + plugin_id=binding.plugin_id, + plugin_unique_identifier=binding.plugin_unique_identifier, + runtime_type=binding.runtime_type, + source=binding.source, + version=binding.version, + ) + for plugin_id, binding in plugin_bindings.items() + } + return provider_summaries, plugin_summaries + def get_models_by_provider(self, tenant_id: str, provider: str) -> list[ModelWithProviderEntityResponse]: """ get provider models. diff --git a/api/services/openapi/license_gate.py b/api/services/openapi/license_gate.py index c4f17e12c0d..7ca7779de6f 100644 --- a/api/services/openapi/license_gate.py +++ b/api/services/openapi/license_gate.py @@ -1,6 +1,6 @@ """License gate for the /openapi/v1/permitted-external-apps* surface. -EE-only. CE deploys (``ENTERPRISE_ENABLED=false``) skip the gate entirely — +Enterprise-edition only. Community and Cloud deployments skip the gate entirely — the EE blueprint chain is what gives CE deploys no callers on this surface in practice, but the explicit short-circuit avoids any test/fixture that flips the surface on without flipping the license. @@ -22,7 +22,9 @@ from functools import wraps from werkzeug.exceptions import Forbidden from configs import dify_config -from services.feature_service import FeatureService, LicenseStatus +from enums import DeploymentEdition +from services.entities.feature_entities import LicenseStatus +from services.feature_service import FeatureService logger = logging.getLogger(__name__) @@ -31,12 +33,12 @@ _VALID_LICENSE_STATUSES: frozenset[LicenseStatus] = frozenset({LicenseStatus.ACT def license_required[**P, R](view: Callable[P, R]) -> Callable[P, R]: """Decorator form. Raises ``Forbidden('license_required')`` when the EE - deployment has no valid license. No-op on CE (``ENTERPRISE_ENABLED=false``). + deployment has no valid license. No-op outside the Enterprise edition. """ @wraps(view) def decorated(*args: P.args, **kwargs: P.kwargs) -> R: - if dify_config.ENTERPRISE_ENABLED and not _is_license_valid(): + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE and not _is_license_valid(): raise Forbidden(description="license_required") return view(*args, **kwargs) diff --git a/api/services/plugin/plugin_parameter_service.py b/api/services/plugin/plugin_parameter_service.py index 786c09b44e1..09d016e381a 100644 --- a/api/services/plugin/plugin_parameter_service.py +++ b/api/services/plugin/plugin_parameter_service.py @@ -12,6 +12,7 @@ from core.tools.utils.encryption import create_tool_provider_encrypter from core.trigger.entities.api_entities import TriggerProviderSubscriptionApiEntity from core.trigger.entities.entities import SubscriptionBuilder from extensions.ext_database import db +from models.provider_ids import TriggerProviderID from models.tools import BuiltinToolProvider from services.trigger.trigger_provider_service import TriggerProviderService from services.trigger.trigger_subscription_builder_service import TriggerSubscriptionBuilderService @@ -85,7 +86,12 @@ class PluginParameterService: case "trigger": subscription: TriggerProviderSubscriptionApiEntity | SubscriptionBuilder | None if credential_id: - subscription = TriggerSubscriptionBuilderService.get_subscription_builder(credential_id) + subscription = TriggerSubscriptionBuilderService.get_subscription_builder( + tenant_id=tenant_id, + user_id=user_id, + provider_id=TriggerProviderID(f"{plugin_id}/{provider}"), + subscription_builder_id=credential_id, + ) if not subscription: trigger_subscription = TriggerProviderService.get_subscription_by_id(tenant_id, credential_id) subscription = trigger_subscription.to_api_entity() if trigger_subscription else None diff --git a/api/services/quota_service.py b/api/services/quota_service.py index 4c784315c75..9804e26e50e 100644 --- a/api/services/quota_service.py +++ b/api/services/quota_service.py @@ -3,12 +3,9 @@ from __future__ import annotations import logging import uuid from dataclasses import dataclass, field -from typing import TYPE_CHECKING from configs import dify_config - -if TYPE_CHECKING: - from enums.quota_type import QuotaType +from enums import DeploymentEdition, QuotaType logger = logging.getLogger(__name__) @@ -88,8 +85,6 @@ class QuotaCharge: def unlimited() -> QuotaCharge: - from enums.quota_type import QuotaType - return QuotaCharge(success=True, charge_id=None, _quota_type=QuotaType.UNLIMITED) @@ -123,8 +118,8 @@ class QuotaService: from services.billing_service import BillingService from services.errors.app import QuotaExceededError - if not dify_config.BILLING_ENABLED: - logger.debug("Billing disabled, allowing request for %s", tenant_id) + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: + logger.debug("Quota billing is unavailable outside the Cloud edition; allowing request for %s", tenant_id) return QuotaCharge(success=True, charge_id=None, _quota_type=quota_type) logger.info("Reserving %d %s quota for tenant %s", amount, quota_type.value, tenant_id) @@ -179,7 +174,7 @@ class QuotaService: @staticmethod def check(quota_type: QuotaType, tenant_id: str, amount: int = 1) -> bool: - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return True if amount <= 0: @@ -198,7 +193,7 @@ class QuotaService: try: from services.billing_service import BillingService - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return if not reservation_id: diff --git a/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py b/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py index d56c239ace2..f411b759f69 100644 --- a/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/built_in/built_in_retrieval.py @@ -30,8 +30,10 @@ class BuiltInPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return result @override - def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: - del session + def get_pipeline_template_detail( + self, template_id: str, current_tenant_id: str, *, session: Session + ) -> dict[str, Any] | None: + del current_tenant_id, session result = self.fetch_pipeline_template_detail_from_builtin(template_id) return result diff --git a/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py b/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py index 3d6baefcc46..195d6d45001 100644 --- a/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/customized/customized_retrieval.py @@ -49,8 +49,10 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): ) @override - def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: - return self.fetch_pipeline_template_detail_from_db(template_id, session=session) + def get_pipeline_template_detail( + self, template_id: str, current_tenant_id: str, *, session: Session + ) -> dict[str, Any] | None: + return self.fetch_pipeline_template_detail_from_db(template_id, current_tenant_id, session=session) @override def get_type(self) -> str: @@ -89,13 +91,20 @@ class CustomizedPipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return {"pipeline_templates": recommended_pipelines_results} @classmethod - def fetch_pipeline_template_detail_from_db(cls, template_id: str, *, session: Session) -> dict[str, Any] | None: + def fetch_pipeline_template_detail_from_db( + cls, template_id: str, current_tenant_id: str, *, session: Session + ) -> dict[str, Any] | None: """ Fetch pipeline template detail from db. :param template_id: Template ID :return: """ - pipeline_template = session.get(PipelineCustomizedTemplate, template_id) + pipeline_template = session.scalar( + select(PipelineCustomizedTemplate).where( + PipelineCustomizedTemplate.id == template_id, + PipelineCustomizedTemplate.tenant_id == current_tenant_id, + ) + ) if not pipeline_template: return None diff --git a/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py b/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py index d5c31ff74b2..5f71440dd1c 100644 --- a/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/database/database_retrieval.py @@ -47,7 +47,10 @@ class DatabasePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): return self.fetch_pipeline_templates_from_db(language, session=session) @override - def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: + def get_pipeline_template_detail( + self, template_id: str, current_tenant_id: str, *, session: Session + ) -> dict[str, Any] | None: + del current_tenant_id return self.fetch_pipeline_template_detail_from_db(template_id, session=session) @override diff --git a/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py b/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py index c61ac6d60f2..c86d013b321 100644 --- a/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py +++ b/api/services/rag_pipeline/pipeline_template/pipeline_template_base.py @@ -10,6 +10,8 @@ class PipelineTemplateRetrievalBase(Protocol): self, language: str, current_tenant_id: str | None = None, *, session: Session ) -> dict[str, Any]: ... - def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: ... + def get_pipeline_template_detail( + self, template_id: str, current_tenant_id: str, *, session: Session + ) -> dict[str, Any] | None: ... def get_type(self) -> str: ... diff --git a/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py b/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py index 29acbd198b6..80ecb0daa57 100644 --- a/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py +++ b/api/services/rag_pipeline/pipeline_template/remote/remote_retrieval.py @@ -18,7 +18,10 @@ class RemotePipelineTemplateRetrieval(PipelineTemplateRetrievalBase): """ @override - def get_pipeline_template_detail(self, template_id: str, *, session: Session) -> dict[str, Any] | None: + def get_pipeline_template_detail( + self, template_id: str, current_tenant_id: str, *, session: Session + ) -> dict[str, Any] | None: + del current_tenant_id try: return self.fetch_pipeline_template_detail_from_dify_official(template_id) except Exception as e: diff --git a/api/services/rag_pipeline/rag_pipeline.py b/api/services/rag_pipeline/rag_pipeline.py index 993f6d79492..452e958ba6c 100644 --- a/api/services/rag_pipeline/rag_pipeline.py +++ b/api/services/rag_pipeline/rag_pipeline.py @@ -31,6 +31,7 @@ from core.helper import marketplace from core.rag.entities import DatasourceCompletedEvent, DatasourceErrorEvent, DatasourceProcessingEvent from core.repositories.factory import DifyCoreRepositoryFactory from core.repositories.sqlalchemy_workflow_node_execution_repository import SQLAlchemyWorkflowNodeExecutionRepository +from core.workflow.llm_environment_variable import validate_llm_environment_model_references from core.workflow.node_factory import LATEST_VERSION, get_node_type_classes_mapping from core.workflow.system_variables import ( SystemVariableKey, @@ -49,6 +50,7 @@ from graphon.errors import WorkflowNodeRunFailedError from graphon.graph_events import GraphNodeEventBase, NodeRunFailedEvent, NodeRunSucceededEvent from graphon.node_events import NodeRunResult from graphon.nodes.base.node import Node +from graphon.nodes.container_effects import ContainerAwaitRequest from graphon.nodes.http_request import HTTP_REQUEST_CONFIG_FILTER_KEY, build_http_request_config from graphon.runtime import VariablePool from graphon.variables.variables import Variable, VariableBase @@ -80,6 +82,7 @@ from services.entities.knowledge_entities.rag_pipeline_entities import ( PipelineTemplateInfoEntity, ) from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError +from services.errors.rag_pipeline import RagPipelineResourceNotFoundError from services.rag_pipeline.pipeline_template.pipeline_template_factory import PipelineTemplateRetrievalFactory from services.tools.builtin_tools_manage_service import BuiltinToolManageService from services.workflow_draft_variable_service import DraftVariableSaver, DraftVarLoader @@ -143,7 +146,12 @@ class RagPipelineService: @classmethod def get_pipeline_template_detail( - cls, template_id: str, type: str = "built-in", *, session: Session + cls, + template_id: str, + current_tenant_id: str, + type: str = "built-in", + *, + session: Session, ) -> dict[str, Any] | None: """ Get pipeline template detail. @@ -156,7 +164,7 @@ class RagPipelineService: mode = dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() built_in_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail( - template_id, session=session + template_id, current_tenant_id, session=session ) if built_in_result is None: logger.warning( @@ -169,10 +177,22 @@ class RagPipelineService: mode = "customized" retrieval_instance = PipelineTemplateRetrievalFactory.get_pipeline_template_factory(mode)() customized_result: dict[str, Any] | None = retrieval_instance.get_pipeline_template_detail( - template_id, session=session + template_id, current_tenant_id, session=session ) return customized_result + @staticmethod + def get_customized_pipeline_template_yaml(template_id: str, current_tenant_id: str, *, session: Session) -> str: + yaml_content = session.scalar( + select(PipelineCustomizedTemplate.yaml_content).where( + PipelineCustomizedTemplate.id == template_id, + PipelineCustomizedTemplate.tenant_id == current_tenant_id, + ) + ) + if yaml_content is None: + raise RagPipelineResourceNotFoundError("Customized pipeline template not found.") + return yaml_content + @classmethod def update_customized_pipeline_template( cls, @@ -440,6 +460,11 @@ class RagPipelineService: if not draft_workflow: raise ValueError("No valid workflow found.") + validate_llm_environment_model_references( + graph=draft_workflow.graph_dict, + environment_variables=draft_workflow.environment_variables, + ) + # create new workflow workflow = Workflow.new( tenant_id=pipeline.tenant_id, @@ -568,7 +593,12 @@ class RagPipelineService: node_id=node_id, user_inputs=user_inputs, user_id=account.id, - variable_pool=_build_seeded_variable_pool(default_system_variables()), + variable_pool=_build_seeded_variable_pool( + build_bootstrap_variables( + system_variables=default_system_variables(), + environment_variables=draft_workflow.environment_variables, + ) + ), variable_loader=DraftVarLoader( engine=db.engine, app_id=pipeline.id, @@ -910,7 +940,10 @@ class RagPipelineService: def _handle_node_run_result( self, - getter: Callable[[], tuple[Node, Generator[GraphNodeEventBase, None, None]]], + getter: Callable[ + [], + tuple[Node, Generator[GraphNodeEventBase | ContainerAwaitRequest, None, None]], + ], start_at: float, tenant_id: str, node_id: str, @@ -1229,45 +1262,48 @@ class RagPipelineService: ) return assemble_workflow_node_execution_traces(node_executions, self._node_execution_service_repo) - @classmethod + @staticmethod def publish_customized_pipeline_template( - cls, - pipeline_id: str, + pipeline: Pipeline, + dataset: Dataset, args: dict[str, Any], - current_user: Account | None = None, - current_tenant_id: str | None = None, + current_user: Account, *, session: Session, - ): - """ - Publish customized pipeline template - """ - current_user, _ = resolve_account_fallback(current_user, current_tenant_id) - pipeline = session.get(Pipeline, pipeline_id) - if not pipeline: - raise ValueError("Pipeline not found") + ) -> None: + """Publish a customized template from a caller-validated pipeline and dataset.""" if not pipeline.workflow_id: - raise ValueError("Pipeline workflow not found") - workflow = session.get(Workflow, pipeline.workflow_id) + raise RagPipelineResourceNotFoundError("Pipeline workflow not found") + workflow = session.scalar( + select(Workflow).where( + Workflow.id == pipeline.workflow_id, + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + ) + ) if not workflow: - raise ValueError("Workflow not found") - dataset = pipeline.retrieve_dataset(session=session) - if not dataset: - raise ValueError("Dataset not found") + raise RagPipelineResourceNotFoundError("Workflow not found") + draft_workflow_id = session.scalar( + select(Workflow.id).where( + Workflow.tenant_id == pipeline.tenant_id, + Workflow.app_id == pipeline.id, + Workflow.version == Workflow.VERSION_DRAFT, + ) + ) + if not draft_workflow_id: + raise RagPipelineResourceNotFoundError("Draft workflow not found") # check template name is exist - template_name = args.get("name") - if template_name: - template = session.scalar( - select(PipelineCustomizedTemplate) - .where( - PipelineCustomizedTemplate.name == template_name, - PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id, - ) - .limit(1) + template = session.scalar( + select(PipelineCustomizedTemplate) + .where( + PipelineCustomizedTemplate.name == args["name"], + PipelineCustomizedTemplate.tenant_id == pipeline.tenant_id, ) - if template: - raise ValueError("Template name is already exists") + .limit(1) + ) + if template: + raise ValueError("Template name is already exists") max_position = session.scalar( select(func.max(PipelineCustomizedTemplate.position)).where( @@ -1279,16 +1315,10 @@ class RagPipelineService: rag_pipeline_dsl_service = RagPipelineDslService(session) dsl = rag_pipeline_dsl_service.export_rag_pipeline_dsl(pipeline=pipeline, include_secret=True) - if args.get("icon_info") is None: - args["icon_info"] = {} - if args.get("description") is None: - raise ValueError("Description is required") - if args.get("name") is None: - raise ValueError("Name is required") pipeline_customized_template = PipelineCustomizedTemplate( - name=args.get("name") or "", - description=args.get("description") or "", - icon=args.get("icon_info") or {}, + name=args["name"], + description=args["description"], + icon=args["icon_info"], tenant_id=pipeline.tenant_id, yaml_content=dsl, install_count=0, diff --git a/api/services/rag_pipeline/rag_pipeline_dsl_service.py b/api/services/rag_pipeline/rag_pipeline_dsl_service.py index d562a4b9adf..6977c28b7a6 100644 --- a/api/services/rag_pipeline/rag_pipeline_dsl_service.py +++ b/api/services/rag_pipeline/rag_pipeline_dsl_service.py @@ -21,6 +21,12 @@ from core.file import remote_fetcher from core.helper.name_generator import generate_incremental_name from core.plugin.entities.plugin import PluginDependency from core.rag.index_processor.constant.index_type import IndexTechniqueType +from core.workflow.llm_environment_variable import ( + LLMEnvironmentVariable, + parse_llm_model_selector, + resolve_llm_model_config, + should_resolve_llm_model_selector, +) from core.workflow.nodes.datasource.entities import DatasourceNodeData from core.workflow.nodes.knowledge_index import KNOWLEDGE_INDEX_NODE_TYPE from core.workflow.nodes.knowledge_retrieval.entities import KnowledgeRetrievalNodeData @@ -28,7 +34,7 @@ from extensions.ext_redis import redis_client from factories import variable_factory from graphon.enums import BuiltinNodeTypes from graphon.model_runtime.utils.encoders import jsonable_encoder -from graphon.nodes.llm.entities import LLMNodeData +from graphon.nodes.llm.entities import LLMNodeData, ModelConfig from graphon.nodes.parameter_extractor.entities import ParameterExtractorNodeData from graphon.nodes.question_classifier.entities import QuestionClassifierNodeData from graphon.nodes.tool.entities import ToolNodeData @@ -38,7 +44,7 @@ from models.enums import CollectionBindingType, DatasetRuntimeMode from models.workflow import Workflow, WorkflowType from services.dsl_content import DSL_MAX_SIZE, dsl_content_size from services.dsl_version import check_version_compatibility -from services.entities.dsl_entities import CheckDependenciesResult, ImportMode, ImportStatus +from services.entities.dsl_entities import CheckDependenciesResult, ImportMode, ImportStatus, PendingImportOwner from services.entities.knowledge_entities.rag_pipeline_entities import ( IconInfo, KnowledgeConfiguration, @@ -64,7 +70,7 @@ class RagPipelineImportInfo(BaseModel): dataset_id: str | None = None -class RagPipelinePendingData(BaseModel): +class RagPipelinePendingData(PendingImportOwner): import_mode: str yaml_content: str pipeline_id: str | None @@ -216,6 +222,8 @@ class RagPipelineDslService: # If major version mismatch, store import info in Redis if status == ImportStatus.PENDING: pending_data = RagPipelinePendingData( + tenant_id=account.current_tenant_id, + account_id=account.id, import_mode=import_mode, yaml_content=content, pipeline_id=pipeline_id, @@ -378,6 +386,15 @@ class RagPipelineDslService: error="Invalid import information", ) pending_data = RagPipelinePendingData.model_validate_json(pending_data) + if not pending_data.is_accessible_by( + tenant_id=account.current_tenant_id, + account_id=account.id, + ): + return RagPipelineImportInfo( + id=import_id, + status=ImportStatus.FAILED, + error="Import information expired or does not exist", + ) data = yaml.safe_load(pending_data.yaml_content) pipeline = None @@ -707,6 +724,34 @@ class RagPipelineDslService: :return: dependencies list format like ["langgenius/google"] """ graph = workflow.graph_dict + referenced_llm_nodes = [ + node.get("data", {}) + for node in graph.get("nodes", []) + if node.get("data", {}).get("type") == BuiltinNodeTypes.LLM + and should_resolve_llm_model_selector(node.get("data", {}).get("model_selector")) + ] + environment_variables = ( + {variable.name: variable for variable in workflow.environment_variables} if referenced_llm_nodes else {} + ) + for node_data in referenced_llm_nodes: + try: + selector = parse_llm_model_selector(node_data["model_selector"]) + variable = environment_variables.get(selector[1]) + if not isinstance(variable, LLMEnvironmentVariable): + raise ValueError( + f"LLM environment variable '{selector[1]}' was not found or is not an LLM variable" + ) + node_data["model"] = resolve_llm_model_config( + node_model=ModelConfig.model_validate(node_data.get("model", {})), + variable_name=selector[1], + variable_value=variable.value, + ).model_dump(mode="json") + except ValueError as exc: + logger.warning( + "Skipping unresolved LLM environment model while extracting dependencies for selector %r: %s", + node_data.get("model_selector"), + exc, + ) dependencies = self._extract_dependencies_from_workflow_graph(graph) return dependencies diff --git a/api/services/rag_pipeline/rag_pipeline_task_proxy.py b/api/services/rag_pipeline/rag_pipeline_task_proxy.py index 52ebbce65a9..4c466e37d0b 100644 --- a/api/services/rag_pipeline/rag_pipeline_task_proxy.py +++ b/api/services/rag_pipeline/rag_pipeline_task_proxy.py @@ -5,7 +5,7 @@ from functools import cached_property from core.app.entities.rag_pipeline_invoke_entities import RagPipelineInvokeEntity from core.rag.pipeline.queue import TenantIsolatedTaskQueue -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from extensions.ext_database import db from services.feature_service import FeatureService from services.file_service import FileService diff --git a/api/services/rag_pipeline/rag_pipeline_transform_service.py b/api/services/rag_pipeline/rag_pipeline_transform_service.py index 6a7902c1908..90dd6e081bd 100644 --- a/api/services/rag_pipeline/rag_pipeline_transform_service.py +++ b/api/services/rag_pipeline/rag_pipeline_transform_service.py @@ -6,7 +6,6 @@ from typing import Any from uuid import uuid4 import yaml -from flask_login import current_user from sqlalchemy import select from sqlalchemy.orm import Session @@ -19,23 +18,34 @@ from core.rag.retrieval.retrieval_methods import RetrievalMethod from factories import variable_factory from models.dataset import Dataset, Document, DocumentPipelineExecutionLog, Pipeline from models.enums import DatasetRuntimeMode, DataSourceType -from models.model import UploadFile from models.workflow import Workflow, WorkflowType from services.entities.knowledge_entities.rag_pipeline_entities import KnowledgeConfiguration, RetrievalSetting +from services.errors.rag_pipeline import RagPipelineResourceNotFoundError +from services.file_service import FileService from services.plugin.plugin_migration import PluginMigration logger = logging.getLogger(__name__) class RagPipelineTransformService: - def transform_dataset(self, dataset_id: str, session: Session): - dataset = session.get(Dataset, dataset_id) - if not dataset: - raise ValueError("Dataset not found") + def transform_dataset(self, dataset: Dataset, account_id: str, session: Session): + """Transform a vendor dataset within the caller-owned transaction. + + The caller must resolve and authorize ``dataset`` before entering this service. Plugin provisioning is an + intentionally non-transactional prerequisite; database changes are committed only after the pipeline and + migrated document metadata have been persisted. + """ if dataset.pipeline_id and dataset.runtime_mode == DatasetRuntimeMode.RAG_PIPELINE: + pipeline = session.scalar( + select(Pipeline) + .where(Pipeline.id == dataset.pipeline_id, Pipeline.tenant_id == dataset.tenant_id) + .limit(1) + ) + if pipeline is None: + raise RagPipelineResourceNotFoundError("Pipeline not found") return { - "pipeline_id": dataset.pipeline_id, - "dataset_id": dataset_id, + "pipeline_id": pipeline.id, + "dataset_id": dataset.id, "status": "success", } if dataset.provider != "vendor": @@ -44,11 +54,11 @@ class RagPipelineTransformService: indexing_technique = dataset.indexing_technique if not datasource_type and not indexing_technique: - return self._transform_to_empty_pipeline(dataset, session=session) + return self._transform_to_empty_pipeline(dataset, account_id=account_id, session=session) - doc_form = dataset.doc_form + doc_form = dataset.get_doc_form(session=session) if not doc_form: - return self._transform_to_empty_pipeline(dataset, session=session) + return self._transform_to_empty_pipeline(dataset, account_id=account_id, session=session) retrieval_model = RetrievalSetting.model_validate(dataset.retrieval_model) if dataset.retrieval_model else None pipeline_yaml = self._get_transform_yaml(doc_form, datasource_type, indexing_technique) # deal dependencies @@ -69,8 +79,6 @@ class RagPipelineTransformService: node = self._deal_file_extensions(node) if node.get("data", {}).get("type") == "knowledge-index": knowledge_configuration = KnowledgeConfiguration.model_validate(node.get("data", {})) - if dataset.tenant_id != current_user.current_tenant_id: - raise ValueError("Unauthorized") node = self._deal_knowledge_index( knowledge_configuration, dataset, indexing_technique, retrieval_model, node ) @@ -80,7 +88,12 @@ class RagPipelineTransformService: workflow_data["graph"] = graph pipeline_yaml["workflow"] = workflow_data # create pipeline - pipeline = self._create_pipeline(pipeline_yaml, session=session) + pipeline = self._create_pipeline( + pipeline_yaml, + tenant_id=dataset.tenant_id, + account_id=account_id, + session=session, + ) # save chunk structure to dataset if doc_form == IndexStructureType.PARENT_CHILD_INDEX: @@ -99,7 +112,7 @@ class RagPipelineTransformService: session.commit() return { "pipeline_id": pipeline.id, - "dataset_id": dataset_id, + "dataset_id": dataset.id, "status": "success", } @@ -195,6 +208,8 @@ class RagPipelineTransformService: self, data: dict[str, Any], *, + tenant_id: str, + account_id: str, session: Session, ) -> Pipeline: """Create a new app or update an existing one.""" @@ -218,11 +233,11 @@ class RagPipelineTransformService: # Create new app pipeline = Pipeline( - tenant_id=current_user.current_tenant_id, + tenant_id=tenant_id, name=pipeline_data.get("name", ""), description=pipeline_data.get("description", ""), - created_by=current_user.id, - updated_by=current_user.id, + created_by=account_id, + updated_by=account_id, is_published=True, is_public=True, ) @@ -238,7 +253,7 @@ class RagPipelineTransformService: type=WorkflowType.RAG_PIPELINE, version="draft", graph=json.dumps(graph), - created_by=current_user.id, + created_by=account_id, environment_variables=environment_variables, conversation_variables=conversation_variables, rag_pipeline_variables=rag_pipeline_variables_list, @@ -250,7 +265,7 @@ class RagPipelineTransformService: type=WorkflowType.RAG_PIPELINE, version=str(datetime.now(UTC).replace(tzinfo=None)), graph=json.dumps(graph), - created_by=current_user.id, + created_by=account_id, environment_variables=environment_variables, conversation_variables=conversation_variables, rag_pipeline_variables=rag_pipeline_variables_list, @@ -292,19 +307,19 @@ class RagPipelineTransformService: logger.debug("Installing missing pipeline plugins %s", package_identifiers_to_install) PluginService.install_from_marketplace_pkg(tenant_id, package_identifiers_to_install) - def _transform_to_empty_pipeline(self, dataset: Dataset, *, session: Session): + def _transform_to_empty_pipeline(self, dataset: Dataset, *, account_id: str, session: Session): pipeline = Pipeline( tenant_id=dataset.tenant_id, name=dataset.name, description=dataset.description, - created_by=current_user.id, + created_by=account_id, ) session.add(pipeline) session.flush() dataset.pipeline_id = pipeline.id dataset.runtime_mode = DatasetRuntimeMode.RAG_PIPELINE - dataset.updated_by = current_user.id + dataset.updated_by = account_id dataset.updated_at = datetime.now(UTC).replace(tzinfo=None) session.add(dataset) session.commit() @@ -320,18 +335,23 @@ class RagPipelineTransformService: jina_node_id = "1752491761974" firecrawl_node_id = "1752565402678" - documents = session.scalars(select(Document).where(Document.dataset_id == dataset.id)).all() + documents = session.scalars( + select(Document).where(Document.dataset_id == dataset.id, Document.tenant_id == dataset.tenant_id) + ).all() for document in documents: data_source_info_dict = document.data_source_info_dict if not data_source_info_dict: continue if document.data_source_type == DataSourceType.UPLOAD_FILE: - document.data_source_type = DataSourceType.LOCAL_FILE file_id = data_source_info_dict.get("upload_file_id") if file_id: - file = session.get(UploadFile, file_id) + file_id = str(file_id) + file = FileService.get_upload_files_by_ids(dataset.tenant_id, [file_id], session=session).get( + file_id + ) if file: + document.data_source_type = DataSourceType.LOCAL_FILE data_source_info = json.dumps( { "real_file_id": file_id, diff --git a/api/services/recommend_app/remote/remote_retrieval.py b/api/services/recommend_app/remote/remote_retrieval.py index c676ec907e0..d30306382ff 100644 --- a/api/services/recommend_app/remote/remote_retrieval.py +++ b/api/services/recommend_app/remote/remote_retrieval.py @@ -1,7 +1,9 @@ import logging +import threading from typing import Any, override import httpx +from cachetools import TTLCache from flask import has_request_context, request from sqlalchemy.orm import Session @@ -13,6 +15,11 @@ from services.recommend_app.recommend_app_type import RecommendAppType logger = logging.getLogger(__name__) +_REMOTE_FETCH_CACHE_MAXSIZE = 64 +_remote_fetch_cache: TTLCache[tuple[str, str], dict[str, Any]] | None = None +_remote_fetch_cache_ttl: int | None = None +_remote_fetch_cache_lock = threading.Lock() + def _current_origin_headers() -> dict[str, str]: origin = request.headers.get("Origin") if has_request_context() else None @@ -25,6 +32,61 @@ def _current_origin_headers() -> dict[str, str]: return {"Origin": console_web_url} +def _remote_fetch_cache_key(url: str, headers: dict[str, str]) -> tuple[str, str]: + return url, headers.get("Origin", "") + + +def _hosted_fetch_cache_ttl() -> int: + ttl = dify_config.HOSTED_FETCH_APP_TEMPLATES_CACHE_TTL + if isinstance(ttl, int) and not isinstance(ttl, bool): + return ttl + return 600 + + +def _get_remote_fetch_cache() -> TTLCache[tuple[str, str], dict[str, Any]] | None: + ttl = _hosted_fetch_cache_ttl() + if ttl <= 0: + return None + + global _remote_fetch_cache, _remote_fetch_cache_ttl + if _remote_fetch_cache is None or _remote_fetch_cache_ttl != ttl: + with _remote_fetch_cache_lock: + if _remote_fetch_cache is None or _remote_fetch_cache_ttl != ttl: + _remote_fetch_cache = TTLCache(maxsize=_REMOTE_FETCH_CACHE_MAXSIZE, ttl=ttl) + _remote_fetch_cache_ttl = ttl + return _remote_fetch_cache + + +def clear_remote_fetch_cache() -> None: + """Reset the in-memory remote fetch cache (used by tests).""" + global _remote_fetch_cache, _remote_fetch_cache_ttl + with _remote_fetch_cache_lock: + _remote_fetch_cache = None + _remote_fetch_cache_ttl = None + + +def _fetch_remote_payload(url: str) -> tuple[int, dict[str, Any] | None]: + headers = _current_origin_headers() + cache_key = _remote_fetch_cache_key(url, headers) + cache = _get_remote_fetch_cache() + if cache is not None: + with _remote_fetch_cache_lock: + cached = cache.get(cache_key) + if cached is not None: + return 200, cached + + response = httpx.get(url, headers=headers, timeout=httpx.Timeout(10.0, connect=3.0)) + status_code = response.status_code + if status_code != 200: + return status_code, None + + result: dict[str, Any] = response.json() + if cache is not None: + with _remote_fetch_cache_lock: + cache[cache_key] = result + return status_code, result + + class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): """ Retrieval recommended app from dify official. @@ -75,10 +137,9 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): """ domain = dify_config.HOSTED_FETCH_APP_TEMPLATES_REMOTE_DOMAIN url = f"{domain}/apps/{app_id}" - response = httpx.get(url, headers=_current_origin_headers(), timeout=httpx.Timeout(10.0, connect=3.0)) - if response.status_code != 200: + status_code, data = _fetch_remote_payload(url) + if status_code != 200: return None - data: dict[str, Any] = response.json() return data @classmethod @@ -90,11 +151,10 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): """ domain = dify_config.HOSTED_FETCH_APP_TEMPLATES_REMOTE_DOMAIN url = f"{domain}/apps?language={language}" - response = httpx.get(url, headers=_current_origin_headers(), timeout=httpx.Timeout(10.0, connect=3.0)) - if response.status_code != 200: - raise ValueError(f"fetch recommended apps failed, status code: {response.status_code}") + status_code, result = _fetch_remote_payload(url) + if status_code != 200: + raise ValueError(f"fetch recommended apps failed, status code: {status_code}") - result: dict[str, Any] = response.json() return result @classmethod @@ -106,9 +166,8 @@ class RemoteRecommendAppRetrieval(RecommendAppRetrievalBase): """ domain = dify_config.HOSTED_FETCH_APP_TEMPLATES_REMOTE_DOMAIN url = f"{domain}/apps/learn-dify?language={language}" - response = httpx.get(url, headers=_current_origin_headers(), timeout=httpx.Timeout(10.0, connect=3.0)) - if response.status_code != 200: - raise ValueError(f"fetch learn dify apps failed, status code: {response.status_code}") + status_code, result = _fetch_remote_payload(url) + if status_code != 200: + raise ValueError(f"fetch learn dify apps failed, status code: {status_code}") - result: dict[str, Any] = response.json() return result diff --git a/api/services/recommended_app_service.py b/api/services/recommended_app_service.py index 3bdfbe6f365..7b09e2b4005 100644 --- a/api/services/recommended_app_service.py +++ b/api/services/recommended_app_service.py @@ -4,12 +4,19 @@ from sqlalchemy import select from sqlalchemy.orm import Session from configs import dify_config +from enums import DeploymentEdition from models.model import AccountTrialAppRecord, App, TrialApp -from services.feature_service import FeatureService from services.recommend_app.recommend_app_factory import RecommendAppRetrievalFactory class RecommendedAppService: + """Own recommended app retrieval and Cloud-only trial eligibility.""" + + @staticmethod + def is_trial_app_enabled() -> bool: + """Return whether trial execution is enabled for this deployment.""" + return dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and dify_config.ENABLE_TRIAL_APP + @classmethod def get_app(cls, app_id: str, *, session: Session) -> App | None: """Return a normal app only when it belongs to the recommended catalog.""" @@ -38,11 +45,12 @@ class RecommendedAppService: ) ) - if FeatureService.get_system_features().enable_trial_app: - apps = result["recommended_apps"] - for app in apps: - app_id = app["app_id"] - app["can_trial"] = cls._can_trial_app(session, app_id) + apps = result["recommended_apps"] + trial_app_ids = ( + cls._get_trial_app_ids(session, [app["app_id"] for app in apps]) if cls.is_trial_app_enabled() else set() + ) + for app in apps: + app["can_trial"] = app["app_id"] in trial_app_ids return result @classmethod @@ -56,11 +64,14 @@ class RecommendedAppService: retrieval_instance = RecommendAppRetrievalFactory.get_recommend_app_factory(mode)() result = retrieval_instance.get_learn_dify_apps(language, session=session) - if FeatureService.get_system_features().enable_trial_app: - for app in result["recommended_apps"]: - app["can_trial"] = cls._can_trial_app(session, app["app_id"]) + apps = result["recommended_apps"] + trial_app_ids = ( + cls._get_trial_app_ids(session, [app["app_id"] for app in apps]) if cls.is_trial_app_enabled() else set() + ) + for app in apps: + app["can_trial"] = app["app_id"] in trial_app_ids - return {"recommended_apps": result["recommended_apps"]} + return {"recommended_apps": apps} @classmethod def get_recommend_app_detail(cls, app_id: str, *, session: Session) -> dict[str, Any] | None: @@ -74,9 +85,7 @@ class RecommendedAppService: result: dict[str, Any] | None = retrieval_instance.get_recommend_app_detail(app_id, session=session) if result is None: return None - if FeatureService.get_system_features().enable_trial_app: - app_id = result["id"] - result["can_trial"] = cls._can_trial_app(session, app_id) + result["can_trial"] = cls.is_trial_app_enabled() and cls._can_trial_app(session, result["id"]) return result @classmethod @@ -102,3 +111,9 @@ class RecommendedAppService: def _can_trial_app(session: Session, app_id: str) -> bool: trial_app_model = session.scalar(select(TrialApp).where(TrialApp.app_id == app_id).limit(1)) return trial_app_model is not None + + @staticmethod + def _get_trial_app_ids(session: Session, app_ids: list[str]) -> set[str]: + if not app_ids: + return set() + return set(session.scalars(select(TrialApp.app_id).where(TrialApp.app_id.in_(app_ids))).all()) diff --git a/api/services/retention/conversation/messages_clean_policy.py b/api/services/retention/conversation/messages_clean_policy.py index 5196344212b..f3e3cd8fb46 100644 --- a/api/services/retention/conversation/messages_clean_policy.py +++ b/api/services/retention/conversation/messages_clean_policy.py @@ -5,7 +5,7 @@ from dataclasses import dataclass from typing import Protocol, override from configs import dify_config -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from services.billing_service import BillingService, SubscriptionPlan logger = logging.getLogger(__name__) @@ -45,7 +45,7 @@ class MessagesCleanPolicy(Protocol): class BillingDisabledPolicy(MessagesCleanPolicy): """ - Policy for community or enterpriseedition (billing disabled). + Policy for self-hosted editions, which do not use Cloud billing plans. No special filter logic, just return all message ids. """ @@ -61,7 +61,7 @@ class BillingDisabledPolicy(MessagesCleanPolicy): class BillingSandboxPolicy(MessagesCleanPolicy): """ - Policy for sandbox plan tenants in cloud edition (billing enabled). + Policy for sandbox-plan tenants in the Cloud edition. Filters messages based on sandbox plan expiration rules: - Skip tenants in the whitelist @@ -186,24 +186,22 @@ def create_message_clean_policy( """ Factory function to create the appropriate message clean policy. - Determines which policy to use based on BILLING_ENABLED configuration: - - If BILLING_ENABLED is True: returns BillingSandboxPolicy - - If BILLING_ENABLED is False: returns BillingDisabledPolicy + Cloud uses BillingSandboxPolicy; self-hosted editions use BillingDisabledPolicy. Args: graceful_period_days: Grace period in days after subscription expiration (default: 21) current_timestamp: Current Unix timestamp for testing (default: None, uses current time) """ - if not dify_config.BILLING_ENABLED: - logger.info("create_message_clean_policy: billing disabled, using BillingDisabledPolicy") + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: + logger.info("create_message_clean_policy: self-hosted edition, using BillingDisabledPolicy") return BillingDisabledPolicy() - # Billing enabled - fetch whitelist from BillingService + # Cloud deployment - fetch whitelist from BillingService. tenant_whitelist = BillingService.get_expired_subscription_cleanup_whitelist() plan_provider = BillingService.get_plan_bulk_with_cache logger.info( - "create_message_clean_policy: billing enabled, using BillingSandboxPolicy " + "create_message_clean_policy: Cloud edition, using BillingSandboxPolicy " "(graceful_period_days=%s, whitelist=%s)", graceful_period_days, tenant_whitelist, diff --git a/api/services/retention/conversation/messages_clean_service.py b/api/services/retention/conversation/messages_clean_service.py index 1e9f0bf1493..ffe130af681 100644 --- a/api/services/retention/conversation/messages_clean_service.py +++ b/api/services/retention/conversation/messages_clean_service.py @@ -169,8 +169,8 @@ class MessagesCleanService: """ Service for cleaning expired messages based on retention policies. - Compatible with non cloud edition (billing disabled): all messages in the time range will be deleted. - If billing is enabled: only sandbox plan tenant messages are deleted (with whitelist and grace period support). + In self-hosted editions, all messages in the time range are deleted. + In the Cloud edition, only sandbox-plan tenant messages are deleted, with whitelist and grace-period support. """ def __init__( diff --git a/api/services/retention/workflow_run/archive_download_preparation.py b/api/services/retention/workflow_run/archive_download_preparation.py index 510a20a5319..cb9200d23da 100644 --- a/api/services/retention/workflow_run/archive_download_preparation.py +++ b/api/services/retention/workflow_run/archive_download_preparation.py @@ -13,7 +13,6 @@ import logging import re import zipfile from collections.abc import Sequence -from typing import cast import pyarrow as pa import pyarrow.compute as pc @@ -310,7 +309,7 @@ def _validate_manifest( if not manifest["tables"]: raise ValueError("manifest tables must not be empty") for table_name, raw_entry in manifest["tables"].items(): - entry = cast(ArchiveBundleTableManifestEntry, raw_entry) + entry = raw_entry expected_object_key = f"{object_prefix}/{table_name}.parquet" if entry["object_key"] != expected_object_key: raise ValueError( diff --git a/api/services/retention/workflow_run/archive_paid_plan_workflow_run.py b/api/services/retention/workflow_run/archive_paid_plan_workflow_run.py index eb69b198202..3890ef58680 100644 --- a/api/services/retention/workflow_run/archive_paid_plan_workflow_run.py +++ b/api/services/retention/workflow_run/archive_paid_plan_workflow_run.py @@ -46,7 +46,7 @@ from sqlalchemy import inspect, select from sqlalchemy.orm import Session, sessionmaker from configs import dify_config -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from extensions.ext_database import db from graphon.enums import WorkflowType from libs.archive_storage import ( @@ -535,8 +535,8 @@ class WorkflowRunArchiver: if self.paid_tenant_ids is not None: return tenant_ids & self.paid_tenant_ids - if not dify_config.BILLING_ENABLED: - # If billing is not enabled, treat all tenants as paid + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: + # Self-hosted editions have no Cloud billing plans, so treat all tenants as paid. return tenant_ids if not tenant_ids: diff --git a/api/services/retention/workflow_run/clear_free_plan_expired_workflow_run_logs.py b/api/services/retention/workflow_run/clear_free_plan_expired_workflow_run_logs.py index 3652997f8af..0e1a9ba0c0b 100644 --- a/api/services/retention/workflow_run/clear_free_plan_expired_workflow_run_logs.py +++ b/api/services/retention/workflow_run/clear_free_plan_expired_workflow_run_logs.py @@ -16,7 +16,7 @@ import click from sqlalchemy.orm import Session, sessionmaker from configs import dify_config -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from extensions.ext_database import db from repositories.api_workflow_run_repository import ( APIWorkflowRunRepository, @@ -478,7 +478,7 @@ class WorkflowRunCleanup: def _filter_free_tenants(self, tenant_ids: Iterable[str]) -> set[str]: tenant_id_list = sorted(set(tenant_ids)) - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return set(tenant_id_list) if not tenant_id_list: @@ -538,7 +538,7 @@ class WorkflowRunCleanup: if self._cleanup_whitelist is not None: return self._cleanup_whitelist - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: self._cleanup_whitelist = set() return self._cleanup_whitelist diff --git a/api/services/schema_definition_service.py b/api/services/schema_definition_service.py new file mode 100644 index 00000000000..8fb61b401d1 --- /dev/null +++ b/api/services/schema_definition_service.py @@ -0,0 +1,24 @@ +"""Application service for querying Console schema definitions.""" + +import logging +from collections.abc import Callable, Mapping +from typing import Any, Protocol + +logger = logging.getLogger(__name__) + + +class SchemaDefinitionSource(Protocol): + def get_all_schema_definitions(self, version: str = "v1") -> list[Mapping[str, Any]]: ... + + +class SchemaDefinitionService: + def __init__(self, *, source_factory: Callable[[], SchemaDefinitionSource]) -> None: + self._source_factory = source_factory + + def list(self) -> tuple[Mapping[str, Any], ...]: + try: + source = self._source_factory() + return tuple(source.get_all_schema_definitions()) + except Exception: + logger.exception("Failed to get schema definitions from local registry") + return () diff --git a/api/services/setup_adapters.py b/api/services/setup_adapters.py new file mode 100644 index 00000000000..e8fc6098905 --- /dev/null +++ b/api/services/setup_adapters.py @@ -0,0 +1,43 @@ +"""Infrastructure adapters for the first-time setup application service.""" + +from contextlib import AbstractContextManager +from typing import override + +from sqlalchemy.orm import Session, sessionmaker + +from extensions.ext_redis import RedisClientWrapper +from services.account_service import RegisterService +from services.setup_service import SetupAccountProvisioner, SetupInput, SetupLock + +_SETUP_LOCK_KEY = "setup:initialize" +_SETUP_LOCK_TIMEOUT_SECONDS = 300 + + +class RegisterServiceAccountProvisioner(SetupAccountProvisioner): + def __init__(self, client: sessionmaker[Session]) -> None: + self._client = client + + @override + def provision(self, setup: SetupInput) -> None: + with self._client() as session: + RegisterService.setup( + email=setup.email, + name=setup.name, + password=setup.password, + ip_address=setup.ip_address, + language=setup.language, + session=session, + ) + + +class RedisSetupLock(SetupLock): + def __init__(self, *, client: RedisClientWrapper) -> None: + self._client = client + + @override + def acquire(self) -> AbstractContextManager[None]: + return self._client.lock( + _SETUP_LOCK_KEY, + timeout=_SETUP_LOCK_TIMEOUT_SECONDS, + blocking_timeout=_SETUP_LOCK_TIMEOUT_SECONDS, + ) diff --git a/api/services/setup_service.py b/api/services/setup_service.py new file mode 100644 index 00000000000..f12f62796f8 --- /dev/null +++ b/api/services/setup_service.py @@ -0,0 +1,83 @@ +"""Application service for first-time Dify setup.""" + +from contextlib import AbstractContextManager +from dataclasses import dataclass +from datetime import datetime +from typing import Protocol + + +@dataclass(frozen=True, slots=True) +class SetupInput: + email: str + name: str + password: str + ip_address: str + language: str | None + + +@dataclass(frozen=True, slots=True) +class SetupStatus: + completed: bool + setup_at: datetime | None = None + + +class SetupState(Protocol): + def get_setup_at(self) -> datetime | None: ... + + def has_tenants(self) -> bool: ... + + +class SetupAccountProvisioner(Protocol): + def provision(self, setup: SetupInput) -> None: ... + + +class SetupLock(Protocol): + def acquire(self) -> AbstractContextManager[None]: ... + + +class SetupAlreadyCompletedError(Exception): + """Raised when setup has already created persistent installation state.""" + + +class InitializationValidationRequiredError(Exception): + """Raised when initialization-password validation has not completed.""" + + +class SetupService: + def __init__( + self, + *, + state: SetupState, + accounts: SetupAccountProvisioner, + lock: SetupLock, + setup_required: bool, + ) -> None: + self._state = state + self._accounts = accounts + self._lock = lock + self._setup_required = setup_required + + def get_status(self) -> SetupStatus: + if not self._setup_required: + return SetupStatus(completed=True) + + setup_at = self._state.get_setup_at() + return SetupStatus(completed=setup_at is not None, setup_at=setup_at) + + def initialize(self, setup: SetupInput, *, initialization_validated: bool) -> None: + with self._lock.acquire(): + if self._state.get_setup_at() is not None or self._state.has_tenants(): + raise SetupAlreadyCompletedError + + if not initialization_validated: + raise InitializationValidationRequiredError + + self._accounts.provision( + SetupInput( + email=setup.email.lower(), + name=setup.name, + password=setup.password, + ip_address=setup.ip_address, + language=setup.language, + ) + ) diff --git a/api/services/snippet_dsl_service.py b/api/services/snippet_dsl_service.py index 14ac9a21f6b..22f495a2370 100644 --- a/api/services/snippet_dsl_service.py +++ b/api/services/snippet_dsl_service.py @@ -19,12 +19,20 @@ from models import Account from models.snippet import CustomizedSnippet, SnippetType from models.workflow import Workflow from services.agent.dsl_service import AgentDslService +from services.agent.retirement_service import WorkflowAgentRetirementService from services.agent.workflow_publish_service import WorkflowAgentPublishService from services.dsl_content import DSL_MAX_SIZE, dsl_content_size from services.dsl_version import check_version_compatibility -from services.entities.dsl_entities import CheckDependenciesResult, DslImportWarning, ImportMode, ImportStatus +from services.entities.dsl_entities import ( + CheckDependenciesResult, + DslImportWarning, + ImportMode, + ImportStatus, + PendingImportOwner, +) from services.plugin.dependencies_analysis import DependenciesAnalysisService from services.snippet_service import SNIPPET_FORBIDDEN_NODE_TYPES, SnippetService +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection logger = logging.getLogger(__name__) @@ -49,7 +57,7 @@ def _check_version_compatibility(imported_version: str) -> ImportStatus: return check_version_compatibility(imported_version, CURRENT_DSL_VERSION) -class SnippetPendingData(BaseModel): +class SnippetPendingData(PendingImportOwner): import_mode: str yaml_content: str name: str | None = None @@ -229,6 +237,8 @@ class SnippetDslService: # If major version mismatch, store import info in Redis if status == ImportStatus.PENDING: pending_data = SnippetPendingData( + tenant_id=account.current_tenant_id, + account_id=account.id, import_mode=import_mode, yaml_content=content, name=name, @@ -313,6 +323,15 @@ class SnippetDslService: pending_data_str = pending_data.decode("utf-8") if isinstance(pending_data, bytes) else pending_data pending = SnippetPendingData.model_validate_json(pending_data_str) + if not pending.is_accessible_by( + tenant_id=account.current_tenant_id, + account_id=account.id, + ): + return SnippetImportInfo( + id=import_id, + status=ImportStatus.FAILED, + error="Import information expired or does not exist", + ) data = yaml.safe_load(pending.yaml_content) if not isinstance(data, dict): @@ -426,6 +445,7 @@ class SnippetDslService: self._session.flush() # Create or update draft workflow + retirement_candidates: set[str] = set() if workflow_data: graph = workflow_data.get("graph", {}) raw_agent_packages = data.get("agent_packages") or {} @@ -444,10 +464,10 @@ class SnippetDslService: unique_hash=unique_hash, account=account, input_fields=input_fields, - sync_agent_bindings=not raw_agent_packages, + sync_agent_bindings=False, ) if raw_agent_packages: - _, warnings = AgentDslService(self._session).import_workflow_packages( + _, warnings, retirement_candidates = AgentDslService(self._session).import_workflow_packages( workflow=draft_workflow, portable_graph=graph, raw_packages=raw_agent_packages, @@ -458,8 +478,29 @@ class SnippetDslService: session=self._session, draft_workflow=draft_workflow, ) + else: + retirement_candidates = WorkflowAgentPublishService.sync_agent_bindings_for_draft( + session=self._session, + draft_workflow=draft_workflow, + account_id=account.id, + ) + WorkflowAgentPublishService.validate_agent_nodes_for_draft_sync( + session=self._session, + draft_workflow=draft_workflow, + ) self._session.commit() + if workflow_data: + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=snippet.tenant_id, + agent_ids=retirement_candidates, + account_id=account.id, + ) + enqueue_agent_resource_collection( + tenant_id=snippet.tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) return snippet def export_snippet_dsl(self, snippet: CustomizedSnippet, include_secret: bool = False) -> str: diff --git a/api/services/snippet_service.py b/api/services/snippet_service.py index c23895a4d60..f34f9789fa5 100644 --- a/api/services/snippet_service.py +++ b/api/services/snippet_service.py @@ -8,7 +8,9 @@ from typing import Any from sqlalchemy import delete, event, func, select from sqlalchemy.orm import Session, sessionmaker +from configs import dify_config from core.workflow.node_factory import LATEST_VERSION, NODE_TYPE_CLASSES_MAPPING +from enums import DeploymentEdition from graphon.enums import BuiltinNodeTypes, NodeType from libs.infinite_scroll_pagination import InfiniteScrollPagination from models import Account, TagBinding @@ -35,6 +37,7 @@ from models.workflow import ( WorkflowType, ) from repositories.factory import DifyAPIRepositoryFactory +from services.agent.retirement_service import WorkflowAgentRetirementService from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError from services.tag_service import TagService from services.workflow_node_execution_trace_service import ( @@ -42,6 +45,7 @@ from services.workflow_node_execution_trace_service import ( assemble_workflow_node_execution_traces, ) from services.workflow_restore import apply_published_workflow_snapshot_to_draft +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection logger = logging.getLogger(__name__) @@ -80,7 +84,7 @@ class SnippetService: @contextmanager def _session_scope(self) -> Generator[Session, None, None]: - current_session = getattr(self, "_session", None) + current_session = self._session if current_session is not None: yield current_session return @@ -89,7 +93,7 @@ class SnippetService: yield session def _commit_if_owned(self, session: Session) -> None: - if getattr(self, "_session", None) is None: + if self._session is None: session.commit() @staticmethod @@ -148,10 +152,9 @@ class SnippetService: @staticmethod def _delete_archived_workflow_run_files(*, snippet: CustomizedSnippet) -> None: - from configs import dify_config from libs.archive_storage import ArchiveStorageNotConfiguredError, get_archive_storage - if not (dify_config.BILLING_ENABLED and dify_config.ARCHIVE_STORAGE_ENABLED): + if not (dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and dify_config.ARCHIVE_STORAGE_ENABLED): return prefix = f"{snippet.tenant_id}/app_id={snippet.id}/" @@ -600,12 +603,13 @@ class SnippetService: from services.agent.workflow_publish_service import WorkflowAgentPublishService + retirement_candidates: set[str] = set() with self._session_scope() as session: session.add(workflow) session.add(snippet) if sync_agent_bindings: session.flush() - WorkflowAgentPublishService.sync_agent_bindings_for_draft( + retirement_candidates = WorkflowAgentPublishService.sync_agent_bindings_for_draft( session=session, draft_workflow=workflow, account_id=account.id, @@ -615,6 +619,17 @@ class SnippetService: draft_workflow=workflow, ) self._commit_if_owned(session) + if self._session is None: + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=snippet.tenant_id, + agent_ids=retirement_candidates, + account_id=account.id, + ) + enqueue_agent_resource_collection( + tenant_id=snippet.tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) return workflow def restore_published_workflow_to_draft( @@ -656,13 +671,24 @@ class SnippetService: session.flush() from services.agent.workflow_publish_service import WorkflowAgentPublishService - WorkflowAgentPublishService.restore_agent_node_bindings_to_draft( + retirement_candidates = WorkflowAgentPublishService.restore_agent_node_bindings_to_draft( session=session, source_workflow=source_workflow, draft_workflow=draft_workflow, account_id=account.id, ) self._commit_if_owned(session) + if self._session is None: + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=snippet.tenant_id, + agent_ids=retirement_candidates, + account_id=account.id, + ) + enqueue_agent_resource_collection( + tenant_id=snippet.tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) return draft_workflow def publish_workflow( @@ -671,7 +697,7 @@ class SnippetService: session: Session, snippet: CustomizedSnippet, account: Account, - ) -> Workflow: + ) -> tuple[Workflow, set[str]]: """ Publish the draft workflow as a new version. @@ -693,6 +719,13 @@ class SnippetService: SnippetService.validate_snippet_graph_forbidden_nodes(draft_workflow.graph_dict) + from core.workflow.llm_environment_variable import validate_llm_environment_model_references + + validate_llm_environment_model_references( + graph=draft_workflow.graph_dict, + environment_variables=draft_workflow.environment_variables, + ) + from services.agent.workflow_publish_service import WorkflowAgentPublishService WorkflowAgentPublishService.validate_agent_nodes_for_publish( @@ -715,7 +748,7 @@ class SnippetService: kind=WorkflowKind.SNIPPET.value, ) session.add(workflow) - WorkflowAgentPublishService.copy_agent_node_bindings_to_published( + retirement_candidates = WorkflowAgentPublishService.copy_agent_node_bindings_to_published( session=session, draft_workflow=draft_workflow, published_workflow=workflow, @@ -728,7 +761,7 @@ class SnippetService: snippet.updated_by = account.id session.add(snippet) - return workflow + return workflow, retirement_candidates def get_all_published_workflows( self, diff --git a/api/services/telemetry_service.py b/api/services/telemetry_service.py new file mode 100644 index 00000000000..adf3565cc0e --- /dev/null +++ b/api/services/telemetry_service.py @@ -0,0 +1,165 @@ +import logging +import platform +import uuid +from datetime import datetime +from typing import Literal + +import httpx +from sqlalchemy import select +from sqlalchemy.orm import Session + +from configs import dify_config +from enums import DeploymentEdition +from libs.datetime_utils import naive_utc_now +from models.model import DifySetup + +logger = logging.getLogger(__name__) + +TelemetryEvent = Literal["install", "heartbeat"] + +SCHEMA_VERSION = 1 + + +class CommunityTelemetryService: + @classmethod + def report_install(cls, *, session: Session) -> bool: + setup = cls._get_setup(session) + if setup is None: + return False + + if setup.instance_id is None: + setup.instance_id = str(uuid.uuid4()) + session.add(setup) + session.commit() + + payload = cls._build_payload(setup, "install") + if not cls._send_event(payload): + return False + + setup.install_reported_at = naive_utc_now() + session.add(setup) + session.commit() + return True + + @classmethod + def report_heartbeat(cls, *, session: Session, now: datetime | None = None) -> bool: + setup = cls._get_setup(session) + if setup is None: + return False + + if setup.instance_id is None: + setup.instance_id = str(uuid.uuid4()) + session.add(setup) + session.commit() + + now = now or naive_utc_now() + if not cls._is_heartbeat_due(setup, now): + return False + + if setup.install_reported_at is None: + cls.report_install(session=session) + + payload = cls._build_payload(setup, "heartbeat") + if not cls._send_event(payload): + return False + + setup.last_heartbeat_at = now + session.add(setup) + session.commit() + return True + + @classmethod + def _get_setup(cls, session: Session) -> DifySetup | None: + return session.scalar(select(DifySetup).order_by(DifySetup.setup_at.asc()).limit(1)) + + @classmethod + def _is_enabled(cls) -> bool: + return ( + dify_config.DEPLOYMENT_EDITION == DeploymentEdition.COMMUNITY + and not dify_config.DISABLE_TELEMETRY + and not dify_config.DO_NOT_TRACK + and not dify_config.CI + and bool(dify_config.TELEMETRY_ENDPOINT) + ) + + @classmethod + def _build_payload(cls, setup: DifySetup, event: TelemetryEvent) -> dict[str, str | int]: + payload: dict[str, str | int] = { + "event": event, + "instance_id": setup.instance_id or "", + "version": setup.version if event == "install" else dify_config.project.version, + "edition": dify_config.DEPLOYMENT_EDITION.value, + "deployment_type": "unknown", + "schema_version": SCHEMA_VERSION, + "os": cls._normalize_os(platform.system()), + "arch": cls._normalize_arch(platform.machine()), + "sent_at": cls._format_datetime(naive_utc_now()), + } + + if event == "install": + payload["installed_at"] = cls._format_datetime(setup.setup_at) + + return payload + + @classmethod + def _send_event(cls, payload: dict[str, str | int]) -> bool: + if not cls._is_enabled(): + return False + + endpoints = [dify_config.TELEMETRY_ENDPOINT] + if dify_config.TELEMETRY_FALLBACK_ENDPOINT not in endpoints: + endpoints.append(dify_config.TELEMETRY_FALLBACK_ENDPOINT) + + for endpoint in endpoints: + if not endpoint: + continue + + try: + response = httpx.post( + endpoint, + json=payload, + timeout=dify_config.TELEMETRY_TIMEOUT_SECONDS, + ) + response.raise_for_status() + return True + except httpx.RequestError: + logger.debug("Failed to send community telemetry event to %s", endpoint, exc_info=True) + except httpx.HTTPStatusError: + logger.debug("Community telemetry endpoint returned an error: %s", endpoint, exc_info=True) + return False + + return False + + @classmethod + def _is_heartbeat_due(cls, setup: DifySetup, now: datetime) -> bool: + if setup.instance_id is None: + return False + + if setup.last_heartbeat_at is not None and setup.last_heartbeat_at.date() >= now.date(): + return False + + return True + + @staticmethod + def _format_datetime(value: datetime) -> str: + return value.replace(microsecond=0).isoformat() + "Z" + + @staticmethod + def _normalize_os(value: str) -> str: + os_name = value.lower() + if os_name in {"linux", "darwin", "windows"}: + return os_name + return "unknown" + + @staticmethod + def _normalize_arch(value: str) -> str: + arch = value.lower() + if arch in {"x86_64", "amd64"}: + return "amd64" + if arch in {"aarch64", "arm64"}: + return "arm64" + if arch.startswith("arm"): + return "arm" + if arch in {"i386", "i686", "x86"}: + return "386" + return "unknown" diff --git a/api/services/tools/builtin_tools_manage_service.py b/api/services/tools/builtin_tools_manage_service.py index 45480f71d1a..dc82f91d3c1 100644 --- a/api/services/tools/builtin_tools_manage_service.py +++ b/api/services/tools/builtin_tools_manage_service.py @@ -31,6 +31,7 @@ from core.tools.utils.system_encryption import decrypt_system_params from extensions.ext_database import db from extensions.ext_redis import redis_client from models.account import Account +from models.enums import PermissionEnum from models.provider_ids import ToolProviderID from models.tools import BuiltinToolProvider, ToolOAuthSystemClient, ToolOAuthTenantClient from services.tools.tools_transform_service import ToolTransformService @@ -282,8 +283,6 @@ class BuiltinToolManageService: cache=NoOpProviderCredentialCache(), ) - from models.enums import PermissionEnum - visibility_enum = PermissionEnum(visibility) if visibility else PermissionEnum.ALL_TEAM # Plugin credentials only expose only_me / all_team_members at creation; # partial-member access is handled by workspace RBAC, not per-credential. diff --git a/api/services/tools/mcp_tools_manage_service.py b/api/services/tools/mcp_tools_manage_service.py index ae184ba4561..6a2d2fea819 100644 --- a/api/services/tools/mcp_tools_manage_service.py +++ b/api/services/tools/mcp_tools_manage_service.py @@ -259,7 +259,7 @@ class MCPToolManageService: mcp_provider.encrypted_credentials = self._process_credentials(authentication, mcp_provider, tenant_id) # Update user-identity forwarding mode. The controller has already - # resolved "leave unchanged" and applied the ENTERPRISE_ENABLED gate, + # resolved "leave unchanged" and applied the Enterprise-edition gate, # so this is always a concrete, vetted value. mcp_provider.identity_mode = identity_mode diff --git a/api/services/trigger/trigger_subscription_builder_service.py b/api/services/trigger/trigger_subscription_builder_service.py index cff735b39d3..0901e5ccf7f 100644 --- a/api/services/trigger/trigger_subscription_builder_service.py +++ b/api/services/trigger/trigger_subscription_builder_service.py @@ -73,102 +73,6 @@ class TriggerSubscriptionBuilderService: with redis_client.lock(lock_key, timeout=cls.__LOCK_EXPIRE_SECONDS__): yield - @classmethod - def verify_trigger_subscription_builder( - cls, - tenant_id: str, - user_id: str, - provider_id: TriggerProviderID, - subscription_builder_id: str, - ) -> Mapping[str, Any]: - """Verify a trigger subscription builder""" - provider_controller = TriggerManager.get_trigger_provider(tenant_id, provider_id) - if not provider_controller: - raise ValueError(f"Provider {provider_id} not found") - - subscription_builder = cls.get_subscription_builder(subscription_builder_id) - if not subscription_builder: - raise ValueError(f"Subscription builder {subscription_builder_id} not found") - - if subscription_builder.credential_type == CredentialType.OAUTH2: - return {"verified": bool(subscription_builder.credentials)} - - if subscription_builder.credential_type == CredentialType.API_KEY: - credentials_to_validate = subscription_builder.credentials - try: - provider_controller.validate_credentials(user_id, credentials_to_validate) - except ToolProviderCredentialValidationError as e: - raise ValueError(f"Invalid credentials: {e}") - return {"verified": True} - - return {"verified": True} - - @classmethod - def build_trigger_subscription_builder( - cls, tenant_id: str, user_id: str, provider_id: TriggerProviderID, subscription_builder_id: str - ) -> None: - """Build a trigger subscription builder""" - provider_controller = TriggerManager.get_trigger_provider(tenant_id, provider_id) - if not provider_controller: - raise ValueError(f"Provider {provider_id} not found") - - # Acquire lock to prevent concurrent build operations - with cls.acquire_builder_lock(subscription_builder_id): - subscription_builder = cls.get_subscription_builder(subscription_builder_id) - if not subscription_builder: - raise ValueError(f"Subscription builder {subscription_builder_id} not found") - - if not subscription_builder.name: - raise ValueError("Subscription builder name is required") - - credential_type = CredentialType.of(subscription_builder.credential_type or CredentialType.UNAUTHORIZED) - if credential_type == CredentialType.UNAUTHORIZED: - # manually create - TriggerProviderService.add_trigger_subscription( - subscription_id=subscription_builder.id, - tenant_id=tenant_id, - user_id=user_id, - name=subscription_builder.name, - provider_id=provider_id, - endpoint_id=subscription_builder.endpoint_id, - parameters=subscription_builder.parameters, - properties=subscription_builder.properties, - credential_expires_at=subscription_builder.credential_expires_at or -1, - expires_at=subscription_builder.expires_at, - credentials=subscription_builder.credentials, - credential_type=credential_type, - ) - else: - # automatically create - subscription: Subscription = TriggerManager.subscribe_trigger( - tenant_id=tenant_id, - user_id=user_id, - provider_id=provider_id, - endpoint=generate_plugin_trigger_endpoint_url(subscription_builder.endpoint_id), - parameters=subscription_builder.parameters, - credentials=subscription_builder.credentials, - credential_type=credential_type, - ) - - TriggerProviderService.add_trigger_subscription( - subscription_id=subscription_builder.id, - tenant_id=tenant_id, - user_id=user_id, - name=subscription_builder.name, - provider_id=provider_id, - endpoint_id=subscription_builder.endpoint_id, - parameters=subscription_builder.parameters, - properties=subscription.properties, - credentials=subscription_builder.credentials, - credential_type=credential_type, - credential_expires_at=subscription_builder.credential_expires_at or -1, - expires_at=subscription_builder.expires_at, - ) - - # Delete the builder after successful subscription creation - cache_key = cls.encode_cache_key(subscription_builder_id) - redis_client.delete(cache_key) - @classmethod def create_trigger_subscription_builder( cls, @@ -208,6 +112,7 @@ class TriggerSubscriptionBuilderService: def update_trigger_subscription_builder( cls, tenant_id: str, + user_id: str, provider_id: TriggerProviderID, subscription_builder_id: str, subscription_builder_updater: SubscriptionBuilderUpdater, @@ -223,9 +128,12 @@ class TriggerSubscriptionBuilderService: # Acquire lock to prevent concurrent updates with cls.acquire_builder_lock(subscription_id): cache_key = cls.encode_cache_key(subscription_id) - subscription_builder_cache = cls.get_subscription_builder(subscription_builder_id) - if not subscription_builder_cache or subscription_builder_cache.tenant_id != tenant_id: - raise ValueError(f"Subscription {subscription_id} expired or not found") + subscription_builder_cache = cls._require_owned_subscription_builder( + tenant_id=tenant_id, + user_id=user_id, + provider_id=provider_id, + subscription_builder_id=subscription_builder_id, + ) subscription_builder_updater.update(subscription_builder_cache) @@ -255,9 +163,12 @@ class TriggerSubscriptionBuilderService: # Acquire lock for the entire update + verify operation with cls.acquire_builder_lock(subscription_id): cache_key = cls.encode_cache_key(subscription_id) - subscription_builder_cache = cls.get_subscription_builder(subscription_builder_id) - if not subscription_builder_cache or subscription_builder_cache.tenant_id != tenant_id: - raise ValueError(f"Subscription {subscription_id} expired or not found") + subscription_builder_cache = cls._require_owned_subscription_builder( + tenant_id=tenant_id, + user_id=user_id, + provider_id=provider_id, + subscription_builder_id=subscription_builder_id, + ) # Update subscription_builder_updater.update(subscription_builder_cache) @@ -300,20 +211,16 @@ class TriggerSubscriptionBuilderService: # Acquire lock for the entire update + build operation with cls.acquire_builder_lock(subscription_id): cache_key = cls.encode_cache_key(subscription_id) - subscription_builder_cache = cls.get_subscription_builder(subscription_builder_id) - if not subscription_builder_cache or subscription_builder_cache.tenant_id != tenant_id: - raise ValueError(f"Subscription {subscription_id} expired or not found") - - # Update - subscription_builder_updater.update(subscription_builder_cache) - redis_client.setex( - cache_key, cls.__BUILDER_CACHE_EXPIRE_SECONDS__, subscription_builder_cache.model_dump_json() + subscription_builder = cls._require_owned_subscription_builder( + tenant_id=tenant_id, + user_id=user_id, + provider_id=provider_id, + subscription_builder_id=subscription_builder_id, ) - # Re-fetch to ensure we have the latest data - subscription_builder = cls.get_subscription_builder(subscription_builder_id) - if not subscription_builder: - raise ValueError(f"Subscription builder {subscription_builder_id} not found") + # Update + subscription_builder_updater.update(subscription_builder) + redis_client.setex(cache_key, cls.__BUILDER_CACHE_EXPIRE_SECONDS__, subscription_builder.model_dump_json()) if not subscription_builder.name: raise ValueError("Subscription builder name is required") @@ -364,7 +271,6 @@ class TriggerSubscriptionBuilderService: ) # Delete the builder after successful subscription creation - cache_key = cls.encode_cache_key(subscription_builder_id) redis_client.delete(cache_key) @classmethod @@ -389,17 +295,57 @@ class TriggerSubscriptionBuilderService: ) @classmethod - def get_subscription_builder(cls, endpoint_id: str) -> SubscriptionBuilder | None: - """ - Get a trigger subscription by the endpoint ID. - """ + def _get_subscription_builder_by_endpoint_id(cls, endpoint_id: str) -> SubscriptionBuilder | None: + """Resolve the public validation capability without authenticated owner context.""" cache_key = cls.encode_cache_key(endpoint_id) subscription_cache = redis_client.get(cache_key) if subscription_cache: - return SubscriptionBuilder.model_validate_json(subscription_cache) + subscription_builder = SubscriptionBuilder.model_validate_json(subscription_cache) + if subscription_builder.endpoint_id == endpoint_id: + return subscription_builder return None + @classmethod + def get_subscription_builder( + cls, + tenant_id: str, + user_id: str, + provider_id: TriggerProviderID, + subscription_builder_id: str, + ) -> SubscriptionBuilder | None: + """Return an owned temporary builder, or None when no temporary builder exists.""" + subscription_builder = cls._get_subscription_builder_by_endpoint_id(subscription_builder_id) + if subscription_builder is None: + return None + if ( + subscription_builder.id != subscription_builder_id + or subscription_builder.tenant_id != tenant_id + or subscription_builder.user_id != user_id + or subscription_builder.provider_id != str(provider_id) + ): + raise ValueError(f"Subscription builder {subscription_builder_id} not found") + return subscription_builder + + @classmethod + def _require_owned_subscription_builder( + cls, + tenant_id: str, + user_id: str, + provider_id: TriggerProviderID, + subscription_builder_id: str, + ) -> SubscriptionBuilder: + """Return an owned temporary builder or reject an absent capability.""" + subscription_builder = cls.get_subscription_builder( + tenant_id=tenant_id, + user_id=user_id, + provider_id=provider_id, + subscription_builder_id=subscription_builder_id, + ) + if subscription_builder is None: + raise ValueError(f"Subscription builder {subscription_builder_id} not found") + return subscription_builder + @classmethod def append_log(cls, endpoint_id: str, request: Request, response: Response) -> None: """Append validation request log to Redis.""" @@ -433,9 +379,22 @@ class TriggerSubscriptionBuilderService: ) @classmethod - def list_logs(cls, endpoint_id: str) -> list[RequestLog]: + def list_logs( + cls, + tenant_id: str, + user_id: str, + provider_id: TriggerProviderID, + subscription_builder_id: str, + ) -> list[RequestLog]: """List request logs for validation endpoint.""" - key = f"trigger:subscription:builder:logs:{endpoint_id}" + subscription_builder = cls._require_owned_subscription_builder( + tenant_id=tenant_id, + user_id=user_id, + provider_id=provider_id, + subscription_builder_id=subscription_builder_id, + ) + + key = f"trigger:subscription:builder:logs:{subscription_builder.endpoint_id}" logs_json = redis_client.get(key) if not logs_json: return [] @@ -451,7 +410,7 @@ class TriggerSubscriptionBuilderService: :return: The Flask response object """ # check if validation endpoint exists - subscription_builder: SubscriptionBuilder | None = cls.get_subscription_builder(endpoint_id) + subscription_builder: SubscriptionBuilder | None = cls._get_subscription_builder_by_endpoint_id(endpoint_id) if not subscription_builder: return None @@ -482,14 +441,21 @@ class TriggerSubscriptionBuilderService: return error_response @classmethod - def get_subscription_builder_by_id(cls, subscription_builder_id: str) -> SubscriptionBuilderApiEntity: + def get_subscription_builder_by_id( + cls, + tenant_id: str, + user_id: str, + provider_id: TriggerProviderID, + subscription_builder_id: str, + ) -> SubscriptionBuilderApiEntity: """Get a trigger subscription builder API entity.""" - subscription_builder = cls.get_subscription_builder(subscription_builder_id) - if not subscription_builder: - raise ValueError(f"Subscription builder {subscription_builder_id} not found") + subscription_builder = cls._require_owned_subscription_builder( + tenant_id=tenant_id, + user_id=user_id, + provider_id=provider_id, + subscription_builder_id=subscription_builder_id, + ) return cls.builder_to_api_entity( - controller=TriggerManager.get_trigger_provider( - subscription_builder.tenant_id, TriggerProviderID(subscription_builder.provider_id) - ), + controller=TriggerManager.get_trigger_provider(tenant_id, provider_id), entity=subscription_builder, ) diff --git a/api/services/trigger/webhook_service.py b/api/services/trigger/webhook_service.py index a921c91d64a..2e4f86c93c8 100644 --- a/api/services/trigger/webhook_service.py +++ b/api/services/trigger/webhook_service.py @@ -23,7 +23,7 @@ from core.workflow.nodes.trigger_webhook.entities import ( WebhookData, WebhookParameter, ) -from enums.quota_type import QuotaType +from enums import QuotaType from extensions.ext_database import db from extensions.ext_redis import redis_client from factories import file_factory diff --git a/api/services/turnstile_service.py b/api/services/turnstile_service.py new file mode 100644 index 00000000000..84634b2310a --- /dev/null +++ b/api/services/turnstile_service.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +import httpx +from pydantic import BaseModel, Field, SecretStr, ValidationError + +from configs import dify_config +from core.helper.http_client_pooling import get_pooled_http_client + +_SITEVERIFY_URL = "https://challenges.cloudflare.com/turnstile/v0/siteverify" +_EXPECTED_ACTION = "signin_code" +_MAX_TOKEN_LENGTH = 2048 +_CLIENT_ERROR_CODES = frozenset( + { + "bad-request", + "invalid-input-response", + "missing-input-response", + "timeout-or-duplicate", + } +) + +_http_client = get_pooled_http_client( + "cloudflare:turnstile", + lambda: httpx.Client( + timeout=httpx.Timeout(5.0, connect=3.0), + limits=httpx.Limits(max_keepalive_connections=20, max_connections=50), + ), +) + + +class TurnstileChallengeRejectedError(Exception): + """The submitted challenge is missing, invalid, expired, or not valid for this site.""" + + +class TurnstileUpstreamError(Exception): + """Turnstile could not be called or returned an unusable response.""" + + +class _TurnstileResponse(BaseModel): + success: bool + hostname: str | None = None + action: str | None = None + error_codes: list[str] = Field(default_factory=list, alias="error-codes") + + +class TurnstileService: + @classmethod + def verify(cls, *, token: str | None, remote_ip: str | None) -> None: + normalized_token = token.strip() if token else "" + if not normalized_token or len(normalized_token) > _MAX_TOKEN_LENGTH: + raise TurnstileChallengeRejectedError + + secret_key = dify_config.TURNSTILE_SECRET_KEY + allowed_hostnames = dify_config.TURNSTILE_ALLOWED_HOSTNAME_SET + if not isinstance(secret_key, SecretStr) or not allowed_hostnames: + raise TurnstileUpstreamError("Turnstile is not configured") + + payload = { + "secret": secret_key.get_secret_value(), + "response": normalized_token, + } + if remote_ip: + payload["remoteip"] = remote_ip + + try: + response = _http_client.post(_SITEVERIFY_URL, data=payload) + response.raise_for_status() + result = _TurnstileResponse.model_validate(response.json()) + except (httpx.HTTPError, ValidationError, ValueError) as exc: + raise TurnstileUpstreamError("Turnstile verification request failed") from exc + + if not result.success: + error_codes = frozenset(result.error_codes) + if error_codes and error_codes.issubset(_CLIENT_ERROR_CODES): + raise TurnstileChallengeRejectedError + raise TurnstileUpstreamError("Turnstile returned a server-side verification error") + + if result.action != _EXPECTED_ACTION or not cls._is_allowed_hostname(result.hostname, allowed_hostnames): + raise TurnstileChallengeRejectedError + + @staticmethod + def _is_allowed_hostname(hostname: str | None, allowed_hostnames: frozenset[str]) -> bool: + normalized_hostname = hostname.lower().strip(".") if hostname else "" + return any( + normalized_hostname == allowed or normalized_hostname.endswith(f".{allowed}") + for allowed in allowed_hostnames + ) diff --git a/api/services/variable_truncator.py b/api/services/variable_truncator.py index 00aa31650c2..78e917b052a 100644 --- a/api/services/variable_truncator.py +++ b/api/services/variable_truncator.py @@ -278,14 +278,14 @@ class VariableTruncator(BaseTruncator): target_length = self._array_element_limit for i, item in enumerate(value): - # Dirty fix: - # The output of `Start` node may contain list of `File` elements, - # causing `AssertionError` while invoking `_truncate_json_primitives`. - # - # This check ensures that `list[File]` are handled separately - if isinstance(item, File): - truncated_value.append(item) - continue + # ``File`` is routed through ``_truncate_json_primitives`` (whose + # dedicated ``File`` branch returns the file as-is with its real + # serialized size). That preserves the count cap + # (``array_element_limit``) and the byte budget (``target_size``) + # for ``list[File]`` — the original "Dirty fix" branch above this + # loop bypassed both guarantees and reported ``used_size=2`` even + # when the returned array serialized to well over the budget. + # See https://github.com/langgenius/dify/issues/39218. if i >= target_length: return _PartResult(truncated_value, used_size, True) if i > 0: @@ -295,7 +295,7 @@ class VariableTruncator(BaseTruncator): break remaining_budget = target_size - used_size - if item is None or isinstance(item, (str, list, dict, bool, int, float, UpdatedVariable)): + if item is None or isinstance(item, (str, list, dict, bool, int, float, File, UpdatedVariable)): part_result = self._truncate_json_primitives(item, remaining_budget) else: raise UnknownTypeError(f"got unknown type {type(item)} in array truncation") diff --git a/api/services/vector_space_admission_service.py b/api/services/vector_space_admission_service.py new file mode 100644 index 00000000000..8d3da92bb1c --- /dev/null +++ b/api/services/vector_space_admission_service.py @@ -0,0 +1,437 @@ +import json +import logging +import math +import re +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +from sqlalchemy.orm import Session + +from configs import dify_config +from core.model_manager import ModelManager +from core.rag.datasource.vdb.vector_factory import Vector +from core.rag.datasource.vdb.vector_type import VectorType +from core.rag.embedding.cached_embedding import CacheEmbedding +from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType +from core.rag.models.document import Document +from enums import CloudPlan, DeploymentEdition +from extensions.ext_redis import redis_client +from graphon.model_runtime.entities.model_entities import ModelType +from models.dataset import Dataset +from services.billing_service import BillingService + +logger = logging.getLogger(__name__) + +_MEBIBYTE = 1024 * 1024 +_FLOAT32_BYTES = 4 +_TIDB_VECTOR_COPIES = 2 +_TIDB_POINT_OVERHEAD_BYTES = 3584 +_WATERMARK_LOCK_TIMEOUT_SECONDS = 5 +_WATERMARK_TTL_SECONDS = 30 * 60 +_ERROR_PATTERN = re.compile( + r"Vector storage is estimated to reach (?P\d+) MB after this upload, " + r"exceeding the (?P\d+) MB limit of the current plan\." +) + +VECTOR_SPACE_ADMISSION_ERROR_CODE = "vector_space_estimate_exceeded" + + +class VectorSpaceAdmissionError(ValueError): + def __init__(self, message: str): + self.description = message + super().__init__(message) + + +@dataclass(frozen=True) +class VectorStorageWorkload: + text_points: int + summary_points: int + probe_text: str | None + + @property + def total_points(self) -> int: + return self.text_points + self.summary_points + + +@dataclass(frozen=True) +class VectorSpaceAdmissionErrorDetails: + estimated_mb: int + plan_limit_mb: int + + +def estimate_tidb_storage_bytes(point_count: int, dimension: int) -> int: + """Estimate TiDB row and columnar storage for vector points.""" + return point_count * (dimension * _FLOAT32_BYTES * _TIDB_VECTOR_COPIES + _TIDB_POINT_OVERHEAD_BYTES) + + +def parse_vector_space_estimate_limits(value: str) -> dict[CloudPlan, int]: + limits: dict[CloudPlan, int] = {} + for item in value.split(","): + plan_name, separator, raw_limit = item.strip().partition(":") + if not separator: + raise ValueError(f"Invalid vector-space estimate limit: {item!r}") + try: + plan = CloudPlan(plan_name) + limit = int(raw_limit) + except (TypeError, ValueError) as error: + raise ValueError(f"Invalid vector-space estimate limit: {item!r}") from error + if limit <= 0 or plan in limits: + raise ValueError(f"Invalid vector-space estimate limit: {item!r}") + limits[plan] = limit + if set(limits) != set(CloudPlan): + raise ValueError(f"Invalid vector-space estimate limits: {value!r}; include sandbox, professional, and team") + return limits + + +def format_vector_space_admission_error(estimated_mb: int, plan_limit_mb: int) -> str: + return ( + f"Vector storage is estimated to reach {estimated_mb} MB after this upload, " + f"exceeding the {plan_limit_mb} MB limit of the current plan." + ) + + +def get_vector_space_admission_error_details(error: str | None) -> VectorSpaceAdmissionErrorDetails | None: + if not error or not (match := _ERROR_PATTERN.fullmatch(error)): + return None + return VectorSpaceAdmissionErrorDetails( + estimated_mb=int(match.group("estimated")), + plan_limit_mb=int(match.group("limit")), + ) + + +def get_vector_space_admission_error_fields(error: str | None) -> dict[str, str | int | None]: + details = get_vector_space_admission_error_details(error) + return { + "error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE if details else None, + "estimated_vector_space_mb": details.estimated_mb if details else None, + "vector_space_limit_mb": details.plan_limit_mb if details else None, + } + + +def build_document_workload( + doc_form: str, + documents: list[Document], + *, + include_summaries: bool, +) -> VectorStorageWorkload: + # V1 estimates text vectors only; attachments are excluded. + texts: list[str] = [] + for document in documents: + if doc_form == IndexStructureType.PARENT_CHILD_INDEX: + texts.extend( + child.page_content + for child in document.children or [] + if child.page_content and child.page_content.strip() + ) + elif document.page_content and document.page_content.strip(): + texts.append(document.page_content) + + summary_points = 0 + if include_summaries and doc_form != IndexStructureType.QA_INDEX: + summary_points = sum(1 for document in documents if document.page_content and document.page_content.strip()) + + return VectorStorageWorkload( + text_points=len(texts), + summary_points=summary_points, + probe_text=texts[0] if texts else None, + ) + + +def build_pipeline_workload( + chunk_structure: str, + chunks: Any, + *, + include_summaries: bool, +) -> VectorStorageWorkload: + # V1 estimates chunk text only; file and image metadata are excluded. + texts: list[str] = [] + summary_points = 0 + + if chunk_structure == IndexStructureType.QA_INDEX: + for chunk in _items(chunks, "qa_chunks"): + question = _field(chunk, "question") + if isinstance(question, str) and question.strip(): + texts.append(question) + elif chunk_structure == IndexStructureType.PARENT_CHILD_INDEX: + for chunk in _items(chunks, "parent_child_chunks"): + parent_content = _field(chunk, "parent_content") + if include_summaries and isinstance(parent_content, str) and parent_content.strip(): + summary_points += 1 + for child in _field(chunk, "child_contents") or []: + if isinstance(child, str) and child.strip(): + texts.append(child) + else: + raw_chunks = chunks if isinstance(chunks, list) else _items(chunks, "general_chunks") + for chunk in raw_chunks: + content = chunk if isinstance(chunk, str) else _field(chunk, "content") + if isinstance(content, str) and content.strip(): + texts.append(content) + if include_summaries: + summary_points += 1 + + return VectorStorageWorkload( + text_points=len(texts), + summary_points=summary_points, + probe_text=texts[0] if texts else None, + ) + + +def _field(value: Any, name: str) -> Any: + if isinstance(value, Mapping): + return value.get(name) + return getattr(value, name, None) # guard-ignore: no-new-getattr -- supports validated chunk models + + +def _items(value: Any, name: str) -> list[Any]: + items = _field(value, name) + return list(items) if items else [] + + +class VectorSpaceAdmissionService: + """Cloud-only pre-write guard for unusually large TiDB vector workloads.""" + + def __init__(self) -> None: + self._dimension_by_dataset: dict[str, int] = {} + self._plan_by_tenant: dict[str, CloudPlan | None] = {} + + def ensure_document_can_be_indexed( + self, + *, + dataset: Dataset, + document_id: str, + doc_form: str, + documents: list[Document], + include_summaries: bool, + session: Session, + ) -> None: + self._ensure_can_write( + dataset=dataset, + document_id=document_id, + workload=build_document_workload( + doc_form, + documents, + include_summaries=include_summaries, + ), + session=session, + ) + + def ensure_pipeline_can_be_indexed( + self, + *, + dataset: Dataset, + document_id: str, + chunk_structure: str, + chunks: Any, + include_summaries: bool, + session: Session, + ) -> None: + self._ensure_can_write( + dataset=dataset, + document_id=document_id, + workload=build_pipeline_workload( + chunk_structure, + chunks, + include_summaries=include_summaries, + ), + session=session, + ) + + def _ensure_can_write( + self, + *, + dataset: Dataset, + document_id: str, + workload: VectorStorageWorkload, + session: Session, + ) -> None: + if ( + dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD + or dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY + or workload.total_points == 0 + or workload.probe_text is None + ): + return + if Vector.resolve_vector_type(dataset, session=session) != VectorType.TIDB_ON_QDRANT: + return + + plan = self._get_plan(dataset.tenant_id) + if plan is None: + return + estimate_limit_mb = parse_vector_space_estimate_limits( + dify_config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB + ).get(plan) + if estimate_limit_mb is None: + return + + current_usage_mb, plan_limit_mb = self._get_usage_and_limit_mb(dataset.tenant_id) + dimension = self._get_embedding_dimension(dataset, workload.probe_text) + estimate_bytes = math.ceil(estimate_tidb_storage_bytes(workload.total_points, dimension)) + document_estimated_mb = estimate_bytes / _MEBIBYTE + base_usage_bytes, projected_usage_bytes = self._reserve_projected_usage( + tenant_id=dataset.tenant_id, + document_id=document_id, + current_usage_bytes=math.ceil(current_usage_mb * _MEBIBYTE), + document_estimate_bytes=estimate_bytes, + estimate_limit_bytes=estimate_limit_mb * _MEBIBYTE, + ) + base_usage_mb = base_usage_bytes / _MEBIBYTE + projected_usage_mb = projected_usage_bytes / _MEBIBYTE + if projected_usage_bytes > estimate_limit_mb * _MEBIBYTE: + logger.warning( + "TiDB vector-space admission rejected tenant_id=%s document_id=%s plan=%s " + "points=%s dimension=%s current_usage_mb=%s document_estimated_mb=%s " + "watermark_base_usage_mb=%s projected_usage_mb=%s plan_limit_mb=%s estimate_limit_mb=%s", + dataset.tenant_id, + document_id, + plan, + workload.total_points, + dimension, + current_usage_mb, + document_estimated_mb, + base_usage_mb, + projected_usage_mb, + plan_limit_mb, + estimate_limit_mb, + ) + raise VectorSpaceAdmissionError( + format_vector_space_admission_error(math.ceil(projected_usage_mb), plan_limit_mb) + ) + + logger.info( + "TiDB vector-space admission allowed tenant_id=%s document_id=%s plan=%s " + "points=%s dimension=%s current_usage_mb=%s document_estimated_mb=%s " + "watermark_base_usage_mb=%s projected_usage_mb=%s estimate_limit_mb=%s", + dataset.tenant_id, + document_id, + plan, + workload.total_points, + dimension, + current_usage_mb, + document_estimated_mb, + base_usage_mb, + projected_usage_mb, + estimate_limit_mb, + ) + + def _get_usage_and_limit_mb(self, tenant_id: str) -> tuple[float, int]: + try: + vector_space = BillingService.get_vector_space(tenant_id) + current_usage_mb = float(vector_space["size"]) + plan_limit_mb = int(vector_space["limit"]) + except Exception as error: + raise VectorSpaceAdmissionError( + "Unable to verify vector storage usage right now. Please try again later." + ) from error + return current_usage_mb, plan_limit_mb + + def _reserve_projected_usage( + self, + *, + tenant_id: str, + document_id: str, + current_usage_bytes: int, + document_estimate_bytes: int, + estimate_limit_bytes: int, + ) -> tuple[int, int]: + watermark_key = f"tenant:{tenant_id}:vector_space_estimate_watermark" + lock_key = f"{watermark_key}:lock" + + try: + with redis_client.lock( + lock_key, + timeout=_WATERMARK_LOCK_TIMEOUT_SECONDS, + blocking_timeout=_WATERMARK_LOCK_TIMEOUT_SECONDS, + ): + raw_state = redis_client.get(watermark_key) + stored_usage_bytes = 0 + document_ids: set[str] = set() + if raw_state: + state = json.loads(raw_state) + stored_usage_bytes = state.get("projected_usage_bytes") + raw_document_ids = state.get("document_ids") + if ( + type(stored_usage_bytes) is not int + or stored_usage_bytes < 0 + or not isinstance(raw_document_ids, list) + or not all(isinstance(item, str) for item in raw_document_ids) + ): + raise ValueError("Invalid vector-space estimate watermark") + document_ids = set(raw_document_ids) + + base_usage_bytes = max(current_usage_bytes, stored_usage_bytes) + projected_usage_bytes = base_usage_bytes + if document_id not in document_ids: + projected_usage_bytes += document_estimate_bytes + + if projected_usage_bytes <= estimate_limit_bytes: + document_ids.add(document_id) + redis_client.setex( + watermark_key, + _WATERMARK_TTL_SECONDS, + json.dumps( + { + "projected_usage_bytes": projected_usage_bytes, + "document_ids": sorted(document_ids), + }, + separators=(",", ":"), + ), + ) + + return base_usage_bytes, projected_usage_bytes + except Exception as error: + raise VectorSpaceAdmissionError( + "Unable to reserve estimated vector storage right now. Please try again later." + ) from error + + def _get_plan(self, tenant_id: str) -> CloudPlan | None: + if tenant_id in self._plan_by_tenant: + return self._plan_by_tenant[tenant_id] + try: + billing_info = BillingService.get_info(tenant_id, exclude_vector_space=True) + except Exception as error: + raise VectorSpaceAdmissionError( + "Unable to verify the subscription plan right now. Please try again later." + ) from error + + plan = None + if billing_info["enabled"]: + try: + plan = CloudPlan(billing_info["subscription"]["plan"]) + except ValueError: + logger.warning( + "Skipping TiDB vector-space admission for unknown plan tenant_id=%s plan=%s", + tenant_id, + billing_info["subscription"]["plan"], + ) + self._plan_by_tenant[tenant_id] = plan + return plan + + def _get_embedding_dimension(self, dataset: Dataset, probe_text: str) -> int: + cached_dimension = self._dimension_by_dataset.get(dataset.id) + if cached_dimension is not None: + return cached_dimension + + model_manager = ModelManager.for_tenant(tenant_id=dataset.tenant_id) + if dataset.embedding_model_provider: + model_instance = model_manager.get_model_instance( + tenant_id=dataset.tenant_id, + provider=dataset.embedding_model_provider, + model_type=ModelType.TEXT_EMBEDDING, + model=dataset.embedding_model, + ) + else: + model_instance = model_manager.get_default_model_instance( + tenant_id=dataset.tenant_id, + model_type=ModelType.TEXT_EMBEDDING, + ) + + embeddings = CacheEmbedding(model_instance).embed_documents([probe_text]) + if not embeddings or not embeddings[0]: + raise VectorSpaceAdmissionError( + "Unable to estimate vector storage for this document. Please try again later." + ) + + dimension = len(embeddings[0]) + self._dimension_by_dataset[dataset.id] = dimension + return dimension diff --git a/api/services/webapp_auth_service.py b/api/services/webapp_auth_service.py index 33267c53d5c..3373a5ddf66 100644 --- a/api/services/webapp_auth_service.py +++ b/api/services/webapp_auth_service.py @@ -116,7 +116,7 @@ class WebAppAuthService: @classmethod def _get_account_jwt_token(cls, account: Account) -> str: - exp_dt = datetime.now(UTC) + timedelta(minutes=dify_config.ACCESS_TOKEN_EXPIRE_MINUTES * 24) + exp_dt = datetime.now(UTC) + timedelta(minutes=dify_config.ACCESS_TOKEN_EXPIRE_MINUTES) exp = int(exp_dt.timestamp()) payload = { diff --git a/api/services/website_service.py b/api/services/website_service.py index ea584088bbc..b833c146216 100644 --- a/api/services/website_service.py +++ b/api/services/website_service.py @@ -17,13 +17,22 @@ from extensions.ext_storage import storage from services.datasource_provider_service import DatasourceProviderService # Reuse pooled HTTP clients to avoid creating new connections per request and ease testing. +# Both clients carry a bounded read/connect timeout so a stalled Jina or +# adaptive-crawl endpoint fails fast instead of pinning a worker. The values +# match the floor used in the WaterCrawl PR (#37512). See #39859. _jina_http_client: httpx.Client = get_pooled_http_client( "website:jinareader", - lambda: httpx.Client(limits=httpx.Limits(max_keepalive_connections=50, max_connections=100)), + lambda: httpx.Client( + timeout=httpx.Timeout(30.0, connect=5.0), + limits=httpx.Limits(max_keepalive_connections=50, max_connections=100), + ), ) _adaptive_http_client: httpx.Client = get_pooled_http_client( "website:adaptivecrawl", - lambda: httpx.Client(limits=httpx.Limits(max_keepalive_connections=50, max_connections=100)), + lambda: httpx.Client( + timeout=httpx.Timeout(30.0, connect=5.0), + limits=httpx.Limits(max_keepalive_connections=50, max_connections=100), + ), ) diff --git a/api/services/workflow/queue_dispatcher.py b/api/services/workflow/queue_dispatcher.py index 0944b20357e..1d915a2cd49 100644 --- a/api/services/workflow/queue_dispatcher.py +++ b/api/services/workflow/queue_dispatcher.py @@ -10,6 +10,7 @@ with appropriate queue routing and priority assignment. from enum import StrEnum from configs import dify_config +from enums import DeploymentEdition from services.billing_service import BillingService @@ -92,7 +93,7 @@ class QueueDispatcherManager: Returns: Appropriate queue dispatcher instance """ - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: try: billing_info = BillingService.get_info(tenant_id) plan = billing_info.get("subscription", {}).get("plan", "sandbox") @@ -100,7 +101,7 @@ class QueueDispatcherManager: # If billing service fails, default to sandbox plan = "sandbox" else: - # If billing is disabled, use team tier as default + # Self-hosted editions use the team queue by default. plan = "team" dispatcher_class = cls.PLAN_DISPATCHER_MAP.get( diff --git a/api/services/workflow_event_snapshot_service.py b/api/services/workflow_event_snapshot_service.py index 1758c13a803..5fb86206ec9 100644 --- a/api/services/workflow_event_snapshot_service.py +++ b/api/services/workflow_event_snapshot_service.py @@ -43,6 +43,7 @@ from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus from graphon.runtime import GraphRuntimeState from graphon.runtime.graph_runtime_state_protocol import ReadOnlyVariablePool from graphon.workflow_type_encoder import WorkflowRuntimeTypeConverter +from libs.datetime_utils import to_utc_timestamp from models.human_input import HumanInputForm from models.model import AppMode, Message from models.workflow import WorkflowNodeExecutionTriggeredFrom, WorkflowRun @@ -451,7 +452,7 @@ def _build_human_input_required_events( ) with session_maker() as session: for form_id, expiration_time, form_definition in session.execute(stmt): - expiration_times_by_form_id[str(form_id)] = int(expiration_time.timestamp()) + expiration_times_by_form_id[str(form_id)] = to_utc_timestamp(expiration_time) try: definition_payload = json.loads(form_definition) if form_definition else {} except (TypeError, json.JSONDecodeError): @@ -564,7 +565,6 @@ def _build_pause_event( variable_pool: ReadOnlyVariablePool | None = None if resumption_context is not None: state = GraphRuntimeState.from_snapshot(resumption_context.serialized_graph_runtime_state) - paused_nodes = state.get_paused_nodes() outputs = dict(WorkflowRuntimeTypeConverter().to_json_encodable(state.outputs or {})) variable_pool = state.variable_pool @@ -572,6 +572,9 @@ def _build_pause_event( pause_entity.get_pause_reasons(), variable_pool=variable_pool, ) + paused_nodes = list( + dict.fromkeys(reason.node_id for reason in resolved_pause_reasons if isinstance(reason, HumanInputRequired)) + ) reasons = [reason.model_dump(mode="json") for reason in resolved_pause_reasons] human_input_form_ids = [ form_id @@ -594,7 +597,7 @@ def _build_pause_event( ) for row in session.execute(stmt): form_id, expiration_time, *_rest = row - expiration_times_by_form_id[str(form_id)] = int(expiration_time.timestamp()) + expiration_times_by_form_id[str(form_id)] = to_utc_timestamp(expiration_time) # Reconnect paths must preserve the same pause-reason contract as live streams; # otherwise clients see schema drift after resume. reasons = enrich_human_input_pause_reasons( diff --git a/api/services/workflow_service.py b/api/services/workflow_service.py index 0d73acfbca2..c821e7b7c80 100644 --- a/api/services/workflow_service.py +++ b/api/services/workflow_service.py @@ -23,6 +23,13 @@ from core.workflow.human_input_adapter import ( adapt_human_input_node_data_for_graph, parse_human_input_delivery_methods, ) +from core.workflow.llm_environment_variable import ( + LLMEnvironmentVariable, + parse_llm_model_selector, + resolve_llm_model_config, + should_resolve_llm_model_selector, + validate_llm_environment_model_references, +) from core.workflow.node_factory import ( LATEST_VERSION, get_node_type_classes_mapping, @@ -43,7 +50,7 @@ from core.workflow.system_variables import build_bootstrap_variables, build_syst from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool from core.workflow.workflow_entry import WorkflowEntry from enterprise.telemetry.draft_trace import enqueue_draft_node_execution_trace -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from events.app_event import app_draft_workflow_was_synced, app_published_workflow_was_updated from extensions.ext_database import db from extensions.ext_storage import storage @@ -62,7 +69,9 @@ from graphon.graph_events import GraphNodeEventBase, NodeRunFailedEvent, NodeRun from graphon.node_events import NodeRunResult from graphon.nodes import BuiltinNodeTypes from graphon.nodes.base.node import Node +from graphon.nodes.container_effects import ContainerAwaitRequest from graphon.nodes.http_request import HTTP_REQUEST_CONFIG_FILTER_KEY, build_http_request_config +from graphon.nodes.llm.entities import ModelConfig from graphon.nodes.start.entities import StartNodeData from graphon.runtime import VariablePool from graphon.variable_loader import load_into_variable_pool @@ -76,6 +85,7 @@ from models.model import App, AppMode from models.tools import WorkflowToolProvider from models.workflow import Workflow, WorkflowNodeExecutionModel, WorkflowNodeExecutionTriggeredFrom, WorkflowType from repositories.factory import DifyAPIRepositoryFactory +from services.agent.retirement_service import WorkflowAgentRetirementService from services.billing_service import BillingService from services.errors.app import ( IsDraftWorkflowError, @@ -83,6 +93,7 @@ from services.errors.app import ( WorkflowHashNotEqualError, WorkflowNotFoundError, ) +from tasks.collect_agent_resources_task import enqueue_agent_resource_collection @dataclass(frozen=True) @@ -131,6 +142,7 @@ HumanInputNode = _DebugHumanInputNode from services.human_input_service import HumanInputService from services.workflow.workflow_converter import WorkflowConverter from services.workflow_ref_service import WorkflowRef +from services.workflow_version_number_service import allocate_version_number from .errors.workflow_service import DraftWorkflowDeletionError, WorkflowInUseError from .human_input_delivery_test_service import ( @@ -146,6 +158,48 @@ from .workflow_restore import apply_published_workflow_snapshot_to_draft _file_access_controller = DatabaseFileAccessController() +def _merge_environment_variable_patch( + current_variables: Sequence[VariableBase], + environment_variable_upserts: Sequence[VariableBase], + deleted_environment_variable_ids: Sequence[str], +) -> list[VariableBase]: + """Merge a per-ID environment-variable patch while preserving untouched server values.""" + upserts_by_id: dict[str, VariableBase] = {} + for variable in environment_variable_upserts: + if not variable.id: + raise ValueError("Patched environment variables require an id.") + if variable.id in upserts_by_id: + raise ValueError(f"Duplicate patched environment variable id: {variable.id}") + upserts_by_id[variable.id] = variable + + deleted_ids = set(deleted_environment_variable_ids) + if len(deleted_ids) != len(deleted_environment_variable_ids): + raise ValueError("Deleted environment variable ids must be unique.") + if "" in deleted_ids: + raise ValueError("Deleted environment variable ids must not be empty.") + if conflicting_ids := deleted_ids.intersection(upserts_by_id): + conflicting_id = min(conflicting_ids) + raise ValueError(f"Environment variable cannot be upserted and deleted in the same patch: {conflicting_id}") + + existing_ids: set[str] = set() + merged_variables: list[VariableBase] = [] + for variable in current_variables: + variable_id = variable.id + if variable_id: + existing_ids.add(variable_id) + if variable_id in deleted_ids: + continue + merged_variables.append(upserts_by_id.get(variable_id, variable)) + + merged_variables.extend( + variable for variable_id, variable in upserts_by_id.items() if variable_id not in existing_ids + ) + names = [variable.name for variable in merged_variables] + if len(set(names)) != len(names): + raise ValueError("Environment variable names must be unique.") + return merged_variables + + class WorkflowService: """ Workflow Service @@ -213,6 +267,19 @@ class WorkflowService: # return draft workflow return workflow + def _get_draft_workflow_for_update(self, app_model: App, *, session: Session) -> Workflow | None: + """Return the app draft while holding its row lock for the caller's transaction.""" + return session.scalar( + select(Workflow) + .where( + Workflow.tenant_id == app_model.tenant_id, + Workflow.app_id == app_model.id, + Workflow.version == Workflow.VERSION_DRAFT, + ) + .limit(1) + .with_for_update() + ) + def get_published_workflow_by_id(self, app_model: App, workflow_id: str, *, session: Session) -> Workflow | None: """ fetch published workflow by workflow_id @@ -290,7 +357,16 @@ class WorkflowService: stmt = ( select(Workflow) .where(Workflow.app_id == app_model.id) - .order_by(Workflow.version.desc()) + # The draft leads the list; its `created_at` is the app's creation time, so it would + # otherwise sort last. Published versions then order by publish time: `version` is a + # stringified timestamp whose microseconds are omitted when zero, so ordering by it + # misplaces versions across second boundaries, and `version_number` is NULL for + # versions published before numbering was introduced. + .order_by( + (Workflow.version == Workflow.VERSION_DRAFT).desc(), + Workflow.created_at.desc(), + Workflow.id.desc(), + ) .limit(limit + 1) .offset((page - 1) * limit) ) @@ -320,6 +396,9 @@ class WorkflowService: environment_variables: Sequence[VariableBase], conversation_variables: Sequence[VariableBase], session: Session, + environment_variable_upserts: Sequence[VariableBase] | None = None, + deleted_environment_variable_ids: Sequence[str] = (), + preserve_environment_variables: bool = False, commit: bool = True, sync_agent_bindings: bool = True, graph_only: bool = False, @@ -331,10 +410,17 @@ class WorkflowService: portable package references can be materialized atomically after the draft workflow has received its target-workspace id. + Existing drafts are row-locked before the hash check. Collaborative + graph-only saves preserve independently persisted draft fields, while + an explicit per-ID environment patch is merged with the graph. + :raises WorkflowHashNotEqualError """ + if environment_variable_upserts is None and deleted_environment_variable_ids: + raise ValueError("Deleted environment variable ids require an environment variable patch.") + # fetch draft workflow by app_model - workflow = self.get_draft_workflow(app_model=app_model, session=session) + workflow = self._get_draft_workflow_for_update(app_model=app_model, session=session) if workflow and workflow.unique_hash != unique_hash: raise WorkflowHashNotEqualError() @@ -349,6 +435,15 @@ class WorkflowService: # create draft workflow if not found if not workflow: + initial_environment_variables = ( + _merge_environment_variable_patch( + environment_variables, + environment_variable_upserts, + deleted_environment_variable_ids, + ) + if environment_variable_upserts is not None + else list(environment_variables) + ) workflow = Workflow( tenant_id=app_model.tenant_id, app_id=app_model.id, @@ -357,7 +452,7 @@ class WorkflowService: graph=json.dumps(graph), features=json.dumps(features), created_by=account.id, - environment_variables=environment_variables, + environment_variables=initial_environment_variables, conversation_variables=conversation_variables, ) session.add(workflow) @@ -368,14 +463,22 @@ class WorkflowService: workflow.updated_at = naive_utc_now() if not graph_only: workflow.features = json.dumps(features) - workflow.environment_variables = environment_variables workflow.conversation_variables = conversation_variables + if environment_variable_upserts is not None: + workflow.environment_variables = _merge_environment_variable_patch( + workflow.environment_variables, + environment_variable_upserts, + deleted_environment_variable_ids, + ) + elif not graph_only and not preserve_environment_variables: + workflow.environment_variables = environment_variables from services.agent.workflow_publish_service import WorkflowAgentPublishService session.flush() + retirement_candidates: set[str] = set() if sync_agent_bindings: - WorkflowAgentPublishService.sync_agent_bindings_for_draft( + retirement_candidates = WorkflowAgentPublishService.sync_agent_bindings_for_draft( session=session, draft_workflow=workflow, account_id=account.id, @@ -388,6 +491,16 @@ class WorkflowService: # commit db session changes if commit: session.commit() + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=app_model.tenant_id, + agent_ids=retirement_candidates, + account_id=account.id, + ) + enqueue_agent_resource_collection( + tenant_id=app_model.tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) # trigger app workflow events if commit: @@ -403,10 +516,8 @@ class WorkflowService: environment_variables: Sequence[VariableBase], account: Account, session: Session, - ): - """ - Update draft workflow environment variables - """ + ) -> None: + """Replace every environment variable on a draft workflow and commit the transaction.""" # fetch draft workflow by app_model workflow = self.get_draft_workflow(app_model=app_model, session=session) @@ -420,6 +531,34 @@ class WorkflowService: # commit db session changes session.commit() + def patch_draft_workflow_environment_variables( + self, + *, + app_model: App, + environment_variables: Sequence[VariableBase], + deleted_environment_variable_ids: Sequence[str], + account: Account, + session: Session, + ) -> None: + """Atomically merge per-ID environment-variable upserts and deletions into a draft workflow. + + The draft row is locked before reading its current variables so concurrent patches preserve + variables they do not touch. Existing variables keep their order and new variables are appended. + The transaction is committed before this method returns. + """ + workflow = self._get_draft_workflow_for_update(app_model=app_model, session=session) + if not workflow: + raise ValueError("No draft workflow found.") + + workflow.environment_variables = _merge_environment_variable_patch( + workflow.environment_variables, + environment_variables, + deleted_environment_variable_ids, + ) + workflow.updated_by = account.id + workflow.updated_at = naive_utc_now() + session.commit() + def update_draft_workflow_conversation_variables( self, *, @@ -509,7 +648,7 @@ class WorkflowService: from services.agent.workflow_publish_service import WorkflowAgentPublishService session.flush() - WorkflowAgentPublishService.restore_agent_node_bindings_to_draft( + retirement_candidates = WorkflowAgentPublishService.restore_agent_node_bindings_to_draft( session=session, source_workflow=source_workflow, draft_workflow=draft_workflow, @@ -517,6 +656,16 @@ class WorkflowService: ) session.commit() + binding_ids, home_snapshot_ids = WorkflowAgentRetirementService.retire_unowned( + tenant_id=app_model.tenant_id, + agent_ids=retirement_candidates, + account_id=account.id, + ) + enqueue_agent_resource_collection( + tenant_id=app_model.tenant_id, + binding_ids=binding_ids, + home_snapshot_ids=home_snapshot_ids, + ) app_draft_workflow_was_synced.send(app_model, synced_draft_workflow=draft_workflow) return draft_workflow @@ -529,7 +678,7 @@ class WorkflowService: account: Account, marked_name: str = "", marked_comment: str = "", - ) -> Workflow: + ) -> tuple[Workflow, set[str]]: draft_workflow_stmt = select(Workflow).where( Workflow.tenant_id == app_model.tenant_id, Workflow.app_id == app_model.id, @@ -539,6 +688,11 @@ class WorkflowService: if not draft_workflow: raise ValueError("No valid workflow found.") + validate_llm_environment_model_references( + graph=draft_workflow.graph_dict, + environment_variables=draft_workflow.environment_variables, + ) + # Validate credentials before publishing, for credential policy check from services.feature_service import FeatureService @@ -556,7 +710,7 @@ class WorkflowService: ) # billing check - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: limit_info = BillingService.get_info(app_model.tenant_id) if limit_info["subscription"]["plan"] == CloudPlan.SANDBOX: # Check trigger node count limit for SANDBOX plan @@ -576,6 +730,7 @@ class WorkflowService: app_id=app_model.id, type=draft_workflow.type, version=Workflow.version_from_datetime(naive_utc_now()), + version_number=allocate_version_number(session=session, app_id=app_model.id), graph=draft_workflow.graph, created_by=account.id, environment_variables=draft_workflow.environment_variables, @@ -588,7 +743,7 @@ class WorkflowService: # commit db session changes session.add(workflow) - WorkflowAgentPublishService.copy_agent_node_bindings_to_published( + retirement_candidates = WorkflowAgentPublishService.copy_agent_node_bindings_to_published( session=session, draft_workflow=draft_workflow, published_workflow=workflow, @@ -598,7 +753,7 @@ class WorkflowService: app_published_workflow_was_updated.send(app_model, published_workflow=workflow) # return new workflow - return workflow + return workflow, retirement_candidates def _validate_workflow_credentials(self, workflow: Workflow, *, session: Session) -> None: """ @@ -609,6 +764,14 @@ class WorkflowService: """ graph_dict = workflow.graph_dict nodes = graph_dict.get("nodes", []) + has_llm_model_reference = any( + node.get("data", {}).get("type") == "llm" + and should_resolve_llm_model_selector(node.get("data", {}).get("model_selector")) + for node in nodes + ) + environment_variables = ( + {variable.name: variable for variable in workflow.environment_variables} if has_llm_model_reference else {} + ) for node in nodes: node_data = node.get("data", {}) @@ -664,7 +827,22 @@ class WorkflowService: self._check_default_tool_credential(workflow.tenant_id, provider, session=session) elif node_type in ["llm", "knowledge_retrieval", "parameter_extractor", "question_classifier"]: + validation_node_data = node_data model_config = node_data.get("model", {}) + if node_type == "llm" and should_resolve_llm_model_selector(node_data.get("model_selector")): + selector = parse_llm_model_selector(node_data["model_selector"]) + variable = environment_variables.get(selector[1]) + if not isinstance(variable, LLMEnvironmentVariable): + raise ValueError( + f"LLM environment variable '{selector[1]}' was not found or is not an LLM variable" + ) + resolved_model = resolve_llm_model_config( + node_model=ModelConfig.model_validate(model_config), + variable_name=selector[1], + variable_value=variable.value, + ) + model_config = resolved_model.model_dump(mode="json") + validation_node_data = {**node_data, "model": model_config} provider = model_config.get("provider") model_name = model_config.get("name") @@ -672,7 +850,9 @@ class WorkflowService: # Validate that the provider+model combination can fetch valid credentials self._validate_llm_model_config(workflow.tenant_id, provider, model_name) # Validate load balancing credentials if load balancing is enabled - self._validate_load_balancing_credentials(workflow, node_data, node_id, session=session) + self._validate_load_balancing_credentials( + workflow, validation_node_data, node_id, session=session + ) else: raise ValueError(f"Node {node_id} ({node_type}): Missing provider or model configuration") @@ -1449,7 +1629,10 @@ class WorkflowService: def _handle_single_step_result( self, - invoke_node_fn: Callable[[], tuple[Node, Generator[GraphNodeEventBase, None, None]]], + invoke_node_fn: Callable[ + [], + tuple[Node, Generator[GraphNodeEventBase | ContainerAwaitRequest, None, None]], + ], start_at: float, node_id: str, ) -> WorkflowNodeExecution: @@ -1485,7 +1668,11 @@ class WorkflowService: return node_execution def _execute_node_safely( - self, invoke_node_fn: Callable[[], tuple[Node, Generator[GraphNodeEventBase, None, None]]] + self, + invoke_node_fn: Callable[ + [], + tuple[Node, Generator[GraphNodeEventBase | ContainerAwaitRequest, None, None]], + ], ) -> tuple[Node, NodeRunResult | None, bool, str | None]: """ Execute node safely and handle errors according to error strategy. diff --git a/api/services/workflow_version_number_service.py b/api/services/workflow_version_number_service.py new file mode 100644 index 00000000000..506702ab58d --- /dev/null +++ b/api/services/workflow_version_number_service.py @@ -0,0 +1,59 @@ +"""Allocation of user-facing workflow version numbers (`#N`). + +Numbers are unique and monotonically increasing within an app, and are never +reused: `workflow_version_counters` keeps the highest number handed out so far, +so deleting a published version does not free its number. + +The counter is keyed by `Workflow.app_id`, which is polymorphic — it holds an app +id, a pipeline id or a snippet id depending on the workflow kind. UUID uniqueness +across those tables is why the same counter table serves all of them. +""" + +from sqlalchemy import select +from sqlalchemy.dialects.mysql import insert as mysql_insert +from sqlalchemy.dialects.postgresql import insert as pg_insert +from sqlalchemy.orm import Session + +from configs import dify_config +from models.workflow import WorkflowVersionCounter + + +def allocate_version_number(*, session: Session, app_id: str) -> int: + """Reserve and return the next version number for `app_id`. + + The upsert acquires a row lock that is held until the caller's transaction + commits, so concurrent publishes of the same app serialize and never receive + the same number. Callers must run inside a transaction; if it rolls back, + both the counter update and workflow creation roll back, leaving the number + available for the next successful publish. + """ + # Dialect-specific upsert, mirroring `workflow_draft_variable_service`: the + # ORM cannot express "insert or increment" and a read-then-write would race. + # PostgreSQL returns the new value inline; MySQL has no RETURNING, so the + # value is read back within the same transaction while the row is still + # locked by the upsert. + if dify_config.SQLALCHEMY_DATABASE_URI_SCHEME == "postgresql": + stmt = ( + pg_insert(WorkflowVersionCounter) + .values(app_id=app_id, last_version_number=1) + .on_conflict_do_update( + index_elements=[WorkflowVersionCounter.app_id], + set_={"last_version_number": WorkflowVersionCounter.last_version_number + 1}, + ) + .returning(WorkflowVersionCounter.last_version_number) + ) + version_number = session.scalar(stmt) + else: + insert_stmt = mysql_insert(WorkflowVersionCounter).values(app_id=app_id, last_version_number=1) + session.execute( + insert_stmt.on_duplicate_key_update( # type: ignore[attr-defined] + last_version_number=WorkflowVersionCounter.last_version_number + 1, + ) + ) + version_number = session.scalar( + select(WorkflowVersionCounter.last_version_number).where(WorkflowVersionCounter.app_id == app_id) + ) + + if version_number is None: + raise ValueError(f"Failed to allocate a workflow version number for app {app_id}.") + return version_number diff --git a/api/services/workspace_member_query_service.py b/api/services/workspace_member_query_service.py new file mode 100644 index 00000000000..76ccb3d25a7 --- /dev/null +++ b/api/services/workspace_member_query_service.py @@ -0,0 +1,96 @@ +"""Application service for listing members of the active Console workspace.""" + +from collections.abc import Mapping, Sequence +from datetime import datetime +from typing import NamedTuple, Protocol + +from machinery.context import RequestContext + + +class WorkspaceMemberRole(NamedTuple): + id: str + name: str + + +class WorkspaceMemberRecord(NamedTuple): + id: str + name: str + email: str + avatar: str | None + last_login_at: datetime | None + last_active_at: datetime + created_at: datetime + status: str + legacy_role: str + + +class WorkspaceMemberQuery(Protocol): + def list_for_workspace(self, workspace_id: str) -> Sequence[WorkspaceMemberRecord]: ... + + +class WorkspaceMemberRoleSubject(NamedTuple): + account_id: str + legacy_role: str + + +class WorkspaceMemberRoleResolver(Protocol): + def resolve_many( + self, + workspace_id: str, + actor_account_id: str, + subjects: Sequence[WorkspaceMemberRoleSubject], + ) -> Mapping[str, Sequence[WorkspaceMemberRole]]: ... + + +class WorkspaceMemberSummary(NamedTuple): + id: str + name: str + email: str + avatar: str | None + last_login_at: datetime | None + last_active_at: datetime + created_at: datetime + role: str + roles: tuple[WorkspaceMemberRole, ...] + status: str + + +class WorkspaceMemberQueryService: + def __init__( + self, + *, + members: WorkspaceMemberQuery, + roles: WorkspaceMemberRoleResolver, + ) -> None: + self._members = members + self._roles = roles + + def list_current(self, context: RequestContext) -> tuple[WorkspaceMemberSummary, ...]: + workspace_id = context.active_workspace_id + if workspace_id is None: + raise RuntimeError("Console account admission did not resolve an active workspace") + + records = tuple(self._members.list_for_workspace(workspace_id)) + role_subjects = tuple( + WorkspaceMemberRoleSubject(account_id=record.id, legacy_role=record.legacy_role) for record in records + ) + + # The repository closes its read Session before role resolution + # performs enterprise I/O. + roles_by_member = self._roles.resolve_many(workspace_id, context.account_id, role_subjects) + + return tuple( + WorkspaceMemberSummary( + id=record.id, + name=record.name, + email=record.email, + avatar=record.avatar, + last_login_at=record.last_login_at, + last_active_at=record.last_active_at, + created_at=record.created_at, + role=record.legacy_role, + roles=tuple(roles_by_member.get(record.id, ())), + status=record.status, + ) + for record in records + ) diff --git a/api/services/workspace_member_role_resolver.py b/api/services/workspace_member_role_resolver.py new file mode 100644 index 00000000000..d2a4ceacec2 --- /dev/null +++ b/api/services/workspace_member_role_resolver.py @@ -0,0 +1,43 @@ +"""Deployment-compatible role resolution for workspace-member queries.""" + +from collections.abc import Mapping, Sequence +from typing import override + +from configs import dify_config +from services.enterprise import rbac_service as enterprise_rbac_service +from services.workspace_member_query_service import ( + WorkspaceMemberRole, + WorkspaceMemberRoleResolver, + WorkspaceMemberRoleSubject, +) + + +class DeploymentWorkspaceMemberRoleResolver(WorkspaceMemberRoleResolver): + """Preserve deployment-specific legacy and enterprise role behavior.""" + + @override + def resolve_many( + self, + workspace_id: str, + actor_account_id: str, + subjects: Sequence[WorkspaceMemberRoleSubject], + ) -> Mapping[str, Sequence[WorkspaceMemberRole]]: + role_subjects = tuple(subjects) + if not role_subjects: + return {} + + if not dify_config.RBAC_ENABLED: + return { + subject.account_id: (WorkspaceMemberRole(id=subject.legacy_role, name=subject.legacy_role),) + for subject in role_subjects + } + + member_roles = enterprise_rbac_service.RBACService.MemberRoles.batch_get( + workspace_id, + actor_account_id, + [subject.account_id for subject in role_subjects], + ) + return { + item.account_id: tuple(WorkspaceMemberRole(id=role.id, name=role.name) for role in item.roles) + for item in member_roles + } diff --git a/api/services/workspace_plan_gateway.py b/api/services/workspace_plan_gateway.py new file mode 100644 index 00000000000..22792b3e38a --- /dev/null +++ b/api/services/workspace_plan_gateway.py @@ -0,0 +1,44 @@ +"""Deployment-aware plan gateway for workspace queries.""" + +import logging +from collections.abc import Mapping, Sequence +from typing import override + +from configs import dify_config +from enums import CloudPlan, DeploymentEdition +from services.billing_service import BillingService +from services.feature_service import FeatureService +from services.workspace_query_service import WorkspacePlanGateway + +logger = logging.getLogger(__name__) + + +class DeploymentWorkspacePlanGateway(WorkspacePlanGateway): + """Resolve workspace plans using deployment-specific Billing and Feature sources.""" + + @override + def resolve_many(self, workspace_ids: Sequence[str]) -> Mapping[str, str]: + ids = tuple(workspace_ids) + if not ids: + return {} + + is_enterprise_only = dify_config.DEPLOYMENT_EDITION == DeploymentEdition.ENTERPRISE + if is_enterprise_only: + return dict.fromkeys(ids, str(CloudPlan.SANDBOX)) + + is_saas = dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD + bulk_plans = BillingService.get_plan_bulk(ids) if is_saas else {} + if is_saas and not bulk_plans: + logger.warning("get_plan_bulk returned empty result, falling back to FeatureService") + + resolved: dict[str, str] = {} + for workspace_id in ids: + tenant_plan = bulk_plans.get(workspace_id) + if tenant_plan: + resolved[workspace_id] = tenant_plan["plan"] or CloudPlan.SANDBOX + continue + + features = FeatureService.get_features(workspace_id, exclude_vector_space=True) + resolved[workspace_id] = features.billing.subscription.plan or CloudPlan.SANDBOX + + return resolved diff --git a/api/services/workspace_query_service.py b/api/services/workspace_query_service.py new file mode 100644 index 00000000000..746454da3c5 --- /dev/null +++ b/api/services/workspace_query_service.py @@ -0,0 +1,65 @@ +"""Application service for listing workspaces visible to a Console account.""" + +from collections.abc import Mapping, Sequence +from datetime import datetime +from typing import NamedTuple, Protocol + +from enums import CloudPlan +from machinery.context import RequestContext + + +class WorkspacePlanGateway(Protocol): + def resolve_many(self, workspace_ids: Sequence[str]) -> Mapping[str, str]: ... + + +class WorkspaceRecord(NamedTuple): + id: str + name: str | None + status: str + created_at: datetime + last_opened_at: datetime | None + + +class WorkspaceQuery(Protocol): + def list_for_account(self, account_id: str) -> Sequence[WorkspaceRecord]: ... + + +class WorkspaceSummary(NamedTuple): + id: str + name: str | None + plan: str + status: str + created_at: datetime + last_opened_at: datetime | None + current: bool + + +class WorkspaceQueryService: + def __init__( + self, + *, + workspaces: WorkspaceQuery, + plans: WorkspacePlanGateway, + ) -> None: + self._workspaces = workspaces + self._plans = plans + + def list_for_account(self, context: RequestContext) -> tuple[WorkspaceSummary, ...]: + records = tuple(self._workspaces.list_for_account(context.account_id)) + + # The repository closes its read Session before plan resolution + # performs Billing/Feature I/O. + plans = self._plans.resolve_many([record.id for record in records]) + + return tuple( + WorkspaceSummary( + id=record.id, + name=record.name, + plan=plans.get(record.id, CloudPlan.SANDBOX), + status=record.status, + created_at=record.created_at, + last_opened_at=record.last_opened_at, + current=record.id == context.active_workspace_id, + ) + for record in records + ) diff --git a/api/services/workspace_service.py b/api/services/workspace_service.py index 5f9003bb755..855dcdd1131 100644 --- a/api/services/workspace_service.py +++ b/api/services/workspace_service.py @@ -1,15 +1,45 @@ +from dataclasses import dataclass +from typing import Literal + from flask_login import current_user from sqlalchemy import select from sqlalchemy.orm import Session from configs import dify_config -from enums.cloud_plan import CloudPlan -from enums.deployment_edition import DeploymentEdition +from enums import CloudPlan, DeploymentEdition from models.account import Tenant, TenantAccountJoin, TenantAccountRole from services.account_service import TenantService +from services.billing_service import BillingService from services.feature_service import FeatureService +@dataclass(frozen=True) +class EffectiveCreditPool: + plan: CloudPlan | None = None + pool_type: Literal["paid", "trial"] | None = None + quota_limit: int | None = None + quota_used: int | None = None + exhausted_at: int | None = None + next_credit_reset_date: int | None = None + + @property + def remaining_credits(self) -> int | None: + if self.quota_limit is None or self.quota_used is None: + return None + if self.is_unlimited: + return -1 + return max(0, self.quota_limit - self.quota_used) + + @property + def is_unlimited(self) -> bool: + return self.quota_limit == -1 + + @property + def is_exhausted(self) -> bool: + remaining_credits = self.remaining_credits + return not self.is_unlimited and (remaining_credits is None or remaining_credits <= 0) + + def _set_credit_pool_info( tenant_info: dict[str, object], *, quota_limit: int, quota_used: int, exhausted_at: int | None = None ) -> None: @@ -20,6 +50,70 @@ def _set_credit_pool_info( class WorkspaceService: + @classmethod + def get_effective_credit_pool(cls, tenant_id: str, *, session: Session) -> EffectiveCreditPool: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: + return EffectiveCreditPool() + + billing_info = BillingService.get_info(tenant_id, exclude_vector_space=True) + subscription_plan = CloudPlan(billing_info["subscription"]["plan"]) + + from services.credit_pool_service import CreditPoolBalance, CreditPoolService + + effective_pool = None + effective_pool_type: Literal["paid", "trial"] = "trial" + if subscription_plan != CloudPlan.SANDBOX: + paid_pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type="paid", session=session) + if paid_pool is not None and (paid_pool.quota_limit == -1 or paid_pool.quota_limit > paid_pool.quota_used): + effective_pool = paid_pool + effective_pool_type = "paid" + + if effective_pool is None: + effective_pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type="trial", session=session) + + if effective_pool is None: + return EffectiveCreditPool( + plan=subscription_plan if billing_info["enabled"] else None, + next_credit_reset_date=billing_info.get("next_credit_reset_date"), + ) + + exhausted_at = effective_pool.exhausted_at if isinstance(effective_pool, CreditPoolBalance) else None + if not ( + isinstance(exhausted_at, int) + and exhausted_at > 0 + and effective_pool.quota_limit > 0 + and effective_pool.quota_used >= effective_pool.quota_limit + ): + exhausted_at = None + + return EffectiveCreditPool( + plan=subscription_plan if billing_info["enabled"] else None, + pool_type=effective_pool_type, + quota_limit=effective_pool.quota_limit, + quota_used=effective_pool.quota_used, + exhausted_at=exhausted_at, + next_credit_reset_date=billing_info.get("next_credit_reset_date"), + ) + + @classmethod + def get_current_workspace_summary(cls, tenant: Tenant, account_id: str, *, session: Session) -> dict[str, object]: + tenant_account_join = session.scalar( + select(TenantAccountJoin) + .where(TenantAccountJoin.tenant_id == tenant.id, TenantAccountJoin.account_id == account_id) + .limit(1) + ) + assert tenant_account_join is not None, "TenantAccountJoin not found" + + effective_pool = cls.get_effective_credit_pool(tenant.id, session=session) + + return { + "id": tenant.id, + "name": tenant.name, + "role": tenant_account_join.role, + "plan": effective_pool.plan, + "credits": effective_pool.remaining_credits, + } + @classmethod def get_tenant_info(cls, tenant: Tenant, session: Session): if not tenant: @@ -27,7 +121,6 @@ class WorkspaceService: tenant_info: dict[str, object] = { "id": tenant.id, "name": tenant.name, - "plan": tenant.plan, "status": tenant.status, "created_at": tenant.created_at, "trial_end_reason": None, @@ -44,6 +137,7 @@ class WorkspaceService: tenant_info["role"] = tenant_account_join.role feature = FeatureService.get_features(tenant.id, exclude_vector_space=True) + tenant_info["plan"] = feature.billing.subscription.plan if feature.billing.enabled else None can_replace_logo = feature.can_replace_logo if can_replace_logo and TenantService.has_roles( diff --git a/api/tasks/agent_backend_session_cleanup_task.py b/api/tasks/agent_backend_session_cleanup_task.py deleted file mode 100644 index f1316266db7..00000000000 --- a/api/tasks/agent_backend_session_cleanup_task.py +++ /dev/null @@ -1,71 +0,0 @@ -"""Celery tasks that execute Agent backend lifecycle-only session cleanup.""" - -from __future__ import annotations - -import logging - -from celery import shared_task - -from clients.agent_backend.factory import create_agent_backend_run_client -from clients.agent_backend.request_builder import AgentBackendRunRequestBuilder -from clients.agent_backend.session_cleanup import ( - AgentBackendSessionCleanupPayload, - cleanup_agent_backend_session, -) -from configs import dify_config - -logger = logging.getLogger(__name__) - - -def _create_agent_backend_client(): - if not (dify_config.AGENT_BACKEND_USE_FAKE or dify_config.AGENT_BACKEND_BASE_URL): - return None - return create_agent_backend_run_client( - base_url=dify_config.AGENT_BACKEND_BASE_URL, - use_fake=dify_config.AGENT_BACKEND_USE_FAKE, - fake_scenario=dify_config.AGENT_BACKEND_FAKE_SCENARIO, - stream_read_timeout_seconds=dify_config.AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS, - stream_max_reconnects=dify_config.AGENT_BACKEND_STREAM_MAX_RECONNECTS, - stream_run_timeout_seconds=dify_config.AGENT_BACKEND_RUN_TIMEOUT_SECONDS, - ) - - -def _run_cleanup_task(payload_dict: dict[str, object]) -> None: - payload = AgentBackendSessionCleanupPayload.model_validate(payload_dict) - result = cleanup_agent_backend_session( - payload=payload, - client=_create_agent_backend_client(), - request_builder=AgentBackendRunRequestBuilder(), - ) - if result.status == "succeeded": - return - - log_fields = { - "tenant_id": payload.metadata.get("tenant_id"), - "app_id": payload.metadata.get("app_id"), - "workflow_run_id": payload.metadata.get("workflow_run_id"), - "node_id": payload.metadata.get("node_id"), - "conversation_id": payload.metadata.get("conversation_id"), - "agent_id": payload.metadata.get("agent_id"), - "previous_agent_backend_run_id": payload.metadata.get("previous_agent_backend_run_id"), - "failed_agent_backend_run_id": payload.metadata.get("failed_agent_backend_run_id"), - "cleanup_run_id": result.cleanup_run_id, - "reason": result.reason, - } - if result.status == "skipped": - logger.info("Agent backend session cleanup skipped: %s", log_fields) - return - - logger.warning("Agent backend session cleanup failed: %s", log_fields) - - -@shared_task(queue="workflow_storage") -def cleanup_workflow_agent_runtime_session(payload_dict: dict[str, object]) -> None: - """Run one workflow-owned Agent backend cleanup payload.""" - _run_cleanup_task(payload_dict) - - -@shared_task(queue="conversation") -def cleanup_conversation_agent_runtime_session(payload_dict: dict[str, object]) -> None: - """Run one conversation-owned Agent backend cleanup payload.""" - _run_cleanup_task(payload_dict) diff --git a/api/tasks/app_generate/resume_agent_app_task.py b/api/tasks/app_generate/resume_agent_app_task.py index 0b884b014bc..eac6d273425 100644 --- a/api/tasks/app_generate/resume_agent_app_task.py +++ b/api/tasks/app_generate/resume_agent_app_task.py @@ -51,6 +51,7 @@ def resume_agent_app_execution(*, conversation_id: str, form_id: str) -> None: app_model=app_model, user=user, conversation_id=conversation_id, + form_id=form_id, invoke_from=_resolve_invoke_from(conversation), session=db.session(), ) diff --git a/api/tasks/app_generate/workflow_execute_task.py b/api/tasks/app_generate/workflow_execute_task.py index d76066a8aa7..8a88ff4dfa6 100644 --- a/api/tasks/app_generate/workflow_execute_task.py +++ b/api/tasks/app_generate/workflow_execute_task.py @@ -311,34 +311,45 @@ def _publish_failed_workflow_terminal_events(exc: Exception, exec_params: AppExe topic.publish(json.dumps(finished_payload.model_dump(mode="json"), ensure_ascii=False).encode()) -def _get_event_name(event: str | Mapping[str, Any] | BaseModel) -> str | None: +def _get_event_data(event: str | Mapping[str, Any] | BaseModel) -> Mapping[str, Any] | None: if isinstance(event, BaseModel): # Temporary compatibility for legacy BaseModel stream events; remove after confirming generators always emit # str / Mapping responses. - event_name = getattr(event, "event", None) - elif isinstance(event, Mapping): - event_name = event.get("event") - else: + return event.model_dump() + if isinstance(event, Mapping): + return event + return None + + +def _get_event_name(event: str | Mapping[str, Any] | BaseModel) -> str | None: + event_data = _get_event_data(event) + if event_data is None: return None + event_name = event_data.get("event") if event_name is None: return None return str(event_name) def _get_task_id(event: str | Mapping[str, Any] | BaseModel) -> str | None: - if isinstance(event, BaseModel): - # Temporary compatibility for legacy BaseModel stream events; remove after confirming generators always emit - # str / Mapping responses. - task_id = getattr(event, "task_id", None) - elif isinstance(event, Mapping): - task_id = event.get("task_id") - else: + event_data = _get_event_data(event) + if event_data is None: return None + task_id = event_data.get("task_id") return task_id if isinstance(task_id, str) and task_id else None +def _get_error_message(event: str | Mapping[str, Any] | BaseModel) -> str | None: + event_data = _get_event_data(event) + if event_data is None: + return None + + message = event_data.get("message") + return message if isinstance(message, str) and message else None + + def _publish_streaming_response( response_stream: Generator[str | Mapping[str, Any] | BaseModel, None, None], workflow_run_id: str | uuid.UUID, @@ -406,6 +417,7 @@ def _publish_streaming_response( started_published = False terminal_published = False last_task_id = normalized_workflow_run_id + stream_error_message: str | None = None try: for event in response_stream: @@ -429,6 +441,8 @@ def _publish_streaming_response( started_published = True elif event_name in terminal_events: terminal_published = True + elif event_name == "error": + stream_error_message = _get_error_message(event) or stream_error_message except Exception as exc: if not terminal_published: logger.exception( @@ -448,7 +462,7 @@ def _publish_streaming_response( normalized_workflow_run_id, ) _publish_failed_terminal_event( - error_message=unexpected_stream_end_message, + error_message=stream_error_message or unexpected_stream_end_message, task_id=last_task_id, publish_started=not started_published, ) @@ -457,7 +471,7 @@ def _publish_streaming_response( @shared_task(queue=WORKFLOW_BASED_APP_EXECUTION_QUEUE) def workflow_based_app_execution_task( payload: str, -) -> Generator[Mapping[str, Any] | str, None, None] | Mapping[str, Any] | None: +) -> Mapping[str, Any] | None: exec_params = AppExecutionParams.model_validate_json(payload) logger.info("workflow_based_app_execution_task run with params: %s", exec_params) diff --git a/api/tasks/collect_agent_resources_task.py b/api/tasks/collect_agent_resources_task.py new file mode 100644 index 00000000000..2b84cc01b79 --- /dev/null +++ b/api/tasks/collect_agent_resources_task.py @@ -0,0 +1,75 @@ +"""Asynchronously collect retired Agent working resources.""" + +from __future__ import annotations + +import logging +from collections.abc import Iterable + +from celery import shared_task + +from services.agent.home_snapshot_service import AgentHomeSnapshotService +from services.agent.workspace_service import AgentWorkspaceService + +logger = logging.getLogger(__name__) + + +@shared_task(queue="retention") +def collect_agent_resources( + *, + tenant_id: str, + binding_ids: list[str], + workspace_ids: list[str], + home_snapshot_ids: list[str], +) -> None: + """Collect only the explicitly identified RETIRED resources.""" + + collectors = ( + (workspace_ids, "workspace_id", AgentWorkspaceService.collect_retired_workspace), + (binding_ids, "binding_id", AgentWorkspaceService.collect_retired_binding), + ( + home_snapshot_ids, + "home_snapshot_id", + AgentHomeSnapshotService.collect_retired_home_snapshot, + ), + ) + for resource_ids, argument_name, collector in collectors: + for resource_id in resource_ids: + try: + collector(tenant_id=tenant_id, **{argument_name: resource_id}) + except Exception: + logger.exception( + "Failed to collect retired Agent resource", + extra={ + "tenant_id": tenant_id, + "resource_type": argument_name.removesuffix("_id"), + "resource_id": resource_id, + }, + ) + + +def enqueue_agent_resource_collection( + *, + tenant_id: str, + binding_ids: Iterable[str] = (), + workspace_ids: Iterable[str] = (), + home_snapshot_ids: Iterable[str] = (), +) -> None: + """Best-effort enqueue of physical collection after retire has committed.""" + + payload = { + "binding_ids": sorted({resource_id for resource_id in binding_ids if resource_id}), + "workspace_ids": sorted({resource_id for resource_id in workspace_ids if resource_id}), + "home_snapshot_ids": sorted({resource_id for resource_id in home_snapshot_ids if resource_id}), + } + if not any(payload.values()): + return + try: + collect_agent_resources.delay(tenant_id=tenant_id, **payload) + except Exception: + logger.exception( + "Failed to enqueue retired Agent resource collection", + extra={"tenant_id": tenant_id, **payload}, + ) + + +__all__ = ["collect_agent_resources", "enqueue_agent_resource_collection"] diff --git a/api/tasks/community_telemetry_task.py b/api/tasks/community_telemetry_task.py new file mode 100644 index 00000000000..c0c6eb46a88 --- /dev/null +++ b/api/tasks/community_telemetry_task.py @@ -0,0 +1,19 @@ +import logging + +from celery import shared_task +from sqlalchemy.orm import sessionmaker + +from extensions.ext_database import db +from services.telemetry_service import CommunityTelemetryService + +logger = logging.getLogger(__name__) + + +@shared_task(name="community_telemetry.send_heartbeat", queue="schedule_executor") +def send_community_telemetry_heartbeat() -> None: + session_factory = sessionmaker(bind=db.engine, expire_on_commit=False) + with session_factory() as session: + try: + CommunityTelemetryService.report_heartbeat(session=session) + except Exception: + logger.debug("Failed to process community telemetry heartbeat", exc_info=True) diff --git a/api/tasks/delete_account_task.py b/api/tasks/delete_account_task.py index 55a99dde7a1..899df12c7ad 100644 --- a/api/tasks/delete_account_task.py +++ b/api/tasks/delete_account_task.py @@ -5,6 +5,7 @@ from sqlalchemy import select from configs import dify_config from core.db.session_factory import session_factory +from enums import DeploymentEdition from models import Account from services.billing_service import BillingService from tasks.mail_account_deletion_task import send_deletion_success_task @@ -17,7 +18,7 @@ def delete_account_task(account_id): with session_factory.create_session() as session: account = session.scalar(select(Account).where(Account.id == account_id).limit(1)) try: - if dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD: BillingService.delete_account(account_id) except Exception: logger.exception("Failed to delete account %s from billing service.", account_id) diff --git a/api/tasks/delete_conversation_task.py b/api/tasks/delete_conversation_task.py index 0b392f60965..6f426fe7fb0 100644 --- a/api/tasks/delete_conversation_task.py +++ b/api/tasks/delete_conversation_task.py @@ -17,7 +17,7 @@ logger = logging.getLogger(__name__) @shared_task(queue="conversation") def delete_conversation_related_data(conversation_id: str): """ - Delete related data conversation in correct order from datatbase to respect foreign key constraints + Delete related data conversation in correct order from database to respect foreign key constraints Args: conversation_id: conversation Id diff --git a/api/tasks/document_indexing_task.py b/api/tasks/document_indexing_task.py index 5d8e6dd701c..2ad17d31577 100644 --- a/api/tasks/document_indexing_task.py +++ b/api/tasks/document_indexing_task.py @@ -13,7 +13,7 @@ from core.entities.document_task import DocumentTask from core.indexing_runner import DocumentIsPausedError, IndexingRunner from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType from core.rag.pipeline.queue import TenantIsolatedTaskQueue -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from libs.datetime_utils import naive_utc_now from models.dataset import Dataset, Document from models.enums import IndexingStatus @@ -107,7 +107,7 @@ def _document_indexing(dataset_id: str, document_ids: Sequence[str]): # Phase 2: Execute indexing without holding locks from the parsing-status update. has_error = False try: - indexing_runner = IndexingRunner() + indexing_runner = IndexingRunner(enforce_vector_space_admission=True) with session_factory.create_session() as session: dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) if not dataset: diff --git a/api/tasks/duplicate_document_indexing_task.py b/api/tasks/duplicate_document_indexing_task.py index 10f37db7494..d062794df33 100644 --- a/api/tasks/duplicate_document_indexing_task.py +++ b/api/tasks/duplicate_document_indexing_task.py @@ -12,7 +12,7 @@ from core.entities.document_task import DocumentTask from core.indexing_runner import DocumentIsPausedError, IndexingRunner from core.rag.index_processor.index_processor_factory import IndexProcessorFactory from core.rag.pipeline.queue import TenantIsolatedTaskQueue -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from libs.datetime_utils import naive_utc_now from models.dataset import Dataset, Document, DocumentSegment from models.enums import IndexingStatus diff --git a/api/tasks/refresh_billing_vector_space_task.py b/api/tasks/refresh_billing_vector_space_task.py index ff3da012e3e..b6d0a377598 100644 --- a/api/tasks/refresh_billing_vector_space_task.py +++ b/api/tasks/refresh_billing_vector_space_task.py @@ -4,6 +4,7 @@ from celery import shared_task from opentelemetry import metrics from configs import dify_config +from enums import DeploymentEdition from services.billing_service import BillingService logger = logging.getLogger(__name__) @@ -19,7 +20,7 @@ _refresh_counter = metrics.get_meter(__name__).create_counter( @shared_task(queue="dataset", bind=True, max_retries=_MAX_RETRIES, default_retry_delay=_RETRY_DELAY_SECONDS) def refresh_billing_vector_space_task(self, tenant_id: str) -> None: """Refresh billing vector-space usage after vector cleanup has completed.""" - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return try: @@ -49,7 +50,7 @@ def refresh_billing_vector_space_task(self, tenant_id: str) -> None: def schedule_billing_vector_space_refresh(tenant_id: str) -> None: """Dispatch a best-effort billing refresh without changing cleanup status.""" - if not dify_config.BILLING_ENABLED: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD: return try: diff --git a/api/tasks/remove_app_and_related_data_task.py b/api/tasks/remove_app_and_related_data_task.py index 2010deb4a80..4562b7d1d90 100644 --- a/api/tasks/remove_app_and_related_data_task.py +++ b/api/tasks/remove_app_and_related_data_task.py @@ -5,25 +5,18 @@ from typing import Any, cast import click import sqlalchemy as sa -from agenton.compositor import CompositorSessionSnapshot from celery import shared_task -from dify_agent.protocol import RuntimeLayerSpec -from pydantic import JsonValue, TypeAdapter from sqlalchemy import delete, select from sqlalchemy.engine import CursorResult from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.orm import sessionmaker -from clients.agent_backend.session_cleanup import AgentBackendSessionCleanupPayload from configs import dify_config from core.db.session_factory import session_factory +from enums import DeploymentEdition from extensions.ext_database import db from libs.archive_storage import ArchiveStorageNotConfiguredError, get_archive_storage -from libs.datetime_utils import naive_utc_now from models import ( - AgentRuntimeSession, - AgentRuntimeSessionOwnerType, - AgentRuntimeSessionStatus, ApiToken, AppAnnotationHitHistory, AppAnnotationSetting, @@ -58,13 +51,8 @@ from models.workflow import ( ) from repositories.factory import DifyAPIRepositoryFactory from services.api_token_service import ApiTokenCache -from tasks.agent_backend_session_cleanup_task import ( - cleanup_conversation_agent_runtime_session, - cleanup_workflow_agent_runtime_session, -) logger = logging.getLogger(__name__) -_RUNTIME_LAYER_SPECS_ADAPTER: TypeAdapter[list[RuntimeLayerSpec]] = TypeAdapter(list[RuntimeLayerSpec]) @shared_task(queue="app_deletion", bind=True, max_retries=3) @@ -72,7 +60,6 @@ def remove_app_and_related_data_task(self, tenant_id: str, app_id: str): logger.info(click.style(f"Start deleting app and related data: {tenant_id}:{app_id}", fg="green")) start_at = time.perf_counter() try: - _cleanup_active_agent_runtime_sessions_for_app(tenant_id, app_id) # Delete related data _delete_app_model_configs(tenant_id, app_id) _delete_app_site(tenant_id, app_id) @@ -87,7 +74,7 @@ def remove_app_and_related_data_task(self, tenant_id: str, app_id: str): _delete_app_workflow_runs(tenant_id, app_id) _delete_app_workflow_node_executions(tenant_id, app_id) _delete_app_workflow_app_logs(tenant_id, app_id) - if dify_config.BILLING_ENABLED and dify_config.ARCHIVE_STORAGE_ENABLED: + if dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD and dify_config.ARCHIVE_STORAGE_ENABLED: _delete_app_workflow_archive_logs(tenant_id, app_id) _delete_archived_workflow_run_files(tenant_id, app_id) _delete_app_conversations(tenant_id, app_id) @@ -113,143 +100,6 @@ def remove_app_and_related_data_task(self, tenant_id: str, app_id: str): raise self.retry(exc=e, countdown=60) # Retry after 60 seconds -def _cleanup_active_agent_runtime_sessions_for_app(tenant_id: str, app_id: str, *, batch_size: int = 100) -> None: - """Best-effort fan-out for ACTIVE Agent runtime sessions during app deletion. - - App deletion must not block on synchronous Agent backend lifecycle work, so - this helper scans ACTIVE ``agent_runtime_sessions`` rows in batches, - dispatches owner-specific cleanup tasks only when enough persisted data - exists to replay a lifecycle-only run, and then marks each visited row - ``CLEANED`` locally regardless of enqueue outcome. The local retirement is - the contract that lets the rest of app deletion continue even when backend - cleanup dispatch is skipped or fails. - """ - if batch_size <= 0: - raise ValueError("batch_size must be positive") - - while True: - with session_factory.create_session() as session: - row_ids = session.scalars( - select(AgentRuntimeSession.id) - .where( - AgentRuntimeSession.tenant_id == tenant_id, - AgentRuntimeSession.app_id == app_id, - AgentRuntimeSession.status == AgentRuntimeSessionStatus.ACTIVE, - ) - .order_by(AgentRuntimeSession.updated_at.asc()) - .limit(batch_size) - ).all() - - if not row_ids: - return - - retired_count = 0 - for row_id in row_ids: - with session_factory.create_session() as session: - row = session.get(AgentRuntimeSession, row_id) - if row is None or row.status != AgentRuntimeSessionStatus.ACTIVE: - retired_count += 1 - continue - - try: - payload = _build_agent_runtime_session_cleanup_payload(row) - if payload is not None: - _enqueue_agent_runtime_session_cleanup(row=row, payload=payload) - except Exception: - logger.warning( - "Failed to enqueue Agent backend cleanup during app deletion: " - "tenant_id=%s app_id=%s owner_type=%s conversation_id=%s workflow_run_id=%s " - "node_id=%s agent_id=%s backend_run_id=%s", - row.tenant_id, - row.app_id, - row.owner_type, - row.conversation_id, - row.workflow_run_id, - row.node_id, - row.agent_id, - row.backend_run_id, - exc_info=True, - ) - finally: - try: - row.status = AgentRuntimeSessionStatus.CLEANED - row.cleaned_at = naive_utc_now() - session.commit() - retired_count += 1 - except Exception: - session.rollback() - logger.warning( - "Failed to retire Agent runtime session during app deletion: " - "tenant_id=%s app_id=%s owner_type=%s conversation_id=%s workflow_run_id=%s " - "node_id=%s agent_id=%s backend_run_id=%s", - row.tenant_id, - row.app_id, - row.owner_type, - row.conversation_id, - row.workflow_run_id, - row.node_id, - row.agent_id, - row.backend_run_id, - exc_info=True, - ) - - if retired_count == 0: - logger.warning( - "Failed to retire any active Agent runtime sessions during app deletion: tenant_id=%s app_id=%s", - tenant_id, - app_id, - ) - return - - -def _build_agent_runtime_session_cleanup_payload( - row: AgentRuntimeSession, -) -> AgentBackendSessionCleanupPayload | None: - runtime_layer_specs = _RUNTIME_LAYER_SPECS_ADAPTER.validate_json(row.composition_layer_specs or "[]") - if not runtime_layer_specs: - return None - - metadata: dict[str, JsonValue] = { - "tenant_id": row.tenant_id, - "app_id": row.app_id, - "agent_id": row.agent_id, - "agent_config_snapshot_id": row.agent_config_snapshot_id, - "previous_agent_backend_run_id": row.backend_run_id, - } - if row.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION: - metadata["conversation_id"] = row.conversation_id - idempotency_key = ( - f"{row.tenant_id}:{row.app_id}:{row.conversation_id}:" - f"{row.agent_id}:app-delete-cleanup:{row.id or row.backend_run_id or 'no-session-id'}" - ) - else: - metadata["workflow_run_id"] = row.workflow_run_id - metadata["node_id"] = row.node_id - idempotency_key = ( - f"{row.tenant_id}:{row.app_id}:{row.workflow_run_id}:{row.node_id}:" - f"{row.agent_id}:app-delete-cleanup:{row.id or row.backend_run_id or 'no-session-id'}" - ) - - return AgentBackendSessionCleanupPayload( - session_snapshot=CompositorSessionSnapshot.model_validate_json(row.session_snapshot), - runtime_layer_specs=runtime_layer_specs, - idempotency_key=idempotency_key, - metadata=metadata, - ) - - -def _enqueue_agent_runtime_session_cleanup( - *, - row: AgentRuntimeSession, - payload: AgentBackendSessionCleanupPayload, -) -> None: - payload_dict = payload.model_dump(mode="json") - if row.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION: - cleanup_conversation_agent_runtime_session.delay(payload_dict) - return - cleanup_workflow_agent_runtime_session.delay(payload_dict) - - def _delete_app_model_configs(tenant_id: str, app_id: str): def del_model_config(session, model_config_id: str): session.execute( diff --git a/api/tasks/retry_document_indexing_task.py b/api/tasks/retry_document_indexing_task.py index f8430cc206a..dcf74737544 100644 --- a/api/tasks/retry_document_indexing_task.py +++ b/api/tasks/retry_document_indexing_task.py @@ -113,7 +113,7 @@ def retry_document_indexing_task(dataset_id: str, document_ids: list[str], user_ rag_pipeline_service = RagPipelineService(rag_session) rag_pipeline_service.retry_error_document(dataset, document, user) else: - indexing_runner = IndexingRunner() + indexing_runner = IndexingRunner(enforce_vector_space_admission=True) indexing_runner.run([document], session) session.commit() redis_client.delete(retry_indexing_cache_key) diff --git a/api/tasks/sync_website_document_indexing_task.py b/api/tasks/sync_website_document_indexing_task.py index dbe029b6016..6cd6d40459a 100644 --- a/api/tasks/sync_website_document_indexing_task.py +++ b/api/tasks/sync_website_document_indexing_task.py @@ -10,8 +10,9 @@ from core.indexing_runner import IndexingRunner from core.rag.index_processor.index_processor_factory import IndexProcessorFactory from extensions.ext_redis import redis_client from libs.datetime_utils import naive_utc_now -from models.dataset import Dataset, Document, DocumentSegment +from models.dataset import Dataset, DocumentSegment from models.enums import IndexingStatus +from services.dataset_ref_service import DatasetRefService from services.feature_service import FeatureService logger = logging.getLogger(__name__) @@ -32,6 +33,15 @@ def sync_website_document_indexing_task(dataset_id: str, document_id: str): dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1)) if dataset is None: raise ValueError("Dataset not found") + tenant_id = dataset.tenant_id + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + document_ref = DatasetRefService.create_document_ref_from_id(dataset_ref, document_id) + document = DatasetRefService.get_document_by_ref(document_ref, session=session) + if document is None: + logger.info(click.style(f"Document not found: {document_id}", fg="yellow")) + return + if document.data_source_type != "website_crawl": + raise ValueError("Document is not a website document") sync_indexing_cache_key = f"document_{document_id}_is_sync" # check document limit @@ -46,40 +56,42 @@ def sync_website_document_indexing_task(dataset_id: str, document_id: str): "your subscription." ) except Exception as e: - document = session.scalar( - select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) - ) - if document: - document.indexing_status = IndexingStatus.ERROR - document.error = str(e) - document.stopped_at = naive_utc_now() - session.add(document) - session.commit() + document.indexing_status = IndexingStatus.ERROR + document.error = str(e) + document.stopped_at = naive_utc_now() + session.add(document) + session.commit() redis_client.delete(sync_indexing_cache_key) return logger.info(click.style(f"Start sync website document: {document_id}", fg="green")) - document = session.scalar( - select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) - ) - if not document: - logger.info(click.style(f"Document not found: {document_id}", fg="yellow")) - return try: # clean old data index_processor = IndexProcessorFactory(document.doc_form).init_index_processor() - segments = session.scalars(select(DocumentSegment).where(DocumentSegment.document_id == document_id)).all() + segments = session.scalars( + select(DocumentSegment).where( + DocumentSegment.tenant_id == tenant_id, + DocumentSegment.dataset_id == dataset_id, + DocumentSegment.document_id == document_id, + ) + ).all() if segments: index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id] # delete from vector index - index_processor.clean( - dataset, index_node_ids, with_keywords=True, delete_child_chunks=True, session=session - ) + if index_node_ids: + index_processor.clean( + dataset, index_node_ids, with_keywords=True, delete_child_chunks=True, session=session + ) segment_ids = [segment.id for segment in segments] if segment_ids: - segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.id.in_(segment_ids)) + segment_delete_stmt = delete(DocumentSegment).where( + DocumentSegment.id.in_(segment_ids), + DocumentSegment.tenant_id == tenant_id, + DocumentSegment.dataset_id == dataset_id, + DocumentSegment.document_id == document_id, + ) session.execute(segment_delete_stmt) session.commit() @@ -95,9 +107,7 @@ def sync_website_document_indexing_task(dataset_id: str, document_id: str): redis_client.delete(sync_indexing_cache_key) except Exception as ex: session.rollback() - document = session.scalar( - select(Document).where(Document.id == document_id, Document.dataset_id == dataset_id).limit(1) - ) + document = DatasetRefService.get_document_by_ref(document_ref, session=session) if document: document.indexing_status = IndexingStatus.ERROR document.error = str(ex) diff --git a/api/tasks/trigger_processing_tasks.py b/api/tasks/trigger_processing_tasks.py index 93f7f01407b..59e2f4e364c 100644 --- a/api/tasks/trigger_processing_tasks.py +++ b/api/tasks/trigger_processing_tasks.py @@ -26,7 +26,7 @@ from core.trigger.entities.entities import TriggerProviderEntity from core.trigger.provider import PluginTriggerProviderController from core.trigger.trigger_manager import TriggerManager from core.workflow.nodes.trigger_plugin.entities import TriggerEventNodeData -from enums.quota_type import QuotaType +from enums import QuotaType from graphon.enums import WorkflowExecutionStatus from models.enums import ( AppTriggerType, diff --git a/api/tasks/workflow_cfs_scheduler/entities.py b/api/tasks/workflow_cfs_scheduler/entities.py index e95d606731c..d308ed0c680 100644 --- a/api/tasks/workflow_cfs_scheduler/entities.py +++ b/api/tasks/workflow_cfs_scheduler/entities.py @@ -1,7 +1,7 @@ from enum import StrEnum from configs import dify_config -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from services.workflow.entities import WorkflowScheduleCFSPlanEntity # Determine queue names based on edition diff --git a/api/tasks/workflow_node_execution_tasks.py b/api/tasks/workflow_node_execution_tasks.py index 0d5475a56d2..e55dcf45330 100644 --- a/api/tasks/workflow_node_execution_tasks.py +++ b/api/tasks/workflow_node_execution_tasks.py @@ -13,6 +13,7 @@ from celery import shared_task from sqlalchemy import select from core.db.session_factory import session_factory +from core.workflow.node_execution_process_data import preserve_workflow_agent_binding_id from graphon.entities.workflow_node_execution import ( WorkflowNodeExecution, ) @@ -144,8 +145,9 @@ def _update_node_execution_from_domain(node_execution: WorkflowNodeExecutionMode # Update serialized data json_converter = WorkflowRuntimeTypeConverter() node_execution.inputs = json.dumps(json_converter.to_json_encodable(execution.inputs)) if execution.inputs else "{}" + process_data = preserve_workflow_agent_binding_id(node_execution.process_data_dict, execution.process_data) node_execution.process_data = ( - json.dumps(json_converter.to_json_encodable(execution.process_data)) if execution.process_data else "{}" + json.dumps(json_converter.to_json_encodable(process_data)) if process_data is not None else "{}" ) node_execution.outputs = ( json.dumps(json_converter.to_json_encodable(execution.outputs)) if execution.outputs else "{}" diff --git a/api/tasks/workflow_schedule_tasks.py b/api/tasks/workflow_schedule_tasks.py index 38737f96e78..36ed189c7c1 100644 --- a/api/tasks/workflow_schedule_tasks.py +++ b/api/tasks/workflow_schedule_tasks.py @@ -8,7 +8,7 @@ from core.workflow.nodes.trigger_schedule.exc import ( ScheduleNotFoundError, TenantOwnerNotFoundError, ) -from enums.quota_type import QuotaType +from enums import QuotaType from models.trigger import WorkflowSchedulePlan from services.async_workflow_service import AsyncWorkflowService from services.errors.app import QuotaExceededError diff --git a/api/tests/integration_tests/.env.example b/api/tests/integration_tests/.env.example index 986ced5f85d..98186d59e83 100644 --- a/api/tests/integration_tests/.env.example +++ b/api/tests/integration_tests/.env.example @@ -95,6 +95,7 @@ HOLOGRES_EF_CONSTRUCTION=400 # Upload configuration UPLOAD_FILE_SIZE_LIMIT=15 +KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN=15 UPLOAD_FILE_BATCH_LIMIT=5 UPLOAD_IMAGE_FILE_SIZE_LIMIT=10 UPLOAD_VIDEO_FILE_SIZE_LIMIT=100 diff --git a/api/tests/integration_tests/controllers/openapi/conftest.py b/api/tests/integration_tests/controllers/openapi/conftest.py index 19a8ab673bc..d9357e81551 100644 --- a/api/tests/integration_tests/controllers/openapi/conftest.py +++ b/api/tests/integration_tests/controllers/openapi/conftest.py @@ -10,6 +10,7 @@ from datetime import UTC, datetime, timedelta import pytest from flask import Flask +from enums import DeploymentEdition from extensions.ext_database import db from extensions.ext_redis import redis_client from models import Account, App, OAuthAccessToken, Tenant, TenantAccountJoin @@ -21,12 +22,12 @@ def _sha256(token: str) -> str: @pytest.fixture(autouse=True) -def disable_enterprise(monkeypatch): +def disable_enterprise(monkeypatch: pytest.MonkeyPatch): """Default to CE behaviour for /openapi/v1 tests. Tests that exercise the EE branch override this with their own monkeypatch in-test.""" from configs import dify_config - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", False) + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @pytest.fixture diff --git a/api/tests/integration_tests/controllers/openapi/test_auth.py b/api/tests/integration_tests/controllers/openapi/test_auth.py index 5f0727fbbec..956a0904dab 100644 --- a/api/tests/integration_tests/controllers/openapi/test_auth.py +++ b/api/tests/integration_tests/controllers/openapi/test_auth.py @@ -12,6 +12,7 @@ import pytest from flask import Flask from flask.testing import FlaskClient +from enums import DeploymentEdition from extensions.ext_database import db from models import App, Tenant @@ -66,7 +67,7 @@ def test_layer0_denies_account_bearer_without_membership( assert res.json.get("message") == "workspace_membership_revoked" -def test_layer0_skipped_when_enterprise_enabled( +def test_layer0_skipped_for_enterprise_edition( test_client: FlaskClient, account_token: str, other_workspace_app: App, @@ -81,7 +82,7 @@ def test_layer0_skipped_when_enterprise_enabled( from configs import dify_config # Override the conftest autouse default for this test only. - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", True) + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) res = test_client.get( f"/openapi/v1/apps/{other_workspace_app.id}/info", diff --git a/api/tests/integration_tests/model_runtime/__mock/plugin_model.py b/api/tests/integration_tests/model_runtime/__mock/plugin_model.py index 9b2bb963d2f..7be1e2744fe 100644 --- a/api/tests/integration_tests/model_runtime/__mock/plugin_model.py +++ b/api/tests/integration_tests/model_runtime/__mock/plugin_model.py @@ -4,6 +4,7 @@ from collections.abc import Generator, Sequence from decimal import Decimal from json import dumps +from core.plugin.entities.plugin import PluginInstallationSource from core.plugin.entities.plugin_daemon import PluginModelProviderEntity from core.plugin.impl.model import PluginModelClient @@ -41,6 +42,7 @@ class MockModelClass(PluginModelClient): tenant_id=tenant_id, plugin_unique_identifier="langgenius/openai/openai", plugin_id="langgenius/openai", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider="openai", label=I18nObject( diff --git a/api/tests/integration_tests/services/retention/test_workflow_run_archiver.py b/api/tests/integration_tests/services/retention/test_workflow_run_archiver.py index 90639ac60d6..098069776f2 100644 --- a/api/tests/integration_tests/services/retention/test_workflow_run_archiver.py +++ b/api/tests/integration_tests/services/retention/test_workflow_run_archiver.py @@ -8,6 +8,7 @@ import pyarrow.parquet as pq import pytest from sqlalchemy.exc import OperationalError +from enums import DeploymentEdition from models.workflow import WorkflowRunArchiveBundle from services.retention.workflow_run.archive_paid_plan_workflow_run import ( ArchiveResult, @@ -281,12 +282,12 @@ class TestGenerateManifest: class TestFilterPaidTenants: - def test_all_tenants_paid_when_billing_disabled(self): + def test_all_tenants_paid_in_community_edition(self): archiver = WorkflowRunArchiver(days=90) tenant_ids = {"t1", "t2", "t3"} with patch("services.retention.workflow_run.archive_paid_plan_workflow_run.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY result = archiver._filter_paid_tenants(tenant_ids) assert result == tenant_ids @@ -295,7 +296,7 @@ class TestFilterPaidTenants: archiver = WorkflowRunArchiver(days=90) with patch("services.retention.workflow_run.archive_paid_plan_workflow_run.dify_config") as cfg: - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD result = archiver._filter_paid_tenants(set()) assert result == set() @@ -313,7 +314,7 @@ class TestFilterPaidTenants: patch("services.retention.workflow_run.archive_paid_plan_workflow_run.dify_config") as cfg, patch("services.retention.workflow_run.archive_paid_plan_workflow_run.BillingService") as billing, ): - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD billing.get_plan_bulk_with_cache.return_value = mock_bulk result = archiver._filter_paid_tenants({"t1", "t2", "t3"}) @@ -328,7 +329,7 @@ class TestFilterPaidTenants: patch("services.retention.workflow_run.archive_paid_plan_workflow_run.dify_config") as cfg, patch("services.retention.workflow_run.archive_paid_plan_workflow_run.BillingService") as billing, ): - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD billing.get_plan_bulk_with_cache.side_effect = RuntimeError("API down") result = archiver._filter_paid_tenants({"t1"}) @@ -341,7 +342,7 @@ class TestFilterPaidTenants: patch("services.retention.workflow_run.archive_paid_plan_workflow_run.dify_config") as cfg, patch("services.retention.workflow_run.archive_paid_plan_workflow_run.BillingService") as billing, ): - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD result = archiver._filter_paid_tenants({"t1", "t2", "t3"}) billing.get_plan_bulk_with_cache.assert_not_called() diff --git a/api/tests/integration_tests/workflow/nodes/test_http.py b/api/tests/integration_tests/workflow/nodes/test_http.py index 7cd7f50b773..86b0cd0fd71 100644 --- a/api/tests/integration_tests/workflow/nodes/test_http.py +++ b/api/tests/integration_tests/workflow/nodes/test_http.py @@ -11,7 +11,7 @@ from core.tools.tool_file_manager import ToolFileManager from core.workflow.node_factory import DifyNodeFactory from core.workflow.node_runtime import DifyFileReferenceFactory from core.workflow.system_variables import build_system_variables -from graphon.enums import WorkflowNodeExecutionStatus +from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.file.file_manager import file_manager from graphon.graph import Graph from graphon.nodes.http_request import HttpRequestNode, HttpRequestNodeConfig, HttpRequestNodeData @@ -193,7 +193,6 @@ def test_custom_authorization_header(setup_http_mock): def test_custom_auth_with_empty_api_key_raises_error(setup_http_mock): """Test: In custom authentication mode, when the api_key is empty, AuthorizationConfigError should be raised.""" from core.workflow.system_variables import build_system_variables - from graphon.enums import BuiltinNodeTypes from graphon.nodes.http_request.entities import ( HttpRequestNodeAuthorization, HttpRequestNodeData, diff --git a/api/tests/integration_tests/workflow/nodes/test_llm.py b/api/tests/integration_tests/workflow/nodes/test_llm.py index d8a0a713f12..2e9b7317f69 100644 --- a/api/tests/integration_tests/workflow/nodes/test_llm.py +++ b/api/tests/integration_tests/workflow/nodes/test_llm.py @@ -97,7 +97,7 @@ def _mock_db_session_close(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(db.session, "close", MagicMock()) -def test_execute_llm(monkeypatch): +def test_execute_llm(monkeypatch: pytest.MonkeyPatch): node = init_llm_node( config={ "id": "llm", @@ -201,7 +201,7 @@ def test_execute_llm(monkeypatch): assert item.node_run_result.outputs.get("usage", {})["total_tokens"] > 0 -def test_execute_llm_with_jinja2(monkeypatch): +def test_execute_llm_with_jinja2(monkeypatch: pytest.MonkeyPatch): """ Test execute LLM node with jinja2 """ diff --git a/api/tests/integration_tests/workflow/nodes/test_template_transform.py b/api/tests/integration_tests/workflow/nodes/test_template_transform.py index 80489e68097..9a7b02597d0 100644 --- a/api/tests/integration_tests/workflow/nodes/test_template_transform.py +++ b/api/tests/integration_tests/workflow/nodes/test_template_transform.py @@ -17,10 +17,11 @@ class _SimpleJinja2Renderer: """Minimal Jinja2-based renderer for integration tests (no code executor).""" def render_template(self, template: str, variables: dict[str, object]) -> str: - from jinja2 import Template + from jinja2.sandbox import SandboxedEnvironment try: - return Template(template).render(**variables) + env = SandboxedEnvironment() + return env.from_string(template).render(**variables) except Exception as exc: raise TemplateRenderError(str(exc)) from exc diff --git a/api/tests/integration_tests/workflow/nodes/test_tool.py b/api/tests/integration_tests/workflow/nodes/test_tool.py index c109be9fae4..d3e248cc28c 100644 --- a/api/tests/integration_tests/workflow/nodes/test_tool.py +++ b/api/tests/integration_tests/workflow/nodes/test_tool.py @@ -70,6 +70,7 @@ def init_tool_node(config: dict): tool_file_manager=tool_file_manager, runtime=DifyToolNodeRuntime(init_params.run_context), ) + node.bind_execution_id(str(uuid.uuid4())) return node diff --git a/api/tests/test_containers_integration_tests/.ruff.toml b/api/tests/test_containers_integration_tests/.ruff.toml index 54bc58a79ad..49244091ffb 100644 --- a/api/tests/test_containers_integration_tests/.ruff.toml +++ b/api/tests/test_containers_integration_tests/.ruff.toml @@ -2,114 +2,94 @@ extend = "../../.ruff.toml" src = ["../.."] [lint] -extend-select = ["ANN401", "ARG", "TID251"] +extend-select = ["ANN401", "ARG"] +# Existing strict-mode debt. Remove a file entry when bringing it under strict checking. [lint.per-file-ignores] -"core/rag/pipeline/test_queue_integration.py" = ["ANN401", "TID251", "ARG"] +"controllers/console/test_apikey.py" = ["ARG002"] +"controllers/openapi/test_app_dsl.py" = ["ARG002"] +"controllers/service_api/dataset/test_dataset.py" = ["ARG002"] +"controllers/web/test_conversation.py" = ["ARG002"] +"controllers/web/test_human_input_form.py" = ["ARG001"] +"controllers/web/test_wraps.py" = ["ARG002"] +"core/app/layers/test_pause_state_persist_layer.py" = ["ARG002"] +"core/rag/pipeline/test_queue_integration.py" = ["ARG002", "TID251"] +"core/rag/retrieval/test_dataset_retrieval_integration.py" = ["ARG002"] +"models/test_conversation_message_inputs.py" = ["ARG001"] "models/test_types_enum_text.py" = ["ANN401", "TID251"] -"services/test_app_dsl_service.py" = ["ANN401", "TID251", "ARG"] -"services/test_file_service_zip_and_lookup.py" = ["ANN401", "TID251", "ARG"] -"trigger/conftest.py" = ["ANN401", "TID251"] -"trigger/test_trigger_e2e.py" = ["ANN401", "TID251", "ARG"] -"controllers/console/app/test_app_apis.py" = ["ARG"] -"controllers/console/app/test_app_import_api.py" = ["ARG"] -"controllers/console/auth/test_oauth.py" = ["ARG"] -"controllers/console/auth/test_password_reset.py" = ["ARG"] -"controllers/console/datasets/test_data_source.py" = ["ARG"] -"controllers/console/test_apikey.py" = ["ARG"] -"controllers/console/workspace/test_tool_provider.py" = ["ARG"] -"controllers/mcp/test_mcp.py" = ["ARG"] -"controllers/openapi/test_app_dsl.py" = ["ARG"] -"controllers/openapi/test_workspaces.py" = ["ARG"] -"controllers/service_api/dataset/test_dataset.py" = ["ARG"] -"controllers/web/test_conversation.py" = ["ARG"] -"controllers/web/test_human_input_form.py" = ["ARG"] -"controllers/web/test_wraps.py" = ["ARG"] -"core/app/layers/test_pause_state_persist_layer.py" = ["ARG"] -"core/rag/retrieval/test_dataset_retrieval_integration.py" = ["ARG"] -"models/test_conversation_message_inputs.py" = ["ARG"] -"models/test_conversation_status_count.py" = ["ARG"] -"repositories/test_sqlalchemy_api_workflow_run_repository.py" = ["ARG"] -"repositories/test_workflow_run_repository.py" = ["ARG"] -"services/auth/test_api_key_auth_service.py" = ["ARG"] -"services/auth/test_auth_integration.py" = ["ARG"] -"services/dataset_collection_binding.py" = ["ARG"] -"services/dataset_service_update_delete.py" = ["ARG"] -"services/document_service_status.py" = ["ARG"] -"services/enterprise/test_account_deletion_sync.py" = ["ARG"] -"services/plugin/test_plugin_parameter_service.py" = ["ARG"] -"services/plugin/test_plugin_service.py" = ["ARG"] -"services/rag_pipeline/test_rag_pipeline_service_db.py" = ["ARG"] -"services/recommend_app/test_database_retrieval.py" = ["ARG"] -"services/test_account_service.py" = ["ARG"] -"services/test_advanced_prompt_template_service.py" = ["ARG"] -"services/test_annotation_service.py" = ["ARG"] -"services/test_api_based_extension_service.py" = ["ARG"] -"services/test_api_token_service.py" = ["ARG"] -"services/test_app_generate_service.py" = ["ARG"] -"services/test_app_service.py" = ["ARG"] -"services/test_attachment_service.py" = ["ARG"] -"services/test_conversation_variable_updater.py" = ["ARG"] -"services/test_dataset_permission_service.py" = ["ARG"] -"services/test_dataset_service_batch_update_document_status.py" = ["ARG"] -"services/test_dataset_service_retrieval.py" = ["ARG"] -"services/test_delete_archived_workflow_run.py" = ["ARG"] -"services/test_document_service_rename_document.py" = ["ARG"] -"services/test_end_user_service.py" = ["ARG"] -"services/test_feature_service.py" = ["ARG"] -"services/test_feedback_service.py" = ["ARG"] -"services/test_file_service.py" = ["ARG"] -"services/test_human_input_delivery_test_service.py" = ["ARG"] -"services/test_message_service.py" = ["ARG"] -"services/test_messages_clean_service.py" = ["ARG", "S110"] -"services/test_metadata_partial_update.py" = ["ARG"] -"services/test_metadata_service.py" = ["ARG"] -"services/test_model_load_balancing_service.py" = ["ARG"] -"services/test_model_provider_service.py" = ["ARG"] -"services/test_oauth_server_service.py" = ["ARG"] -"services/test_ops_service.py" = ["ARG"] -"services/test_saved_message_service.py" = ["ARG"] -"services/test_web_conversation_service.py" = ["ARG"] -"services/test_webapp_auth_service.py" = ["ARG"] -"services/test_webhook_service.py" = ["ARG"] -"services/test_workflow_app_service.py" = ["ARG"] -"services/test_workflow_draft_variable_service.py" = ["ARG"] -"services/test_workflow_run_service.py" = ["ARG"] -"services/test_workflow_service.py" = ["ARG"] -"services/test_workspace_service.py" = ["ARG"] -"services/tools/test_api_tools_manage_service.py" = ["ARG"] -"services/tools/test_mcp_tools_manage_service.py" = ["ARG"] -"services/tools/test_tools_transform_service.py" = ["ARG"] -"services/workflow/test_workflow_converter.py" = ["ARG"] -"tasks/test_add_document_to_index_task.py" = ["ARG"] -"tasks/test_batch_clean_document_task.py" = ["ARG"] -"tasks/test_batch_create_segment_to_index_task.py" = ["ARG"] +"repositories/test_sqlalchemy_api_workflow_run_repository.py" = ["ARG002", "ARG005"] +"repositories/test_workflow_run_repository.py" = ["ARG002"] +"services/auth/test_auth_integration.py" = ["ARG002"] +"services/dataset_collection_binding.py" = ["ARG002"] +"services/document_service_status.py" = ["ARG002"] +"services/rag_pipeline/test_rag_pipeline_service_db.py" = ["ARG002"] +"services/recommend_app/test_database_retrieval.py" = ["ARG002"] +"services/test_account_service.py" = ["ARG002"] +"services/test_advanced_prompt_template_service.py" = ["ARG002"] +"services/test_app_dsl_service.py" = ["ANN401", "ARG001", "ARG002", "ARG005", "TID251"] +"services/test_app_service.py" = ["ARG002"] +"services/test_attachment_service.py" = ["ARG002"] +"services/test_conversation_variable_updater.py" = ["ARG002"] +"services/test_dataset_service_batch_update_document_status.py" = ["ARG002"] +"services/test_delete_archived_workflow_run.py" = ["ARG002"] +"services/test_document_service_rename_document.py" = ["ARG001"] +"services/test_end_user_service.py" = ["ARG002"] +"services/test_feature_service.py" = ["ARG002"] +"services/test_file_service.py" = ["ARG002"] +"services/test_file_service_zip_and_lookup.py" = ["TID251"] +"services/test_messages_clean_service.py" = ["ARG002", "S110"] +"services/test_metadata_partial_update.py" = ["ARG002"] +"services/test_metadata_service.py" = ["ARG002"] +"services/test_model_load_balancing_service.py" = ["ARG002"] +"services/test_model_provider_service.py" = ["ARG002"] +"services/test_ops_service.py" = ["ARG002"] +"services/test_webapp_auth_service.py" = ["ARG002"] +"services/test_webhook_service.py" = ["ARG002"] +"services/test_workflow_draft_variable_service.py" = ["ARG002"] +"services/test_workflow_run_service.py" = ["ARG002"] +"services/test_workflow_service.py" = ["ARG002"] +"services/test_workspace_service.py" = ["ARG002"] +"services/tools/test_api_tools_manage_service.py" = ["ARG002"] +"services/tools/test_mcp_tools_manage_service.py" = ["ARG002", "ARG005"] +"services/tools/test_tools_transform_service.py" = ["ARG002"] +"services/workflow/test_workflow_converter.py" = ["ARG002"] +"tasks/test_add_document_to_index_task.py" = ["ARG002"] +"tasks/test_batch_clean_document_task.py" = ["ARG002"] +"tasks/test_batch_create_segment_to_index_task.py" = ["ARG001", "ARG002"] "tasks/test_clean_dataset_task.py" = ["T201"] -"tasks/test_clean_notion_document_task.py" = ["ARG"] -"tasks/test_create_segment_to_index_task.py" = ["ARG"] -"tasks/test_dataset_indexing_task.py" = ["ARG"] -"tasks/test_deal_dataset_vector_index_task.py" = ["ARG"] -"tasks/test_delete_segment_from_index_task.py" = ["ARG"] -"tasks/test_disable_segment_from_index_task.py" = ["ARG"] -"tasks/test_disable_segments_from_index_task.py" = ["ARG"] -"tasks/test_document_indexing_sync_task.py" = ["ARG"] -"tasks/test_document_indexing_task.py" = ["ARG"] -"tasks/test_document_indexing_update_task.py" = ["ARG"] -"tasks/test_duplicate_document_indexing_task.py" = ["ARG"] -"tasks/test_enable_segments_to_index_task.py" = ["ARG"] -"tasks/test_mail_change_mail_task.py" = ["ARG"] -"tasks/test_mail_email_code_login_task.py" = ["ARG"] -"tasks/test_mail_human_input_delivery_task.py" = ["ARG"] -"tasks/test_mail_inner_task.py" = ["ARG"] -"tasks/test_mail_invite_member_task.py" = ["ARG"] -"tasks/test_mail_owner_transfer_task.py" = ["ARG"] -"tasks/test_mail_register_task.py" = ["ARG"] -"tasks/test_rag_pipeline_run_tasks.py" = ["ARG"] +"tasks/test_clean_notion_document_task.py" = ["ARG002"] +"tasks/test_create_segment_to_index_task.py" = ["ARG002"] +"tasks/test_dataset_indexing_task.py" = ["ARG002"] +"tasks/test_deal_dataset_vector_index_task.py" = ["ARG002"] +"tasks/test_delete_segment_from_index_task.py" = ["ARG002"] +"tasks/test_disable_segment_from_index_task.py" = ["ARG002"] +"tasks/test_disable_segments_from_index_task.py" = ["ARG002"] +"tasks/test_document_indexing_sync_task.py" = ["ARG002"] +"tasks/test_document_indexing_task.py" = ["ARG002"] +"tasks/test_document_indexing_update_task.py" = ["ARG002"] +"tasks/test_duplicate_document_indexing_task.py" = ["ARG002"] +"tasks/test_enable_segments_to_index_task.py" = ["ARG002"] +"tasks/test_mail_change_mail_task.py" = ["ARG002"] +"tasks/test_mail_email_code_login_task.py" = ["ARG002"] +"tasks/test_mail_human_input_delivery_task.py" = ["ARG001"] +"tasks/test_mail_inner_task.py" = ["ARG002"] +"tasks/test_mail_invite_member_task.py" = ["ARG002"] +"tasks/test_mail_owner_transfer_task.py" = ["ARG002"] +"tasks/test_mail_register_task.py" = ["ARG002"] +"tasks/test_rag_pipeline_run_tasks.py" = ["ARG002"] "test_workflow_pause_integration.py" = ["T201"] -"workflow/nodes/code_executor/test_code_javascript.py" = ["ARG"] -"workflow/nodes/code_executor/test_code_jinja2.py" = ["ARG"] -"workflow/nodes/code_executor/test_code_python3.py" = ["ARG"] +"trigger/conftest.py" = ["ANN401", "TID251"] +"trigger/test_trigger_e2e.py" = ["ANN401", "ARG001", "TID251"] +"workflow/nodes/code_executor/test_code_javascript.py" = ["ARG002"] +"workflow/nodes/code_executor/test_code_jinja2.py" = ["ARG002"] +"workflow/nodes/code_executor/test_code_python3.py" = ["ARG002"] "workflow/nodes/code_executor/test_utils.py" = ["T201"] +[lint.flake8-tidy-imports.banned-api."flask_restx.reqparse"] +msg = "Use Pydantic payload/query models instead of reqparse." + +[lint.flake8-tidy-imports.banned-api."flask_restx.reqparse.RequestParser"] +msg = "Use Pydantic payload/query models instead of reqparse." + [lint.flake8-tidy-imports.banned-api."typing.Any"] msg = "Use object, Protocol, TypedDict, TypeVar, ParamSpec, or a localized cast instead." diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py index e6af51d6022..7852315b671 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py @@ -36,7 +36,7 @@ def test_export_customized_pipeline_template_from_database( db_session_with_containers.expire_all() with flask_app_with_containers.test_request_context("/"): - response, status = method(api, template.id) + response, status = method(api, db_session_with_containers, template.tenant_id, template.id) assert status == 200 assert response == {"data": "yaml-data"} diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py b/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py index 7c63cee4107..bc0a205ecb1 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py +++ b/api/tests/test_containers_integration_tests/controllers/console/datasets/test_data_source.py @@ -1,494 +1,89 @@ -"""Testcontainers integration tests for controllers.console.datasets.data_source endpoints.""" +"""Integration coverage for Notion page bindings backed by persisted documents.""" -from __future__ import annotations - -import inspect -from collections.abc import Iterator -from datetime import UTC, datetime -from unittest.mock import MagicMock, PropertyMock, patch +from inspect import unwrap +from unittest.mock import MagicMock, patch from uuid import uuid4 -import pytest from flask import Flask from sqlalchemy.orm import Session -from werkzeug.exceptions import NotFound -from controllers.console.datasets import data_source -from controllers.console.datasets.data_source import ( - DataSourceApi, - DataSourceNotionDatasetSyncApi, - DataSourceNotionDocumentSyncApi, - DataSourceNotionIndexingEstimateApi, - DataSourceNotionListApi, - DataSourceNotionPreviewApi, -) -from core.rag.index_processor.constant.index_type import IndexStructureType -from models import Account, DataSourceOauthBinding +from controllers.console.datasets.data_source import DataSourceNotionListApi, DataSourceNotionListQuery +from models import Account from models.dataset import Document from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus -@pytest.fixture -def current_user() -> Account: - account = Account(name="Test User", email="u1@example.com") - account.id = "u1" - return account +def test_notion_page_is_marked_bound_from_persisted_document( + flask_app_with_containers: Flask, + db_session_with_containers: Session, +) -> None: + tenant_id = str(uuid4()) + dataset_id = str(uuid4()) + account = Account(name="Test User", email="user@example.com") + account.id = str(uuid4()) + document = Document( + tenant_id=tenant_id, + dataset_id=dataset_id, + position=1, + data_source_type=DataSourceType.NOTION_IMPORT, + data_source_info='{"notion_page_id": "page-1"}', + batch=f"batch-{uuid4()}", + name="Notion Page", + created_from=DocumentCreatedFrom.WEB, + created_by=str(uuid4()), + indexing_status=IndexingStatus.COMPLETED, + enabled=True, + ) + db_session_with_containers.add(document) + db_session_with_containers.commit() + runtime = MagicMock( + get_online_document_pages=lambda **_kwargs: iter( + [ + MagicMock( + result=[ + MagicMock( + workspace_id="workspace-1", + workspace_name="Workspace", + workspace_icon=None, + pages=[ + MagicMock( + page_id="page-1", + page_name="Page", + type="page", + parent_id="parent", + page_icon=None, + ) + ], + ) + ] + ) + ] + ), + datasource_provider_type=lambda: None, + ) - -@pytest.fixture -def mock_engine() -> Iterator[None]: - with patch.object( - type(data_source.db), - "engine", - new_callable=PropertyMock, - return_value=MagicMock(), + with ( + flask_app_with_containers.test_request_context(f"/?credential_id=c1&dataset_id={dataset_id}"), + patch( + "controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials", + return_value={"token": "token"}, + ), + patch( + "controllers.console.datasets.data_source.DatasetService.get_dataset", + return_value=MagicMock(data_source_type="notion_import"), + ), + patch( + "core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime", + return_value=runtime, + ), ): - yield - - -class TestDataSourceApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_get_success(self, app: Flask) -> None: - api = DataSourceApi() - method = inspect.unwrap(api.get) - - binding = DataSourceOauthBinding( - tenant_id="tenant-1", - access_token="token", - provider="notion", - source_info={ - "workspace_name": "Workspace", - "workspace_id": "workspace-1", - "workspace_icon": None, - "total": 1, - "pages": [ - { - "page_id": "page-1", - "page_name": "Page", - "page_icon": {"type": "emoji", "emoji": "P", "url": None}, - "parent_id": "parent-1", - "type": "page", - } - ], - }, - ) - binding.id = "b1" - binding.created_at = datetime(2026, 5, 25, 1, 2, 3, tzinfo=UTC) - binding.disabled = False - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.data_source.db.session.scalars", - return_value=MagicMock(all=lambda: [binding]), - ), - ): - response, status = method(api, "tenant-1") - - assert status == 200 - assert response["data"][0] == { - "id": "b1", - "provider": "notion", - "created_at": 1779670923, - "is_bound": True, - "disabled": False, - "source_info": { - "workspace_name": "Workspace", - "workspace_id": "workspace-1", - "workspace_icon": None, - "pages": [ - { - "page_name": "Page", - "page_id": "page-1", - "page_icon": {"type": "emoji", "url": None, "emoji": "P"}, - "parent_id": "parent-1", - "type": "page", - } - ], - "total": 1, - }, - "link": "http://localhost/console/api/oauth/data-source/notion", - } - - def test_get_no_bindings(self, app: Flask) -> None: - api = DataSourceApi() - method = inspect.unwrap(api.get) - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.data_source.db.session.scalars", - return_value=MagicMock(all=lambda: []), - ), - ): - response, status = method(api, "tenant-1") - - assert status == 200 - assert response["data"] == [] - - def test_patch_enable_binding(self, app: Flask) -> None: - api = DataSourceApi() - method = inspect.unwrap(api.patch) - - binding = MagicMock(id="b1", disabled=True) - session = MagicMock() - session.scalar.return_value = binding - - with app.test_request_context("/"): - response, status = method(api, session, "tenant-1", "b1", "enable") - - assert status == 200 - assert binding.disabled is False - - def test_patch_disable_binding(self, app: Flask) -> None: - api = DataSourceApi() - method = inspect.unwrap(api.patch) - - binding = MagicMock(id="b1", disabled=False) - session = MagicMock() - session.scalar.return_value = binding - - with app.test_request_context("/"): - response, status = method(api, session, "tenant-1", "b1", "disable") - - assert status == 200 - assert binding.disabled is True - - def test_patch_binding_not_found(self, app: Flask) -> None: - api = DataSourceApi() - method = inspect.unwrap(api.patch) - session = MagicMock() - session.scalar.return_value = None - - with app.test_request_context("/"): - with pytest.raises(NotFound): - method(api, session, "tenant-1", "b1", "enable") - - def test_patch_enable_already_enabled(self, app: Flask) -> None: - api = DataSourceApi() - method = inspect.unwrap(api.patch) - - binding = MagicMock(id="b1", disabled=False) - session = MagicMock() - session.scalar.return_value = binding - - with app.test_request_context("/"): - with pytest.raises(ValueError): - method(api, session, "tenant-1", "b1", "enable") - - def test_patch_disable_already_disabled(self, app: Flask) -> None: - api = DataSourceApi() - method = inspect.unwrap(api.patch) - - binding = MagicMock(id="b1", disabled=True) - session = MagicMock() - session.scalar.return_value = binding - - with app.test_request_context("/"): - with pytest.raises(ValueError): - method(api, session, "tenant-1", "b1", "disable") - - -class TestDataSourceNotionListApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_get_credential_not_found(self, app: Flask, current_user: Account) -> None: - api = DataSourceNotionListApi() - method = inspect.unwrap(api.get) - - with ( - app.test_request_context("/?credential_id=c1"), - patch( - "controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials", - return_value=None, - ), - ): - with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", current_user) - - def test_get_success_no_dataset_id(self, app: Flask, current_user: Account, mock_engine: None) -> None: - api = DataSourceNotionListApi() - method = inspect.unwrap(api.get) - - page = MagicMock( - page_id="p1", - page_name="Page 1", - type="page", - parent_id="parent", - page_icon=None, + response, status = unwrap(DataSourceNotionListApi().get)( + DataSourceNotionListApi(), + DataSourceNotionListQuery(credential_id="c1", dataset_id=dataset_id), + db_session_with_containers, + tenant_id, + account, ) - online_document_message = MagicMock( - result=[ - MagicMock( - workspace_id="w1", - workspace_name="My Workspace", - workspace_icon="icon", - pages=[page], - ) - ] - ) - - with ( - app.test_request_context("/?credential_id=c1"), - patch( - "controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials", - return_value={"token": "t"}, - ), - patch( - "core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime", - return_value=MagicMock( - get_online_document_pages=lambda **kw: iter([online_document_message]), - datasource_provider_type=lambda: None, - ), - ), - ): - response, status = method(api, MagicMock(), "tenant-1", current_user) - - assert status == 200 - - def test_get_success_with_dataset_id( - self, app: Flask, current_user: Account, mock_engine: None, db_session_with_containers: Session - ) -> None: - api = DataSourceNotionListApi() - method = inspect.unwrap(api.get) - tenant_id = str(uuid4()) - dataset_id = str(uuid4()) - - page = MagicMock( - page_id="p1", - page_name="Page 1", - type="page", - parent_id="parent", - page_icon=None, - ) - - online_document_message = MagicMock( - result=[ - MagicMock( - workspace_id="w1", - workspace_name="My Workspace", - workspace_icon="icon", - pages=[page], - ) - ] - ) - - dataset = MagicMock(data_source_type="notion_import") - document = Document( - tenant_id=tenant_id, - dataset_id=dataset_id, - position=1, - data_source_type=DataSourceType.NOTION_IMPORT, - data_source_info='{"notion_page_id": "p1"}', - batch=f"batch-{uuid4()}", - name="Notion Page", - created_from=DocumentCreatedFrom.WEB, - created_by=str(uuid4()), - indexing_status=IndexingStatus.COMPLETED, - enabled=True, - ) - db_session_with_containers.add(document) - db_session_with_containers.commit() - - with ( - app.test_request_context(f"/?credential_id=c1&dataset_id={dataset_id}"), - patch( - "controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials", - return_value={"token": "t"}, - ), - patch( - "controllers.console.datasets.data_source.DatasetService.get_dataset", - return_value=dataset, - ), - patch( - "core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime", - return_value=MagicMock( - get_online_document_pages=lambda **kw: iter([online_document_message]), - datasource_provider_type=lambda: None, - ), - ), - ): - response, status = method(api, db_session_with_containers, tenant_id, current_user) - - assert status == 200 - - def test_get_invalid_dataset_type(self, app: Flask, current_user: Account, mock_engine: None) -> None: - api = DataSourceNotionListApi() - method = inspect.unwrap(api.get) - - dataset = MagicMock(data_source_type="other_type") - - with ( - app.test_request_context("/?credential_id=c1&dataset_id=ds1"), - patch( - "controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials", - return_value={"token": "t"}, - ), - patch( - "controllers.console.datasets.data_source.DatasetService.get_dataset", - return_value=dataset, - ), - ): - with pytest.raises(ValueError): - method(api, MagicMock(), "tenant-1", current_user) - - -class TestDataSourceNotionPreviewApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_get_preview_success(self, app: Flask) -> None: - api = DataSourceNotionPreviewApi() - method = inspect.unwrap(api.get) - - extractor = MagicMock(extract=lambda: [MagicMock(page_content="hello")]) - - with ( - app.test_request_context("/?credential_id=c1"), - patch( - "controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials", - return_value={"integration_secret": "t"}, - ), - patch( - "controllers.console.datasets.data_source.NotionExtractor", - return_value=extractor, - ), - ): - response, status = method(api, "tenant-1", "p1", "page") - - assert status == 200 - - -class TestDataSourceNotionIndexingEstimateApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_post_indexing_estimate_success(self, app: Flask) -> None: - api = DataSourceNotionIndexingEstimateApi() - method = inspect.unwrap(api.post) - - empty_rules: dict[str, object] = {} - payload: dict[str, object] = { - "notion_info_list": [ - { - "workspace_id": "w1", - "credential_id": "c1", - "pages": [{"page_id": "p1", "type": "page"}], - } - ], - "process_rule": {"rules": empty_rules}, - "doc_form": IndexStructureType.PARAGRAPH_INDEX, - "doc_language": "English", - } - - with ( - app.test_request_context("/", method="POST", json=payload, headers={"Content-Type": "application/json"}), - patch( - "controllers.console.datasets.data_source.DocumentService.estimate_args_validate", - ), - patch( - "controllers.console.datasets.data_source.IndexingRunner.indexing_estimate", - return_value=MagicMock(model_dump=lambda: {"total_pages": 1}), - ), - ): - response, status = method(api, MagicMock(), "tenant-1") - - assert status == 200 - - -class TestDataSourceNotionDatasetSyncApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_get_success(self, app: Flask) -> None: - api = DataSourceNotionDatasetSyncApi() - method = inspect.unwrap(api.get) - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.data_source.DatasetService.get_dataset", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.data_source.DocumentService.get_document_by_dataset_id", - return_value=[MagicMock(id="d1")], - ), - patch( - "controllers.console.datasets.data_source.document_indexing_sync_task.delay", - return_value=None, - ), - ): - response, status = method(api, MagicMock(), "ds-1") - - assert status == 200 - - def test_get_dataset_not_found(self, app: Flask) -> None: - api = DataSourceNotionDatasetSyncApi() - method = inspect.unwrap(api.get) - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.data_source.DatasetService.get_dataset", - return_value=None, - ), - ): - with pytest.raises(NotFound): - method(api, MagicMock(), "ds-1") - - -class TestDataSourceNotionDocumentSyncApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_get_success(self, app: Flask) -> None: - api = DataSourceNotionDocumentSyncApi() - method = inspect.unwrap(api.get) - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.data_source.DatasetService.get_dataset", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.data_source.DocumentService.get_document", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.data_source.document_indexing_sync_task.delay", - return_value=None, - ), - ): - response, status = method(api, MagicMock(), "ds-1", "doc-1") - - assert status == 200 - - def test_get_document_not_found(self, app: Flask) -> None: - api = DataSourceNotionDocumentSyncApi() - method = inspect.unwrap(api.get) - - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.data_source.DatasetService.get_dataset", - return_value=MagicMock(), - ), - patch( - "controllers.console.datasets.data_source.DocumentService.get_document", - return_value=None, - ), - ): - with pytest.raises(NotFound): - method(api, MagicMock(), "ds-1", "doc-1") + assert status == 200 + assert response["notion_info"][0]["pages"][0]["is_bound"] is True diff --git a/api/tests/test_containers_integration_tests/controllers/console/explore/test_conversation.py b/api/tests/test_containers_integration_tests/controllers/console/explore/test_conversation.py index 14e34689ebb..d6efdb9a1c4 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/explore/test_conversation.py +++ b/api/tests/test_containers_integration_tests/controllers/console/explore/test_conversation.py @@ -198,7 +198,13 @@ class TestConversationRenameApi: return_value=conversation, ), ): - result = method(api, user, chat_app, "cid") + result = method( + api, + conversation_module.ConversationRenamePayload.model_validate({"name": "new"}), + user, + chat_app, + "cid", + ) assert result["id"] == "cid" @@ -215,7 +221,13 @@ class TestConversationRenameApi: ), ): with pytest.raises(NotFound): - method(api, user, chat_app, "cid") + method( + api, + conversation_module.ConversationRenamePayload.model_validate({"name": "new"}), + user, + chat_app, + "cid", + ) class TestConversationPinApi: diff --git a/api/tests/test_containers_integration_tests/controllers/console/test_feature.py b/api/tests/test_containers_integration_tests/controllers/console/test_feature.py index 9eb76c81520..e0e2a46ffc5 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/test_feature.py +++ b/api/tests/test_containers_integration_tests/controllers/console/test_feature.py @@ -8,7 +8,8 @@ from unittest.mock import patch from flask.testing import FlaskClient from sqlalchemy.orm import Session -from services.feature_service import FeatureModel, FeatureService, LimitationModel +from services.entities.feature_entities import FeatureModel, LimitationModel +from services.feature_service import FeatureService from tests.test_containers_integration_tests.controllers.console.helpers import ( authenticate_console_client, create_console_account_and_tenant, diff --git a/api/tests/test_containers_integration_tests/controllers/console/test_files.py b/api/tests/test_containers_integration_tests/controllers/console/test_files.py index 8985c1ba66a..e8f1f3ec778 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/test_files.py +++ b/api/tests/test_containers_integration_tests/controllers/console/test_files.py @@ -32,11 +32,13 @@ def test_file_upload_config_returns_console_limits( assert response.status_code == 200 assert response.json == { "file_size_limit": dify_config.UPLOAD_FILE_SIZE_LIMIT, + "knowledge_file_size_limit": dify_config.UPLOAD_FILE_SIZE_LIMIT, "batch_count_limit": dify_config.UPLOAD_FILE_BATCH_LIMIT, "file_upload_limit": dify_config.BATCH_UPLOAD_LIMIT, "image_file_size_limit": dify_config.UPLOAD_IMAGE_FILE_SIZE_LIMIT, "video_file_size_limit": dify_config.UPLOAD_VIDEO_FILE_SIZE_LIMIT, "audio_file_size_limit": dify_config.UPLOAD_AUDIO_FILE_SIZE_LIMIT, + "skill_file_size_limit": dify_config.UPLOAD_SKILL_FILE_SIZE_LIMIT, "workflow_file_upload_limit": dify_config.WORKFLOW_FILE_UPLOAD_LIMIT, "image_file_batch_limit": dify_config.IMAGE_FILE_BATCH_LIMIT, "single_chunk_attachment_limit": dify_config.SINGLE_CHUNK_ATTACHMENT_LIMIT, diff --git a/api/tests/test_containers_integration_tests/controllers/console/test_setup.py b/api/tests/test_containers_integration_tests/controllers/console/test_setup.py new file mode 100644 index 00000000000..248b662fd0b --- /dev/null +++ b/api/tests/test_containers_integration_tests/controllers/console/test_setup.py @@ -0,0 +1,159 @@ +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor +from contextlib import AbstractContextManager +from threading import Event, Lock +from unittest.mock import MagicMock, patch + +import pytest +from faker import Faker +from flask import Flask +from flask.testing import FlaskClient +from sqlalchemy import func, select +from sqlalchemy.orm import Session + +from models.account import Account, Tenant, TenantAccountJoin +from models.model import DifySetup +from services.account_service import RegisterService +from services.setup_adapters import RedisSetupLock +from tests.test_containers_integration_tests.helpers import generate_valid_password + + +@pytest.fixture +def setup_dependencies() -> Iterator[MagicMock]: + with ( + patch("services.account_service.FeatureService") as feature_service, + patch("services.account_service.BillingService") as billing_service, + patch("services.account_service.CommunityTelemetryService.report_install") as report_install, + ): + feature_service.get_system_features.return_value.is_allow_register = True + feature_service.get_license.return_value.seats.is_available.return_value = True + feature_service.get_license.return_value.workspaces.is_available.return_value = True + feature_service.is_workspace_creation_allowed.return_value = True + billing_service.is_email_in_freeze.return_value = False + yield report_install + + +def _setup_payload(*, email: str, password: str) -> dict[str, str]: + return { + "email": email, + "name": "Admin", + "password": password, + "language": "en-US", + } + + +def test_setup_endpoint_persists_bootstrap_state_and_rejects_repeat( + db_session_with_containers: Session, + test_client_with_containers: FlaskClient, + monkeypatch: pytest.MonkeyPatch, + setup_dependencies: MagicMock, +) -> None: + monkeypatch.delenv("INIT_PASSWORD", raising=False) + password = generate_valid_password(Faker()) + + response = test_client_with_containers.post( + "/console/api/setup", + json=_setup_payload(email="admin@example.com", password=password), + headers={"CF-Connecting-IP": "203.0.113.7"}, + ) + + assert response.status_code == 201 + assert response.get_json() == {"result": "success"} + setup_dependencies.assert_called_once() + + repeated_response = test_client_with_containers.post( + "/console/api/setup", + json=_setup_payload(email="other@example.com", password=password), + ) + + assert repeated_response.status_code == 403 + + db_session_with_containers.expire_all() + assert db_session_with_containers.scalar(select(func.count()).select_from(DifySetup)) == 1 + assert db_session_with_containers.scalar(select(func.count()).select_from(Account)) == 1 + assert db_session_with_containers.scalar(select(func.count()).select_from(Tenant)) == 1 + assert db_session_with_containers.scalar(select(func.count()).select_from(TenantAccountJoin)) == 1 + + account = db_session_with_containers.scalar(select(Account)) + assert account is not None + assert account.email == "admin@example.com" + assert account.last_login_ip == "203.0.113.7" + + +def test_concurrent_setup_requests_create_only_one_bootstrap_identity( + flask_app_with_containers: Flask, + db_session_with_containers: Session, + monkeypatch: pytest.MonkeyPatch, + setup_dependencies: MagicMock, +) -> None: + monkeypatch.delenv("INIT_PASSWORD", raising=False) + password = generate_valid_password(Faker()) + provision_started = Event() + allow_provision_to_finish = Event() + second_lock_attempted = Event() + attempt_guard = Lock() + lock_attempts = 0 + original_setup = RegisterService.setup + original_acquire = RedisSetupLock.acquire + + def blocking_setup( + email: str, + name: str, + password: str, + ip_address: str, + language: str | None, + *, + session: Session, + ) -> None: + if not provision_started.is_set(): + provision_started.set() + assert allow_provision_to_finish.wait(timeout=10) + original_setup( + email=email, + name=name, + password=password, + ip_address=ip_address, + language=language, + session=session, + ) + + def tracked_acquire(self: RedisSetupLock) -> AbstractContextManager[None]: + nonlocal lock_attempts + with attempt_guard: + lock_attempts += 1 + if lock_attempts == 2: + second_lock_attempted.set() + return original_acquire(self) + + def post_setup(email: str) -> int: + with flask_app_with_containers.test_client() as client: + response = client.post( + "/console/api/setup", + json=_setup_payload(email=email, password=password), + ) + return response.status_code + + with ( + patch.object(RegisterService, "setup", side_effect=blocking_setup), + patch.object(RedisSetupLock, "acquire", tracked_acquire), + ThreadPoolExecutor(max_workers=2) as executor, + ): + first_request = executor.submit(post_setup, "admin-1@example.com") + try: + assert provision_started.wait(timeout=10) + second_request = executor.submit(post_setup, "admin-2@example.com") + assert second_lock_attempted.wait(timeout=10) + assert not second_request.done() + finally: + allow_provision_to_finish.set() + + results = [first_request.result(timeout=30), second_request.result(timeout=30)] + + assert sorted(results) == [201, 403] + setup_dependencies.assert_called_once() + + db_session_with_containers.expire_all() + assert db_session_with_containers.scalar(select(func.count()).select_from(DifySetup)) == 1 + assert db_session_with_containers.scalar(select(func.count()).select_from(Account)) == 1 + assert db_session_with_containers.scalar(select(func.count()).select_from(Tenant)) == 1 + assert db_session_with_containers.scalar(select(func.count()).select_from(TenantAccountJoin)) == 1 diff --git a/api/tests/test_containers_integration_tests/controllers/console/test_spec.py b/api/tests/test_containers_integration_tests/controllers/console/test_spec.py new file mode 100644 index 00000000000..798cf852ebd --- /dev/null +++ b/api/tests/test_containers_integration_tests/controllers/console/test_spec.py @@ -0,0 +1,37 @@ +from flask.testing import FlaskClient +from sqlalchemy.orm import Session + +from tests.test_containers_integration_tests.controllers.console.helpers import ( + authenticate_console_client, + create_console_account_and_tenant, +) + + +def test_schema_definitions_endpoint_uses_admission_and_builtin_registry( + db_session_with_containers: Session, + test_client_with_containers: FlaskClient, +) -> None: + account, _ = create_console_account_and_tenant(db_session_with_containers) + headers = authenticate_console_client(test_client_with_containers, account) + + response = test_client_with_containers.get( + "/console/api/spec/schema-definitions", + headers=headers, + ) + + assert response.status_code == 200 + definitions = response.get_json() + assert isinstance(definitions, list) + assert definitions + assert all({"name", "label", "schema"} <= definition.keys() for definition in definitions) + + +def test_schema_definitions_endpoint_rejects_unauthenticated_request( + db_session_with_containers: Session, + test_client_with_containers: FlaskClient, +) -> None: + create_console_account_and_tenant(db_session_with_containers) + + response = test_client_with_containers.get("/console/api/spec/schema-definitions") + + assert response.status_code == 401 diff --git a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py index e58bed04270..d72b6b736cc 100644 --- a/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py +++ b/api/tests/test_containers_integration_tests/controllers/service_api/dataset/test_dataset.py @@ -1,1020 +1,24 @@ -""" -Integration tests for Service API Dataset controllers. +"""Container integration coverage for Service API dataset controllers.""" -Migrated from unit_tests/controllers/service_api/dataset/test_dataset.py. - -Tests coverage for: -- DatasetCreatePayload, DatasetUpdatePayload Pydantic models -- Tag-related payloads (create, update, delete, binding) -- DatasetListQuery model -- API endpoint error handling and controller behavior - -Services (DatasetService, TagService, DocumentService) remain mocked -since these test controller-level behavior. -""" - -import uuid -from contextlib import ExitStack -from datetime import UTC, datetime -from unittest.mock import ANY, Mock, patch +from unittest.mock import patch import pytest from flask import Flask -from sqlalchemy.orm import Session, scoped_session -from werkzeug.exceptions import Forbidden, NotFound +from sqlalchemy.orm import Session - -class SessionMatcher: - def __eq__(self, other): - return isinstance(other, (Session, scoped_session)) - - -import services -from controllers.service_api.dataset.dataset import ( - DatasetCreatePayload, - DatasetListQuery, - DatasetUpdatePayload, - TagBindingPayload, - TagCreatePayload, - TagDeletePayload, - TagUnbindingPayload, - TagUpdatePayload, -) -from controllers.service_api.dataset.error import DatasetInUseError, DatasetNameDuplicateError, InvalidActionError from models.account import Account -from models.dataset import Dataset, DatasetPermissionEnum from models.enums import TagType from models.model import Tag - -# --------------------------------------------------------------------------- -# Pydantic model validation tests -# --------------------------------------------------------------------------- - - -class TestDatasetCreatePayload: - """Test suite for DatasetCreatePayload Pydantic model.""" - - def test_payload_with_required_name(self): - payload = DatasetCreatePayload(name="Test Dataset") - assert payload.name == "Test Dataset" - assert payload.description == "" - assert payload.permission == DatasetPermissionEnum.ONLY_ME - - def test_payload_with_all_fields(self): - payload = DatasetCreatePayload( - name="Full Dataset", - description="A comprehensive dataset description", - indexing_technique="high_quality", - permission=DatasetPermissionEnum.ALL_TEAM, - provider="vendor", - embedding_model="text-embedding-ada-002", - embedding_model_provider="openai", - ) - assert payload.name == "Full Dataset" - assert payload.description == "A comprehensive dataset description" - assert payload.indexing_technique == "high_quality" - assert payload.permission == DatasetPermissionEnum.ALL_TEAM - assert payload.provider == "vendor" - assert payload.embedding_model == "text-embedding-ada-002" - assert payload.embedding_model_provider == "openai" - - def test_payload_name_length_validation_min(self): - with pytest.raises(ValueError): - DatasetCreatePayload(name="") - - def test_payload_name_length_validation_max(self): - with pytest.raises(ValueError): - DatasetCreatePayload(name="A" * 41) - - def test_payload_description_max_length(self): - with pytest.raises(ValueError): - DatasetCreatePayload(name="Dataset", description="A" * 401) - - @pytest.mark.parametrize("technique", ["high_quality", "economy"]) - def test_payload_valid_indexing_techniques(self, technique): - payload = DatasetCreatePayload(name="Dataset", indexing_technique=technique) - assert payload.indexing_technique == technique - - def test_payload_with_external_knowledge_settings(self): - payload = DatasetCreatePayload( - name="External Dataset", external_knowledge_api_id="api_123", external_knowledge_id="knowledge_456" - ) - assert payload.external_knowledge_api_id == "api_123" - assert payload.external_knowledge_id == "knowledge_456" - - -class TestDatasetUpdatePayload: - """Test suite for DatasetUpdatePayload Pydantic model.""" - - def test_payload_all_optional(self): - payload = DatasetUpdatePayload() - assert payload.name is None - assert payload.description is None - assert payload.permission is None - - def test_payload_with_partial_update(self): - payload = DatasetUpdatePayload(name="Updated Name", description="Updated description") - assert payload.name == "Updated Name" - assert payload.description == "Updated description" - - def test_payload_with_permission_change(self): - payload = DatasetUpdatePayload( - permission=DatasetPermissionEnum.PARTIAL_TEAM, - partial_member_list=[{"user_id": "user_123", "role": "editor"}], - ) - assert payload.permission == DatasetPermissionEnum.PARTIAL_TEAM - assert payload.partial_member_list is not None - assert len(payload.partial_member_list) == 1 - - def test_payload_name_length_validation(self): - with pytest.raises(ValueError): - DatasetUpdatePayload(name="") - with pytest.raises(ValueError): - DatasetUpdatePayload(name="A" * 41) - - -class TestDatasetListQuery: - """Test suite for DatasetListQuery Pydantic model.""" - - def test_query_with_defaults(self): - query = DatasetListQuery() - assert query.page == 1 - assert query.limit == 20 - assert query.keyword is None - assert query.include_all is False - assert query.tag_ids == [] - - def test_query_with_all_filters(self): - query = DatasetListQuery( - page=3, limit=50, keyword="machine learning", include_all=True, tag_ids=["tag1", "tag2", "tag3"] - ) - assert query.page == 3 - assert query.limit == 50 - assert query.keyword == "machine learning" - assert query.include_all is True - assert len(query.tag_ids) == 3 - - def test_query_with_tag_filter(self): - query = DatasetListQuery(tag_ids=["tag_abc", "tag_def"]) - assert query.tag_ids == ["tag_abc", "tag_def"] - - -class TestTagCreatePayload: - """Test suite for TagCreatePayload Pydantic model.""" - - def test_payload_with_name(self): - payload = TagCreatePayload(name="New Tag") - assert payload.name == "New Tag" - - def test_payload_name_length_min(self): - with pytest.raises(ValueError): - TagCreatePayload(name="") - - def test_payload_name_length_max(self): - with pytest.raises(ValueError): - TagCreatePayload(name="A" * 51) - - def test_payload_with_unicode_name(self): - payload = TagCreatePayload(name="标签 🏷️ Тег") - assert payload.name == "标签 🏷️ Тег" - - -class TestTagUpdatePayload: - """Test suite for TagUpdatePayload Pydantic model.""" - - def test_payload_with_name_and_id(self): - payload = TagUpdatePayload(name="Updated Tag", tag_id="tag_123") - assert payload.name == "Updated Tag" - assert payload.tag_id == "tag_123" - - def test_payload_requires_tag_id(self): - with pytest.raises(ValueError): - TagUpdatePayload.model_validate({"name": "Updated Tag"}) - - -class TestTagDeletePayload: - """Test suite for TagDeletePayload Pydantic model.""" - - def test_payload_with_tag_id(self): - payload = TagDeletePayload(tag_id="tag_to_delete") - assert payload.tag_id == "tag_to_delete" - - def test_payload_requires_tag_id(self): - with pytest.raises(ValueError): - TagDeletePayload.model_validate({}) - - -class TestTagBindingPayload: - """Test suite for TagBindingPayload Pydantic model.""" - - def test_payload_with_valid_data(self): - payload = TagBindingPayload(tag_ids=["tag1", "tag2"], target_id="dataset_123") - assert len(payload.tag_ids) == 2 - assert payload.target_id == "dataset_123" - - def test_payload_rejects_empty_tag_ids(self): - with pytest.raises(ValueError) as exc_info: - TagBindingPayload(tag_ids=[], target_id="dataset_123") - assert "Tag IDs is required" in str(exc_info.value) - - def test_payload_single_tag_id(self): - payload = TagBindingPayload(tag_ids=["single_tag"], target_id="dataset_456") - assert payload.tag_ids == ["single_tag"] - - -class TestTagUnbindingPayload: - """Test suite for TagUnbindingPayload Pydantic model.""" - - def test_payload_with_valid_data(self): - payload = TagUnbindingPayload(tag_ids=["tag_123"], target_id="dataset_456") - assert payload.tag_ids == ["tag_123"] - assert payload.target_id == "dataset_456" - - def test_payload_normalizes_legacy_tag_id(self): - payload = TagUnbindingPayload(tag_id="tag_123", target_id="dataset_456") - assert payload.tag_ids == ["tag_123"] - assert payload.target_id == "dataset_456" - - def test_payload_rejects_empty_tag_ids(self): - with pytest.raises(ValueError) as exc_info: - TagUnbindingPayload(tag_ids=[], target_id="dataset_456") - assert "Tag IDs is required" in str(exc_info.value) - - -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - -from inspect import unwrap +from tests.test_containers_integration_tests.controllers.console.helpers import create_console_account_and_tenant @pytest.fixture -def app(flask_app_with_containers: Flask): - # Uses the full containerised app so that Flask config, extensions, and - # blueprint registrations match production. Most tests mock the service - # layer to isolate controller logic; a few (e.g. test_list_tags_from_db) - # exercise the real DB-backed path to validate end-to-end behaviour. +def app(flask_app_with_containers: Flask) -> Flask: return flask_app_with_containers -@pytest.fixture -def mock_tenant(): - tenant = Mock() - tenant.id = str(uuid.uuid4()) - return tenant - - -@pytest.fixture -def mock_dataset(): - return make_dataset(id=str(uuid.uuid4()), tenant_id=str(uuid.uuid4())) - - -@pytest.fixture(autouse=True) -def dataset_model_getter_defaults(): - getters: dict[str, object] = { - "get_app_count": 0, - "get_document_count": 0, - "get_word_count": 0, - "get_author_name": None, - "get_tags": [], - "get_doc_form": None, - "get_external_knowledge_info": None, - "get_doc_metadata": [], - "get_is_published": False, - "get_total_documents": 0, - "get_total_available_documents": 0, - } - - with ExitStack() as stack: - for name, value in getters.items(): - getter_mock = stack.enter_context(patch.object(Dataset, name, autospec=True)) - getter_mock.return_value = value - yield - - -def make_dataset(**overrides) -> Dataset: - base = { - "id": "ds-1", - "tenant_id": "tenant-1", - "name": "Dataset", - "description": "desc", - "provider": "vendor", - "permission": "only_me", - "data_source_type": None, - "indexing_technique": "economy", - "created_by": "account-1", - "created_at": datetime(2024, 1, 1, 12, 0, 0, tzinfo=UTC), - "updated_by": None, - "updated_at": datetime(2024, 1, 1, 12, 0, 0, tzinfo=UTC), - "embedding_model": None, - "embedding_model_provider": None, - "retrieval_model": None, - "summary_index_setting": None, - "built_in_field_enabled": False, - "pipeline_id": None, - "runtime_mode": "general", - "chunk_structure": None, - "icon_info": None, - "enable_api": False, - "is_multimodal": False, - } - base.update(overrides) - return Dataset(**base) - - -def make_tag(*, id: str, name: str, binding_count: int | None = None) -> Tag: - tag = Tag(tenant_id="tenant-1", type=TagType.KNOWLEDGE, name=name, created_by="account-1") - tag.id = id - if binding_count is not None: - tag.__dict__["binding_count"] = binding_count - return tag - - -DATASET_DETAIL_KEYS = { - "id", - "name", - "description", - "provider", - "permission", - "data_source_type", - "indexing_technique", - "app_count", - "document_count", - "word_count", - "created_by", - "author_name", - "created_at", - "updated_by", - "updated_at", - "embedding_model", - "embedding_model_provider", - "embedding_available", - "retrieval_model_dict", - "summary_index_setting", - "tags", - "doc_form", - "external_knowledge_info", - "external_retrieval_model", - "doc_metadata", - "built_in_field_enabled", - "pipeline_id", - "runtime_mode", - "chunk_structure", - "icon_info", - "is_published", - "total_documents", - "total_available_documents", - "enable_api", - "is_multimodal", - "maintainer", -} - - -def assert_dataset_detail_shape(response: dict, *, with_partial_members: bool = False) -> None: - expected_keys = set(DATASET_DETAIL_KEYS) - if with_partial_members: - expected_keys.add("partial_member_list") - assert set(response) == expected_keys - assert isinstance(response["created_at"], int) - assert isinstance(response["updated_at"], int) - assert set(response["retrieval_model_dict"]) == { - "search_method", - "reranking_enable", - "reranking_mode", - "reranking_model", - "weights", - "top_k", - "score_threshold_enabled", - "score_threshold", - } - if response["external_retrieval_model"] is not None: - assert set(response["external_retrieval_model"]) == { - "top_k", - "score_threshold", - "score_threshold_enabled", - } - if not with_partial_members: - assert "partial_member_list" not in response - - -# --------------------------------------------------------------------------- -# API endpoint tests — DatasetListApi -# --------------------------------------------------------------------------- - - -class TestDatasetListApiGet: - """Test suite for DatasetListApi.get() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_list_datasets_success( - self, - mock_dataset_svc, - mock_current_user, - mock_provider_mgr, - app: Flask, - mock_tenant, - ): - from controllers.service_api.dataset.dataset import DatasetListApi - - mock_current_user.__class__ = Account - mock_current_user.current_tenant_id = mock_tenant.id - mock_dataset_svc.get_datasets.return_value = ([make_dataset()], 1) - - mock_configs = Mock() - mock_configs.get_models.return_value = [] - mock_provider_mgr.return_value.get_configurations.return_value = mock_configs - - with app.test_request_context("/datasets?page=1&limit=20", method="GET"): - api = DatasetListApi() - response, status = api.get(tenant_id=mock_tenant.id) - - assert status == 200 - assert set(response) == {"data", "has_more", "limit", "total", "page"} - assert response["has_more"] is False - assert response["limit"] == 20 - assert response["total"] == 1 - assert response["page"] == 1 - assert len(response["data"]) == 1 - assert_dataset_detail_shape(response["data"][0]) - - @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_list_datasets_preserves_repeated_tag_ids( - self, - mock_dataset_svc, - mock_current_user, - mock_provider_mgr, - app: Flask, - mock_tenant, - ): - from controllers.service_api.dataset.dataset import DatasetListApi - - mock_current_user.__class__ = Account - mock_current_user.current_tenant_id = mock_tenant.id - mock_dataset_svc.get_datasets.return_value = ([make_dataset()], 1) - - mock_configs = Mock() - mock_configs.get_models.return_value = [] - mock_provider_mgr.return_value.get_configurations.return_value = mock_configs - - with app.test_request_context("/datasets?tag_ids=tag-a&tag_ids=tag-b", method="GET"): - api = DatasetListApi() - response, status = api.get(tenant_id=mock_tenant.id) - - assert status == 200 - assert response["total"] == 1 - mock_dataset_svc.get_datasets.assert_called_once_with( - 1, - 20, - SessionMatcher(), - mock_tenant.id, - mock_current_user, - None, - ["tag-a", "tag-b"], - False, - ) - - -class TestDatasetListApiPost: - """Test suite for DatasetListApi.post() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_create_dataset_success( - self, - mock_dataset_svc, - mock_current_user, - app: Flask, - mock_tenant, - ): - from controllers.service_api.dataset.dataset import DatasetListApi - - mock_current_user.__class__ = Account - mock_dataset_svc.create_empty_dataset.return_value = make_dataset(name="New Dataset") - - with app.test_request_context( - "/datasets", - method="POST", - json={"name": "New Dataset"}, - ): - api = DatasetListApi() - response, status = unwrap(api.post)(api, Mock(spec=Session), tenant_id=mock_tenant.id) - - assert status == 200 - assert_dataset_detail_shape(response) - assert response["name"] == "New Dataset" - mock_dataset_svc.create_empty_dataset.assert_called_once() - - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_create_dataset_duplicate_name( - self, - mock_dataset_svc, - mock_current_user, - app: Flask, - mock_tenant, - ): - from controllers.service_api.dataset.dataset import DatasetListApi - - mock_current_user.__class__ = Account - mock_dataset_svc.create_empty_dataset.side_effect = services.errors.dataset.DatasetNameDuplicateError() - - with app.test_request_context( - "/datasets", - method="POST", - json={"name": "Existing Dataset"}, - ): - api = DatasetListApi() - with pytest.raises(DatasetNameDuplicateError): - unwrap(api.post)(api, Mock(spec=Session), tenant_id=mock_tenant.id) - - -# --------------------------------------------------------------------------- -# API endpoint tests — DatasetApi -# --------------------------------------------------------------------------- - - -class TestDatasetApiGet: - """Test suite for DatasetApi.get() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.DatasetPermissionService") - @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_get_dataset_success( - self, - mock_dataset_svc, - mock_current_user, - mock_provider_mgr, - mock_perm_svc, - app: Flask, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DatasetApi - - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.check_dataset_permission.return_value = None - mock_current_user.__class__ = Account - mock_current_user.current_tenant_id = mock_dataset.tenant_id - - mock_configs = Mock() - mock_configs.get_models.return_value = [] - mock_provider_mgr.return_value.get_configurations.return_value = mock_configs - - with app.test_request_context( - f"/datasets/{mock_dataset.id}", - method="GET", - ): - api = DatasetApi() - response, status = api.get(_=mock_dataset.tenant_id, dataset_id=mock_dataset.id) - - assert status == 200 - assert_dataset_detail_shape(response) - assert response["embedding_available"] is True - assert response["retrieval_model_dict"]["search_method"] == "keyword_search" - - @patch("controllers.service_api.dataset.dataset.DatasetPermissionService") - @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_get_dataset_partial_members_shape( - self, - mock_dataset_svc, - mock_current_user, - mock_provider_mgr, - mock_perm_svc, - app: Flask, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DatasetApi - - mock_dataset.permission = "partial_members" - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.check_dataset_permission.return_value = None - mock_current_user.__class__ = Account - mock_current_user.current_tenant_id = mock_dataset.tenant_id - mock_perm_svc.get_dataset_partial_member_list.return_value = ["user-1", "user-2"] - - mock_configs = Mock() - mock_configs.get_models.return_value = [] - mock_provider_mgr.return_value.get_configurations.return_value = mock_configs - - with app.test_request_context( - f"/datasets/{mock_dataset.id}", - method="GET", - ): - api = DatasetApi() - response, status = api.get(_=mock_dataset.tenant_id, dataset_id=mock_dataset.id) - - assert status == 200 - assert_dataset_detail_shape(response, with_partial_members=True) - assert response["partial_member_list"] == ["user-1", "user-2"] - - @patch("controllers.service_api.dataset.dataset.DatasetPermissionService") - @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_get_dataset_uses_default_external_retrieval_model( - self, - mock_dataset_svc, - mock_current_user, - mock_provider_mgr, - mock_perm_svc, - app: Flask, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DatasetApi - - mock_dataset.retrieval_model = None - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.check_dataset_permission.return_value = None - mock_current_user.__class__ = Account - mock_current_user.current_tenant_id = mock_dataset.tenant_id - - mock_configs = Mock() - mock_configs.get_models.return_value = [] - mock_provider_mgr.return_value.get_configurations.return_value = mock_configs - - with app.test_request_context(f"/datasets/{mock_dataset.id}", method="GET"): - api = DatasetApi() - response, status = api.get(_=mock_dataset.tenant_id, dataset_id=mock_dataset.id) - - assert status == 200 - assert_dataset_detail_shape(response) - assert response["external_retrieval_model"] == { - "top_k": 2, - "score_threshold": 0.0, - "score_threshold_enabled": None, - } - - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_get_dataset_not_found(self, mock_dataset_svc, app, mock_dataset): - from controllers.service_api.dataset.dataset import DatasetApi - - mock_dataset_svc.get_dataset.return_value = None - - with app.test_request_context( - f"/datasets/{mock_dataset.id}", - method="GET", - ): - api = DatasetApi() - with pytest.raises(NotFound): - api.get(_=mock_dataset.tenant_id, dataset_id=mock_dataset.id) - - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_get_dataset_no_permission( - self, - mock_dataset_svc, - mock_current_user, - app: Flask, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DatasetApi - - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.check_dataset_permission.side_effect = services.errors.account.NoPermissionError() - - with app.test_request_context( - f"/datasets/{mock_dataset.id}", - method="GET", - ): - api = DatasetApi() - with pytest.raises(Forbidden): - api.get(_=mock_dataset.tenant_id, dataset_id=mock_dataset.id) - - -class TestDatasetApiPatch: - """Test suite for DatasetApi.patch() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.DatasetPermissionService") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_patch_dataset_success_shape( - self, - mock_dataset_svc, - mock_current_user, - mock_perm_svc, - app: Flask, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DatasetApi - - updated_dataset = make_dataset(id=mock_dataset.id, tenant_id=mock_dataset.tenant_id, name="Updated Dataset") - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.update_dataset.return_value = updated_dataset - mock_perm_svc.check_permission.return_value = None - mock_perm_svc.get_dataset_partial_member_list.return_value = ["user-1"] - mock_current_user.__class__ = Account - mock_current_user.current_tenant_id = mock_dataset.tenant_id - - payload = { - "name": "Updated Dataset", - "permission": "partial_members", - "partial_member_list": [{"user_id": "user-1", "role": "editor"}], - } - with app.test_request_context( - f"/datasets/{mock_dataset.id}", - method="PATCH", - json=payload, - ): - api = DatasetApi() - response, status = unwrap(api.patch)( - api, - Mock(spec=Session), - _=mock_dataset.tenant_id, - dataset_id=mock_dataset.id, - ) - - assert status == 200 - assert_dataset_detail_shape(response, with_partial_members=True) - assert response["name"] == "Updated Dataset" - assert response["partial_member_list"] == ["user-1"] - mock_dataset_svc.update_dataset.assert_called_once() - _, update_data, _ = mock_dataset_svc.update_dataset.call_args.args - session = mock_dataset_svc.update_dataset.call_args.kwargs["session"] - assert isinstance(session, (Session, scoped_session)) - assert update_data["name"] == "Updated Dataset" - assert update_data["permission"] == "partial_members" - mock_perm_svc.update_partial_member_list.assert_called_once_with( - mock_dataset.tenant_id, - mock_dataset.id, - [{"user_id": "user-1", "role": "editor"}], - SessionMatcher(), - ) - - -class TestDatasetApiDelete: - """Test suite for DatasetApi.delete() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.DatasetPermissionService") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_delete_dataset_success( - self, - mock_dataset_svc, - mock_current_user, - mock_perm_svc, - app: Flask, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DatasetApi - - mock_dataset_svc.delete_dataset.return_value = True - - with app.test_request_context( - f"/datasets/{mock_dataset.id}", - method="DELETE", - ): - api = DatasetApi() - result = unwrap(api.delete)(api, Mock(), _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) - - assert result == ("", 204) - - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_delete_dataset_not_found( - self, - mock_dataset_svc, - mock_current_user, - app: Flask, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DatasetApi - - mock_dataset_svc.delete_dataset.return_value = False - - with app.test_request_context( - f"/datasets/{mock_dataset.id}", - method="DELETE", - ): - api = DatasetApi() - with pytest.raises(NotFound): - unwrap(api.delete)(api, Mock(), _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) - - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_delete_dataset_in_use( - self, - mock_dataset_svc, - mock_current_user, - app: Flask, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DatasetApi - - mock_dataset_svc.delete_dataset.side_effect = services.errors.dataset.DatasetInUseError() - - with app.test_request_context( - f"/datasets/{mock_dataset.id}", - method="DELETE", - ): - api = DatasetApi() - with pytest.raises(DatasetInUseError): - unwrap(api.delete)(api, Mock(), _=mock_dataset.tenant_id, dataset_id=mock_dataset.id) - - -# --------------------------------------------------------------------------- -# API endpoint tests — DocumentStatusApi -# --------------------------------------------------------------------------- - - -class TestDocumentStatusApiPatch: - """Test suite for DocumentStatusApi.patch() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.DocumentService") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_batch_update_status_success( - self, - mock_dataset_svc, - mock_current_user, - mock_doc_svc, - app: Flask, - mock_tenant, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DocumentStatusApi - - mock_current_user.__class__ = Account - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.check_dataset_permission.return_value = None - mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc_svc.batch_update_document_status.return_value = None - - with app.test_request_context( - f"/datasets/{mock_dataset.id}/documents/status/enable", - method="PATCH", - json={"document_ids": ["doc-1", "doc-2"]}, - ): - api = DocumentStatusApi() - response, status = api.patch( - tenant_id=mock_tenant.id, - dataset_id=mock_dataset.id, - action="enable", - ) - - assert status == 200 - assert response["result"] == "success" - - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_batch_update_status_dataset_not_found( - self, - mock_dataset_svc, - app: Flask, - mock_tenant, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DocumentStatusApi - - mock_dataset_svc.get_dataset.return_value = None - - with app.test_request_context( - f"/datasets/{mock_dataset.id}/documents/status/enable", - method="PATCH", - json={"document_ids": ["doc-1"]}, - ): - api = DocumentStatusApi() - with pytest.raises(NotFound): - api.patch( - tenant_id=mock_tenant.id, - dataset_id=mock_dataset.id, - action="enable", - ) - - @patch("controllers.service_api.dataset.dataset.DocumentService") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_batch_update_status_permission_error( - self, - mock_dataset_svc, - mock_current_user, - mock_doc_svc, - app: Flask, - mock_tenant, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DocumentStatusApi - - mock_current_user.__class__ = Account - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.check_dataset_permission.side_effect = services.errors.account.NoPermissionError( - "No permission" - ) - - with app.test_request_context( - f"/datasets/{mock_dataset.id}/documents/status/enable", - method="PATCH", - json={"document_ids": ["doc-1"]}, - ): - api = DocumentStatusApi() - with pytest.raises(Forbidden): - api.patch( - tenant_id=mock_tenant.id, - dataset_id=mock_dataset.id, - action="enable", - ) - - @patch("controllers.service_api.dataset.dataset.DocumentService") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_batch_update_status_indexing_error( - self, - mock_dataset_svc, - mock_current_user, - mock_doc_svc, - app: Flask, - mock_tenant, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DocumentStatusApi - - mock_current_user.__class__ = Account - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.check_dataset_permission.return_value = None - mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc_svc.batch_update_document_status.side_effect = services.errors.document.DocumentIndexingError() - - with app.test_request_context( - f"/datasets/{mock_dataset.id}/documents/status/enable", - method="PATCH", - json={"document_ids": ["doc-1"]}, - ): - api = DocumentStatusApi() - with pytest.raises(InvalidActionError): - api.patch( - tenant_id=mock_tenant.id, - dataset_id=mock_dataset.id, - action="enable", - ) - - @patch("controllers.service_api.dataset.dataset.DocumentService") - @patch("controllers.service_api.dataset.dataset.current_user") - @patch("controllers.service_api.dataset.dataset.DatasetService") - def test_batch_update_status_value_error( - self, - mock_dataset_svc, - mock_current_user, - mock_doc_svc, - app: Flask, - mock_tenant, - mock_dataset, - ): - from controllers.service_api.dataset.dataset import DocumentStatusApi - - mock_current_user.__class__ = Account - mock_dataset_svc.get_dataset.return_value = mock_dataset - mock_dataset_svc.check_dataset_permission.return_value = None - mock_dataset_svc.check_dataset_model_setting.return_value = None - mock_doc_svc.batch_update_document_status.side_effect = ValueError("Invalid action") - - with app.test_request_context( - f"/datasets/{mock_dataset.id}/documents/status/enable", - method="PATCH", - json={"document_ids": ["doc-1"]}, - ): - api = DocumentStatusApi() - with pytest.raises(InvalidActionError): - api.patch( - tenant_id=mock_tenant.id, - dataset_id=mock_dataset.id, - action="enable", - ) - - -# --------------------------------------------------------------------------- -# API endpoint tests — Tags -# --------------------------------------------------------------------------- - - class TestDatasetTagsApiGet: - """Test suite for DatasetTagsApi.get() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.TagService") - @patch("controllers.service_api.dataset.dataset.current_user") - def test_list_tags_success( - self, - mock_current_user, - mock_tag_svc, - app: Flask, - ): - from controllers.service_api.dataset.dataset import DatasetTagsApi - - mock_current_user.__class__ = Account - mock_current_user.current_tenant_id = "tenant-1" - mock_tag = make_tag(id="tag-1", name="Test Tag", binding_count=0) - mock_tag_svc.get_tags.return_value = [mock_tag] - - with app.test_request_context("/datasets/tags", method="GET"): - api = DatasetTagsApi() - response, status = api.get(_=None) - - assert status == 200 - assert response == [{"id": "tag-1", "name": "Test Tag", "type": "knowledge", "binding_count": "0"}] - mock_tag_svc.get_tags.assert_called_once_with("knowledge", "tenant-1", session=SessionMatcher()) + """Exercise the unmocked tag query against the container database.""" @patch("controllers.service_api.dataset.dataset.current_user") def test_list_tags_from_db( @@ -1022,15 +26,8 @@ class TestDatasetTagsApiGet: mock_current_user, app: Flask, db_session_with_containers: Session, - ): - """Integration test: creates real Tag rows and retrieves them - through the controller without mocking TagService.""" - from tests.test_containers_integration_tests.controllers.console.helpers import ( - create_console_account_and_tenant, - ) - + ) -> None: account, tenant = create_console_account_and_tenant(db_session_with_containers) - tag = Tag( name="Integration Tag", type=TagType.KNOWLEDGE, @@ -1046,336 +43,9 @@ class TestDatasetTagsApiGet: from controllers.service_api.dataset.dataset import DatasetTagsApi with app.test_request_context("/datasets/tags", method="GET"): - api = DatasetTagsApi() - response, status = api.get(_=None) + response, status = DatasetTagsApi().get(_=None) assert status == 200 - assert any(t["name"] == "Integration Tag" for t in response) - assert all(set(t) == {"id", "name", "type", "binding_count"} for t in response) - assert all(isinstance(t["binding_count"], str) for t in response) - - -class TestDatasetTagsApiPost: - """Test suite for DatasetTagsApi.post() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.TagService") - @patch("controllers.service_api.dataset.dataset.current_user") - def test_create_tag_success( - self, - mock_current_user, - mock_tag_svc, - app: Flask, - ): - from controllers.service_api.dataset.dataset import DatasetTagsApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = True - mock_current_user.is_dataset_editor = True - mock_tag = make_tag(id="tag-new", name="New Tag") - mock_tag_svc.save_tags.return_value = mock_tag - - with app.test_request_context( - "/datasets/tags", - method="POST", - json={"name": "New Tag"}, - ): - api = DatasetTagsApi() - response, status = api.post(_=None) - - assert status == 200 - assert response == {"id": "tag-new", "name": "New Tag", "type": "knowledge", "binding_count": "0"} - mock_tag_svc.save_tags.assert_called_once() - - @patch("controllers.service_api.dataset.dataset.current_user") - def test_create_tag_forbidden(self, mock_current_user, app: Flask): - from controllers.service_api.dataset.dataset import DatasetTagsApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = False - mock_current_user.is_dataset_editor = False - - with app.test_request_context( - "/datasets/tags", - method="POST", - json={"name": "New Tag"}, - ): - api = DatasetTagsApi() - with pytest.raises(Forbidden): - api.post(_=None) - - -class TestDatasetTagsApiPatch: - """Test suite for DatasetTagsApi.patch() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.TagService") - @patch("controllers.service_api.dataset.dataset.service_api_ns") - @patch("controllers.service_api.dataset.dataset.current_user") - def test_update_tag_success( - self, - mock_current_user, - mock_service_api_ns, - mock_tag_svc, - app: Flask, - ): - from controllers.service_api.dataset.dataset import DatasetTagsApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = True - mock_current_user.is_dataset_editor = True - - mock_tag = make_tag(id="tag-1", name="Updated Tag") - mock_tag_svc.update_tags.return_value = mock_tag - mock_tag_svc.get_tag_binding_count.return_value = 5 - mock_service_api_ns.payload = {"name": "Updated Tag", "tag_id": "tag-1"} - - with app.test_request_context( - "/datasets/tags", - method="PATCH", - json={"name": "Updated Tag", "tag_id": "tag-1"}, - ): - api = DatasetTagsApi() - response, status = api.patch(_=None) - - assert status == 200 - assert response == {"id": "tag-1", "name": "Updated Tag", "type": "knowledge", "binding_count": "5"} - mock_tag_svc.update_tags.assert_called_once() - update_payload, tag_id, session = mock_tag_svc.update_tags.call_args.args - assert update_payload.name == "Updated Tag" - assert tag_id == "tag-1" - - @patch("controllers.service_api.dataset.dataset.current_user") - def test_update_tag_forbidden(self, mock_current_user, app: Flask): - from controllers.service_api.dataset.dataset import DatasetTagsApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = False - mock_current_user.is_dataset_editor = False - - with app.test_request_context( - "/datasets/tags", - method="PATCH", - json={"name": "Updated Tag", "tag_id": "tag-1"}, - ): - api = DatasetTagsApi() - with pytest.raises(Forbidden): - api.patch(_=None) - - -class TestDatasetTagsApiDelete: - """Test suite for DatasetTagsApi.delete() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.TagService") - @patch("controllers.service_api.dataset.dataset.service_api_ns") - @patch("libs.login.current_user") - def test_delete_tag_success( - self, - mock_current_user, - mock_service_api_ns, - mock_tag_svc, - app: Flask, - ): - from controllers.service_api.dataset.dataset import DatasetTagsApi - - user_obj = Mock(spec=Account) - user_obj.has_edit_permission = True - mock_current_user.has_edit_permission = True - # Assign as plain lambda to avoid AsyncMock returning a coroutine - mock_current_user._get_current_object = lambda: user_obj - - mock_tag_svc.delete_tag.return_value = None - mock_service_api_ns.payload = {"tag_id": "tag-1"} - - with app.test_request_context( - "/datasets/tags", - method="DELETE", - json={"tag_id": "tag-1"}, - ): - api = DatasetTagsApi() - result = api.delete(_=None) - - assert result == ("", 204) - mock_tag_svc.delete_tag.assert_called_once_with("tag-1", ANY, tag_type=TagType.KNOWLEDGE) - - @patch("libs.login.current_user") - def test_delete_tag_forbidden(self, mock_current_user, app: Flask): - from controllers.service_api.dataset.dataset import DatasetTagsApi - - user_obj = Mock(spec=Account) - user_obj.has_edit_permission = False - mock_current_user.has_edit_permission = False - # Assign as plain lambda to avoid AsyncMock returning a coroutine - mock_current_user._get_current_object = lambda: user_obj - - with app.test_request_context( - "/datasets/tags", - method="DELETE", - json={"tag_id": "tag-1"}, - ): - api = DatasetTagsApi() - with pytest.raises(Forbidden): - api.delete(_=None) - - -class TestDatasetTagsBindingStatusApi: - """Test suite for DatasetTagsBindingStatusApi endpoints.""" - - @patch("controllers.service_api.dataset.dataset.TagService") - @patch("controllers.service_api.dataset.dataset.current_user") - def test_get_dataset_tags_binding_status( - self, - mock_current_user, - mock_tag_svc, - app: Flask, - ): - from controllers.service_api.dataset.dataset import DatasetTagsBindingStatusApi - - mock_current_user.__class__ = Account - mock_current_user.current_tenant_id = "tenant_123" - mock_tag = Mock() - mock_tag.id = "tag_1" - mock_tag.name = "Test Tag" - mock_tag_svc.get_tags_by_target_id.return_value = [mock_tag] - - with app.test_request_context("/", method="GET"): - api = DatasetTagsBindingStatusApi() - response, status_code = api.get("tenant_123", dataset_id="dataset_123") - - assert status_code == 200 - assert response["data"] == [{"id": "tag_1", "name": "Test Tag"}] - assert response["total"] == 1 - mock_tag_svc.get_tags_by_target_id.assert_called_once_with("knowledge", "tenant_123", "dataset_123", ANY) - - -class TestDatasetTagBindingApiPost: - """Test suite for DatasetTagBindingApi.post() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.TagService") - @patch("controllers.service_api.dataset.dataset.current_user") - def test_bind_tags_success( - self, - mock_current_user, - mock_tag_svc, - app: Flask, - ): - from controllers.service_api.dataset.dataset import DatasetTagBindingApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = True - mock_current_user.is_dataset_editor = True - mock_tag_svc.save_tag_binding.return_value = None - - with app.test_request_context( - "/datasets/tags/binding", - method="POST", - json={"tag_ids": ["tag-1"], "target_id": "ds-1"}, - ): - api = DatasetTagBindingApi() - result = api.post(_=None) - - assert result == ("", 204) - from services.tag_service import TagBindingCreatePayload - - mock_tag_svc.save_tag_binding.assert_called_once_with( - TagBindingCreatePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE), - ANY, - ) - - @patch("controllers.service_api.dataset.dataset.current_user") - def test_bind_tags_forbidden(self, mock_current_user, app: Flask): - from controllers.service_api.dataset.dataset import DatasetTagBindingApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = False - mock_current_user.is_dataset_editor = False - - with app.test_request_context( - "/datasets/tags/binding", - method="POST", - json={"tag_ids": ["tag-1"], "target_id": "ds-1"}, - ): - api = DatasetTagBindingApi() - with pytest.raises(Forbidden): - api.post(_=None) - - -class TestDatasetTagUnbindingApiPost: - """Test suite for DatasetTagUnbindingApi.post() endpoint.""" - - @patch("controllers.service_api.dataset.dataset.TagService") - @patch("controllers.service_api.dataset.dataset.current_user") - def test_unbind_tag_success( - self, - mock_current_user, - mock_tag_svc, - app: Flask, - ): - from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = True - mock_current_user.is_dataset_editor = True - mock_tag_svc.delete_tag_binding.return_value = None - - with app.test_request_context( - "/datasets/tags/unbinding", - method="POST", - json={"tag_ids": ["tag-1"], "target_id": "ds-1"}, - ): - api = DatasetTagUnbindingApi() - result = api.post(_=None) - - assert result == ("", 204) - from services.tag_service import TagBindingDeletePayload - - mock_tag_svc.delete_tag_binding.assert_called_once_with( - TagBindingDeletePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE), - ANY, - ) - - @patch("controllers.service_api.dataset.dataset.TagService") - @patch("controllers.service_api.dataset.dataset.current_user") - def test_unbind_legacy_tag_id_success( - self, - mock_current_user, - mock_tag_svc, - app: Flask, - ): - from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = True - mock_current_user.is_dataset_editor = True - mock_tag_svc.delete_tag_binding.return_value = None - - with app.test_request_context( - "/datasets/tags/unbinding", - method="POST", - json={"tag_id": "tag-1", "target_id": "ds-1"}, - ): - api = DatasetTagUnbindingApi() - result = api.post(_=None) - - assert result == ("", 204) - from services.tag_service import TagBindingDeletePayload - - mock_tag_svc.delete_tag_binding.assert_called_once_with( - TagBindingDeletePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE), - ANY, - ) - - @patch("controllers.service_api.dataset.dataset.current_user") - def test_unbind_tag_forbidden(self, mock_current_user, app: Flask): - from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi - - mock_current_user.__class__ = Account - mock_current_user.has_edit_permission = False - mock_current_user.is_dataset_editor = False - - with app.test_request_context( - "/datasets/tags/unbinding", - method="POST", - json={"tag_ids": ["tag-1"], "target_id": "ds-1"}, - ): - api = DatasetTagUnbindingApi() - with pytest.raises(Forbidden): - api.post(_=None) + assert any(item["name"] == "Integration Tag" for item in response) + assert all(set(item) == {"id", "name", "type", "binding_count"} for item in response) + assert all(isinstance(item["binding_count"], str) for item in response) diff --git a/api/tests/test_containers_integration_tests/controllers/web/test_human_input_form.py b/api/tests/test_containers_integration_tests/controllers/web/test_human_input_form.py index 4ec706c1f10..5c3ad5e199b 100644 --- a/api/tests/test_containers_integration_tests/controllers/web/test_human_input_form.py +++ b/api/tests/test_containers_integration_tests/controllers/web/test_human_input_form.py @@ -37,7 +37,7 @@ from models.human_input import ( from models.model import App, AppMode, CustomizeTokenStrategy, Site from models.workflow import WorkflowRun, WorkflowType from repositories.sqlalchemy_api_workflow_run_repository import DifyAPISQLAlchemyWorkflowRunRepository -from services.feature_service import FeatureModel +from services.entities.feature_entities import FeatureModel class _TestWorkflowRunRepository(DifyAPISQLAlchemyWorkflowRunRepository): diff --git a/api/tests/test_containers_integration_tests/controllers/web/test_site.py b/api/tests/test_containers_integration_tests/controllers/web/test_site.py index 7f4fd45d037..cdba83851fc 100644 --- a/api/tests/test_containers_integration_tests/controllers/web/test_site.py +++ b/api/tests/test_containers_integration_tests/controllers/web/test_site.py @@ -11,11 +11,12 @@ from werkzeug.exceptions import Forbidden from configs import dify_config from controllers.web.site import AppSiteApi, WebAppSiteResponse, WebModelConfigResponse +from enums import DeploymentEdition 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 -from services.feature_service import FeatureModel +from services.entities.feature_entities import FeatureModel @pytest.fixture @@ -97,6 +98,7 @@ class TestAppSiteApi: assert result["end_user_id"] == end_user.id assert result["plan"] == "basic" assert result["enable_site"] is True + assert result["mode"] == AppMode.CHAT @patch("controllers.web.site.FileService.get_file_presigned_url") @patch("controllers.web.site.FeatureService.get_features") @@ -119,7 +121,7 @@ class TestAppSiteApi: 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, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch.object(dify_config, "STORAGE_TYPE", StorageType.S3), app.test_request_context("/site"), ): @@ -178,6 +180,7 @@ class TestWebAppSiteResponse: response = WebAppSiteResponse.from_app_site( tenant=tenant, app_model=app_model, + mode=AppMode.CHAT, site=_site_model(app_id=app_model.id), end_user_id="eu-1", features=FeatureModel(can_replace_logo=False, webapp_copyright_enabled=True), @@ -185,6 +188,7 @@ class TestWebAppSiteResponse: ) assert response.app_id == app_model.id + assert response.mode == AppMode.CHAT assert response.end_user_id == "eu-1" assert response.enable_site is True assert response.plan == "basic" @@ -209,6 +213,7 @@ class TestWebAppSiteResponse: response = WebAppSiteResponse.from_app_site( tenant=tenant, app_model=app_model, + mode=AppMode.CHAT, site=site, end_user_id=None, features=FeatureModel(can_replace_logo=False, webapp_copyright_enabled=True), @@ -236,6 +241,7 @@ class TestWebAppSiteResponse: response = WebAppSiteResponse.from_app_site( tenant=tenant, app_model=app_model, + mode=AppMode.CHAT, site=_site_model(app_id=app_model.id), end_user_id="eu-1", features=FeatureModel(can_replace_logo=True, webapp_copyright_enabled=True), diff --git a/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py b/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py index aa85ac2ca7b..01e241a9b94 100644 --- a/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py +++ b/api/tests/test_containers_integration_tests/controllers/web/test_wraps.py @@ -18,6 +18,7 @@ from controllers.web.wraps import ( _validate_webapp_token, decode_jwt_token, ) +from models.enums import EndUserType pytestmark = pytest.mark.usefixtures("db_session_with_containers") @@ -189,7 +190,6 @@ class TestDecodeJwtToken: return flask_app_with_containers def _create_app_site_enduser(self, db_session: Session, *, enable_site: bool = True): - from models.enums import EndUserType from models.model import App, AppMode, CustomizeTokenStrategy, EndUser, Site tenant_id = str(uuid4()) diff --git a/api/tests/test_containers_integration_tests/core/app/layers/test_pause_state_persist_layer.py b/api/tests/test_containers_integration_tests/core/app/layers/test_pause_state_persist_layer.py index 84f01ea52ee..2d5e3ebcca9 100644 --- a/api/tests/test_containers_integration_tests/core/app/layers/test_pause_state_persist_layer.py +++ b/api/tests/test_containers_integration_tests/core/app/layers/test_pause_state_persist_layer.py @@ -17,7 +17,6 @@ These tests use TestContainers to spin up real services for integration testing, providing more reliable and realistic test scenarios than mocks. """ -import json import uuid from time import time from unittest.mock import Mock @@ -237,12 +236,12 @@ class TestPauseStatePersistenceLayerTestContainers: # Create LLM usage llm_usage = LLMUsage.empty_usage() + llm_usage.total_tokens = total_tokens # Create graph runtime state graph_runtime_state = GraphRuntimeState( variable_pool=variable_pool, start_at=start_at, - total_tokens=total_tokens, llm_usage=llm_usage, outputs=outputs or {}, node_run_steps=node_run_steps, @@ -366,9 +365,6 @@ class TestPauseStatePersistenceLayerTestContainers: resumption_context = WorkflowResumptionContext.loads(storage_content) assert resumption_context.version == "1" assert resumption_context.serialized_graph_runtime_state == graph_runtime_state.dumps() - expected_state = json.loads(graph_runtime_state.dumps()) - actual_state = json.loads(resumption_context.serialized_graph_runtime_state) - assert actual_state == expected_state persisted_entity = resumption_context.get_generate_entity() assert isinstance(persisted_entity, WorkflowAppGenerateEntity) assert persisted_entity.workflow_execution_id == self.test_workflow_run_id @@ -414,13 +410,11 @@ class TestPauseStatePersistenceLayerTestContainers: state_bytes = pause_entity.get_state() resumption_context = WorkflowResumptionContext.loads(state_bytes.decode()) - retrieved_state = json.loads(resumption_context.serialized_graph_runtime_state) - expected_state = json.loads(graph_runtime_state.dumps()) + retrieved_state = GraphRuntimeState.from_snapshot(resumption_context.serialized_graph_runtime_state) - assert retrieved_state == expected_state - assert retrieved_state["outputs"] == complex_outputs - assert retrieved_state["total_tokens"] == 250 - assert retrieved_state["node_run_steps"] == 10 + assert retrieved_state.outputs == complex_outputs + assert retrieved_state.total_tokens == 250 + assert retrieved_state.node_run_steps == 10 assert resumption_context.get_generate_entity().workflow_execution_id == self.test_workflow_run_id def test_database_transaction_handling(self, db_session_with_containers: Session): diff --git a/api/tests/test_containers_integration_tests/models/test_types_enum_text.py b/api/tests/test_containers_integration_tests/models/test_types_enum_text.py index b325c97f7d1..cb114537e63 100644 --- a/api/tests/test_containers_integration_tests/models/test_types_enum_text.py +++ b/api/tests/test_containers_integration_tests/models/test_types_enum_text.py @@ -210,7 +210,7 @@ class TestEnumText: assert str(exc.value) == "'invalid' is not a valid _UserType" - def test_select_legacy_model_type_values(self, engine_with_containers: Engine): + def test_select_rejects_legacy_model_type_values(self, engine_with_containers: Engine): insertion_sql = """ INSERT INTO enum_text_legacy_model_type_test (id, model_type) VALUES (1, 'text-generation'), @@ -221,11 +221,9 @@ class TestEnumText: session.execute(sa.text(insertion_sql)) session.commit() - with Session(engine_with_containers) as session: - records = session.scalars(select(_LegacyModelTypeRecord).order_by(_LegacyModelTypeRecord.id)).all() + for record_id, legacy_value in enumerate(("text-generation", "embeddings", "reranking"), 1): + with pytest.raises(ValueError) as exc: + with Session(engine_with_containers) as session: + session.scalar(select(_LegacyModelTypeRecord).where(_LegacyModelTypeRecord.id == record_id)) - assert [record.model_type for record in records] == [ - ModelType.LLM, - ModelType.TEXT_EMBEDDING, - ModelType.RERANK, - ] + assert str(exc.value) == f"'{legacy_value}' is not a valid ModelType" diff --git a/api/tests/test_containers_integration_tests/pyrefly.toml b/api/tests/test_containers_integration_tests/pyrefly.toml index e73d5bbe117..6bdd09b3057 100644 --- a/api/tests/test_containers_integration_tests/pyrefly.toml +++ b/api/tests/test_containers_integration_tests/pyrefly.toml @@ -1,42 +1,24 @@ preset = "strict" -strict-callable-subtyping = true project-includes = ["."] search-path = ["../.."] +python-platform = "linux" +python-version = "3.12.0" +infer-with-first-use = true +min-severity = "warn" -# Verify project-excludes from the repo root: -# tmp_config=$(mktemp --tmpdir=api/tests/test_containers_integration_tests pyrefly-no-excludes.XXXXXX.toml) -# awk 'BEGIN {skip=0} /^project-excludes = \[/ {skip=1; next} skip && /^\]/ {skip=0; next} !skip {print}' api/tests/test_containers_integration_tests/pyrefly.toml > "$tmp_config" -# tmp_name=$(basename "$tmp_config") -# comm -3 <(sed -n 's/^ "\(.*\)",$/\1/p' api/tests/test_containers_integration_tests/pyrefly.toml | sort) <(uv --directory api run pyrefly check --config "tests/test_containers_integration_tests/$tmp_name" --summary=none --output-format=min-text 2>/dev/null | rg '^ERROR ' | sed -E 's#^ERROR (tests/test_containers_integration_tests/[^:]+):.*#\1#' | sed 's#^tests/test_containers_integration_tests/##' | sort -u) -# rm --force "$tmp_config" +# Existing strict-mode debt. Remove a file when bringing it under strict checking. project-excludes = [ "commands/test_legacy_model_type_migration.py", - "controllers/console/app/test_app_apis.py", - "controllers/console/app/test_app_import_api.py", "controllers/console/app/test_chat_conversation_status_count_api.py", "controllers/console/app/test_conversation_read_timestamp.py", "controllers/console/app/test_workflow_draft_variable.py", - "controllers/console/auth/test_email_register.py", - "controllers/console/auth/test_forgot_password.py", - "controllers/console/auth/test_oauth.py", - "controllers/console/auth/test_password_reset.py", - "controllers/console/datasets/rag_pipeline/test_rag_pipeline.py", - "controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py", - "controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py", - "controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py", - "controllers/console/datasets/test_data_source.py", - "controllers/console/explore/test_conversation.py", "controllers/console/test_api_based_extension.py", "controllers/console/test_apikey.py", - "controllers/console/workspace/test_tool_provider.py", - "controllers/console/workspace/test_trigger_providers.py", "controllers/console/workspace/test_workspace_wraps.py", - "controllers/mcp/test_mcp.py", "controllers/service_api/dataset/test_dataset.py", "controllers/service_api/test_site.py", "controllers/web/test_conversation.py", "controllers/web/test_site.py", - "controllers/web/test_web_forgot_password.py", "controllers/web/test_wraps.py", "core/app/layers/test_pause_state_persist_layer.py", "core/rag/pipeline/test_queue_integration.py", @@ -57,22 +39,16 @@ project-excludes = [ "repositories/test_sqlalchemy_execution_extra_content_repository.py", "repositories/test_sqlalchemy_workflow_node_execution_repository.py", "repositories/test_workflow_run_repository.py", - "services/auth/test_api_key_auth_service.py", "services/auth/test_auth_integration.py", "services/dataset_collection_binding.py", "services/dataset_service_update_delete.py", "services/document_service_status.py", - "services/enterprise/test_account_deletion_sync.py", - "services/plugin/test_plugin_parameter_service.py", - "services/plugin/test_plugin_service.py", - "services/rag_pipeline/test_rag_pipeline_service_db.py", "services/recommend_app/test_database_retrieval.py", "services/test_account_service.py", "services/test_advanced_prompt_template_service.py", "services/test_agent_service.py", "services/test_annotation_service.py", "services/test_api_based_extension_service.py", - "services/test_api_token_service.py", "services/test_app_dsl_service.py", "services/test_app_generate_service.py", "services/test_app_service.py", @@ -97,20 +73,15 @@ project-excludes = [ "services/test_document_service_rename_document.py", "services/test_end_user_service.py", "services/test_feature_service.py", - "services/test_feedback_service.py", "services/test_file_service.py", "services/test_human_input_delivery_test.py", - "services/test_human_input_delivery_test_service.py", "services/test_message_export_service.py", "services/test_message_service.py", "services/test_message_service_execution_extra_content.py", "services/test_message_service_extra_contents.py", "services/test_messages_clean_service.py", - "services/test_metadata_partial_update.py", - "services/test_metadata_service.py", "services/test_model_load_balancing_service.py", "services/test_model_provider_service.py", - "services/test_oauth_server_service.py", "services/test_ops_service.py", "services/test_restore_archived_workflow_run.py", "services/test_saved_message_service.py", @@ -124,7 +95,6 @@ project-excludes = [ "services/test_workflow_run_service.py", "services/test_workflow_service.py", "services/test_workspace_service.py", - "services/tools/test_api_tools_manage_service.py", "services/tools/test_mcp_tools_manage_service.py", "services/tools/test_tools_transform_service.py", "services/tools/test_workflow_tools_manage_service.py", @@ -161,7 +131,6 @@ project-excludes = [ "test_workflow_pause_integration.py", "trigger/conftest.py", "trigger/test_trigger_e2e.py", - "workflow/nodes/code_executor/test_code_executor.py", "workflow/nodes/code_executor/test_code_javascript.py", "workflow/nodes/code_executor/test_code_jinja2.py", "workflow/nodes/code_executor/test_code_python3.py", @@ -169,7 +138,9 @@ project-excludes = [ ] [errors] +missing-override-decorator = "error" redundant-cast = true unannotated-return = true -unnecessary-type-conversion = true unused-ignore = true +implicit-any-lambda = "info" +unnecessary-type-conversion = "info" diff --git a/api/tests/test_containers_integration_tests/repositories/test_sqlalchemy_workflow_trigger_log_repository.py b/api/tests/test_containers_integration_tests/repositories/test_sqlalchemy_workflow_trigger_log_repository.py index 0c4d75359e4..e7ac10fc986 100644 --- a/api/tests/test_containers_integration_tests/repositories/test_sqlalchemy_workflow_trigger_log_repository.py +++ b/api/tests/test_containers_integration_tests/repositories/test_sqlalchemy_workflow_trigger_log_repository.py @@ -125,8 +125,7 @@ def test_delete_by_run_ids_empty_short_circuits(db_session_with_containers: Sess remaining_count = db_session_with_containers.scalar( select(func.count()) .select_from(WorkflowTriggerLog) - .where(WorkflowTriggerLog.tenant_id == tenant_id) - .where(WorkflowTriggerLog.workflow_run_id == run_id) + .where(WorkflowTriggerLog.tenant_id == tenant_id, WorkflowTriggerLog.workflow_run_id == run_id) ) assert remaining_count == 1 finally: diff --git a/api/tests/test_containers_integration_tests/services/dataset_service_update_delete.py b/api/tests/test_containers_integration_tests/services/dataset_service_update_delete.py index 7d4cd63267a..3acf8e5cd8e 100644 --- a/api/tests/test_containers_integration_tests/services/dataset_service_update_delete.py +++ b/api/tests/test_containers_integration_tests/services/dataset_service_update_delete.py @@ -11,13 +11,13 @@ from uuid import uuid4 import pytest from sqlalchemy.orm import Session -from werkzeug.exceptions import NotFound from core.rag.index_processor.constant.index_type import IndexTechniqueType from models import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus from models.dataset import AppDatasetJoin, Dataset, DatasetPermissionEnum from models.enums import DataSourceType from models.model import App +from services.dataset_ref_service import DatasetRefService from services.dataset_service import DatasetService from services.errors.account import NoPermissionError @@ -228,9 +228,10 @@ class TestDatasetServiceDatasetUseCheck: dataset = DatasetUpdateDeleteTestDataFactory.create_dataset(db_session_with_containers, tenant.id, owner.id) app = DatasetUpdateDeleteTestDataFactory.create_app(db_session_with_containers, tenant.id, owner.id) DatasetUpdateDeleteTestDataFactory.create_app_dataset_join(db_session_with_containers, app.id, dataset.id) + dataset_ref = DatasetRefService.create_dataset_ref(dataset) # Act - result = DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) + result = DatasetService.dataset_use_check(dataset_ref, session=db_session_with_containers) # Assert assert result is True @@ -252,9 +253,10 @@ class TestDatasetServiceDatasetUseCheck: db_session_with_containers, role=TenantAccountRole.OWNER ) dataset = DatasetUpdateDeleteTestDataFactory.create_dataset(db_session_with_containers, tenant.id, owner.id) + dataset_ref = DatasetRefService.create_dataset_ref(dataset) # Act - result = DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) + result = DatasetService.dataset_use_check(dataset_ref, session=db_session_with_containers) # Assert assert result is False @@ -288,11 +290,8 @@ class TestDatasetServiceUpdateDatasetApiStatus: current_time = datetime.datetime(2023, 1, 1, 12, 0, 0) # Act - with ( - patch("services.dataset_service.current_user", owner), - patch("services.dataset_service.naive_utc_now", return_value=current_time), - ): - DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers) + with patch("services.dataset_service.naive_utc_now", return_value=current_time): + DatasetService.update_dataset_api_status(dataset, True, owner, session=db_session_with_containers) # Assert db_session_with_containers.refresh(dataset) @@ -323,36 +322,14 @@ class TestDatasetServiceUpdateDatasetApiStatus: current_time = datetime.datetime(2023, 1, 1, 12, 0, 0) # Act - with ( - patch("services.dataset_service.current_user", owner), - patch("services.dataset_service.naive_utc_now", return_value=current_time), - ): - DatasetService.update_dataset_api_status(dataset.id, False, session=db_session_with_containers) + with patch("services.dataset_service.naive_utc_now", return_value=current_time): + DatasetService.update_dataset_api_status(dataset, False, owner, session=db_session_with_containers) # Assert db_session_with_containers.refresh(dataset) assert dataset.enable_api is False assert dataset.updated_by == owner.id - def test_update_dataset_api_status_not_found_error(self, db_session_with_containers: Session): - """ - Test error handling when dataset is not found. - - Verifies that when the dataset ID doesn't exist, a NotFound - exception is raised. - - This test ensures: - - NotFound exception is raised - - No updates are performed - - Error message is appropriate - """ - # Arrange - dataset_id = str(uuid4()) - - # Act & Assert - with pytest.raises(NotFound, match="Dataset not found"): - DatasetService.update_dataset_api_status(dataset_id, True, session=db_session_with_containers) - def test_update_dataset_api_status_missing_current_user_error(self, db_session_with_containers: Session): """ Test error handling when current_user is missing. @@ -372,13 +349,11 @@ class TestDatasetServiceUpdateDatasetApiStatus: dataset = DatasetUpdateDeleteTestDataFactory.create_dataset( db_session_with_containers, tenant.id, owner.id, enable_api=False ) - + actor = Account(name="missing-id", email="missing-id@example.com") + actor.id = "" # Act & Assert - with ( - patch("services.dataset_service.current_user", None), - pytest.raises(ValueError, match="Current user or current user id not found"), - ): - DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers) + with pytest.raises(ValueError, match="Current user or current user id not found"): + DatasetService.update_dataset_api_status(dataset, True, actor, session=db_session_with_containers) # Verify no commit was attempted db_session_with_containers.rollback() diff --git a/api/tests/test_containers_integration_tests/services/document_service_status.py b/api/tests/test_containers_integration_tests/services/document_service_status.py index 7e78cef1db3..6ffaf49d947 100644 --- a/api/tests/test_containers_integration_tests/services/document_service_status.py +++ b/api/tests/test_containers_integration_tests/services/document_service_status.py @@ -8,7 +8,7 @@ pause, recover, retry, batch updates, and renaming. import datetime import json -from unittest.mock import create_autospec, patch +from unittest.mock import MagicMock, create_autospec, patch from uuid import uuid4 import pytest @@ -562,9 +562,9 @@ class TestDocumentServiceRetryDocument: The retry_document method: 1. Validates documents are not already being retried - 2. Sets retry flag in Redis cache - 3. Resets document indexing_status to waiting - 4. Commits changes to database + 2. Atomically reserves retry flags in Redis cache + 3. Resets all document indexing statuses to waiting + 4. Commits all changes together 5. Triggers retry task Test scenarios include: @@ -595,12 +595,22 @@ class TestDocumentServiceRetryDocument: ): user_id = str(uuid4()) mock_current_user.id = user_id + retry_locks = [] + + def create_retry_lock(*_args, **_kwargs): + retry_lock = MagicMock() + retry_lock.acquire.return_value = True + retry_locks.append(retry_lock) + return retry_lock + + mock_redis.lock.side_effect = create_retry_lock yield { "current_user": mock_current_user, "redis_client": mock_redis, "retry_task": mock_task, "user_id": user_id, + "retry_locks": retry_locks, } def test_retry_document_single_success( @@ -629,8 +639,6 @@ class TestDocumentServiceRetryDocument: indexing_status=IndexingStatus.ERROR, ) - mock_document_service_dependencies["redis_client"].get.return_value = None - # Act DocumentService.retry_document(dataset.id, [document], session=db_session_with_containers) @@ -639,7 +647,11 @@ class TestDocumentServiceRetryDocument: assert document.indexing_status == IndexingStatus.WAITING expected_cache_key = f"document_{document.id}_is_retried" - mock_document_service_dependencies["redis_client"].setex.assert_called_once_with(expected_cache_key, 600, 1) + mock_document_service_dependencies["redis_client"].lock.assert_called_once_with( + expected_cache_key, timeout=600, thread_local=False + ) + retry_lock = mock_document_service_dependencies["retry_locks"][0] + retry_lock.acquire.assert_called_once_with(blocking=False) mock_document_service_dependencies["retry_task"].delay.assert_called_once_with( dataset.id, [document.id], mock_document_service_dependencies["user_id"] ) @@ -676,8 +688,6 @@ class TestDocumentServiceRetryDocument: position=2, ) - mock_document_service_dependencies["redis_client"].get.return_value = None - # Act DocumentService.retry_document(dataset.id, [document1, document2], session=db_session_with_containers) @@ -715,7 +725,10 @@ class TestDocumentServiceRetryDocument: indexing_status=IndexingStatus.ERROR, ) - mock_document_service_dependencies["redis_client"].get.return_value = "1" + retry_lock = MagicMock() + retry_lock.acquire.return_value = False + mock_document_service_dependencies["redis_client"].lock.side_effect = None + mock_document_service_dependencies["redis_client"].lock.return_value = retry_lock # Act & Assert with pytest.raises(ValueError, match="Document is being retried, please try again later"): @@ -724,6 +737,42 @@ class TestDocumentServiceRetryDocument: db_session_with_containers.refresh(document) assert document.indexing_status == IndexingStatus.ERROR + def test_retry_document_later_conflict_leaves_batch_unchanged( + self, db_session_with_containers: Session, mock_document_service_dependencies + ): + dataset = DocumentStatusTestDataFactory.create_dataset(db_session_with_containers) + document1 = DocumentStatusTestDataFactory.create_document( + db_session_with_containers, + dataset_id=dataset.id, + tenant_id=dataset.tenant_id, + document_id=str(uuid4()), + indexing_status=IndexingStatus.ERROR, + ) + document2 = DocumentStatusTestDataFactory.create_document( + db_session_with_containers, + dataset_id=dataset.id, + tenant_id=dataset.tenant_id, + document_id=str(uuid4()), + indexing_status=IndexingStatus.ERROR, + position=2, + ) + first_retry_lock = MagicMock() + first_retry_lock.acquire.return_value = True + second_retry_lock = MagicMock() + second_retry_lock.acquire.return_value = False + mock_document_service_dependencies["redis_client"].lock.side_effect = [first_retry_lock, second_retry_lock] + + with pytest.raises(ValueError, match="Document is being retried, please try again later"): + DocumentService.retry_document(dataset.id, [document1, document2], session=db_session_with_containers) + + db_session_with_containers.refresh(document1) + db_session_with_containers.refresh(document2) + assert document1.indexing_status == IndexingStatus.ERROR + assert document2.indexing_status == IndexingStatus.ERROR + first_retry_lock.release.assert_called_once_with() + second_retry_lock.release.assert_not_called() + mock_document_service_dependencies["retry_task"].delay.assert_not_called() + def test_retry_document_missing_current_user_error( self, db_session_with_containers: Session, mock_document_service_dependencies ): @@ -748,13 +797,16 @@ class TestDocumentServiceRetryDocument: indexing_status=IndexingStatus.ERROR, ) - mock_document_service_dependencies["redis_client"].get.return_value = None mock_document_service_dependencies["current_user"].id = None # Act & Assert with pytest.raises(ValueError, match="Current user or current user id not found"): DocumentService.retry_document(dataset.id, [document], session=db_session_with_containers) + db_session_with_containers.refresh(document) + assert document.indexing_status == IndexingStatus.ERROR + mock_document_service_dependencies["redis_client"].lock.assert_not_called() + class TestDocumentServiceBatchUpdateDocumentStatus: """ diff --git a/api/tests/test_containers_integration_tests/services/test_account_service.py b/api/tests/test_containers_integration_tests/services/test_account_service.py index 26b20e83a9b..e6d148c0ffe 100644 --- a/api/tests/test_containers_integration_tests/services/test_account_service.py +++ b/api/tests/test_containers_integration_tests/services/test_account_service.py @@ -9,6 +9,7 @@ from werkzeug.exceptions import Unauthorized from configs import dify_config from controllers.console.error import AccountNotFound, NotAllowedCreateWorkspace +from enums import DeploymentEdition from models import AccountStatus, App, Dataset, TenantAccountJoin, TenantStatus from services.account_service import AccountService, RegisterService, TenantService, TokenPair from services.errors.account import ( @@ -156,7 +157,7 @@ class TestAccountService: # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = True - dify_config.BILLING_ENABLED = True + dify_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD with pytest.raises(AccountRegisterError): AccountService.create_account( @@ -167,7 +168,7 @@ class TestAccountService: session=db_session_with_containers, ) - dify_config.BILLING_ENABLED = False # Reset config for other tests + dify_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY # Reset config for other tests def test_authenticate_account_not_found( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -1105,14 +1106,14 @@ class TestAccountService: fake = Faker() email_in_freeze = fake.email() # Setup mocks - dify_config.BILLING_ENABLED = True + dify_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = True with pytest.raises(AccountRegisterError): AccountService.get_user_through_email(email_in_freeze, session=db_session_with_containers) # Reset config - dify_config.BILLING_ENABLED = False + dify_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY def test_delete_account(self, db_session_with_containers: Session, mock_external_service_dependencies): """ diff --git a/api/tests/test_containers_integration_tests/services/test_agent_service.py b/api/tests/test_containers_integration_tests/services/test_agent_service.py index 00b4a1563ff..445f2704641 100644 --- a/api/tests/test_containers_integration_tests/services/test_agent_service.py +++ b/api/tests/test_containers_integration_tests/services/test_agent_service.py @@ -843,7 +843,6 @@ class TestAgentService: conversation, message = self._create_test_conversation_and_message(db_session_with_containers, app, account) from graphon.file import FileTransferMethod, FileType - from models.enums import CreatorUserRole # Add files to message from models.model import MessageFile diff --git a/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py b/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py index 32591c13312..8ea3a7b8c81 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_dsl_service.py @@ -142,7 +142,6 @@ class TestAppDslService: mock_feature_service.get_system_features.return_value.webapp_auth.enabled = False mock_enterprise_service.WebAppAuth.update_app_access_mode.return_value = None mock_enterprise_service.WebAppAuth.cleanup_webapp.return_value = None - yield { "workflow_service": mock_workflow_service, "dependencies_service": mock_dependencies_service, @@ -492,6 +491,9 @@ class TestAppDslService: redis_key = f"{IMPORT_INFO_REDIS_KEY_PREFIX}{result.id}" stored = redis_client.get(redis_key) assert stored is not None + pending = PendingData.model_validate_json(stored) + assert pending.tenant_id == _DEFAULT_TENANT_ID + assert pending.account_id == _DEFAULT_ACCOUNT_ID def test_import_app_completed_uses_declared_dependencies( self, db_session_with_containers: Session, mock_external_service_dependencies @@ -607,7 +609,11 @@ class TestAppDslService: icon_background="#fff", app_id=None, ) - redis_client.setex(redis_key, IMPORT_INFO_REDIS_EXPIRY, pending.model_dump_json()) + redis_client.setex( + redis_key, + IMPORT_INFO_REDIS_EXPIRY, + pending.model_dump_json(exclude={"tenant_id", "account_id"}), + ) created_app = SimpleNamespace( id=str(uuid4()), @@ -1034,7 +1040,9 @@ class TestAppDslService: ) ) - imported_graph, warnings = AgentDslService(db_session_with_containers).import_workflow_packages( + imported_graph, warnings, retirement_candidates = AgentDslService( + db_session_with_containers + ).import_workflow_packages( workflow=workflow, portable_graph=graph, raw_packages={"agent_1": package.model_dump(mode="json")}, @@ -1043,6 +1051,7 @@ class TestAppDslService: db_session_with_containers.commit() assert warnings == [] + assert retirement_candidates == set() graph_bindings = [node["data"]["agent_binding"] for node in imported_graph["nodes"]] assert all(binding["binding_type"] == WorkflowAgentBindingType.INLINE_AGENT.value for binding in graph_bindings) assert len({binding["agent_id"] for binding in graph_bindings}) == 2 @@ -1199,6 +1208,10 @@ class TestAppDslService: ) assert imported_agent is not None assert imported_agent.active_config_is_published is False + imported_app = db_session_with_containers.get(App, result.app_id) + assert imported_app is not None + assert imported_app.enable_site is False + assert imported_app.enable_api is False draft = db_session_with_containers.scalar( select(AgentConfigDraft).where( AgentConfigDraft.agent_id == imported_agent.id, @@ -1305,7 +1318,7 @@ class TestAppDslService: with pytest.raises( WorkflowNotFoundError, - match="Missing draft workflow configuration, please check.", + match="Workflow version not found. Workflow ID:", ): AppDslService.export_dsl( app, include_secret=False, workflow_id=str(uuid4()), session=db_session_with_containers diff --git a/api/tests/test_containers_integration_tests/services/test_app_generate_service.py b/api/tests/test_containers_integration_tests/services/test_app_generate_service.py index 89cc7715d1c..d9f81caf01c 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_generate_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_generate_service.py @@ -8,6 +8,7 @@ from faker import Faker from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom +from enums import DeploymentEdition from models import App from models.enums import EndUserType from models.model import EndUser @@ -106,14 +107,14 @@ class TestAppGenerateService: mock_account_feature_service.get_system_features.return_value.is_allow_register = True # Setup dify_config mock returns - mock_dify_config.BILLING_ENABLED = False + mock_dify_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY mock_dify_config.APP_MAX_ACTIVE_REQUESTS = 100 mock_dify_config.APP_DEFAULT_ACTIVE_REQUESTS = 100 mock_dify_config.APP_DAILY_RATE_LIMIT = 1000 - mock_quota_dify_config.BILLING_ENABLED = False + mock_quota_dify_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY - mock_global_dify_config.BILLING_ENABLED = False + mock_global_dify_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY mock_global_dify_config.APP_MAX_ACTIVE_REQUESTS = 100 mock_global_dify_config.APP_DAILY_RATE_LIMIT = 1000 mock_global_dify_config.HOSTED_POOL_CREDITS = 1000 @@ -514,21 +515,21 @@ class TestAppGenerateService: # Verify the result assert result == ["test_response"] - def test_generate_with_billing_enabled_sandbox_plan( + def test_generate_in_cloud_sandbox_plan( self, db_session_with_containers: Session, mock_external_service_dependencies ): """ - Test generation with billing enabled and sandbox plan. + Test generation in the Cloud edition with a sandbox plan. """ fake = Faker() app, account = self._create_test_app_and_account( db_session_with_containers, mock_external_service_dependencies, mode="completion" ) - # Set BILLING_ENABLED to True for this test - mock_external_service_dependencies["dify_config"].BILLING_ENABLED = True - mock_external_service_dependencies["quota_dify_config"].BILLING_ENABLED = True - mock_external_service_dependencies["global_dify_config"].BILLING_ENABLED = True + # Billing services are available in the Cloud deployment edition. + mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD + mock_external_service_dependencies["quota_dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD + mock_external_service_dependencies["global_dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.CLOUD # Setup test arguments args = {"inputs": {"query": fake.text(max_nb_chars=50)}, "response_mode": "streaming"} diff --git a/api/tests/test_containers_integration_tests/services/test_app_service.py b/api/tests/test_containers_integration_tests/services/test_app_service.py index 2875f014380..c87914131d2 100644 --- a/api/tests/test_containers_integration_tests/services/test_app_service.py +++ b/api/tests/test_containers_integration_tests/services/test_app_service.py @@ -42,7 +42,6 @@ class TestAppService: mock_model_instance = mock_model_manager.return_value mock_model_instance.get_default_model_instance.return_value = None mock_model_instance.get_default_provider_model_name.return_value = ("openai", "gpt-3.5-turbo") - yield { "feature_service": mock_feature_service, "enterprise_service": mock_enterprise_service, diff --git a/api/tests/test_containers_integration_tests/services/test_conversation_service.py b/api/tests/test_containers_integration_tests/services/test_conversation_service.py index 09411f0b0e1..75f017fbc50 100644 --- a/api/tests/test_containers_integration_tests/services/test_conversation_service.py +++ b/api/tests/test_containers_integration_tests/services/test_conversation_service.py @@ -6,13 +6,11 @@ from unittest.mock import patch from uuid import uuid4 import pytest -from agenton.compositor import CompositorSessionSnapshot from sqlalchemy import select from sqlalchemy.orm import Session -from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore from core.app.entities.app_invoke_entities import InvokeFrom -from models import AgentRuntimeSession, AgentRuntimeSessionOwnerType, AgentRuntimeSessionStatus, TenantAccountRole +from models import TenantAccountRole from models.account import Account, Tenant, TenantAccountJoin from models.enums import ConversationFromSource, EndUserType from models.model import App, Conversation, EndUser, Message, MessageAnnotation @@ -1077,9 +1075,8 @@ class TestConversationServiceExport: # Assert assert result == conversation - @patch("services.conversation_service.cleanup_conversation_agent_runtime_session") @patch("services.conversation_service.delete_conversation_related_data") - def test_delete_conversation(self, mock_delete_task, mock_cleanup_task, db_session_with_containers: Session): + def test_delete_conversation(self, mock_delete_task, db_session_with_containers: Session): """ Test conversation deletion with async cleanup. @@ -1098,20 +1095,6 @@ class TestConversationServiceExport: user, ) conversation_id = conversation.id - runtime_session = AgentRuntimeSession( - tenant_id=app_model.tenant_id, - app_id=app_model.id, - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id=str(uuid4()), - agent_config_snapshot_id=str(uuid4()), - backend_run_id="backend-run-1", - session_snapshot=CompositorSessionSnapshot(layers=[]).model_dump_json(), - composition_layer_specs='[{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]', - conversation_id=conversation.id, - status=AgentRuntimeSessionStatus.ACTIVE, - ) - db_session_with_containers.add(runtime_session) - db_session_with_containers.commit() # Act - Delete the conversation ConversationService.delete( @@ -1126,27 +1109,11 @@ class TestConversationServiceExport: # Step 2: Async cleanup task triggered # The Celery task will handle cleanup of messages, annotations, etc. mock_delete_task.delay.assert_called_once_with(conversation_id) - mock_cleanup_task.delay.assert_called_once() - cleanup_payload = mock_cleanup_task.delay.call_args.args[0] - assert cleanup_payload["metadata"]["conversation_id"] == conversation_id - assert ( - cleanup_payload["idempotency_key"] - == f"{app_model.tenant_id}:{app_model.id}:{conversation_id}:agent-runtime-session-cleanup:" - f"{runtime_session.agent_id}:{runtime_session.agent_config_snapshot_id}:{runtime_session.backend_run_id}" - ) - runtime_session_row = db_session_with_containers.scalar( - select(AgentRuntimeSession).where(AgentRuntimeSession.id == runtime_session.id) - ) - assert runtime_session_row is not None - assert runtime_session_row.status == AgentRuntimeSessionStatus.CLEANED - - @patch("services.conversation_service.cleanup_conversation_agent_runtime_session") @patch("services.conversation_service.delete_conversation_related_data") def test_delete_conversation_not_owned_by_account( self, mock_delete_task, - mock_cleanup_task, db_session_with_containers: Session, ): """ @@ -1178,22 +1145,17 @@ class TestConversationServiceExport: not_deleted = db_session_with_containers.scalar(select(Conversation).where(Conversation.id == conversation.id)) assert not_deleted is not None mock_delete_task.delay.assert_not_called() - mock_cleanup_task.delay.assert_not_called() - @patch("services.conversation_service.cleanup_conversation_agent_runtime_session") @patch("services.conversation_service.delete_conversation_related_data") def test_delete_handles_exception_and_rollback( self, mock_delete_task, - mock_cleanup_task, db_session_with_containers: Session, ): """ Test that delete propagates exceptions and does not trigger the cleanup task. - When a DB error occurs during deletion, the conversation row stays in - place, but any already-enqueued Agent backend cleanup remains a - best-effort terminal lifecycle action. + When a DB error occurs during deletion, the conversation row stays in place. """ # Arrange app_model, user = ConversationServiceIntegrationTestDataFactory.create_app_and_account( @@ -1203,20 +1165,6 @@ class TestConversationServiceExport: db_session_with_containers, app_model, user ) conversation_id = conversation.id - runtime_session = AgentRuntimeSession( - tenant_id=app_model.tenant_id, - app_id=app_model.id, - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id=str(uuid4()), - agent_config_snapshot_id=str(uuid4()), - backend_run_id="backend-run-rollback", - session_snapshot=CompositorSessionSnapshot(layers=[]).model_dump_json(), - composition_layer_specs='[{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]', - conversation_id=conversation.id, - status=AgentRuntimeSessionStatus.ACTIVE, - ) - db_session_with_containers.add(runtime_session) - db_session_with_containers.commit() # Act — force an error during the delete to exercise the rollback path with patch.object(db_session_with_containers, "delete", side_effect=Exception("DB error")): @@ -1228,111 +1176,9 @@ class TestConversationServiceExport: session=db_session_with_containers, ) - # Assert — related-data deletion is not scheduled, but the backend - # cleanup task was already enqueued before the row delete failed. + # Assert — related-data deletion is not scheduled. mock_delete_task.delay.assert_not_called() - mock_cleanup_task.delay.assert_called_once() - cleanup_payload = mock_cleanup_task.delay.call_args.args[0] - assert ( - cleanup_payload["idempotency_key"] - == f"{app_model.tenant_id}:{app_model.id}:{conversation_id}:agent-runtime-session-cleanup:" - f"{runtime_session.agent_id}:{runtime_session.agent_config_snapshot_id}:{runtime_session.backend_run_id}" - ) # Conversation is still present because the deletion was never committed still_there = db_session_with_containers.scalar(select(Conversation).where(Conversation.id == conversation_id)) assert still_there is not None - - @patch("services.conversation_service.cleanup_conversation_agent_runtime_session") - @patch("services.conversation_service.delete_conversation_related_data") - def test_delete_ignores_mark_cleaned_failure( - self, - mock_delete_task, - mock_cleanup_task, - db_session_with_containers: Session, - ): - app_model, user = ConversationServiceIntegrationTestDataFactory.create_app_and_account( - db_session_with_containers - ) - conversation = ConversationServiceIntegrationTestDataFactory.create_conversation( - db_session_with_containers, - app_model, - user, - ) - runtime_session = AgentRuntimeSession( - tenant_id=app_model.tenant_id, - app_id=app_model.id, - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id=str(uuid4()), - agent_config_snapshot_id=str(uuid4()), - backend_run_id="backend-run-cleanup-failure", - session_snapshot=CompositorSessionSnapshot(layers=[]).model_dump_json(), - composition_layer_specs='[{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]', - conversation_id=conversation.id, - status=AgentRuntimeSessionStatus.ACTIVE, - ) - db_session_with_containers.add(runtime_session) - db_session_with_containers.commit() - - with patch.object(AgentAppRuntimeSessionStore, "mark_cleaned", side_effect=RuntimeError("cleanup failed")): - ConversationService.delete( - app_model=app_model, - conversation_id=conversation.id, - user=user, - session=db_session_with_containers, - ) - - deleted = db_session_with_containers.scalar(select(Conversation).where(Conversation.id == conversation.id)) - assert deleted is None - mock_delete_task.delay.assert_called_once_with(conversation.id) - mock_cleanup_task.delay.assert_called_once() - - @patch("services.conversation_service.cleanup_conversation_agent_runtime_session") - @patch("services.conversation_service.delete_conversation_related_data") - def test_delete_ignores_cleanup_enqueue_failure_and_still_retires_runtime_session( - self, - mock_delete_task, - mock_cleanup_task, - db_session_with_containers: Session, - ): - app_model, user = ConversationServiceIntegrationTestDataFactory.create_app_and_account( - db_session_with_containers - ) - conversation = ConversationServiceIntegrationTestDataFactory.create_conversation( - db_session_with_containers, - app_model, - user, - ) - conversation_id = conversation.id - runtime_session = AgentRuntimeSession( - tenant_id=app_model.tenant_id, - app_id=app_model.id, - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id=str(uuid4()), - agent_config_snapshot_id=str(uuid4()), - backend_run_id="backend-run-enqueue-failure", - session_snapshot=CompositorSessionSnapshot(layers=[]).model_dump_json(), - composition_layer_specs='[{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]', - conversation_id=conversation.id, - status=AgentRuntimeSessionStatus.ACTIVE, - ) - db_session_with_containers.add(runtime_session) - db_session_with_containers.commit() - mock_cleanup_task.delay.side_effect = RuntimeError("queue down") - - ConversationService.delete( - app_model=app_model, - conversation_id=conversation_id, - user=user, - session=db_session_with_containers, - ) - - deleted = db_session_with_containers.scalar(select(Conversation).where(Conversation.id == conversation_id)) - assert deleted is None - mock_delete_task.delay.assert_called_once_with(conversation_id) - mock_cleanup_task.delay.assert_called_once() - runtime_session_row = db_session_with_containers.scalar( - select(AgentRuntimeSession).where(AgentRuntimeSession.id == runtime_session.id) - ) - assert runtime_session_row is not None - assert runtime_session_row.status == AgentRuntimeSessionStatus.CLEANED diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py index e722f943820..f9a176c1f57 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_document.py @@ -142,7 +142,9 @@ def test_get_document_queries_by_dataset_and_document_id(db_session_with_contain def test_get_documents_by_ids_returns_empty_for_empty_input(db_session_with_containers: Session): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) - result = DocumentService.get_documents_by_ids(dataset.id, [], session=db_session_with_containers) + result = DocumentService.get_documents_by_ids( + DatasetRefService.create_dataset_ref(dataset), [], session=db_session_with_containers + ) assert result == [] @@ -157,7 +159,9 @@ def test_get_documents_by_ids_uses_single_batch_query(db_session_with_containers position=2, ) - result = DocumentService.get_documents_by_ids(dataset.id, [doc_a.id, doc_b.id], db_session_with_containers) + result = DocumentService.get_documents_by_ids( + DatasetRefService.create_dataset_ref(dataset), [doc_a.id, doc_b.id], db_session_with_containers + ) assert {document.id for document in result} == {doc_a.id, doc_b.id} @@ -319,7 +323,7 @@ def test_get_upload_files_by_document_id_for_zip_download_raises_for_missing_doc ) -def test_get_upload_files_by_document_id_for_zip_download_rejects_cross_tenant_access( +def test_get_upload_files_by_document_id_for_zip_download_hides_cross_tenant_documents( db_session_with_containers: Session, ): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) @@ -335,7 +339,7 @@ def test_get_upload_files_by_document_id_for_zip_download_rejects_cross_tenant_a data_source_info={"upload_file_id": upload_file.id}, ) - with pytest.raises(Forbidden, match="No permission"): + with pytest.raises(NotFound, match="Document not found"): DocumentService._get_upload_files_by_document_id_for_zip_download( dataset_id=dataset.id, document_ids=[document.id], @@ -527,7 +531,7 @@ def test_get_working_documents_by_dataset_id_returns_completed_enabled_unarchive assert [document.id for document in result] == [available_document.id] -def test_get_error_documents_by_dataset_id_returns_error_and_paused_documents(db_session_with_containers: Session): +def test_get_error_documents_by_dataset_ref_returns_error_and_paused_documents(db_session_with_containers: Session): dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers) error_document = DocumentServiceIntegrationFactory.create_document( db_session_with_containers, @@ -547,7 +551,8 @@ def test_get_error_documents_by_dataset_id_returns_error_and_paused_documents(db indexing_status=IndexingStatus.COMPLETED, ) - result = DocumentService.get_error_documents_by_dataset_id(dataset.id, session=db_session_with_containers) + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + result = DocumentService.get_error_documents_by_dataset_ref(dataset_ref, session=db_session_with_containers) assert {document.id for document in result} == {error_document.id, paused_document.id} diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py index 6b32273624b..58a129a15e4 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py @@ -8,7 +8,6 @@ from uuid import uuid4 import pytest from flask import Flask from sqlalchemy.orm import Session -from werkzeug.exceptions import NotFound from core.rag.index_processor.constant.index_type import IndexTechniqueType from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole @@ -21,6 +20,7 @@ from models.dataset import ( DatasetPermissionEnum, ) from models.enums import DataSourceType +from services.dataset_ref_service import DatasetRef, DatasetRefService from services.dataset_service import DatasetCollectionBindingService, DatasetPermissionService, DatasetService from services.errors.account import NoPermissionError @@ -213,7 +213,9 @@ class TestDatasetServicePermissionsAndLifecycle: dataset_id=dataset.id, ) - assert DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) is True + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + + assert DatasetService.dataset_use_check(dataset_ref, session=db_session_with_containers) is True def test_dataset_use_check_returns_false_when_join_missing(self, db_session_with_containers: Session): owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) @@ -223,7 +225,9 @@ class TestDatasetServicePermissionsAndLifecycle: created_by=owner.id, ) - assert DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) is False + dataset_ref = DatasetRefService.create_dataset_ref(dataset) + + assert DatasetService.dataset_use_check(dataset_ref, session=db_session_with_containers) is False def test_check_dataset_permission_rejects_cross_tenant_access(self, db_session_with_containers: Session): owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) @@ -371,13 +375,6 @@ class TestDatasetServicePermissionsAndLifecycle: user=operator, dataset=dataset, session=db_session_with_containers ) - def test_update_dataset_api_status_raises_not_found_for_missing_dataset( - self, flask_app_with_containers: Flask, db_session_with_containers: Session - ): - with flask_app_with_containers.app_context(): - with pytest.raises(NotFound, match="Dataset not found"): - DatasetService.update_dataset_api_status(str(uuid4()), True, session=db_session_with_containers) - def test_update_dataset_api_status_requires_current_user_id(self, db_session_with_containers: Session): owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) dataset = DatasetPermissionIntegrationFactory.create_dataset( @@ -386,10 +383,10 @@ class TestDatasetServicePermissionsAndLifecycle: created_by=owner.id, enable_api=False, ) - - with patch("services.dataset_service.current_user", SimpleNamespace(id=None)): - with pytest.raises(ValueError, match="Current user or current user id not found"): - DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers) + actor = Account(name="missing-id", email="missing-id@example.com") + actor.id = "" + with pytest.raises(ValueError, match="Current user or current user id not found"): + DatasetService.update_dataset_api_status(dataset, True, actor, session=db_session_with_containers) def test_update_dataset_api_status_updates_fields_and_commits(self, db_session_with_containers: Session): owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers) @@ -401,11 +398,8 @@ class TestDatasetServicePermissionsAndLifecycle: ) now = datetime(2026, 4, 14, 18, 0, 0) - with ( - patch("services.dataset_service.current_user", owner), - patch("services.dataset_service.naive_utc_now", return_value=now), - ): - DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers) + with patch("services.dataset_service.naive_utc_now", return_value=now): + DatasetService.update_dataset_api_status(dataset, True, owner, session=db_session_with_containers) db_session_with_containers.refresh(dataset) assert dataset.enable_api is True @@ -419,12 +413,10 @@ class TestDatasetServicePermissionsAndLifecycle: features = SimpleNamespace( billing=SimpleNamespace(enabled=False, subscription=SimpleNamespace(plan="professional")) ) + dataset_ref = DatasetRef(tenant_id=tenant.id, dataset_id=str(uuid4())) - with ( - patch("services.dataset_service.current_user", owner), - patch("services.dataset_service.FeatureService.get_features", return_value=features), - ): - result = DatasetService.get_dataset_auto_disable_logs(str(uuid4()), session=db_session_with_containers) + with patch("services.dataset_service.FeatureService.get_features", return_value=features): + result = DatasetService.get_dataset_auto_disable_logs(dataset_ref, session=db_session_with_containers) assert result == {"document_ids": [], "count": 0} @@ -450,12 +442,10 @@ class TestDatasetServicePermissionsAndLifecycle: features = SimpleNamespace( billing=SimpleNamespace(enabled=True, subscription=SimpleNamespace(plan="professional")) ) + dataset_ref = DatasetRefService.create_dataset_ref(dataset) - with ( - patch("services.dataset_service.current_user", owner), - patch("services.dataset_service.FeatureService.get_features", return_value=features), - ): - result = DatasetService.get_dataset_auto_disable_logs(dataset.id, session=db_session_with_containers) + with patch("services.dataset_service.FeatureService.get_features", return_value=features): + result = DatasetService.get_dataset_auto_disable_logs(dataset_ref, session=db_session_with_containers) assert result["count"] == 2 assert len(result["document_ids"]) == 2 diff --git a/api/tests/test_containers_integration_tests/services/test_feature_service.py b/api/tests/test_containers_integration_tests/services/test_feature_service.py index 7a86af9f410..ca933a08462 100644 --- a/api/tests/test_containers_integration_tests/services/test_feature_service.py +++ b/api/tests/test_containers_integration_tests/services/test_feature_service.py @@ -4,16 +4,16 @@ import pytest from faker import Faker from sqlalchemy.orm import Session -from enums.cloud_plan import CloudPlan -from enums.deployment_edition import DeploymentEdition -from services.feature_service import ( +from enums import CloudPlan, DeploymentEdition +from services.entities.feature_entities import ( FeatureModel, - FeatureService, KnowledgeRateLimitModel, LicenseModel, LicenseStatus, + SSOProtocol, SystemFeatureModel, ) +from services.feature_service import FeatureService class TestFeatureService: @@ -29,7 +29,7 @@ class TestFeatureService: # Setup default mock returns for BillingService mock_billing_service.get_info.return_value = { "enabled": True, - "subscription": {"plan": "pro", "interval": "monthly", "education": True}, + "subscription": {"plan": CloudPlan.PROFESSIONAL, "interval": "monthly", "education": True}, "members": {"size": 5, "limit": 10}, "apps": {"size": 3, "limit": 20}, "vector_space": {"size": 2, "limit": 10}, @@ -41,7 +41,10 @@ class TestFeatureService: "knowledge_rate_limit": {"limit": 100}, } - mock_billing_service.get_knowledge_rate_limit.return_value = {"limit": 100, "subscription_plan": "pro"} + mock_billing_service.get_knowledge_rate_limit.return_value = { + "limit": 100, + "subscription_plan": CloudPlan.PROFESSIONAL, + } # Setup default mock returns for EnterpriseService mock_enterprise_service.get_workspace_info.return_value = { @@ -87,20 +90,19 @@ class TestFeatureService: def test_get_features_success(self, db_session_with_containers: Session, mock_external_service_dependencies): """ - Test successful feature retrieval with billing and enterprise enabled. + Test successful feature retrieval for the Cloud edition. This test verifies: - Proper feature model creation with all required fields - Correct integration with billing service - - Proper enterprise workspace information handling + - Enterprise workspace information remains isolated from Cloud - Return value correctness and structure """ # Arrange: Setup test data with proper config mocking tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = True mock_config.DATASET_OPERATOR_ENABLED = True @@ -115,7 +117,7 @@ class TestFeatureService: # Verify billing features assert result.billing.enabled is True - assert result.billing.subscription.plan == "pro" + assert result.billing.subscription.plan == CloudPlan.PROFESSIONAL assert result.billing.subscription.interval == "monthly" assert result.education.activated is True @@ -145,10 +147,10 @@ class TestFeatureService: assert result.model_load_balancing_enabled is True assert result.knowledge_rate_limit == 100 - # Verify enterprise features - assert result.workspace_members.enabled is True - assert result.workspace_members.size == 5 - assert result.workspace_members.limit == 10 + # Enterprise workspace features are not loaded in Cloud. + assert result.workspace_members.enabled is False + assert result.workspace_members.size == 0 + assert result.workspace_members.limit == 0 # Verify webapp copyright is enabled for non-sandbox plans assert result.webapp_copyright_enabled is True @@ -156,9 +158,7 @@ class TestFeatureService: # Verify mock interactions mock_external_service_dependencies["billing_service"].get_info.assert_called_once_with(tenant_id) - mock_external_service_dependencies["enterprise_service"].get_workspace_info.assert_called_once_with( - tenant_id - ) + mock_external_service_dependencies["enterprise_service"].get_workspace_info.assert_not_called() def test_get_features_sandbox_plan(self, db_session_with_containers: Session, mock_external_service_dependencies): """ @@ -174,8 +174,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = False mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = False @@ -230,7 +229,7 @@ class TestFeatureService: self, db_session_with_containers: Session, mock_external_service_dependencies ): """ - Test successful knowledge rate limit retrieval with billing enabled. + Test successful knowledge rate limit retrieval in the Cloud edition. This test verifies: - Proper knowledge rate limit model creation @@ -242,7 +241,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD # Act: Execute the method under test result = FeatureService.get_knowledge_rate_limit(tenant_id) @@ -254,7 +253,7 @@ class TestFeatureService: # Verify rate limit configuration assert result.enabled is True assert result.limit == 100 - assert result.subscription_plan == "pro" + assert result.subscription_plan == CloudPlan.PROFESSIONAL # Verify mock interactions mock_external_service_dependencies["billing_service"].get_knowledge_rate_limit.assert_called_once_with( @@ -276,7 +275,6 @@ class TestFeatureService: with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = True mock_config.ENABLE_EMAIL_CODE_LOGIN = True mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -301,7 +299,7 @@ class TestFeatureService: # Verify SSO configuration assert result.sso_enforced_for_signin is True - assert result.sso_enforced_for_signin_protocol == "saml" + assert result.sso_enforced_for_signin_protocol is SSOProtocol.SAML # Verify authentication settings assert result.enable_email_code_login is True @@ -349,7 +347,6 @@ class TestFeatureService: # Arrange: Setup test data with exact same config as success test with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = True mock_config.ENABLE_EMAIL_CODE_LOGIN = True mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -383,7 +380,7 @@ class TestFeatureService: # SSO settings should be visible for login page rendering assert result.sso_enforced_for_signin is True - assert result.sso_enforced_for_signin_protocol == "saml" + assert result.sso_enforced_for_signin_protocol is SSOProtocol.SAML # General auth settings should be visible assert result.enable_email_code_login is True @@ -403,7 +400,7 @@ class TestFeatureService: """ # Arrange with patch("services.feature_service.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE # Act result = FeatureService.get_license() @@ -422,7 +419,7 @@ class TestFeatureService: ): """Non-enterprise deployments have no license, so limits are unconstrained.""" with patch("services.feature_service.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY result = FeatureService.get_license() @@ -447,7 +444,6 @@ class TestFeatureService: # Arrange: Setup basic config mock (no enterprise) with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY - mock_config.ENTERPRISE_ENABLED = False mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = True mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -483,11 +479,11 @@ class TestFeatureService: # Verify marketplace configuration assert result.enable_marketplace is False - def test_get_features_billing_disabled( + def test_get_features_community_edition( self, db_session_with_containers: Session, mock_external_service_dependencies ): """ - Test feature retrieval when billing is disabled. + Test feature retrieval for the Community edition. This test verifies: - Proper feature model creation without billing @@ -495,10 +491,9 @@ class TestFeatureService: - Default configuration values - Return value correctness and structure """ - # Arrange: Setup billing disabled mock + # Arrange: Use the Community edition. with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = False - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = True mock_config.DATASET_OPERATOR_ENABLED = True @@ -540,20 +535,20 @@ class TestFeatureService: assert result.workspace_members.enabled is False assert result.webapp_copyright_enabled is False - def test_get_knowledge_rate_limit_billing_disabled( + def test_get_knowledge_rate_limit_community_edition( self, db_session_with_containers: Session, mock_external_service_dependencies ): """ - Test knowledge rate limit retrieval when billing is disabled. + Test knowledge rate limit retrieval for the Community edition. This test verifies: - Proper knowledge rate limit model creation without billing - Default rate limit configuration - Return value correctness and structure """ - # Arrange: Setup billing disabled mock + # Arrange: Use the Community edition. with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY tenant_id = self._create_test_tenant_id() @@ -567,7 +562,7 @@ class TestFeatureService: # Verify default configuration assert result.enabled is False assert result.limit == 10 - assert result.subscription_plan == "" # Empty string when billing is disabled + assert result.subscription_plan == "" # Verify no billing service calls mock_external_service_dependencies["billing_service"].get_knowledge_rate_limit.assert_not_called() @@ -576,7 +571,7 @@ class TestFeatureService: self, db_session_with_containers: Session, mock_external_service_dependencies ): """ - Test feature retrieval with enterprise enabled but billing disabled. + Test feature retrieval for the Enterprise edition. This test verifies: - Proper feature model creation with enterprise only @@ -586,8 +581,7 @@ class TestFeatureService: """ # Arrange: Setup enterprise only mock with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = False - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_config.CAN_REPLACE_LOGO = False mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = False @@ -602,7 +596,7 @@ class TestFeatureService: assert result is not None assert isinstance(result, FeatureModel) - # Verify billing is disabled + # Cloud billing is not loaded in the Enterprise edition. assert result.billing.enabled is False # Verify enterprise features @@ -633,11 +627,11 @@ class TestFeatureService: ) mock_external_service_dependencies["billing_service"].get_info.assert_not_called() - def test_get_system_features_enterprise_disabled( + def test_get_system_features_community_edition( self, db_session_with_containers: Session, mock_external_service_dependencies ): """ - Test system features retrieval when enterprise is disabled. + Test system features retrieval for the Community edition. This test verifies: - Proper system feature model creation without enterprise @@ -645,10 +639,9 @@ class TestFeatureService: - Default configuration values - Return value correctness and structure """ - # Arrange: Setup enterprise disabled mock + # Arrange: Use the Community edition. with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY - mock_config.ENTERPRISE_ENABLED = False mock_config.MARKETPLACE_ENABLED = True mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -692,18 +685,17 @@ class TestFeatureService: def test_get_features_no_tenant_id(self, db_session_with_containers: Session, mock_external_service_dependencies): """ - Test feature retrieval without tenant ID (billing disabled). + Test Cloud feature retrieval without a tenant ID. This test verifies: - Proper feature model creation without tenant ID - - Correct handling when billing is disabled + - Billing data is not loaded without a tenant ID - Default configuration values - Return value correctness and structure """ # Arrange: Setup no tenant ID scenario with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = True @@ -716,7 +708,7 @@ class TestFeatureService: assert result is not None assert isinstance(result, FeatureModel) - # Verify billing is disabled due to no tenant ID + # Billing data is not loaded without a tenant ID. assert result.billing.enabled is False # Verify environment-based features @@ -752,8 +744,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = True @@ -761,7 +752,7 @@ class TestFeatureService: mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, - "subscription": {"plan": "basic", "interval": "yearly"}, + "subscription": {"plan": CloudPlan.PROFESSIONAL, "interval": "yearly"}, # Missing members, apps, vector_space, etc. } @@ -774,7 +765,7 @@ class TestFeatureService: # Verify billing features assert result.billing.enabled is True - assert result.billing.subscription.plan == "basic" + assert result.billing.subscription.plan == CloudPlan.PROFESSIONAL assert result.billing.subscription.interval == "yearly" # Verify default values for missing billing info @@ -791,7 +782,7 @@ class TestFeatureService: assert result.knowledge_rate_limit == 10 assert result.docs_processing == "standard" - # Verify basic plan restrictions (non-sandbox plans have webapp copyright enabled) + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True @@ -814,8 +805,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = True @@ -823,7 +813,7 @@ class TestFeatureService: mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, - "subscription": {"plan": "pro", "interval": "monthly"}, + "subscription": {"plan": CloudPlan.PROFESSIONAL, "interval": "monthly"}, "vector_space": {"size": 0, "limit": 0}, "apps": {"size": 5, "limit": 10}, } @@ -843,7 +833,7 @@ class TestFeatureService: assert result.apps.size == 5 assert result.apps.limit == 10 - # Verify pro plan features + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True @@ -875,7 +865,6 @@ class TestFeatureService: # Arrange: Setup edge case webapp auth mock with proper config with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -900,7 +889,7 @@ class TestFeatureService: assert result.webapp_auth.allow_sso is False assert result.webapp_auth.allow_email_code_login is True assert result.webapp_auth.allow_email_password_login is False - assert result.webapp_auth.sso_config.protocol == "" + assert result.webapp_auth.sso_config.protocol is None # Verify enterprise features assert result.branding.enabled is True @@ -909,7 +898,7 @@ class TestFeatureService: # Verify default values for missing enterprise info assert result.sso_enforced_for_signin is False - assert result.sso_enforced_for_signin_protocol == "" + assert result.sso_enforced_for_signin_protocol is None assert result.enable_email_code_login is False assert result.enable_email_password_login is True assert result.is_allow_register is False @@ -933,8 +922,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = True @@ -942,7 +930,7 @@ class TestFeatureService: mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, - "subscription": {"plan": "basic", "interval": "yearly"}, + "subscription": {"plan": CloudPlan.PROFESSIONAL, "interval": "yearly"}, "members": {"size": 10, "limit": 10}, "vector_space": {"size": 3, "limit": 5}, } @@ -962,7 +950,7 @@ class TestFeatureService: assert result.vector_space.size == 3 assert result.vector_space.limit == 5 - # Verify basic plan features (non-sandbox plans have webapp copyright enabled) + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True @@ -995,7 +983,6 @@ class TestFeatureService: # Test case 1: Official only scope with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1019,7 +1006,6 @@ class TestFeatureService: # Test case 2: All plugins scope with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1040,7 +1026,6 @@ class TestFeatureService: # Test case 3: Specific partners scope with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1064,7 +1049,6 @@ class TestFeatureService: # Test case 4: None scope with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1101,8 +1085,7 @@ class TestFeatureService: } with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = False - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE # Act: Execute the method under test result = FeatureService.get_features(tenant_id) @@ -1138,7 +1121,7 @@ class TestFeatureService: """ # Arrange: Setup inactive license mock with proper config with patch("services.feature_service.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1188,7 +1171,6 @@ class TestFeatureService: # Arrange: Setup partial enterprise info mock with proper config with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1218,7 +1200,7 @@ class TestFeatureService: # Verify SSO configuration assert result.sso_enforced_for_signin is True - assert result.sso_enforced_for_signin_protocol == "" + assert result.sso_enforced_for_signin_protocol is None # Verify branding configuration (partial) assert result.branding.application_title == "Partial Enterprise" @@ -1230,7 +1212,7 @@ class TestFeatureService: assert result.webapp_auth.allow_sso is False assert result.webapp_auth.allow_email_code_login is False assert result.webapp_auth.allow_email_password_login is False - assert result.webapp_auth.sso_config.protocol == "" + assert result.webapp_auth.sso_config.protocol is None # Verify default license status assert result.license.status == "none" @@ -1260,8 +1242,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = True @@ -1269,7 +1250,7 @@ class TestFeatureService: mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, - "subscription": {"plan": "enterprise", "interval": "yearly"}, + "subscription": {"plan": CloudPlan.TEAM, "interval": "yearly"}, "members": {"size": 0, "limit": 0}, "apps": {"size": 0, "limit": -1}, "vector_space": {"size": 0, "limit": 999999}, @@ -1296,7 +1277,7 @@ class TestFeatureService: assert result.annotation_quota_limit.size == 0 assert result.annotation_quota_limit.limit == 1 - # Verify enterprise plan features + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True @@ -1318,7 +1299,6 @@ class TestFeatureService: # Arrange: Setup edge case protocols mock with proper config with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1342,8 +1322,8 @@ class TestFeatureService: assert isinstance(result, SystemFeatureModel) # Verify edge case protocols - assert result.sso_enforced_for_signin_protocol == "" - assert result.webapp_auth.sso_config.protocol == " " + assert result.sso_enforced_for_signin_protocol is None + assert result.webapp_auth.sso_config.protocol is None # Verify webapp auth configuration assert result.webapp_auth.allow_sso is True @@ -1374,7 +1354,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, - "subscription": {"plan": "education", "interval": "semester", "education": True}, + "subscription": {"plan": CloudPlan.PROFESSIONAL, "interval": "semester", "education": True}, "members": {"size": 100, "limit": 200}, "apps": {"size": 50, "limit": 100}, "vector_space": {"size": 20, "limit": 50}, @@ -1383,6 +1363,7 @@ class TestFeatureService: } with patch("services.feature_service.dify_config") as mock_config: + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.EDUCATION_ENABLED = True # Act: Execute the method under test @@ -1396,7 +1377,7 @@ class TestFeatureService: assert result.education.enabled is True assert result.education.activated is True - # Verify education plan limits + # Verify education subscription limits. assert result.members.size == 100 assert result.members.limit == 200 assert result.apps.size == 50 @@ -1408,54 +1389,13 @@ class TestFeatureService: assert result.annotation_quota_limit.size == 200 assert result.annotation_quota_limit.limit == 500 - # Verify education plan features + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True # Verify mock interactions mock_external_service_dependencies["billing_service"].get_info.assert_called_once_with(tenant_id) - def test_license_limitation_model_is_available( - self, db_session_with_containers: Session, mock_external_service_dependencies - ): - """ - Test LicenseLimitationModel.is_available method with various scenarios. - - This test verifies: - - Proper quota availability calculation - - Correct handling of unlimited limits - - Proper handling of disabled limits - - Return value correctness for different scenarios - """ - from services.feature_service import LicenseLimitationModel - - # Test case 1: Limit disabled - disabled_limit = LicenseLimitationModel(enabled=False, size=5, limit=10) - assert disabled_limit.is_available(3) is True - assert disabled_limit.is_available(10) is True - - # Test case 2: Unlimited limit - unlimited_limit = LicenseLimitationModel(enabled=True, size=5, limit=0) - assert unlimited_limit.is_available(3) is True - assert unlimited_limit.is_available(100) is True - - # Test case 3: Available quota - available_limit = LicenseLimitationModel(enabled=True, size=5, limit=10) - assert available_limit.is_available(3) is True - assert available_limit.is_available(5) is True - assert available_limit.is_available(1) is True - - # Test case 4: Insufficient quota - insufficient_limit = LicenseLimitationModel(enabled=True, size=8, limit=10) - assert insufficient_limit.is_available(3) is False - assert insufficient_limit.is_available(2) is True - assert insufficient_limit.is_available(1) is True - - # Test case 5: Exact quota usage - exact_limit = LicenseLimitationModel(enabled=True, size=7, limit=10) - assert exact_limit.is_available(3) is True - assert exact_limit.is_available(3) is True - def test_get_features_workspace_members_disabled( self, db_session_with_containers: Session, mock_external_service_dependencies ): @@ -1475,8 +1415,7 @@ class TestFeatureService: } with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = False - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE # Act: Execute the method under test result = FeatureService.get_features(tenant_id) @@ -1510,7 +1449,7 @@ class TestFeatureService: """ # Arrange: Setup expired license mock with proper config with patch("services.feature_service.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1561,8 +1500,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = True mock_config.DATASET_OPERATOR_ENABLED = True @@ -1570,7 +1508,7 @@ class TestFeatureService: mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, - "subscription": {"plan": "premium", "interval": "monthly"}, + "subscription": {"plan": CloudPlan.TEAM, "interval": "monthly"}, "docs_processing": "advanced", "can_replace_logo": True, "model_load_balancing_enabled": True, @@ -1588,7 +1526,7 @@ class TestFeatureService: assert result.can_replace_logo is True assert result.model_load_balancing_enabled is True - # Verify premium plan features + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True @@ -1618,7 +1556,6 @@ class TestFeatureService: # Arrange: Setup edge case branding mock with proper config with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1657,7 +1594,7 @@ class TestFeatureService: # Verify default values for missing enterprise info assert result.sso_enforced_for_signin is False - assert result.sso_enforced_for_signin_protocol == "" + assert result.sso_enforced_for_signin_protocol is None assert result.enable_email_code_login is False assert result.enable_email_password_login is True assert result.is_allow_register is False @@ -1681,8 +1618,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = True @@ -1690,7 +1626,7 @@ class TestFeatureService: mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, - "subscription": {"plan": "enterprise", "interval": "yearly"}, + "subscription": {"plan": CloudPlan.TEAM, "interval": "yearly"}, "annotation_quota_limit": {"size": 999, "limit": 1000}, "knowledge_rate_limit": {"limit": 500}, } @@ -1709,7 +1645,7 @@ class TestFeatureService: # Verify knowledge rate limit assert result.knowledge_rate_limit == 500 - # Verify enterprise plan features + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True @@ -1743,8 +1679,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = True @@ -1752,7 +1687,7 @@ class TestFeatureService: mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, - "subscription": {"plan": "pro", "interval": "monthly"}, + "subscription": {"plan": CloudPlan.PROFESSIONAL, "interval": "monthly"}, "documents_upload_quota": { "size": 0, # Edge case: zero current size "limit": 0, # Edge case: zero limit @@ -1774,7 +1709,7 @@ class TestFeatureService: # Verify knowledge rate limit assert result.knowledge_rate_limit == 100 - # Verify pro plan features + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True @@ -1807,7 +1742,6 @@ class TestFeatureService: # Arrange: Setup lost license mock with proper config with patch("services.feature_service.dify_config") as mock_config: mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - mock_config.ENTERPRISE_ENABLED = True mock_config.MARKETPLACE_ENABLED = False mock_config.ENABLE_EMAIL_CODE_LOGIN = False mock_config.ENABLE_EMAIL_PASSWORD_LOGIN = True @@ -1835,7 +1769,7 @@ class TestFeatureService: # Verify default values for missing enterprise info assert result.sso_enforced_for_signin is False - assert result.sso_enforced_for_signin_protocol == "" + assert result.sso_enforced_for_signin_protocol is None assert result.enable_email_code_login is False assert result.enable_email_password_login is True assert result.is_allow_register is False @@ -1860,8 +1794,7 @@ class TestFeatureService: tenant_id = self._create_test_tenant_id() with patch("services.feature_service.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = True mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = True @@ -1870,7 +1803,7 @@ class TestFeatureService: mock_external_service_dependencies["billing_service"].get_info.return_value = { "enabled": True, "subscription": { - "plan": "pro", + "plan": CloudPlan.PROFESSIONAL, "interval": "monthly", "education": False, # Education explicitly disabled }, @@ -1890,7 +1823,7 @@ class TestFeatureService: # Verify knowledge rate limit assert result.knowledge_rate_limit == 100 - # Verify pro plan features + # Verify paid plan behavior. assert result.webapp_copyright_enabled is True assert result.is_allow_transfer_workspace is True diff --git a/api/tests/test_containers_integration_tests/services/test_messages_clean_service.py b/api/tests/test_containers_integration_tests/services/test_messages_clean_service.py index 1a1efe03371..2af90d0ef2a 100644 --- a/api/tests/test_containers_integration_tests/services/test_messages_clean_service.py +++ b/api/tests/test_containers_integration_tests/services/test_messages_clean_service.py @@ -4,13 +4,13 @@ import datetime import json import uuid from decimal import Decimal -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest from faker import Faker from sqlalchemy.orm import Session -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from extensions.ext_redis import redis_client from graphon.file import FileType from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole @@ -100,15 +100,21 @@ class TestMessagesCleanServiceIntegration: yield mock @pytest.fixture - def mock_billing_enabled(self): - """Mock BILLING_ENABLED to be True.""" - with patch("services.retention.conversation.messages_clean_policy.dify_config.BILLING_ENABLED", True): + def cloud_edition(self): + """Use the Cloud deployment edition.""" + with patch( + "services.retention.conversation.messages_clean_policy.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.CLOUD, + ): yield @pytest.fixture - def mock_billing_disabled(self): - """Mock BILLING_ENABLED to be False.""" - with patch("services.retention.conversation.messages_clean_policy.dify_config.BILLING_ENABLED", False): + def non_cloud_edition(self): + """Use a non-Cloud deployment edition.""" + with patch( + "services.retention.conversation.messages_clean_policy.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ): yield def _create_account_and_tenant(self, db_session_with_containers: Session, plan: str = CloudPlan.SANDBOX): @@ -311,11 +317,11 @@ class TestMessagesCleanServiceIntegration: ) db_session_with_containers.add(resource) - def test_billing_disabled_deletes_all_messages_in_time_range( - self, db_session_with_containers: Session, mock_billing_disabled + def test_non_cloud_edition_deletes_all_messages_in_time_range( + self, db_session_with_containers: Session, non_cloud_edition ): """Test that BillingDisabledPolicy deletes all messages within time range regardless of tenant plan.""" - # Arrange - Create tenant with messages (plan doesn't matter for billing disabled) + # Arrange - Create tenant with messages; plans do not apply outside Cloud. account, tenant = self._create_account_and_tenant(db_session_with_containers, plan=CloudPlan.SANDBOX) app = self._create_app(db_session_with_containers, tenant, account) conv = self._create_conversation(db_session_with_containers, app) @@ -378,9 +384,7 @@ class TestMessagesCleanServiceIntegration: == 1 ) - def test_no_messages_returns_empty_stats( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist - ): + def test_no_messages_returns_empty_stats(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test cleaning when there are no messages to delete (B1).""" # Arrange end_before = datetime.datetime.now() - datetime.timedelta(days=30) @@ -405,9 +409,7 @@ class TestMessagesCleanServiceIntegration: assert stats["filtered_messages"] == 0 assert stats["total_deleted"] == 0 - def test_mixed_sandbox_and_paid_tenants( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist - ): + def test_mixed_sandbox_and_paid_tenants(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test cleaning with mixed sandbox and paid tenants (B2).""" # Arrange - Create sandbox tenants with expired messages sandbox_tenants = [] @@ -501,7 +503,7 @@ class TestMessagesCleanServiceIntegration: ) def test_cursor_pagination_multiple_batches( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist + self, db_session_with_containers: Session, cloud_edition, mock_whitelist ): """Test cursor pagination works correctly across multiple batches (B3).""" # Arrange - Create sandbox tenant with messages that will span multiple batches @@ -550,7 +552,7 @@ class TestMessagesCleanServiceIntegration: # All messages should be deleted assert db_session_with_containers.query(Message).where(Message.id.in_(message_ids)).count() == 0 - def test_dry_run_does_not_delete(self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist): + def test_dry_run_does_not_delete(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test dry_run mode does not delete messages (B4).""" # Arrange account, tenant = self._create_account_and_tenant(db_session_with_containers, plan=CloudPlan.SANDBOX) @@ -599,9 +601,7 @@ class TestMessagesCleanServiceIntegration: == 3 ) - def test_partial_plan_data_safe_default( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist - ): + def test_partial_plan_data_safe_default(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test when billing returns partial data, unknown tenants are preserved (B5).""" # Arrange - Create 3 tenants tenants_data = [] @@ -668,9 +668,7 @@ class TestMessagesCleanServiceIntegration: db_session_with_containers.query(Message).where(Message.id == tenants_data[2]["message_id"]).count() == 1 ) # Unknown tenant's message preserved (safe default) - def test_empty_plan_data_skips_deletion( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist - ): + def test_empty_plan_data_skips_deletion(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test when billing returns empty data, skip deletion entirely (B6).""" # Arrange account, tenant = self._create_account_and_tenant(db_session_with_containers, plan=CloudPlan.SANDBOX) @@ -705,9 +703,7 @@ class TestMessagesCleanServiceIntegration: # Message should still exist (safe default - don't delete if plan is unknown) assert db_session_with_containers.query(Message).where(Message.id == msg_id).count() == 1 - def test_time_range_boundary_behavior( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist - ): + def test_time_range_boundary_behavior(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test that messages are correctly filtered by [start_from, end_before) time range (B7).""" # Arrange account, tenant = self._create_account_and_tenant(db_session_with_containers, plan=CloudPlan.SANDBOX) @@ -798,7 +794,7 @@ class TestMessagesCleanServiceIntegration: # After range, kept assert db_session_with_containers.query(Message).where(Message.id == msg_after_id).count() == 1 - def test_grace_period_scenarios(self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist): + def test_grace_period_scenarios(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test cleaning with different graceful period scenarios (B8).""" # Arrange - Create 5 different tenants with different plan and expiration scenarios now_timestamp = int(datetime.datetime.now(datetime.UTC).timestamp()) @@ -920,7 +916,7 @@ class TestMessagesCleanServiceIntegration: ) # Professional plan, kept assert db_session_with_containers.query(Message).where(Message.id == msg5_id).count() == 1 # At boundary, kept - def test_tenant_whitelist(self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist): + def test_tenant_whitelist(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test that whitelisted tenants' messages are not deleted (B9).""" # Arrange - Create 3 sandbox tenants with expired messages tenants_data = [] @@ -989,9 +985,7 @@ class TestMessagesCleanServiceIntegration: # Verify tenant2's message was deleted (not whitelisted) assert db_session_with_containers.query(Message).where(Message.id == tenants_data[2]["message_id"]).count() == 0 - def test_from_days_cleans_old_messages( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist - ): + def test_from_days_cleans_old_messages(self, db_session_with_containers: Session, cloud_edition, mock_whitelist): """Test from_days correctly cleans messages older than N days (B11).""" # Arrange account, tenant = self._create_account_and_tenant(db_session_with_containers, plan=CloudPlan.SANDBOX) @@ -1054,7 +1048,7 @@ class TestMessagesCleanServiceIntegration: assert db_session_with_containers.query(Message).where(Message.id.in_(recent_msg_ids)).count() == 2 def test_whitelist_precedence_over_grace_period( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist + self, db_session_with_containers: Session, cloud_edition, mock_whitelist ): """Test that whitelist takes precedence over grace period logic.""" # Arrange - Create 2 sandbox tenants @@ -1123,7 +1117,7 @@ class TestMessagesCleanServiceIntegration: ) # Within grace period def test_empty_whitelist_deletes_eligible_messages( - self, db_session_with_containers: Session, mock_billing_enabled, mock_whitelist + self, db_session_with_containers: Session, cloud_edition, mock_whitelist ): """Test that empty whitelist behaves as no whitelist (all eligible messages deleted).""" # Arrange - Create sandbox tenant with expired messages @@ -1172,65 +1166,8 @@ class TestMessagesCleanServiceIntegration: # Verify all messages were deleted assert db_session_with_containers.query(Message).where(Message.id.in_(msg_ids)).count() == 0 - def test_from_time_range_validation(self): - """Test that from_time_range raises ValueError for invalid inputs.""" - policy = MagicMock(spec=BillingDisabledPolicy) - now = datetime.datetime.now() - - with pytest.raises(ValueError, match="start_from .* must be less than end_before"): - MessagesCleanService.from_time_range(policy, now, now) - - with pytest.raises(ValueError, match="batch_size .* must be greater than 0"): - MessagesCleanService.from_time_range(policy, now - datetime.timedelta(days=1), now, batch_size=0) - - def test_from_time_range_success(self): - """Test that from_time_range creates a service with correct parameters.""" - policy = MagicMock(spec=BillingDisabledPolicy) - start = datetime.datetime(2024, 1, 1) - end = datetime.datetime(2024, 2, 1) - - service = MessagesCleanService.from_time_range(policy, start, end) - assert service._start_from == start - assert service._end_before == end - - def test_from_days_validation(self): - """Test that from_days raises ValueError for invalid inputs.""" - policy = MagicMock(spec=BillingDisabledPolicy) - - with pytest.raises(ValueError, match="days .* must be greater than or equal to 0"): - MessagesCleanService.from_days(policy, days=-1) - - with pytest.raises(ValueError, match="batch_size .* must be greater than 0"): - MessagesCleanService.from_days(policy, days=30, batch_size=0) - - def test_from_days_success(self): - """Test that from_days creates a service with correct parameters.""" - policy = MagicMock(spec=BillingDisabledPolicy) - - with patch("services.retention.conversation.messages_clean_service.naive_utc_now") as mock_now: - fixed_now = datetime.datetime(2024, 6, 1) - mock_now.return_value = fixed_now - - service = MessagesCleanService.from_days(policy, days=10) - assert service._start_from is None - assert service._end_before == fixed_now - datetime.timedelta(days=10) - def test_batch_delete_message_relations_empty(self, db_session_with_containers: Session): """Test that batch_delete_message_relations with empty list does nothing.""" # Get execute call count before MessagesCleanService._batch_delete_message_relations(db_session_with_containers, []) # No exception means success — empty list is a no-op - - def test_run_calls_clean_messages(self): - """Test that run() delegates to _clean_messages_by_time_range.""" - policy = MagicMock(spec=BillingDisabledPolicy) - service = MessagesCleanService( - policy=policy, - end_before=datetime.datetime.now(), - batch_size=10, - ) - with patch.object(service, "_clean_messages_by_time_range") as mock_clean: - mock_clean.return_value = {"total_deleted": 5} - result = service.run() - assert result == {"total_deleted": 5} - mock_clean.assert_called_once() diff --git a/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py b/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py index a9399985307..b38bc6f4862 100644 --- a/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py +++ b/api/tests/test_containers_integration_tests/services/test_metadata_partial_update.py @@ -1,16 +1,18 @@ from __future__ import annotations +from concurrent.futures import ThreadPoolExecutor +from threading import Barrier from unittest.mock import patch from uuid import uuid4 import pytest from flask import Flask -from sqlalchemy import select -from sqlalchemy.orm import Session +from sqlalchemy import event, select +from sqlalchemy.orm import Session, sessionmaker from models import Account, Tenant -from models.dataset import Dataset, DatasetMetadataBinding, Document -from models.enums import DataSourceType, DocumentCreatedFrom +from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding, Document +from models.enums import DatasetMetadataType, DataSourceType, DocumentCreatedFrom from services.entities.knowledge_entities.knowledge_entities import ( DocumentMetadataOperation, MetadataDetail, @@ -54,6 +56,19 @@ def _create_document( return document +def _create_metadata(db_session: Session, *, dataset: Dataset, created_by: str, name: str) -> DatasetMetadata: + metadata = DatasetMetadata( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + type=DatasetMetadataType.STRING, + name=name, + created_by=created_by, + ) + db_session.add(metadata) + db_session.commit() + return metadata + + class TestMetadataPartialUpdate: @pytest.fixture def tenant_id(self) -> str: @@ -87,10 +102,12 @@ class TestMetadataPartialUpdate: doc_metadata={"existing_key": "existing_value"}, ) - meta_id = str(uuid4()) + metadata = _create_metadata( + db_session_with_containers, dataset=dataset, created_by=current_account.id, name="new_key" + ) operation = DocumentMetadataOperation( document_id=document.id, - metadata_list=[MetadataDetail(id=meta_id, name="new_key", value="new_value")], + metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value="new_value")], partial_update=True, ) metadata_args = MetadataOperationData(operation_data=[operation]) @@ -105,6 +122,86 @@ class TestMetadataPartialUpdate: assert updated_doc.doc_metadata["existing_key"] == "existing_value" assert updated_doc.doc_metadata["new_key"] == "new_value" + def test_concurrent_partial_updates_keep_both_values( + self, + db_session_with_containers: Session, + tenant_id: str, + current_account: Account, + ) -> None: + dataset = _create_dataset(db_session_with_containers, tenant_id=tenant_id) + document = _create_document( + db_session_with_containers, + dataset_id=dataset.id, + tenant_id=tenant_id, + doc_metadata={}, + ) + metadatas = [ + _create_metadata( + db_session_with_containers, + dataset=dataset, + created_by=current_account.id, + name=name, + ) + for name in ("author", "region") + ] + operations = [ + MetadataOperationData( + operation_data=[ + DocumentMetadataOperation( + document_id=document.id, + metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value=value)], + partial_update=True, + ) + ] + ) + for metadata, value in zip(metadatas, ("Alice", "EU"), strict=True) + ] + dataset_id = dataset.id + document_id = document.id + session_factory = sessionmaker(bind=db_session_with_containers.get_bind()) + engine = db_session_with_containers.get_bind() + locked_select_barrier = Barrier(2) + locked_selects: list[None] = [] + + def synchronize_locked_select( + _conn: object, + _cursor: object, + statement: str, + _parameters: object, + _context: object, + _executemany: bool, + ) -> None: + if "FROM documents" in statement and "FOR UPDATE" in statement: + locked_selects.append(None) + locked_select_barrier.wait(timeout=10) + + def update(metadata_args: MetadataOperationData) -> None: + with session_factory() as session: + owned_dataset = session.get(Dataset, dataset_id) + assert owned_dataset is not None + MetadataService.update_documents_metadata( + owned_dataset, metadata_args, current_account, session=session + ) + + event.listen(engine, "before_cursor_execute", synchronize_locked_select) + try: + with ( + patch.object(MetadataService, "knowledge_base_metadata_lock_check"), + patch("services.metadata_service.redis_client.delete"), + ThreadPoolExecutor(max_workers=2) as executor, + ): + futures = [executor.submit(update, operation) for operation in operations] + for future in futures: + future.result(timeout=20) + finally: + event.remove(engine, "before_cursor_execute", synchronize_locked_select) + + db_session_with_containers.expire_all() + updated_doc = db_session_with_containers.get(Document, document_id) + assert updated_doc is not None + assert len(locked_selects) == 2 + assert updated_doc.doc_metadata == {"author": "Alice", "region": "EU"} + def test_full_update_replaces_metadata( self, flask_app_with_containers: Flask, @@ -120,10 +217,12 @@ class TestMetadataPartialUpdate: doc_metadata={"existing_key": "existing_value"}, ) - meta_id = str(uuid4()) + metadata = _create_metadata( + db_session_with_containers, dataset=dataset, created_by=current_account.id, name="new_key" + ) operation = DocumentMetadataOperation( document_id=document.id, - metadata_list=[MetadataDetail(id=meta_id, name="new_key", value="new_value")], + metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value="new_value")], partial_update=False, ) metadata_args = MetadataOperationData(operation_data=[operation]) @@ -154,12 +253,14 @@ class TestMetadataPartialUpdate: doc_metadata={"existing_key": "existing_value"}, ) - meta_id = str(uuid4()) + metadata = _create_metadata( + db_session_with_containers, dataset=dataset, created_by=current_account.id, name="existing_key" + ) existing_binding = DatasetMetadataBinding( tenant_id=tenant_id, dataset_id=dataset.id, document_id=document.id, - metadata_id=meta_id, + metadata_id=metadata.id, created_by=user_id, ) db_session_with_containers.add(existing_binding) @@ -167,7 +268,7 @@ class TestMetadataPartialUpdate: operation = DocumentMetadataOperation( document_id=document.id, - metadata_list=[MetadataDetail(id=meta_id, name="existing_key", value="existing_value")], + metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value="existing_value")], partial_update=True, ) metadata_args = MetadataOperationData(operation_data=[operation]) @@ -180,7 +281,7 @@ class TestMetadataPartialUpdate: bindings = db_session_with_containers.scalars( select(DatasetMetadataBinding).where( DatasetMetadataBinding.document_id == document.id, - DatasetMetadataBinding.metadata_id == meta_id, + DatasetMetadataBinding.metadata_id == metadata.id, ) ).all() assert len(bindings) == 1 @@ -200,10 +301,12 @@ class TestMetadataPartialUpdate: doc_metadata={"existing_key": "existing_value"}, ) - meta_id = str(uuid4()) + metadata = _create_metadata( + db_session_with_containers, dataset=dataset, created_by=current_account.id, name="key" + ) operation = DocumentMetadataOperation( document_id=document.id, - metadata_list=[MetadataDetail(id=meta_id, name="key", value="value")], + metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value="value")], partial_update=True, ) metadata_args = MetadataOperationData(operation_data=[operation]) diff --git a/api/tests/test_containers_integration_tests/services/test_metadata_service.py b/api/tests/test_containers_integration_tests/services/test_metadata_service.py index 00afe7f8467..d0b8a71cd56 100644 --- a/api/tests/test_containers_integration_tests/services/test_metadata_service.py +++ b/api/tests/test_containers_integration_tests/services/test_metadata_service.py @@ -10,8 +10,14 @@ from core.rag.index_processor.constant.built_in_field import BuiltInField from core.rag.index_processor.constant.index_type import IndexStructureType from models import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding, Document -from models.enums import DataSourceType, DocumentCreatedFrom -from services.entities.knowledge_entities.knowledge_entities import MetadataArgs +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus +from services.entities.knowledge_entities.knowledge_entities import ( + DocumentMetadataOperation, + MetadataArgs, + MetadataDetail, + MetadataOperationData, +) +from services.errors.metadata import MetadataResourceNotFoundError from services.metadata_service import MetadataService @@ -300,7 +306,7 @@ class TestMetadataService: # Act: Execute the method under test new_name = "new_name" result = MetadataService.update_metadata_name( - dataset.id, metadata.id, new_name, account, tenant.id, session=db_session_with_containers + dataset, metadata.id, new_name, account, session=db_session_with_containers ) # Assert: Verify the expected outcomes @@ -340,7 +346,7 @@ class TestMetadataService: # Act & Assert: Verify proper error handling with pytest.raises(ValueError, match="Metadata name cannot exceed 255 characters."): MetadataService.update_metadata_name( - dataset.id, metadata.id, long_name, account, tenant.id, session=db_session_with_containers + dataset, metadata.id, long_name, account, session=db_session_with_containers ) def test_update_metadata_name_already_exists( @@ -371,7 +377,7 @@ class TestMetadataService: # Try to update first metadata with second metadata's name with pytest.raises(ValueError, match="Metadata name already exists."): MetadataService.update_metadata_name( - dataset.id, first_metadata.id, "second_metadata", account, tenant.id, session=db_session_with_containers + dataset, first_metadata.id, "second_metadata", account, session=db_session_with_containers ) def test_update_metadata_name_conflicts_with_built_in_field( @@ -399,7 +405,7 @@ class TestMetadataService: with pytest.raises(ValueError, match="Metadata name already exists in Built-in fields."): MetadataService.update_metadata_name( - dataset.id, metadata.id, built_in_field_name, account, tenant.id, session=db_session_with_containers + dataset, metadata.id, built_in_field_name, account, session=db_session_with_containers ) def test_update_metadata_name_not_found( @@ -424,7 +430,7 @@ class TestMetadataService: # Act: Execute the method under test result = MetadataService.update_metadata_name( - dataset.id, fake_metadata_id, new_name, account, tenant.id, session=db_session_with_containers + dataset, fake_metadata_id, new_name, account, session=db_session_with_containers ) # Assert: Verify the method returns None when metadata is not found @@ -451,7 +457,7 @@ class TestMetadataService: ) # Act: Execute the method under test - result = MetadataService.delete_metadata(dataset.id, metadata.id, session=db_session_with_containers) + result = MetadataService.delete_metadata(dataset, metadata.id, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -482,7 +488,7 @@ class TestMetadataService: fake_metadata_id = str(uuid.uuid4()) # Use valid UUID format # Act: Execute the method under test - result = MetadataService.delete_metadata(dataset.id, fake_metadata_id, session=db_session_with_containers) + result = MetadataService.delete_metadata(dataset, fake_metadata_id, session=db_session_with_containers) # Assert: Verify the method returns None when metadata is not found assert result is None @@ -528,7 +534,7 @@ class TestMetadataService: db_session_with_containers.commit() # Act: Execute the method under test - result = MetadataService.delete_metadata(dataset.id, metadata.id, session=db_session_with_containers) + result = MetadataService.delete_metadata(dataset, metadata.id, session=db_session_with_containers) # Assert: Verify the expected outcomes assert result is not None @@ -540,6 +546,69 @@ class TestMetadataService: # Note: The service attempts to update document metadata but may not succeed # due to mock configuration. The main functionality (metadata deletion) is verified. + @pytest.mark.parametrize("operation", ["rename", "delete"]) + @pytest.mark.parametrize("binding_owner", ["metadata", "document"]) + def test_metadata_changes_ignore_historical_foreign_document_binding( + self, + operation: str, + binding_owner: str, + db_session_with_containers: Session, + mock_external_service_dependencies: MetadataServiceDeps, + ) -> None: + account, tenant = self._create_test_account_and_tenant( + db_session_with_containers, mock_external_service_dependencies + ) + dataset = self._create_test_dataset( + db_session_with_containers, mock_external_service_dependencies, account, tenant + ) + foreign_account, foreign_tenant = self._create_test_account_and_tenant( + db_session_with_containers, mock_external_service_dependencies + ) + foreign_dataset = self._create_test_dataset( + db_session_with_containers, + mock_external_service_dependencies, + foreign_account, + foreign_tenant, + ) + foreign_document = self._create_test_document( + db_session_with_containers, + mock_external_service_dependencies, + foreign_dataset, + foreign_account, + ) + foreign_document.enabled = True + foreign_document.archived = False + foreign_document.indexing_status = IndexingStatus.COMPLETED + foreign_document.doc_metadata = {"old_name": "foreign-value"} + + metadata = MetadataService.create_metadata( + dataset.id, + MetadataArgs(type="string", name="old_name"), + account, + tenant.id, + session=db_session_with_containers, + ) + db_session_with_containers.add( + DatasetMetadataBinding( + tenant_id=dataset.tenant_id if binding_owner == "metadata" else foreign_dataset.tenant_id, + dataset_id=dataset.id if binding_owner == "metadata" else foreign_dataset.id, + metadata_id=metadata.id, + document_id=foreign_document.id, + created_by=account.id, + ) + ) + db_session_with_containers.commit() + + if operation == "rename": + MetadataService.update_metadata_name( + dataset, metadata.id, "new_name", account, session=db_session_with_containers + ) + else: + MetadataService.delete_metadata(dataset, metadata.id, session=db_session_with_containers) + + db_session_with_containers.refresh(foreign_document) + assert foreign_document.doc_metadata == {"old_name": "foreign-value"} + def test_get_built_in_fields_success( self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps ) -> None: @@ -799,13 +868,6 @@ class TestMetadataService: # Mock DocumentService.get_document mock_external_service_dependencies["document_service"].get_document.return_value = document - # Create metadata operation data - from services.entities.knowledge_entities.knowledge_entities import ( - DocumentMetadataOperation, - MetadataDetail, - MetadataOperationData, - ) - metadata_detail = MetadataDetail(id=metadata.id, name=metadata.name, value="test_value") operation = DocumentMetadataOperation(document_id=document.id, metadata_list=[metadata_detail]) @@ -833,6 +895,77 @@ class TestMetadataService: assert binding.tenant_id == tenant.id assert binding.dataset_id == dataset.id + @pytest.mark.parametrize("foreign_resource", ["metadata", "document"]) + def test_update_documents_metadata_rejects_foreign_owner_before_writes( + self, + foreign_resource: str, + db_session_with_containers: Session, + mock_external_service_dependencies: MetadataServiceDeps, + ) -> None: + account, tenant = self._create_test_account_and_tenant( + db_session_with_containers, mock_external_service_dependencies + ) + dataset = self._create_test_dataset( + db_session_with_containers, mock_external_service_dependencies, account, tenant + ) + document = self._create_test_document( + db_session_with_containers, mock_external_service_dependencies, dataset, account + ) + metadata = MetadataService.create_metadata( + dataset.id, + MetadataArgs(type="string", name="owned"), + account, + tenant.id, + session=db_session_with_containers, + ) + + foreign_account, foreign_tenant = self._create_test_account_and_tenant( + db_session_with_containers, mock_external_service_dependencies + ) + foreign_dataset = self._create_test_dataset( + db_session_with_containers, + mock_external_service_dependencies, + foreign_account, + foreign_tenant, + ) + foreign_document = self._create_test_document( + db_session_with_containers, + mock_external_service_dependencies, + foreign_dataset, + foreign_account, + ) + foreign_metadata = MetadataService.create_metadata( + foreign_dataset.id, + MetadataArgs(type="string", name="foreign"), + foreign_account, + foreign_tenant.id, + session=db_session_with_containers, + ) + operation = DocumentMetadataOperation( + document_id=foreign_document.id if foreign_resource == "document" else document.id, + metadata_list=[ + MetadataDetail( + id=foreign_metadata.id if foreign_resource == "metadata" else metadata.id, + name="ignored", + value="value", + ) + ], + ) + + with pytest.raises(MetadataResourceNotFoundError, match=f"{foreign_resource.capitalize()} not found"): + MetadataService.update_documents_metadata( + dataset, + MetadataOperationData(operation_data=[operation]), + account, + session=db_session_with_containers, + ) + + db_session_with_containers.refresh(document) + db_session_with_containers.refresh(foreign_document) + assert document.doc_metadata is None + assert foreign_document.doc_metadata is None + assert db_session_with_containers.query(DatasetMetadataBinding).count() == 0 + def test_update_documents_metadata_with_built_in_fields_enabled( self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps ) -> None: @@ -865,13 +998,6 @@ class TestMetadataService: # Mock DocumentService.get_document mock_external_service_dependencies["document_service"].get_document.return_value = document - # Create metadata operation data - from services.entities.knowledge_entities.knowledge_entities import ( - DocumentMetadataOperation, - MetadataDetail, - MetadataOperationData, - ) - metadata_detail = MetadataDetail(id=metadata.id, name=metadata.name, value="test_value") operation = DocumentMetadataOperation(document_id=document.id, metadata_list=[metadata_detail]) @@ -911,13 +1037,6 @@ class TestMetadataService: dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers ) - # Create metadata operation data - from services.entities.knowledge_entities.knowledge_entities import ( - DocumentMetadataOperation, - MetadataDetail, - MetadataOperationData, - ) - metadata_detail = MetadataDetail(id=metadata.id, name=metadata.name, value="test_value") # Use a valid UUID format that does not exist in the database @@ -927,9 +1046,7 @@ class TestMetadataService: operation_data = MetadataOperationData(operation_data=[operation]) - # Act & Assert: The method should raise ValueError("Document not found.") - # because the exception is now re-raised after rollback - with pytest.raises(ValueError, match="Document not found"): + with pytest.raises(MetadataResourceNotFoundError, match="Document not found"): MetadataService.update_documents_metadata( dataset, operation_data, account, session=db_session_with_containers ) diff --git a/api/tests/test_containers_integration_tests/services/test_webhook_service.py b/api/tests/test_containers_integration_tests/services/test_webhook_service.py index b4022e44f56..a7a06362cad 100644 --- a/api/tests/test_containers_integration_tests/services/test_webhook_service.py +++ b/api/tests/test_containers_integration_tests/services/test_webhook_service.py @@ -1,577 +1,171 @@ +"""Database integration tests for webhook trigger and workflow lookup.""" + import json -from io import BytesIO -from unittest.mock import MagicMock, patch +from typing import TypedDict import pytest from faker import Faker from flask import Flask from sqlalchemy.orm import Session -from werkzeug.datastructures import FileStorage +from enums import DeploymentEdition +from models.account import Account, Tenant from models.enums import AppTriggerStatus, AppTriggerType from models.model import App from models.trigger import AppTrigger, WorkflowWebhookTrigger from models.workflow import Workflow from services.account_service import AccountService, TenantService +from services.entities.feature_entities import SystemFeatureModel from services.trigger.webhook_service import WebhookService from tests.test_containers_integration_tests.helpers import generate_valid_password -class TestWebhookService: - """Integration tests for WebhookService using testcontainers.""" +class WebhookIntegrationData(TypedDict): + tenant: Tenant + account: Account + app: App + workflow: Workflow + webhook_trigger: WorkflowWebhookTrigger + webhook_id: str + app_trigger: AppTrigger - @pytest.fixture - def mock_external_dependencies(self): - """Mock external service dependencies.""" - with ( - patch("services.trigger.webhook_service.AsyncWorkflowService", autospec=True) as mock_async_service, - patch("services.trigger.webhook_service.ToolFileManager", autospec=True) as mock_tool_file_manager, - patch("services.trigger.webhook_service.file_factory", autospec=True) as mock_file_factory, - patch("services.account_service.FeatureService", autospec=True) as mock_feature_service, - ): - # Mock ToolFileManager - mock_tool_file_instance = mock_tool_file_manager.return_value # Mock file creation - mock_tool_file = MagicMock() - mock_tool_file.id = "test_file_id" - mock_tool_file_instance.create_file_by_raw.return_value = mock_tool_file - # Mock file factory - mock_file_obj = MagicMock() - mock_file_factory.build_from_mapping.return_value = mock_file_obj +@pytest.fixture +def test_data( + db_session_with_containers: Session, + monkeypatch: pytest.MonkeyPatch, +) -> WebhookIntegrationData: + """Persist the webhook graph with account and workspace creation enabled.""" - # Mock feature service - mock_feature_service.get_system_features.return_value.is_allow_register = True - mock_feature_service.is_workspace_creation_allowed.return_value = True - - yield { - "async_service": mock_async_service, - "tool_file_manager": mock_tool_file_manager, - "file_factory": mock_file_factory, - "tool_file": mock_tool_file, - "file_obj": mock_file_obj, - "feature_service": mock_feature_service, - } - - @pytest.fixture - def test_data(self, db_session_with_containers: Session, mock_external_dependencies): - """Create test data for webhook service tests.""" - fake = Faker() - - # Create account and tenant - account = AccountService.create_account( - email=fake.email(), - name=fake.name(), - interface_language="en-US", - password=generate_valid_password(fake), - session=db_session_with_containers, - ) - TenantService.create_owner_tenant_if_not_exist(account, name=fake.company(), session=db_session_with_containers) - tenant = account.current_tenant - assert tenant is not None - - # Create app - app = App( - tenant_id=tenant.id, - name=fake.company(), - description=fake.text(), - mode="workflow", - icon="", - icon_background="", - enable_site=True, - enable_api=True, - ) - db_session_with_containers.add(app) - db_session_with_containers.flush() - - # Create workflow - workflow_data = { - "nodes": [ - { - "id": "webhook_node", - "type": "webhook", - "data": { - "type": "trigger-webhook", - "title": "Test Webhook", - "method": "post", - "content_type": "application/json", - "headers": [ - {"name": "Authorization", "required": True}, - {"name": "Content-Type", "required": False}, - ], - "params": [{"name": "version", "required": True}, {"name": "format", "required": False}], - "body": [ - {"name": "message", "type": "string", "required": True}, - {"name": "count", "type": "number", "required": False}, - {"name": "upload", "type": "file", "required": False}, - ], - "status_code": 200, - "response_body": '{"status": "success"}', - "timeout": 30, - }, - } - ], - "edges": [], - } - - workflow = Workflow( - tenant_id=tenant.id, - app_id=app.id, - type="workflow", - graph=json.dumps(workflow_data), - features=json.dumps({}), - created_by=account.id, - environment_variables=[], - conversation_variables=[], - version="1.0", - ) - db_session_with_containers.add(workflow) - db_session_with_containers.flush() - - app.workflow_id = workflow.id - db_session_with_containers.flush() - - # Create webhook trigger - webhook_id = fake.uuid4()[:16] - webhook_trigger = WorkflowWebhookTrigger( - app_id=app.id, - node_id="webhook_node", - tenant_id=tenant.id, - webhook_id=str(webhook_id), - created_by=account.id, - ) - db_session_with_containers.add(webhook_trigger) - db_session_with_containers.flush() - - # Create app trigger (required for non-debug mode) - app_trigger = AppTrigger( - tenant_id=tenant.id, - app_id=app.id, - node_id="webhook_node", - trigger_type=AppTriggerType.TRIGGER_WEBHOOK, - provider_name="webhook", - title="Test Webhook", - status=AppTriggerStatus.ENABLED, - ) - db_session_with_containers.add(app_trigger) - db_session_with_containers.commit() - - return { - "tenant": tenant, - "account": account, - "app": app, - "workflow": workflow, - "webhook_trigger": webhook_trigger, - "webhook_id": webhook_id, - "app_trigger": app_trigger, - } - - def test_get_webhook_trigger_and_workflow_success(self, test_data, flask_app_with_containers: Flask): - """Test successful retrieval of webhook trigger and workflow.""" - webhook_id = test_data["webhook_id"] - - with flask_app_with_containers.app_context(): - webhook_trigger, workflow, node_config = WebhookService.get_webhook_trigger_and_workflow(webhook_id) - - assert webhook_trigger is not None - assert webhook_trigger.webhook_id == webhook_id - assert workflow is not None - assert workflow.app_id == test_data["app"].id - assert node_config is not None - assert node_config["id"] == "webhook_node" - assert node_config["data"].title == "Test Webhook" - - def test_get_webhook_trigger_and_workflow_not_found(self, flask_app_with_containers: Flask): - """Test webhook trigger not found scenario.""" - with flask_app_with_containers.app_context(): - with pytest.raises(ValueError, match="Webhook not found"): - WebhookService.get_webhook_trigger_and_workflow("nonexistent_webhook") - - def test_extract_webhook_data_json(self): - """Test webhook data extraction from JSON request.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="POST", - headers={"Content-Type": "application/json", "Authorization": "Bearer token"}, - query_string="version=1&format=json", - json={"message": "hello", "count": 42}, - ): - webhook_trigger = MagicMock() - webhook_data = WebhookService.extract_webhook_data(webhook_trigger) - - assert webhook_data["method"] == "POST" - assert webhook_data["headers"]["Authorization"] == "Bearer token" - assert webhook_data["query_params"]["version"] == "1" - assert webhook_data["query_params"]["format"] == "json" - assert webhook_data["body"]["message"] == "hello" - assert webhook_data["body"]["count"] == 42 - assert webhook_data["files"] == {} - - def test_extract_webhook_data_form_urlencoded(self): - """Test webhook data extraction from form URL encoded request.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="POST", - headers={"Content-Type": "application/x-www-form-urlencoded"}, - data={"username": "test", "password": "secret"}, - ): - webhook_trigger = MagicMock() - webhook_data = WebhookService.extract_webhook_data(webhook_trigger) - - assert webhook_data["method"] == "POST" - assert webhook_data["body"]["username"] == "test" - assert webhook_data["body"]["password"] == "secret" - - def test_extract_webhook_data_multipart_with_files(self, mock_external_dependencies): - """Test webhook data extraction from multipart form with files.""" - app = Flask(__name__) - - # Create a mock file - file_content = b"test file content" - file_storage = FileStorage(stream=BytesIO(file_content), filename="test.txt", content_type="text/plain") - - with app.test_request_context( - "/webhook", - method="POST", - headers={"Content-Type": "multipart/form-data"}, - data={"message": "test", "file": file_storage}, - ): - webhook_trigger = MagicMock() - webhook_trigger.tenant_id = "test_tenant" - - webhook_data = WebhookService.extract_webhook_data(webhook_trigger) - - assert webhook_data["method"] == "POST" - assert webhook_data["body"]["message"] == "test" - assert "file" in webhook_data["files"] - - # Verify file processing was called - mock_external_dependencies["tool_file_manager"].assert_called_once() - mock_external_dependencies["file_factory"].build_from_mapping.assert_called_once() - - def test_extract_webhook_data_raw_text(self): - """Test webhook data extraction from raw text request.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", method="POST", headers={"Content-Type": "text/plain"}, data="raw text content" - ): - webhook_trigger = MagicMock() - webhook_data = WebhookService.extract_webhook_data(webhook_trigger) - - assert webhook_data["method"] == "POST" - assert webhook_data["body"]["raw"] == "raw text content" - - def test_extract_and_validate_webhook_request_success(self): - """Test successful webhook request validation and type conversion.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="POST", - headers={"Content-Type": "application/json", "Authorization": "Bearer token"}, - query_string="version=1", - json={"message": "hello"}, - ): - webhook_trigger = MagicMock() - node_config = { + fake = Faker() + system_features = SystemFeatureModel( + deployment_edition=DeploymentEdition.COMMUNITY, + is_allow_register=True, + ) + monkeypatch.setattr( + "services.account_service.FeatureService.get_system_features", + lambda: system_features, + ) + monkeypatch.setattr( + "services.account_service.FeatureService.is_workspace_creation_allowed", + lambda: True, + ) + account = AccountService.create_account( + email=fake.email(), + name=fake.name(), + interface_language="en-US", + password=generate_valid_password(fake), + session=db_session_with_containers, + ) + TenantService.create_owner_tenant_if_not_exist(account, name=fake.company(), session=db_session_with_containers) + tenant = account.current_tenant + assert tenant is not None + app = App( + tenant_id=tenant.id, + name=fake.company(), + description=fake.text(), + mode="workflow", + icon="", + icon_background="", + enable_site=True, + enable_api=True, + ) + db_session_with_containers.add(app) + db_session_with_containers.flush() + workflow_data = { + "nodes": [ + { + "id": "webhook_node", + "type": "webhook", "data": { + "type": "trigger-webhook", + "title": "Test Webhook", "method": "post", "content_type": "application/json", "headers": [ {"name": "Authorization", "required": True}, {"name": "Content-Type", "required": False}, ], - "params": [{"name": "version", "required": True}], - "body": [{"name": "message", "type": "string", "required": True}], - } + "params": [ + {"name": "version", "required": True}, + {"name": "format", "required": False}, + ], + "body": [ + {"name": "message", "type": "string", "required": True}, + {"name": "count", "type": "number", "required": False}, + {"name": "upload", "type": "file", "required": False}, + ], + "status_code": 200, + "response_body": '{"status": "success"}', + "timeout": 30, + }, } + ], + "edges": [], + } + workflow = Workflow( + tenant_id=tenant.id, + app_id=app.id, + type="workflow", + graph=json.dumps(workflow_data), + features=json.dumps({}), + created_by=account.id, + environment_variables=[], + conversation_variables=[], + version="1.0", + ) + db_session_with_containers.add(workflow) + db_session_with_containers.flush() + app.workflow_id = workflow.id + webhook_id = fake.uuid4()[:16] + webhook_trigger = WorkflowWebhookTrigger( + app_id=app.id, + node_id="webhook_node", + tenant_id=tenant.id, + webhook_id=webhook_id, + created_by=account.id, + ) + db_session_with_containers.add(webhook_trigger) + db_session_with_containers.flush() + app_trigger = AppTrigger( + tenant_id=tenant.id, + app_id=app.id, + node_id="webhook_node", + trigger_type=AppTriggerType.TRIGGER_WEBHOOK, + provider_name="webhook", + title="Test Webhook", + status=AppTriggerStatus.ENABLED, + ) + db_session_with_containers.add(app_trigger) + db_session_with_containers.commit() + return { + "tenant": tenant, + "account": account, + "app": app, + "workflow": workflow, + "webhook_trigger": webhook_trigger, + "webhook_id": webhook_id, + "app_trigger": app_trigger, + } - result = WebhookService.extract_and_validate_webhook_data(webhook_trigger, node_config) - - assert result["headers"]["Authorization"] == "Bearer token" - assert result["query_params"]["version"] == "1" - assert result["body"]["message"] == "hello" - - def test_extract_and_validate_webhook_request_method_mismatch(self): - """Test webhook validation with HTTP method mismatch.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="GET", - headers={"Content-Type": "application/json"}, - ): - webhook_trigger = MagicMock() - node_config = {"data": {"method": "post", "content_type": "application/json"}} - - with pytest.raises(ValueError, match="HTTP method mismatch"): - WebhookService.extract_and_validate_webhook_data(webhook_trigger, node_config) - - def test_extract_and_validate_webhook_request_missing_required_header(self): - """Test webhook validation with missing required header.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="POST", - headers={"Content-Type": "application/json"}, - ): - webhook_trigger = MagicMock() - node_config = { - "data": { - "method": "post", - "content_type": "application/json", - "headers": [{"name": "Authorization", "required": True}], - } - } - - with pytest.raises(ValueError, match="Required header missing: Authorization"): - WebhookService.extract_and_validate_webhook_data(webhook_trigger, node_config) - - def test_extract_and_validate_webhook_request_case_insensitive_headers(self): - """Test webhook validation with case-insensitive header matching.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="POST", - headers={"Content-Type": "application/json", "authorization": "Bearer token"}, - json={"message": "hello"}, - ): - webhook_trigger = MagicMock() - node_config = { - "data": { - "method": "post", - "content_type": "application/json", - "headers": [{"name": "Authorization", "required": True}], - "body": [{"name": "message", "type": "string", "required": True}], - } - } - - result = WebhookService.extract_and_validate_webhook_data(webhook_trigger, node_config) - - assert result["headers"].get("Authorization") == "Bearer token" - - def test_extract_and_validate_webhook_request_missing_required_param(self): - """Test webhook validation with missing required query parameter.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="POST", - headers={"Content-Type": "application/json"}, - json={"message": "hello"}, - ): - webhook_trigger = MagicMock() - node_config = { - "data": { - "method": "post", - "content_type": "application/json", - "params": [{"name": "version", "required": True}], - "body": [{"name": "message", "type": "string", "required": True}], - } - } - - with pytest.raises(ValueError, match="Required parameter missing: version"): - WebhookService.extract_and_validate_webhook_data(webhook_trigger, node_config) - - def test_extract_and_validate_webhook_request_missing_required_body_param(self): - """Test webhook validation with missing required body parameter.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="POST", - headers={"Content-Type": "application/json"}, - json={}, - ): - webhook_trigger = MagicMock() - node_config = { - "data": { - "method": "post", - "content_type": "application/json", - "body": [{"name": "message", "type": "string", "required": True}], - } - } - - with pytest.raises(ValueError, match="Required body parameter missing: message"): - WebhookService.extract_and_validate_webhook_data(webhook_trigger, node_config) - - def test_extract_and_validate_webhook_request_missing_required_file(self): - """Test webhook validation when required file is missing from multipart request.""" - app = Flask(__name__) - - with app.test_request_context( - "/webhook", - method="POST", - data={"note": "test"}, - content_type="multipart/form-data", - ): - webhook_trigger = MagicMock() - webhook_trigger.tenant_id = "tenant" - webhook_trigger.created_by = "user" - node_config = { - "data": { - "method": "post", - "content_type": "multipart/form-data", - "body": [{"name": "file", "type": "file", "required": True}], - } - } - - result = WebhookService.extract_and_validate_webhook_data(webhook_trigger, node_config) - - assert result["files"] == {} - - def test_trigger_workflow_execution_success( - self, test_data, mock_external_dependencies, flask_app_with_containers: Flask - ): - """Test successful workflow execution trigger.""" - webhook_data = { - "method": "POST", - "headers": {"Authorization": "Bearer token"}, - "query_params": {"version": "1"}, - "body": {"message": "hello"}, - "files": {}, - } +class TestWebhookServiceDatabaseLookup: + def test_get_webhook_trigger_and_workflow_success( + self, + test_data: WebhookIntegrationData, + flask_app_with_containers: Flask, + ) -> None: with flask_app_with_containers.app_context(): - # Mock tenant owner lookup to return the test account - with patch("services.trigger.webhook_service.select", autospec=True) as mock_select: - mock_query = MagicMock() - mock_select.return_value.join.return_value.where.return_value = mock_query + webhook_trigger, workflow, node_config = WebhookService.get_webhook_trigger_and_workflow( + test_data["webhook_id"] + ) - # Mock the session to return our test account - with patch("services.trigger.webhook_service.Session", autospec=True) as mock_session: - mock_session_instance = MagicMock() - mock_session.return_value.__enter__.return_value = mock_session_instance - mock_session_instance.scalar.return_value = test_data["account"] - - # Should not raise any exceptions - WebhookService.trigger_workflow_execution( - test_data["webhook_trigger"], webhook_data, test_data["workflow"] - ) - - # Verify AsyncWorkflowService was called - mock_external_dependencies["async_service"].trigger_workflow_async.assert_called_once() - - def test_trigger_workflow_execution_end_user_service_failure( - self, test_data, mock_external_dependencies, flask_app_with_containers: Flask - ): - """Test workflow execution trigger when EndUserService fails.""" - webhook_data = {"method": "POST", "headers": {}, "query_params": {}, "body": {}, "files": {}} + assert webhook_trigger.webhook_id == test_data["webhook_id"] + assert workflow.app_id == test_data["app"].id + assert node_config["id"] == "webhook_node" + assert node_config["data"].title == "Test Webhook" + def test_get_webhook_trigger_and_workflow_not_found(self, flask_app_with_containers: Flask) -> None: with flask_app_with_containers.app_context(): - # Mock EndUserService to raise an exception - with patch( - "services.trigger.webhook_service.EndUserService.get_or_create_end_user_by_type", autospec=True - ) as mock_end_user: - mock_end_user.side_effect = ValueError("Failed to create end user") - - with pytest.raises(ValueError, match="Failed to create end user"): - WebhookService.trigger_workflow_execution( - test_data["webhook_trigger"], webhook_data, test_data["workflow"] - ) - - def test_generate_webhook_response_default(self): - """Test webhook response generation with default values.""" - node_config = {"data": {}} - - response_data, status_code = WebhookService.generate_webhook_response(node_config) - - assert status_code == 200 - assert response_data["status"] == "success" - assert "Webhook processed successfully" in response_data["message"] - - def test_generate_webhook_response_custom_json(self): - """Test webhook response generation with custom JSON response.""" - node_config = {"data": {"status_code": 201, "response_body": '{"result": "created", "id": 123}'}} - - response_data, status_code = WebhookService.generate_webhook_response(node_config) - - assert status_code == 201 - assert response_data["result"] == "created" - assert response_data["id"] == 123 - - def test_generate_webhook_response_custom_text(self): - """Test webhook response generation with custom text response.""" - node_config = {"data": {"status_code": 202, "response_body": "Request accepted for processing"}} - - response_data, status_code = WebhookService.generate_webhook_response(node_config) - - assert status_code == 202 - assert response_data["message"] == "Request accepted for processing" - - def test_generate_webhook_response_invalid_json(self): - """Test webhook response generation with invalid JSON response.""" - node_config = {"data": {"status_code": 400, "response_body": '{"invalid": json}'}} - - response_data, status_code = WebhookService.generate_webhook_response(node_config) - - assert status_code == 400 - assert response_data["message"] == '{"invalid": json}' - - def test_process_file_uploads_success(self, mock_external_dependencies): - """Test successful file upload processing.""" - # Create mock files - files = { - "file1": MagicMock(filename="test1.txt", content_type="text/plain"), - "file2": MagicMock(filename="test2.jpg", content_type="image/jpeg"), - } - - # Mock file reads - files["file1"].read.return_value = b"content1" - files["file2"].read.return_value = b"content2" - - webhook_trigger = MagicMock() - webhook_trigger.tenant_id = "test_tenant" - - result = WebhookService._process_file_uploads(files, webhook_trigger) - - assert len(result) == 2 - assert "file1" in result - assert "file2" in result - - # Verify file processing was called for each file - assert mock_external_dependencies["tool_file_manager"].call_count == 2 - assert mock_external_dependencies["file_factory"].build_from_mapping.call_count == 2 - - def test_process_file_uploads_with_errors(self, mock_external_dependencies): - """Test file upload processing with errors.""" - # Create mock files, one will fail - files = { - "good_file": MagicMock(filename="test.txt", content_type="text/plain"), - "bad_file": MagicMock(filename="test.bad", content_type="text/plain"), - } - - files["good_file"].stream.read.return_value = b"content" - files["bad_file"].stream.read.side_effect = Exception("Read error") - - webhook_trigger = MagicMock() - webhook_trigger.tenant_id = "test_tenant" - - result = WebhookService._process_file_uploads(files, webhook_trigger) - - # Should process the good file and skip the bad one - assert len(result) == 1 - assert "good_file" in result - assert "bad_file" not in result - - def test_process_file_uploads_empty_filename(self, mock_external_dependencies): - """Test file upload processing with empty filename.""" - files = { - "no_filename": MagicMock(filename="", content_type="text/plain"), - "none_filename": MagicMock(filename=None, content_type="text/plain"), - } - - webhook_trigger = MagicMock() - webhook_trigger.tenant_id = "test_tenant" - - result = WebhookService._process_file_uploads(files, webhook_trigger) - - # Should skip files without filenames - assert len(result) == 0 - mock_external_dependencies["tool_file_manager"].assert_not_called() + with pytest.raises(ValueError, match="Webhook not found"): + WebhookService.get_webhook_trigger_and_workflow("nonexistent_webhook") diff --git a/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py b/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py index ddaee0cf051..432a483b8c0 100644 --- a/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py +++ b/api/tests/test_containers_integration_tests/services/test_webhook_service_relationships.py @@ -13,7 +13,7 @@ from sqlalchemy import select from sqlalchemy.orm import Session from core.trigger.constants import TRIGGER_WEBHOOK_NODE_TYPE -from enums.quota_type import QuotaType +from enums import QuotaType from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole from models.enums import AppTriggerStatus, AppTriggerType from models.model import App diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py index f553b0f72a0..71c9c81e0d2 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_app_service.py @@ -3,7 +3,6 @@ from __future__ import annotations import json import uuid from datetime import UTC, datetime, timedelta -from types import SimpleNamespace from unittest.mock import patch import pytest @@ -12,13 +11,13 @@ from sqlalchemy.orm import Session from graphon.enums import WorkflowExecutionStatus from models import EndUser, Workflow, WorkflowAppLog, WorkflowArchiveLog, WorkflowRun -from models.enums import AppTriggerType, CreatorUserRole, EndUserType, WorkflowRunTriggeredFrom +from models.enums import CreatorUserRole, EndUserType, WorkflowRunTriggeredFrom from models.workflow import WorkflowAppLogCreatedFrom from services.account_service import AccountService, TenantService # Delay import of AppService to avoid circular dependency # from services.app_service import AppService, CreateAppParams -from services.workflow_app_service import LogView, WorkflowAppService +from services.workflow_app_service import WorkflowAppService from tests.test_containers_integration_tests.helpers import generate_valid_password @@ -1627,73 +1626,3 @@ class TestWorkflowAppService: end_user_item = next(d for d in result["data"] if d["created_by_end_user"] is not None) assert account_item["created_by_account"].id == account.id assert end_user_item["created_by_end_user"].id == end_user.id - - -class TestLogView: - def test_details_and_proxy_attributes(self): - log = SimpleNamespace(id="log-1", status="succeeded") - view = LogView(log=log, details={"trigger_metadata": {"type": "plugin"}}) - - assert view.details == {"trigger_metadata": {"type": "plugin"}} - assert view.status == "succeeded" - - -class TestHandleTriggerMetadata: - def test_returns_empty_dict_when_metadata_missing(self): - service = WorkflowAppService() - assert service.handle_trigger_metadata("tenant-1", None) == {} - - def test_enriches_plugin_icons(self): - service = WorkflowAppService() - meta = { - "type": AppTriggerType.TRIGGER_PLUGIN.value, - "icon_filename": "light.png", - "icon_dark_filename": "dark.png", - } - with patch( - "services.workflow_app_service.PluginService.get_plugin_icon_url", - side_effect=["https://cdn/light.png", "https://cdn/dark.png"], - ) as mock_icon: - result = service.handle_trigger_metadata("tenant-1", json.dumps(meta)) - - assert result["icon"] == "https://cdn/light.png" - assert result["icon_dark"] == "https://cdn/dark.png" - assert mock_icon.call_count == 2 - - def test_non_plugin_metadata_without_icon_lookup(self): - service = WorkflowAppService() - meta = {"type": AppTriggerType.TRIGGER_WEBHOOK.value} - with patch("services.workflow_app_service.PluginService.get_plugin_icon_url") as mock_icon: - result = service.handle_trigger_metadata("tenant-1", json.dumps(meta)) - - assert result["type"] == AppTriggerType.TRIGGER_WEBHOOK.value - mock_icon.assert_not_called() - - -class TestSafeJsonLoads: - @pytest.mark.parametrize( - ("value", "expected"), - [ - (None, None), - ("", None), - ('{"k":"v"}', {"k": "v"}), - ("not-json", None), - ({"raw": True}, {"raw": True}), - ], - ) - def test_handles_various_inputs(self, value, expected): - assert WorkflowAppService._safe_json_loads(value) == expected - - -class TestSafeParseUuid: - def test_returns_none_for_short_or_invalid_values(self): - service = WorkflowAppService() - assert service._safe_parse_uuid("short") is None - assert service._safe_parse_uuid("x" * 40) is None - - def test_returns_uuid_for_valid_string(self): - service = WorkflowAppService() - raw = str(uuid.uuid4()) - result = service._safe_parse_uuid(raw) - assert result is not None - assert str(result) == raw diff --git a/api/tests/test_containers_integration_tests/services/test_workflow_service.py b/api/tests/test_containers_integration_tests/services/test_workflow_service.py index 6531ed4fbb0..05697997511 100644 --- a/api/tests/test_containers_integration_tests/services/test_workflow_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workflow_service.py @@ -12,7 +12,9 @@ import pytest from faker import Faker from sqlalchemy.orm import Session +from graphon.enums import BuiltinNodeTypes, ErrorStrategy, WorkflowNodeExecutionStatus from models import Account, AccountStatus, App, TenantStatus, Workflow +from models.enums import CreatorUserRole from models.model import AppMode from models.workflow import WorkflowType from services.workflow_ref_service import WorkflowRef @@ -158,7 +160,6 @@ class TestWorkflowService: workflow = self._create_test_workflow(db_session_with_containers, app, account, fake) # Create a mock node execution record - from models.enums import CreatorUserRole from models.workflow import WorkflowNodeExecutionModel node_execution = WorkflowNodeExecutionModel() @@ -870,12 +871,13 @@ class TestWorkflowService: from unittest.mock import patch with patch("flask_login.utils._get_user", return_value=account, autospec=True): - result = workflow_service.publish_workflow( + result, retirement_candidates = workflow_service.publish_workflow( session=db_session_with_containers, app_model=app, account=account ) # Assert assert result is not None + assert retirement_candidates == set() assert result.version != Workflow.VERSION_DRAFT # Version should be a timestamp format like '2025-08-22 00:10:24.722051' assert isinstance(result.version, str) @@ -1636,7 +1638,6 @@ class TestWorkflowService: import uuid from datetime import datetime - from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.graph_events import NodeRunSucceededEvent from graphon.node_events import NodeRunResult from graphon.nodes.base.node import Node @@ -1681,12 +1682,10 @@ class TestWorkflowService: # Assert assert result is not None assert result.node_id == node_id - from graphon.enums import BuiltinNodeTypes assert result.node_type == BuiltinNodeTypes.START # Should match the mock node type assert result.title == "Test Node" # Import the enum for comparison - from graphon.enums import WorkflowNodeExecutionStatus assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED assert result.inputs is not None @@ -1711,7 +1710,6 @@ class TestWorkflowService: import uuid from datetime import datetime - from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus from graphon.graph_events import NodeRunFailedEvent from graphon.node_events import NodeRunResult from graphon.nodes.base.node import Node @@ -1756,7 +1754,6 @@ class TestWorkflowService: assert result is not None assert result.node_id == node_id # Import the enum for comparison - from graphon.enums import WorkflowNodeExecutionStatus assert result.status == WorkflowNodeExecutionStatus.FAILED assert result.error is not None @@ -1780,7 +1777,6 @@ class TestWorkflowService: import uuid from datetime import datetime - from graphon.enums import BuiltinNodeTypes, ErrorStrategy, WorkflowNodeExecutionStatus from graphon.graph_events import NodeRunFailedEvent from graphon.node_events import NodeRunResult from graphon.nodes.base.node import Node @@ -1826,7 +1822,6 @@ class TestWorkflowService: assert result is not None assert result.node_id == node_id # Import the enum for comparison - from graphon.enums import WorkflowNodeExecutionStatus assert result.status == WorkflowNodeExecutionStatus.EXCEPTION # Should be EXCEPTION, not FAILED assert result.outputs is not None diff --git a/api/tests/test_containers_integration_tests/services/test_workspace_service.py b/api/tests/test_containers_integration_tests/services/test_workspace_service.py index f22775a13eb..d6edcd25f85 100644 --- a/api/tests/test_containers_integration_tests/services/test_workspace_service.py +++ b/api/tests/test_containers_integration_tests/services/test_workspace_service.py @@ -1,12 +1,13 @@ from __future__ import annotations +import json from unittest.mock import MagicMock, patch import pytest from faker import Faker from sqlalchemy.orm import Session -from enums.deployment_edition import DeploymentEdition +from enums import CloudPlan, DeploymentEdition from models import Account, Tenant, TenantAccountJoin, TenantAccountRole from services.credit_pool_service import CreditPoolBalance from services.workspace_service import WorkspaceService @@ -24,7 +25,10 @@ class TestWorkspaceService: patch("services.workspace_service.dify_config") as mock_dify_config, ): # Setup default mock returns - mock_feature_service.get_features.return_value.can_replace_logo = True + feature = mock_feature_service.get_features.return_value + feature.can_replace_logo = True + feature.billing.enabled = True + feature.billing.subscription.plan = "professional" mock_tenant_service.has_roles.return_value = True mock_dify_config.FILES_URL = "https://example.com/files" @@ -112,7 +116,7 @@ class TestWorkspaceService: assert result is not None assert result["id"] == tenant.id assert result["name"] == tenant.name - assert result["plan"] == tenant.plan + assert result["plan"] == "professional" assert result["status"] == tenant.status assert result["role"] == TenantAccountRole.OWNER assert result["created_at"] == tenant.created_at @@ -159,7 +163,7 @@ class TestWorkspaceService: assert result is not None assert result["id"] == tenant.id assert result["name"] == tenant.name - assert result["plan"] == tenant.plan + assert result["plan"] == "professional" assert result["status"] == tenant.status assert result["role"] == TenantAccountRole.OWNER assert result["created_at"] == tenant.created_at @@ -214,7 +218,7 @@ class TestWorkspaceService: assert result is not None assert result["id"] == tenant.id assert result["name"] == tenant.name - assert result["plan"] == tenant.plan + assert result["plan"] == "professional" assert result["status"] == tenant.status assert result["role"] == TenantAccountRole.NORMAL assert result["created_at"] == tenant.created_at @@ -330,8 +334,6 @@ class TestWorkspaceService: for config in test_configs: # Update tenant custom config - import json - tenant.custom_config = json.dumps(config) db_session_with_containers.commit() @@ -502,8 +504,6 @@ class TestWorkspaceService: for config in test_configs: # Update tenant custom config - import json - tenant.custom_config = json.dumps(config) db_session_with_containers.commit() @@ -561,8 +561,6 @@ class TestWorkspaceService: self, db_session_with_containers: Session, mock_external_service_dependencies ): """replace_webapp_logo should be None when custom_config_dict does not have the key.""" - import json - fake = Faker() account, tenant = self._create_test_account_and_tenant( db_session_with_containers, mock_external_service_dependencies @@ -583,8 +581,6 @@ class TestWorkspaceService: self, db_session_with_containers: Session, mock_external_service_dependencies ): """The logo URL should use dify_config.FILES_URL as the base.""" - import json - fake = Faker() account, tenant = self._create_test_account_and_tenant( db_session_with_containers, mock_external_service_dependencies @@ -606,20 +602,23 @@ class TestWorkspaceService: def test_get_tenant_info_should_not_include_cloud_fields_in_self_hosted( self, db_session_with_containers: Session, mock_external_service_dependencies ): - """next_credit_reset_date and trial_credits should NOT appear in SELF_HOSTED mode.""" + """Cloud-only billing data should not appear in SELF_HOSTED mode.""" fake = Faker() account, tenant = self._create_test_account_and_tenant( db_session_with_containers, mock_external_service_dependencies ) mock_external_service_dependencies["dify_config"].DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY - mock_external_service_dependencies["feature_service"].get_features.return_value.can_replace_logo = False + feature = mock_external_service_dependencies["feature_service"].get_features.return_value + feature.can_replace_logo = False + feature.billing.enabled = False mock_external_service_dependencies["tenant_service"].has_roles.return_value = False with patch("services.workspace_service.current_user", account): result = WorkspaceService.get_tenant_info(tenant, db_session_with_containers) assert result is not None + assert result["plan"] is None assert "next_credit_reset_date" not in result assert "trial_credits" not in result assert "trial_credits_used" not in result @@ -773,8 +772,6 @@ class TestWorkspaceService: self, db_session_with_containers: Session, mock_external_service_dependencies ): """When plan is SANDBOX, skip paid pool and use trial pool.""" - from enums.cloud_plan import CloudPlan - fake = Faker() account, tenant = self._create_test_account_and_tenant( db_session_with_containers, mock_external_service_dependencies diff --git a/api/tests/test_containers_integration_tests/tasks/test_batch_clean_document_task.py b/api/tests/test_containers_integration_tests/tasks/test_batch_clean_document_task.py index 436c8f11b05..4193223f311 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_batch_clean_document_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_batch_clean_document_task.py @@ -19,7 +19,7 @@ from extensions.storage.storage_type import StorageType from libs.datetime_utils import naive_utc_now from models import Account, Tenant, TenantAccountJoin, TenantAccountRole from models.dataset import Dataset, Document, DocumentSegment -from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus, SegmentStatus +from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, IndexingStatus, SegmentStatus from models.model import UploadFile from tasks.batch_clean_document_task import batch_clean_document_task @@ -207,8 +207,6 @@ class TestBatchCleanDocumentTask: """ fake = Faker() - from models.enums import CreatorUserRole - upload_file = UploadFile( tenant_id=account.current_tenant.id, storage_type=StorageType.LOCAL, diff --git a/api/tests/test_containers_integration_tests/tasks/test_dataset_indexing_task.py b/api/tests/test_containers_integration_tests/tasks/test_dataset_indexing_task.py index 7f102f7375f..0c36794cb8c 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_dataset_indexing_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_dataset_indexing_task.py @@ -11,7 +11,7 @@ from sqlalchemy.orm import Session from core.indexing_runner import DocumentIsPausedError from core.rag.index_processor.constant.index_type import IndexTechniqueType -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from models import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus from models.dataset import Dataset, Document from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus diff --git a/api/tests/test_containers_integration_tests/tasks/test_delete_account_task.py b/api/tests/test_containers_integration_tests/tasks/test_delete_account_task.py index 9dfc6325d01..bbb95b0d233 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_delete_account_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_delete_account_task.py @@ -14,6 +14,7 @@ from _pytest.logging import LogCaptureFixture from pytest_mock import MockerFixture from sqlalchemy.orm import Session +from enums import DeploymentEdition from models.account import Account from tasks.delete_account_task import delete_account_task @@ -35,14 +36,14 @@ def mock_external_dependencies(mocker: MockerFixture) -> tuple[MagicMock, MagicM return billing_service, mail_task -def test_billing_enabled_account_exists_calls_billing_and_sends_email( +def test_cloud_account_exists_calls_billing_and_sends_email( db_session_with_containers: Session, mock_external_dependencies: tuple[MagicMock, MagicMock], mocker: MockerFixture, ) -> None: billing_service, mail_task = mock_external_dependencies account = _create_account(db_session_with_containers, email="a@b.com") - mocker.patch("tasks.delete_account_task.dify_config.BILLING_ENABLED", True) + mocker.patch("tasks.delete_account_task.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) delete_account_task(account.id) @@ -50,14 +51,14 @@ def test_billing_enabled_account_exists_calls_billing_and_sends_email( mail_task.delay.assert_called_once_with(account.email) -def test_billing_disabled_account_exists_sends_email_only( +def test_community_account_exists_sends_email_only( db_session_with_containers: Session, mock_external_dependencies: tuple[MagicMock, MagicMock], mocker: MockerFixture, ) -> None: billing_service, mail_task = mock_external_dependencies account = _create_account(db_session_with_containers, email="x@y.com") - mocker.patch("tasks.delete_account_task.dify_config.BILLING_ENABLED", False) + mocker.patch("tasks.delete_account_task.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) delete_account_task(account.id) @@ -65,12 +66,12 @@ def test_billing_disabled_account_exists_sends_email_only( mail_task.delay.assert_called_once_with(account.email) -def test_billing_enabled_account_not_found_calls_billing_no_email( +def test_cloud_account_not_found_calls_billing_no_email( mock_external_dependencies: tuple[MagicMock, MagicMock], mocker: MockerFixture, caplog: LogCaptureFixture ) -> None: billing_service, mail_task = mock_external_dependencies account_id = str(uuid4()) - mocker.patch("tasks.delete_account_task.dify_config.BILLING_ENABLED", True) + mocker.patch("tasks.delete_account_task.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) delete_account_task(account_id) @@ -87,7 +88,7 @@ def test_billing_delete_raises_propagates_and_no_email( billing_service, mail_task = mock_external_dependencies account = _create_account(db_session_with_containers, email="err@example.com") billing_service.delete_account.side_effect = RuntimeError("billing down") - mocker.patch("tasks.delete_account_task.dify_config.BILLING_ENABLED", True) + mocker.patch("tasks.delete_account_task.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) with pytest.raises(RuntimeError, match="billing down"): delete_account_task(account.id) diff --git a/api/tests/test_containers_integration_tests/tasks/test_document_indexing_task.py b/api/tests/test_containers_integration_tests/tasks/test_document_indexing_task.py index b58b2f01da4..bb45e2be34a 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_document_indexing_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_document_indexing_task.py @@ -7,7 +7,7 @@ from sqlalchemy.orm import Session from core.entities.document_task import DocumentTask from core.rag.index_processor.constant.index_type import IndexTechniqueType -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from models import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus from models.dataset import Dataset, Document from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus diff --git a/api/tests/test_containers_integration_tests/tasks/test_duplicate_document_indexing_task.py b/api/tests/test_containers_integration_tests/tasks/test_duplicate_document_indexing_task.py index 74199255e96..d5393435bfa 100644 --- a/api/tests/test_containers_integration_tests/tasks/test_duplicate_document_indexing_task.py +++ b/api/tests/test_containers_integration_tests/tasks/test_duplicate_document_indexing_task.py @@ -7,7 +7,7 @@ from sqlalchemy.orm import Session from core.indexing_runner import DocumentIsPausedError from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from models import Account, Tenant, TenantAccountJoin, TenantAccountRole from models.dataset import Dataset, Document, DocumentSegment from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus, SegmentStatus diff --git a/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py b/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py index 90dd2bcfc84..c08b0be6a04 100644 --- a/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py +++ b/api/tests/test_containers_integration_tests/trigger/test_trigger_e2e.py @@ -24,6 +24,7 @@ from core.trigger.debug import event_selectors from core.trigger.debug.event_bus import TriggerDebugEventBus from core.trigger.debug.event_selectors import PluginTriggerDebugEventPoller, WebhookTriggerDebugEventPoller from core.trigger.debug.events import PluginTriggerDebugEvent, build_plugin_pool_key +from enums import DeploymentEdition from graphon.enums import BuiltinNodeTypes from libs.datetime_utils import naive_utc_now from models.account import Account, Tenant @@ -115,7 +116,9 @@ def test_publish_blocks_start_and_trigger_coexistence( "is_plugin_manager_enabled", classmethod(lambda _cls: False), ) - monkeypatch.setattr("services.workflow_service.dify_config", SimpleNamespace(BILLING_ENABLED=False)) + monkeypatch.setattr( + "services.workflow_service.dify_config", SimpleNamespace(DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY) + ) with pytest.raises(ValueError, match="Start node and trigger nodes cannot coexist"): workflow_service.publish_workflow(session=db_session_with_containers, app_model=app_model, account=account) diff --git a/api/tests/unit_tests/.ruff.toml b/api/tests/unit_tests/.ruff.toml new file mode 100644 index 00000000000..2e9b34cf208 --- /dev/null +++ b/api/tests/unit_tests/.ruff.toml @@ -0,0 +1,429 @@ +extend = "../../.ruff.toml" +src = ["../.."] + +[lint] +extend-select = ["ANN401", "ARG"] + +# Existing strict-mode debt. Remove a file entry when bringing it under strict checking. +[lint.per-file-ignores] +"clients/agent_backend/test_request_builder.py" = ["TID251"] +"commands/test_archive_workflow_runs.py" = ["ARG005"] +"commands/test_data_migration_wizard.py" = ["ARG005"] +"commands/test_legacy_model_type_migration.py" = ["ARG001", "ARG002", "ARG005"] +"controllers/common/test_agent_app_parameters.py" = ["ARG005", "TID251"] +"controllers/common/test_app_access.py" = ["ARG005"] +"controllers/console/agent/test_agent_controllers.py" = ["ARG001", "ARG002", "ARG003", "ARG005", "TID251"] +"controllers/console/app/test_agent_app_sandbox.py" = ["ARG002", "ARG005"] +"controllers/console/app/test_agent_config_inspector.py" = ["ARG005"] +"controllers/console/app/test_agent_drive_inspector.py" = ["ARG005"] +"controllers/console/app/test_agent_manage_guard.py" = ["ARG001"] +"controllers/console/app/test_agent_skills.py" = ["ARG005"] +"controllers/console/app/test_annotation_security.py" = ["ARG002"] +"controllers/console/app/test_app_apis.py" = ["ARG001", "ARG002"] +"controllers/console/app/test_app_import_api.py" = ["ARG001", "ARG002", "ARG005"] +"controllers/console/app/test_app_response_models.py" = ["ARG002", "ARG004", "ARG005"] +"controllers/console/app/test_conversation_api.py" = ["ARG001"] +"controllers/console/app/test_generator_api.py" = ["ARG001"] +"controllers/console/app/test_mcp_server_response.py" = ["ARG002"] +"controllers/console/app/test_message_api.py" = ["ARG001"] +"controllers/console/app/test_statistic_api.py" = ["ANN401", "ARG001", "TID251"] +"controllers/console/app/test_workflow.py" = ["ARG001", "ARG005"] +"controllers/console/app/test_workflow_convert_api.py" = ["ARG005"] +"controllers/console/app/test_workflow_node_output_inspector.py" = ["ANN401", "ARG001", "TID251"] +"controllers/console/app/test_workflow_run_api.py" = ["ANN401", "TID251"] +"controllers/console/app/workflow_draft_variables_test.py" = ["TID251"] +"controllers/console/auth/test_account_activation.py" = ["ARG002"] +"controllers/console/auth/test_authentication_security.py" = ["ARG002"] +"controllers/console/auth/test_email_verification.py" = ["ARG002"] +"controllers/console/auth/test_login_logout.py" = ["ARG002"] +"controllers/console/auth/test_oauth.py" = ["ARG002"] +"controllers/console/auth/test_oauth_timezone.py" = ["ARG001"] +"controllers/console/auth/test_password_reset.py" = ["ARG002"] +"controllers/console/auth/test_token_refresh.py" = ["ARG002"] +"controllers/console/billing/test_billing.py" = ["ARG002"] +"controllers/console/datasets/test_datasets.py" = ["ARG005"] +"controllers/console/datasets/test_datasets_document.py" = ["ARG002", "ARG005"] +"controllers/console/datasets/test_datasets_document_download.py" = ["ARG005"] +"controllers/console/datasets/test_datasets_segments.py" = ["TID251"] +"controllers/console/datasets/test_external.py" = ["TID251"] +"controllers/console/datasets/test_wraps.py" = ["ARG001"] +"controllers/console/explore/test_trial.py" = ["ANN401", "ARG001", "ARG005", "TID251"] +"controllers/console/explore/test_wraps.py" = ["ARG001"] +"controllers/console/snippets/test_snippet_workflow.py" = ["ARG001"] +"controllers/console/tag/test_tags.py" = ["ARG002"] +"controllers/console/test_files.py" = ["ARG001", "ARG002"] +"controllers/console/test_human_input_form.py" = ["ARG001", "ARG005"] +"controllers/console/test_init_validate.py" = ["ARG005"] +"controllers/console/test_workspace_account.py" = ["ARG002"] +"controllers/console/test_workspace_members.py" = ["ARG002", "ARG005"] +"controllers/console/test_wraps.py" = ["ARG001", "ARG002"] +"controllers/console/workspace/test_load_balancing_config.py" = ["ARG001"] +"controllers/console/workspace/test_plugin.py" = ["ARG002", "TID251"] +"controllers/console/workspace/test_snippets.py" = ["ARG002"] +"controllers/console/workspace/test_tool_providers.py" = ["ARG001", "ARG005"] +"controllers/console/workspace/test_trigger_providers.py" = ["ARG001"] +"controllers/console/workspace/test_workspace.py" = ["ARG005"] +"controllers/files/test_image_preview.py" = ["ARG005"] +"controllers/files/test_tool_files.py" = ["ARG002", "ARG005"] +"controllers/files/test_upload.py" = ["ARG002", "ARG005"] +"controllers/inner_api/plugin/test_plugin.py" = ["ARG002"] +"controllers/inner_api/plugin/test_plugin_wraps.py" = ["ARG001", "ARG002", "ARG003", "TID251"] +"controllers/inner_api/test_runtime_credentials.py" = ["ARG001"] +"controllers/mcp/test_mcp.py" = ["ARG002"] +"controllers/openapi/auth/test_conditions.py" = ["ARG005"] +"controllers/openapi/auth/test_flow.py" = ["ARG005"] +"controllers/openapi/auth/test_pipeline.py" = ["ARG001"] +"controllers/openapi/conftest.py" = ["ARG001"] +"controllers/openapi/test_account.py" = ["ARG005"] +"controllers/openapi/test_app_describe_builder.py" = ["ARG001"] +"controllers/openapi/test_app_run_streaming.py" = ["ARG001"] +"controllers/openapi/test_contract.py" = ["ARG001", "TID251"] +"controllers/openapi/test_error_contract.py" = ["ARG002"] +"controllers/openapi/test_human_input_form.py" = ["ARG002"] +"controllers/openapi/test_oauth_sso_claims.py" = ["ARG002"] +"controllers/openapi/test_workflow_events_openapi.py" = ["ARG002", "ARG005"] +"controllers/openapi/test_workspaces_members.py" = ["ARG001"] +"controllers/service_api/app/test_app.py" = ["ARG002"] +"controllers/service_api/app/test_completion.py" = ["ARG002"] +"controllers/service_api/app/test_file.py" = ["ARG002"] +"controllers/service_api/app/test_hitl_service_api.py" = ["ARG002", "ARG005"] +"controllers/service_api/app/test_workflow_events.py" = ["ARG005"] +"controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py" = ["ARG002"] +"controllers/service_api/dataset/test_dataset_segment.py" = ["ARG002"] +"controllers/service_api/dataset/test_document.py" = ["ARG001", "ARG002"] +"controllers/service_api/dataset/test_metadata.py" = ["ARG002"] +"controllers/service_api/test_trace_session_id_parsing.py" = ["ARG001"] +"controllers/service_api/test_wraps.py" = ["ARG001", "ARG002"] +"controllers/trigger/test_trigger.py" = ["ARG002"] +"controllers/trigger/test_webhook.py" = ["ARG002"] +"controllers/web/conftest.py" = ["ANN401", "TID251"] +"controllers/web/test_app.py" = ["ARG002", "ARG005"] +"controllers/web/test_audio.py" = ["ARG002"] +"controllers/web/test_completion.py" = ["ARG002"] +"controllers/web/test_feature.py" = ["ARG002"] +"controllers/web/test_human_input_form.py" = ["ARG001", "ARG002", "ARG005"] +"controllers/web/test_message_endpoints.py" = ["ARG002"] +"controllers/web/test_remote_files.py" = ["ARG002"] +"controllers/web/test_saved_message.py" = ["ARG002"] +"controllers/web/test_web_login.py" = ["ARG002"] +"controllers/web/test_web_passport.py" = ["ARG002"] +"controllers/web/test_workflow.py" = ["ARG002"] +"core/agent/test_base_agent_runner.py" = ["ARG002"] +"core/agent/test_cot_agent_runner.py" = ["ARG001"] +"core/agent/test_cot_chat_agent_runner.py" = ["ARG002"] +"core/agent/test_fc_agent_runner.py" = ["TID251"] +"core/app/app_config/common/test_parameters_mapping.py" = ["ARG002"] +"core/app/app_config/easy_ui_based_app/test_dataset_manager.py" = ["ARG001", "ARG002"] +"core/app/app_config/easy_ui_based_app/test_model_config_converter.py" = ["ARG002"] +"core/app/app_config/easy_ui_based_app/test_variables_manager.py" = ["ARG002"] +"core/app/apps/advanced_chat/test_app_generator.py" = ["ARG001", "ARG002", "ARG005"] +"core/app/apps/advanced_chat/test_app_runner_input_moderation.py" = ["ARG001", "ARG005"] +"core/app/apps/advanced_chat/test_generate_task_pipeline.py" = ["ARG005"] +"core/app/apps/advanced_chat/test_generate_task_pipeline_core.py" = ["ARG001", "ARG002", "ARG005"] +"core/app/apps/agent_app/test_app_generator.py" = ["ARG001", "ARG005"] +"core/app/apps/agent_app/test_app_runner.py" = ["ANN401", "ARG002", "ARG005", "TID251"] +"core/app/apps/agent_app/test_input_guards.py" = ["ANN401", "ARG002", "TID251"] +"core/app/apps/agent_app/test_resolve_agent.py" = ["ANN401", "TID251"] +"core/app/apps/agent_app/test_runtime_request_builder.py" = ["ARG002", "TID251"] +"core/app/apps/agent_chat/test_agent_chat_app_config_manager.py" = ["ARG005"] +"core/app/apps/agent_chat/test_agent_chat_app_generator.py" = ["ARG001"] +"core/app/apps/chat/test_app_config_manager.py" = ["ARG001"] +"core/app/apps/chat/test_app_generator_and_runner.py" = ["ARG001", "ARG002", "ARG005"] +"core/app/apps/common/test_workflow_response_converter_truncation.py" = ["TID251"] +"core/app/apps/completion/test_app_runner.py" = ["ARG001", "ARG002", "ARG005"] +"core/app/apps/pipeline/test_pipeline_generator.py" = ["ARG001", "ARG002", "ARG005"] +"core/app/apps/pipeline/test_pipeline_runner.py" = ["ARG001", "ARG002"] +"core/app/apps/test_advanced_chat_app_generator.py" = ["ARG001"] +"core/app/apps/test_base_app_generator.py" = ["ARG005"] +"core/app/apps/test_base_app_runner.py" = ["ARG001", "ARG002", "ARG005"] +"core/app/apps/test_pause_resume.py" = ["ANN401", "TID251"] +"core/app/apps/test_streaming_utils.py" = ["ARG001"] +"core/app/apps/test_workflow_app_generator.py" = ["ARG005"] +"core/app/apps/test_workflow_app_runner_core.py" = ["ARG001", "ARG002", "ARG004", "ARG005"] +"core/app/apps/test_workflow_app_runner_single_node.py" = ["ANN401", "TID251"] +"core/app/apps/test_workflow_pause_events.py" = ["ARG005"] +"core/app/apps/workflow/test_app_generator_extra.py" = ["ARG005"] +"core/app/apps/workflow/test_generate_task_pipeline_core.py" = ["ARG002", "ARG005"] +"core/app/features/rate_limiting/test_rate_limit.py" = ["ARG001"] +"core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py" = ["ARG002"] +"core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py" = ["ARG001", "ARG002", "ARG005"] +"core/app/task_pipeline/test_message_cycle_manager_optimization.py" = ["ARG002"] +"core/app/test_easy_ui_model_config_manager.py" = ["ARG005"] +"core/app/workflow/layers/test_persistence_inspector_publish.py" = ["ANN401", "ARG005", "TID251"] +"core/app/workflow/test_file_runtime.py" = ["ARG001", "ARG005"] +"core/app/workflow/test_observability_layer_extra.py" = ["ARG005"] +"core/app/workflow/test_persistence_layer.py" = ["ARG001"] +"core/base/test_app_generator_tts_publisher.py" = ["ARG002"] +"core/callback_handler/test_agent_tool_callback_handler.py" = ["ARG002"] +"core/callback_handler/test_workflow_tool_callback_handler.py" = ["ARG002"] +"core/datasource/__base/test_datasource_provider.py" = ["ARG002"] +"core/datasource/test_datasource_file_manager.py" = ["ARG001", "ARG002"] +"core/datasource/test_notion_provider.py" = ["ARG002", "TID251"] +"core/datasource/test_website_crawl.py" = ["ARG002"] +"core/datasource/utils/test_message_transformer.py" = ["ARG002"] +"core/entities/test_entities_mcp_provider.py" = ["ARG001"] +"core/entities/test_entities_provider_configuration.py" = ["ANN401", "ARG001", "ARG005", "TID251"] +"core/extension/test_extensible.py" = ["ARG002", "ARG005"] +"core/external_data_tool/api/test_api.py" = ["ARG001"] +"core/external_data_tool/test_base.py" = ["TID251"] +"core/external_data_tool/test_external_data_fetch.py" = ["ARG001"] +"core/helper/code_executor/test_code_executor.py" = ["TID251"] +"core/helper/code_executor/test_template_transformer.py" = ["ANN401", "TID251"] +"core/llm_generator/test_llm_generator.py" = ["ARG002"] +"core/mcp/auth/test_auth_flow.py" = ["ARG002"] +"core/mcp/client/test_session.py" = ["ARG001", "TID251"] +"core/mcp/client/test_sse.py" = ["ARG001", "TID251"] +"core/mcp/client/test_streamable_http.py" = ["ARG001", "ARG005", "S110", "TID251"] +"core/mcp/session/test_base_session.py" = ["S110"] +"core/mcp/session/test_client_session.py" = ["ARG005"] +"core/mcp/test_mcp_client.py" = ["ARG002"] +"core/memory/test_token_buffer_memory.py" = ["ARG002"] +"core/moderation/test_content_moderation.py" = ["TID251"] +"core/moderation/test_output_moderation.py" = ["ARG001", "ARG002"] +"core/ops/test_base_trace_instance.py" = ["ARG001"] +"core/ops/test_lookup_helpers.py" = ["ARG002"] +"core/ops/test_ops_trace_manager.py" = ["ARG001", "ARG002", "ARG005"] +"core/ops/test_trace_queue_manager.py" = ["ARG004"] +"core/ops/test_trace_session_metadata.py" = ["ARG001", "ARG005"] +"core/plugin/impl/test_agent_client.py" = ["ARG001"] +"core/plugin/impl/test_datasource_manager.py" = ["ARG001"] +"core/plugin/impl/test_oauth_handler.py" = ["ARG001"] +"core/plugin/impl/test_tool_manager.py" = ["ARG001"] +"core/plugin/impl/test_trigger_client.py" = ["ARG001"] +"core/plugin/test_endpoint_client.py" = ["ARG002"] +"core/plugin/test_model_runtime_adapter.py" = ["ARG002"] +"core/plugin/test_plugin_runtime.py" = ["ARG001", "ARG002", "TID251"] +"core/prompt/test_advanced_prompt_transform.py" = ["ARG005"] +"core/prompt/test_prompt_transform.py" = ["ARG005"] +"core/rag/datasource/keyword/jieba/test_jieba.py" = ["ARG001", "TID251"] +"core/rag/datasource/keyword/jieba/test_jieba_keyword_table_handler.py" = ["ARG004"] +"core/rag/datasource/keyword/test_keyword_factory.py" = ["ARG005"] +"core/rag/datasource/test_datasource_retrieval.py" = ["ARG001", "ARG002", "ARG005", "TID251"] +"core/rag/datasource/test_retrieval_attachment_access.py" = ["ARG005"] +"core/rag/datasource/vdb/test_vector_factory.py" = ["ARG002"] +"core/rag/embedding/test_embedding_base.py" = ["TID251"] +"core/rag/embedding/test_embedding_service.py" = ["ARG001"] +"core/rag/extractor/firecrawl/test_firecrawl.py" = ["TID251"] +"core/rag/extractor/test_csv_extractor.py" = ["ARG001", "ARG005"] +"core/rag/extractor/test_excel_extractor.py" = ["ARG002", "ARG005"] +"core/rag/extractor/test_extract_processor.py" = ["ARG001", "ARG005"] +"core/rag/extractor/test_helpers.py" = ["ARG002"] +"core/rag/extractor/test_markdown_extractor.py" = ["ARG001"] +"core/rag/extractor/test_notion_extractor.py" = ["ARG002", "ARG005"] +"core/rag/extractor/test_pdf_extractor.py" = ["ARG001", "ARG005"] +"core/rag/extractor/test_text_extractor.py" = ["ARG001"] +"core/rag/extractor/test_word_extractor.py" = ["ARG001", "ARG002", "ARG005"] +"core/rag/extractor/unstructured/test_unstructured_extractors.py" = ["ARG001", "ARG005"] +"core/rag/extractor/watercrawl/test_watercrawl.py" = ["ARG001", "ARG005", "TID251"] +"core/rag/indexing/processor/conftest.py" = ["ANN401", "ARG002", "TID251"] +"core/rag/indexing/processor/test_paragraph_index_processor.py" = ["ARG002", "TID251"] +"core/rag/indexing/processor/test_qa_index_processor.py" = ["ARG001", "TID251"] +"core/rag/indexing/test_index_processor.py" = ["ARG005"] +"core/rag/indexing/test_indexing_runner.py" = ["ARG005", "TID251"] +"core/rag/pipeline/test_queue.py" = ["ARG002"] +"core/rag/retrieval/test_dataset_retrieval.py" = ["ARG001", "ARG002", "ARG005", "TID251"] +"core/rag/splitter/test_text_splitter.py" = ["ARG002", "ARG005"] +"core/repositories/test_celery_workflow_execution_repository.py" = ["ARG002"] +"core/repositories/test_celery_workflow_node_execution_repository.py" = ["ARG002"] +"core/repositories/test_human_input_form_repository_impl.py" = ["ARG001", "ARG005"] +"core/repositories/test_human_input_repository.py" = ["ANN401", "ARG001", "ARG005", "TID251"] +"core/repositories/test_sqlalchemy_workflow_node_execution_repository.py" = ["ANN401", "ARG002", "ARG005", "TID251"] +"core/repositories/test_workflow_node_execution_truncation.py" = ["TID251"] +"core/schemas/test_resolver.py" = ["ARG005", "T201"] +"core/telemetry/test_facade.py" = ["ARG002", "ARG004"] +"core/telemetry/test_gateway_integration.py" = ["ARG002"] +"core/test_model_manager.py" = ["ARG001"] +"core/test_trigger_debug_event_selectors.py" = ["ARG002"] +"core/tools/test_base_tool.py" = ["ANN401", "ARG002", "TID251"] +"core/tools/test_builtin_tool_base.py" = ["ARG001", "ARG002", "TID251"] +"core/tools/test_builtin_tool_provider.py" = ["ARG001", "ARG005", "TID251"] +"core/tools/test_builtin_tools_extra.py" = ["ARG005"] +"core/tools/test_custom_tool.py" = ["ARG001", "ARG005", "TID251"] +"core/tools/test_dataset_retriever_tool.py" = ["ARG005"] +"core/tools/test_mcp_tool.py" = ["S110"] +"core/tools/test_tool_engine.py" = ["ANN401", "ARG002", "ARG005", "TID251"] +"core/tools/test_tool_file_manager.py" = ["ARG001"] +"core/tools/test_tool_label_manager.py" = ["TID251"] +"core/tools/test_tool_manager.py" = ["ANN401", "ARG001", "ARG005", "TID251"] +"core/tools/test_tool_provider_controller.py" = ["TID251"] +"core/tools/utils/test_configuration.py" = ["ARG001", "ARG002", "TID251"] +"core/tools/utils/test_encryption.py" = ["ANN401", "TID251"] +"core/tools/utils/test_message_transformer.py" = ["TID251"] +"core/tools/utils/test_misc_utils_extra.py" = ["ARG002"] +"core/tools/utils/test_model_invocation_utils.py" = ["ARG005", "TID251"] +"core/tools/utils/test_parser.py" = ["TID251"] +"core/tools/utils/test_web_reader_tool.py" = ["ARG001", "ARG002", "ARG005"] +"core/tools/workflow_as_tool/test_provider.py" = ["TID251"] +"core/tools/workflow_as_tool/test_tool.py" = ["ANN401", "ARG001", "ARG005", "TID251"] +"core/trigger/conftest.py" = ["ANN401", "TID251"] +"core/trigger/debug/test_debug_event_selectors.py" = ["ARG002", "TID251"] +"core/variables/test_segment_type_validation.py" = ["TID251"] +"core/workflow/context/test_execution_context.py" = ["ANN401", "ARG002", "S110", "TID251"] +"core/workflow/context/test_flask_app_context.py" = ["ARG002"] +"core/workflow/generator/test_runner.py" = ["ARG001", "ARG002", "TID251"] +"core/workflow/generator/test_runner_missing.py" = ["ARG003"] +"core/workflow/generator/test_tool_catalogue.py" = ["ARG002"] +"core/workflow/graph_engine/layers/test_observability.py" = ["ARG002"] +"core/workflow/graph_engine/test_mock_config.py" = ["TID251"] +"core/workflow/graph_engine/test_mock_factory.py" = ["TID251"] +"core/workflow/graph_engine/test_mock_nodes.py" = ["ANN401", "S110", "TID251"] +"core/workflow/graph_engine/test_parallel_human_input_join_resume.py" = ["ARG002", "TID251"] +"core/workflow/graph_engine/test_table_runner.py" = ["ARG001", "TID251"] +"core/workflow/nodes/agent_v2/test_agent_node.py" = ["ARG001", "ARG002", "ARG005"] +"core/workflow/nodes/agent_v2/test_ask_human_hitl.py" = ["ANN401", "TID251"] +"core/workflow/nodes/agent_v2/test_dify_tools_builder.py" = ["ANN401", "ARG001", "ARG002", "ARG005", "TID251"] +"core/workflow/nodes/agent_v2/test_output_adapter.py" = ["ARG005"] +"core/workflow/nodes/agent_v2/test_runtime_request_builder.py" = ["ARG002"] +"core/workflow/nodes/agent_v2/test_validators.py" = ["ARG001"] +"core/workflow/nodes/http_request/test_http_request_node.py" = ["ANN401", "ARG002", "TID251"] +"core/workflow/nodes/human_input/test_entities.py" = ["TID251"] +"core/workflow/nodes/human_input/test_human_input_form_filled_event.py" = ["TID251"] +"core/workflow/nodes/iteration/test_iteration_child_engine_errors.py" = ["ARG002", "TID251"] +"core/workflow/nodes/knowledge_index/test_knowledge_index_node.py" = ["ARG002"] +"core/workflow/nodes/knowledge_retrieval/test_knowledge_retrieval_node.py" = ["ARG002"] +"core/workflow/nodes/llm/test_node.py" = ["ARG002"] +"core/workflow/nodes/parameter_extractor/test_parameter_extractor_node.py" = ["TID251"] +"core/workflow/nodes/test_document_extractor_node.py" = ["ARG001"] +"core/workflow/nodes/tool/test_tool_node.py" = ["ANN401", "ARG002", "TID251"] +"core/workflow/nodes/webhook/test_webhook_file_conversion.py" = ["TID251"] +"core/workflow/nodes/webhook/test_webhook_node.py" = ["TID251"] +"core/workflow/test_form_input_serialization_compat.py" = ["ANN401", "TID251"] +"core/workflow/test_human_input_adapter.py" = ["ARG005"] +"core/workflow/test_node_factory.py" = ["ARG002"] +"core/workflow/test_workflow_entry.py" = ["ARG001"] +"enterprise/telemetry/test_enterprise_trace.py" = ["ARG002", "TID251"] +"enterprise/telemetry/test_exporter.py" = ["ARG001"] +"enterprise/telemetry/test_gateway.py" = ["ARG002"] +"events/event_handlers/test_delete_tool_parameters_cache_when_sync_draft_workflow.py" = ["ARG005"] +"extensions/logstore/test_sql_escape.py" = ["ARG001", "ARG002"] +"extensions/otel/decorators/handlers/test_generate_handler.py" = ["ARG001", "ARG002"] +"extensions/otel/decorators/handlers/test_workflow_app_runner_handler.py" = ["ARG001"] +"extensions/otel/decorators/test_base.py" = ["ARG002"] +"extensions/otel/decorators/test_handler.py" = ["ARG002"] +"extensions/otel/test_retrieval_tracing.py" = ["ARG001"] +"extensions/test_ext_request_logging.py" = ["ARG002"] +"extensions/test_redis.py" = ["ARG001"] +"factories/test_build_from_mapping.py" = ["ARG001"] +"factories/test_file_factory.py" = ["ARG001"] +"factories/test_variable_factory.py" = ["TID251"] +"fields/test_file_fields.py" = ["ARG005"] +"libs/_human_input/support.py" = ["TID251"] +"libs/broadcast_channel/redis/test_channel_unit_tests.py" = ["ARG002", "ARG005"] +"libs/broadcast_channel/redis/test_streams_channel_unit_tests.py" = ["ARG001", "ARG002", "TID251"] +"libs/test_cron_compatibility.py" = ["S110"] +"libs/test_email_i18n.py" = ["ANN401", "TID251"] +"libs/test_oauth_bearer_rate_limit_ordering.py" = ["ARG001"] +"libs/test_pyrefly_type_coverage.py" = ["TID251"] +"libs/test_schedule_utils_enhanced.py" = ["S110"] +"libs/test_sendgrid_client.py" = ["ARG001", "TID251"] +"libs/test_smtp_client.py" = ["TID251"] +"models/test_dataset_models.py" = ["ARG005"] +"models/test_plugin_entities.py" = ["TID251"] +"models/test_snippet.py" = ["ARG001"] +"oss/__mock/aliyun_oss.py" = ["ARG002"] +"oss/__mock/baidu_obs.py" = ["ARG002"] +"oss/__mock/base.py" = ["ARG002"] +"oss/__mock/tencent_cos.py" = ["ARG002"] +"oss/__mock/volcengine_tos.py" = ["ARG002"] +"oss/aliyun_oss/aliyun_oss/test_aliyun_oss.py" = ["ARG002"] +"oss/baidu_obs/test_baidu_obs.py" = ["ARG002"] +"oss/opendal/test_opendal.py" = ["ARG002"] +"oss/tencent_cos/test_tencent_cos.py" = ["ARG002"] +"oss/volcengine_tos/test_volcengine_tos.py" = ["ARG002"] +"services/agent/test_agent_observability_service.py" = ["ARG002", "ARG005"] +"services/agent/test_agent_services.py" = ["ARG001", "ARG002", "ARG003", "ARG005"] +"services/agent/test_composer_candidates.py" = ["ARG005"] +"services/agent/test_prompt_mentions.py" = ["ARG005"] +"services/agent/test_skill_tool_inference_service.py" = ["ARG001", "ARG005"] +"services/auth/test_jina_auth_standalone_module.py" = ["TID251"] +"services/controller_api.py" = ["ARG002"] +"services/data_migration/test_import_service.py" = ["ARG002", "ARG005"] +"services/dataset_service_test_helpers.py" = ["TID251"] +"services/enterprise/test_account_deletion_sync.py" = ["ARG001"] +"services/enterprise/test_rbac_service.py" = ["ARG002"] +"services/enterprise/test_traceparent_propagation.py" = ["ARG002"] +"services/hit_service.py" = ["TID251"] +"services/plugin/test_plugin_parameter_service.py" = ["ARG002"] +"services/rag_pipeline/pipeline_template/test_built_in_retrieval.py" = ["ARG001"] +"services/rag_pipeline/test_rag_pipeline_dsl_service.py" = ["ARG001", "ARG005", "T201", "TID251"] +"services/rag_pipeline/test_rag_pipeline_service.py" = ["ARG001", "ARG005"] +"services/rag_pipeline/test_rag_pipeline_task_proxy.py" = ["ARG001", "ARG005"] +"services/rag_pipeline/test_rag_pipeline_transform_service.py" = ["ARG001"] +"services/recommend_app/test_remote_retrieval.py" = ["ARG002"] +"services/retention/workflow_run/test_archive_download_preparation.py" = ["ARG002"] +"services/retention/workflow_run/test_archive_log_service.py" = ["ARG001", "ARG002"] +"services/retention/workflow_run/test_bundle_archive_maintenance.py" = ["TID251"] +"services/retention/workflow_run/test_restore_archived_workflow_run.py" = ["ARG002"] +"services/test_annotation_service.py" = ["ANN401", "TID251"] +"services/test_api_token_service.py" = ["ARG002"] +"services/test_app_generate_service.py" = ["ARG001", "ARG002", "ARG004"] +"services/test_app_generate_service_streaming_integration.py" = ["ARG002", "TID251"] +"services/test_archive_workflow_run_logs.py" = ["ARG002"] +"services/test_audio_service.py" = ["ARG002", "TID251"] +"services/test_batch_indexing_base.py" = ["ANN401", "TID251"] +"services/test_billing_service.py" = ["ARG001", "ARG002"] +"services/test_clear_free_plan_expired_workflow_run_logs.py" = ["ANN401", "ARG002", "ARG005", "TID251"] +"services/test_clear_free_plan_tenant_expired_logs.py" = ["ARG002", "ARG003"] +"services/test_dataset_service_document.py" = ["ARG002"] +"services/test_dataset_service_lock_not_owned.py" = ["ARG001", "ARG005"] +"services/test_dataset_service_segment.py" = ["ARG002"] +"services/test_datasource_provider_service.py" = ["ARG002"] +"services/test_external_dataset_service.py" = ["ARG002", "TID251"] +"services/test_feature_service_human_input_email_delivery.py" = ["ARG005"] +"services/test_feedback_service.py" = ["ARG002"] +"services/test_human_input_delivery_test_service.py" = ["ARG005"] +"services/test_knowledge_service.py" = ["TID251"] +"services/test_message_service.py" = ["ARG002"] +"services/test_messages_clean_service.py" = ["TID251"] +"services/test_model_load_balancing_service.py" = ["ANN401", "ARG001", "ARG005", "TID251"] +"services/test_model_provider_service.py" = ["ANN401", "TID251"] +"services/test_model_provider_service_sanitization.py" = ["ARG002", "ARG005"] +"services/test_oauth_server_service.py" = ["ARG002"] +"services/test_operation_service.py" = ["TID251"] +"services/test_rag_pipeline_task_proxy.py" = ["ARG002"] +"services/test_recommended_app_service.py" = ["ARG001"] +"services/test_schedule_service.py" = ["ANN401", "TID251"] +"services/test_snippet_service.py" = ["ARG001", "ARG002"] +"services/test_summary_index_service.py" = ["ARG001"] +"services/test_telemetry_service.py" = ["ARG001", "ARG005"] +"services/test_variable_truncator.py" = ["ARG002", "TID251"] +"services/test_variable_truncator_additional.py" = ["ANN401", "TID251"] +"services/test_vector_service.py" = ["ARG001", "TID251"] +"services/test_webhook_service_additional.py" = ["ANN401", "ARG002", "TID251"] +"services/test_website_service.py" = ["TID251"] +"services/test_workflow_comment_service.py" = ["ARG001", "ARG002"] +"services/test_workflow_run_service.py" = ["ANN401", "ARG002", "TID251"] +"services/test_workflow_service.py" = ["ANN401", "ARG002", "TID251"] +"services/tools/test_builtin_tools_manage_service.py" = ["ARG001", "ARG002"] +"services/tools/test_tools_manage_service.py" = ["ARG002"] +"services/workflow/test_inspector_events.py" = ["ANN401", "TID251"] +"services/workflow/test_node_output_inspector_service.py" = ["TID251"] +"services/workflow/test_workflow_converter_additional.py" = ["ANN401", "ARG001", "ARG005", "TID251"] +"services/workflow/test_workflow_event_snapshot_service.py" = ["ANN401", "ARG002", "ARG005", "TID251"] +"services/workflow/test_workflow_event_snapshot_service_additional.py" = ["ANN401", "ARG002", "ARG005", "TID251"] +"tasks/test_agent_backend_session_cleanup_task.py" = ["ARG005"] +"tasks/test_clean_dataset_task.py" = ["ARG002"] +"tasks/test_clean_document_task.py" = ["ARG002"] +"tasks/test_dataset_indexing_task.py" = ["ARG001", "ARG002"] +"tasks/test_document_indexing_sync_task.py" = ["ARG002"] +"tasks/test_duplicate_document_indexing_task.py" = ["ARG002"] +"tasks/test_human_input_timeout_tasks.py" = ["ARG001", "ARG002", "ARG005", "TID251"] +"tasks/test_initialize_created_app_rbac_access_task.py" = ["ARG005"] +"tasks/test_mail_send_task.py" = ["ARG002"] +"tasks/test_ops_trace_task.py" = ["ARG004"] +"tasks/test_process_tenant_plugin_autoupgrade_check_task.py" = ["ARG001"] +"tasks/test_remove_app_and_related_data_task.py" = ["ARG002"] +"tasks/test_trigger_processing_tasks.py" = ["ARG002"] +"tasks/test_workflow_execute_task.py" = ["ARG005"] +"test_app_factory.py" = ["ARG001"] +"test_pytest_dify.py" = ["ARG001"] +"tools/test_mcp_tool.py" = ["TID251"] + +[lint.flake8-tidy-imports.banned-api."flask_restx.reqparse"] +msg = "Use Pydantic payload/query models instead of reqparse." + +[lint.flake8-tidy-imports.banned-api."flask_restx.reqparse.RequestParser"] +msg = "Use Pydantic payload/query models instead of reqparse." + +[lint.flake8-tidy-imports.banned-api."typing.Any"] +msg = "Use object, Protocol, TypedDict, TypeVar, ParamSpec, or a localized cast instead." diff --git a/api/tests/unit_tests/clients/agent_backend/test_cleanup_composition_compositor_integration.py b/api/tests/unit_tests/clients/agent_backend/test_cleanup_composition_compositor_integration.py deleted file mode 100644 index 00e74b1b659..00000000000 --- a/api/tests/unit_tests/clients/agent_backend/test_cleanup_composition_compositor_integration.py +++ /dev/null @@ -1,134 +0,0 @@ -"""Integration test for the cleanup request against the real agenton compositor. - -The bug fixed by A+D was invisible to unit tests that use ``FakeAgentBackendRunClient`` -because the fake client never runs agenton's ``_validate_session_snapshot``. This -test plugs a cleanup request through the real ``Compositor`` (with the same -providers the agent backend wires in production) so that the snapshot-vs- -composition name-order check would fail loudly if the cleanup builder ever -regressed back to the empty-composition shape. -""" - -from __future__ import annotations - -from typing import cast - -import pytest -from agenton.compositor import Compositor, CompositorSessionSnapshot, LayerProvider -from agenton.compositor.schemas import LayerSessionSnapshot -from agenton.layers.base import LifecycleState -from agenton_collections.layers.plain import PLAIN_PROMPT_LAYER_TYPE_ID -from agenton_collections.layers.plain.basic import PromptLayer -from agenton_collections.layers.pydantic_ai import PYDANTIC_AI_HISTORY_LAYER_TYPE_ID, PydanticAIHistoryLayer - -from clients.agent_backend import AgentBackendRunRequestBuilder, RuntimeLayerSpec - - -def test_cleanup_request_passes_agenton_snapshot_validation(): - """The cleanup request's composition layer names must match the (filtered) - snapshot's layer names exactly — agenton's compositor enforces this and - the agent backend rejects mismatches as ``run_failed`` asynchronously, - which is the trap A/D fixed.""" - # Persisted (non-plugin) layer specs — these are what cleanup will replay. - # We exclude the dify.execution_context layer from this integration check - # because its real provider needs a plugin-daemon HTTP client; the cleanup - # validation we are exercising is the snapshot-vs-composition name check, - # which is purely structural and does not depend on which non-plugin layer - # types appear. - persisted_specs = [ - RuntimeLayerSpec( - name="workflow_node_job_prompt", - type=PLAIN_PROMPT_LAYER_TYPE_ID, - config={"prefix": "Do the cleanup."}, - ), - RuntimeLayerSpec(name="history", type=PYDANTIC_AI_HISTORY_LAYER_TYPE_ID), - ] - # Saved snapshot still carries the LLM layer entry — cleanup's - # ``_filter_snapshot_to_specs`` must drop it so names match. - full_snapshot = CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot( - name="workflow_node_job_prompt", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ), - LayerSessionSnapshot( - name="history", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={"messages": []}, - ), - LayerSessionSnapshot( - name="llm", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ), - ] - ) - - cleanup_request = AgentBackendRunRequestBuilder().build_cleanup_request( - session_snapshot=full_snapshot, - runtime_layer_specs=persisted_specs, - ) - - # Drive the real agenton compositor through ``from_config`` + ``_create_run`` - # the same way the agent backend's RunScheduler does. ``_create_run`` is the - # private path that calls ``_validate_session_snapshot``; we use it directly - # to keep the test synchronous (no async ``enter()`` lifecycle needed — - # validation is the only thing under test). - config = { - "schema_version": 1, - "layers": [ - {"name": layer.name, "type": layer.type, "deps": dict(layer.deps), "metadata": dict(layer.metadata)} - for layer in cleanup_request.composition.layers - ], - } - compositor = Compositor.from_config( - config, - providers=[ - LayerProvider.from_layer_type(PromptLayer), - LayerProvider.from_layer_type(PydanticAIHistoryLayer), - ], - ) - - layer_configs = {layer.name: layer.config for layer in cleanup_request.composition.layers} - # This is the call that would raise ``ValueError`` if the cleanup snapshot - # and composition disagreed on layer names — the exact failure mode the - # original ``layers=[]`` cleanup hit. - run = compositor._create_run( # type: ignore[reportPrivateUsage] - configs=cast(dict[str, object], layer_configs), - session_snapshot=cleanup_request.session_snapshot, - ) - assert list(run.slots.keys()) == ["workflow_node_job_prompt", "history"] - - -def test_cleanup_request_with_mismatched_specs_would_be_rejected_by_agenton(): - """Regression sentinel: if a future refactor stops filtering the snapshot, - agenton would reject the request — and that rejection is what the runtime - fix is preventing. We confirm the validator does fail when given the - pre-fix shape so the previous test's success is not a coincidence.""" - snapshot_with_extra = CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot( - name="history", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ), - LayerSessionSnapshot( - name="llm", # extra layer not in composition - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ), - ] - ) - compositor = Compositor.from_config( - { - "schema_version": 1, - "layers": [{"name": "history", "type": PYDANTIC_AI_HISTORY_LAYER_TYPE_ID, "deps": {}, "metadata": {}}], - }, - providers=[LayerProvider.from_layer_type(PydanticAIHistoryLayer)], - ) - - with pytest.raises(ValueError, match="layer names must match"): - compositor._create_run( # type: ignore[reportPrivateUsage] - configs={}, - session_snapshot=snapshot_with_extra, - ) diff --git a/api/tests/unit_tests/clients/agent_backend/test_client.py b/api/tests/unit_tests/clients/agent_backend/test_client.py index 421ffad9092..cf13067d239 100644 --- a/api/tests/unit_tests/clients/agent_backend/test_client.py +++ b/api/tests/unit_tests/clients/agent_backend/test_client.py @@ -40,6 +40,7 @@ def _request(): agent_mode="workflow_run", invoke_from="debugger", ), + backend_binding_ref="binding-ref-1", workflow_node_job_prompt="Do the task.", user_prompt="hello", ) diff --git a/api/tests/unit_tests/clients/agent_backend/test_event_adapter.py b/api/tests/unit_tests/clients/agent_backend/test_event_adapter.py index 14c51895df0..961e40fec7c 100644 --- a/api/tests/unit_tests/clients/agent_backend/test_event_adapter.py +++ b/api/tests/unit_tests/clients/agent_backend/test_event_adapter.py @@ -10,6 +10,7 @@ from dify_agent.protocol import ( RunCancelledEventData, RunFailedEvent, RunFailedEventData, + RunFailureType, RunStartedEvent, RunSucceededEvent, RunSucceededEventData, @@ -133,7 +134,11 @@ def test_event_adapter_maps_run_failed_to_failed_result(): RunFailedEvent( id="4-0", run_id="run-1", - data=RunFailedEventData(error="boom", reason="runtime"), + data=RunFailedEventData( + error="boom", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + reason="runtime", + ), ) ) @@ -142,6 +147,7 @@ def test_event_adapter_maps_run_failed_to_failed_result(): run_id="run-1", source_event_id="4-0", error="boom", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, reason="runtime", ) ] diff --git a/api/tests/unit_tests/clients/agent_backend/test_fake_client.py b/api/tests/unit_tests/clients/agent_backend/test_fake_client.py index 5862117f622..6ad82633036 100644 --- a/api/tests/unit_tests/clients/agent_backend/test_fake_client.py +++ b/api/tests/unit_tests/clients/agent_backend/test_fake_client.py @@ -23,6 +23,7 @@ def _request(): agent_mode="workflow_run", invoke_from="debugger", ), + backend_binding_ref="binding-ref-1", workflow_node_job_prompt="Do the task.", user_prompt="hello", ) diff --git a/api/tests/unit_tests/clients/agent_backend/test_request_builder.py b/api/tests/unit_tests/clients/agent_backend/test_request_builder.py index b2b77275fa7..6f2a10452c5 100644 --- a/api/tests/unit_tests/clients/agent_backend/test_request_builder.py +++ b/api/tests/unit_tests/clients/agent_backend/test_request_builder.py @@ -1,10 +1,7 @@ from typing import Any, cast import pytest -from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot from agenton.layers import ExitIntent -from agenton.layers.base import LifecycleState from agenton_collections.layers.plain import PLAIN_PROMPT_LAYER_TYPE_ID, PromptLayerConfig from agenton_collections.layers.pydantic_ai import PYDANTIC_AI_HISTORY_LAYER_TYPE_ID from dify_agent.layers.dify_core_tools import ( @@ -45,8 +42,6 @@ from clients.agent_backend import ( AgentBackendOutputConfig, AgentBackendRunRequestBuilder, AgentBackendWorkflowNodeRunInput, - RuntimeLayerSpec, - extract_runtime_layer_specs, redact_for_agent_backend_log, ) from clients.agent_backend.request_builder import DIFY_DRIVE_LAYER_ID, DIFY_SHELL_LAYER_ID @@ -71,6 +66,7 @@ def _run_input() -> AgentBackendWorkflowNodeRunInput: agent_mode="workflow_run", invoke_from="debugger", ), + backend_binding_ref="binding-ref-1", idempotency_key="workflow-run-1:node-execution-1", agent_soul_prompt="You are a careful reviewer.", workflow_node_job_prompt="Review the previous node output.", @@ -123,6 +119,29 @@ def test_request_builder_separates_agent_soul_and_workflow_job_prompt(): assert dumped["composition"]["layers"][2]["config"]["user"] == "Summarize the report." +def test_request_builder_forwards_plugin_specific_model_settings_via_extra_body(): + run_input = _run_input().model_copy( + update={ + "model": AgentBackendModelConfig( + plugin_id="langgenius/tongyi", + model_provider="tongyi", + model="qwen-plus-latest", + credentials={"api_key": "secret-key"}, + model_settings={"temperature": 0.7, "enable_thinking": True, "thinking_budget": 4096}, + ) + } + ) + + request = AgentBackendRunRequestBuilder().build_for_workflow_node(run_input) + layers = {layer.name: layer for layer in request.composition.layers} + model_config = cast(DifyPluginLLMLayerConfig, layers[DIFY_AGENT_MODEL_LAYER_ID].config) + + assert model_config.model_settings == { + "temperature": 0.7, + "extra_body": {"enable_thinking": True, "thinking_budget": 4096}, + } + + @pytest.mark.parametrize("agent_config_version_kind", ["snapshot", "draft"]) def test_agent_app_request_builder_keeps_agent_soul_prompt_for_snapshot_and_draft( agent_config_version_kind: str, @@ -271,89 +290,6 @@ def test_request_builder_adds_knowledge_layer_when_configured(): assert knowledge_config.sets[0].dataset_ids == ["dataset-1"] -def test_request_builder_can_delete_on_exit_for_cleanup_paths(): - run_input = _run_input() - run_input.suspend_on_exit = False - - request = AgentBackendRunRequestBuilder().build_for_workflow_node(run_input) - - assert request.on_exit.default is ExitIntent.DELETE - - -def test_request_builder_builds_cleanup_request_replays_persisted_layer_specs(): - """The cleanup request must replay the persisted (non-plugin) layer specs - and filter the snapshot to match so the agenton compositor's - snapshot-vs-composition name-order validator passes.""" - session_snapshot = CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot(name="history", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={"k": 1}), - LayerSessionSnapshot(name="llm", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={}), - ] - ) - specs = [RuntimeLayerSpec(name="history", type="pydantic_ai.history")] - - request = AgentBackendRunRequestBuilder().build_cleanup_request( - session_snapshot=session_snapshot, - runtime_layer_specs=specs, - idempotency_key="run-1:node-1:binding-1:agent-session-cleanup", - metadata={"workflow_run_id": "run-1"}, - ) - - assert [layer.name for layer in request.composition.layers] == ["history"] - assert request.session_snapshot is not None - assert [layer.name for layer in request.session_snapshot.layers] == ["history"] - assert request.on_exit.default is ExitIntent.DELETE - assert request.idempotency_key == "run-1:node-1:binding-1:agent-session-cleanup" - assert request.metadata["agent_backend_lifecycle"] == "session_cleanup" - assert "purpose" not in request.model_dump(mode="json") - - -def test_request_builder_rejects_empty_runtime_layer_specs(): - """Empty specs would put us back in the original ``layers=[]`` trap that - fails on agenton's snapshot-vs-composition validation.""" - with pytest.raises(ValueError, match="runtime_layer_specs"): - AgentBackendRunRequestBuilder().build_cleanup_request( - session_snapshot=CompositorSessionSnapshot(layers=[]), - runtime_layer_specs=[], - ) - - -def test_extract_runtime_layer_specs_drops_plugin_layers_keeps_configs(): - from dify_agent.protocol import RunComposition, RunLayerSpec - - composition = RunComposition( - layers=[ - RunLayerSpec( - name="agent_soul_prompt", - type="plain.prompt", - config=PromptLayerConfig(prefix="hello"), - ), - RunLayerSpec( - name="llm", - type="dify.plugin.llm", - config=None, # protocol allows None; the redacted config is what matters - ), - RunLayerSpec( - name="tools", - type="dify.plugin.tools", - ), - RunLayerSpec( - name="history", - type="pydantic_ai.history", - ), - ] - ) - - specs = extract_runtime_layer_specs(composition) - - assert [spec.name for spec in specs] == ["agent_soul_prompt", "history"] - # Non-plugin configs are dumped as JSON-compatible dicts so the persisted - # row can be replayed without holding live pydantic instances. - soul_config = specs[0].config - assert isinstance(soul_config, dict) - assert soul_config.get("prefix") == "hello" - - def test_request_builder_rejects_blank_prompts(): with pytest.raises(ValidationError): AgentBackendWorkflowNodeRunInput( @@ -397,6 +333,7 @@ def _agent_app_input(*, include_shell: bool = False) -> AgentBackendAgentAppRunI agent_mode="agent_app", invoke_from="web-app", ), + backend_binding_ref="binding-ref-1", agent_soul_prompt="You are Iris.", user_prompt="List files.", include_shell=include_shell, @@ -422,7 +359,7 @@ def test_workflow_request_builder_adds_shell_layer_when_include_shell(): assert shell.type == DIFY_SHELL_LAYER_TYPE_ID # The shell layer depends on execution_context so the agent server can mint # per-command Agent Stub env for sandbox CLI forwarding. - assert shell.deps == {"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID} + assert shell.deps == {"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID, "runtime": "runtime"} shell_config = cast(DifyShellLayerConfig, shell.config) assert shell_config.env[0].name == "PROJECT_NAME" @@ -436,7 +373,10 @@ def test_workflow_request_builder_binds_drive_to_shell_when_configured(): layers = {layer.name: layer for layer in request.composition.layers} layer_names = [layer.name for layer in request.composition.layers] - assert layers[DIFY_SHELL_LAYER_ID].deps == {"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID} + assert layers[DIFY_SHELL_LAYER_ID].deps == { + "execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID, + "runtime": "runtime", + } shell_config = cast(DifyShellLayerConfig, layers[DIFY_SHELL_LAYER_ID].config) assert shell_config.agent_stub_drive_ref == "agent-agent-1" assert layers[DIFY_DRIVE_LAYER_ID].deps == {"shell": DIFY_SHELL_LAYER_ID} @@ -470,7 +410,10 @@ def test_agent_app_request_builder_adds_shell_layer_when_include_shell(): assert DIFY_SHELL_LAYER_ID in layers assert layers[DIFY_SHELL_LAYER_ID].type == DIFY_SHELL_LAYER_TYPE_ID - assert layers[DIFY_SHELL_LAYER_ID].deps == {"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID} + assert layers[DIFY_SHELL_LAYER_ID].deps == { + "execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID, + "runtime": "runtime", + } shell_config = cast(DifyShellLayerConfig, layers[DIFY_SHELL_LAYER_ID].config) assert shell_config.env[0].name == "APP_ENV" @@ -483,7 +426,10 @@ def test_agent_app_request_builder_binds_drive_to_shell_when_configured(): layers = {layer.name: layer for layer in request.composition.layers} layer_names = [layer.name for layer in request.composition.layers] - assert layers[DIFY_SHELL_LAYER_ID].deps == {"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID} + assert layers[DIFY_SHELL_LAYER_ID].deps == { + "execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID, + "runtime": "runtime", + } shell_config = cast(DifyShellLayerConfig, layers[DIFY_SHELL_LAYER_ID].config) assert shell_config.agent_stub_drive_ref == "agent-agent-1" assert layers[DIFY_DRIVE_LAYER_ID].deps == {"shell": DIFY_SHELL_LAYER_ID} diff --git a/api/tests/unit_tests/clients/agent_backend/test_session_cleanup.py b/api/tests/unit_tests/clients/agent_backend/test_session_cleanup.py deleted file mode 100644 index 6b72850aeca..00000000000 --- a/api/tests/unit_tests/clients/agent_backend/test_session_cleanup.py +++ /dev/null @@ -1,124 +0,0 @@ -from datetime import UTC, datetime - -from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot -from agenton.layers.base import LifecycleState -from dify_agent.protocol import RunStatusResponse - -from clients.agent_backend import ( - AgentBackendError, - AgentBackendSessionCleanupPayload, - FakeAgentBackendRunClient, - RuntimeLayerSpec, - cleanup_agent_backend_session, -) - - -def _payload() -> AgentBackendSessionCleanupPayload: - return AgentBackendSessionCleanupPayload( - session_snapshot=CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot( - name="history", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ) - ] - ), - runtime_layer_specs=[RuntimeLayerSpec(name="history", type="pydantic_ai.history")], - idempotency_key="cleanup-1", - metadata={"tenant_id": "tenant-1"}, - timeout_seconds=15.0, - ) - - -def test_cleanup_agent_backend_session_runs_create_and_wait_until_success(): - client = FakeAgentBackendRunClient(run_id="cleanup-run-1") - - result = cleanup_agent_backend_session(payload=_payload(), client=client) - - assert result.status == "succeeded" - assert result.cleanup_run_id == "cleanup-run-1" - assert client.request is not None - assert [layer.name for layer in client.request.composition.layers] == ["history"] - - -def test_cleanup_agent_backend_session_skips_when_client_is_missing(): - result = cleanup_agent_backend_session( - payload=_payload(), - client=None, - ) - - assert result.status == "skipped" - assert result.reason == "no_agent_backend_client" - - -def test_cleanup_agent_backend_session_skips_when_session_snapshot_is_missing(): - payload = _payload().model_copy(update={"session_snapshot": None}) - - result = cleanup_agent_backend_session(payload=payload, client=FakeAgentBackendRunClient()) - - assert result.status == "skipped" - assert result.reason == "missing_session_snapshot" - - -def test_cleanup_agent_backend_session_skips_when_runtime_layer_specs_are_missing(): - payload = _payload().model_copy(update={"runtime_layer_specs": []}) - - result = cleanup_agent_backend_session(payload=payload, client=FakeAgentBackendRunClient()) - - assert result.status == "skipped" - assert result.reason == "missing_runtime_layer_specs" - - -class _FailedStatusClient(FakeAgentBackendRunClient): - def wait_run(self, run_id: str, *, timeout_seconds: float | None = None) -> RunStatusResponse: - del timeout_seconds - return RunStatusResponse( - run_id=run_id, - status="failed", - created_at=datetime(2026, 1, 1, tzinfo=UTC), - updated_at=datetime(2026, 1, 1, tzinfo=UTC), - error="snapshot mismatch", - ) - - -def test_cleanup_agent_backend_session_reports_failed_terminal_status(): - client = _FailedStatusClient(run_id="cleanup-run-2") - - result = cleanup_agent_backend_session(payload=_payload(), client=client) - - assert result.status == "failed" - assert result.reason == "snapshot mismatch" - assert result.cleanup_run_id == "cleanup-run-2" - - -class _CreateRunFailureClient(FakeAgentBackendRunClient): - def create_run(self, request): # type: ignore[override] - del request - raise AgentBackendError("create run failed") - - -class _WaitRunFailureClient(FakeAgentBackendRunClient): - def wait_run(self, run_id: str, *, timeout_seconds: float | None = None) -> RunStatusResponse: - del run_id, timeout_seconds - raise AgentBackendError("wait run failed") - - -def test_cleanup_agent_backend_session_returns_failed_when_create_run_raises(): - result = cleanup_agent_backend_session(payload=_payload(), client=_CreateRunFailureClient()) - - assert result.status == "failed" - assert result.reason == "create run failed" - assert result.cleanup_run_id is None - - -def test_cleanup_agent_backend_session_returns_failed_with_cleanup_run_id_when_wait_run_raises(): - result = cleanup_agent_backend_session( - payload=_payload(), - client=_WaitRunFailureClient(run_id="cleanup-run-3"), - ) - - assert result.status == "failed" - assert result.reason == "wait run failed" - assert result.cleanup_run_id == "cleanup-run-3" diff --git a/api/tests/unit_tests/commands/test_account_commands.py b/api/tests/unit_tests/commands/test_account_commands.py new file mode 100644 index 00000000000..bbcc5d3e704 --- /dev/null +++ b/api/tests/unit_tests/commands/test_account_commands.py @@ -0,0 +1,38 @@ +from unittest.mock import Mock + +import pytest +from click.testing import CliRunner + +from commands.account import reset_email, reset_password + + +def test_reset_password_does_not_swallow_keyboard_interrupt(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + "commands.account.AccountService.get_account_by_email_with_case_fallback", + Mock(return_value=Mock()), + ) + monkeypatch.setattr("commands.account.valid_password", Mock(side_effect=KeyboardInterrupt)) + + result = CliRunner().invoke( + reset_password, + ["--email", "a@example.com", "--new-password", "whatever", "--password-confirm", "whatever"], + ) + + assert not isinstance(result.exception, SystemExit) or result.exception.code != 0 + assert "Invalid password" not in result.output + + +def test_reset_email_does_not_swallow_keyboard_interrupt(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + "commands.account.AccountService.get_account_by_email_with_case_fallback", + Mock(return_value=Mock()), + ) + monkeypatch.setattr("commands.account.email_validate", Mock(side_effect=KeyboardInterrupt)) + + result = CliRunner().invoke( + reset_email, + ["--email", "a@example.com", "--new-email", "b@example.com", "--email-confirm", "b@example.com"], + ) + + assert not isinstance(result.exception, SystemExit) or result.exception.code != 0 + assert "Invalid email" not in result.output diff --git a/api/tests/unit_tests/commands/test_archive_workflow_runs.py b/api/tests/unit_tests/commands/test_archive_workflow_runs.py index 60766b71ca8..c92cd8657da 100644 --- a/api/tests/unit_tests/commands/test_archive_workflow_runs.py +++ b/api/tests/unit_tests/commands/test_archive_workflow_runs.py @@ -1,12 +1,35 @@ +"""Tests for workflow-run archive command database boundaries. + +Planning deliberately creates a fresh session for every tenant prefix and for +every database retry. SQLite-backed tests keep the query, filtering, counting, +and session lifecycle real while clocks, billing lookup, and command failures +remain narrow external-boundary substitutions. +""" + import datetime -from unittest.mock import MagicMock +import logging +from dataclasses import dataclass +from types import SimpleNamespace +from unittest.mock import MagicMock, Mock +from uuid import uuid4 import click import pytest from click.testing import CliRunner +from sqlalchemy import Engine, event, text from sqlalchemy.exc import OperationalError +from sqlalchemy.orm import Session, scoped_session, sessionmaker from commands import retention +from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus +from models.base import TypeBase +from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom +from models.workflow import ( + WorkflowNodeExecutionModel, + WorkflowNodeExecutionTriggeredFrom, + WorkflowRun, + WorkflowType, +) from services.retention.workflow_run import bundle_archive_maintenance from services.retention.workflow_run.bundle_archive_maintenance import ( BundleOperationResult, @@ -27,11 +50,74 @@ def _db_disconnect_error() -> OperationalError: ) -def _session_context(session): - context = MagicMock() - context.__enter__.return_value = session - context.__exit__.return_value = False - return context +@dataclass(frozen=True) +class ArchiveDatabase: + """Creates archive candidates and their node executions in SQLite.""" + + session_maker: sessionmaker[Session] + end_before: datetime.datetime + + def add_run( + self, + tenant_id: str, + *, + created_at: datetime.datetime | None = None, + status: WorkflowExecutionStatus = WorkflowExecutionStatus.SUCCEEDED, + run_type: WorkflowType = WorkflowType.WORKFLOW, + ) -> str: + run_id = str(uuid4()) + with self.session_maker.begin() as session: + session.add( + WorkflowRun( + id=run_id, + tenant_id=tenant_id, + app_id=str(uuid4()), + workflow_id=str(uuid4()), + type=run_type, + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, + version="1", + graph="{}", + inputs="{}", + status=status, + outputs="{}", + error=None, + elapsed_time=0, + total_tokens=0, + total_steps=1, + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + created_at=created_at or self.end_before - datetime.timedelta(days=1), + ) + ) + return run_id + + def add_node(self, run_id: str, tenant_id: str, *, index: int) -> None: + with self.session_maker.begin() as session: + session.add( + WorkflowNodeExecutionModel( + id=str(uuid4()), + tenant_id=tenant_id, + app_id=str(uuid4()), + workflow_id=str(uuid4()), + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + workflow_run_id=run_id, + index=index, + predecessor_node_id=None, + node_execution_id=None, + node_id=f"node-{index}", + node_type="start", + title="Start", + inputs="{}", + process_data="{}", + outputs="{}", + status=WorkflowNodeExecutionStatus.SUCCEEDED, + error=None, + elapsed_time=0, + execution_metadata="{}", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + ) + ) def _delete_summary( @@ -86,100 +172,178 @@ def test_archive_tenant_id_parser_keeps_omitted_scope_unset(): assert retention._parse_comma_separated_ids(None, param_name="tenant-ids") is None -def test_resolve_archive_tenant_ids_from_plan_uses_explicit_sessions(monkeypatch): - end_before = datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC) - sessions = [MagicMock(name="session-a"), MagicMock(name="session-b")] - session_maker = MagicMock(side_effect=[_session_context(sessions[0]), _session_context(sessions[1])]) - calls = [] +@pytest.fixture +def archive_db(sqlite_engine: Engine) -> ArchiveDatabase: + """Create only the workflow tables used by archive planning.""" - def get_candidate_tenants(session, prefix, *, start_from, end_before): - calls.append((session, prefix, start_from, end_before)) - return [f"{prefix}-paid", f"{prefix}-free"] + TypeBase.metadata.create_all( + sqlite_engine, + tables=[WorkflowRun.__table__, WorkflowNodeExecutionModel.__table__], + ) + return ArchiveDatabase( + session_maker=sessionmaker(bind=sqlite_engine, expire_on_commit=False), + end_before=datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC), + ) - monkeypatch.setattr(retention, "_get_archive_candidate_tenant_ids_by_prefix", get_candidate_tenants) + +def _tenant_id(prefix: str, suffix: int) -> str: + return f"{prefix}{suffix:07x}-0000-0000-0000-000000000000" + + +def test_resolve_archive_tenant_ids_from_plan_uses_fresh_real_sessions( + archive_db: ArchiveDatabase, monkeypatch: pytest.MonkeyPatch +) -> None: + paid_a = _tenant_id("a", 1) + free_a = _tenant_id("a", 2) + paid_b = _tenant_id("b", 1) + free_b = _tenant_id("b", 2) + for tenant_id in (paid_a, free_a, paid_b, free_b): + archive_db.add_run(tenant_id) + + # These decoys verify the real candidate query's time, status, and type filters. + archive_db.add_run( + _tenant_id("a", 3), + created_at=archive_db.end_before + datetime.timedelta(seconds=1), + ) + archive_db.add_run(_tenant_id("a", 4), status=WorkflowExecutionStatus.RUNNING) + archive_db.add_run(_tenant_id("a", 5), run_type=WorkflowType.CHAT) + + opened_sessions: list[Session] = [] + + def record_session(session: Session, _transaction: object, _connection: object) -> None: + opened_sessions.append(session) + + event.listen(archive_db.session_maker.class_, "after_begin", record_session) monkeypatch.setattr( retention, "_filter_paid_workflow_archive_tenant_ids", - lambda tenant_ids: (["a-paid", "b-paid"], ["a-free", "b-free"]), + lambda tenant_ids: ([paid_a, paid_b], sorted(set(tenant_ids) - {paid_a, paid_b})), ) + try: + tenant_plan = retention._resolve_archive_tenant_ids_from_plan( + session_maker=archive_db.session_maker, + tenant_ids=None, + tenant_prefixes=["a", "b"], + start_from=None, + end_before=archive_db.end_before, + ) + finally: + event.remove(archive_db.session_maker.class_, "after_begin", record_session) - tenant_plan = retention._resolve_archive_tenant_ids_from_plan( - session_maker=session_maker, - tenant_ids=None, - tenant_prefixes=["a", "b"], - start_from=None, - end_before=end_before, - ) - - assert tenant_plan["archive_tenant_ids"] == ["a-paid", "b-paid"] - assert tenant_plan["paid_tenant_ids"] == ["a-paid", "b-paid"] - assert tenant_plan["unpaid_tenant_ids"] == ["a-free", "b-free"] - assert calls == [ - (sessions[0], "a", None, end_before), - (sessions[1], "b", None, end_before), - ] + assert tenant_plan == { + "archive_tenant_ids": [paid_a, paid_b], + "paid_tenant_ids": [paid_a, paid_b], + "unpaid_tenant_ids": [free_a, free_b], + } + assert len(opened_sessions) == 2 + assert opened_sessions[0] is not opened_sessions[1] -def test_safe_remove_scoped_session_discards_registry_and_disposes_after_remove_error(monkeypatch): - fake_db = MagicMock() - fake_db.session.remove.side_effect = RuntimeError("server closed the connection unexpectedly") - monkeypatch.setattr(retention, "db", fake_db) +def test_safe_remove_scoped_session_recovers_from_real_closed_connection( + sqlite_engine: Engine, + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + maker = sessionmaker(bind=sqlite_engine) + registry = scoped_session(maker) + registry().execute(text("select 1")) + dbapi_connection = registry().connection().connection.dbapi_connection + assert dbapi_connection is not None + dbapi_connection.close() + monkeypatch.setattr(retention, "db", SimpleNamespace(session=registry, engine=sqlite_engine)) - retention._safe_remove_scoped_session("archive workflow run command") + with caplog.at_level(logging.WARNING, logger="commands.retention"): + retention._safe_remove_scoped_session("archive workflow run command") - fake_db.session.remove.assert_called_once() - fake_db.session.registry.clear.assert_called_once() - fake_db.engine.dispose.assert_called_once() + assert not registry.registry.has() + assert any("Ignoring DB scoped-session cleanup error" in message for message in caplog.messages) -def test_archive_command_db_retry_retries_retryable_db_disconnect(monkeypatch): - operation = MagicMock(side_effect=[_db_disconnect_error(), "ok"]) - sleep = MagicMock() +def test_archive_command_db_retry_retries_retryable_db_disconnect(monkeypatch: pytest.MonkeyPatch) -> None: + attempts = iter([_db_disconnect_error(), "ok"]) + sleep = Mock() monkeypatch.setattr("services.retention.workflow_run.db_retry.time.sleep", sleep) - result = retention._run_archive_command_db_retry("archive plan", operation) + def operation() -> str: + result = next(attempts) + if isinstance(result, Exception): + raise result + return result - assert result == "ok" - assert operation.call_count == 2 + assert retention._run_archive_command_db_retry("archive plan", operation) == "ok" sleep.assert_called_once_with(1.0) -def test_archive_plan_prefix_stats_retries_count_query_with_fresh_session(monkeypatch): - end_before = datetime.datetime(2025, 4, 1, tzinfo=datetime.UTC) - sessions = [MagicMock(name="session-1"), MagicMock(name="session-2")] - sessions[0].scalar.side_effect = _db_disconnect_error() - sessions[1].scalar.side_effect = [7, 9] - session_maker = MagicMock(side_effect=[_session_context(sessions[0]), _session_context(sessions[1])]) - sleep = MagicMock() +def test_archive_plan_prefix_stats_retries_with_fresh_session_and_real_counts( + archive_db: ArchiveDatabase, sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch +) -> None: + tenant_id = _tenant_id("a", 1) + run_ids = [archive_db.add_run(tenant_id) for _ in range(7)] + for index in range(9): + archive_db.add_node(run_ids[index % len(run_ids)], tenant_id, index=index) - monkeypatch.setattr( - retention, - "_get_archive_candidate_tenant_ids_by_prefix", - lambda session, prefix, *, start_from, end_before: [f"{prefix}-tenant"], + # Decoys outside the selected prefix and archive window must not affect counts. + decoy_run_id = archive_db.add_run(_tenant_id("b", 1)) + archive_db.add_node(decoy_run_id, _tenant_id("b", 1), index=99) + archive_db.add_run( + tenant_id, + created_at=archive_db.end_before + datetime.timedelta(seconds=1), ) + + fail_next_query = True + + def disconnect_once( + _connection: object, + _cursor: object, + _statement: str, + _parameters: object, + _context: object, + _executemany: bool, + ) -> None: + nonlocal fail_next_query + if fail_next_query: + fail_next_query = False + raise _db_disconnect_error() + + opened_sessions: list[Session] = [] + + def record_session(session: Session, _transaction: object, _connection: object) -> None: + opened_sessions.append(session) + + sleep = Mock() monkeypatch.setattr("services.retention.workflow_run.db_retry.time.sleep", sleep) + event.listen(sqlite_engine, "before_cursor_execute", disconnect_once) + event.listen(archive_db.session_maker.class_, "after_begin", record_session) + try: + stats = retention._get_archive_plan_prefix_stats( + archive_db.session_maker, + "a", + start_from=None, + end_before=archive_db.end_before, + ) + finally: + event.remove(sqlite_engine, "before_cursor_execute", disconnect_once) + event.remove(archive_db.session_maker.class_, "after_begin", record_session) - stats = retention._get_archive_plan_prefix_stats( - session_maker, - "a", - start_from=None, - end_before=end_before, - ) - - assert stats["tenant_ids"] == ["a-tenant"] - assert stats["workflow_runs"] == 7 - assert stats["workflow_node_executions"] == 9 - assert session_maker.call_count == 2 + assert stats == { + "tenant_ids": [tenant_id], + "workflow_runs": 7, + "workflow_node_executions": 9, + } + assert len(opened_sessions) == 2 + assert opened_sessions[0] is not opened_sessions[1] sleep.assert_called_once_with(1.0) -def test_archive_workflow_runs_raises_click_exception_when_tenant_plan_fails(monkeypatch): - fake_db = MagicMock() - monkeypatch.setattr(retention, "db", fake_db) +def test_archive_workflow_runs_raises_click_exception_when_tenant_plan_fails( + sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch +) -> None: + registry = scoped_session(sessionmaker(bind=sqlite_engine)) + monkeypatch.setattr(retention, "db", SimpleNamespace(engine=sqlite_engine, session=registry)) monkeypatch.setattr( retention, "_resolve_archive_tenant_ids_from_plan", - MagicMock(side_effect=RuntimeError("tenant plan failed")), + Mock(side_effect=RuntimeError("tenant plan failed")), ) with pytest.raises(click.ClickException, match="Failed to resolve workflow archive tenant plan"): @@ -201,7 +365,7 @@ def test_archive_workflow_runs_raises_click_exception_when_tenant_plan_fails(mon ) -def test_delete_archived_workflow_runs_keeps_single_page_behavior_without_all_pages(monkeypatch): +def test_delete_archived_workflow_runs_keeps_single_page_behavior_without_all_pages(monkeypatch: pytest.MonkeyPatch): deleter = _patch_bundle_deleter( monkeypatch, [_delete_summary(processed=2, succeeded=2, next_catalog_id=_CURSOR_1)], @@ -220,7 +384,7 @@ def test_delete_archived_workflow_runs_keeps_single_page_behavior_without_all_pa assert deleter.delete_batch.call_args.kwargs["limit"] == 2 -def test_delete_archived_workflow_runs_all_pages_continues_until_empty_page(monkeypatch): +def test_delete_archived_workflow_runs_all_pages_continues_until_empty_page(monkeypatch: pytest.MonkeyPatch): deleter = _patch_bundle_deleter( monkeypatch, [ @@ -243,7 +407,9 @@ def test_delete_archived_workflow_runs_all_pages_continues_until_empty_page(monk ] -def test_delete_archived_workflow_runs_all_pages_fetches_empty_page_after_exact_full_page(monkeypatch): +def test_delete_archived_workflow_runs_all_pages_fetches_empty_page_after_exact_full_page( + monkeypatch: pytest.MonkeyPatch, +): deleter = _patch_bundle_deleter( monkeypatch, [ @@ -262,7 +428,7 @@ def test_delete_archived_workflow_runs_all_pages_fetches_empty_page_after_exact_ assert deleter.delete_batch.call_args_list[1].kwargs["after_catalog_id"] == _CURSOR_1 -def test_delete_archived_workflow_runs_all_pages_stops_at_first_failed_page(monkeypatch): +def test_delete_archived_workflow_runs_all_pages_stops_at_first_failed_page(monkeypatch: pytest.MonkeyPatch): failed_result = BundleOperationResult( catalog_id=_CURSOR_2, bundle_id="bundle-failed", @@ -291,7 +457,7 @@ def test_delete_archived_workflow_runs_all_pages_stops_at_first_failed_page(monk assert f"resume_after_catalog_id={_CURSOR_1}" in result.output -def test_delete_archived_workflow_runs_all_pages_fails_when_cursor_does_not_advance(monkeypatch): +def test_delete_archived_workflow_runs_all_pages_fails_when_cursor_does_not_advance(monkeypatch: pytest.MonkeyPatch): deleter = _patch_bundle_deleter( monkeypatch, [_delete_summary(processed=1, succeeded=1, next_catalog_id=None)], @@ -307,7 +473,7 @@ def test_delete_archived_workflow_runs_all_pages_fails_when_cursor_does_not_adva assert "cursor did not advance" in result.output.lower() -def test_delete_archived_workflow_runs_all_pages_uses_preview_cursor_for_dry_run(monkeypatch): +def test_delete_archived_workflow_runs_all_pages_uses_preview_cursor_for_dry_run(monkeypatch: pytest.MonkeyPatch): deleter = _patch_bundle_deleter( monkeypatch, [ @@ -328,7 +494,9 @@ def test_delete_archived_workflow_runs_all_pages_uses_preview_cursor_for_dry_run ] -def test_delete_archived_workflow_runs_dry_run_failure_separates_preview_and_destructive_cursors(monkeypatch): +def test_delete_archived_workflow_runs_dry_run_failure_separates_preview_and_destructive_cursors( + monkeypatch: pytest.MonkeyPatch, +): failed_result = BundleOperationResult( catalog_id=_CURSOR_2, bundle_id="bundle-failed", @@ -371,7 +539,7 @@ def test_delete_archived_workflow_runs_dry_run_failure_separates_preview_and_des assert f"destructive_retry_after_catalog_id={_CURSOR_0}" in result.output -def test_delete_archived_workflow_runs_all_pages_starts_after_explicit_cursor(monkeypatch): +def test_delete_archived_workflow_runs_all_pages_starts_after_explicit_cursor(monkeypatch: pytest.MonkeyPatch): deleter = _patch_bundle_deleter(monkeypatch, [_delete_summary(processed=0)]) result = CliRunner().invoke( @@ -413,7 +581,7 @@ def test_delete_archived_workflow_runs_rejects_invalid_run_shard_options(monkeyp deleter.delete_batch.assert_not_called() -def test_delete_archived_workflow_runs_passes_formatted_run_shard_to_service(monkeypatch): +def test_delete_archived_workflow_runs_passes_formatted_run_shard_to_service(monkeypatch: pytest.MonkeyPatch): deleter = _patch_bundle_deleter(monkeypatch, [_delete_summary(processed=0)]) result = CliRunner().invoke( @@ -441,7 +609,7 @@ def test_delete_archived_workflow_runs_passes_formatted_run_shard_to_service(mon assert deleter.delete_batch.call_args.kwargs["shard"] == "03-of-16" -def test_delete_archived_workflow_runs_rejects_mixed_catalog_shards_before_delete(monkeypatch): +def test_delete_archived_workflow_runs_rejects_mixed_catalog_shards_before_delete(monkeypatch: pytest.MonkeyPatch): deleter = _patch_bundle_deleter(monkeypatch, [_delete_summary(processed=0)]) deleter.validate_catalog_shards.side_effect = ValueError("unexpected shards: 00-of-01") diff --git a/api/tests/unit_tests/commands/test_check_no_new_getattr.py b/api/tests/unit_tests/commands/test_check_no_new_getattr.py index 4efcb37da92..a63569c706f 100644 --- a/api/tests/unit_tests/commands/test_check_no_new_getattr.py +++ b/api/tests/unit_tests/commands/test_check_no_new_getattr.py @@ -85,6 +85,42 @@ def main_branch_rev(repo: Path) -> str: return git(repo, "rev-parse", "main") +@pytest.mark.parametrize( + ("source_line", "rule_id"), + [ + ( + "value = getattr(module, name) # guard-ignore: no-new-getattr -- lazy export proxy", + "no-new-getattr", + ), + ( + "session.rollback() # guard-ignore: no-new-controller-sqlalchemy -- decorator owns rollback", + "no-new-controller-sqlalchemy", + ), + ], +) +def test_has_reasoned_guard_ignore_accepts_custom_rules(source_line: str, rule_id: str) -> None: + module = load_guard_module() + + assert module.has_reasoned_guard_ignore(source_line, rule_id) + + +@pytest.mark.parametrize( + ("source_line", "rule_id"), + [ + ("value = getattr(module, name) # noqa: no-new-getattr legacy marker", "no-new-getattr"), + ("value = getattr(module, name) # guard-ignore: no-new-getattr", "no-new-getattr"), + ( + "value = getattr(module, name) # guard-ignore: another-rule -- wrong rule", + "no-new-getattr", + ), + ], +) +def test_has_reasoned_guard_ignore_rejects_invalid_markers(source_line: str, rule_id: str) -> None: + module = load_guard_module() + + assert not module.has_reasoned_guard_ignore(source_line, rule_id) + + def test_resolve_ast_grep_command_prefers_ast_grep(monkeypatch: pytest.MonkeyPatch) -> None: module = load_guard_module() monkeypatch.setattr( @@ -776,7 +812,7 @@ def test_modified_hunk_with_increased_getattr_count_fails(tmp_path: Path) -> Non assert "net-new getattr" in result.stderr -def test_inline_noqa_suppression_with_explanatory_text_skips_added_getattr(tmp_path: Path) -> None: +def test_inline_guard_ignore_with_explanatory_text_skips_added_getattr(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -795,20 +831,20 @@ def test_inline_noqa_suppression_with_explanatory_text_skips_added_getattr(tmp_p "pkg/existing.py", """ def read_value(obj): - return getattr(obj, "dynamic_name", None) # noqa: no-new-getattr needed for plugin-defined attributes + return getattr(obj, "dynamic_name", None) # guard-ignore: no-new-getattr -- plugin-defined attributes """, ) commit_all(tmp_path, "add suppressed getattr") result = run_script(tmp_path, "--base-rev", base_rev) - assert "no-new-getattr needed for plugin-defined attributes" in (tmp_path / "pkg/existing.py").read_text( + assert "guard-ignore: no-new-getattr -- plugin-defined attributes" in (tmp_path / "pkg/existing.py").read_text( encoding="utf-8" ) assert result.returncode == 0, stderr_lines(result) -def test_inline_noqa_without_explanatory_text_is_not_sufficient(tmp_path: Path) -> None: +def test_inline_guard_ignore_without_explanatory_text_is_not_sufficient(tmp_path: Path) -> None: init_repo(tmp_path) write_repo_file( tmp_path, @@ -827,10 +863,10 @@ def test_inline_noqa_without_explanatory_text_is_not_sufficient(tmp_path: Path) "pkg/existing.py", """ def read_value(obj): - return getattr(obj, "dynamic_name", None) # noqa: no-new-getattr + return getattr(obj, "dynamic_name", None) # guard-ignore: no-new-getattr """, ) - commit_all(tmp_path, "add bare noqa getattr") + commit_all(tmp_path, "add bare guard ignore getattr") result = run_script(tmp_path, "--base-rev", base_rev) diff --git a/api/tests/unit_tests/commands/test_data_migration_wizard.py b/api/tests/unit_tests/commands/test_data_migration_wizard.py index 992a18080b3..37b72060f42 100644 --- a/api/tests/unit_tests/commands/test_data_migration_wizard.py +++ b/api/tests/unit_tests/commands/test_data_migration_wizard.py @@ -31,7 +31,7 @@ def test_parse_index_selection_supports_comma_indexes(): assert parse_index_selection("1, 3", ["a", "b", "c"]) == ["a", "c"] -def test_print_wizard_step_adds_separator(monkeypatch): +def test_print_wizard_step_adds_separator(monkeypatch: pytest.MonkeyPatch): output_lines = [] monkeypatch.setattr("commands.data_migration.click.echo", output_lines.append) @@ -45,7 +45,7 @@ def test_conflict_strategy_choices_exclude_replace(): assert CONFLICT_STRATEGY_CHOICES == ["fail", "skip", "update"] -def test_prompt_app_ids_explains_comma_selection_and_default(monkeypatch): +def test_prompt_app_ids_explains_comma_selection_and_default(monkeypatch: pytest.MonkeyPatch): from commands.data_migration import _prompt_app_ids prompts = [] @@ -67,7 +67,7 @@ def test_prompt_app_ids_explains_comma_selection_and_default(monkeypatch): assert "Currently supported app types: workflow and chatflow." in output_lines -def test_prompt_tool_category_marks_auto_discovered_tools(monkeypatch): +def test_prompt_tool_category_marks_auto_discovered_tools(monkeypatch: pytest.MonkeyPatch): output_lines = [] monkeypatch.setattr("commands.data_migration.click.echo", output_lines.append) @@ -85,7 +85,7 @@ def test_prompt_tool_category_marks_auto_discovered_tools(monkeypatch): assert output_lines[:2] == ["", "==== Custom API tools ===="] -def test_prompt_tool_category_explains_comma_selection_and_default(monkeypatch): +def test_prompt_tool_category_explains_comma_selection_and_default(monkeypatch: pytest.MonkeyPatch): prompts = [] def capture_prompt(text, **kwargs): @@ -110,7 +110,7 @@ def test_prompt_tool_category_explains_comma_selection_and_default(monkeypatch): ] -def test_prompt_output_file_shows_default(monkeypatch): +def test_prompt_output_file_shows_default(monkeypatch: pytest.MonkeyPatch): prompts = [] def capture_prompt(text, **kwargs): @@ -124,7 +124,7 @@ def test_prompt_output_file_shows_default(monkeypatch): assert prompts[0][1]["show_default"] is True -def test_prompt_tool_category_marks_auto_by_detail_and_supports_multi_select(monkeypatch): +def test_prompt_tool_category_marks_auto_by_detail_and_supports_multi_select(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr("commands.data_migration.click.echo", lambda *_args, **_kwargs: None) monkeypatch.setattr("commands.data_migration.click.prompt", lambda *args, **kwargs: "1,2") @@ -159,7 +159,7 @@ def test_prompt_tool_category_marks_auto_by_value(): assert "1. [auto] embedded_workflow_as_tool (tool-1)" in output_lines -def test_print_auto_tools_lists_each_category(monkeypatch): +def test_print_auto_tools_lists_each_category(monkeypatch: pytest.MonkeyPatch): output_lines = [] monkeypatch.setattr("commands.data_migration.click.echo", output_lines.append) @@ -208,7 +208,7 @@ def test_resolve_mcp_tool_names_does_not_compare_non_uuid_identifier_to_uuid_id( assert resolved == {provider.name: provider.id} -def test_prompt_additional_tools_prints_final_selection_when_skipped(monkeypatch): +def test_prompt_additional_tools_prints_final_selection_when_skipped(monkeypatch: pytest.MonkeyPatch): output_lines = [] confirm_prompts = [] @@ -236,7 +236,7 @@ def test_prompt_additional_tools_prints_final_selection_when_skipped(monkeypatch assert "- [auto] weather: 3bac3aa9-dd87-4351-9459-a7099137b028" in output_lines -def test_final_tool_selection_deduplicates_manual_tool_already_auto(monkeypatch): +def test_final_tool_selection_deduplicates_manual_tool_already_auto(monkeypatch: pytest.MonkeyPatch): output_lines = [] monkeypatch.setattr("commands.data_migration.click.echo", output_lines.append) @@ -259,7 +259,7 @@ def test_final_tool_selection_deduplicates_manual_tool_already_auto(monkeypatch) assert not any(line.startswith("- [manual]") for line in output_lines) -def test_prompt_output_file_rejects_yes_no_typo(monkeypatch): +def test_prompt_output_file_rejects_yes_no_typo(monkeypatch: pytest.MonkeyPatch): import click import pytest @@ -269,7 +269,7 @@ def test_prompt_output_file_rejects_yes_no_typo(monkeypatch): _prompt_output_file() -def test_confirm_wizard_summary_shows_conflict_strategy(monkeypatch): +def test_confirm_wizard_summary_shows_conflict_strategy(monkeypatch: pytest.MonkeyPatch): output_lines = [] confirm_prompts = [] @@ -298,7 +298,7 @@ def test_confirm_wizard_summary_shows_conflict_strategy(monkeypatch): assert confirm_prompts == [("Write migration package? [y/n, default: y]", {"default": True, "show_default": False})] -def test_confirm_wizard_summary_shows_final_deduplicated_tool_selection(monkeypatch): +def test_confirm_wizard_summary_shows_final_deduplicated_tool_selection(monkeypatch: pytest.MonkeyPatch): output_lines = [] monkeypatch.setattr("commands.data_migration.click.echo", output_lines.append) @@ -343,7 +343,7 @@ def test_confirm_wizard_summary_shows_final_deduplicated_tool_selection(monkeypa assert "- [manual] weather-id" not in output_lines -def test_import_options_prompts_explain_secrets_reuse_and_conflicts(monkeypatch): +def test_import_options_prompts_explain_secrets_reuse_and_conflicts(monkeypatch: pytest.MonkeyPatch): from commands.data_migration import _prompt_import_options output_lines = [] diff --git a/api/tests/unit_tests/commands/test_fix_app_site_missing.py b/api/tests/unit_tests/commands/test_fix_app_site_missing.py index a7b05e3bbb1..c69f4eb6463 100644 --- a/api/tests/unit_tests/commands/test_fix_app_site_missing.py +++ b/api/tests/unit_tests/commands/test_fix_app_site_missing.py @@ -1,77 +1,132 @@ +import uuid from types import SimpleNamespace from unittest.mock import MagicMock import pytest -from sqlalchemy.orm import Session +from sqlalchemy import event +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker from commands import system as system_commands +from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole +from models.enums import CustomizeTokenStrategy +from models.model import App, AppMode, IconType, Site -def test_fix_app_site_missing_passes_loaded_session_to_signal(monkeypatch: pytest.MonkeyPatch) -> None: - account = object() - tenant = MagicMock() - tenant.get_accounts.return_value = [account] - app = SimpleNamespace(id="app-1", tenant_id="tenant-1") +def _persist_missing_site_owner(session: Session) -> tuple[Account, App]: + """Persist an app without a Site and its complete tenant owner chain.""" + tenant = Tenant(name="Command workspace") + account = Account(name="Owner", email=f"owner-{uuid.uuid4()}@example.com") + membership = TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + current=True, + role=TenantAccountRole.OWNER, + ) + app = App( + id=str(uuid.uuid4()), + tenant_id=tenant.id, + name="Missing Site App", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#FFFFFF", + enable_site=True, + enable_api=False, + created_by=account.id, + ) + session.add_all([tenant, account, membership, app]) + session.commit() + return account, app - session = Session() + +def _site_for(app: App) -> Site: + return Site( + app_id=app.id, + title=app.name, + default_language="en-US", + customize_token_strategy=CustomizeTokenStrategy.UUID, + code=f"site-{app.id}", + ) + + +def _bind_command_database( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], +) -> tuple[scoped_session[Session], Session]: + """Expose a callable real scoped session through the Flask-SQLAlchemy shape.""" + command_sessions = scoped_session(sqlite_session_factory) + command_session = command_sessions() + monkeypatch.setattr( + system_commands, + "db", + SimpleNamespace(engine=sqlite_engine, session=command_sessions), + ) + return command_sessions, command_session + + +def test_fix_app_site_missing_passes_loaded_session_to_signal( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + account, app = _persist_missing_site_owner(sqlite_session) + command_sessions, command_session = _bind_command_database(monkeypatch, sqlite_engine, sqlite_session_factory) phase_events: list[str] = [] - scalar = MagicMock(return_value=app) - get = MagicMock(return_value=tenant) - commit = MagicMock(side_effect=lambda: phase_events.append("commit")) - monkeypatch.setattr(session, "scalar", scalar) - monkeypatch.setattr(session, "get", get) - monkeypatch.setattr(session, "commit", commit) + event.listen(command_session, "after_commit", lambda _session: phase_events.append("commit")) - scoped_session = MagicMock(return_value=session) - scoped_session.scalar.return_value = app + def create_site(sender: App, *, account: Account, session: Session) -> None: + phase_events.append("signal") + assert sender.id == app.id + assert account.id == account_id + assert session is command_session + session.add(_site_for(sender)) - connection = MagicMock() - connection.execute.side_effect = [[SimpleNamespace(id=app.id)], []] - engine = MagicMock() - engine.begin.return_value.__enter__.return_value = connection - - monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=scoped_session)) - send = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("signal")) + account_id = account.id + send = MagicMock(side_effect=create_site) monkeypatch.setattr(system_commands.app_was_created, "send", send) - system_commands.fix_app_site_missing.callback() + try: + system_commands.fix_app_site_missing.callback() + finally: + command_sessions.remove() - scoped_session.assert_called_once_with() - scalar.assert_called_once() - get.assert_called_once_with(system_commands.Tenant, app.tenant_id) - tenant.get_accounts.assert_called_once_with(session=session) - send.assert_called_once_with(app, account=account, session=session) - commit.assert_called_once_with() + send.assert_called_once() assert phase_events == ["signal", "commit"] - assert isinstance(send.call_args.kwargs["session"], Session) + sqlite_session.expire_all() + persisted_site = sqlite_session.query(Site).filter_by(app_id=app.id).one() + assert persisted_site.title == app.name -def test_fix_app_site_missing_rolls_back_when_signal_fails(monkeypatch: pytest.MonkeyPatch) -> None: - account = object() - tenant = MagicMock() - tenant.get_accounts.return_value = [account] - app = SimpleNamespace(id="app-1", tenant_id="tenant-1") - session = MagicMock() +def test_fix_app_site_missing_rolls_back_when_signal_fails( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + _account, app = _persist_missing_site_owner(sqlite_session) + command_sessions, command_session = _bind_command_database(monkeypatch, sqlite_engine, sqlite_session_factory) phase_events: list[str] = [] - session.scalar.return_value = app - session.get.return_value = tenant - session.rollback.side_effect = lambda: phase_events.append("rollback") + event.listen(command_session, "after_rollback", lambda _session: phase_events.append("rollback")) - connection = MagicMock() - connection.execute.side_effect = [[SimpleNamespace(id=app.id)], []] - engine = MagicMock() - engine.begin.return_value.__enter__.return_value = connection - - monkeypatch.setattr(system_commands, "db", SimpleNamespace(engine=engine, session=MagicMock(return_value=session))) - - def fail_signal(*_args, **_kwargs) -> None: + def fail_signal(sender: App, **_kwargs: object) -> None: phase_events.append("signal") + # Ensure the command's next raw scan terminates while its own transaction + # still exercises the rollback path. + with sqlite_session_factory() as observer: + observer.add(_site_for(sender)) + observer.commit() raise RuntimeError("failed") monkeypatch.setattr(system_commands.app_was_created, "send", MagicMock(side_effect=fail_signal)) - system_commands.fix_app_site_missing.callback() + try: + system_commands.fix_app_site_missing.callback() + finally: + command_sessions.remove() - session.rollback.assert_called_once_with() - session.commit.assert_not_called() assert phase_events == ["signal", "rollback"] + sqlite_session.expire_all() + assert sqlite_session.get(App, app.id) is not None diff --git a/api/tests/unit_tests/commands/test_generate_swagger_specs.py b/api/tests/unit_tests/commands/test_generate_swagger_specs.py index 7ec832c526c..8726067822b 100644 --- a/api/tests/unit_tests/commands/test_generate_swagger_specs.py +++ b/api/tests/unit_tests/commands/test_generate_swagger_specs.py @@ -115,6 +115,7 @@ def test_system_features_specs_exclude_backend_only_fields(tmp_path): written_paths = module.generate_specs(tmp_path) excluded_fields = { + "enable_trial_app", "is_allow_create_workspace", "max_plugin_package_size", "plugin_manager", @@ -223,7 +224,13 @@ def test_generate_specs_include_console_contract_shapes_for_schema_migration(tmp app_detail_schema = schemas["RecommendedAppDetailResponse"] assert app_detail_schema["properties"]["id"]["type"] == "string" assert app_detail_schema["properties"]["export_data"]["type"] == "string" - assert {"type": "boolean"} in app_detail_schema["properties"]["can_trial"]["anyOf"] + assert app_detail_schema["properties"]["can_trial"]["type"] == "boolean" + assert "anyOf" not in app_detail_schema["properties"]["can_trial"] + assert "can_trial" in app_detail_schema["required"] + app_list_item_schema = schemas["RecommendedAppResponse"] + assert app_list_item_schema["properties"]["can_trial"]["type"] == "boolean" + assert "anyOf" not in app_list_item_schema["properties"]["can_trial"] + assert "can_trial" in app_list_item_schema["required"] app_detail_nullable_schema = schemas["RecommendedAppDetailNullableResponse"] assert _response_schema(paths["/explore/apps/{app_id}"]["get"])["$ref"] == ( "#/components/schemas/RecommendedAppDetailNullableResponse" diff --git a/api/tests/unit_tests/commands/test_legacy_model_type_migration.py b/api/tests/unit_tests/commands/test_legacy_model_type_migration.py index 9eb73b82516..c5402df700e 100644 --- a/api/tests/unit_tests/commands/test_legacy_model_type_migration.py +++ b/api/tests/unit_tests/commands/test_legacy_model_type_migration.py @@ -6,6 +6,7 @@ import json import os import threading import time +from collections.abc import Iterator from datetime import datetime, timedelta from pathlib import Path from types import SimpleNamespace @@ -15,9 +16,12 @@ import pytest import sqlalchemy as sa from click.testing import CliRunner from sqlalchemy.exc import OperationalError +from sqlalchemy.orm import Session, SessionTransaction, sessionmaker from graphon.model_runtime.entities.model_entities import ModelType +from models import Dataset, DatasetPermission, DatasetPermissionEnum from models.account import Tenant +from models.base import TypeBase from models.enums import CredentialSourceType from models.provider import ProviderModel from tests.helpers.legacy_model_type_migration import ( @@ -59,6 +63,40 @@ def command_module(): ) +@pytest.fixture +def rbac_session(sqlite_engine: sa.Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]: + """Bind RBAC command reads to persisted SQLite dataset rows.""" + + TypeBase.metadata.create_all( + sqlite_engine, + tables=[Dataset.__table__, DatasetPermission.__table__], + ) + factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr("commands.rbac.session_factory.create_session", factory) + with factory() as session: + yield session + + +def _persist_dataset( + session: Session, + *, + dataset_id: str = "dataset-1", + tenant_id: str = "tenant-1", + permission: DatasetPermissionEnum = DatasetPermissionEnum.ONLY_ME, + created_by: str = "creator-account-1", +) -> Dataset: + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name=f"Dataset {dataset_id}", + permission=permission, + created_by=created_by, + ) + session.add(dataset) + session.commit() + return dataset + + def _parse_json_lines(output: io.StringIO) -> list[dict[str, object]]: return [json.loads(line) for line in output.getvalue().splitlines() if line.strip()] @@ -363,56 +401,35 @@ def test_dataset_permission_rbac_migration_maps_legacy_permissions_to_enum_scope def test_dataset_permission_rbac_migration_uses_dataset_creator_as_operator( command_module, + rbac_session: Session, monkeypatch: pytest.MonkeyPatch, ) -> None: rbac_module = importlib.import_module("commands.rbac") - dataset_row = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - permission="only_me", - created_by="creator-account-1", - ) - execute_results = [[dataset_row], [], []] + _persist_dataset(rbac_session) calls: list[dict[str, object]] = [] - session_closed = False - - class FakeExecuteResult: - def __init__(self, rows: list[object]) -> None: - self._rows = rows - - def all(self) -> list[object]: - return self._rows - - class FakeSession: - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, traceback) -> None: - nonlocal session_closed - session_closed = True - pass - - def execute(self, stmt): - return FakeExecuteResult(execute_results.pop(0)) - - class FakeSessionFactory: - @staticmethod - def create_session() -> FakeSession: - return FakeSession() + read_transaction_ended = False def fake_replace_whitelist(**kwargs): - assert session_closed is True + assert read_transaction_ended is True calls.append(kwargs) - monkeypatch.setattr(rbac_module, "session_factory", FakeSessionFactory) - monkeypatch.setattr(rbac_module.RBACService.DatasetAccess, "replace_whitelist", fake_replace_whitelist) + def _record_transaction_end(session: Session, transaction: object) -> None: + nonlocal read_transaction_ended + del transaction + if session.get_bind() is rbac_session.get_bind(): + read_transaction_ended = True - command_module.migrate_dataset_permissions_to_rbac.callback( - tenant_id=None, - dataset_id=None, - batch_size=500, - dry_run=False, - ) + sa.event.listen(Session, "after_transaction_end", _record_transaction_end) + monkeypatch.setattr(rbac_module.RBACService.DatasetAccess, "replace_whitelist", fake_replace_whitelist) + try: + command_module.migrate_dataset_permissions_to_rbac.callback( + tenant_id=None, + dataset_id=None, + batch_size=500, + dry_run=False, + ) + finally: + sa.event.remove(Session, "after_transaction_end", _record_transaction_end) assert calls[0]["tenant_id"] == "tenant-1" assert calls[0]["account_id"] == "creator-account-1" @@ -422,41 +439,19 @@ def test_dataset_permission_rbac_migration_uses_dataset_creator_as_operator( def test_dataset_permission_rbac_migration_dry_run_outputs_structured_proposed_changes( command_module, + rbac_session: Session, monkeypatch: pytest.MonkeyPatch, ) -> None: rbac_module = importlib.import_module("commands.rbac") - dataset_row = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - permission="partial_members", - created_by="creator-account-1", + dataset = _persist_dataset(rbac_session, permission=DatasetPermissionEnum.PARTIAL_TEAM) + rbac_session.add( + DatasetPermission( + dataset_id=dataset.id, + account_id="member-account-1", + tenant_id=dataset.tenant_id, + ) ) - permission_row = SimpleNamespace(dataset_id="dataset-1", account_id="member-account-1") - execute_results = [[dataset_row], [permission_row], []] - - class FakeExecuteResult: - def __init__(self, rows: list[object]) -> None: - self._rows = rows - - def all(self) -> list[object]: - return self._rows - - class FakeSession: - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, traceback) -> None: - pass - - def execute(self, stmt): - return FakeExecuteResult(execute_results.pop(0)) - - class FakeSessionFactory: - @staticmethod - def create_session() -> FakeSession: - return FakeSession() - - monkeypatch.setattr(rbac_module, "session_factory", FakeSessionFactory) + rbac_session.commit() monkeypatch.setattr( rbac_module.RBACService.DatasetAccess, "replace_whitelist", @@ -1306,50 +1301,36 @@ def test_provider_models_processing_uses_same_plan_locking_and_transaction_entry begin_calls: list[str] = [] configure_calls: list[str] = [] - class _FakeBeginContext: - def __init__(self, phase: str) -> None: - self._phase = phase + def _record_begin(session: Session, transaction: SessionTransaction) -> None: + if session.get_bind() is sqlite_engine and transaction.parent is None: + begin_calls.append(current_phase["name"]) - def __enter__(self) -> None: - begin_calls.append(self._phase) - - def __exit__(self, exc_type, exc, tb) -> bool: - return False - - class _FakeSession: - def __init__(self, phase: str) -> None: - self._phase = phase - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb) -> bool: - return False - - def begin(self) -> _FakeBeginContext: - return _FakeBeginContext(self._phase) - - def _fake_session_factory(engine: sa.Engine) -> _FakeSession: - return _FakeSession(current_phase["name"]) - - def _fake_build_plan(self, session, candidate, *, lock_rows: bool): + def _fake_build_plan(self, session: Session, candidate, *, lock_rows: bool): + assert session.get_bind() is sqlite_engine lock_rows_seen.append((current_phase["name"], lock_rows)) - return SimpleNamespace(group_row_ids=[str(candidate.row.id)], winner=None, loser_rows=[]) + return migration_module._ProviderModelGroupPlan( + group_row_ids=[str(candidate.row.id)], + winner=None, + loser_rows=[], + ) def _fake_emit_plan(self, plan, *, session, tx_id: str, business_key: dict[str, object]) -> None: return None - def _fake_configure(self, session) -> None: + def _fake_configure(self, session: Session) -> None: + assert session.get_bind() is sqlite_engine configure_calls.append(current_phase["name"]) - monkeypatch.setattr(migration_module, "_session_factory", _fake_session_factory) monkeypatch.setattr(migration_module.Migration, "_build_provider_model_group_plan", _fake_build_plan) monkeypatch.setattr(migration_module.Migration, "_emit_provider_model_group_plan", _fake_emit_plan) monkeypatch.setattr(migration_module.Migration, "_configure_lock_timeout", _fake_configure) - - dry_migration._process_provider_model_group(candidate, business_key) - current_phase["name"] = "apply" - apply_migration._process_provider_model_group(candidate, business_key) + sa.event.listen(Session, "after_transaction_create", _record_begin) + try: + dry_migration._process_provider_model_group(candidate, business_key) + current_phase["name"] = "apply" + apply_migration._process_provider_model_group(candidate, business_key) + finally: + sa.event.remove(Session, "after_transaction_create", _record_begin) assert [phase for phase, _ in lock_rows_seen] == ["dry", "apply"] assert lock_rows_seen[0][1] == lock_rows_seen[1][1] @@ -1392,6 +1373,22 @@ def test_process_load_balancing_model_config_row_logs_stacktrace_for_lock_timeou sqlite_engine: sa.Engine, monkeypatch: pytest.MonkeyPatch, ) -> None: + create_minimal_legacy_model_type_schema(sqlite_engine) + created_at = datetime(2025, 1, 1, 12, 0, 0) + _insert_load_balancing_model_config( + sqlite_engine, + row_id="40000000-0000-0000-0000-000000000001", + tenant_id="tenant-1", + provider_name="openai", + model_name="gpt-4o-mini", + model_type="text-generation", + name="credential", + encrypted_config="{}", + credential_id="50000000-0000-0000-0000-000000000001", + enabled=True, + created_at=created_at, + updated_at=created_at, + ) output = io.StringIO() migration = migration_module.Migration( tenant_id="tenant-1", @@ -1401,37 +1398,18 @@ def test_process_load_balancing_model_config_row_logs_stacktrace_for_lock_timeou model_types=(ModelType.LLM,), orm_models=(migration_module.LoadBalancingModelConfig,), ) - candidate = migration_module._RowWithRawModelType( - row=SimpleNamespace(id="lb-row-1"), - raw_model_type="text-generation", - canonical_model_type=ModelType.LLM, - ) + candidate = migration._load_load_balancing_model_config_candidates(None)[0] lock_timeout_exc = OperationalError("SELECT 1", {}, SimpleNamespace(pgcode="55P03")) + transaction_begins = 0 - class _FakeBeginContext: - def __enter__(self) -> None: - return None - - def __exit__(self, exc_type, exc, tb) -> bool: - return False - - class _FakeSession: - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb) -> bool: - return False - - def begin(self) -> _FakeBeginContext: - return _FakeBeginContext() - - def _fake_session_factory(engine: sa.Engine) -> _FakeSession: - return _FakeSession() + def _record_begin(session: Session, transaction: SessionTransaction) -> None: + nonlocal transaction_begins + if session.get_bind() is sqlite_engine and transaction.parent is None: + transaction_begins += 1 def _fake_reload(self, session, original_candidate, *, lock_rows: bool): raise lock_timeout_exc - monkeypatch.setattr(migration_module, "_session_factory", _fake_session_factory) monkeypatch.setattr(migration_module.Migration, "_configure_lock_timeout", lambda self, session: None) monkeypatch.setattr( migration_module.Migration, @@ -1439,17 +1417,22 @@ def test_process_load_balancing_model_config_row_logs_stacktrace_for_lock_timeou _fake_reload, ) - migration._process_load_balancing_model_config_row(candidate) + sa.event.listen(Session, "after_transaction_create", _record_begin) + try: + migration._process_load_balancing_model_config_row(candidate) + finally: + sa.event.remove(Session, "after_transaction_create", _record_begin) lines = _parse_json_lines(output) assert len(lines) == 1 assert lines[0]["event"] == "lock_timeout_skipped" attrs = cast(dict[str, object], lines[0]["attrs"]) assert attrs["table_name"] == "load_balancing_model_configs" - assert attrs["id"] == "lb-row-1" + assert attrs["id"] == str(candidate.row.id) assert attrs["error"] == str(lock_timeout_exc) assert isinstance(attrs["stacktrace"], str) assert "OperationalError" in attrs["stacktrace"] + assert transaction_begins == 1 def test_process_load_balancing_model_config_row_logs_update_after_sql_execution( @@ -1457,6 +1440,23 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution sqlite_engine: sa.Engine, monkeypatch: pytest.MonkeyPatch, ) -> None: + create_minimal_legacy_model_type_schema(sqlite_engine) + created_at = datetime(2025, 1, 1, 12, 0, 0) + row_id = "40000000-0000-0000-0000-000000000002" + _insert_load_balancing_model_config( + sqlite_engine, + row_id=row_id, + tenant_id="tenant-1", + provider_name="openai", + model_name="gpt-4o-mini", + model_type="text-generation", + name="credential", + encrypted_config="{}", + credential_id="50000000-0000-0000-0000-000000000002", + enabled=True, + created_at=created_at, + updated_at=created_at, + ) migration = migration_module.Migration( tenant_id="tenant-1", engine=sqlite_engine, @@ -1465,42 +1465,33 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution model_types=(ModelType.LLM,), orm_models=(migration_module.LoadBalancingModelConfig,), ) - candidate = migration_module._RowWithRawModelType( - row=SimpleNamespace(id="lb-row-1"), - raw_model_type="text-generation", - canonical_model_type=ModelType.LLM, - ) + candidate = migration._load_load_balancing_model_config_candidates(None)[0] action_log: list[str] = [] - class _FakeBeginContext: - def __enter__(self) -> None: + def _record_begin(session: Session, transaction: SessionTransaction) -> None: + if session.get_bind() is sqlite_engine and transaction.parent is None: action_log.append("begin") - def __exit__(self, exc_type, exc, tb) -> bool: - return False - - class _FakeSession: - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb) -> bool: - return False - - def begin(self) -> _FakeBeginContext: - return _FakeBeginContext() - - def execute(self, stmt) -> None: + def _record_sql( + connection: sa.Connection, + cursor: object, + statement: str, + parameters: object, + context: object, + executemany: bool, + ) -> None: + del connection, cursor, parameters, context, executemany + if statement.lstrip().upper().startswith("UPDATE"): action_log.append("sql_execute") - def _fake_session_factory(engine: sa.Engine) -> _FakeSession: - return _FakeSession() - def _fake_configure(self, session) -> None: action_log.append("configure_lock_timeout") - def _fake_reload(self, session, original_candidate, *, lock_rows: bool): + original_reload = migration_module.Migration._reload_load_balancing_model_config_candidate + + def _record_reload(self, session: Session, original_candidate, *, lock_rows: bool): action_log.append(f"reload_candidate:{lock_rows}") - return candidate + return original_reload(self, session, original_candidate, lock_rows=lock_rows) def _fake_log_row_updated(self, *args, **kwargs) -> None: action_log.append("log_row_updated") @@ -1508,12 +1499,11 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution def _fake_cache_cleanup(self, *, row_id: str, tx_id: str) -> None: action_log.append("cache_cleanup") - monkeypatch.setattr(migration_module, "_session_factory", _fake_session_factory) monkeypatch.setattr(migration_module.Migration, "_configure_lock_timeout", _fake_configure) monkeypatch.setattr( migration_module.Migration, "_reload_load_balancing_model_config_candidate", - _fake_reload, + _record_reload, ) monkeypatch.setattr(migration_module.Migration, "_log_row_updated", _fake_log_row_updated) monkeypatch.setattr( @@ -1522,7 +1512,13 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution _fake_cache_cleanup, ) - migration._process_load_balancing_model_config_row(candidate) + sa.event.listen(Session, "after_transaction_create", _record_begin) + sa.event.listen(sqlite_engine, "before_cursor_execute", _record_sql) + try: + migration._process_load_balancing_model_config_row(candidate) + finally: + sa.event.remove(sqlite_engine, "before_cursor_execute", _record_sql) + sa.event.remove(Session, "after_transaction_create", _record_begin) assert action_log == [ "begin", @@ -1532,6 +1528,10 @@ def test_process_load_balancing_model_config_row_logs_update_after_sql_execution "log_row_updated", "cache_cleanup", ] + with Session(sqlite_engine) as session: + persisted = session.get(migration_module.LoadBalancingModelConfig, row_id) + assert persisted is not None + assert persisted.model_type == ModelType.LLM def test_load_balancing_model_config_cache_delete_failure_logs_stacktrace( diff --git a/api/tests/unit_tests/commands/test_reset_encrypt_key_pair.py b/api/tests/unit_tests/commands/test_reset_encrypt_key_pair.py index 59a12a616b5..5e2f8ffaa09 100644 --- a/api/tests/unit_tests/commands/test_reset_encrypt_key_pair.py +++ b/api/tests/unit_tests/commands/test_reset_encrypt_key_pair.py @@ -16,6 +16,7 @@ from sqlalchemy.orm import Session import commands from commands import system as system_commands from core.tools.entities.tool_entities import ApiProviderSchemaType +from enums import DeploymentEdition from graphon.model_runtime.entities.model_entities import ModelType from models import Tenant from models.provider import Provider, ProviderModel, ProviderType @@ -87,7 +88,7 @@ def _bind_command_to_sqlite(monkeypatch: pytest.MonkeyPatch, session: Session) - def test_reset_aborts_when_not_self_hosted(monkeypatch, capsys): - monkeypatch.setattr(system_commands.dify_config, "EDITION", "CLOUD") + monkeypatch.setattr(system_commands.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) exit_code = _invoke_reset() captured = capsys.readouterr() @@ -106,7 +107,7 @@ def test_reset_purges_provider_and_tool_tables_for_each_tenant( ) -> 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.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}") _bind_command_to_sqlite(monkeypatch, sqlite_session) @@ -146,7 +147,7 @@ def test_reset_purges_provider_and_tool_tables_for_each_tenant( ) 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.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) monkeypatch.setattr(system_commands, "generate_key_pair", lambda tenant_id: f"new-key-{tenant_id}") _bind_command_to_sqlite(monkeypatch, sqlite_session) diff --git a/api/tests/unit_tests/configs/test_dify_config.py b/api/tests/unit_tests/configs/test_dify_config.py index 3588f6beef1..af671c77600 100644 --- a/api/tests/unit_tests/configs/test_dify_config.py +++ b/api/tests/unit_tests/configs/test_dify_config.py @@ -3,9 +3,11 @@ import os import pytest from flask import Flask from packaging.version import Version +from pydantic import SecretStr from yarl import URL from configs.app_config import DifyConfig +from enums import DeploymentEdition def _clear_environment(monkeypatch: pytest.MonkeyPatch) -> None: @@ -78,11 +80,12 @@ def test_dify_config(monkeypatch: pytest.MonkeyPatch): assert config.COMMIT_SHA == "" # default values - assert config.EDITION == "SELF_HOSTED" + assert config.DEPLOYMENT_EDITION is DeploymentEdition.COMMUNITY assert config.API_COMPRESSION_ENABLED is False assert config.AGENT_SHELL_ENABLED is True assert config.SENTRY_TRACES_SAMPLE_RATE == 1.0 assert config.TEMPLATE_TRANSFORM_MAX_LENGTH == 400_000 + assert config.GRAPH_ENGINE_SCALE_UP_THRESHOLD == 0 # annotated field with custom configured value assert config.HTTP_REQUEST_MAX_READ_TIMEOUT == 300 @@ -94,6 +97,44 @@ def test_dify_config(monkeypatch: pytest.MonkeyPatch): assert Version(config.project.version) >= Version("1.0.0") +@pytest.mark.parametrize( + ("environment_value", "expected"), + [ + pytest.param(None, "", id="unset"), + pytest.param("", "", id="empty"), + pytest.param("expected", "expected", id="ascii"), + pytest.param("pässwörd-🔐", "pässwörd-🔐", id="unicode"), + ], +) +def test_init_password_defaults_to_empty_and_preserves_environment_value( + monkeypatch: pytest.MonkeyPatch, + environment_value: str | None, + expected: str, +) -> None: + _set_basic_config_env(monkeypatch) + if environment_value is None: + monkeypatch.delenv("INIT_PASSWORD", raising=False) + else: + monkeypatch.setenv("INIT_PASSWORD", environment_value) + + config = DifyConfig(_env_file=None) + + assert expected == config.INIT_PASSWORD + + +@pytest.mark.parametrize("edition", list(DeploymentEdition)) +def test_deployment_edition_is_loaded_from_environment( + monkeypatch: pytest.MonkeyPatch, + edition: DeploymentEdition, +) -> None: + _set_basic_config_env(monkeypatch) + monkeypatch.setenv("DEPLOYMENT_EDITION", edition.value) + + config = DifyConfig(_env_file=None) + + assert config.DEPLOYMENT_EDITION is edition + + def test_new_user_default_plugin_ids_are_parsed_from_env(monkeypatch: pytest.MonkeyPatch) -> None: _set_basic_config_env(monkeypatch) monkeypatch.setenv( @@ -109,6 +150,18 @@ def test_new_user_default_plugin_ids_are_parsed_from_env(monkeypatch: pytest.Mon ] +def test_turnstile_config_is_parsed_from_env(monkeypatch: pytest.MonkeyPatch) -> None: + _set_basic_config_env(monkeypatch) + monkeypatch.setenv("TURNSTILE_SECRET_KEY", " test-secret ") + monkeypatch.setenv("TURNSTILE_ALLOWED_HOSTNAMES", "dify.dev, Login.Example.COM. ") + + config = DifyConfig(_env_file=None) + + assert isinstance(config.TURNSTILE_SECRET_KEY, SecretStr) + assert config.TURNSTILE_SECRET_KEY.get_secret_value() == "test-secret" + assert frozenset({"dify.dev", "login.example.com"}) == config.TURNSTILE_ALLOWED_HOSTNAME_SET + + def test_plugin_remote_install_port_rejects_host_port_spec(monkeypatch: pytest.MonkeyPatch) -> None: """A 'host:port' compose publish spec must produce an actionable error, not an opaque int_parsing traceback.""" _set_basic_config_env(monkeypatch) @@ -201,6 +254,16 @@ def test_internal_files_url_prefers_explicit_value(monkeypatch: pytest.MonkeyPat assert config.INTERNAL_FILES_URL == "http://files-internal:5001" +def test_empty_files_url_overrides_console_api_url_for_relative_browser_uris(monkeypatch: pytest.MonkeyPatch): + _clear_environment(monkeypatch) + monkeypatch.setenv("FILES_URL", "") + monkeypatch.setenv("CONSOLE_API_URL", "http://api:5001") + + config = DifyConfig(_env_file=None) + + assert config.FILES_URL == "" + + # NOTE: If there is a `.env` file in your Workspace, this test might not succeed as expected. # This is due to `pymilvus` loading all the variables from the `.env` file into `os.environ`. def test_flask_configs(monkeypatch: pytest.MonkeyPatch): @@ -227,7 +290,7 @@ def test_flask_configs(monkeypatch: pytest.MonkeyPatch): # configs read from pydantic-settings assert config["LOG_LEVEL"] == "INFO" assert config["COMMIT_SHA"] == "" - assert config["EDITION"] == "SELF_HOSTED" + assert config["DEPLOYMENT_EDITION"] is DeploymentEdition.COMMUNITY assert config["API_COMPRESSION_ENABLED"] is False assert config["SENTRY_TRACES_SAMPLE_RATE"] == 1.0 diff --git a/api/tests/unit_tests/configs/test_env_consistency.py b/api/tests/unit_tests/configs/test_env_consistency.py index 81e08638145..1afdd307b8c 100644 --- a/api/tests/unit_tests/configs/test_env_consistency.py +++ b/api/tests/unit_tests/configs/test_env_consistency.py @@ -4,6 +4,7 @@ from dotenv import dotenv_values BASE_API_AND_DOCKER_CONFIG_SET_DIFF: frozenset[str] = frozenset( ( + "AGENT_BACKEND_API_TOKEN", "APP_MAX_EXECUTION_TIME", "BATCH_UPLOAD_LIMIT", "CELERY_BEAT_SCHEDULER_TIME", @@ -43,6 +44,7 @@ BASE_API_AND_DOCKER_CONFIG_SET_DIFF: frozenset[str] = frozenset( BASE_API_AND_DOCKER_COMPOSE_CONFIG_SET_DIFF: frozenset[str] = frozenset( ( + "AGENT_BACKEND_API_TOKEN", "BATCH_UPLOAD_LIMIT", "CELERY_BEAT_SCHEDULER_TIME", "HTTP_REQUEST_MAX_CONNECT_TIMEOUT", diff --git a/api/tests/unit_tests/configs/test_file_upload_config.py b/api/tests/unit_tests/configs/test_file_upload_config.py new file mode 100644 index 00000000000..666ff3f666e --- /dev/null +++ b/api/tests/unit_tests/configs/test_file_upload_config.py @@ -0,0 +1,23 @@ +import pytest + +from configs.feature import FileUploadConfig + + +def test_paid_plan_file_size_limit_uses_its_default(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("UPLOAD_FILE_SIZE_LIMIT", "23") + monkeypatch.delenv("KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN", raising=False) + + config = FileUploadConfig() + + assert config.UPLOAD_FILE_SIZE_LIMIT == 23 + assert config.KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN == 15 + + +def test_paid_plan_file_size_limit_can_be_configured_separately(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("UPLOAD_FILE_SIZE_LIMIT", "23") + monkeypatch.setenv("KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN", "50") + + config = FileUploadConfig() + + assert config.UPLOAD_FILE_SIZE_LIMIT == 23 + assert config.KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN == 50 diff --git a/api/tests/unit_tests/configs/test_tidb_on_qdrant_config.py b/api/tests/unit_tests/configs/test_tidb_on_qdrant_config.py new file mode 100644 index 00000000000..5e7ff57dde3 --- /dev/null +++ b/api/tests/unit_tests/configs/test_tidb_on_qdrant_config.py @@ -0,0 +1,19 @@ +import pytest + +from configs.middleware.vdb.tidb_on_qdrant_config import TidbOnQdrantConfig + + +def test_estimated_storage_limits_default(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB", raising=False) + + config = TidbOnQdrantConfig() + + assert config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB == "sandbox:60,professional:6400,team:25600" + + +def test_estimated_storage_limits_custom(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB", "sandbox:61,professional:6500,team:26000") + + config = TidbOnQdrantConfig() + + assert config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB == "sandbox:61,professional:6500,team:26000" diff --git a/api/tests/unit_tests/conftest.py b/api/tests/unit_tests/conftest.py index 0714ef1bd89..3870d9b1d94 100644 --- a/api/tests/unit_tests/conftest.py +++ b/api/tests/unit_tests/conftest.py @@ -1,11 +1,13 @@ import os +import shutil from collections.abc import Iterator +from pathlib import Path from unittest.mock import MagicMock, patch import pytest from flask import Flask from sqlalchemy import create_engine -from sqlalchemy.engine import Engine +from sqlalchemy.engine import URL, Engine from sqlalchemy.orm import Session, sessionmaker # Getting the absolute path of the current file's directory @@ -35,13 +37,13 @@ os.environ.setdefault("OPENDAL_SCHEME", "fs") os.environ.setdefault("OPENDAL_FS_ROOT", "/tmp/dify-storage") os.environ.setdefault("STORAGE_TYPE", "opendal") -from core.db.session_factory import configure_session_factory, session_factory +import core.db.session_factory as session_factory_module from extensions import ext_redis from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole from models.base import TypeBase -def _patch_redis_clients_on_loaded_modules(): +def _patch_redis_clients_on_loaded_modules() -> None: """Ensure any module-level redis_client references point to the shared redis_mock.""" import sys @@ -49,10 +51,9 @@ def _patch_redis_clients_on_loaded_modules(): for module in list(sys.modules.values()): if module is None: continue - if hasattr(module, "redis_client"): - module.redis_client = redis_mock - if hasattr(module, "_pubsub_redis_client"): - module.pubsub_redis_client = redis_mock + for client_attribute in ("redis_client", "_pubsub_redis_client"): + if hasattr(module, client_attribute): + setattr(module, client_attribute, redis_mock) @pytest.fixture @@ -61,13 +62,13 @@ def app() -> Flask: @pytest.fixture(autouse=True) -def _provide_app_context(app: Flask): +def _provide_app_context(app: Flask) -> Iterator[None]: with app.app_context(): yield @pytest.fixture(autouse=True) -def _patch_redis_clients(): +def _patch_redis_clients() -> Iterator[None]: """Patch redis_client to MagicMock only for unit test executions.""" with ( @@ -79,7 +80,7 @@ def _patch_redis_clients(): @pytest.fixture(autouse=True) -def reset_redis_mock(): +def reset_redis_mock() -> None: """reset the Redis mock before each test""" redis_mock.reset_mock() redis_mock.get.return_value = None @@ -89,7 +90,7 @@ def reset_redis_mock(): redis_mock.exists.return_value = False redis_mock.set.return_value = None redis_mock.expire.return_value = None - redis_mock.hgetall.return_value = {} + redis_mock.hgetall.return_value = dict[bytes, bytes]() redis_mock.hdel.return_value = None redis_mock.incr.return_value = 1 @@ -98,7 +99,7 @@ def reset_redis_mock(): @pytest.fixture(autouse=True) -def reset_secret_key(): +def reset_secret_key() -> Iterator[None]: """Ensure SECRET_KEY-dependent logic sees an empty config value by default.""" from configs import dify_config @@ -111,42 +112,101 @@ def reset_secret_key(): dify_config.SECRET_KEY = original -@pytest.fixture(scope="session") -def _unit_test_engine(): - engine = create_engine("sqlite:///:memory:") - yield engine - engine.dispose() - - @pytest.fixture -def sqlite_engine() -> Iterator[Engine]: - """Create an isolated in-memory SQLite engine for tests that need a disposable database.""" +def _sqlite_engine(_sqlite_database_template: Path, tmp_path: Path) -> Iterator[Engine]: + """Create an engine over a pristine per-test copy of the SQLite schema.""" + + database_path = tmp_path / "unit-tests.sqlite3" + shutil.copyfile(_sqlite_database_template, database_path) + engine = create_engine(URL.create("sqlite", database=str(database_path))) - engine = create_engine("sqlite:///:memory:") try: yield engine finally: engine.dispose() + database_path.unlink(missing_ok=True) -@pytest.fixture -def sqlite_session(request: pytest.FixtureRequest, sqlite_engine: Engine) -> Iterator[Session]: - """Yield a SQLite session after creating the model tables passed through ``request.param``.""" +@pytest.fixture(scope="session") +def _sqlite_database_template(tmp_path_factory: pytest.TempPathFactory) -> Path: + """Create one empty full-schema SQLite database per pytest worker.""" - models: tuple[type[TypeBase], ...] = request.param - tables = [model.metadata.tables[model.__tablename__] for model in models] - TypeBase.metadata.create_all(sqlite_engine, tables=tables) - session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) - with session_factory() as session: - yield session + database_path = tmp_path_factory.mktemp("sqlite-template") / "unit-tests.sqlite3" + engine = create_engine(URL.create("sqlite", database=str(database_path))) + try: + TypeBase.metadata.create_all(engine) + finally: + engine.dispose() + return database_path @pytest.fixture(autouse=True) -def _configure_session_factory(_unit_test_engine): - try: - session_factory.get_session_maker() - except RuntimeError: - configure_session_factory(_unit_test_engine, expire_on_commit=False) +def _sqlite_session_factory( + _sqlite_engine: Engine, + monkeypatch: pytest.MonkeyPatch, +) -> sessionmaker[Session]: + """Bind all unit-test Sessions to the pristine full-schema SQLite database.""" + + factory = sessionmaker(bind=_sqlite_engine, expire_on_commit=False) + monkeypatch.setattr(session_factory_module, "_session_maker", factory) + return factory + + +@pytest.fixture +def _unbound_session_factory( + _sqlite_session_factory: sessionmaker[Session], + monkeypatch: pytest.MonkeyPatch, +) -> sessionmaker[Session]: + """Create one unbound factory and install it as the global test factory.""" + + factory = sessionmaker() + monkeypatch.setattr(session_factory_module, "_session_maker", factory) + return factory + + +@pytest.fixture +def sqlite_engine(_sqlite_engine: Engine) -> Engine: + """Expose the pristine full-schema SQLite engine to tests.""" + + return _sqlite_engine + + +@pytest.fixture +def sqlite_session_factory(_sqlite_session_factory: sessionmaker[Session]) -> sessionmaker[Session]: + """Expose the shared SQLite session factory to tests.""" + + return _sqlite_session_factory + + +@pytest.fixture +def sqlite_session(_sqlite_session_factory: sessionmaker[Session]) -> Iterator[Session]: + """Yield a session over the pristine full-schema SQLite database. + + Legacy indirect model parameters remain accepted by pytest but are ignored. + Remove those decorators as their test files receive individual review. + """ + + with _sqlite_session_factory() as session: + yield session + + +@pytest.fixture +def unbound_session_factory(_unbound_session_factory: sessionmaker[Session]) -> sessionmaker[Session]: + """Expose an unbound factory for paths that must not require persistence.""" + + return _unbound_session_factory + + +@pytest.fixture +def unbound_session(_unbound_session_factory: sessionmaker[Session]) -> Iterator[Session]: + """Yield an unbound Session for paths that must not require persistence. + + Bind-requiring database access fails, while bind-free Session operations can + still succeed. + """ + + with _unbound_session_factory() as session: + yield session def persist_service_api_tenant_owner(session: Session, tenant: Tenant, owner: Account) -> TenantAccountJoin: diff --git a/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py b/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py index 85d79606b7f..7b9e43143bc 100644 --- a/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py +++ b/api/tests/unit_tests/controllers/common/test_agent_app_parameters.py @@ -8,7 +8,15 @@ from sqlalchemy.orm import Session from controllers.common.agent_app_parameters import get_published_agent_app_feature_dict_and_user_input_form from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict from core.app.apps.agent_app.errors import AgentAppGeneratorError, AgentAppNotPublishedError -from models.agent import Agent, AgentConfigSnapshot, AgentScope, AgentSource, AgentStatus +from models.agent import ( + Agent, + AgentConfigRevision, + AgentConfigRevisionOperation, + AgentConfigSnapshot, + AgentScope, + AgentSource, + AgentStatus, +) from models.model import AppAnnotationSetting @@ -55,22 +63,53 @@ def _persist_snapshot( tenant_id: str, agent_id: str, config_snapshot: dict[str, Any], + publish_visible: bool = True, ) -> AgentConfigSnapshot: snapshot = AgentConfigSnapshot( id=snapshot_id, tenant_id=tenant_id, agent_id=agent_id, version=1, + home_snapshot_id=_stable_uuid(f"home-snapshot:{snapshot_id}"), config_snapshot=config_snapshot, ) session.add(snapshot) + if publish_visible: + _persist_publish_revision( + session, + snapshot_id=snapshot_id, + tenant_id=tenant_id, + agent_id=agent_id, + commit=False, + ) session.commit() return snapshot +def _persist_publish_revision( + session: Session, + *, + snapshot_id: str, + tenant_id: str, + agent_id: str, + commit: bool = True, +) -> None: + session.add( + AgentConfigRevision( + tenant_id=tenant_id, + agent_id=agent_id, + current_snapshot_id=snapshot_id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + ) + ) + if commit: + session.commit() + + @pytest.mark.parametrize( "sqlite_session", - [(Agent, AgentConfigSnapshot, AppAnnotationSetting)], + [(Agent, AgentConfigSnapshot, AgentConfigRevision, AppAnnotationSetting)], indirect=True, ) def test_published_agent_app_parameters_use_soul_file_upload(sqlite_session: Session): @@ -136,7 +175,7 @@ def test_published_agent_app_parameters_use_soul_file_upload(sqlite_session: Ses assert parameters["user_input_form"] == [{"text-input": {"label": "topic", "variable": "topic", "required": True}}] -@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True) +@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot, AgentConfigRevision)], indirect=True) def test_published_agent_app_parameters_requires_bound_agent(sqlite_session: Session): tenant_id = _stable_uuid("tenant:unbound") app_model = _app_model(tenant_id=tenant_id, bound_agent_id=None) @@ -145,7 +184,7 @@ def test_published_agent_app_parameters_requires_bound_agent(sqlite_session: Ses get_published_agent_app_feature_dict_and_user_input_form(app_model, session=sqlite_session) -@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True) +@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot, AgentConfigRevision)], indirect=True) def test_published_agent_app_parameters_requires_existing_active_agent(sqlite_session: Session): requested_tenant_id = _stable_uuid("tenant:requested") agent_id = _stable_uuid("agent:cross-tenant") @@ -169,7 +208,7 @@ def test_published_agent_app_parameters_requires_existing_active_agent(sqlite_se False, ], ) -@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True) +@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot, AgentConfigRevision)], indirect=True) def test_published_agent_app_parameters_requires_published_agent( active_config_is_published: bool, sqlite_session: Session ): @@ -188,7 +227,7 @@ def test_published_agent_app_parameters_requires_published_agent( get_published_agent_app_feature_dict_and_user_input_form(app_model, session=sqlite_session) -@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True) +@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot, AgentConfigRevision)], indirect=True) def test_published_agent_app_parameters_allows_unpublished_draft_with_active_snapshot(sqlite_session: Session): tenant_id = _stable_uuid("tenant:unpublished-draft") agent_id = _stable_uuid("agent:unpublished-draft") @@ -218,7 +257,33 @@ def test_published_agent_app_parameters_allows_unpublished_draft_with_active_sna assert user_input_form == [] -@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True) +@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot, AgentConfigRevision)], indirect=True) +def test_published_agent_app_parameters_rejects_seeded_unpublished_snapshot(sqlite_session: Session): + tenant_id = _stable_uuid("tenant:never-published") + agent_id = _stable_uuid("agent:never-published") + snapshot_id = _stable_uuid("snapshot:never-published") + app_model = _app_model(tenant_id=tenant_id, bound_agent_id=agent_id) + _persist_agent( + sqlite_session, + tenant_id=tenant_id, + agent_id=agent_id, + active_config_snapshot_id=snapshot_id, + active_config_is_published=False, + ) + _persist_snapshot( + sqlite_session, + snapshot_id=snapshot_id, + tenant_id=tenant_id, + agent_id=agent_id, + config_snapshot={}, + publish_visible=False, + ) + + with pytest.raises(AgentAppNotPublishedError, match="not been published"): + get_published_agent_app_feature_dict_and_user_input_form(app_model, session=sqlite_session) + + +@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot, AgentConfigRevision)], indirect=True) def test_published_agent_app_parameters_requires_published_snapshot(sqlite_session: Session): tenant_id = _stable_uuid("tenant:missing-snapshot") agent_id = _stable_uuid("agent:missing-snapshot") @@ -230,12 +295,18 @@ def test_published_agent_app_parameters_requires_published_snapshot(sqlite_sessi active_config_snapshot_id=_stable_uuid("snapshot:missing"), active_config_is_published=True, ) + _persist_publish_revision( + sqlite_session, + snapshot_id=_stable_uuid("snapshot:missing"), + tenant_id=tenant_id, + agent_id=agent_id, + ) with pytest.raises(AgentAppGeneratorError, match="published version not found"): get_published_agent_app_feature_dict_and_user_input_form(app_model, session=sqlite_session) -@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot)], indirect=True) +@pytest.mark.parametrize("sqlite_session", [(Agent, AgentConfigSnapshot, AgentConfigRevision)], indirect=True) def test_published_agent_app_parameters_allows_missing_legacy_app_model_config(sqlite_session: Session): tenant_id = _stable_uuid("tenant:no-legacy-config") agent_id = _stable_uuid("agent:no-legacy-config") diff --git a/api/tests/unit_tests/controllers/common/test_app_access.py b/api/tests/unit_tests/controllers/common/test_app_access.py index df5debb24d6..2a3212b5940 100644 --- a/api/tests/unit_tests/controllers/common/test_app_access.py +++ b/api/tests/unit_tests/controllers/common/test_app_access.py @@ -2,9 +2,8 @@ from __future__ import annotations -from unittest.mock import MagicMock - import pytest +from sqlalchemy.orm import Session from controllers.common.app_access import ( APP_LIST_PERMISSION_KEYS, @@ -106,28 +105,32 @@ class TestResolveAppAccessFilter: lambda tenant_id, account_id: whitelist, ) - def test_default_preview_is_unrestricted(self, monkeypatch: pytest.MonkeyPatch): + def test_default_preview_is_unrestricted(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=True)) permissions = _permissions(app_default_keys=["app.preview"]) - flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions) assert flt.accessible_app_ids is None assert flt.can_manage_own_apps is False - def test_default_preview_overrides_whitelist_restriction(self, monkeypatch: pytest.MonkeyPatch): + def test_default_preview_overrides_whitelist_restriction( + self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=["app-9"])) permissions = _permissions( workspace_keys=["app.full_access", "app.create_and_management"], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions) # Workspace-level preview grant defeats the whitelist restriction. assert flt.accessible_app_ids is None assert flt.can_manage_own_apps is True - def test_override_apps_collected_without_default_preview(self, monkeypatch: pytest.MonkeyPatch): + def test_override_apps_collected_without_default_preview( + self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=True)) permissions = _permissions( app_overrides=[ @@ -136,23 +139,23 @@ class TestResolveAppAccessFilter: ], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions) assert flt.accessible_app_ids == {"app-1"} - def test_whitelist_union_with_override_apps(self, monkeypatch: pytest.MonkeyPatch): + def test_whitelist_union_with_override_apps(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=["app-5"])) permissions = _permissions( app_overrides=[ResourcePermissionKeys(resource_id="app-1", permission_keys=["app.acl.preview"])], ) - flt = resolve_app_access_filter("tenant-1", "acc-1", session=MagicMock(), permissions=permissions) + flt = resolve_app_access_filter("tenant-1", "acc-1", session=unbound_session, permissions=permissions) assert flt.accessible_app_ids == {"app-1", "app-5"} - def test_fetches_permissions_when_not_supplied(self, monkeypatch: pytest.MonkeyPatch): + def test_fetches_permissions_when_not_supplied(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): self._patch_whitelist(monkeypatch, ResourceWhitelistResources(unrestricted=False, resource_ids=[])) - session = MagicMock() + session = unbound_session captured: dict[str, object] = {} def get_permissions(tenant_id: str, account_id: str, *, session: object): diff --git a/api/tests/unit_tests/controllers/common/test_session.py b/api/tests/unit_tests/controllers/common/test_session.py index 9da06133a6c..888738bdfd6 100644 --- a/api/tests/unit_tests/controllers/common/test_session.py +++ b/api/tests/unit_tests/controllers/common/test_session.py @@ -1,152 +1,140 @@ from __future__ import annotations +from contextlib import contextmanager +from unittest.mock import patch + import pytest -from sqlalchemy import Engine, literal, select -from sqlalchemy.orm import Session +from sqlalchemy import event, literal, select +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, sessionmaker from controllers.common import session as session_module +from models import Tenant -class FakeSession: - committed: bool - rolled_back: bool - closed: bool - - def __init__(self) -> None: - self.committed = False - self.rolled_back = False - self.closed = False - - def commit(self) -> None: - self.committed = True - - def rollback(self) -> None: - self.rolled_back = True +@contextmanager +def _bind_session_factory(session: Session): + database_session_factory = sessionmaker( + bind=session.get_bind(), + expire_on_commit=False, + ) + with patch("core.db.session_factory._session_maker", database_session_factory): + yield -class FakeSessionContext: - session: FakeSession - entered: bool - exited: bool - exc_type: object | None - - def __init__(self, session: FakeSession) -> None: - self.session = session - self.entered = False - self.exited = False - self.exc_type = None - - def __enter__(self) -> FakeSession: - self.entered = True - return self.session - - def __exit__(self, exc_type: object | None, *_args: object) -> None: - self.exited = True - self.exc_type = exc_type - self.session.closed = True +def _tenant_names(session: Session) -> list[str]: + session.expire_all() + return list(session.scalars(select(Tenant.name).order_by(Tenant.name)).all()) -def test_with_session_write_commits_on_success(monkeypatch: pytest.MonkeyPatch) -> None: - session = FakeSession() - session_context = FakeSessionContext(session) - monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) +@pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) +def test_with_session_write_commits_on_success(sqlite_session: Session) -> None: + commit_observed = False + injected_session: Session | None = None class Handler: @session_module.with_session(write=True) - def post(self, injected_session): - assert injected_session is session + def post(self, session: Session): + nonlocal commit_observed, injected_session + injected_session = session + + def observe_commit(_session: Session) -> None: + nonlocal commit_observed + commit_observed = True + + event.listen(session, "after_commit", observe_commit) + session.add(Tenant(name="committed tenant")) return "ok" - assert Handler().post() == "ok" + with _bind_session_factory(sqlite_session): + assert Handler().post() == "ok" - assert session.closed - assert session.committed - assert not session.rolled_back - assert session_context.entered - assert session_context.exited - assert session_context.exc_type is None + assert commit_observed + assert injected_session is not None + assert not injected_session.in_transaction() + assert _tenant_names(sqlite_session) == ["committed tenant"] -def test_with_session_default_write_commits_on_success(monkeypatch: pytest.MonkeyPatch) -> None: - session = FakeSession() - session_context = FakeSessionContext(session) - monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) - - class Handler: - @session_module.with_session - def post(self, injected_session): - assert injected_session is session - return "ok" - - assert Handler().post() == "ok" - assert session.committed - assert not session.rolled_back - - -def test_with_session_write_rolls_back_on_error(monkeypatch: pytest.MonkeyPatch) -> None: - session = FakeSession() - session_context = FakeSessionContext(session) - monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) - - class Handler: - @session_module.with_session(write=True) - def get(self, _session): - raise RuntimeError("boom") - - with pytest.raises(RuntimeError, match="boom"): - Handler().get() - - assert session.closed - assert not session.committed - assert session.rolled_back - assert session_context.entered - assert session_context.exited - assert session_context.exc_type is RuntimeError - - -def test_with_session_write_allows_commit_then_more_database_work( - monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine -) -> None: - monkeypatch.setattr(session_module.session_factory, "create_session", lambda: Session(sqlite_engine)) - +@pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) +def test_with_session_default_write_commits_on_success(sqlite_session: Session) -> None: class Handler: @session_module.with_session def post(self, session: Session): - session.commit() - return session.scalar(select(literal(1))) + session.add(Tenant(name="default write tenant")) + return "ok" - assert Handler().post() == 1 + with _bind_session_factory(sqlite_session): + assert Handler().post() == "ok" + + assert _tenant_names(sqlite_session) == ["default write tenant"] -def test_with_session_read_mode_does_not_commit(monkeypatch: pytest.MonkeyPatch) -> None: - session = FakeSession() - session_context = FakeSessionContext(session) - monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) +@pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) +def test_with_session_write_rolls_back_on_error(sqlite_session: Session) -> None: + rollback_observed = False + injected_session: Session | None = None + + class Handler: + @session_module.with_session(write=True) + def get(self, session: Session): + nonlocal rollback_observed, injected_session + injected_session = session + + def observe_rollback(_session: Session) -> None: + nonlocal rollback_observed + rollback_observed = True + + event.listen(session, "after_rollback", observe_rollback) + session.add(Tenant(name="rolled back tenant")) + session.flush() + raise RuntimeError("boom") + + with _bind_session_factory(sqlite_session), pytest.raises(RuntimeError, match="boom"): + Handler().get() + + assert rollback_observed + assert injected_session is not None + assert not injected_session.in_transaction() + assert _tenant_names(sqlite_session) == [] + + +def test_with_session_write_allows_commit_then_more_database_work(sqlite_engine: Engine) -> None: + with Session(sqlite_engine) as sqlite_session, _bind_session_factory(sqlite_session): + + class Handler: + @session_module.with_session + def post(self, session: Session): + session.commit() + return session.scalar(select(literal(1))) + + assert Handler().post() == 1 + + +@pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) +def test_with_session_read_mode_does_not_commit(sqlite_session: Session) -> None: + injected_session: Session | None = None class Handler: @session_module.with_session(write=False) - def get(self, injected_session): - assert injected_session is session + def get(self, session: Session): + nonlocal injected_session + injected_session = session + session.add(Tenant(name="uncommitted read tenant")) + session.flush() return "ok" - assert Handler().get() == "ok" + with _bind_session_factory(sqlite_session): + assert Handler().get() == "ok" - assert session.closed - assert not session.committed - assert not session.rolled_back - assert session_context.entered - assert session_context.exited - assert session_context.exc_type is None + assert injected_session is not None + assert not injected_session.in_transaction() + assert _tenant_names(sqlite_session) == [] -def test_with_session_preserves_wrapped_metadata(monkeypatch: pytest.MonkeyPatch) -> None: - session = FakeSession() - session_context = FakeSessionContext(session) - monkeypatch.setattr(session_module.session_factory, "create_session", lambda: session_context) - +def test_with_session_preserves_wrapped_metadata() -> None: class Handler: @session_module.with_session - def get(self, _session): + def get(self, _session: Session): """handler docs""" return "ok" diff --git a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py index 14628b2614d..48135e8d35d 100644 --- a/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py +++ b/api/tests/unit_tests/controllers/console/agent/test_agent_controllers.py @@ -1,3 +1,4 @@ +from datetime import datetime from inspect import getsource, unwrap from types import SimpleNamespace from typing import Any, cast @@ -5,6 +6,7 @@ from unittest.mock import MagicMock, Mock, call import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError, NotFound from controllers.console import console_ns @@ -26,21 +28,29 @@ from controllers.console.agent.roster import ( AgentApiKeyApi, AgentApiKeyListApi, AgentApiStatusApi, + AgentApiStatusPayload, AgentAppApi, AgentAppCopyApi, + AgentAppCopyPayload, + AgentAppCreatePayload, AgentAppListApi, + AgentAppUpdatePayload, AgentBuildDraftApi, AgentBuildDraftApplyApi, AgentBuildDraftCheckoutApi, + AgentBuildDraftCheckoutPayload, AgentDebugConversationRefreshApi, AgentInviteOptionsApi, + AgentInviteOptionsQuery, AgentLogMessagesApi, AgentLogsApi, AgentLogSourcesApi, AgentPublishApi, + AgentPublishPayload, AgentRosterVersionDetailApi, AgentRosterVersionRestoreApi, AgentRosterVersionsApi, + AgentStatisticsQuery, AgentStatisticsSummaryApi, ) from controllers.console.app import completion as completion_controller @@ -53,8 +63,77 @@ from controllers.console.app.message import ( AgentMessageFeedbackApi, AgentMessageSuggestedQuestionApi, ) -from models.agent import AgentConfigDraftType -from services.entities.agent_entities import ComposerSaveStrategy, ComposerVariant +from core.app.entities.app_invoke_entities import InvokeFrom +from models.agent import Agent, AgentConfigDraftType, AgentScope, AgentSource, AgentStatus +from models.enums import ApiTokenType, ConversationFromSource +from models.model import ApiToken, App, AppMode, Conversation, Message +from services.entities.agent_entities import ( + ComposerSavePayload, + ComposerSaveStrategy, + ComposerVariant, + WorkflowAgentComposerQuery, + WorkflowComposerCopyFromRosterPayload, +) + + +def _persist_conversation_message( + session: Session, + *, + app_id: str, + conversation_id: str, + message_id: str, + created_at: datetime, +) -> tuple[Conversation, Message]: + conversation = session.get(Conversation, conversation_id) + if conversation is None: + conversation = Conversation( + app_id=app_id, + app_model_config_id=None, + model_provider=None, + override_model_configs=None, + model_id=None, + mode=AppMode.CHAT, + name="Conversation", + inputs={}, + introduction="", + system_instruction="", + system_instruction_tokens=0, + status="normal", + invoke_from=InvokeFrom.DEBUGGER, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id=None, + from_account_id="00000000-0000-0000-0000-000000000021", + ) + conversation.id = conversation_id + session.add(conversation) + session.flush() + message = Message( + app_id=app_id, + conversation_id=conversation.id, + inputs={}, + query="query", + 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, + provider_response_latency=0, + total_price=0, + currency="USD", + invoke_from=InvokeFrom.DEBUGGER, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id=None, + from_account_id="00000000-0000-0000-0000-000000000021", + app_mode=AppMode.CHAT, + created_at=created_at, + ) + message.id = message_id + session.add(message) + session.flush() + return conversation, message def _version_response(version_id: str = "version-1") -> dict: @@ -115,6 +194,7 @@ def _agent_app_composer_response() -> dict: "active_config_snapshot_id": "version-1", }, "active_config_snapshot": _version_response(), + "active_config_is_published": True, "agent_soul": {}, "save_options": ["save_to_current_version"], } @@ -278,12 +358,7 @@ def test_agent_app_list_and_create_use_agent_route( ) monkeypatch.setattr( roster_controller.AgentRosterService, - "get_or_create_agent_app_debug_conversation_id", - lambda _self, **kwargs: "debug-conversation-detail", - ) - monkeypatch.setattr( - roster_controller.AgentRosterService, - "get_or_create_agent_app_debug_conversation_id", + "get_or_create_build_conversation", lambda _self, **kwargs: "debug-conversation-detail", ) monkeypatch.setattr( @@ -313,7 +388,7 @@ def test_agent_app_list_and_create_use_agent_route( ) monkeypatch.setattr( roster_controller.AgentRosterService, - "load_or_create_agent_app_debug_conversation_ids_by_agent_id", + "load_or_create_build_conversation_ids_by_agent_id", lambda _self, **kwargs: {"agent-list": "debug-conversation-list"}, ) monkeypatch.setattr( @@ -326,7 +401,7 @@ def test_agent_app_list_and_create_use_agent_route( monkeypatch.setattr( roster_controller.AgentRosterService, - "get_or_create_agent_app_debug_conversation_id", + "get_or_create_build_conversation", get_or_create_debug_conversation, ) monkeypatch.setattr( @@ -369,14 +444,20 @@ def test_agent_app_list_and_create_use_agent_route( json={"name": "Iris", "description": "Agent app", "role": "Coordinator", "icon_type": "emoji", "icon": "robot"}, ): created, status = unwrap(AgentAppListApi.post)( - AgentAppListApi(), MagicMock(), "tenant-1", SimpleNamespace(id=account_id) + AgentAppListApi(), + AgentAppCreatePayload( + name="Iris", description="Agent app", role="Coordinator", icon_type="emoji", icon="robot" + ), + MagicMock(), + "tenant-1", + SimpleNamespace(id=account_id), ) assert status == 201 assert created["id"] == "agent-created" assert created["app_id"] == "app-created" assert created["debug_conversation_id"] == "debug-conversation-created" assert created["role"] == "Created role" - assert created["active_config_is_published"] is False + assert "active_config_is_published" not in created assert "bound_agent_id" not in created create_call = cast(dict[str, object], captured["create"]) create_params = cast(Any, create_call["params"]) @@ -386,7 +467,6 @@ def test_agent_app_list_and_create_use_agent_route( "tenant_id": "tenant-1", "agent_id": "agent-created", "account_id": account_id, - "draft_type": AgentConfigDraftType.DEBUG_BUILD, "commit": False, } @@ -424,7 +504,13 @@ def test_agent_app_create_omits_optional_role_as_empty_string( "/console/api/agent", json={"name": "No-role Iris", "description": "Agent app", "icon_type": "emoji", "icon": "robot"}, ): - created, status = unwrap(AgentAppListApi.post)(AgentAppListApi(), MagicMock(), "tenant-1", current_user) + created, status = unwrap(AgentAppListApi.post)( + AgentAppListApi(), + AgentAppCreatePayload(name="No-role Iris", description="Agent app", icon_type="emoji", icon="robot"), + MagicMock(), + "tenant-1", + current_user, + ) assert status == 201 assert created == {"id": "agent-created", "app_id": "app-created"} create_call = cast(dict[str, object], captured["create"]) @@ -435,25 +521,32 @@ def test_agent_app_create_omits_optional_role_as_empty_string( def test_agent_app_detail_update_delete_resolve_app_from_agent_id( - app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str + app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str, sqlite_session: Session ) -> None: agent_id = "00000000-0000-0000-0000-000000000001" - app_model = _app_detail_obj(id="app-1", bound_agent_id=agent_id) - agent = SimpleNamespace( - id=agent_id, - app_id="app-1", - backing_app_id=None, + tenant_id = "00000000-0000-0000-0000-000000000002" + app_id = "00000000-0000-0000-0000-000000000003" + app_model = _app_detail_obj(id=app_id, tenant_id=tenant_id, bound_agent_id=agent_id) + agent = Agent( + tenant_id=tenant_id, + name="Resolved agent", + description="", role="Resolved role", - debug_conversation_id="debug-conversation-detail", - active_config_snapshot_id=None, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + app_id=app_id, + status=AgentStatus.ACTIVE, ) + agent.id = agent_id + sqlite_session.add(agent) + sqlite_session.flush() captured: dict[str, object] = {} monkeypatch.setattr(roster_controller.AgentRosterService, "get_agent_app_model", lambda _self, **kwargs: app_model) monkeypatch.setattr(roster_controller, "_resolve_agent_runtime_app_model", lambda _session, **kwargs: app_model) monkeypatch.setattr(roster_controller.AgentRosterService, "get_app_backing_agent", lambda _self, **kwargs: agent) monkeypatch.setattr( roster_controller.AgentRosterService, - "get_or_create_agent_app_debug_conversation_id", + "get_or_create_build_conversation", lambda _self, **kwargs: "debug-conversation-detail", ) monkeypatch.setattr( @@ -464,6 +557,11 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( "get_system_features", lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ) + monkeypatch.setattr( + roster_controller, + "agent_has_workflow_callable_active_snapshot", + lambda **_kwargs: False, + ) class FakeAppService: def get_app(self, app_obj: object, *, session: object) -> object: @@ -472,42 +570,49 @@ def test_agent_app_detail_update_delete_resolve_app_from_agent_id( def update_app(self, app_obj: object, args: dict[str, object], *, session: object) -> object: captured["update"] = {"app": app_obj, "args": args} - return _app_detail_obj(id="app-1", name=args["name"], bound_agent_id=agent_id) + return _app_detail_obj(id=app_id, tenant_id=tenant_id, name=args["name"], bound_agent_id=agent_id) def delete_app(self, app_obj: object, *, session: object) -> None: captured["delete"] = app_obj monkeypatch.setattr(roster_controller, "AppService", FakeAppService) - session = Mock() - session.scalar.return_value = agent - detail = unwrap(AgentAppApi.get)(AgentAppApi(), session, "tenant-1", SimpleNamespace(id=account_id), agent_id) + session = sqlite_session + detail = unwrap(AgentAppApi.get)(AgentAppApi(), session, tenant_id, SimpleNamespace(id=account_id), agent_id) assert detail["id"] == agent_id - assert detail["app_id"] == "app-1" + assert detail["app_id"] == app_id assert detail["debug_conversation_id"] == "debug-conversation-detail" assert detail["debug_conversation_has_messages"] is True assert detail["debug_conversation_message_count"] == 2 assert detail["role"] == "Resolved role" - assert detail["active_config_is_published"] is False + assert detail["access_ready"] is False + assert "active_config_is_published" not in detail assert "bound_agent_id" not in detail assert captured["get_app"] == {"app": app_model, "session": session} with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001", json={"name": "Renamed", "description": "", "role": "Reviewer", "icon_type": "emoji", "icon": "R"}, ): - updated = unwrap(AgentAppApi.put)(AgentAppApi(), session, "tenant-1", SimpleNamespace(id=account_id), agent_id) + updated = unwrap(AgentAppApi.put)( + AgentAppApi(), + AgentAppUpdatePayload(name="Renamed", description="", role="Reviewer", icon_type="emoji", icon="R"), + session, + tenant_id, + SimpleNamespace(id=account_id), + agent_id, + ) assert updated["name"] == "Renamed" assert updated["id"] == agent_id - assert updated["app_id"] == "app-1" + assert updated["app_id"] == app_id assert updated["debug_conversation_id"] == "debug-conversation-detail" assert updated["debug_conversation_has_messages"] is True assert updated["debug_conversation_message_count"] == 2 assert updated["role"] == "Resolved role" - assert updated["active_config_is_published"] is False + assert "active_config_is_published" not in updated assert "bound_agent_id" not in updated update_call = cast(dict[str, object], captured["update"]) assert update_call["app"] is app_model assert cast(dict[str, object], update_call["args"])["role"] == "Reviewer" - deleted, status = unwrap(AgentAppApi.delete)(AgentAppApi(), session, "tenant-1", agent_id) + deleted, status = unwrap(AgentAppApi.delete)(AgentAppApi(), session, tenant_id, agent_id) assert (deleted, status) == ("", 204) assert captured["delete"] is app_model @@ -543,7 +648,19 @@ def test_agent_app_copy_uses_agent_id_and_returns_agent_detail( }, ): copied, status = unwrap(AgentAppCopyApi.post)( - AgentAppCopyApi(), MagicMock(), "tenant-1", current_user, agent_id + AgentAppCopyApi(), + AgentAppCopyPayload( + name="Iris copy", + description="Copied", + role="Copied role", + icon_type="emoji", + icon="sparkles", + icon_background="#fff", + ), + MagicMock(), + "tenant-1", + current_user, + agent_id, ) assert status == 201 assert copied == {"id": "copied-agent", "app_id": "copied-app", "name": "Iris"} @@ -560,25 +677,16 @@ def test_agent_app_copy_uses_agent_id_and_returns_agent_detail( } -@pytest.mark.parametrize( - ("payload", "expected_draft_type"), - [ - (None, AgentConfigDraftType.DEBUG_BUILD), - ({"draft_type": "draft"}, AgentConfigDraftType.DRAFT), - ], -) -def test_agent_debug_conversation_refresh_uses_current_user_and_draft_type( +def test_agent_debug_conversation_refresh_resets_build_for_current_user( 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] = {} class FakeRosterService: - def refresh_agent_app_debug_conversation_id(self, **kwargs: object) -> str: + def reset_build_conversation(self, **kwargs: object) -> str: captured.update(kwargs) return "new-debug-conversation-id" @@ -586,7 +694,6 @@ def test_agent_debug_conversation_refresh_uses_current_user_and_draft_type( with app.test_request_context( "/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 @@ -600,7 +707,6 @@ def test_agent_debug_conversation_refresh_uses_current_user_and_draft_type( "tenant_id": "tenant-1", "agent_id": agent_id, "account_id": account_id, - "draft_type": expected_draft_type, } @@ -654,7 +760,14 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/publish", json={"version_note": "publish v1"} ): - published = unwrap(AgentPublishApi.post)(AgentPublishApi(), MagicMock(), "tenant-1", current_user, agent_id) + published = unwrap(AgentPublishApi.post)( + AgentPublishApi(), + AgentPublishPayload(version_note="publish v1"), + MagicMock(), + "tenant-1", + current_user, + agent_id, + ) assert published["active_config_snapshot_id"] == "version-1" captured["publish"].pop("session", None) assert captured["publish"] == { @@ -667,7 +780,12 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft/checkout", json={"force": True} ): checked_out = unwrap(AgentBuildDraftCheckoutApi.post)( - AgentBuildDraftCheckoutApi(), MagicMock(), "tenant-1", current_user, agent_id + AgentBuildDraftCheckoutApi(), + AgentBuildDraftCheckoutPayload(force=True), + MagicMock(), + "tenant-1", + current_user, + agent_id, ) assert checked_out["draft"]["id"] == "build-draft-1" captured["checkout"].pop("session", None) @@ -686,7 +804,17 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( "/console/api/agent/00000000-0000-0000-0000-000000000001/build-draft", json={"variant": "agent_app", "save_strategy": "save_to_current_version", "agent_soul": {}}, ): - saved = unwrap(AgentBuildDraftApi.put)(AgentBuildDraftApi(), MagicMock(), "tenant-1", current_user, agent_id) + saved = unwrap(AgentBuildDraftApi.put)( + AgentBuildDraftApi(), + ComposerSavePayload( + variant=ComposerVariant.AGENT_APP, + save_strategy=ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION, + ), + MagicMock(), + "tenant-1", + current_user, + agent_id, + ) assert saved["draft"]["id"] == "build-draft-1" assert captured["save"]["tenant_id"] == "tenant-1" assert captured["save"]["agent_id"] == agent_id @@ -715,12 +843,19 @@ def test_agent_publish_and_build_draft_routes_call_composer_service( def test_agent_api_access_uses_agent_id_and_returns_service_api_metadata(monkeypatch: pytest.MonkeyPatch) -> None: agent_id = "00000000-0000-0000-0000-000000000001" app_model = SimpleNamespace( - id="app-1", enable_api=True, api_base_url="https://api.example.test/v1", api_rpm=60, api_rph=600 + id="app-1", + tenant_id="tenant-1", + enable_api=True, + api_base_url="https://api.example.test/v1", + api_rpm=60, + api_rph=600, ) monkeypatch.setattr(roster_controller, "_resolve_agent_app_model", lambda _session, **kwargs: app_model) - monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda _session, app_id: 2) + monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda _session, _app: 2) + monkeypatch.setattr(roster_controller, "_agent_app_access_ready", lambda _session, _app: True) response = unwrap(AgentApiAccessApi.get)(AgentApiAccessApi(), MagicMock(), "tenant-1", agent_id) assert response == { + "access_ready": True, "enabled": True, "service_api_base_url": "https://api.example.test/v1", "streaming_only": True, @@ -738,17 +873,38 @@ def test_agent_api_access_uses_agent_id_and_returns_service_api_metadata(monkeyp } -def test_agent_api_status_and_key_routes_resolve_backing_app(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_agent_api_key_count_scopes_tenant_and_keeps_legacy_tokens(sqlite_session: Session) -> None: + app_model = cast(App, _app_detail_obj()) + sqlite_session.add_all( + [ + ApiToken(type=ApiTokenType.APP, token="owned", app_id=app_model.id, tenant_id=app_model.tenant_id), + ApiToken(type=ApiTokenType.APP, token="legacy", app_id=app_model.id, tenant_id=None), + ApiToken(type=ApiTokenType.APP, token="foreign", app_id=app_model.id, tenant_id="tenant-2"), + ] + ) + sqlite_session.commit() + + assert roster_controller._agent_api_key_count(sqlite_session, app_model) == 2 + + +def test_agent_api_status_and_key_routes_resolve_backing_app( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: agent_id = "00000000-0000-0000-0000-000000000001" api_key_id = "00000000-0000-0000-0000-000000000002" app_model = SimpleNamespace( - id="app-1", enable_api=False, api_base_url="https://api.example.test/v1", api_rpm=0, api_rph=0 + id="app-1", + tenant_id="tenant-1", + enable_api=False, + api_base_url="https://api.example.test/v1", + api_rpm=0, + api_rph=0, ) captured: dict[str, object] = {} - session = MagicMock() resolve_app = Mock(return_value=app_model) monkeypatch.setattr(roster_controller, "_resolve_agent_app_model", resolve_app) - monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda _session, app_id: 1) + monkeypatch.setattr(roster_controller, "_agent_api_key_count", lambda _session, _app: 1) + monkeypatch.setattr(roster_controller, "_agent_app_access_ready", lambda _session, _app: True) class FakeAppService: def update_app_api_status(self, app_obj: object, enable_api: bool, *, session: object) -> object: @@ -789,34 +945,44 @@ def test_agent_api_status_and_key_routes_resolve_backing_app(app: Flask, monkeyp with app.test_request_context( "/console/api/agent/00000000-0000-0000-0000-000000000001/api-enable", json={"enable_api": True} ): - enabled = unwrap(AgentApiStatusApi.post)(AgentApiStatusApi(), session, "tenant-1", agent_id) + enabled = unwrap(AgentApiStatusApi.post)( + AgentApiStatusApi(), AgentApiStatusPayload(enable_api=True), unbound_session, "tenant-1", agent_id + ) assert enabled["enabled"] is True assert captured["enable"] == {"app": app_model, "enable_api": True} - keys = unwrap(AgentApiKeyListApi.get)(AgentApiKeyListApi(), session, "tenant-1", agent_id) + keys = unwrap(AgentApiKeyListApi.get)(AgentApiKeyListApi(), unbound_session, "tenant-1", agent_id) assert keys == {"data": []} - assert captured["list_keys"] == {"session": session, "resource_id": "app-1", "tenant_id": "tenant-1"} - created, status = unwrap(AgentApiKeyListApi.post)(AgentApiKeyListApi(), session, "tenant-1", agent_id) + assert captured["list_keys"] == { + "session": unbound_session, + "resource_id": "app-1", + "tenant_id": "tenant-1", + } + created, status = unwrap(AgentApiKeyListApi.post)(AgentApiKeyListApi(), unbound_session, "tenant-1", agent_id) assert status == 201 assert created["id"] == api_key_id assert created["token"] == "app-test-token" - assert captured["create_key"] == {"session": session, "resource_id": "app-1", "tenant_id": "tenant-1"} + assert captured["create_key"] == { + "session": unbound_session, + "resource_id": "app-1", + "tenant_id": "tenant-1", + } current_user = SimpleNamespace(id="account-1", is_admin_or_owner=True) deleted, delete_status = unwrap(AgentApiKeyApi.delete)( - AgentApiKeyApi(), session, "tenant-1", current_user, agent_id, api_key_id + AgentApiKeyApi(), unbound_session, "tenant-1", current_user, agent_id, api_key_id ) assert (deleted, delete_status) == ("", 204) assert captured["delete_key"] == { - "session": session, + "session": unbound_session, "resource_id": "app-1", "api_key_id": api_key_id, "tenant_id": "tenant-1", "current_user": current_user, } assert resolve_app.call_args_list == [ - call(session, tenant_id="tenant-1", agent_id=agent_id), - call(session, tenant_id="tenant-1", agent_id=agent_id), - call(session, tenant_id="tenant-1", agent_id=agent_id), - call(session, tenant_id="tenant-1", agent_id=agent_id), + call(unbound_session, tenant_id="tenant-1", agent_id=agent_id), + call(unbound_session, tenant_id="tenant-1", agent_id=agent_id), + call(unbound_session, tenant_id="tenant-1", agent_id=agent_id), + call(unbound_session, tenant_id="tenant-1", agent_id=agent_id), ] @@ -839,15 +1005,12 @@ def test_agent_app_update_allows_empty_role(app: Flask, monkeypatch: pytest.Monk ) monkeypatch.setattr( roster_controller.AgentRosterService, - "get_or_create_agent_app_debug_conversation_id", + "get_or_create_build_conversation", lambda _self, **kwargs: "debug-conversation-detail", ) monkeypatch.setattr( roster_controller.AgentRosterService, "count_agent_app_debug_conversation_messages", lambda _self, **kwargs: 0 ) - monkeypatch.setattr( - roster_controller.AgentRosterService, "active_config_is_published", lambda _self, **kwargs: False - ) monkeypatch.setattr( roster_controller.FeatureService, "get_system_features", @@ -868,7 +1031,12 @@ def test_agent_app_update_allows_empty_role(app: Flask, monkeypatch: pytest.Monk json={"name": "Renamed", "description": "", "role": "", "icon_type": "emoji", "icon": "R"}, ): updated = unwrap(AgentAppApi.put)( - AgentAppApi(), MagicMock(), "tenant-1", SimpleNamespace(id="account-1"), agent_id + AgentAppApi(), + AgentAppUpdatePayload(name="Renamed", description="", role="", icon_type="emoji", icon="R"), + MagicMock(), + "tenant-1", + SimpleNamespace(id="account-1"), + agent_id, ) assert updated["role"] == "" update_call = cast(dict[str, object], captured["update"]) @@ -884,7 +1052,9 @@ def test_invite_options_get_parses_app_id(app: Flask, monkeypatch: pytest.Monkey monkeypatch.setattr(roster_controller.AgentRosterService, "list_invite_options", list_invite_options) with app.test_request_context("/console/api/agent/invite-options?page=1&limit=10&app_id=app-1"): - result = unwrap(AgentInviteOptionsApi.get)(AgentInviteOptionsApi(), MagicMock(), "tenant-1") + result = unwrap(AgentInviteOptionsApi.get)( + AgentInviteOptionsApi(), AgentInviteOptionsQuery(page=1, limit=10, app_id="app-1"), MagicMock(), "tenant-1" + ) assert result == {"data": [], "page": 1, "limit": 10, "total": 0, "has_more": False} assert captured == {"tenant_id": "tenant-1", "page": 1, "limit": 10, "keyword": None, "app_id": "app-1"} @@ -1121,7 +1291,7 @@ def test_agent_observability_routes_resolve_app_from_agent_id( "/console/api/agent/00000000-0000-0000-0000-000000000001/statistics/summary?source=api" ): statistics = unwrap(AgentStatisticsSummaryApi.get)( - AgentStatisticsSummaryApi(), MagicMock(), "tenant-1", account, agent_id + AgentStatisticsSummaryApi(), AgentStatisticsQuery(source="api"), MagicMock(), "tenant-1", account, agent_id ) assert statistics["summary"]["total_messages"] == 1 stats_call = cast(dict[str, object], captured["statistics"]) @@ -1173,18 +1343,40 @@ def test_workflow_composer_get_put_validate_candidates_impact_and_save( ) with app.test_request_context("?snapshot_id=preview-version"): workflow_state = unwrap(WorkflowAgentComposerApi.get)( - WorkflowAgentComposerApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerApi(), + WorkflowAgentComposerQuery(snapshot_id="preview-version"), + MagicMock(), + "tenant-1", + account_id, + app_model, + "node-1", ) assert workflow_state["node_id"] == "node-1" assert captured_load["account_id"] == account_id assert captured_load["snapshot_id"] == "preview-version" + composer_save_payload = ComposerSavePayload( + variant=ComposerVariant.WORKFLOW, + save_strategy=ComposerSaveStrategy.NODE_JOB_ONLY, + binding={"binding_type": "roster_agent", "current_snapshot_id": "version-1"}, + ) with app.test_request_context(json=payload): saved_state = unwrap(WorkflowAgentComposerApi.put)( - WorkflowAgentComposerApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerApi(), + composer_save_payload, + MagicMock(), + "tenant-1", + account_id, + app_model, + "node-1", ) assert saved_state["save_options"] == ["node_job_only"] assert unwrap(WorkflowAgentComposerValidateApi.post)( - WorkflowAgentComposerValidateApi(), MagicMock(), "tenant-1", app_model, "node-1" + WorkflowAgentComposerValidateApi(), + composer_save_payload, + MagicMock(), + "tenant-1", + app_model, + "node-1", ) == {"result": "success", "errors": [], "warnings": [], "knowledge_retrieval_placeholder": []} assert ( unwrap(WorkflowAgentComposerCandidatesApi.get)( @@ -1194,10 +1386,21 @@ def test_workflow_composer_get_put_validate_candidates_impact_and_save( ) with app.test_request_context(json=payload): assert unwrap(WorkflowAgentComposerImpactApi.post)( - WorkflowAgentComposerImpactApi(), MagicMock(), "tenant-1", app_model, "node-1" + WorkflowAgentComposerImpactApi(), + composer_save_payload, + MagicMock(), + "tenant-1", + app_model, + "node-1", ) == {"current_snapshot_id": "version-1", "workflow_node_count": 1, "bindings": []} assert unwrap(WorkflowAgentComposerSaveToRosterApi.post)( - WorkflowAgentComposerSaveToRosterApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerSaveToRosterApi(), + composer_save_payload, + MagicMock(), + "tenant-1", + account_id, + app_model, + "node-1", )["save_options"] == ["node_job_only"] @@ -1205,6 +1408,10 @@ def test_workflow_composer_get_uses_write_transaction() -> None: assert "@with_session\n @get_app_model" in getsource(WorkflowAgentComposerApi) +def test_build_draft_apply_leaves_transaction_ownership_to_service() -> None: + assert "@with_session(write=False)\n def post" in getsource(AgentBuildDraftApplyApi) + + def test_workflow_composer_copy_from_roster(app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str) -> None: app_model = SimpleNamespace(id="app-1") captured: dict[str, object] = {} @@ -1241,7 +1448,17 @@ def test_workflow_composer_copy_from_roster(app: Flask, monkeypatch: pytest.Monk } ): result = unwrap(WorkflowAgentComposerCopyFromRosterApi.post)( - WorkflowAgentComposerCopyFromRosterApi(), MagicMock(), "tenant-1", account_id, app_model, "node-1" + WorkflowAgentComposerCopyFromRosterApi(), + WorkflowComposerCopyFromRosterPayload( + source_agent_id="roster-agent-1", + source_snapshot_id="roster-version-1", + idempotency_key="copy-1", + ), + MagicMock(), + "tenant-1", + account_id, + app_model, + "node-1", ) assert result["binding"]["binding_type"] == "inline_agent" captured.pop("session", None) @@ -1260,7 +1477,15 @@ def test_workflow_impact_returns_empty_without_version(app: Flask) -> None: payload = {"variant": ComposerVariant.WORKFLOW.value, "save_strategy": ComposerSaveStrategy.NODE_JOB_ONLY.value} with app.test_request_context(json=payload): result = unwrap(WorkflowAgentComposerImpactApi.post)( - WorkflowAgentComposerImpactApi(), MagicMock(), "tenant-1", SimpleNamespace(id="app-1"), "node-1" + WorkflowAgentComposerImpactApi(), + ComposerSavePayload( + variant=ComposerVariant.WORKFLOW, + save_strategy=ComposerSaveStrategy.NODE_JOB_ONLY, + ), + MagicMock(), + "tenant-1", + SimpleNamespace(id="app-1"), + "node-1", ) assert result == {"current_snapshot_id": None, "workflow_node_count": 0, "bindings": []} @@ -1299,15 +1524,25 @@ def test_agent_composer_routes_resolve_app_from_agent_id( composer_controller.AgentComposerService, "collect_validation_findings", collect_validation_findings ) monkeypatch.setattr(composer_controller.AgentComposerService, "get_agent_app_candidates", get_agent_app_candidates) - assert unwrap(AgentComposerApi.get)(AgentComposerApi(), MagicMock(), "tenant-1", agent_id)["variant"] == "agent_app" + composer = unwrap(AgentComposerApi.get)(AgentComposerApi(), MagicMock(), "tenant-1", agent_id) + assert composer["variant"] == "agent_app" + assert composer["active_config_is_published"] is True assert cast(dict[str, object], captured["load"])["agent_id"] == agent_id + composer_save_payload = ComposerSavePayload( + variant=ComposerVariant.AGENT_APP, + save_strategy=ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION, + agent_soul={"prompt": {"system_prompt": "x"}}, + ) with app.test_request_context(json=payload): - assert ( - unwrap(AgentComposerApi.put)(AgentComposerApi(), MagicMock(), "tenant-1", account_id, agent_id)["variant"] - == "agent_app" + saved_composer = unwrap(AgentComposerApi.put)( + AgentComposerApi(), composer_save_payload, MagicMock(), "tenant-1", account_id, agent_id ) + assert saved_composer["variant"] == "agent_app" + assert saved_composer["active_config_is_published"] is True assert cast(dict[str, object], captured["save"])["agent_id"] == agent_id - assert unwrap(AgentComposerValidateApi.post)(AgentComposerValidateApi(), MagicMock(), "tenant-1", agent_id) == { + assert unwrap(AgentComposerValidateApi.post)( + AgentComposerValidateApi(), composer_save_payload, MagicMock(), "tenant-1", agent_id + ) == { "result": "success", "errors": [], "warnings": [], @@ -1322,7 +1557,7 @@ def test_agent_composer_routes_resolve_app_from_agent_id( def test_agent_chat_generate_and_stop_routes_resolve_app_from_agent_id( - app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str + app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str, unbound_session: Session ) -> None: agent_id = "00000000-0000-0000-0000-000000000001" app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent") @@ -1353,22 +1588,21 @@ def test_agent_chat_generate_and_stop_routes_resolve_app_from_agent_id( ) monkeypatch.setattr(completion_controller, "_create_chat_message", create_chat_message) monkeypatch.setattr(completion_controller, "_stop_chat_message", stop_chat_message) - session = Mock() with app.test_request_context(json={"inputs": {}, "query": "hello"}): assert unwrap(AgentChatMessageApi.post)( - AgentChatMessageApi(), session, "tenant-1", SimpleNamespace(id=account_id), agent_id + AgentChatMessageApi(), unbound_session, "tenant-1", SimpleNamespace(id=account_id), agent_id ) == {"result": "generated"} assert cast(dict[str, object], captured["resolve"]) == {"tenant_id": "tenant-1", "agent_id": agent_id} - assert captured["resolve_session"] is session + assert captured["resolve_session"] is unbound_session create_call = cast(dict[str, object], captured["create"]) - assert create_call["session"] is session + assert create_call["session"] is unbound_session assert create_call["app_model"] is app_model assert cast(SimpleNamespace, create_call["current_user"]).id == account_id assert unwrap(AgentChatMessageStopApi.post)( - AgentChatMessageStopApi(), session, "tenant-1", account_id, agent_id, "task-1" + AgentChatMessageStopApi(), unbound_session, "tenant-1", account_id, agent_id, "task-1" ) == ({"result": "success"}, 200) assert captured["stop_resolve"] == { - "session": session, + "session": unbound_session, "tenant_id": "tenant-1", "agent_id": agent_id, } @@ -1377,7 +1611,6 @@ def test_agent_chat_generate_and_stop_routes_resolve_app_from_agent_id( def test_agent_chat_stream_preflight_raises_first_error_event() -> None: - class ClosableStream: def __init__(self) -> None: self.closed = False @@ -1420,7 +1653,7 @@ def test_agent_chat_stream_preflight_preserves_first_normal_event() -> None: def test_agent_build_chat_finalize_route_resolves_app_from_agent_id( - app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str + app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str, unbound_session: Session ) -> None: agent_id = "00000000-0000-0000-0000-000000000001" app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent") @@ -1440,14 +1673,13 @@ def test_agent_build_chat_finalize_route_resolves_app_from_agent_id( lambda _self, **kwargs: resolve_agent_app_model(**kwargs), ) monkeypatch.setattr(completion_controller, "_create_build_chat_finalization_message", create_finalization_message) - session = Mock() with app.test_request_context(): assert unwrap(AgentBuildChatFinalizeApi.post)( - AgentBuildChatFinalizeApi(), session, "tenant-1", SimpleNamespace(id=account_id), agent_id + AgentBuildChatFinalizeApi(), unbound_session, "tenant-1", SimpleNamespace(id=account_id), agent_id ) == {"result": "generated"} assert cast(dict[str, object], captured["resolve"]) == {"tenant_id": "tenant-1", "agent_id": agent_id} finalize_call = cast(dict[str, object], captured["finalize"]) - assert finalize_call["session"] is session + assert finalize_call["session"] is unbound_session assert finalize_call["app_model"] is app_model assert finalize_call["current_tenant_id"] == "tenant-1" assert finalize_call["agent_id"] == agent_id @@ -1455,7 +1687,7 @@ def test_agent_build_chat_finalize_route_resolves_app_from_agent_id( def test_build_chat_finalization_helper_forces_debug_build_and_push_prompt( - app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str + app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str, unbound_session: Session ) -> None: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent") captured: dict[str, object] = {} @@ -1474,18 +1706,17 @@ def test_build_chat_finalization_helper_forces_debug_build_and_push_prompt( completion_controller, "_resolve_current_user_agent_debug_conversation_id", resolve_debug_conversation ) monkeypatch.setattr(completion_controller.AppGenerateService, "generate", generate) - session = Mock() with app.test_request_context(headers={"X-Trace-Id": "trace-1"}): result = completion_controller._create_build_chat_finalization_message( current_tenant_id="tenant-1", current_user=SimpleNamespace(id=account_id), app_model=app_model, agent_id="agent-1", - session=session, + session=unbound_session, ) assert result == ({"result": "success"}, 200) assert captured["resolve_debug_conversation"] == { - "session": session, + "session": unbound_session, "current_tenant_id": "tenant-1", "current_user": SimpleNamespace(id=account_id), "app_model": app_model, @@ -1501,12 +1732,11 @@ def test_build_chat_finalization_helper_forces_debug_build_and_push_prompt( assert args["conversation_id"] == "debug-conversation-1" assert args["inputs"] == {} assert args["auto_generate_name"] is False - assert args[completion_controller.AGENT_RUNTIME_EXIT_INTENT_ARG] == "delete" + assert "_agent_runtime_exit_intent" not in args assert args["external_trace_id"] == "trace-1" def test_drain_streaming_generate_response_returns_on_message_end() -> None: - class ClosableResponse: def __init__(self) -> None: self._chunks = iter( @@ -1573,6 +1803,7 @@ def test_agent_chat_helper_resolves_scoped_conversation_and_forces_streaming( payload_extra: dict[str, str | None], expected_draft_type: AgentConfigDraftType, expected_start_new: bool, + unbound_session: Session, ) -> None: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent") current_user = SimpleNamespace(id=account_id) @@ -1600,7 +1831,7 @@ def test_agent_chat_helper_resolves_scoped_conversation_and_forces_streaming( headers={"X-Trace-Id": "trace-1"}, ): result = completion_controller._create_chat_message( - current_user=current_user, app_model=app_model, session=Mock() + current_user=current_user, app_model=app_model, session=unbound_session ) assert result == {"response": {"answer": "ok"}} assert captured["app_model"] is app_model @@ -1617,7 +1848,7 @@ def test_agent_chat_helper_resolves_scoped_conversation_and_forces_streaming( def test_agent_chat_helper_ignores_private_exit_intent_payload_key( - app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str + app: Flask, monkeypatch: pytest.MonkeyPatch, account_id: str, unbound_session: Session ) -> None: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent") current_user = SimpleNamespace(id=account_id) @@ -1644,13 +1875,13 @@ def test_agent_chat_helper_ignores_private_exit_intent_payload_key( "inputs": {}, "query": "hello", "response_mode": "streaming", - completion_controller.AGENT_RUNTIME_EXIT_INTENT_ARG: "delete", + "_agent_runtime_exit_intent": "delete", } ): result = completion_controller._create_chat_message( current_user=current_user, app_model=app_model, - session=Mock(), + session=unbound_session, ) assert result == {"response": {"answer": "ok"}} @@ -1658,7 +1889,7 @@ def test_agent_chat_helper_ignores_private_exit_intent_payload_key( args = cast(dict[str, object], captured["args"]) assert args["response_mode"] == "streaming" assert args["conversation_id"] == "debug-conversation-1" - assert completion_controller.AGENT_RUNTIME_EXIT_INTENT_ARG not in args + assert "_agent_runtime_exit_intent" not in args @pytest.mark.parametrize( @@ -1674,6 +1905,7 @@ def test_agent_chat_helper_rejects_foreign_debug_conversation_before_generation( account_id: str, payload_extra: dict[str, str], expected_draft_type: AgentConfigDraftType, + unbound_session: Session, ) -> None: app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode="agent") generate = MagicMock() @@ -1699,7 +1931,7 @@ def test_agent_chat_helper_rejects_foreign_debug_conversation_before_generation( current_user=SimpleNamespace(id=account_id), app_model=app_model, agent_id="agent-1", - session=Mock(), + session=unbound_session, ) resolve_debug_conversation.assert_called_once() @@ -1718,14 +1950,18 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app def __init__(self, session: object) -> None: calls.append({"session": session}) - def get_or_create_agent_app_debug_conversation_id(self, **kwargs: object) -> str: - calls.append({"get_or_create": kwargs}) + def get_or_create_build_conversation(self, **kwargs: object) -> str: + calls.append({"get_build": kwargs}) return f"debug-{kwargs['agent_id']}" - def refresh_agent_app_debug_conversation_id(self, **kwargs: object) -> str: - calls.append({"refresh": kwargs}) + def rotate_preview_conversation(self, **kwargs: object) -> str: + calls.append({"rotate_preview": kwargs}) return f"new-{kwargs['agent_id']}" + def get_current_preview_conversation(self, **kwargs: object) -> str: + calls.append({"get_preview": kwargs}) + return f"preview-{kwargs['agent_id']}" + def get_app_backing_agent(self, **kwargs: object) -> object: calls.append({"get_app_backing_agent": kwargs}) return SimpleNamespace(id="backing-agent") @@ -1757,33 +1993,46 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app draft_type=AgentConfigDraftType.DRAFT, start_new=True, ) + current_preview_id = completion_controller._resolve_current_user_agent_debug_conversation_id( + session="session-1", # type: ignore[arg-type] + current_tenant_id="tenant-1", + current_user=SimpleNamespace(id="account-1"), + app_model=SimpleNamespace(id="app-1"), + agent_id="agent-1", + draft_type=AgentConfigDraftType.DRAFT, + ) assert explicit_id == "new-agent-1" assert fallback_id == "debug-backing-agent" assert fallback_preview_id == "new-backing-agent" + assert current_preview_id == "preview-agent-1" assert calls[1] == { - "refresh": { + "rotate_preview": { "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": { + "get_build": { "tenant_id": "tenant-1", "agent_id": "backing-agent", "account_id": "account-1", - "draft_type": AgentConfigDraftType.DEBUG_BUILD, } } assert calls[6] == {"get_app_backing_agent": {"tenant_id": "tenant-1", "app_id": "app-1"}} assert calls[7] == { - "refresh": { + "rotate_preview": { "tenant_id": "tenant-1", "agent_id": "backing-agent", "account_id": "account-1", - "draft_type": AgentConfigDraftType.DRAFT, + } + } + assert calls[9] == { + "get_preview": { + "tenant_id": "tenant-1", + "agent_id": "agent-1", + "account_id": "account-1", } } @@ -1815,25 +2064,30 @@ def test_resolve_current_user_agent_debug_conversation_uses_agent_or_backing_app ], ) def test_agent_chat_helper_maps_generation_errors( - app: Flask, monkeypatch: pytest.MonkeyPatch, error: Exception, expected: type[Exception] + app: Flask, + monkeypatch: pytest.MonkeyPatch, + error: Exception, + expected: type[Exception], + unbound_session: Session, ) -> None: app_model = SimpleNamespace(id="app-1", mode="chat") monkeypatch.setattr(completion_controller.AppGenerateService, "generate", lambda **_: (_ for _ in ()).throw(error)) with app.test_request_context(json={"inputs": {}, "query": "hello"}): with pytest.raises(expected): completion_controller._create_chat_message( - current_user=SimpleNamespace(id="account-1"), app_model=app_model, session=Mock() + current_user=SimpleNamespace(id="account-1"), app_model=app_model, session=unbound_session ) -def test_agent_chat_message_routes_resolve_app_from_agent_id(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_agent_chat_message_routes_resolve_app_from_agent_id( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: agent_id = "00000000-0000-0000-0000-000000000001" message_id = "00000000-0000-0000-0000-000000000002" app_model = SimpleNamespace(id="app-1", mode="agent") current_user = SimpleNamespace(id="account-1") captured: dict[str, object] = {} resolver_calls: list[dict[str, object]] = [] - session = Mock() def resolve_agent_app_model(**kwargs: object) -> object: resolver_calls.append(kwargs) @@ -1861,50 +2115,68 @@ def test_agent_chat_message_routes_resolve_app_from_agent_id(app: Flask, monkeyp monkeypatch.setattr(message_controller, "_get_message_suggested_questions", get_message_suggested_questions) monkeypatch.setattr(message_controller, "_get_message_detail", get_message_detail) assert unwrap(AgentChatMessageListApi.get)( - AgentChatMessageListApi(), session, "tenant-1", current_user, agent_id + AgentChatMessageListApi(), unbound_session, "tenant-1", current_user, agent_id ) == {"data": []} list_call = cast(dict[str, object], captured["list"]) - assert list_call["session"] is session + assert list_call["session"] is unbound_session assert list_call["app_model"] is app_model with app.test_request_context(json={"message_id": message_id, "rating": "like"}): assert unwrap(AgentMessageFeedbackApi.post)( - AgentMessageFeedbackApi(), session, "tenant-1", current_user, agent_id + AgentMessageFeedbackApi(), unbound_session, "tenant-1", current_user, agent_id ) == {"result": "success"} feedback_call = cast(dict[str, object], captured["feedback"]) - assert feedback_call["session"] is session + assert feedback_call["session"] is unbound_session assert feedback_call["app_model"] is app_model assert feedback_call["current_user"] is current_user assert unwrap(AgentMessageSuggestedQuestionApi.get)( - AgentMessageSuggestedQuestionApi(), session, "tenant-1", current_user, agent_id, message_id + AgentMessageSuggestedQuestionApi(), unbound_session, "tenant-1", current_user, agent_id, message_id ) == {"data": ["next"]} suggested_call = cast(dict[str, object], captured["suggested"]) - assert suggested_call["session"] is session + assert suggested_call["session"] is unbound_session assert suggested_call["app_model"] is app_model assert suggested_call["current_user"] is current_user assert suggested_call["message_id"] == message_id - assert unwrap(AgentMessageApi.get)(AgentMessageApi(), session, "tenant-1", agent_id, message_id) == { + assert unwrap(AgentMessageApi.get)(AgentMessageApi(), unbound_session, "tenant-1", agent_id, message_id) == { "id": message_id } detail_call = cast(dict[str, object], captured["detail"]) - assert detail_call == {"session": session, "app_model": app_model, "message_id": message_id} + assert detail_call == {"session": unbound_session, "app_model": app_model, "message_id": message_id} assert resolver_calls == [ - {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id}, - {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id}, - {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id}, - {"session": session, "tenant_id": "tenant-1", "agent_id": agent_id}, + {"session": unbound_session, "tenant_id": "tenant-1", "agent_id": agent_id}, + {"session": unbound_session, "tenant_id": "tenant-1", "agent_id": agent_id}, + {"session": unbound_session, "tenant_id": "tenant-1", "agent_id": agent_id}, + {"session": unbound_session, "tenant_id": "tenant-1", "agent_id": agent_id}, ] -def test_list_chat_messages_supports_first_id_pagination(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_list_chat_messages_supports_first_id_pagination( + app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + app_id = "00000000-0000-0000-0000-000000000001" conversation_id = "00000000-0000-0000-0000-000000000010" first_message_id = "00000000-0000-0000-0000-000000000011" older_message_id = "00000000-0000-0000-0000-000000000012" - conversation = SimpleNamespace(id=conversation_id) - first_message = SimpleNamespace(id=first_message_id, created_at=2) - older_message = SimpleNamespace(id=older_message_id, created_at=1) - scalar_values = iter([conversation, first_message, True]) - scalars_result = SimpleNamespace(all=lambda: [older_message]) - session = SimpleNamespace(scalar=lambda _stmt: next(scalar_values), scalars=lambda _stmt: scalars_result) + _persist_conversation_message( + sqlite_session, + app_id=app_id, + conversation_id=conversation_id, + message_id="00000000-0000-0000-0000-000000000013", + created_at=datetime(2025, 1, 1), + ) + _persist_conversation_message( + sqlite_session, + app_id=app_id, + conversation_id=conversation_id, + message_id=older_message_id, + created_at=datetime(2025, 1, 2), + ) + _persist_conversation_message( + sqlite_session, + app_id=app_id, + conversation_id=conversation_id, + message_id=first_message_id, + created_at=datetime(2025, 1, 3), + ) class FakeMessagePaginationResponse: @classmethod @@ -1923,20 +2195,27 @@ def test_list_chat_messages_supports_first_id_pagination(app: Flask, monkeypatch f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}&first_id={first_message_id}&limit=1" ): result = message_controller._list_chat_messages( - session=session, app_model=SimpleNamespace(id="app-1", mode="chat") + session=sqlite_session, app_model=SimpleNamespace(id=app_id, mode="chat") ) assert result == {"data": [older_message_id], "limit": 1, "has_more": True} -def test_list_agent_chat_messages_uses_current_user_conversation(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_list_agent_chat_messages_uses_current_user_conversation( + app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + app_id = "00000000-0000-0000-0000-000000000001" conversation_id = "00000000-0000-0000-0000-000000000010" message_id = "00000000-0000-0000-0000-000000000011" - conversation = SimpleNamespace(id=conversation_id) - message = SimpleNamespace(id=message_id, created_at=1) + conversation, _ = _persist_conversation_message( + sqlite_session, + app_id=app_id, + conversation_id=conversation_id, + message_id=message_id, + created_at=datetime(2025, 1, 1), + ) current_user = SimpleNamespace(id="account-1") - app_model = SimpleNamespace(id="app-1", mode="agent") + app_model = SimpleNamespace(id=app_id, mode="agent") captured: dict[str, object] = {} - session = SimpleNamespace(scalar=lambda _stmt: False, scalars=lambda _stmt: SimpleNamespace(all=lambda: [message])) class FakeMessagePaginationResponse: @classmethod @@ -1957,13 +2236,17 @@ def test_list_agent_chat_messages_uses_current_user_conversation(app: Flask, mon monkeypatch.setattr(message_controller, "attach_message_extra_contents", lambda messages: None) monkeypatch.setattr(message_controller, "MessageInfiniteScrollPaginationResponse", FakeMessagePaginationResponse) with app.test_request_context(f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}"): - result = message_controller._list_chat_messages(session=session, app_model=app_model, current_user=current_user) + result = message_controller._list_chat_messages( + session=sqlite_session, app_model=app_model, current_user=current_user + ) assert result == {"data": [message_id], "limit": 20, "has_more": False} - assert captured.pop("session") is session + assert captured.pop("session") is sqlite_session assert captured == {"app_model": app_model, "conversation_id": conversation_id, "user": current_user} -def test_list_agent_chat_messages_rejects_foreign_conversation(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_list_agent_chat_messages_rejects_foreign_conversation( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: conversation_id = "00000000-0000-0000-0000-000000000010" monkeypatch.setattr( message_controller.ConversationService, @@ -1973,32 +2256,33 @@ def test_list_agent_chat_messages_rejects_foreign_conversation(app: Flask, monke with app.test_request_context(f"/console/api/agent/agent-1/chat-messages?conversation_id={conversation_id}"): with pytest.raises(NotFound): message_controller._list_chat_messages( - session=Mock(), + session=unbound_session, app_model=SimpleNamespace(id="app-1", mode="agent"), current_user=SimpleNamespace(id="account-1"), ) def test_update_message_feedback_rejects_empty_rating_without_existing_feedback( - app: Flask, + app: Flask, sqlite_session: Session ) -> None: + app_id = "00000000-0000-0000-0000-000000000001" message_id = "00000000-0000-0000-0000-000000000002" - message = SimpleNamespace( - id=message_id, - app_id="app-1", - admin_feedback_with_session=MagicMock(return_value=None), + _, message = _persist_conversation_message( + sqlite_session, + app_id=app_id, + conversation_id="00000000-0000-0000-0000-000000000010", + message_id=message_id, + created_at=datetime(2025, 1, 1), ) - session = MagicMock() - session.scalar.return_value = message with app.test_request_context(json={"message_id": message_id, "rating": None}): with pytest.raises(ValueError, match="rating cannot be None"): message_controller._update_message_feedback( - session=session, + session=sqlite_session, current_user=SimpleNamespace(id="account-1"), - app_model=SimpleNamespace(id="app-1"), + app_model=SimpleNamespace(id=app_id), ) - message.admin_feedback_with_session.assert_called_once_with(session=session) + assert message.admin_feedback_with_session(session=sqlite_session) is None @pytest.mark.parametrize( @@ -2021,12 +2305,10 @@ def test_update_message_feedback_rejects_empty_rating_without_existing_feedback( ], ) def test_get_message_suggested_questions_maps_service_errors( - monkeypatch: pytest.MonkeyPatch, error: Exception, expected: type[Exception] + monkeypatch: pytest.MonkeyPatch, error: Exception, expected: type[Exception], unbound_session: Session ) -> None: - session = Mock() - def raise_error(**kwargs: object) -> None: - assert kwargs["session"] is session + assert kwargs["session"] is unbound_session raise error monkeypatch.setattr( @@ -2036,7 +2318,7 @@ def test_get_message_suggested_questions_maps_service_errors( ) with pytest.raises(expected): message_controller._get_message_suggested_questions( - session=session, + session=unbound_session, current_user=SimpleNamespace(id="account-1"), app_model=SimpleNamespace(id="app-1"), message_id="00000000-0000-0000-0000-000000000002", diff --git a/api/tests/unit_tests/controllers/console/agent/test_app_helpers.py b/api/tests/unit_tests/controllers/console/agent/test_app_helpers.py index e05d57bd73a..b396c9bfd1b 100644 --- a/api/tests/unit_tests/controllers/console/agent/test_app_helpers.py +++ b/api/tests/unit_tests/controllers/console/agent/test_app_helpers.py @@ -2,12 +2,15 @@ from unittest.mock import MagicMock from uuid import UUID import pytest +from sqlalchemy.orm import Session from controllers.console.agent import app_helpers -def test_resolve_agent_app_model_reuses_caller_session(monkeypatch: pytest.MonkeyPatch) -> None: - session = MagicMock() +def test_resolve_agent_app_model_reuses_caller_session( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: + session = unbound_session app = MagicMock() service = MagicMock() service.get_agent_app_model.return_value = app @@ -28,8 +31,10 @@ def test_resolve_agent_app_model_reuses_caller_session(monkeypatch: pytest.Monke ) -def test_resolve_agent_runtime_app_model_reuses_caller_session(monkeypatch: pytest.MonkeyPatch) -> None: - session = MagicMock() +def test_resolve_agent_runtime_app_model_reuses_caller_session( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: + session = unbound_session app = MagicMock() service = MagicMock() service.get_agent_runtime_app_model.return_value = app diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py b/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py index 265c542e635..ed227fd09d6 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_app_sandbox.py @@ -2,43 +2,85 @@ from __future__ import annotations from inspect import unwrap from types import SimpleNamespace -from unittest.mock import MagicMock import pytest from dify_agent.client import DifyAgentClientError, DifyAgentHTTPError, DifyAgentTimeoutError -from dify_agent.protocol import SandboxListResponse, SandboxReadResponse +from dify_agent.protocol import BindingFileListResponse, BindingFileReadResponse from controllers.console import agent_app_sandbox as module from models.model import App, AppMode, IconType -from services.agent_app_sandbox_service import AgentSandboxInfo, AgentSandboxInspectorError, AgentSandboxUploadDownload +from services.agent_app_sandbox_service import AgentSandboxDownload, AgentSandboxInfo, AgentSandboxInspectorError class _AgentAppService: def __init__(self) -> None: - self.calls: list[tuple[str, str, str, str, str]] = [] + self.calls: list[tuple[str, str, str, str, str, str, str, str]] = [] - def get_info(self, *, tenant_id: str, app_id: str, conversation_id: str) -> AgentSandboxInfo: - self.calls.append(("info", tenant_id, app_id, conversation_id, "")) - return AgentSandboxInfo(session_id="abc1234", workspace_cwd="~/workspace/abc1234") + def resolve_app_id(self, *, tenant_id: str, agent_id: str) -> str: + return "app-1" - def list_files(self, *, tenant_id: str, app_id: str, conversation_id: str, path: str) -> SandboxListResponse: - self.calls.append(("list", tenant_id, app_id, conversation_id, path)) - return SandboxListResponse(path=path, entries=[], truncated=False) + def get_info( + self, + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: str, + caller_id: str, + account_id: str, + ) -> AgentSandboxInfo: + self.calls.append(("info", tenant_id, app_id, agent_id, caller_type, caller_id, account_id, "")) + return AgentSandboxInfo(workspace_cwd=".") - def read_file(self, *, tenant_id: str, app_id: str, conversation_id: str, path: str) -> SandboxReadResponse: - self.calls.append(("read", tenant_id, app_id, conversation_id, path)) - return SandboxReadResponse(path=path, size=5, truncated=False, binary=False, text="hello") + def list_files( + self, + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: str, + caller_id: str, + account_id: str, + path: str, + ) -> BindingFileListResponse: + self.calls.append(("list", tenant_id, app_id, agent_id, caller_type, caller_id, account_id, path)) + return BindingFileListResponse(path=path, entries=[], truncated=False) - def upload_file( - self, *, tenant_id: str, app_id: str, conversation_id: str, path: str - ) -> AgentSandboxUploadDownload: - self.calls.append(("upload", tenant_id, app_id, conversation_id, path)) - return AgentSandboxUploadDownload(url="https://files.example/report.txt") + def read_file( + self, + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: str, + caller_id: str, + account_id: str, + path: str, + ) -> BindingFileReadResponse: + self.calls.append(("read", tenant_id, app_id, agent_id, caller_type, caller_id, account_id, path)) + return BindingFileReadResponse(path=path, size=5, truncated=False, binary=False, text="hello") + + def download_file( + self, + *, + tenant_id: str, + app_id: str, + agent_id: str, + caller_type: str, + caller_id: str, + account_id: str, + path: str, + ) -> AgentSandboxDownload: + self.calls.append(("download", tenant_id, app_id, agent_id, caller_type, caller_id, account_id, path)) + return AgentSandboxDownload(url="https://files.example/report.txt") class _WorkflowService: def __init__(self) -> None: - self.calls: list[tuple[str, str, str, str, str, str | None, str]] = [] + self.calls: list[tuple[str, ...]] = [] + + def resolve_app_id(self, *, tenant_id: str, app_id: str) -> str: + return app_id def list_files( self, @@ -47,12 +89,12 @@ class _WorkflowService: app_id: str, workflow_run_id: str, node_id: str, - node_execution_id: str | None, + node_execution_id: str, path: str, session, - ) -> SandboxListResponse: - self.calls.append(("list", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) - return SandboxListResponse(path=path, entries=[], truncated=False) + ) -> BindingFileListResponse: + self.calls.append(("list", tenant_id, app_id, workflow_run_id, node_id, path)) + return BindingFileListResponse(path=path, entries=[], truncated=False) def read_file( self, @@ -61,26 +103,26 @@ class _WorkflowService: app_id: str, workflow_run_id: str, node_id: str, - node_execution_id: str | None, + node_execution_id: str, path: str, session, - ) -> SandboxReadResponse: - self.calls.append(("read", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) - return SandboxReadResponse(path=path, size=5, truncated=False, binary=False, text="hello") + ) -> BindingFileReadResponse: + self.calls.append(("read", tenant_id, app_id, workflow_run_id, node_id, path)) + return BindingFileReadResponse(path=path, size=5, truncated=False, binary=False, text="hello") - def upload_file( + def download_file( self, *, tenant_id: str, app_id: str, workflow_run_id: str, node_id: str, - node_execution_id: str | None, + node_execution_id: str, + account_id: str, path: str, - session, - ) -> AgentSandboxUploadDownload: - self.calls.append(("upload", tenant_id, app_id, workflow_run_id, node_id, node_execution_id, path)) - return AgentSandboxUploadDownload(url="https://files.example/upload.txt") + ) -> AgentSandboxDownload: + self.calls.append(("download", tenant_id, app_id, workflow_run_id, node_id, account_id, path)) + return AgentSandboxDownload(url="https://files.example/download.txt") def _app_model(app_id: str = "app-1") -> App: @@ -124,60 +166,58 @@ def test_handle_maps_sandbox_and_agent_backend_errors() -> None: def test_agent_app_sandbox_resources_proxy_service(monkeypatch: pytest.MonkeyPatch) -> None: service = _AgentAppService() - session = MagicMock() - resolver = MagicMock(return_value=_app_model()) + account = SimpleNamespace(id="account-1") monkeypatch.setattr(module, "AgentAppSandboxService", lambda: service) - monkeypatch.setattr(module, "resolve_agent_runtime_app_model", resolver) monkeypatch.setattr( module, "query_params_from_request", - lambda model: SimpleNamespace(conversation_id="conv-1", path="sub/report.txt"), + lambda model: SimpleNamespace(caller_type="build_draft", caller_id="build-1", path="sub/report.txt"), ) - monkeypatch.setattr( - module, - "request", - SimpleNamespace(get_json=lambda silent=True: {"conversation_id": "conv-1", "path": "report.txt"}), + info = unwrap(module.AgentAppSandboxInfoResource.get)(object(), account, "tenant-1", "agent-1") + listing = unwrap(module.AgentAppSandboxListResource.get)(object(), account, "tenant-1", "agent-1") + preview = unwrap(module.AgentAppSandboxReadResource.get)(object(), account, "tenant-1", "agent-1") + req_data = module.AgentSandboxDownloadPayload.model_validate( + {"caller_type": "build_draft", "caller_id": "build-1", "path": "report.txt"} ) + download = unwrap(module.AgentAppSandboxDownloadResource.post)(object(), req_data, account, "tenant-1", "agent-1") - info = unwrap(module.AgentAppSandboxInfoResource.get)(object(), session, "tenant-1", "agent-1") - listing = unwrap(module.AgentAppSandboxListResource.get)(object(), session, "tenant-1", "agent-1") - preview = unwrap(module.AgentAppSandboxReadResource.get)(object(), session, "tenant-1", "agent-1") - upload = unwrap(module.AgentAppSandboxUploadResource.post)(object(), session, "tenant-1", "agent-1") - - assert info == {"session_id": "abc1234", "workspace_cwd": "~/workspace/abc1234"} + assert info == {"workspace_cwd": "."} assert listing["path"] == "sub/report.txt" assert preview["text"] == "hello" - assert upload == {"url": "https://files.example/report.txt"} + assert download == {"url": "https://files.example/report.txt"} assert service.calls == [ - ("info", "tenant-1", "app-1", "conv-1", ""), - ("list", "tenant-1", "app-1", "conv-1", "sub/report.txt"), - ("read", "tenant-1", "app-1", "conv-1", "sub/report.txt"), - ("upload", "tenant-1", "app-1", "conv-1", "report.txt"), + ("info", "tenant-1", "app-1", "agent-1", "build_draft", "build-1", "account-1", ""), + ("list", "tenant-1", "app-1", "agent-1", "build_draft", "build-1", "account-1", "sub/report.txt"), + ("read", "tenant-1", "app-1", "agent-1", "build_draft", "build-1", "account-1", "sub/report.txt"), + ("download", "tenant-1", "app-1", "agent-1", "build_draft", "build-1", "account-1", "report.txt"), ] - assert all(call.kwargs["session"] is session for call in resolver.call_args_list) def test_agent_app_sandbox_resource_returns_normalized_errors(monkeypatch: pytest.MonkeyPatch) -> None: class FailingService: + def resolve_app_id(self, **kwargs): + return "app-1" + def get_info(self, **kwargs): - raise AgentSandboxInspectorError("no_active_session", "no active session", status_code=404) + raise AgentSandboxInspectorError("no_active_binding", "no active binding", status_code=404) def list_files(self, **kwargs): - raise AgentSandboxInspectorError("no_active_session", "no active session", status_code=404) + raise AgentSandboxInspectorError("no_active_binding", "no active binding", status_code=404) monkeypatch.setattr(module, "AgentAppSandboxService", FailingService) - session = MagicMock() - monkeypatch.setattr(module, "resolve_agent_runtime_app_model", MagicMock(return_value=_app_model())) + account = SimpleNamespace(id="account-1") monkeypatch.setattr( - module, "query_params_from_request", lambda model: SimpleNamespace(conversation_id="conv-1", path=".") + module, + "query_params_from_request", + lambda model: SimpleNamespace(caller_type="conversation", caller_id="conv-1", path="."), ) - assert unwrap(module.AgentAppSandboxInfoResource.get)(object(), session, "tenant-1", "agent-1") == ( - {"code": "no_active_session", "message": "no active session"}, + assert unwrap(module.AgentAppSandboxInfoResource.get)(object(), account, "tenant-1", "agent-1") == ( + {"code": "no_active_binding", "message": "no active binding"}, 404, ) - assert unwrap(module.AgentAppSandboxListResource.get)(object(), session, "tenant-1", "agent-1") == ( - {"code": "no_active_session", "message": "no active session"}, + assert unwrap(module.AgentAppSandboxListResource.get)(object(), account, "tenant-1", "agent-1") == ( + {"code": "no_active_binding", "message": "no active binding"}, 404, ) @@ -188,12 +228,7 @@ def test_workflow_agent_sandbox_resources_proxy_service(monkeypatch: pytest.Monk monkeypatch.setattr( module, "query_params_from_request", - lambda model: SimpleNamespace(node_execution_id="exec-1", path="out.txt"), - ) - monkeypatch.setattr( - module, - "request", - SimpleNamespace(get_json=lambda silent=True: {"node_execution_id": "exec-1", "path": "upload.txt"}), + lambda model: SimpleNamespace(node_execution_id="execution-1", path="out.txt"), ) app_model = _app_model() @@ -203,15 +238,19 @@ def test_workflow_agent_sandbox_resources_proxy_service(monkeypatch: pytest.Monk preview = unwrap(module.WorkflowAgentSandboxReadResource.get)( object(), "tenant-1", app_model, "run-1", "agent-node" ) - upload = unwrap(module.WorkflowAgentSandboxUploadResource.post)( - object(), "tenant-1", app_model, "run-1", "agent-node" + req_data = module.WorkflowAgentSandboxDownloadPayload.model_validate( + {"node_execution_id": "execution-1", "path": "download.txt"} + ) + account = SimpleNamespace(id="account-1") + download = unwrap(module.WorkflowAgentSandboxDownloadResource.post)( + object(), req_data, "tenant-1", account, "app-1", "run-1", "agent-node" ) assert listing["path"] == "out.txt" assert preview["text"] == "hello" - assert upload == {"url": "https://files.example/upload.txt"} + assert download == {"url": "https://files.example/download.txt"} assert service.calls == [ - ("list", "tenant-1", "app-1", "run-1", "agent-node", "exec-1", "out.txt"), - ("read", "tenant-1", "app-1", "run-1", "agent-node", "exec-1", "out.txt"), - ("upload", "tenant-1", "app-1", "run-1", "agent-node", "exec-1", "upload.txt"), + ("list", "tenant-1", "app-1", "run-1", "agent-node", "out.txt"), + ("read", "tenant-1", "app-1", "run-1", "agent-node", "out.txt"), + ("download", "tenant-1", "app-1", "run-1", "agent-node", "account-1", "download.txt"), ] diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_config_inspector.py b/api/tests/unit_tests/controllers/console/app/test_agent_config_inspector.py index 1d857e7284c..d21d0e2618a 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_config_inspector.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_config_inspector.py @@ -8,9 +8,11 @@ from __future__ import annotations from inspect import unwrap from types import SimpleNamespace -from unittest.mock import MagicMock, PropertyMock, patch +from unittest.mock import MagicMock, patch from flask import Flask +from sqlalchemy import event +from sqlalchemy.orm import Session from controllers.console.app import agent_config_inspector as inspector from controllers.console.app.agent_config_inspector import ( @@ -26,7 +28,6 @@ from controllers.console.app.agent_config_inspector import ( AgentConfigSkillInspectByAgentApi, AgentConfigSkillsApi, AgentConfigSkillUploadByAgentApi, - console_ns, ) from services.agent_config_service import AgentConfigServiceError @@ -46,8 +47,8 @@ _APP = SimpleNamespace( _USER = SimpleNamespace(id="acct-1") -def test_resolve_bound_agent_uses_injected_session(): - session = MagicMock() +def test_resolve_bound_agent_uses_injected_session(unbound_session: Session): + session = unbound_session resolver = MagicMock(return_value="agent-1") app_model = SimpleNamespace(bound_agent_id_with_session=resolver) result = inspector._resolve_agent_id(session, app_model, None) @@ -100,8 +101,10 @@ def test_manifest_resolves_workflow_node_agent_and_normal_draft(): assert config_service.return_value.manifest.call_args.kwargs["config_version_kind"].value == "draft" -def test_normal_draft_resolution_commits_created_draft_before_service_session() -> None: - session = MagicMock() +def test_normal_draft_resolution_commits_created_draft_before_service_session(sqlite_session: Session) -> None: + session = sqlite_session + commits: list[str] = [] + event.listen(session, "after_commit", lambda _session: commits.append("commit")) with patch(f"{_MOD}.AgentComposerService") as composer: composer.load_agent_composer.return_value = {"draft": {"id": "draft-1"}} version_id, version_kind = inspector._resolve_console_version( @@ -114,7 +117,7 @@ def test_normal_draft_resolution_commits_created_draft_before_service_session() ) assert version_id == "draft-1" assert version_kind.value == "draft" - session.commit.assert_called_once() + assert commits == ["commit"] def test_skill_inspect_by_agent_returns_strict_json_response(): @@ -208,13 +211,10 @@ def test_skill_upload_by_agent_delegates_after_version_resolution(): def test_file_upload_by_agent_delegates_to_service_owned_upload_lookup(): raw = _raw(AgentConfigFilesByAgentApi.post) - with app.test_request_context("/?draft_type=debug_build"): + with app.test_request_context("/?draft_type=debug_build", json={"upload_file_id": "upload-1"}): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP), patch(f"{_MOD}.AgentComposerService") as composer, - patch.object( - type(console_ns), "payload", new_callable=PropertyMock, return_value={"upload_file_id": "upload-1"} - ), patch(f"{_MOD}.AgentConfigService") as config_service, ): composer.load_agent_app_build_draft.return_value = {"draft": {"id": "build-draft-1"}} @@ -222,7 +222,14 @@ def test_file_upload_by_agent_delegates_to_service_owned_upload_lookup(): "file": {"id": "guide.txt", "name": "guide.txt", "file_id": "upload-1"}, "config_version": {"id": "build-draft-1", "kind": "build_draft", "writable": True}, } - body, status = raw(AgentConfigFilesByAgentApi(), MagicMock(), "tenant-1", _USER, "agent-1") + body, status = raw( + AgentConfigFilesByAgentApi(), + inspector.AgentConfigFileUploadPayload(upload_file_id="upload-1"), + MagicMock(), + "tenant-1", + _USER, + "agent-1", + ) assert status == 201 assert body["file"]["name"] == "guide.txt" assert config_service.return_value.push_file_for_console.call_args.kwargs["upload_file_id"] == "upload-1" diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_drive_inspector.py b/api/tests/unit_tests/controllers/console/app/test_agent_drive_inspector.py index 453e2f29ff0..ad38a927fcd 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_drive_inspector.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_drive_inspector.py @@ -12,6 +12,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch from flask import Flask +from sqlalchemy.orm import Session from controllers.console.app import agent_drive_inspector as inspector from controllers.console.app.agent_drive_inspector import ( @@ -43,18 +44,17 @@ _APP = SimpleNamespace( ) -def test_resolve_bound_agent_uses_injected_session(): - session = MagicMock() +def test_resolve_bound_agent_uses_injected_session(unbound_session: Session): resolver = MagicMock(return_value="agent-1") app_model = SimpleNamespace(bound_agent_id_with_session=resolver) - result = inspector._resolve_agent_id(session, app_model, None) + result = inspector._resolve_agent_id(unbound_session, app_model, None) assert result == "agent-1" - resolver.assert_called_once_with(session=session) - assert resolver.call_args.kwargs["session"] is session + resolver.assert_called_once_with(session=unbound_session) + assert resolver.call_args.kwargs["session"] is unbound_session -def test_list_filters_value_pointers_out_of_console_payload(): +def test_list_filters_value_pointers_out_of_console_payload(unbound_session: Session): raw = _raw(AgentDriveListApi.get) with app.test_request_context("/?prefix=pdf-toolkit/"): with patch(f"{_MOD}.AgentDriveService") as drive: @@ -69,15 +69,14 @@ def test_list_filters_value_pointers_out_of_console_payload(): "created_at": 1718000000, } ] - body = raw(AgentDriveListApi(), MagicMock(), _APP) + body = raw(AgentDriveListApi(), unbound_session, _APP) assert body["items"][0]["key"] == "pdf-toolkit/SKILL.md" assert "file_id" not in body["items"][0] assert drive.return_value.manifest.call_args.kwargs["prefix"] == "pdf-toolkit/" -def test_list_by_agent_filters_value_pointers_out_of_console_payload(): +def test_list_by_agent_filters_value_pointers_out_of_console_payload(unbound_session: Session): raw = _raw(AgentDriveListByAgentApi.get) - session = MagicMock() with app.test_request_context("/?prefix=pdf-toolkit/"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, @@ -94,28 +93,27 @@ def test_list_by_agent_filters_value_pointers_out_of_console_payload(): "created_at": 1718000000, } ] - body = raw(AgentDriveListByAgentApi(), session, "tenant-1", "agent-1") + body = raw(AgentDriveListByAgentApi(), unbound_session, "tenant-1", "agent-1") assert body["items"][0]["key"] == "pdf-toolkit/SKILL.md" assert "file_id" not in body["items"][0] - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=unbound_session, tenant_id="tenant-1", agent_id="agent-1") assert drive.return_value.manifest.call_args.kwargs["agent_id"] == "agent-1" - assert drive.return_value.manifest.call_args.kwargs["session"] is session + assert drive.return_value.manifest.call_args.kwargs["session"] is unbound_session -def test_list_resolves_workflow_node_binding_agent(): +def test_list_resolves_workflow_node_binding_agent(unbound_session: Session): raw = _raw(AgentDriveListApi.get) with app.test_request_context("/?node_id=agent-node-1"): with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentDriveService") as drive: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" drive.return_value.manifest.return_value = [] - raw(AgentDriveListApi(), MagicMock(), _APP) + raw(AgentDriveListApi(), unbound_session, _APP) assert drive.return_value.manifest.call_args.kwargs["agent_id"] == "wf-agent-9" assert composer.resolve_workflow_node_agent_id.call_args.kwargs["node_id"] == "agent-node-1" -def test_skill_list_by_agent_calls_service(): +def test_skill_list_by_agent_calls_service(unbound_session: Session): raw = _raw(AgentDriveSkillListByAgentApi.get) - session = MagicMock() with app.test_request_context("/"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, @@ -134,27 +132,26 @@ def test_skill_list_by_agent_calls_service(): "created_at": 1718000000, } ] - body = raw(AgentDriveSkillListByAgentApi(), session, "tenant-1", "agent-1") + body = raw(AgentDriveSkillListByAgentApi(), unbound_session, "tenant-1", "agent-1") assert body["items"][0]["path"] == "pdf-toolkit" - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=unbound_session, tenant_id="tenant-1", agent_id="agent-1") assert drive.return_value.list_skills.call_args.kwargs["agent_id"] == "agent-1" - assert drive.return_value.list_skills.call_args.kwargs["session"] is session + assert drive.return_value.list_skills.call_args.kwargs["session"] is unbound_session -def test_skill_list_resolves_workflow_node_binding_agent(): +def test_skill_list_resolves_workflow_node_binding_agent(unbound_session: Session): raw = _raw(AgentDriveSkillListApi.get) with app.test_request_context("/?node_id=agent-node-1"): with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentDriveService") as drive: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" drive.return_value.list_skills.return_value = [] - body = raw(AgentDriveSkillListApi(), MagicMock(), _APP) + body = raw(AgentDriveSkillListApi(), unbound_session, _APP) assert body == {"items": []} assert drive.return_value.list_skills.call_args.kwargs["agent_id"] == "wf-agent-9" -def test_skill_inspect_by_agent_returns_strict_json_response(): +def test_skill_inspect_by_agent_returns_strict_json_response(unbound_session: Session): raw = _raw(AgentDriveSkillInspectByAgentApi.get) - session = MagicMock() payload = { "path": "pdf-toolkit", "skill_md_key": "pdf-toolkit/SKILL.md", @@ -191,14 +188,14 @@ def test_skill_inspect_by_agent_returns_strict_json_response(): patch(f"{_MOD}.AgentDriveService") as drive, ): drive.return_value.inspect_skill.return_value = payload - response = raw(AgentDriveSkillInspectByAgentApi(), session, "tenant-1", "agent-1", "pdf-toolkit") + response = raw(AgentDriveSkillInspectByAgentApi(), unbound_session, "tenant-1", "agent-1", "pdf-toolkit") assert response.status_code == 200 assert response.get_json()["skill_md"]["text"] == "# PDF Toolkit\nUse it.\n" assert b"# PDF Toolkit\\nUse it.\\n" in response.get_data() - assert drive.return_value.inspect_skill.call_args.kwargs["session"] is session + assert drive.return_value.inspect_skill.call_args.kwargs["session"] is unbound_session -def test_skill_inspect_resolves_workflow_node_binding_agent(): +def test_skill_inspect_resolves_workflow_node_binding_agent(unbound_session: Session): raw = _raw(AgentDriveSkillInspectApi.get) payload = { "path": "pdf-toolkit", @@ -220,24 +217,23 @@ def test_skill_inspect_resolves_workflow_node_binding_agent(): with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentDriveService") as drive: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-9" drive.return_value.inspect_skill.return_value = payload - response = raw(AgentDriveSkillInspectApi(), MagicMock(), _APP, "pdf-toolkit") + response = raw(AgentDriveSkillInspectApi(), unbound_session, _APP, "pdf-toolkit") assert response.get_json()["path"] == "pdf-toolkit" assert drive.return_value.inspect_skill.call_args.kwargs["agent_id"] == "wf-agent-9" -def test_list_400_when_no_agent_bound(): +def test_list_400_when_no_agent_bound(unbound_session: Session): raw = _raw(AgentDriveListApi.get) resolver = MagicMock(return_value=None) app_without_agent = SimpleNamespace(bound_agent_id_with_session=resolver) - session = MagicMock() with app.test_request_context("/"): - body, status = raw(AgentDriveListApi(), session, app_without_agent) + body, status = raw(AgentDriveListApi(), unbound_session, app_without_agent) assert status == 400 assert body["code"] == "agent_not_bound" - resolver.assert_called_once_with(session=session) + resolver.assert_called_once_with(session=unbound_session) -def test_preview_passes_through_and_maps_errors(): +def test_preview_passes_through_and_maps_errors(unbound_session: Session): raw = _raw(AgentDrivePreviewApi.get) with app.test_request_context("/?key=pdf-toolkit/SKILL.md"): with patch(f"{_MOD}.AgentDriveService") as drive: @@ -248,21 +244,20 @@ def test_preview_passes_through_and_maps_errors(): "binary": False, "text": "# hi", } - body = raw(AgentDrivePreviewApi(), MagicMock(), _APP) + body = raw(AgentDrivePreviewApi(), unbound_session, _APP) assert body["text"] == "# hi" with app.test_request_context("/?key=ghost/SKILL.md"): with patch(f"{_MOD}.AgentDriveService") as drive: drive.return_value.preview.side_effect = AgentDriveError( "drive_key_not_found", "no drive entry", status_code=404 ) - body, status = raw(AgentDrivePreviewApi(), MagicMock(), _APP) + body, status = raw(AgentDrivePreviewApi(), unbound_session, _APP) assert status == 404 assert body["code"] == "drive_key_not_found" -def test_preview_by_agent_passes_through_and_maps_errors(): +def test_preview_by_agent_passes_through_and_maps_errors(unbound_session: Session): raw = _raw(AgentDrivePreviewByAgentApi.get) - session = MagicMock() with app.test_request_context("/?key=pdf-toolkit/SKILL.md"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, @@ -275,10 +270,10 @@ def test_preview_by_agent_passes_through_and_maps_errors(): "binary": False, "text": "# hi", } - body = raw(AgentDrivePreviewByAgentApi(), session, "tenant-1", "agent-1") + body = raw(AgentDrivePreviewByAgentApi(), unbound_session, "tenant-1", "agent-1") assert body["text"] == "# hi" - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") - assert drive.return_value.preview.call_args.kwargs["session"] is session + resolve_app.assert_called_once_with(session=unbound_session, tenant_id="tenant-1", agent_id="agent-1") + assert drive.return_value.preview.call_args.kwargs["session"] is unbound_session with app.test_request_context("/?key=ghost/SKILL.md"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP), @@ -287,30 +282,29 @@ def test_preview_by_agent_passes_through_and_maps_errors(): drive.return_value.preview.side_effect = AgentDriveError( "drive_key_not_found", "no drive entry", status_code=404 ) - body, status = raw(AgentDrivePreviewByAgentApi(), session, "tenant-1", "agent-1") + body, status = raw(AgentDrivePreviewByAgentApi(), unbound_session, "tenant-1", "agent-1") assert status == 404 assert body["code"] == "drive_key_not_found" -def test_download_returns_signed_url_json(): +def test_download_returns_signed_url_json(unbound_session: Session): raw = _raw(AgentDriveDownloadApi.get) with app.test_request_context("/?key=pdf-toolkit/.DIFY-SKILL-FULL.zip"): with patch(f"{_MOD}.AgentDriveService") as drive: drive.return_value.download_url.return_value = "https://signed.example/zip" - body = raw(AgentDriveDownloadApi(), MagicMock(), _APP) + body = raw(AgentDriveDownloadApi(), unbound_session, _APP) assert body == {"url": "https://signed.example/zip"} -def test_download_by_agent_returns_signed_url_json(): +def test_download_by_agent_returns_signed_url_json(unbound_session: Session): raw = _raw(AgentDriveDownloadByAgentApi.get) - session = MagicMock() with app.test_request_context("/?key=pdf-toolkit/.DIFY-SKILL-FULL.zip"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.AgentDriveService") as drive, ): drive.return_value.download_url.return_value = "https://signed.example/zip" - body = raw(AgentDriveDownloadByAgentApi(), session, "tenant-1", "agent-1") + body = raw(AgentDriveDownloadByAgentApi(), unbound_session, "tenant-1", "agent-1") assert body == {"url": "https://signed.example/zip"} - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") - assert drive.return_value.download_url.call_args.kwargs["session"] is session + resolve_app.assert_called_once_with(session=unbound_session, tenant_id="tenant-1", agent_id="agent-1") + assert drive.return_value.download_url.call_args.kwargs["session"] is unbound_session diff --git a/api/tests/unit_tests/controllers/console/app/test_agent_skills.py b/api/tests/unit_tests/controllers/console/app/test_agent_skills.py index e555e9c5094..f0496fda088 100644 --- a/api/tests/unit_tests/controllers/console/app/test_agent_skills.py +++ b/api/tests/unit_tests/controllers/console/app/test_agent_skills.py @@ -8,12 +8,15 @@ bare Flask request context with the services mocked — covering request handlin from __future__ import annotations import io +from datetime import UTC, datetime from inspect import unwrap from types import SimpleNamespace from unittest.mock import MagicMock, patch +from uuid import uuid4 import pytest from flask import Flask +from sqlalchemy.orm import Session from controllers.console.app import agent as agent_controller from controllers.console.app.agent import ( @@ -23,12 +26,16 @@ from controllers.console.app.agent import ( AgentSkillUploadApi, AgentSkillUploadByAgentApi, ) -from models.model import AppMode +from extensions.storage.storage_type import StorageType +from models.enums import CreatorUserRole +from models.model import AppMode, UploadFile from services.agent.skill_package_service import SkillPackageError from services.agent_drive_service import AgentDriveError _MOD = "controllers.console.app.agent" app = Flask(__name__) +_TENANT_ID = "00000000-0000-0000-0000-000000000010" +_UPLOAD_FILE_ID = "0fa6f9bc-3416-4476-8857-a13129704dd9" def _raw(method): @@ -43,30 +50,49 @@ def _file_ctx(*, files: dict[str, bytes] | None = None): _USER = SimpleNamespace(id="user-1") _APP = SimpleNamespace( id="app-1", - tenant_id="tenant-1", + tenant_id=_TENANT_ID, mode=AppMode.AGENT, bound_agent_id_with_session=lambda *, session: "agent-1", ) _WORKFLOW_APP = SimpleNamespace( id="app-1", - tenant_id="tenant-1", + tenant_id=_TENANT_ID, mode=AppMode.WORKFLOW, bound_agent_id_with_session=lambda *, session: None, ) -def test_resolve_bound_agent_uses_injected_session(): - session = MagicMock() +def _persist_upload(session: Session, *, name: str = "sample.pdf") -> UploadFile: + upload = UploadFile( + tenant_id=_TENANT_ID, + storage_type=StorageType.LOCAL, + key=f"uploads/{name}", + name=name, + size=5, + extension="pdf", + mime_type="application/pdf", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + created_at=datetime.now(UTC), + used=False, + ) + upload.id = _UPLOAD_FILE_ID + session.add(upload) + session.commit() + return upload + + +def test_resolve_bound_agent_uses_injected_session(unbound_session: Session): resolver = MagicMock(return_value="agent-1") app_model = SimpleNamespace(bound_agent_id_with_session=resolver) - result = agent_controller._resolve_agent_id(session, app_model, None) + result = agent_controller._resolve_agent_id(unbound_session, app_model, None) assert result == "agent-1" - resolver.assert_called_once_with(session=session) - assert resolver.call_args.kwargs["session"] is session + resolver.assert_called_once_with(session=unbound_session) + assert resolver.call_args.kwargs["session"] is unbound_session -def test_upload_standardizes_into_drive_and_returns_skill_ref(): +def test_upload_standardizes_into_drive_and_returns_skill_ref(unbound_session: Session): raw = _raw(AgentSkillUploadApi.post) with _file_ctx(files={"file": b"zip-bytes"}): with patch(f"{_MOD}.SkillStandardizeService") as svc: @@ -74,61 +100,59 @@ def test_upload_standardizes_into_drive_and_returns_skill_ref(): "skill": {"path": "skill-a", "skill_md_key": "skill-a/SKILL.md"}, "manifest": {"name": "Skill A"}, } - body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _APP) + body, status = raw(AgentSkillUploadApi(), unbound_session, _USER, _APP) assert status == 201 assert body["skill"] == {"path": "skill-a", "skill_md_key": "skill-a/SKILL.md"} assert svc.return_value.standardize.call_args.kwargs["agent_id"] == "agent-1" -def test_upload_by_agent_resolves_app_and_standardizes_into_drive(): +def test_upload_by_agent_resolves_app_and_standardizes_into_drive(unbound_session: Session): raw = _raw(AgentSkillUploadByAgentApi.post) with _file_ctx(files={"file": b"zip-bytes"}): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.SkillStandardizeService") as svc, ): - session = MagicMock() svc.return_value.standardize.return_value = {"skill": {"path": "skill-a"}, "manifest": {}} - body, status = raw(AgentSkillUploadByAgentApi(), session, "tenant-1", _USER, "agent-1") + body, status = raw(AgentSkillUploadByAgentApi(), unbound_session, "tenant-1", _USER, "agent-1") assert status == 201 assert body["skill"] == {"path": "skill-a"} - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=unbound_session, tenant_id="tenant-1", agent_id="agent-1") assert svc.return_value.standardize.call_args.kwargs["agent_id"] == "agent-1" -def test_upload_no_file_is_400(): +def test_upload_no_file_is_400(unbound_session: Session): raw = _raw(AgentSkillUploadApi.post) with _file_ctx(files={}): - body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _APP) + body, status = raw(AgentSkillUploadApi(), unbound_session, _USER, _APP) assert status == 400 assert body["code"] == "no_file" -def test_upload_maps_package_error(): +def test_upload_maps_package_error(unbound_session: Session): raw = _raw(AgentSkillUploadApi.post) with _file_ctx(files={"file": b"bad"}): with patch(f"{_MOD}.SkillStandardizeService") as svc: svc.return_value.standardize.side_effect = SkillPackageError( "missing_skill_md", "no SKILL.md", status_code=400 ) - body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _APP) + body, status = raw(AgentSkillUploadApi(), unbound_session, _USER, _APP) assert status == 400 assert body["code"] == "missing_skill_md" -def test_upload_no_bound_agent_is_400(): +def test_upload_no_bound_agent_is_400(unbound_session: Session): raw = _raw(AgentSkillUploadApi.post) resolver = MagicMock(return_value=None) app_without_agent = SimpleNamespace(bound_agent_id_with_session=resolver) - session = MagicMock() with _file_ctx(files={"file": b"zip"}): - body, status = raw(AgentSkillUploadApi(), session, _USER, app_without_agent) + body, status = raw(AgentSkillUploadApi(), unbound_session, _USER, app_without_agent) assert status == 400 assert body["code"] == "agent_not_bound" - resolver.assert_called_once_with(session=session) + resolver.assert_called_once_with(session=unbound_session) -def test_upload_resolves_workflow_node_agent(): +def test_upload_resolves_workflow_node_agent(unbound_session: Session): raw = _raw(AgentSkillUploadApi.post) with app.test_request_context( "/?node_id=agent-node-1", method="POST", data={"file": (io.BytesIO(b"zip"), "skill.zip")} @@ -136,18 +160,18 @@ def test_upload_resolves_workflow_node_agent(): with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.SkillStandardizeService") as svc: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-1" svc.return_value.standardize.return_value = {"skill": {"path": "s"}, "manifest": {}} - body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _WORKFLOW_APP) + body, status = raw(AgentSkillUploadApi(), unbound_session, _USER, _WORKFLOW_APP) assert status == 201 assert body["skill"] == {"path": "s"} assert svc.return_value.standardize.call_args.kwargs["agent_id"] == "wf-agent-1" -def test_upload_maps_drive_error(): +def test_upload_maps_drive_error(unbound_session: Session): raw = _raw(AgentSkillUploadApi.post) with _file_ctx(files={"file": b"zip"}): with patch(f"{_MOD}.SkillStandardizeService") as svc: svc.return_value.standardize.side_effect = AgentDriveError("source_not_found", "nope", status_code=404) - body, status = raw(AgentSkillUploadApi(), MagicMock(), _USER, _APP) + body, status = raw(AgentSkillUploadApi(), unbound_session, _USER, _APP) assert status == 404 assert body["code"] == "source_not_found" @@ -156,86 +180,81 @@ def _json_ctx(payload: dict | None = None, *, method: str = "POST", query_string return app.test_request_context(f"/?{query_string}", method=method, json=payload or {}) -def test_files_commit_validates_upload_and_returns_drive_ref(): +def test_files_commit_validates_upload_and_returns_drive_ref(sqlite_session: Session): from controllers.console.app.agent import AgentDriveFilesApi raw = _raw(AgentDriveFilesApi.post) - upload = SimpleNamespace(id="uf-1", name="sample qna.pdf") - with _json_ctx({"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"}): + upload = _persist_upload(sqlite_session, name="sample qna.pdf") + with _json_ctx({"upload_file_id": _UPLOAD_FILE_ID}): with patch(f"{_MOD}.console_ns") as ns, patch(f"{_MOD}.AgentDriveService") as drive: - session = MagicMock() - ns.payload = {"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"} - session.scalar.return_value = upload + ns.payload = {"upload_file_id": _UPLOAD_FILE_ID} drive.return_value.commit.return_value = [ {"key": "files/sample qna.pdf", "size": 5, "mime_type": "application/pdf"} ] - body, status = raw(AgentDriveFilesApi(), session, _USER, _APP) + body, status = raw(AgentDriveFilesApi(), sqlite_session, _USER, _APP) assert status == 201 assert body["file"]["drive_key"] == "files/sample qna.pdf" - assert body["file"]["file_id"] == "uf-1" + assert body["file"]["file_id"] == upload.id item = drive.return_value.commit.call_args.kwargs["items"][0] assert item.value_owned_by_drive is True assert item.file_ref.kind == "upload_file" -def test_files_by_agent_commit_uses_agent_route_and_ignores_node_id(): +def test_files_by_agent_commit_uses_agent_route_and_ignores_node_id(sqlite_session: Session): raw = _raw(AgentDriveFilesByAgentApi.post) - upload = SimpleNamespace(id="uf-1", name="sample.pdf") - with _json_ctx({"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"}, query_string="node_id=ignored"): + _persist_upload(sqlite_session) + with _json_ctx({"upload_file_id": _UPLOAD_FILE_ID}, query_string="node_id=ignored"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.console_ns") as ns, patch(f"{_MOD}.AgentDriveService") as drive, ): - session = MagicMock() - ns.payload = {"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"} - session.scalar.return_value = upload + ns.payload = {"upload_file_id": _UPLOAD_FILE_ID} drive.return_value.commit.return_value = [ {"key": "files/sample.pdf", "size": 5, "mime_type": "application/pdf"} ] - body, status = raw(AgentDriveFilesByAgentApi(), session, "tenant-1", _USER, "agent-1") + body, status = raw(AgentDriveFilesByAgentApi(), sqlite_session, "tenant-1", _USER, "agent-1") assert status == 201 - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=sqlite_session, tenant_id="tenant-1", agent_id="agent-1") -def test_files_commit_404_when_upload_not_in_tenant(): +def test_files_commit_404_when_upload_not_in_tenant(sqlite_session: Session): from controllers.console.app.agent import AgentDriveFilesApi raw = _raw(AgentDriveFilesApi.post) - with _json_ctx({"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"}): + other_upload = _persist_upload(sqlite_session) + other_upload.tenant_id = str(uuid4()) + sqlite_session.commit() + with _json_ctx({"upload_file_id": _UPLOAD_FILE_ID}): with patch(f"{_MOD}.console_ns") as ns: - session = MagicMock() - ns.payload = {"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"} - session.scalar.return_value = None - body, status = raw(AgentDriveFilesApi(), session, _USER, _APP) + ns.payload = {"upload_file_id": _UPLOAD_FILE_ID} + body, status = raw(AgentDriveFilesApi(), sqlite_session, _USER, _APP) assert status == 404 assert body["code"] == "upload_file_not_found" -def test_files_commit_resolves_workflow_node_agent(): +def test_files_commit_resolves_workflow_node_agent(sqlite_session: Session): from controllers.console.app.agent import AgentDriveFilesApi raw = _raw(AgentDriveFilesApi.post) - upload = SimpleNamespace(id="uf-1", name="sample.pdf") - with _json_ctx({"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"}, query_string="node_id=agent-node-1"): + _persist_upload(sqlite_session) + with _json_ctx({"upload_file_id": _UPLOAD_FILE_ID}, query_string="node_id=agent-node-1"): with ( patch(f"{_MOD}.console_ns") as ns, patch(f"{_MOD}.AgentDriveService") as drive, patch(f"{_MOD}.AgentComposerService") as composer, ): - session = MagicMock() - ns.payload = {"upload_file_id": "0fa6f9bc-3416-4476-8857-a13129704dd9"} - session.scalar.return_value = upload + ns.payload = {"upload_file_id": _UPLOAD_FILE_ID} composer.resolve_workflow_node_agent_id.return_value = "wf-agent-1" drive.return_value.commit.return_value = [ {"key": "files/sample.pdf", "size": 5, "mime_type": "application/pdf"} ] - body, status = raw(AgentDriveFilesApi(), session, _USER, _WORKFLOW_APP) + body, status = raw(AgentDriveFilesApi(), sqlite_session, _USER, _WORKFLOW_APP) assert status == 201 assert drive.return_value.commit.call_args.kwargs["agent_id"] == "wf-agent-1" -def test_files_delete_updates_soul_then_drive(): +def test_files_delete_updates_soul_then_drive(unbound_session: Session): from controllers.console.app.agent import AgentDriveFilesApi raw = _raw(AgentDriveFilesApi.delete) @@ -245,26 +264,25 @@ def test_files_delete_updates_soul_then_drive(): drive.return_value.commit.side_effect = lambda **kw: ( calls.append("drive") or [{"key": "files/sample.pdf", "removed": True}] ) - body = raw(AgentDriveFilesApi(), MagicMock(), _USER, _APP) + body = raw(AgentDriveFilesApi(), unbound_session, _USER, _APP) assert calls == ["drive"] assert body == {"result": "success", "removed_keys": ["files/sample.pdf"]} -def test_files_by_agent_delete_uses_agent_route_and_ignores_node_id(): +def test_files_by_agent_delete_uses_agent_route_and_ignores_node_id(unbound_session: Session): raw = _raw(AgentDriveFilesByAgentApi.delete) with _json_ctx(method="DELETE", query_string="key=files/sample.pdf&node_id=ignored"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.AgentDriveService") as drive, ): - session = MagicMock() drive.return_value.commit.return_value = [{"key": "files/sample.pdf", "removed": True}] - body = raw(AgentDriveFilesByAgentApi(), session, "tenant-1", _USER, "agent-1") + body = raw(AgentDriveFilesByAgentApi(), unbound_session, "tenant-1", _USER, "agent-1") assert body == {"result": "success", "removed_keys": ["files/sample.pdf"]} - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=unbound_session, tenant_id="tenant-1", agent_id="agent-1") -def test_files_delete_resolves_workflow_node_agent(): +def test_files_delete_resolves_workflow_node_agent(unbound_session: Session): from controllers.console.app.agent import AgentDriveFilesApi raw = _raw(AgentDriveFilesApi.delete) @@ -272,12 +290,12 @@ def test_files_delete_resolves_workflow_node_agent(): with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.AgentDriveService") as drive: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-1" drive.return_value.commit.return_value = [{"key": "files/sample.pdf", "removed": True}] - body = raw(AgentDriveFilesApi(), MagicMock(), _USER, _WORKFLOW_APP) + body = raw(AgentDriveFilesApi(), unbound_session, _USER, _WORKFLOW_APP) assert body == {"result": "success", "removed_keys": ["files/sample.pdf"]} assert drive.return_value.commit.call_args.kwargs["agent_id"] == "wf-agent-1" -def test_files_delete_survives_drive_failure(): +def test_files_delete_survives_drive_failure(unbound_session: Session): from controllers.console.app.agent import AgentDriveFilesApi raw = _raw(AgentDriveFilesApi.delete) @@ -285,10 +303,10 @@ def test_files_delete_survives_drive_failure(): with patch(f"{_MOD}.AgentDriveService") as drive: drive.return_value.commit.side_effect = RuntimeError("storage down") with pytest.raises(RuntimeError, match="storage down"): - raw(AgentDriveFilesApi(), MagicMock(), _USER, _APP) + raw(AgentDriveFilesApi(), unbound_session, _USER, _APP) -def test_skill_delete_uses_slug_prefix_and_is_idempotent(): +def test_skill_delete_uses_slug_prefix_and_is_idempotent(unbound_session: Session): from controllers.console.app.agent import AgentSkillApi raw = _raw(AgentSkillApi.delete) @@ -298,38 +316,37 @@ def test_skill_delete_uses_slug_prefix_and_is_idempotent(): {"key": "tender-analyzer/SKILL.md", "removed": True}, {"key": "tender-analyzer/.DIFY-SKILL-FULL.zip", "removed": True}, ] - body = raw(AgentSkillApi(), MagicMock(), _USER, _APP, "tender-analyzer") + body = raw(AgentSkillApi(), unbound_session, _USER, _APP, "tender-analyzer") assert body == { "result": "success", "removed_keys": ["tender-analyzer/SKILL.md", "tender-analyzer/.DIFY-SKILL-FULL.zip"], } -def test_skill_delete_by_agent_uses_agent_route(): +def test_skill_delete_by_agent_uses_agent_route(unbound_session: Session): raw = _raw(AgentSkillByAgentApi.delete) with _json_ctx(method="DELETE", query_string="node_id=ignored"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.AgentDriveService") as drive, ): - session = MagicMock() drive.return_value.commit.return_value = [{"key": "tender-analyzer/SKILL.md", "removed": True}] - body = raw(AgentSkillByAgentApi(), session, "tenant-1", _USER, "agent-1", "tender-analyzer") + body = raw(AgentSkillByAgentApi(), unbound_session, "tenant-1", _USER, "agent-1", "tender-analyzer") assert body == {"result": "success", "removed_keys": ["tender-analyzer/SKILL.md"]} - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=unbound_session, tenant_id="tenant-1", agent_id="agent-1") -def test_skill_delete_rejects_path_like_slug(): +def test_skill_delete_rejects_path_like_slug(unbound_session: Session): from controllers.console.app.agent import AgentSkillApi raw = _raw(AgentSkillApi.delete) with _json_ctx(method="DELETE"): - body, status = raw(AgentSkillApi(), MagicMock(), _USER, _APP, "a/b") + body, status = raw(AgentSkillApi(), unbound_session, _USER, _APP, "a/b") assert status == 400 assert body["code"] == "drive_key_invalid" -def test_infer_tools_returns_draft_suggestions(): +def test_infer_tools_returns_draft_suggestions(unbound_session: Session): from controllers.console.app.agent import AgentSkillInferToolsApi raw = _raw(AgentSkillInferToolsApi.post) @@ -340,27 +357,32 @@ def test_infer_tools_returns_draft_suggestions(): "cli_tools": [{"name": "ffmpeg", "inferred_from": "audio-transcribe"}], "reason": None, } - body = raw(AgentSkillInferToolsApi(), MagicMock(), _APP, "audio-transcribe") + body = raw(AgentSkillInferToolsApi(), unbound_session, _APP, "audio-transcribe") assert body["inferable"] is True assert svc.return_value.infer.call_args.kwargs["slug"] == "audio-transcribe" -def test_infer_tools_by_agent_uses_agent_route(): +def test_infer_tools_by_agent_uses_agent_route(unbound_session: Session): raw = _raw(AgentSkillInferToolsByAgentApi.post) with _json_ctx(query_string="node_id=ignored"): with ( patch(f"{_MOD}.resolve_agent_runtime_app_model", return_value=_APP) as resolve_app, patch(f"{_MOD}.SkillToolInferenceService") as svc, ): - session = MagicMock() svc.return_value.infer.return_value = {"inferable": True, "cli_tools": [], "reason": None} - body = raw(AgentSkillInferToolsByAgentApi(), session, "tenant-1", "agent-1", "audio-transcribe") + body = raw( + AgentSkillInferToolsByAgentApi(), + unbound_session, + "tenant-1", + "agent-1", + "audio-transcribe", + ) assert body["inferable"] is True - resolve_app.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") + resolve_app.assert_called_once_with(session=unbound_session, tenant_id="tenant-1", agent_id="agent-1") assert svc.return_value.infer.call_args.kwargs["agent_id"] == "agent-1" -def test_infer_tools_resolves_workflow_node_agent(): +def test_infer_tools_resolves_workflow_node_agent(unbound_session: Session): from controllers.console.app.agent import AgentSkillInferToolsApi raw = _raw(AgentSkillInferToolsApi.post) @@ -368,12 +390,12 @@ def test_infer_tools_resolves_workflow_node_agent(): with patch(f"{_MOD}.AgentComposerService") as composer, patch(f"{_MOD}.SkillToolInferenceService") as svc: composer.resolve_workflow_node_agent_id.return_value = "wf-agent-1" svc.return_value.infer.return_value = {"inferable": False, "cli_tools": [], "reason": "none"} - body = raw(AgentSkillInferToolsApi(), MagicMock(), _WORKFLOW_APP, "audio-transcribe") + body = raw(AgentSkillInferToolsApi(), unbound_session, _WORKFLOW_APP, "audio-transcribe") assert body["inferable"] is False assert svc.return_value.infer.call_args.kwargs["agent_id"] == "wf-agent-1" -def test_infer_tools_maps_inference_errors(): +def test_infer_tools_maps_inference_errors(unbound_session: Session): from controllers.console.app.agent import AgentSkillInferToolsApi from services.agent.skill_tool_inference_service import SkillToolInferenceError @@ -383,19 +405,19 @@ def test_infer_tools_maps_inference_errors(): svc.return_value.infer.side_effect = SkillToolInferenceError( "default_model_not_configured", "no model", status_code=400 ) - body, status = raw(AgentSkillInferToolsApi(), MagicMock(), _APP, "audio-transcribe") + body, status = raw(AgentSkillInferToolsApi(), unbound_session, _APP, "audio-transcribe") assert status == 400 assert body["code"] == "default_model_not_configured" -def test_infer_tools_rejects_path_like_slug_and_unbound_app(): +def test_infer_tools_rejects_path_like_slug_and_unbound_app(unbound_session: Session): from controllers.console.app.agent import AgentSkillInferToolsApi raw = _raw(AgentSkillInferToolsApi.post) with _json_ctx(): - body, status = raw(AgentSkillInferToolsApi(), MagicMock(), _APP, "a/b") + body, status = raw(AgentSkillInferToolsApi(), unbound_session, _APP, "a/b") assert (status, body["code"]) == (400, "drive_key_invalid") app_without_agent = SimpleNamespace(bound_agent_id_with_session=MagicMock(return_value=None)) with _json_ctx(): - body, status = raw(AgentSkillInferToolsApi(), MagicMock(), app_without_agent, "x") + body, status = raw(AgentSkillInferToolsApi(), unbound_session, app_without_agent, "x") assert (status, body["code"]) == (400, "agent_not_bound") diff --git a/api/tests/unit_tests/controllers/console/app/test_annotation_api.py b/api/tests/unit_tests/controllers/console/app/test_annotation_api.py index 33f19ae3550..028e5005583 100644 --- a/api/tests/unit_tests/controllers/console/app/test_annotation_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_annotation_api.py @@ -2,18 +2,33 @@ from __future__ import annotations from inspect import unwrap from types import SimpleNamespace -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.console.app import annotation as annotation_module +from models.model import App, AppMode, IconType from services.app_ref_service import AnnotationRef, AppRef -def _app_model() -> SimpleNamespace: - return SimpleNamespace(id="app-1", tenant_id="tenant-1", status="normal") +def _persist_app(session: Session) -> App: + app = App( + id="app-1", + tenant_id="tenant-1", + name="Annotation app", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#ffffff", + enable_site=False, + enable_api=True, + ) + session.add(app) + session.commit() + return app def _annotation_model(annotation_id: str = "ann-1") -> SimpleNamespace: @@ -109,27 +124,25 @@ def test_annotation_file_payload_valid(): assert payload.message_id == "550e8400-e29b-41d4-a716-446655440000" -def test_get_app_ref_raises_not_found_when_app_is_not_in_current_tenant(): - session = MagicMock() - session.scalar.return_value = None +def test_get_app_ref_raises_not_found_when_app_is_not_in_current_tenant(sqlite_session: Session): + _persist_app(sqlite_session) with ( patch.object( annotation_module, "current_account_with_tenant", - return_value=(SimpleNamespace(id="account-1"), "tenant-1"), + return_value=(SimpleNamespace(id="account-1"), "tenant-2"), ), ): with pytest.raises(NotFound): - annotation_module._get_app_ref(session, "app-1") + annotation_module._get_app_ref(sqlite_session, "app-1") class TestConsoleAnnotationRefBoundaries: - def test_batch_delete_uses_app_ref(self, app: Flask): + def test_batch_delete_uses_app_ref(self, app: Flask, sqlite_session: Session): api = annotation_module.AnnotationApi() handler = unwrap(api.delete) delete_mock = Mock() - session = MagicMock() - session.scalar.return_value = _app_model() + _persist_app(sqlite_session) with ( app.test_request_context("/?annotation_id=ann-1&annotation_id=ann-2", method="DELETE"), @@ -140,23 +153,20 @@ class TestConsoleAnnotationRefBoundaries: ), patch.object(annotation_module.AppAnnotationService, "delete_app_annotations_in_batch", delete_mock), ): - response, status = handler(api, session, "app-1") + response, status = handler(api, sqlite_session, "app-1") assert response == "" assert status == 204 - delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"], session) + delete_mock.assert_called_once_with(AppRef("tenant-1", "app-1"), ["ann-1", "ann-2"], sqlite_session) - def test_update_uses_annotation_ref(self, app: Flask): + def test_update_uses_annotation_ref(self, app: Flask, sqlite_session: Session): api = annotation_module.AnnotationUpdateDeleteApi() handler = unwrap(api.post) update_mock = Mock(return_value=_annotation_model()) - payload = {"question": "updated"} - session = MagicMock() - session.scalar.return_value = _app_model() + _persist_app(sqlite_session) with ( - app.test_request_context("/annotations/ann-1", method="POST", json=payload), - patch.object(type(annotation_module.console_ns), "payload", payload), + app.test_request_context("/annotations/ann-1", method="POST", json={"question": "updated"}), patch.object( annotation_module, "current_account_with_tenant", @@ -164,19 +174,24 @@ class TestConsoleAnnotationRefBoundaries: ), patch.object(annotation_module.AppAnnotationService, "update_app_annotation_directly", update_mock), ): - response = handler(api, session, "app-1", "ann-1") + response = handler( + api, + annotation_module.UpdateAnnotationPayload(question="updated"), + sqlite_session, + "app-1", + "ann-1", + ) assert response["question"] == "q" update_mock.assert_called_once() assert update_mock.call_args.args[1] == AnnotationRef(AppRef("tenant-1", "app-1"), "ann-1") - assert update_mock.call_args.args[2] is session + assert update_mock.call_args.args[2] is sqlite_session - def test_delete_uses_annotation_ref(self, app: Flask): + def test_delete_uses_annotation_ref(self, app: Flask, sqlite_session: Session): api = annotation_module.AnnotationUpdateDeleteApi() handler = unwrap(api.delete) delete_mock = Mock() - session = MagicMock() - session.scalar.return_value = _app_model() + _persist_app(sqlite_session) with ( app.test_request_context("/annotations/ann-1", method="DELETE"), @@ -187,15 +202,15 @@ class TestConsoleAnnotationRefBoundaries: ), patch.object(annotation_module.AppAnnotationService, "delete_app_annotation", delete_mock), ): - response, status = handler(api, session, "app-1", "ann-1") + response, status = handler(api, sqlite_session, "app-1", "ann-1") assert response == "" assert status == 204 delete_mock.assert_called_once() assert delete_mock.call_args.args[0] == AnnotationRef(AppRef("tenant-1", "app-1"), "ann-1") - assert delete_mock.call_args.args[1] is session + assert delete_mock.call_args.args[1] is sqlite_session - def test_hit_history_uses_annotation_ref(self, app: Flask): + def test_hit_history_uses_annotation_ref(self, app: Flask, sqlite_session: Session): api = annotation_module.AnnotationHitHistoryListApi() handler = unwrap(api.get) history = SimpleNamespace( @@ -208,8 +223,7 @@ class TestConsoleAnnotationRefBoundaries: created_at=None, ) hit_history_mock = Mock(return_value=([history], 1)) - session = MagicMock() - session.scalar.return_value = _app_model() + _persist_app(sqlite_session) with ( app.test_request_context("/hit-histories?page=2&limit=5", method="GET"), @@ -220,7 +234,9 @@ class TestConsoleAnnotationRefBoundaries: ), patch.object(annotation_module.AppAnnotationService, "get_annotation_hit_histories", hit_history_mock), ): - response = handler(api, session, "app-1", "ann-1") + response = handler(api, sqlite_session, "app-1", "ann-1") assert response["total"] == 1 - hit_history_mock.assert_called_once_with(AnnotationRef(AppRef("tenant-1", "app-1"), "ann-1"), 2, 5, session) + hit_history_mock.assert_called_once_with( + AnnotationRef(AppRef("tenant-1", "app-1"), "ann-1"), 2, 5, sqlite_session + ) diff --git a/api/tests/unit_tests/controllers/console/app/test_app_apis.py b/api/tests/unit_tests/controllers/console/app/test_app_apis.py index cbefd35b96e..165eaad2633 100644 --- a/api/tests/unit_tests/controllers/console/app/test_app_apis.py +++ b/api/tests/unit_tests/controllers/console/app/test_app_apis.py @@ -57,7 +57,11 @@ from controllers.console.app.ops_trace import TraceConfigPayload, TraceProviderQ from controllers.console.app.site import AppSiteUpdatePayload from controllers.console.app.workflow import AdvancedChatWorkflowRunPayload, SyncDraftWorkflowPayload from controllers.console.app.workflow_app_log import WorkflowAppLogQuery -from controllers.console.app.workflow_draft_variable import WorkflowDraftVariableUpdatePayload +from controllers.console.app.workflow_draft_variable import ( + EnvironmentVariableUpdatePayload, + WorkflowDraftVariableListQuery, + WorkflowDraftVariableUpdatePayload, +) from controllers.console.app.workflow_statistic import WorkflowStatisticQuery from controllers.console.app.workflow_trigger import Parser, ParserEnable from models import App, Site @@ -141,7 +145,13 @@ class TestCompletionEndpoints: Session(sqlite_engine) as session, app.test_request_context("/", json={"inputs": {}, "model_config": {}, "query": "hi"}), ): - resp = method(api, session, _make_account(), app_model=MagicMock(id=APP_ID)) + resp = method( + api, + CompletionMessagePayload(inputs={}, model_config={}, query="hi"), + session, + _make_account(), + app_model=MagicMock(id=APP_ID), + ) assert resp == {"result": {"text": "ok"}} @@ -164,7 +174,13 @@ class TestCompletionEndpoints: app.test_request_context("/", json={"inputs": {}, "model_config": {}, "query": "hi"}), pytest.raises(NotFound), ): - method(api, session, _make_account(), app_model=MagicMock(id=APP_ID)) + method( + api, + CompletionMessagePayload(inputs={}, model_config={}, query="hi"), + session, + _make_account(), + app_model=MagicMock(id=APP_ID), + ) def test_completion_api_provider_not_initialized( self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine @@ -183,7 +199,13 @@ class TestCompletionEndpoints: app.test_request_context("/", json={"inputs": {}, "model_config": {}, "query": "hi"}), pytest.raises(completion_module.ProviderNotInitializeError), ): - method(api, session, _make_account(), app_model=MagicMock(id=APP_ID)) + method( + api, + CompletionMessagePayload(inputs={}, model_config={}, query="hi"), + session, + _make_account(), + app_model=MagicMock(id=APP_ID), + ) def test_completion_api_quota_exceeded( self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine @@ -202,11 +224,19 @@ class TestCompletionEndpoints: app.test_request_context("/", json={"inputs": {}, "model_config": {}, "query": "hi"}), pytest.raises(completion_module.ProviderQuotaExceededError), ): - method(api, session, _make_account(), app_model=MagicMock(id=APP_ID)) + method( + api, + CompletionMessagePayload(inputs={}, model_config={}, query="hi"), + session, + _make_account(), + app_model=MagicMock(id=APP_ID), + ) class TestAppEndpoints: - def test_app_put_should_preserve_icon_type_when_payload_omits_it(self, app: Flask, monkeypatch: pytest.MonkeyPatch): + def test_app_put_should_preserve_icon_type_when_payload_omits_it( + self, app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): api = app_module.AppApi() method = unwrap(api.put) payload = { @@ -227,7 +257,17 @@ class TestAppEndpoints: app.test_request_context("/console/api/apps/app-1", method="PUT", json=payload), patch.object(type(console_ns), "payload", payload), ): - response = method(api, MagicMock(spec=Session), app_model=_make_app(icon_type=app_module.IconType.EMOJI)) + response = method( + api, + app_module.UpdateAppPayload( + name="Updated App", + description="Updated description", + icon="🤖", + icon_background="#FFFFFF", + ), + unbound_session, + app_model=_make_app(icon_type=app_module.IconType.EMOJI), + ) assert response == {"id": "app-1"} assert app_service.update_app.call_args.args[1]["icon_type"] is None @@ -244,7 +284,9 @@ class TestAppEndpoints: } ) - def test_app_icon_post_should_forward_icon_type(self, app: Flask, monkeypatch: pytest.MonkeyPatch): + def test_app_icon_post_should_forward_icon_type( + self, app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): api = app_module.AppIconApi() method = unwrap(api.post) payload = { @@ -264,7 +306,16 @@ class TestAppEndpoints: app.test_request_context("/console/api/apps/app-1/icon", method="POST", json=payload), patch.object(type(console_ns), "payload", payload), ): - response = method(api, MagicMock(spec=Session), app_model=_make_app()) + response = method( + api, + app_module.AppIconPayload( + icon="https://example.com/icon.png", + icon_type=app_module.IconType.IMAGE, + icon_background="#FFFFFF", + ), + unbound_session, + app_model=_make_app(), + ) assert response == {"id": "app-1"} assert app_service.update_app_icon.call_args.args[1:] == ( @@ -294,7 +345,7 @@ class TestOpsTraceEndpoints: ) with app.test_request_context("/?tracing_provider=langfuse"): - result = method(api, app_model=MagicMock(id="app-1")) + result = method(api, TraceProviderQuery(tracing_provider="langfuse"), MagicMock(id="app-1")) assert result == {"has_not_configured": True} @@ -313,7 +364,11 @@ class TestOpsTraceEndpoints: json={"tracing_provider": "langfuse", "tracing_config": {"api_key": "k"}}, ): with pytest.raises(BadRequest): - method(api, app_model=MagicMock(id="app-1")) + method( + api, + TraceConfigPayload(tracing_provider="langfuse", tracing_config={"api_key": "k"}), + MagicMock(id="app-1"), + ) def test_trace_app_config_delete_not_found(self, app: Flask, monkeypatch: pytest.MonkeyPatch): api = ops_trace_module.TraceAppConfigApi() @@ -327,7 +382,7 @@ class TestOpsTraceEndpoints: with app.test_request_context("/?tracing_provider=langfuse"): with pytest.raises(BadRequest): - method(api, app_model=MagicMock(id="app-1")) + method(api, TraceProviderQuery(tracing_provider="langfuse"), MagicMock(id="app-1")) class TestSiteEndpoints: @@ -367,7 +422,13 @@ class TestSiteEndpoints: site = self._add_site(db.session) with database_app.test_request_context("/", json={"title": "My Site", "input_placeholder": "Ask me anything"}): - result = method(api, db.session, _make_account(), app_model=_make_app()) + result = method( + api, + AppSiteUpdatePayload(title="My Site", input_placeholder="Ask me anything"), + db.session, + _make_account(), + app_model=_make_app(), + ) db.session.refresh(site) assert isinstance(result, dict) @@ -434,7 +495,7 @@ class TestWorkflowAppLogEndpoints: ) with database_app.test_request_context("/?page=1&limit=20"): - result = method(api, app_model=_make_app("app-1")) + result = method(api, WorkflowAppLogQuery(page=1, limit=20), app_model=_make_app("app-1")) assert result == {"page": 1, "limit": 20, "total": 0, "has_more": False, "data": []} @@ -464,10 +525,82 @@ class TestWorkflowDraftVariableEndpoints: monkeypatch.setattr(workflow_draft_variable_module, "WorkflowService", DummyWorkflowService) with database_app.test_request_context("/?page=1&limit=20"): - result = method(api, _make_account(), app_model=_make_app("app-1")) + result = method( + api, + WorkflowDraftVariableListQuery(page=1, limit=20), + _make_account(), + app_model=_make_app("app-1"), + ) assert result == {"items": [], "total": 0} + def test_environment_variable_update_payload_preserves_full_replace_default(self) -> None: + payload = EnvironmentVariableUpdatePayload(environment_variables=[]) + + assert payload.patch is False + assert payload.deleted_environment_variable_ids == [] + + @pytest.mark.parametrize( + "payload", + [ + {"environment_variables": [], "deleted_environment_variable_ids": ["env-a"]}, + { + "environment_variables": [{"name": "a", "value_type": "string", "value": "a"}], + "patch": True, + }, + { + "environment_variables": [{"id": "env-a", "name": "a", "value_type": "string", "value": "a"}], + "patch": True, + "deleted_environment_variable_ids": ["env-a"], + }, + ], + ) + def test_environment_variable_patch_payload_rejects_ambiguous_mutations(self, payload: dict) -> None: + with pytest.raises(ValidationError): + EnvironmentVariableUpdatePayload.model_validate(payload) + + def test_environment_variable_collection_post_routes_patch_to_service( + self, database_app: Flask, monkeypatch: pytest.MonkeyPatch + ) -> None: + api = workflow_draft_variable_module.EnvironmentVariableCollectionApi() + method = unwrap(api.post) + captured: dict = {} + + class DummyWorkflowService: + def patch_draft_workflow_environment_variables(self, **kwargs) -> None: + captured.update(kwargs) + + def update_draft_workflow_environment_variables(self, **_kwargs) -> None: + raise AssertionError("patch request must not use full replacement") + + monkeypatch.setattr(workflow_draft_variable_module, "WorkflowService", DummyWorkflowService) + + with database_app.test_request_context( + "/", + json={ + "environment_variables": [{"id": "env-a", "name": "a", "value_type": "string", "value": "new-a"}], + "patch": True, + "deleted_environment_variable_ids": ["env-b"], + }, + ): + result = method( + api, + EnvironmentVariableUpdatePayload( + environment_variables=[{"id": "env-a", "name": "a", "value_type": "string", "value": "new-a"}], + patch=True, + deleted_environment_variable_ids=["env-b"], + ), + _make_account(), + app_model=_make_app(), + ) + + assert result == {"result": "success"} + assert [(variable.id, variable.value) for variable in captured["environment_variables"]] == [("env-a", "new-a")] + assert captured["deleted_environment_variable_ids"] == ["env-b"] + assert captured["app_model"].id == APP_ID + assert captured["account"].id == USER_ID + assert captured["session"].get_bind() is db.engine + class TestWorkflowStatisticEndpoints: def test_workflow_statistic_time_range(self): @@ -505,7 +638,7 @@ class TestWorkflowStatisticEndpoints: with database_app.test_request_context("/"): account = _make_account() account.timezone = "UTC" - response = method(api, account, app_model=_make_app("app-1", tenant_id="t1")) + response = method(api, WorkflowStatisticQuery(), account, app_model=_make_app("app-1", tenant_id="t1")) assert response.get_json() == {"data": [{"date": "2024-01-01"}]} @@ -535,7 +668,7 @@ class TestWorkflowStatisticEndpoints: with database_app.test_request_context("/"): account = _make_account() account.timezone = "UTC" - response = method(api, account, app_model=_make_app("app-1", tenant_id="t1")) + response = method(api, WorkflowStatisticQuery(), account, app_model=_make_app("app-1", tenant_id="t1")) assert response.get_json() == {"data": [{"date": "2024-01-02"}]} @@ -565,7 +698,7 @@ class TestWorkflowTriggerEndpoints: db.session.commit() with database_app.test_request_context("/?node_id=node-1"): - result = method(api, app_model=_make_app()) + result = method(api, Parser(node_id="node-1"), app_model=_make_app()) assert isinstance(result, dict) assert {"id", "webhook_id", "webhook_url", "webhook_debug_url", "node_id", "created_at"} <= set(result.keys()) diff --git a/api/tests/unit_tests/controllers/console/app/test_app_import_api.py b/api/tests/unit_tests/controllers/console/app/test_app_import_api.py index e60656201ff..b53fa4676dd 100644 --- a/api/tests/unit_tests/controllers/console/app/test_app_import_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_app_import_api.py @@ -13,14 +13,14 @@ from sqlalchemy import Engine, event from sqlalchemy.orm import Session from controllers.console.app import app_import as app_import_module -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from models.account import Account from models.base import TypeBase from models.engine import db from models.model import App, AppMode from services.app_dsl_service import ImportStatus from services.entities.dsl_entities import CheckDependenciesResult -from services.feature_service import SystemFeatureModel, WebAppAuthModel +from services.entities.feature_entities import SystemFeatureModel, WebAppAuthModel def _unwrap(func): @@ -58,6 +58,7 @@ def _install_features(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None: def _make_account(account_id: str = "u1") -> Account: account = Account(name="Test User", email="test@example.com") account.id = account_id + account._current_tenant = MagicMock(id="tenant-1") return account @@ -158,7 +159,7 @@ class TestAppImportApi: ) with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}): - response, status = method(api, _make_account()) + response, status = method(api, app_import_module.AppImportPayload(mode="yaml-content"), _make_account()) assert transaction_events.rollbacks == 1 assert transaction_events.commits == 0 @@ -184,7 +185,7 @@ class TestAppImportApi: ) with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}): - response, status = method(api, _make_account()) + response, status = method(api, app_import_module.AppImportPayload(mode="yaml-content"), _make_account()) assert transaction_events.commits == 1 assert transaction_events.rollbacks == 0 @@ -212,7 +213,7 @@ class TestAppImportApi: monkeypatch.setattr(app_import_module.EnterpriseService.WebAppAuth, "update_app_access_mode", update_access) with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}): - response, status = method(api, _make_account()) + response, status = method(api, app_import_module.AppImportPayload(mode="yaml-content"), _make_account()) assert transaction_events.commits == 1 assert transaction_events.rollbacks == 0 @@ -250,7 +251,7 @@ class TestAppImportApi: ) with app.test_request_context("/console/api/apps/imports", method="POST", json={"mode": "yaml-content"}): - response, status = method() + response, status = method(app_import_module.AppImportPayload(mode="yaml-content")) assert transaction_events.commits == 1 _assert_app_persistence(sqlite_app_engine, app_id, persisted=True) @@ -290,7 +291,7 @@ class TestAppImportApi: method="POST", json={"mode": "yaml-content", "app_id": "existing-app"}, ): - response, status = method() + response, status = method(app_import_module.AppImportPayload(mode="yaml-content", app_id="existing-app")) assert transaction_events.commits == 1 _assert_app_persistence(sqlite_app_engine, app_id, persisted=True) @@ -343,14 +344,13 @@ class TestAppImportConfirmApi: "current_account_with_tenant", lambda: (_make_account(), "tenant-1"), ) - monkeypatch.setattr( - app_import_module.redis_client, - "get", - lambda *_args, **_kwargs: ( + redis_get = MagicMock( + return_value=( b'{"import_mode":"yaml-content","yaml_content":"app: {}","app_id":null,' b'"name":null,"description":null,"icon_type":null,"icon":null,"icon_background":null}' - ), + ) ) + monkeypatch.setattr(app_import_module.redis_client, "get", redis_get) monkeypatch.setattr(app_import_module.dify_config, "RBAC_ENABLED", True) app_id = _install_persisting_service_result( monkeypatch, @@ -370,6 +370,7 @@ class TestAppImportConfirmApi: _assert_app_persistence(sqlite_app_engine, app_id, persisted=True) assert status == 200 assert response["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"] + redis_get.assert_called_once_with("app_import_info:import-1") def test_import_confirm_does_not_attach_permission_keys_when_overwriting_existing_app( self, diff --git a/api/tests/unit_tests/controllers/console/app/test_app_response_models.py b/api/tests/unit_tests/controllers/console/app/test_app_response_models.py index 6b1dce534ed..bf953ff3b4f 100644 --- a/api/tests/unit_tests/controllers/console/app/test_app_response_models.py +++ b/api/tests/unit_tests/controllers/console/app/test_app_response_models.py @@ -1,6 +1,7 @@ from __future__ import annotations import builtins +import json import sys from datetime import datetime from importlib import util @@ -12,9 +13,13 @@ import pytest from flask import Flask from flask.views import MethodView from pydantic import ValidationError +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from werkzeug.datastructures import MultiDict from configs import dify_config +from models.model import App, AppMode, IconType +from models.workflow import Workflow, WorkflowType # kombu references MethodView as a global when importing celery/kombu pools. if not hasattr(builtins, "MethodView"): @@ -326,19 +331,9 @@ def test_app_list_query_accepts_single_repeated_tag_id(app_module): assert query.tag_ids == [tag_id] -def test_create_app_endpoint_rejects_agent_mode(app_module, monkeypatch: pytest.MonkeyPatch): - payload = {"name": "Iris", "mode": "agent", "description": "Agent app"} - app_service = MagicMock() - monkeypatch.setattr(app_module, "AppService", lambda: app_service) - - app_module.console_ns.payload = payload - try: - with pytest.raises(ValidationError): - _unwrap(app_module.AppListApi().post)(MagicMock(), "tenant-1", SimpleNamespace(id="account-1")) - finally: - app_module.console_ns.payload = None - - app_service.create_app.assert_not_called() +def test_create_app_endpoint_rejects_agent_mode(app_module): + with pytest.raises(ValidationError): + app_module.CreateAppPayload.model_validate({"name": "Iris", "mode": "agent", "description": "Agent app"}) def test_app_partial_serialization_uses_aliases(app_models): @@ -439,8 +434,9 @@ def test_app_detail_with_site_includes_nested_serialization(app_models): assert "role" not in serialized -def test_app_response_view_uses_the_caller_session_for_query_backed_fields(app_module, monkeypatch): - session = MagicMock() +def test_app_response_view_uses_the_caller_session_for_query_backed_fields( + app_module, monkeypatch, unbound_session: Session +): app_obj = MagicMock() app_model_config = SimpleNamespace(app_id="app-1") app_obj.desc_or_prompt_with_session.return_value = "Description" @@ -455,7 +451,7 @@ def test_app_response_view_uses_the_caller_session_for_query_backed_fields(app_m load_annotation_reply = MagicMock(return_value={"enabled": False}) monkeypatch.setattr("services.app_service.load_annotation_reply_config", load_annotation_reply) - view = app_module.AppResponseView(app_obj, session=session) + view = app_module.AppResponseView(app_obj, session=unbound_session) site = view.site workflow = view.workflow model_config = view.app_model_config @@ -483,8 +479,8 @@ def test_app_response_view_uses_the_caller_session_for_query_backed_fields(app_m app_obj.tags_with_session, app_obj.author_name_with_session, ): - method.assert_called_once_with(session=session) - load_annotation_reply.assert_called_once_with(session, "app-1") + method.assert_called_once_with(session=unbound_session) + load_annotation_reply.assert_called_once_with(unbound_session, "app-1") def test_app_pagination_aliases_per_page_and_has_next(app_models): @@ -529,7 +525,11 @@ def test_app_pagination_aliases_per_page_and_has_next(app_models): def test_app_list_uses_injected_session_for_draft_workflows( - app: Flask, app_module: ModuleType, monkeypatch: pytest.MonkeyPatch + app: Flask, + app_module: ModuleType, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + unbound_session: Session, ) -> None: api = app_module.AppListApi() method = _unwrap(api.get) @@ -541,15 +541,20 @@ def test_app_list_uses_injected_session_for_draft_workflows( mode_compatible_with_agent="workflow", ) app_pagination = SimpleNamespace(page=1, per_page=20, total=1, has_next=False, items=[app_item]) - workflow = SimpleNamespace( + workflow = Workflow( id="workflow-1", + tenant_id="tenant-1", app_id="app-1", - walk_nodes=lambda: iter([("trigger-1", {"type": "trigger-webhook"})]), + type=WorkflowType.WORKFLOW, + version=Workflow.VERSION_DRAFT, + graph=json.dumps({"nodes": [{"id": "trigger-1", "data": {"type": "trigger-webhook"}}], "edges": []}), + features=json.dumps({}), + created_by="user-1", + environment_variables=[], + conversation_variables=[], ) - session = MagicMock() - session.execute.return_value.scalars.return_value.all.return_value = [workflow] - scoped_session = MagicMock() - scoped_session.execute.side_effect = AssertionError("db.session should not be used") + sqlite_session.add(workflow) + sqlite_session.commit() monkeypatch.setattr( app_module, @@ -578,20 +583,18 @@ def test_app_list_uses_injected_session_for_draft_workflows( "get", get_permissions, ) - monkeypatch.setattr(app_module, "db", SimpleNamespace(session=scoped_session)) + monkeypatch.setattr(app_module, "db", SimpleNamespace(session=unbound_session)) with app.test_request_context("/console/api/apps?page=1&limit=20", method="GET"): - response, status = method("tenant-1", "user-1", session) + response, status = method("tenant-1", "user-1", sqlite_session) assert status == 200 assert response["data"][0]["has_draft_trigger"] is True - session.execute.assert_called_once() - scoped_session.execute.assert_not_called() - get_permissions.assert_called_once_with("tenant-1", "user-1", session=session) + get_permissions.assert_called_once_with("tenant-1", "user-1", session=sqlite_session) assert response["data"][0]["permission_keys"] == ["app.acl.edit"] -def test_app_create_api_attaches_permission_keys(app, app_module): +def test_app_create_api_attaches_permission_keys(app, app_module, unbound_session: Session): method = app_module.AppListApi.post while hasattr(method, "__wrapped__"): method = method.__wrapped__ @@ -637,7 +640,17 @@ def test_app_create_api_attaches_permission_keys(app, app_module): replace_whitelist, ) - resp, status = method(app_module.AppListApi(), MagicMock(), "tenant-1", SimpleNamespace(id="acct-1")) + resp, status = method( + app_module.AppListApi(), + app_module.CreateAppPayload( + name="Created App", + description="Summary", + mode="advanced-chat", + ), + unbound_session, + "tenant-1", + SimpleNamespace(id="acct-1"), + ) assert status == 201 assert resp["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"] @@ -645,7 +658,7 @@ def test_app_create_api_attaches_permission_keys(app, app_module): initialize_rbac_task.delay.assert_called_once_with("tenant-1", "acct-1", app_id="app-new") -def test_app_list_api_attaches_permission_keys(app, app_module): +def test_app_list_api_attaches_permission_keys(app, app_module, sqlite_session: Session): method = app_module.AppListApi.get while hasattr(method, "__wrapped__"): method = method.__wrapped__ @@ -697,9 +710,7 @@ def test_app_list_api_attaches_permission_keys(app, app_module): lambda tenant_id, account_id: SimpleNamespace(unrestricted=True, resource_ids=[]), ) - session = MagicMock() - session.execute.return_value.scalars.return_value.all.return_value = [] - resp, status = method(app_module.AppListApi(), "tenant-1", "acct-1", session) + resp, status = method(app_module.AppListApi(), "tenant-1", "acct-1", sqlite_session) assert status == 200 params = get_paginate_apps.call_args.args[2] @@ -708,7 +719,124 @@ def test_app_list_api_attaches_permission_keys(app, app_module): assert resp["data"][0]["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"] -def test_app_list_api_limits_to_apps_created_by_current_user_without_view_permission(app, app_module): +def test_recent_app_list_api_returns_only_home_card_fields(app, app_module, unbound_session: Session): + method = app_module.RecentAppListApi.get + while hasattr(method, "__wrapped__"): + method = method.__wrapped__ + + recent_app = SimpleNamespace( + id="app-1", + name="Recent App", + icon_type="emoji", + icon="🚀", + icon_background="#FFFFFF", + mode="chat", + author_name="Recent Author", + updated_at=_ts(15), + maintainer="acct-1", + ) + get_recent_apps = MagicMock(return_value=[recent_app]) + + with app.test_request_context("/apps/recent?limit=8"): + with pytest.MonkeyPatch.context() as monkeypatch: + monkeypatch.setattr(dify_config, "RBAC_ENABLED", False) + monkeypatch.setattr(app_module.AppService, "get_recent_apps", get_recent_apps) + monkeypatch.setattr( + app_module.enterprise_rbac_service.RBACService.MyPermissions, + "get", + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( + app=app_module.enterprise_rbac_service.ResourcePermissionSnapshot( + overrides=[ + app_module.enterprise_rbac_service.ResourcePermissionKeys( + resource_id="app-1", + permission_keys=["app.acl.monitor"], + ) + ] + ) + ), + ) + + resp, status = method(app_module.RecentAppListApi(), "tenant-1", "acct-1", unbound_session) + + assert status == 200 + assert resp == { + "data": [ + { + "id": "app-1", + "name": "Recent App", + "icon_type": "emoji", + "icon": "🚀", + "icon_background": "#FFFFFF", + "mode": "chat", + "author_name": "Recent Author", + "updated_at": int(_ts(15).timestamp()), + "permission_keys": ["app.acl.monitor"], + "maintainer": "acct-1", + "icon_url": None, + } + ] + } + params = get_recent_apps.call_args.args[2] + assert params.limit == 8 + assert "total" not in resp + assert "description" not in resp["data"][0] + assert "tags" not in resp["data"][0] + assert "workflow" not in resp["data"][0] + + +@pytest.mark.parametrize("mode", ["channel", "rag-pipeline", "agent"]) +def test_recent_app_response_rejects_non_home_app_modes(app_module, mode: str) -> None: + with pytest.raises(ValidationError): + app_module.RecentAppResponse.model_validate( + { + "id": "app-1", + "name": "Recent App", + "mode": mode, + "updated_at": _ts(), + } + ) + + +def test_recent_app_list_api_applies_rbac_visibility_filter(app, app_module, unbound_session: Session): + method = app_module.RecentAppListApi.get + while hasattr(method, "__wrapped__"): + method = method.__wrapped__ + + get_recent_apps = MagicMock(return_value=[]) + with app.test_request_context("/apps/recent"): + with pytest.MonkeyPatch.context() as monkeypatch: + monkeypatch.setattr(dify_config, "RBAC_ENABLED", True) + monkeypatch.setattr(app_module.AppService, "get_recent_apps", get_recent_apps) + monkeypatch.setattr( + app_module.enterprise_rbac_service.RBACService.MyPermissions, + "get", + lambda tenant_id, account_id, session: app_module.enterprise_rbac_service.MyPermissionsResponse( + workspace=app_module.enterprise_rbac_service.WorkspacePermissionSnapshot( + permission_keys=["app.create_and_management"] + ) + ), + ) + monkeypatch.setattr( + app_module.enterprise_rbac_service.RBACService.AppAccess, + "whitelist_resources", + lambda tenant_id, account_id: SimpleNamespace( + unrestricted=False, + resource_ids=["app-shared"], + ), + ) + + resp, status = method(app_module.RecentAppListApi(), "tenant-1", "acct-1", unbound_session) + + assert status == 200 + assert resp == {"data": []} + params = get_recent_apps.call_args.args[2] + assert params.accessible_app_ids == ["app-shared"] + assert params.include_own_apps is True + + +def test_app_list_api_limits_to_apps_created_by_current_user_without_view_permission( + app, app_module, unbound_session: Session +): method = app_module.AppListApi.get while hasattr(method, "__wrapped__"): method = method.__wrapped__ @@ -740,8 +868,7 @@ def test_app_list_api_limits_to_apps_created_by_current_user_without_view_permis lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ) - session = MagicMock() - resp, status = method(app_module.AppListApi(), "tenant-1", "acct-1", session) + resp, status = method(app_module.AppListApi(), "tenant-1", "acct-1", unbound_session) assert status == 200 assert resp["data"] == [] @@ -751,7 +878,9 @@ def test_app_list_api_limits_to_apps_created_by_current_user_without_view_permis assert params.is_created_by_me is None -def test_app_list_api_limits_to_preview_overrides_without_manage_own_permission(app, app_module): +def test_app_list_api_limits_to_preview_overrides_without_manage_own_permission( + app, app_module, unbound_session: Session +): method = app_module.AppListApi.get while hasattr(method, "__wrapped__"): method = method.__wrapped__ @@ -798,8 +927,7 @@ def test_app_list_api_limits_to_preview_overrides_without_manage_own_permission( lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ) - session = MagicMock() - method(app_module.AppListApi(), "tenant-1", "acct-1", session) + method(app_module.AppListApi(), "tenant-1", "acct-1", unbound_session) params = get_paginate_apps.call_args.args[2] assert params.accessible_app_ids == ["app-acl-shared", "app-full", "app-shared", "app-whitelist-only"] @@ -807,7 +935,9 @@ def test_app_list_api_limits_to_preview_overrides_without_manage_own_permission( assert params.is_created_by_me is None -def test_app_list_api_returns_no_apps_without_workspace_or_resource_view_permission(app, app_module): +def test_app_list_api_returns_no_apps_without_workspace_or_resource_view_permission( + app, app_module, unbound_session: Session +): method = app_module.AppListApi.get while hasattr(method, "__wrapped__"): method = method.__wrapped__ @@ -835,8 +965,7 @@ def test_app_list_api_returns_no_apps_without_workspace_or_resource_view_permiss lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ) - session = MagicMock() - method(app_module.AppListApi(), "tenant-1", "acct-1", session) + method(app_module.AppListApi(), "tenant-1", "acct-1", unbound_session) params = get_paginate_apps.call_args.args[2] assert params.accessible_app_ids == ["app-not-permitted"] @@ -844,7 +973,7 @@ def test_app_list_api_returns_no_apps_without_workspace_or_resource_view_permiss assert params.is_created_by_me is None -def test_app_detail_api_attaches_current_user_permission_keys(app, app_module): +def test_app_detail_api_attaches_current_user_permission_keys(app, app_module, unbound_session: Session): method = app_module.AppApi.get while hasattr(method, "__wrapped__"): method = method.__wrapped__ @@ -875,7 +1004,12 @@ def test_app_detail_api_attaches_current_user_permission_keys(app, app_module): overrides=[ app_module.enterprise_rbac_service.ResourcePermissionKeys( resource_id="app-1", - permission_keys=["app.acl.view_layout", "app.acl.edit", "app.acl.monitor"], + permission_keys=[ + "app.acl.view_layout", + "app.acl.edit", + "app.acl.deploy", + "app.acl.monitor", + ], ) ] ) @@ -887,40 +1021,45 @@ def test_app_detail_api_attaches_current_user_permission_keys(app, app_module): get_permissions, ) - session = MagicMock() resp = method( app_module.AppApi(), - session, + unbound_session, "tenant-1", SimpleNamespace(id="acct-1"), app_model=app_obj, ) - get_app.assert_called_once_with(app_obj, session=session) - get_permissions.assert_called_once_with("tenant-1", "acct-1", app_id="app-1", session=session) - assert resp["permission_keys"] == ["app.acl.view_layout", "app.acl.edit", "app.acl.monitor"] + get_app.assert_called_once_with(app_obj, session=unbound_session) + get_permissions.assert_called_once_with("tenant-1", "acct-1", app_id="app-1", session=unbound_session) + assert resp["permission_keys"] == [ + "app.acl.view_layout", + "app.acl.edit", + "app.acl.deploy", + "app.acl.monitor", + ] -def test_app_copy_api_attaches_permission_keys(app, app_module): +def test_app_copy_api_attaches_permission_keys(app, app_module, sqlite_session: Session, sqlite_engine: Engine): method = app_module.AppCopyApi.post while hasattr(method, "__wrapped__"): method = method.__wrapped__ - app_obj = SimpleNamespace( - id="app-new", + app_obj = App( + id="00000000-0000-0000-0000-000000000101", + tenant_id="00000000-0000-0000-0000-000000000102", name="Copied App", description="Summary", - mode_compatible_with_agent="workflow", + mode=AppMode.WORKFLOW, + icon_type=IconType.EMOJI, + icon="copy", + icon_background="#ffffff", enable_site=True, enable_api=True, - permission_keys=[], ) + sqlite_session.add(app_obj) + sqlite_session.commit() - import_result = SimpleNamespace(status=app_module.ImportStatus.COMPLETED, app_id="app-new") - fake_session = MagicMock() - fake_session.__enter__.return_value = fake_session - fake_session.__exit__.return_value = None - fake_session.scalar.return_value = app_obj + import_result = SimpleNamespace(status=app_module.ImportStatus.COMPLETED, app_id=app_obj.id) with app.test_request_context("/apps/app-original/copy", method="POST", json={}): with pytest.MonkeyPatch.context() as monkeypatch: @@ -938,25 +1077,21 @@ def test_app_copy_api_attaches_permission_keys(app, app_module): "get_system_features", lambda: SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ) - monkeypatch.setattr(app_module, "db", SimpleNamespace(engine=object(), session=lambda: MagicMock())) - monkeypatch.setattr( - app_module, - "Session", - lambda *_args, **_kwargs: fake_session, - ) + monkeypatch.setattr(app_module, "db", SimpleNamespace(engine=sqlite_engine)) monkeypatch.setattr( app_module.enterprise_rbac_service.RBACService.AppPermissions, "batch_get", - lambda tenant_id, account_id, app_ids, session: {"app-new": ["app.acl.view_layout", "app.acl.edit"]}, + lambda tenant_id, account_id, app_ids, session: {app_obj.id: ["app.acl.view_layout", "app.acl.edit"]}, ) resp, status = method( app_module.AppCopyApi(), + app_module.CopyAppPayload(), "tenant-1", SimpleNamespace(id="acct-1"), app_model=SimpleNamespace(id="app-original"), ) assert status == 201 - assert fake_session.scalar.called + assert sqlite_session.get(App, app_obj.id) is not None assert resp["permission_keys"] == ["app.acl.view_layout", "app.acl.edit"] diff --git a/api/tests/unit_tests/controllers/console/app/test_audio.py b/api/tests/unit_tests/controllers/console/app/test_audio.py index 8b661aa2646..347627809ff 100644 --- a/api/tests/unit_tests/controllers/console/app/test_audio.py +++ b/api/tests/unit_tests/controllers/console/app/test_audio.py @@ -17,6 +17,8 @@ from controllers.console.app.audio import ( ChatMessageAudioApi, ChatMessageTextApi, TextModesApi, + TextToSpeechPayload, + TextToSpeechVoiceQuery, ) from controllers.console.app.error import ( AppUnavailableError, @@ -290,7 +292,7 @@ def test_console_text_api_success(app: Flask, monkeypatch: pytest.MonkeyPatch) - method="POST", json={"text": "hello", "voice": "v"}, ): - response = handler(api, app_model=app_model) + response = handler(api, TextToSpeechPayload(text="hello"), app_model=app_model) assert response == {"audio": "ok"} @@ -315,7 +317,7 @@ def test_console_text_api_builds_message_ref(app: Flask, monkeypatch: pytest.Mon ), patch("controllers.console.app.audio.current_user", SimpleNamespace(id="account-1")), ): - response = handler(api, app_model=app_model) + response = handler(api, TextToSpeechPayload(text="hello", message_id="message-1"), app_model=app_model) assert response == {"audio": "ok"} assert calls["message_ref"] == MessageRef(AppRef("tenant-1", "app-1"), "message-1", account_id="account-1") @@ -334,7 +336,7 @@ def test_console_text_api_error_mapping(app: Flask, monkeypatch: pytest.MonkeyPa json={"text": "hello"}, ): with pytest.raises(ProviderQuotaExceededError): - handler(api, app_model=app_model) + handler(api, TextToSpeechPayload(text="hello"), app_model=app_model) def test_console_text_modes_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -346,7 +348,7 @@ def test_console_text_modes_success(app: Flask, monkeypatch: pytest.MonkeyPatch) app_model = SimpleNamespace(tenant_id="t1") with app.test_request_context("/console/api/apps/app/text-to-audio/voices?language=en", method="GET"): - response = handler(api, app_model=app_model) + response = handler(api, TextToSpeechVoiceQuery(language="en-US"), app_model=app_model) assert response == expected_voices @@ -364,7 +366,7 @@ def test_console_text_modes_language_error(app: Flask, monkeypatch: pytest.Monke with app.test_request_context("/console/api/apps/app/text-to-audio/voices?language=en", method="GET"): with pytest.raises(AppUnavailableError): - handler(api, app_model=app_model) + handler(api, TextToSpeechVoiceQuery(language="en-US"), app_model=app_model) def test_audio_to_text_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -424,7 +426,7 @@ def test_text_to_audio_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -> N method="POST", json={"text": "hello"}, ): - response = method(api, app_model=app_model) + response = method(api, TextToSpeechPayload(text="hello"), app_model=app_model) assert response == {"audio": "ok"} @@ -443,7 +445,7 @@ def test_text_to_audio_voices_success(app: Flask, monkeypatch: pytest.MonkeyPatc method="GET", query_string={"language": "en-US"}, ): - response = method(api, app_model=app_model) + response = method(api, TextToSpeechVoiceQuery(language="en-US"), app_model=app_model) assert response == expected_voices @@ -481,7 +483,7 @@ def test_text_to_audio_with_language_param(app: Flask, monkeypatch: pytest.Monke method="POST", json={"text": "hello", "language": "en-US"}, ): - response = method(api, app_model=app_model) + response = method(api, TextToSpeechPayload(text="hello"), app_model=app_model) assert response == {"audio": "test"} @@ -501,5 +503,5 @@ def test_text_to_audio_voices_with_language_filter(app: Flask, monkeypatch: pyte "/console/api/apps/app-1/text-to-audio/voices?language=en-US", method="GET", ): - response = method(api, app_model=app_model) + response = method(api, TextToSpeechVoiceQuery(language="en-US"), app_model=app_model) assert isinstance(response, list) diff --git a/api/tests/unit_tests/controllers/console/app/test_conversation_api.py b/api/tests/unit_tests/controllers/console/app/test_conversation_api.py index 1f2819e6d39..01cc0c1edab 100644 --- a/api/tests/unit_tests/controllers/console/app/test_conversation_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_conversation_api.py @@ -6,10 +6,13 @@ from unittest.mock import MagicMock import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest, NotFound from controllers.console.app import conversation as conversation_module -from models.model import AppMode +from core.app.entities.app_invoke_entities import InvokeFrom +from models.enums import ConversationFromSource +from models.model import AppMode, Conversation from services.errors.conversation import ConversationNotExistsError @@ -17,7 +20,32 @@ def _make_account(): return SimpleNamespace(timezone="UTC", id="u1") -def test_completion_conversation_list_returns_paginated_result(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def _conversation(*, conversation_id: str = "c1", app_id: str = "app-1") -> Conversation: + conversation = Conversation( + app_id=app_id, + app_model_config_id=None, + model_provider=None, + override_model_configs=None, + model_id=None, + mode=AppMode.CHAT, + name="Conversation", + inputs={}, + introduction="", + system_instruction="", + system_instruction_tokens=0, + status="normal", + invoke_from=InvokeFrom.EXPLORE, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id=None, + from_account_id="u1", + ) + conversation.id = conversation_id + return conversation + + +def test_completion_conversation_list_returns_paginated_result( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: api = conversation_module.CompletionConversationApi() method = unwrap(api.get) account = _make_account() @@ -30,11 +58,19 @@ def test_completion_conversation_list_returns_paginated_result(app: Flask, monke paginate_result.items = [] monkeypatch.setattr(conversation_module, "paginate_query", lambda *_args, **_kwargs: paginate_result) with app.test_request_context("/console/api/apps/app-1/completion-conversations", method="GET"): - response = method(api, MagicMock(), account, app_model=SimpleNamespace(id="app-1")) + response = method( + api, + conversation_module.CompletionConversationQuery(), + unbound_session, + account, + app_model=SimpleNamespace(id="app-1"), + ) assert response == {"page": 1, "limit": 20, "total": 0, "has_more": False, "data": []} -def test_completion_conversation_list_invalid_time_range(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_completion_conversation_list_invalid_time_range( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: api = conversation_module.CompletionConversationApi() method = unwrap(api.get) account = _make_account() @@ -47,10 +83,18 @@ def test_completion_conversation_list_invalid_time_range(app: Flask, monkeypatch "/console/api/apps/app-1/completion-conversations", method="GET", query_string={"start": "bad"} ): with pytest.raises(BadRequest): - method(api, MagicMock(), account, app_model=SimpleNamespace(id="app-1")) + method( + api, + conversation_module.CompletionConversationQuery(), + unbound_session, + account, + app_model=SimpleNamespace(id="app-1"), + ) -def test_chat_conversation_list_advanced_chat_calls_paginate(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_chat_conversation_list_advanced_chat_calls_paginate( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: api = conversation_module.ChatConversationApi() method = unwrap(api.get) account = _make_account() @@ -63,30 +107,35 @@ def test_chat_conversation_list_advanced_chat_calls_paginate(app: Flask, monkeyp paginate_result.items = [] monkeypatch.setattr(conversation_module, "paginate_query", lambda *_args, **_kwargs: paginate_result) with app.test_request_context("/console/api/apps/app-1/chat-conversations", method="GET"): - response = method(api, MagicMock(), account, app_model=SimpleNamespace(id="app-1", mode=AppMode.ADVANCED_CHAT)) + response = method( + api, + conversation_module.ChatConversationQuery(), + unbound_session, + account, + app_model=SimpleNamespace(id="app-1", mode=AppMode.ADVANCED_CHAT), + ) assert response == {"page": 1, "limit": 20, "total": 0, "has_more": False, "data": []} -def test_get_conversation_updates_read_at(monkeypatch: pytest.MonkeyPatch) -> None: - conversation = SimpleNamespace(id="c1", app_id="app-1") - session = MagicMock() - session.scalar.return_value = conversation +def test_get_conversation_updates_read_at(sqlite_session: Session) -> None: + conversation = _conversation() + sqlite_session.add(conversation) + sqlite_session.flush() + session = sqlite_session result = conversation_module._get_conversation(session, _make_account(), SimpleNamespace(id="app-1"), "c1") assert result is conversation - session.execute.assert_called_once() - session.flush.assert_called_once() - session.refresh.assert_called_once_with(conversation) + assert conversation.read_at is not None + assert conversation.read_account_id == "u1" -def test_get_conversation_missing_raises_not_found(monkeypatch: pytest.MonkeyPatch) -> None: - session = MagicMock() - session.scalar.return_value = None +def test_get_conversation_missing_raises_not_found(sqlite_session: Session) -> None: + session = sqlite_session with pytest.raises(NotFound): conversation_module._get_conversation(session, _make_account(), SimpleNamespace(id="app-1"), "missing") -def test_conversation_response_source_uses_caller_session() -> None: - session = MagicMock() +def test_conversation_response_source_uses_caller_session(unbound_session: Session) -> None: + session = unbound_session account = object() annotation = MagicMock() annotation.account_with_session.return_value = account @@ -136,7 +185,9 @@ def test_conversation_response_source_uses_caller_session() -> None: annotation.account_with_session.assert_called_once_with(session=session) -def test_completion_conversation_delete_maps_not_found(monkeypatch: pytest.MonkeyPatch) -> None: +def test_completion_conversation_delete_maps_not_found( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: api = conversation_module.CompletionConversationDetailApi() method = unwrap(api.delete) monkeypatch.setattr( @@ -144,6 +195,6 @@ def test_completion_conversation_delete_maps_not_found(monkeypatch: pytest.Monke "delete", lambda *_args, **_kwargs: (_ for _ in ()).throw(ConversationNotExistsError()), ) - session = MagicMock() + session = unbound_session with pytest.raises(NotFound): method(api, session, _make_account(), app_model=SimpleNamespace(id="app-1"), conversation_id="c1") diff --git a/api/tests/unit_tests/controllers/console/app/test_conversation_variables_api.py b/api/tests/unit_tests/controllers/console/app/test_conversation_variables_api.py index ab3eacd03c7..f44358e6897 100644 --- a/api/tests/unit_tests/controllers/console/app/test_conversation_variables_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_conversation_variables_api.py @@ -1,6 +1,5 @@ from __future__ import annotations -from contextlib import nullcontext from datetime import UTC, datetime from inspect import unwrap from types import SimpleNamespace @@ -8,87 +7,98 @@ from types import SimpleNamespace import pytest from flask import Flask from pydantic import ValidationError +from sqlalchemy import Engine +from sqlalchemy.orm import Session from controllers.console.app import conversation_variables as conversation_variables_module +from factories import variable_factory from graphon.variables.types import SegmentType +from models import ConversationVariable -def test_get_conversation_variables_returns_paginated_response(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_conversation_variables_returns_paginated_response( + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, +) -> None: api = conversation_variables_module.ConversationVariablesApi() method = unwrap(api.get) created_at = datetime(2026, 1, 1, tzinfo=UTC) updated_at = datetime(2026, 1, 2, tzinfo=UTC) - row = SimpleNamespace( - created_at=created_at, - updated_at=updated_at, - to_variable=lambda: SimpleNamespace( - model_dump=lambda: { - "id": "var-1", - "name": "my_var", - "value_type": "string", - "value": "value", - "description": "desc", - } - ), - ) - session = SimpleNamespace(scalars=lambda _stmt: SimpleNamespace(all=lambda: [row])) - monkeypatch.setattr(conversation_variables_module, "db", SimpleNamespace(engine=object())) - monkeypatch.setattr( - conversation_variables_module, - "sessionmaker", - lambda *_args, **_kwargs: SimpleNamespace(begin=lambda: nullcontext(session)), + variable = variable_factory.build_conversation_variable_from_mapping( + { + "id": "var-1", + "name": "my_var", + "value_type": SegmentType.STRING, + "value": "value", + "description": "desc", + } ) + row = ConversationVariable.from_variable(app_id="app-1", conversation_id="conv-1", variable=variable) + row.created_at = created_at + row.updated_at = updated_at + sqlite_session.add(row) + sqlite_session.commit() + sqlite_session.expire(row) + expected_created_at = int(row.created_at.timestamp()) + expected_updated_at = int(row.updated_at.timestamp()) + monkeypatch.setattr(conversation_variables_module, "db", SimpleNamespace(engine=sqlite_engine)) with app.test_request_context( "/console/api/apps/app-1/conversation-variables", method="GET", query_string={"conversation_id": "conv-1"}, ): - response = method(api, app_model=SimpleNamespace(id="app-1")) + response = method( + api, + conversation_variables_module.ConversationVariablesQuery(conversation_id="conv-1"), + app_model=SimpleNamespace(id="app-1"), + ) assert response["page"] == 1 assert response["limit"] == 100 assert response["total"] == 1 assert response["has_more"] is False assert response["data"][0]["id"] == "var-1" - assert response["data"][0]["created_at"] == int(created_at.timestamp()) - assert response["data"][0]["updated_at"] == int(updated_at.timestamp()) + assert response["data"][0]["created_at"] == expected_created_at + assert response["data"][0]["updated_at"] == expected_updated_at +@pytest.mark.parametrize("sqlite_session", [(ConversationVariable,)], indirect=True) def test_get_conversation_variables_normalizes_value_type_and_value( - app: Flask, monkeypatch: pytest.MonkeyPatch + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, ) -> None: api = conversation_variables_module.ConversationVariablesApi() method = unwrap(api.get) - row = SimpleNamespace( - created_at=None, - updated_at=None, - to_variable=lambda: SimpleNamespace( - model_dump=lambda: { - "id": "var-2", - "name": "my_var_2", - "value_type": SegmentType.INTEGER, - "value": 42, - "description": None, - } - ), - ) - session = SimpleNamespace(scalars=lambda _stmt: SimpleNamespace(all=lambda: [row])) - monkeypatch.setattr(conversation_variables_module, "db", SimpleNamespace(engine=object())) - monkeypatch.setattr( - conversation_variables_module, - "sessionmaker", - lambda *_args, **_kwargs: SimpleNamespace(begin=lambda: nullcontext(session)), + variable = variable_factory.build_conversation_variable_from_mapping( + { + "id": "var-2", + "name": "my_var_2", + "value_type": SegmentType.INTEGER, + "value": 42, + "description": "", + } ) + sqlite_session.add(ConversationVariable.from_variable(app_id="app-1", conversation_id="conv-1", variable=variable)) + sqlite_session.commit() + monkeypatch.setattr(conversation_variables_module, "db", SimpleNamespace(engine=sqlite_engine)) with app.test_request_context( "/console/api/apps/app-1/conversation-variables", method="GET", query_string={"conversation_id": "conv-1"}, ): - response = method(api, app_model=SimpleNamespace(id="app-1")) + response = method( + api, + conversation_variables_module.ConversationVariablesQuery(conversation_id="conv-1"), + app_model=SimpleNamespace(id="app-1"), + ) assert response["data"][0]["value_type"] == "number" assert response["data"][0]["value"] == "42" @@ -100,4 +110,4 @@ def test_get_conversation_variables_requires_conversation_id(app) -> None: with app.test_request_context("/console/api/apps/app-1/conversation-variables", method="GET"): with pytest.raises(ValidationError): - method(api, app_model=SimpleNamespace(id="app-1")) + conversation_variables_module.ConversationVariablesQuery.model_validate({}) diff --git a/api/tests/unit_tests/controllers/console/app/test_generator_api.py b/api/tests/unit_tests/controllers/console/app/test_generator_api.py index 62747a96fd2..c0bcf61ab87 100644 --- a/api/tests/unit_tests/controllers/console/app/test_generator_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_generator_api.py @@ -6,11 +6,19 @@ from types import SimpleNamespace from unittest.mock import MagicMock import pytest -from flask import Flask +from flask import Flask, request from sqlalchemy.orm import Session from controllers.console.app import generator as generator_module from controllers.console.app.error import ProviderNotInitializeError +from controllers.console.app.generator import ( + InstructionGeneratePayload, + InstructionTemplatePayload, + RuleCodeGeneratePayload, + RuleGeneratePayload, + WorkflowGeneratePayload, + WorkflowInstructionSuggestionsPayload, +) from core.errors.error import ProviderTokenNotInitError from models.model import App, AppMode @@ -61,7 +69,7 @@ def test_rule_generate_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -> N method="POST", json={"instruction": "do it", "model_config": _model_config_payload()}, ): - response = method(api, "t1") + response = method(api, RuleGeneratePayload.model_validate(request.get_json()), "t1") assert response == {"rules": []} @@ -81,7 +89,7 @@ def test_rule_code_generate_maps_token_error(app: Flask, monkeypatch: pytest.Mon json={"instruction": "do it", "model_config": _model_config_payload()}, ): with pytest.raises(ProviderNotInitializeError): - method(api, "t1") + method(api, RuleCodeGeneratePayload.model_validate(request.get_json()), "t1") @pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) @@ -100,7 +108,9 @@ def test_instruction_generate_app_not_found(app: Flask, sqlite_session: Session) "model_config": _model_config_payload(), }, ): - response, status = method(api, sqlite_session, "t1") + response, status = method( + api, InstructionGeneratePayload.model_validate(request.get_json()), sqlite_session, "t1" + ) assert status == 400 assert response["error"] == "app app-1 not found" @@ -127,7 +137,9 @@ def test_instruction_generate_workflow_not_found( "model_config": _model_config_payload(), }, ): - response, status = method(api, sqlite_session, "t1") + response, status = method( + api, InstructionGeneratePayload.model_validate(request.get_json()), sqlite_session, "t1" + ) assert status == 400 assert response["error"] == "workflow app-1 not found" @@ -155,7 +167,9 @@ def test_instruction_generate_node_missing( "model_config": _model_config_payload(), }, ): - response, status = method(api, sqlite_session, "t1") + response, status = method( + api, InstructionGeneratePayload.model_validate(request.get_json()), sqlite_session, "t1" + ) assert status == 400 assert response["error"] == "node node-1 not found" @@ -188,7 +202,7 @@ def test_instruction_generate_code_node(app: Flask, monkeypatch: pytest.MonkeyPa "model_config": _model_config_payload(), }, ): - response = method(api, sqlite_session, "t1") + response = method(api, InstructionGeneratePayload.model_validate(request.get_json()), sqlite_session, "t1") assert response == {"code": "x"} assert workflow_service.app_model is app_model @@ -218,7 +232,7 @@ def test_instruction_generate_legacy_modify( "model_config": _model_config_payload(), }, ): - response = method(api, sqlite_session, "t1") + response = method(api, InstructionGeneratePayload.model_validate(request.get_json()), sqlite_session, "t1") assert response == {"instruction": "ok"} @@ -238,7 +252,9 @@ def test_instruction_generate_incompatible_params(app: Flask, sqlite_session: Se "model_config": _model_config_payload(), }, ): - response, status = method(api, sqlite_session, "t1") + response, status = method( + api, InstructionGeneratePayload.model_validate(request.get_json()), sqlite_session, "t1" + ) assert status == 400 assert response["error"] == "incompatible parameters" @@ -253,7 +269,7 @@ def test_instruction_template_prompt(app: Flask) -> None: method="POST", json={"type": "prompt"}, ): - response = method(api) + response = method(api, InstructionTemplatePayload.model_validate(request.get_json())) assert "data" in response @@ -268,7 +284,7 @@ def test_instruction_template_invalid_type(app: Flask) -> None: json={"type": "unknown"}, ): with pytest.raises(ValueError): - method(api) + method(api, InstructionTemplatePayload.model_validate(request.get_json())) # ─ /workflow-generate ───────────────────────────────────────────────────────── @@ -331,7 +347,7 @@ def test_workflow_generate_returns_service_result(app: Flask, monkeypatch: pytes method="POST", json=_workflow_generate_payload(), ): - response = method(api, "t1") + response = method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") assert response == expected @@ -361,7 +377,7 @@ def test_workflow_generate_maps_provider_token_error(app: Flask, monkeypatch: py json=_workflow_generate_payload(), ): with pytest.raises(ProviderNotInitializeError): - method(api, "t1") + method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") def test_workflow_generate_maps_quota_error(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -379,7 +395,7 @@ def test_workflow_generate_maps_quota_error(app: Flask, monkeypatch: pytest.Monk json=_workflow_generate_payload(), ): with pytest.raises(ProviderQuotaExceededError): - method(api, "t1") + method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") def test_workflow_generate_maps_model_not_support_error(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -397,7 +413,7 @@ def test_workflow_generate_maps_model_not_support_error(app: Flask, monkeypatch: json=_workflow_generate_payload(), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(api, "t1") + method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") def test_workflow_generate_maps_invoke_error(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -415,7 +431,7 @@ def test_workflow_generate_maps_invoke_error(app: Flask, monkeypatch: pytest.Mon json=_workflow_generate_payload(), ): with pytest.raises(CompletionRequestError): - method(api, "t1") + method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") def test_workflow_generate_accepts_advanced_chat_mode(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -442,7 +458,7 @@ def test_workflow_generate_accepts_advanced_chat_mode(app: Flask, monkeypatch: p method="POST", json=payload, ): - method(api, "t1") + method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") assert captured["mode"] == "advanced-chat" assert captured["instruction"] == "Summarize a URL" @@ -474,7 +490,7 @@ def test_workflow_generate_forwards_current_graph_for_refine(app: Flask, monkeyp method="POST", json=payload, ): - method(api, "t1") + method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") assert captured["current_graph"] == graph @@ -501,7 +517,7 @@ def test_workflow_generate_current_graph_defaults_to_none(app: Flask, monkeypatc method="POST", json=_workflow_generate_payload(), ): - method(api, "t1") + method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") assert captured["current_graph"] is None @@ -528,7 +544,7 @@ def test_workflow_generate_accepts_auto_mode(app: Flask, monkeypatch: pytest.Mon payload = _workflow_generate_payload() payload["mode"] = "auto" with app.test_request_context("/console/api/workflow-generate", method="POST", json=payload): - response = method(api, "t1") + response = method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") assert captured["mode"] == "auto" assert response["mode"] == "advanced-chat" @@ -608,7 +624,7 @@ def test_workflow_instruction_suggestions_route_returns_list(app: Flask, monkeyp method="POST", json={"mode": "workflow", "language": "French", "count": 3}, ): - response = method(api, "t1") + response = method(api, WorkflowInstructionSuggestionsPayload.model_validate(request.get_json()), "t1") assert response == {"suggestions": ["Summarize a URL", "Translate text"]} assert captured["mode"] == "workflow" @@ -632,7 +648,7 @@ def test_workflow_instruction_suggestions_route_empty_is_valid_200(app: Flask, m method="POST", json={"mode": "advanced-chat"}, ): - response = method(api, "t1") + response = method(api, WorkflowInstructionSuggestionsPayload.model_validate(request.get_json()), "t1") assert response == {"suggestions": []} @@ -667,7 +683,7 @@ def test_workflow_generate_stream_emits_plan_then_result(app: Flask, monkeypatch method="POST", json=_workflow_generate_payload(), ): - response = method(api, "t1") + response = method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") assert response.mimetype == "text/event-stream" frames = _read_sse_frames(response) @@ -694,7 +710,7 @@ def test_workflow_generate_stream_provider_error_emits_result_event( method="POST", json=_workflow_generate_payload(), ): - response = method(api, "t1") + response = method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") frames = _read_sse_frames(response) assert len(frames) == 1 @@ -711,7 +727,7 @@ def test_workflow_generate_stream_rejects_empty_instruction(app: Flask, monkeypa payload = _workflow_generate_payload() payload["instruction"] = " " with app.test_request_context("/console/api/workflow-generate/stream", method="POST", json=payload): - response, status = method(api, "t1") + response, status = method(api, WorkflowGeneratePayload.model_validate(request.get_json()), "t1") assert status == 400 assert response["errors"][0]["code"] == "EMPTY_INSTRUCTION" diff --git a/api/tests/unit_tests/controllers/console/app/test_generator_api_missing.py b/api/tests/unit_tests/controllers/console/app/test_generator_api_missing.py index d6f6bd703f4..6ce20305bdc 100644 --- a/api/tests/unit_tests/controllers/console/app/test_generator_api_missing.py +++ b/api/tests/unit_tests/controllers/console/app/test_generator_api_missing.py @@ -1,8 +1,14 @@ import pytest -from flask import Flask +from flask import Flask, request from sqlalchemy.orm import Session from controllers.console.app import generator as generator_module +from controllers.console.app.generator import ( + InstructionGeneratePayload, + RuleCodeGeneratePayload, + RuleGeneratePayload, + RuleStructuredOutputPayload, +) from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError from graphon.model_runtime.errors.invoke import InvokeError @@ -47,7 +53,7 @@ def test_rule_generate_exceptions(app: Flask, monkeypatch: pytest.MonkeyPatch) - json={"instruction": "do it", "model_config": _model_config_payload()}, ): with pytest.raises(expected_exception): - method(api, "t1") + method(api, RuleGeneratePayload.model_validate(request.get_json()), "t1") def test_rule_code_generate_exceptions(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -73,7 +79,7 @@ def test_rule_code_generate_exceptions(app: Flask, monkeypatch: pytest.MonkeyPat json={"instruction": "do it", "model_config": _model_config_payload()}, ): with pytest.raises(expected_exception): - method(api, "t1") + method(api, RuleCodeGeneratePayload.model_validate(request.get_json()), "t1") def test_structured_output_generate_exceptions(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -100,7 +106,7 @@ def test_structured_output_generate_exceptions(app: Flask, monkeypatch: pytest.M json={"instruction": "do it", "model_config": _model_config_payload()}, ): with pytest.raises(expected_exception): - method(api, "t1") + method(api, RuleStructuredOutputPayload.model_validate(request.get_json()), "t1") @pytest.mark.parametrize("sqlite_session", [()], indirect=True) @@ -138,4 +144,4 @@ def test_instruction_generate_exceptions( }, ): with pytest.raises(expected_exception): - method(api, sqlite_session, "t1") + method(api, InstructionGeneratePayload.model_validate(request.get_json()), sqlite_session, "t1") diff --git a/api/tests/unit_tests/controllers/console/app/test_mcp_server_response.py b/api/tests/unit_tests/controllers/console/app/test_mcp_server_response.py index 1b392d5185e..0978a79b27c 100644 --- a/api/tests/unit_tests/controllers/console/app/test_mcp_server_response.py +++ b/api/tests/unit_tests/controllers/console/app/test_mcp_server_response.py @@ -1,25 +1,59 @@ import datetime from inspect import unwrap from types import SimpleNamespace -from unittest.mock import PropertyMock, patch +from unittest.mock import patch +import pytest from flask import Flask +from sqlalchemy import select +from sqlalchemy.orm import Session +from werkzeug.exceptions import NotFound from controllers.console import console_ns -from controllers.console.app.mcp_server import AppMCPServerController, AppMCPServerResponse +from controllers.console.app.mcp_server import ( + AppMCPServerController, + AppMCPServerRefreshController, + AppMCPServerResponse, + MCPServerCreatePayload, + MCPServerUpdatePayload, +) +from controllers.console.wraps import RBACPermission, RBACResourceScope +from models.enums import AppMCPServerStatus +from models.model import AppMCPServer class _ValidatedResponse: - def __init__(self, payload): + def __init__(self, payload: dict[str, str]) -> None: self._payload = payload - def model_dump(self, mode="json"): + def model_dump(self, mode: str = "json") -> dict[str, str]: return self._payload +def _server( + *, + tenant_id: str = "tenant-1", + app_id: str = "app-1", + name: str = "Demo App", + description: str = "Description", + parameters: str = "{}", + status: AppMCPServerStatus = AppMCPServerStatus.ACTIVE, + server_code: str = "server-code", +) -> AppMCPServer: + return AppMCPServer( + tenant_id=tenant_id, + app_id=app_id, + name=name, + description=description, + parameters=parameters, + status=status, + server_code=server_code, + ) + + class TestAppMCPServerResponse: - def test_parameters_json_string_parsed(self): - data = { + def test_parameters_json_string_parsed(self) -> None: + data: dict[str, object] = { "id": "s1", "name": "test", "server_code": "code", @@ -30,8 +64,8 @@ class TestAppMCPServerResponse: resp = AppMCPServerResponse.model_validate(data) assert resp.parameters == {"key": "value"} - def test_parameters_invalid_json_returns_original(self): - data = { + def test_parameters_invalid_json_returns_original(self) -> None: + data: dict[str, object] = { "id": "s1", "name": "test", "server_code": "code", @@ -42,8 +76,8 @@ class TestAppMCPServerResponse: resp = AppMCPServerResponse.model_validate(data) assert resp.parameters == "not-valid-json" - def test_parameters_dict_passthrough(self): - data = { + def test_parameters_dict_passthrough(self) -> None: + data: dict[str, object] = { "id": "s1", "name": "test", "server_code": "code", @@ -54,8 +88,8 @@ class TestAppMCPServerResponse: resp = AppMCPServerResponse.model_validate(data) assert resp.parameters == {"already": "parsed"} - def test_parameters_json_array_parsed(self): - data = { + def test_parameters_json_array_parsed(self) -> None: + data: dict[str, object] = { "id": "s1", "name": "test", "server_code": "code", @@ -66,9 +100,9 @@ class TestAppMCPServerResponse: resp = AppMCPServerResponse.model_validate(data) assert resp.parameters == ["a", "b"] - def test_timestamps_normalized(self): + def test_timestamps_normalized(self) -> None: dt = datetime.datetime(2024, 1, 1, 0, 0, 0, tzinfo=datetime.UTC) - data = { + data: dict[str, object] = { "id": "s1", "name": "test", "server_code": "code", @@ -82,8 +116,8 @@ class TestAppMCPServerResponse: assert resp.created_at == int(dt.timestamp()) assert resp.updated_at == int(dt.timestamp()) - def test_timestamps_none(self): - data = { + def test_timestamps_none(self) -> None: + data: dict[str, object] = { "id": "s1", "name": "test", "server_code": "code", @@ -97,83 +131,186 @@ class TestAppMCPServerResponse: class TestAppMCPServerController: - def test_get_returns_empty_dict_when_server_missing(self): + def test_get_returns_empty_dict_when_server_missing(self, sqlite_session: Session) -> None: api = AppMCPServerController() method = unwrap(api.get) - with patch("controllers.console.app.mcp_server.db.session.scalar", return_value=None): + with patch("controllers.console.app.mcp_server.db.session", sqlite_session): response = method(api, app_model=SimpleNamespace(id="app-1")) assert response == {} - def test_post_returns_201(self): + def test_post_returns_201(self, sqlite_session: Session) -> None: api = AppMCPServerController() method = unwrap(api.post) payload = {"parameters": {"timeout": 30}} + req_data = MCPServerCreatePayload.model_validate(payload) app = Flask(__name__) app.config["TESTING"] = True with ( app.test_request_context("/", json=payload), - patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), - patch("controllers.console.app.mcp_server.db.session.add"), - patch("controllers.console.app.mcp_server.db.session.commit"), + patch("controllers.console.app.mcp_server.db.session", sqlite_session), patch("controllers.console.app.mcp_server.AppMCPServer.generate_server_code", return_value="server-code"), - patch( - "controllers.console.app.mcp_server.AppMCPServerResponse.model_validate", - return_value=_ValidatedResponse({"id": "server-1"}), - ), ): response, status_code = method( - api, "tenant-1", app_model=SimpleNamespace(id="app-1", name="Demo App", description="App description") + api, + req_data, + "tenant-1", + app_model=SimpleNamespace(id="app-1", name="Demo App", description="App description"), ) - assert response == {"id": "server-1"} + server = sqlite_session.scalar(select(AppMCPServer)) + assert server is not None + assert response["server_code"] == "server-code" + assert response["parameters"] == {"timeout": 30} assert status_code == 201 - def test_put_binds_server_lookup_to_app_ref(self): + def test_put_updates_server_for_app(self, sqlite_session: Session) -> None: api = AppMCPServerController() method = unwrap(api.put) payload = {"id": "server-1", "description": "Updated", "parameters": {"timeout": 30}, "status": "active"} + req_data = MCPServerUpdatePayload.model_validate(payload) app = Flask(__name__) app.config["TESTING"] = True - server = SimpleNamespace( - id="server-1", - tenant_id="tenant-1", - app_id="app-1", - name="Old", - description="Old", - parameters="{}", - status="active", - ) + server = _server(name="Old", description="Old") + server.id = "server-1" + sqlite_session.add(server) + sqlite_session.commit() with ( app.test_request_context("/", json=payload), - patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), - patch("controllers.console.app.mcp_server.db.session.scalar", return_value=server) as scalar, - patch("controllers.console.app.mcp_server.db.session.get") as get_mock, - patch("controllers.console.app.mcp_server.db.session.commit") as commit, - patch( - "controllers.console.app.mcp_server.AppMCPServerResponse.model_validate", - return_value=_ValidatedResponse({"id": "server-1"}), - ), + patch("controllers.console.app.mcp_server.db.session", sqlite_session), ): response = method( api, + req_data, app_model=SimpleNamespace( id="app-1", tenant_id="tenant-1", name="Demo App", description="App description" ), ) + sqlite_session.expire_all() + updated_server = sqlite_session.get(AppMCPServer, "server-1") + assert updated_server is not None + assert response["id"] == "server-1" + assert updated_server.description == "Updated" + + @pytest.mark.parametrize( + ("foreign_tenant_id", "foreign_app_id"), + [ + ("tenant-2", "app-1"), + ("tenant-1", "app-2"), + ], + ) + def test_put_scopes_server_lookup_to_complete_app_ref( + self, + sqlite_session: Session, + foreign_tenant_id: str, + foreign_app_id: str, + ) -> None: + api = AppMCPServerController() + method = unwrap(api.put) + payload = {"id": "server-1", "description": "Updated", "parameters": {"timeout": 30}, "status": "active"} + req_data = MCPServerUpdatePayload.model_validate(payload) + app = Flask(__name__) + app.config["TESTING"] = True + foreign_server = _server( + tenant_id=foreign_tenant_id, + app_id=foreign_app_id, + name="Other", + server_code="other-code", + ) + foreign_server.id = "server-1" + sqlite_session.add(foreign_server) + sqlite_session.commit() + + with ( + app.test_request_context("/", json=payload), + patch("controllers.console.app.mcp_server.db.session", sqlite_session), + pytest.raises(NotFound), + ): + method( + api, + req_data, + app_model=SimpleNamespace( + id="app-1", tenant_id="tenant-1", name="Demo App", description="App description" + ), + ) + + sqlite_session.expire_all() + unchanged_server = sqlite_session.get(AppMCPServer, "server-1") + assert unchanged_server is not None + assert unchanged_server.description == "Description" + + +class TestAppMCPServerRefreshController: + def test_post_refreshes_server_bound_to_app_and_tenant(self): + api = AppMCPServerRefreshController() + method = unwrap(api.post) + server = SimpleNamespace(server_code="old-code") + + with ( + patch("controllers.console.app.mcp_server.db.session.scalar", return_value=server) as scalar, + patch("controllers.console.app.mcp_server.db.session.commit") as commit, + patch("controllers.console.app.mcp_server.AppMCPServer.generate_server_code", return_value="new-code"), + patch( + "controllers.console.app.mcp_server.AppMCPServerResponse.model_validate", + return_value=_ValidatedResponse({"id": "server-1", "server_code": "new-code"}), + ), + ): + response = method(api, "tenant-1", app_model=SimpleNamespace(id="app-1")) + stmt = scalar.call_args.args[0] compiled = stmt.compile() statement = str(compiled) - assert "app_mcp_servers.id" in statement assert "app_mcp_servers.tenant_id" in statement assert "app_mcp_servers.app_id" in statement - assert payload["id"] in compiled.params.values() assert "tenant-1" in compiled.params.values() assert "app-1" in compiled.params.values() - get_mock.assert_not_called() + assert server.server_code == "new-code" commit.assert_called_once() - assert response == {"id": "server-1"} + assert response == {"id": "server-1", "server_code": "new-code"} + + def test_route_is_app_scoped_post(self): + route_map = { + resource.__name__: urls + for resource, urls, _route_doc, _kwargs in console_ns.resources + if resource.__name__ == "AppMCPServerRefreshController" + } + + assert route_map["AppMCPServerRefreshController"] == ("/apps//server/refresh",) + assert hasattr(AppMCPServerRefreshController, "post") + assert not hasattr(AppMCPServerRefreshController, "get") + + def test_post_requires_app_view_layout_permission(self): + method = AppMCPServerRefreshController.post + while "rbac_permission_required" not in method.__code__.co_qualname: + method = method.__wrapped__ + + class PermissionCheckedError(Exception): + pass + + current_user = SimpleNamespace(id="account-1") + with ( + patch("controllers.common.wraps.dify_config.RBAC_ENABLED", True), + patch( + "controllers.common.wraps.current_account_with_tenant", + return_value=(current_user, "tenant-1"), + ), + patch( + "controllers.common.wraps.enforce_rbac_access", + side_effect=PermissionCheckedError, + ) as enforce_rbac_access, + pytest.raises(PermissionCheckedError), + ): + method(AppMCPServerRefreshController(), app_id="app-1") + + enforce_rbac_access.assert_called_once_with( + tenant_id="tenant-1", + account_id="account-1", + resource_type=RBACResourceScope.APP, + scene=RBACPermission.APP_VIEW_LAYOUT, + resource_required=True, + path_args={"app_id": "app-1"}, + ) diff --git a/api/tests/unit_tests/controllers/console/app/test_message_api.py b/api/tests/unit_tests/controllers/console/app/test_message_api.py index 0c2efc12198..ae6e994b6d5 100644 --- a/api/tests/unit_tests/controllers/console/app/test_message_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_message_api.py @@ -7,12 +7,67 @@ from unittest.mock import MagicMock import pytest from flask import Flask +from sqlalchemy import event +from sqlalchemy.orm import Session from controllers.console.app import message as message_module +from core.app.entities.app_invoke_entities import InvokeFrom +from models.enums import ConversationFromSource +from models.model import AppMode, Conversation, Message -def test_app_message_routes_pass_injected_session(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - session = MagicMock() +def _persist_message(session: Session, *, message_id: str, app_id: str = "app-1") -> Message: + conversation = Conversation( + app_id=app_id, + app_model_config_id=None, + model_provider=None, + override_model_configs=None, + model_id=None, + mode=AppMode.CHAT, + name="Conversation", + inputs={}, + introduction="", + system_instruction="", + system_instruction_tokens=0, + status="normal", + invoke_from=InvokeFrom.DEBUGGER, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id=None, + from_account_id="account-1", + ) + conversation.id = "conversation-1" + message = Message( + app_id=app_id, + conversation_id=conversation.id, + inputs={}, + query="query", + 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, + provider_response_latency=0, + total_price=0, + currency="USD", + invoke_from=InvokeFrom.DEBUGGER, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id=None, + from_account_id="account-1", + app_mode=AppMode.CHAT, + ) + message.id = message_id + session.add_all([conversation, message]) + session.flush() + return message + + +def test_app_message_routes_pass_injected_session( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: + session = unbound_session current_user = SimpleNamespace(id="account-1") app_model = SimpleNamespace(id="app-1", mode="chat") message_id = "550e8400-e29b-41d4-a716-446655440000" @@ -44,17 +99,15 @@ def test_app_message_routes_pass_injected_session(app: Flask, monkeypatch: pytes assert get_message_detail.call_args.kwargs["session"] is session -def test_update_message_feedback_commits_injected_session(app: Flask) -> None: +def test_update_message_feedback_commits_injected_session(app: Flask, sqlite_session: Session) -> None: message_id = "550e8400-e29b-41d4-a716-446655440000" feedback = SimpleNamespace(rating="dislike", content=None) get_admin_feedback = MagicMock(return_value=feedback) - message = SimpleNamespace( - id=message_id, - conversation_id="conversation-1", - admin_feedback_with_session=get_admin_feedback, - ) - session = MagicMock() - session.scalar.return_value = message + message = _persist_message(sqlite_session, message_id=message_id) + message.admin_feedback_with_session = get_admin_feedback + session = sqlite_session + commits: list[str] = [] + event.listen(session, "after_commit", lambda _session: commits.append("commit")) with app.test_request_context(json={"message_id": message_id, "rating": "like", "content": "helpful"}): result = message_module._update_message_feedback( @@ -67,16 +120,15 @@ def test_update_message_feedback_commits_injected_session(app: Flask) -> None: assert feedback.rating == "like" assert feedback.content == "helpful" get_admin_feedback.assert_called_once_with(session=session) - session.commit.assert_called_once_with() + assert commits == ["commit"] -def test_get_message_detail_uses_injected_session(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_message_detail_uses_injected_session(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: message_id = "550e8400-e29b-41d4-a716-446655440000" - message = SimpleNamespace(id=message_id) + message = _persist_message(sqlite_session, message_id=message_id) response_source = object() response_source_factory = MagicMock(return_value=response_source) - session = MagicMock() - session.scalar.return_value = message + session = sqlite_session monkeypatch.setattr(message_module, "attach_message_extra_contents", MagicMock()) monkeypatch.setattr(message_module, "MessageResponseSource", response_source_factory) monkeypatch.setattr(message_module, "dump_response", lambda _model, value: value) @@ -89,11 +141,10 @@ def test_get_message_detail_uses_injected_session(monkeypatch: pytest.MonkeyPatc assert result is response_source response_source_factory.assert_called_once_with(message, session=session) - session.scalar.assert_called_once() -def test_message_response_source_uses_caller_session_for_nested_fields() -> None: - session = MagicMock() +def test_message_response_source_uses_caller_session_for_nested_fields(unbound_session: Session) -> None: + session = unbound_session account = object() feedback = MagicMock() feedback.from_account_with_session.return_value = account diff --git a/api/tests/unit_tests/controllers/console/app/test_ops_trace_api.py b/api/tests/unit_tests/controllers/console/app/test_ops_trace_api.py index d08106aa5f2..a1f2147a895 100644 --- a/api/tests/unit_tests/controllers/console/app/test_ops_trace_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_ops_trace_api.py @@ -6,6 +6,7 @@ from unittest.mock import MagicMock, PropertyMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden from controllers.common import wraps as common_wraps @@ -13,8 +14,10 @@ from controllers.console import console_ns from controllers.console import wraps as console_wraps from controllers.console.app import ops_trace as ops_trace_module from controllers.console.app import wraps as app_wraps +from enums import DeploymentEdition from libs import login as login_lib from models.account import Account, AccountStatus, TenantAccountRole +from models.model import App, AppMode, IconType def _make_account(role: TenantAccountRole) -> Account: @@ -40,7 +43,7 @@ def _patch_console_guards( ) -> None: monkeypatch.setattr(login_lib.dify_config, "LOGIN_DISABLED", True) monkeypatch.setattr(login_lib.dify_config, "RBAC_ENABLED", rbac_enabled) - monkeypatch.setattr(console_wraps.dify_config, "EDITION", "CLOUD") + monkeypatch.setattr(console_wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr(login_lib, "current_user", account) monkeypatch.setattr(login_lib, "current_account_with_tenant", lambda: (account, account.current_tenant_id)) monkeypatch.setattr(console_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id)) @@ -137,11 +140,33 @@ def test_trace_config_mutations_require_rbac_permission( payload: dict[str, object] | None, service_method_name: str, service_result: object, + sqlite_session: Session, ) -> None: app.config.setdefault("RESTX_MASK_HEADER", "X-Fields") account = _make_account(TenantAccountRole.NORMAL) _patch_console_guards(monkeypatch, account, _make_app(), rbac_enabled=True) - monkeypatch.setattr(common_wraps.db, "session", SimpleNamespace(scalar=lambda _stmt: "other-account")) + owned_app = App() + owned_app.id = "app-123" + owned_app.tenant_id = "tenant-123" + owned_app.name = "Trace app" + owned_app.description = "" + owned_app.mode = AppMode.CHAT + owned_app.icon_type = IconType.EMOJI + owned_app.icon = "robot" + owned_app.icon_background = "#ffffff" + owned_app.enable_site = False + owned_app.enable_api = False + owned_app.api_rpm = 0 + owned_app.api_rph = 0 + owned_app.is_demo = False + owned_app.is_public = False + owned_app.is_universal = False + owned_app.max_active_requests = None + owned_app.maintainer = "other-account" + owned_app.use_icon_as_answer_icon = False + sqlite_session.add(owned_app) + sqlite_session.commit() + monkeypatch.setattr(common_wraps.db, "session", sqlite_session) monkeypatch.setattr(common_wraps.RBACService.CheckAccess, "check", MagicMock(return_value=False)) service_mock = MagicMock(return_value=service_result) monkeypatch.setattr(ops_trace_module.OpsService, service_method_name, service_mock) diff --git a/api/tests/unit_tests/controllers/console/app/test_statistic_api.py b/api/tests/unit_tests/controllers/console/app/test_statistic_api.py index c51a38ad798..8afb4762e8c 100644 --- a/api/tests/unit_tests/controllers/console/app/test_statistic_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_statistic_api.py @@ -53,7 +53,12 @@ def test_daily_message_statistic_returns_rows(app: Flask, monkeypatch: pytest.Mo _install_db(monkeypatch, rows) with app.test_request_context("/console/api/apps/app-1/statistics/daily-messages", method="GET"): - response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + response = method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) assert _json_payload(response) == {"data": [{"date": "2024-01-01", "message_count": 3}]} @@ -67,7 +72,12 @@ def test_daily_conversation_statistic_returns_rows(app: Flask, monkeypatch: pyte _install_db(monkeypatch, rows) with app.test_request_context("/console/api/apps/app-1/statistics/daily-conversations", method="GET"): - response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + response = method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) assert _json_payload(response) == {"data": [{"date": "2024-01-02", "conversation_count": 5}]} @@ -81,7 +91,12 @@ def test_daily_token_cost_statistic_returns_rows(app: Flask, monkeypatch: pytest _install_db(monkeypatch, rows) with app.test_request_context("/console/api/apps/app-1/statistics/token-costs", method="GET"): - response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + response = method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) data = _json_payload(response) assert len(data["data"]) == 1 @@ -99,7 +114,12 @@ def test_daily_terminals_statistic_returns_rows(app: Flask, monkeypatch: pytest. _install_db(monkeypatch, rows) with app.test_request_context("/console/api/apps/app-1/statistics/daily-end-users", method="GET"): - response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + response = method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) assert _json_payload(response) == {"data": [{"date": "2024-01-04", "terminal_count": 7}]} @@ -126,7 +146,12 @@ def test_daily_message_statistic_with_invalid_time_range(app: Flask, monkeypatch with app.test_request_context("/console/api/apps/app-1/statistics/daily-messages", method="GET"): with pytest.raises(BadRequest): - method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) def test_daily_message_statistic_multiple_rows(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -142,7 +167,12 @@ def test_daily_message_statistic_multiple_rows(app: Flask, monkeypatch: pytest.M _install_db(monkeypatch, rows) with app.test_request_context("/console/api/apps/app-1/statistics/daily-messages", method="GET"): - response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + response = method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) data = _json_payload(response) assert len(data["data"]) == 3 @@ -156,7 +186,12 @@ def test_daily_message_statistic_empty_result(app: Flask, monkeypatch: pytest.Mo _install_db(monkeypatch, []) with app.test_request_context("/console/api/apps/app-1/statistics/daily-messages", method="GET"): - response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + response = method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) assert _json_payload(response) == {"data": []} @@ -175,7 +210,12 @@ def test_daily_conversation_statistic_with_time_range(app: Flask, monkeypatch: p monkeypatch.setattr(statistic_module, "convert_datetime_to_date", lambda field: field) with app.test_request_context("/console/api/apps/app-1/statistics/daily-conversations", method="GET"): - response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + response = method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) assert _json_payload(response) == {"data": [{"date": "2024-01-02", "conversation_count": 5}]} @@ -192,7 +232,12 @@ def test_daily_token_cost_with_multiple_currencies(app: Flask, monkeypatch: pyte _install_db(monkeypatch, rows) with app.test_request_context("/console/api/apps/app-1/statistics/token-costs", method="GET"): - response = method(api, SimpleNamespace(timezone="UTC"), app_model=SimpleNamespace(id="app-1")) + response = method( + api, + SimpleNamespace(start=None, end=None), + SimpleNamespace(timezone="UTC"), + app_model=SimpleNamespace(id="app-1"), + ) data = _json_payload(response) assert len(data["data"]) == 2 diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow.py b/api/tests/unit_tests/controllers/console/app/test_workflow.py index fd07c42f752..82d6add152b 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow.py @@ -10,11 +10,14 @@ from unittest.mock import Mock import pytest from flask import Flask from pydantic import ValidationError +from sqlalchemy import Engine +from sqlalchemy.orm import Session from werkzeug.exceptions import HTTPException, NotFound from controllers.common.errors import InvalidArgumentError from controllers.console.app import workflow as workflow_module from controllers.console.app.error import DraftWorkflowNotExist, DraftWorkflowNotSync +from core.workflow.llm_environment_variable import LLMEnvironmentVariable from graphon.file import File, FileTransferMethod, FileType from graphon.variables import SecretVariable, StringVariable from graphon.variables.variables import RAGPipelineVariable @@ -145,7 +148,8 @@ def test_sync_draft_workflow_success(app: Flask, monkeypatch: pytest.MonkeyPatch workflow_module.variable_factory, "build_conversation_variable_from_mapping", lambda *_args: "conv" ) - service = SimpleNamespace(sync_draft_workflow=lambda **_kwargs: workflow) + sync_draft_workflow = Mock(return_value=workflow) + service = SimpleNamespace(sync_draft_workflow=sync_draft_workflow) monkeypatch.setattr(workflow_module, "WorkflowService", lambda: service) api = workflow_module.DraftWorkflowApi() @@ -159,6 +163,89 @@ def test_sync_draft_workflow_success(app: Flask, monkeypatch: pytest.MonkeyPatch response = handler(api, "t1", app_model=SimpleNamespace(id="app")) assert response["result"] == "success" + assert sync_draft_workflow.call_args.kwargs["environment_variables"] == [] + assert sync_draft_workflow.call_args.kwargs["preserve_environment_variables"] is True + + +def test_sync_draft_workflow_passes_environment_patch(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + workflow = SimpleNamespace( + unique_hash="next-hash", + updated_at=None, + created_at=datetime(2024, 1, 1), + ) + patched_variable = StringVariable( + id="env-model", + name="shared_model", + value="model", + selector=["env", "shared_model"], + ) + build_environment_variable = Mock(return_value=patched_variable) + sync_draft_workflow = Mock(return_value=workflow) + monkeypatch.setattr( + workflow_module.variable_factory, + "build_environment_variable_from_mapping", + build_environment_variable, + ) + monkeypatch.setattr( + workflow_module, + "WorkflowService", + lambda: SimpleNamespace(sync_draft_workflow=sync_draft_workflow), + ) + + api = workflow_module.DraftWorkflowApi() + handler = inspect.unwrap(api.post) + with app.test_request_context( + "/apps/app/workflows/draft", + method="POST", + json={ + "graph": {}, + "features": {}, + "hash": "current-hash", + "environment_variable_patch": { + "environment_variables": [ + { + "id": "env-model", + "name": "shared_model", + "value": "model", + "value_type": "string", + } + ], + "deleted_environment_variable_ids": ["env-old"], + }, + }, + ): + response = handler(api, "t1", app_model=SimpleNamespace(id="app")) + + assert response["result"] == "success" + assert sync_draft_workflow.call_args.kwargs["preserve_environment_variables"] is True + assert sync_draft_workflow.call_args.kwargs["environment_variable_upserts"] == [patched_variable] + assert sync_draft_workflow.call_args.kwargs["deleted_environment_variable_ids"] == ["env-old"] + build_environment_variable.assert_called_once() + + +def test_sync_draft_workflow_rejects_overlapping_environment_patch_ids() -> None: + with pytest.raises(ValidationError, match="cannot be upserted and deleted"): + workflow_module.SyncDraftWorkflowPayload.model_validate( + { + "graph": {}, + "features": {}, + "environment_variable_patch": { + "environment_variables": [{"id": "env-model"}], + "deleted_environment_variable_ids": ["env-model"], + }, + } + ) + + +def test_sync_draft_workflow_rejects_legacy_environment_variables() -> None: + with pytest.raises(ValidationError, match="Extra inputs are not permitted"): + workflow_module.SyncDraftWorkflowPayload.model_validate( + { + "graph": {}, + "features": {}, + "environment_variables": [], + } + ) def test_sync_draft_workflow_hash_mismatch(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: @@ -331,45 +418,34 @@ def test_restore_published_workflow_to_draft_returns_400_for_invalid_structure( def test_get_published_workflows_serializes_items_before_session_closes( - app: Flask, monkeypatch: pytest.MonkeyPatch + app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine ) -> None: api = workflow_module.PublishedAllWorkflowApi() handler = inspect.unwrap(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() - base_workflow = _make_workflow() class _Workflow: + def __init__(self, session: Session) -> None: + self._session = session + def __getattr__(self, name): return getattr(base_workflow, name) @property def id(self): - assert session_state["open"] is True + assert self._session.in_transaction() return "w1" - monkeypatch.setattr(workflow_module, "db", SimpleNamespace(engine=object())) - monkeypatch.setattr(workflow_module, "sessionmaker", lambda *_args, **_kwargs: _SessionMaker()) + def get_all_published_workflow(*, session: Session, **_kwargs): + assert isinstance(session, Session) + return [_Workflow(session)], False + + monkeypatch.setattr(workflow_module, "db", SimpleNamespace(engine=sqlite_engine)) monkeypatch.setattr( workflow_module, "WorkflowService", - lambda: SimpleNamespace( - get_all_published_workflow=lambda **_kwargs: ([_Workflow()], False), - ), + lambda: SimpleNamespace(get_all_published_workflow=get_all_published_workflow), ) with app.test_request_context( @@ -521,6 +597,31 @@ def test_workflow_response_masks_secret_environment_variables() -> None: ] +def test_workflow_response_preserves_llm_environment_variable_type() -> None: + workflow = _make_workflow( + environment_variables=[ + LLMEnvironmentVariable( + id="env-llm", + name="for_summarize", + value={"provider": "provider", "name": "model", "mode": "chat"}, + selector=["env", "for_summarize"], + ) + ] + ) + + response = workflow_module.WorkflowResponse.model_validate(workflow, from_attributes=True).model_dump(mode="json") + + assert response["environment_variables"] == [ + { + "id": "env-llm", + "name": "for_summarize", + "value": {"provider": "provider", "name": "model", "mode": "chat"}, + "value_type": "llm", + "description": "", + } + ] + + def test_workflow_response_rejects_invalid_environment_variable_dict() -> None: workflow = _make_workflow(environment_variables=[{"value_type": "not-a-segment-type"}]) @@ -623,6 +724,7 @@ def test_advanced_chat_run_conversation_not_exists(app: Flask, monkeypatch: pyte def test_trigger_run_loads_draft_with_request_session( app: Flask, monkeypatch: pytest.MonkeyPatch, + unbound_session: Session, resource: type, payload: dict[str, object], ) -> None: @@ -632,7 +734,7 @@ def test_trigger_run_loads_draft_with_request_session( "WorkflowService", lambda: SimpleNamespace(get_draft_workflow=get_draft_workflow), ) - session = Mock() + session = unbound_session app_model = SimpleNamespace(id="app-1") handler = inspect.unwrap(resource.post) diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_comment_api.py b/api/tests/unit_tests/controllers/console/app/test_workflow_comment_api.py index 8c9c9f9d562..d568744292c 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_comment_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_comment_api.py @@ -14,6 +14,7 @@ from controllers.console import console_ns from controllers.console import wraps as console_wraps from controllers.console.app import workflow_comment as workflow_comment_module from controllers.console.app import wraps as app_wraps +from enums import DeploymentEdition from libs import login as login_lib from models.account import Account, AccountStatus, TenantAccountRole @@ -47,7 +48,7 @@ def _patch_console_guards(monkeypatch: pytest.MonkeyPatch, account: Account, app monkeypatch.setattr(login_lib, "current_account_with_tenant", lambda: (account, account.current_tenant_id)) monkeypatch.setattr(login_lib, "check_csrf_token", lambda *_, **__: None) monkeypatch.setattr(console_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id)) - monkeypatch.setattr(console_wraps.dify_config, "EDITION", "CLOUD") + monkeypatch.setattr(console_wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr(app_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id)) monkeypatch.setattr(app_wraps, "_load_app_model_from_scoped_session", lambda _app_id: app_model) diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py b/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py index 956706eafb6..507a38e0155 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_human_input_debug_api.py @@ -11,6 +11,7 @@ from pydantic import ValidationError from controllers.console import wraps as console_wraps from controllers.console.app import workflow as workflow_module from controllers.console.app import wraps as app_wraps +from enums import DeploymentEdition from libs import login as login_lib from models.account import Account, AccountStatus, TenantAccountRole from models.model import AppMode @@ -32,15 +33,14 @@ def _make_app(mode: AppMode) -> SimpleNamespace: def _patch_console_guards(monkeypatch: pytest.MonkeyPatch, account: Account, app_model: SimpleNamespace) -> None: # Skip setup and auth guardrails - monkeypatch.setattr("configs.dify_config.EDITION", "CLOUD") + monkeypatch.setattr("configs.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr(login_lib.dify_config, "LOGIN_DISABLED", True) monkeypatch.setattr(login_lib, "current_user", account) monkeypatch.setattr(login_lib, "current_account_with_tenant", lambda: (account, account.current_tenant_id)) monkeypatch.setattr(login_lib, "check_csrf_token", lambda *_, **__: None) monkeypatch.setattr(console_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id)) monkeypatch.setattr(app_wraps, "current_account_with_tenant", lambda: (account, account.current_tenant_id)) - monkeypatch.setattr(console_wraps.dify_config, "EDITION", "CLOUD") - monkeypatch.delenv("INIT_PASSWORD", raising=False) + monkeypatch.setattr(console_wraps.dify_config, "INIT_PASSWORD", "") # Avoid hitting the database when resolving the app model monkeypatch.setattr(app_wraps, "_load_app_model_from_scoped_session", lambda _app_id: app_model) diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py b/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py index dfe35a89f57..8cca6ab9f41 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_node_output_inspector.py @@ -31,6 +31,7 @@ from uuid import UUID import pytest from controllers.console.app import workflow_node_output_inspector as ctrl +from graphon.enums import WorkflowExecutionStatus from services.workflow.inspector_events import InspectorMessage from services.workflow.node_output_inspector_service import ( NodeOutputInspectorError, @@ -61,8 +62,6 @@ def run_id() -> UUID: def _snapshot_view(*, status: str, node_id: str = "agent-1") -> WorkflowRunSnapshotView: - from graphon.enums import WorkflowExecutionStatus - return WorkflowRunSnapshotView( workflow_run_id="00000000-0000-0000-0000-0000000000aa", workflow_run_status=WorkflowExecutionStatus(status), diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_pause_details_api.py b/api/tests/unit_tests/controllers/console/app/test_workflow_pause_details_api.py index 502fa99fd75..88aed318fb2 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_pause_details_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_pause_details_api.py @@ -1,19 +1,68 @@ +"""Console workflow pause-detail tests backed by persisted workflow execution state.""" + from __future__ import annotations import inspect +from dataclasses import dataclass from datetime import datetime -from types import SimpleNamespace from unittest.mock import Mock +from uuid import uuid4 import pytest from flask import Flask +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from controllers.common.errors import NotFoundError from controllers.console.app import workflow_run as workflow_run_module from core.workflow.nodes.human_input.entities import ParagraphInputConfig, UserActionConfig from core.workflow.nodes.human_input.pause_reason import HumanInputRequired from graphon.enums import WorkflowExecutionStatus -from models.workflow import WorkflowRun +from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom +from models.workflow import WorkflowPause, WorkflowRun, WorkflowType + + +@dataclass(frozen=True) +class _Database: + engine: Engine + session: Session + + +def _persist_run( + session: Session, + *, + run_id: str, + tenant_id: str, + status: WorkflowExecutionStatus, + paused: bool = False, +) -> WorkflowRun: + workflow_id = str(uuid4()) + workflow_run = WorkflowRun( + id=run_id, + tenant_id=tenant_id, + app_id=str(uuid4()), + workflow_id=workflow_id, + type=WorkflowType.WORKFLOW, + triggered_from=WorkflowRunTriggeredFrom.DEBUGGING, + version="draft", + graph="{}", + inputs="{}", + status=status, + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + created_at=datetime(2024, 1, 1, 12, 0, 0), + ) + session.add(workflow_run) + if paused: + session.add( + WorkflowPause( + workflow_id=workflow_id, + workflow_run_id=run_id, + state_object_key="workflow-pauses/state.json", + ) + ) + session.commit() + return workflow_run class _PauseEntity: @@ -25,15 +74,25 @@ class _PauseEntity: return self._reasons -def test_pause_details_returns_backstage_input_url(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_pause_details_returns_backstage_input_url( + app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: monkeypatch.setattr(workflow_run_module.dify_config, "APP_WEB_URL", "https://web.example.com") - workflow_run = Mock(spec=WorkflowRun) - workflow_run.tenant_id = "tenant-123" - workflow_run.status = WorkflowExecutionStatus.PAUSED - workflow_run.created_at = datetime(2024, 1, 1, 12, 0, 0) - fake_db = SimpleNamespace(engine=Mock(), session=SimpleNamespace(get=lambda *_: workflow_run)) - monkeypatch.setattr(workflow_run_module, "db", fake_db) + tenant_id = str(uuid4()) + run_id = str(uuid4()) + _persist_run( + sqlite_session, + run_id=run_id, + tenant_id=tenant_id, + status=WorkflowExecutionStatus.PAUSED, + paused=True, + ) + monkeypatch.setattr( + workflow_run_module, + "db", + _Database(engine=sqlite_session.get_bind(), session=sqlite_session), + ) reason = HumanInputRequired( form_id="form-1", @@ -58,12 +117,12 @@ def test_pause_details_returns_backstage_input_url(app: Flask, monkeypatch: pyte lambda _form_ids: {"form-1": "backstage-token"}, ) - with app.test_request_context("/console/api/workflow/run-1/pause-details", method="GET"): + with app.test_request_context(f"/console/api/workflow/{run_id}/pause-details", method="GET"): handler = inspect.unwrap(workflow_run_module.ConsoleWorkflowPauseDetailsApi.get) response, status = handler( workflow_run_module.ConsoleWorkflowPauseDetailsApi(), - "tenant-123", - workflow_run_id="run-1", + tenant_id, + workflow_run_id=run_id, ) assert status == 200 @@ -77,39 +136,56 @@ def test_pause_details_returns_backstage_input_url(app: Flask, monkeypatch: pyte assert "pending_human_inputs" not in response -def test_pause_details_tenant_isolation(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: +def test_pause_details_tenant_isolation(app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: monkeypatch.setattr(workflow_run_module.dify_config, "APP_WEB_URL", "https://web.example.com") - workflow_run = Mock(spec=WorkflowRun) - workflow_run.tenant_id = "tenant-456" - workflow_run.status = WorkflowExecutionStatus.PAUSED - workflow_run.created_at = datetime(2024, 1, 1, 12, 0, 0) - fake_db = SimpleNamespace(engine=Mock(), session=SimpleNamespace(get=lambda *_: workflow_run)) - monkeypatch.setattr(workflow_run_module, "db", fake_db) + run_id = str(uuid4()) + _persist_run( + sqlite_session, + run_id=run_id, + tenant_id=str(uuid4()), + status=WorkflowExecutionStatus.PAUSED, + paused=True, + ) + monkeypatch.setattr( + workflow_run_module, + "db", + _Database(engine=sqlite_session.get_bind(), session=sqlite_session), + ) handler = inspect.unwrap(workflow_run_module.ConsoleWorkflowPauseDetailsApi.get) - with app.test_request_context("/console/api/workflow/run-1/pause-details", method="GET"): + with app.test_request_context(f"/console/api/workflow/{run_id}/pause-details", method="GET"): with pytest.raises(NotFoundError): handler( workflow_run_module.ConsoleWorkflowPauseDetailsApi(), - "tenant-123", - workflow_run_id="run-1", + str(uuid4()), + workflow_run_id=run_id, ) -def test_pause_details_returns_empty_response_for_non_paused_run(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - workflow_run = Mock(spec=WorkflowRun) - workflow_run.tenant_id = "tenant-123" - workflow_run.status = WorkflowExecutionStatus.RUNNING - fake_db = SimpleNamespace(engine=Mock(), session=SimpleNamespace(get=lambda *_: workflow_run)) - monkeypatch.setattr(workflow_run_module, "db", fake_db) +def test_pause_details_returns_empty_response_for_non_paused_run( + app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + tenant_id = str(uuid4()) + run_id = str(uuid4()) + _persist_run( + sqlite_session, + run_id=run_id, + tenant_id=tenant_id, + status=WorkflowExecutionStatus.RUNNING, + ) + monkeypatch.setattr( + workflow_run_module, + "db", + _Database(engine=sqlite_session.get_bind(), session=sqlite_session), + ) - with app.test_request_context("/console/api/workflow/run-1/pause-details", method="GET"): + with app.test_request_context(f"/console/api/workflow/{run_id}/pause-details", method="GET"): handler = inspect.unwrap(workflow_run_module.ConsoleWorkflowPauseDetailsApi.get) response, status = handler( workflow_run_module.ConsoleWorkflowPauseDetailsApi(), - "tenant-123", - workflow_run_id="run-1", + tenant_id, + workflow_run_id=run_id, ) assert status == 200 diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_run_api.py b/api/tests/unit_tests/controllers/console/app/test_workflow_run_api.py index 71034ebd405..4986331e5e7 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_run_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_run_api.py @@ -98,7 +98,11 @@ def test_workflow_run_list_returns_frontend_history_contract(app: Flask, monkeyp handler = unwrap(api.get) with app.test_request_context("/apps/app-1/workflow-runs?limit=10", method="GET"): - payload = handler(api, app_model=SimpleNamespace(id="app-1", tenant_id="tenant-1")) + payload = handler( + api, + workflow_run_module.WorkflowRunListQuery(limit=10), + app_model=SimpleNamespace(id="app-1", tenant_id="tenant-1"), + ) response = _serialize_200_response(api.get, payload) @@ -139,7 +143,11 @@ def test_advanced_chat_workflow_run_list_keeps_message_fields(app: Flask, monkey handler = unwrap(api.get) with app.test_request_context("/apps/app-1/advanced-chat/workflow-runs?limit=1", method="GET"): - payload = handler(api, app_model=SimpleNamespace(id="app-1", tenant_id="tenant-1")) + payload = handler( + api, + workflow_run_module.WorkflowRunListQuery(limit=1), + app_model=SimpleNamespace(id="app-1", tenant_id="tenant-1"), + ) response = _serialize_200_response(api.get, payload) diff --git a/api/tests/unit_tests/controllers/console/app/test_workflow_trigger_api.py b/api/tests/unit_tests/controllers/console/app/test_workflow_trigger_api.py index 14386efda37..cb1ce0cf87f 100644 --- a/api/tests/unit_tests/controllers/console/app/test_workflow_trigger_api.py +++ b/api/tests/unit_tests/controllers/console/app/test_workflow_trigger_api.py @@ -1,14 +1,64 @@ from __future__ import annotations import inspect +import uuid from datetime import UTC, datetime 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 import Engine +from sqlalchemy.orm import Session from controllers.console import console_ns from controllers.console.app import workflow_trigger as workflow_trigger_module +from models.base import TypeBase +from models.enums import AppTriggerStatus, AppTriggerType +from models.model import App, AppMode, IconType +from models.trigger import AppTrigger + + +@pytest.fixture +def database_session(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch): + models = (App, AppTrigger) + tables = [model.metadata.tables[model.__tablename__] for model in models] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + monkeypatch.setattr(workflow_trigger_module, "db", SimpleNamespace(engine=sqlite_engine)) + with Session(sqlite_engine, expire_on_commit=False) as session: + yield session + + +def _persist_app_trigger( + session: Session, + *, + tenant_id: str | None = None, + status: AppTriggerStatus = AppTriggerStatus.ENABLED, +) -> tuple[App, AppTrigger]: + tenant_id = tenant_id or str(uuid.uuid4()) + app_model = App( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Workflow App", + mode=AppMode.WORKFLOW, + icon_type=IconType.EMOJI, + icon="workflow", + icon_background="#FFFFFF", + enable_site=False, + enable_api=True, + ) + trigger = AppTrigger( + tenant_id=tenant_id, + app_id=app_model.id, + node_id="node-1", + trigger_type=AppTriggerType.TRIGGER_PLUGIN, + title="Trigger", + provider_name="provider", + status=status, + ) + session.add_all([app_model, trigger]) + session.commit() + return app_model, trigger def test_parser_models_validate(): @@ -23,22 +73,20 @@ def test_parser_models_validate(): def test_workflow_trigger_response_serializes_datetime(): created_at = datetime(2026, 1, 2, 3, 4, 5, tzinfo=UTC) - trigger = SimpleNamespace( - id="trigger-1", - trigger_type="trigger-plugin", + response = workflow_trigger_module.WorkflowTriggerResponse( + id=str(uuid.uuid4()), + trigger_type=AppTriggerType.TRIGGER_PLUGIN, title="Trigger", node_id="node-1", provider_name="provider", icon="https://example.com/icon", - status="enabled", + status=AppTriggerStatus.ENABLED, created_at=created_at, updated_at=created_at, ) - payload = workflow_trigger_module.WorkflowTriggerResponse.model_validate(trigger, from_attributes=True).model_dump( - mode="json" - ) - assert payload["id"] == "trigger-1" + payload = response.model_dump(mode="json") + assert payload["id"] == response.id assert payload["created_at"] == "2026-01-02T03:04:05Z" assert payload["updated_at"] == "2026-01-02T03:04:05Z" @@ -59,51 +107,33 @@ def test_webhook_trigger_response_serializes_datetime(): assert payload["created_at"] == "2026-01-02T03:04:05Z" -def test_app_triggers_get_uses_injected_tenant_id(app: Flask) -> None: - trigger = SimpleNamespace( - id="trigger-1", - trigger_type="trigger-plugin", - title="Trigger", +def test_app_triggers_get_uses_injected_tenant_id(app: Flask, database_session: Session) -> None: + app_model, trigger = _persist_app_trigger(database_session) + other_tenant_trigger = AppTrigger( + tenant_id=str(uuid.uuid4()), + app_id=app_model.id, node_id="node-1", + trigger_type=AppTriggerType.TRIGGER_PLUGIN, + title="Other Tenant Trigger", provider_name="provider", - icon="", - status="enabled", - created_at=None, - updated_at=None, + status=AppTriggerStatus.ENABLED, ) - session = MagicMock() - session.execute.return_value.scalars.return_value.all.return_value = [trigger] + database_session.add(other_tenant_trigger) + database_session.commit() api = workflow_trigger_module.AppTriggersApi() method = inspect.unwrap(api.get) - with ( - app.test_request_context("/"), - patch.object(type(workflow_trigger_module.db), "engine", new_callable=PropertyMock, return_value=MagicMock()), - patch("controllers.console.app.workflow_trigger.sessionmaker") as sessionmaker_mock, - ): - sessionmaker_mock.return_value.begin.return_value.__enter__.return_value = session - response = method(api, "tenant-1", SimpleNamespace(id="app-1")) + with app.test_request_context("/"): + response = method(api, app_model.tenant_id, app_model) - assert response["data"][0]["id"] == "trigger-1" + assert [item["id"] for item in response["data"]] == [trigger.id] assert response["data"][0]["icon"].endswith("/provider/icon") -def test_app_trigger_enable_uses_injected_tenant_id(app: Flask) -> None: - trigger = SimpleNamespace( - id="trigger-1", - trigger_type="trigger-plugin", - title="Trigger", - node_id="node-1", - provider_name="provider", - icon="", - status="disabled", - created_at=None, - updated_at=None, - ) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = trigger - payload = {"trigger_id": "trigger-1", "enable_trigger": True} +def test_app_trigger_enable_uses_injected_tenant_id(app: Flask, database_session: Session) -> None: + app_model, trigger = _persist_app_trigger(database_session, status=AppTriggerStatus.DISABLED) + payload = {"trigger_id": trigger.id, "enable_trigger": True} api = workflow_trigger_module.AppTriggerEnableApi() method = inspect.unwrap(api.post) @@ -111,11 +141,17 @@ def test_app_trigger_enable_uses_injected_tenant_id(app: Flask) -> None: with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), - patch.object(type(workflow_trigger_module.db), "engine", new_callable=PropertyMock, return_value=MagicMock()), - patch("controllers.console.app.workflow_trigger.sessionmaker") as sessionmaker_mock, ): - sessionmaker_mock.return_value.begin.return_value.__enter__.return_value = session - response = method(api, "tenant-1", SimpleNamespace(id="app-1")) + response = method( + api, + workflow_trigger_module.ParserEnable(trigger_id=trigger.id, enable_trigger=True), + app_model.tenant_id, + app_model, + ) - assert response["id"] == "trigger-1" + assert response["id"] == trigger.id assert response["status"] == "enabled" + database_session.expire_all() + persisted_trigger = database_session.get(AppTrigger, trigger.id) + assert persisted_trigger is not None + assert persisted_trigger.status == AppTriggerStatus.ENABLED diff --git a/api/tests/unit_tests/controllers/console/app/test_wraps.py b/api/tests/unit_tests/controllers/console/app/test_wraps.py index a92a9dee433..2f94aedaf52 100644 --- a/api/tests/unit_tests/controllers/console/app/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/app/test_wraps.py @@ -7,7 +7,6 @@ from unittest.mock import MagicMock from uuid import uuid4 import pytest -from sqlalchemy import Select from sqlalchemy.orm import Session from controllers.common import session as session_module @@ -33,7 +32,6 @@ def _persist_app(sqlite_session: Session, *, mode: AppMode = AppMode.CHAT) -> Ap return app_model -@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) def test_get_app_model_injects_model(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: app_model = _persist_app(sqlite_session) monkeypatch.setattr(wraps_module, "current_account_with_tenant", lambda: (None, app_model.tenant_id)) @@ -46,7 +44,6 @@ def test_get_app_model_injects_model(monkeypatch: pytest.MonkeyPatch, sqlite_ses assert handler(app_id=app_model.id) == app_model.id -@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) def test_get_app_model_rejects_wrong_mode(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: app_model = _persist_app(sqlite_session) monkeypatch.setattr(wraps_module, "current_account_with_tenant", lambda: (None, app_model.tenant_id)) @@ -60,17 +57,10 @@ def test_get_app_model_rejects_wrong_mode(monkeypatch: pytest.MonkeyPatch, sqlit handler(app_id=app_model.id) -def test_get_app_model_with_trial_requires_trial_app_registration(monkeypatch: pytest.MonkeyPatch) -> None: - app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1") - session = MagicMock(spec=Session) - - def scalar(statement: Select[tuple[App]]) -> object | None: - has_trial_app_join = any( - from_clause.is_derived_from(TrialApp.__table__) for from_clause in statement.get_final_froms() - ) - return None if has_trial_app_join else app_model - - monkeypatch.setattr(session, "scalar", scalar) +def test_get_app_model_with_trial_requires_trial_app_registration( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + app_model = _persist_app(sqlite_session) recommended_get_app = MagicMock(return_value=None) monkeypatch.setattr(wraps_module.RecommendedAppService, "get_app", recommended_get_app) @@ -80,14 +70,15 @@ def test_get_app_model_with_trial_requires_trial_app_registration(monkeypatch: p return app_model.id with pytest.raises(AppNotFoundError): - Handler().get(session, app_id="app-1") + Handler().get(sqlite_session, app_id=app_model.id) - recommended_get_app.assert_called_once_with("app-1", session=session) + recommended_get_app.assert_called_once_with(app_model.id, session=sqlite_session) -def test_get_app_model_with_trial_falls_back_to_recommended_app(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_app_model_with_trial_falls_back_to_recommended_app( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1") - session = MagicMock(spec=Session) trial_app_loader = MagicMock(return_value=None) recommended_get_app = MagicMock(return_value=app_model) monkeypatch.setattr(wraps_module, "_load_app_model_with_trial", trial_app_loader) @@ -98,14 +89,15 @@ def test_get_app_model_with_trial_falls_back_to_recommended_app(monkeypatch: pyt def get(self, _injected_session, app_model): return app_model.id - assert Handler().get(session, app_id="app-1") == "app-1" - trial_app_loader.assert_called_once_with(session, "app-1") - recommended_get_app.assert_called_once_with("app-1", session=session) + assert Handler().get(unbound_session, app_id="app-1") == "app-1" + trial_app_loader.assert_called_once_with(unbound_session, "app-1") + recommended_get_app.assert_called_once_with("app-1", session=unbound_session) -def test_get_app_model_with_trial_prefers_trial_registration(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_app_model_with_trial_prefers_trial_registration( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal", tenant_id="t1") - session = MagicMock(spec=Session) trial_app_loader = MagicMock(return_value=app_model) recommended_get_app = MagicMock() monkeypatch.setattr(wraps_module, "_load_app_model_with_trial", trial_app_loader) @@ -116,8 +108,8 @@ def test_get_app_model_with_trial_prefers_trial_registration(monkeypatch: pytest def get(self, _injected_session, app_model): return app_model.id - assert Handler().get(session, app_id="app-1") == "app-1" - trial_app_loader.assert_called_once_with(session, "app-1") + assert Handler().get(unbound_session, app_id="app-1") == "app-1" + trial_app_loader.assert_called_once_with(unbound_session, "app-1") recommended_get_app.assert_not_called() @@ -134,7 +126,6 @@ def test_wraps_with_session_reexports_common_session_decorator() -> None: assert wraps_module.with_session is with_session -@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) def test_get_app_model_prefers_injected_session( monkeypatch: pytest.MonkeyPatch, sqlite_session: Session, @@ -154,26 +145,27 @@ def test_get_app_model_prefers_injected_session( assert Handler().get(sqlite_session, app_id=app_model.id) == app_model.id -def test_get_app_model_with_trial_prefers_injected_session(monkeypatch: pytest.MonkeyPatch) -> None: - app_model = SimpleNamespace(id="app-1", mode=AppMode.CHAT.value, status="normal") - session = MagicMock(spec=Session) - session.scalar.return_value = app_model +def test_get_app_model_with_trial_prefers_injected_session( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + app_model = _persist_app(sqlite_session) + sqlite_session.add(TrialApp(app_id=app_model.id, tenant_id=app_model.tenant_id)) + sqlite_session.commit() monkeypatch.setattr( wraps_module.db, "session", SimpleNamespace(scalar=lambda *_args, **_kwargs: pytest.fail("db.session should not be used")), ) - monkeypatch.setattr(session_module.session_factory, "create_session", lambda: nullcontext(session)) + monkeypatch.setattr(session_module.session_factory, "create_session", lambda: nullcontext(sqlite_session)) class Handler: @with_session(write=False) @wraps_module.get_app_model_with_trial(None) def get(self, injected_session, app_model): - assert injected_session is session + assert injected_session is sqlite_session return app_model.id - assert Handler().get(app_id="app-1") == "app-1" - session.scalar.assert_called_once() + assert Handler().get(app_id=app_model.id) == app_model.id def test_get_app_model_with_trial_requires_injected_session() -> None: diff --git a/api/tests/unit_tests/controllers/console/auth/test_account_activation.py b/api/tests/unit_tests/controllers/console/auth/test_account_activation.py index 7669eed8d2a..e9afc527bcd 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_account_activation.py +++ b/api/tests/unit_tests/controllers/console/auth/test_account_activation.py @@ -1,443 +1,311 @@ -""" -Test suite for account activation flows. +"""SQLite-backed tests for account invitation and activation flows.""" -This module tests the account activation mechanism including: -- Invitation token validation -- Account activation with user preferences -- Workspace member onboarding -- Initial login after activation -""" +from __future__ import annotations -from unittest.mock import ANY, MagicMock, patch +from types import SimpleNamespace +from unittest.mock import ANY, Mock, patch import pytest from flask import Flask +from sqlalchemy import func, select +from sqlalchemy.orm import Session, scoped_session +from controllers.console.auth import activate as activate_module 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 +from enums import DeploymentEdition +from models.account import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole + + +@pytest.fixture +def app(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Flask: + session_proxy = scoped_session(lambda: sqlite_session) + monkeypatch.setattr(activate_module, "db", SimpleNamespace(session=session_proxy)) + app = Flask(__name__) + app.config["TESTING"] = True + return app + + +@pytest.fixture +def invitation(sqlite_session: Session) -> dict[str, object]: + account = Account(name="Invited user", email="invitee@example.com", status=AccountStatus.PENDING) + account.id = "account-123" + tenant = Tenant(name="Test Workspace") + tenant.id = "workspace-123" + sqlite_session.add_all([account, tenant]) + sqlite_session.commit() + return { + "data": {"email": account.email}, + "tenant": tenant, + "account": account, + } + + +@pytest.fixture +def switch_tenant(monkeypatch: pytest.MonkeyPatch) -> Mock: + switch = Mock() + monkeypatch.setattr(activate_module.TenantService, "switch_tenant", switch) + return switch + + +def _post(app: Flask, payload: dict[str, object]) -> dict[str, str]: + with app.test_request_context("/activate", method="POST", json=payload): + return ActivateApi().post() + + +def _setup_payload(**overrides: object) -> dict[str, object]: + payload: dict[str, object] = { + "workspace_id": "workspace-123", + "email": "invitee@example.com", + "token": "valid_token", + "name": "John Doe", + "interface_language": "en-US", + "timezone": "UTC", + } + payload.update(overrides) + return payload class TestActivateCheckApi: - """Test cases for checking activation token validity.""" - - @pytest.fixture - def app(self): - """Create Flask test application.""" - app = Flask(__name__) - app.config["TESTING"] = True - return app - - @pytest.fixture - def mock_invitation(self): - """Create mock invitation object.""" - tenant = MagicMock() - tenant.id = "workspace-123" - tenant.name = "Test Workspace" - - return { - "data": {"email": "invitee@example.com"}, - "tenant": tenant, - } - - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - def test_check_valid_invitation_token(self, mock_get_invitation: MagicMock, app: Flask, mock_invitation: MagicMock): - """ - Test checking valid invitation token. - - Verifies that: - - Valid token returns invitation data - - Workspace information is included - - Invitee email is returned - """ - # Arrange - mock_get_invitation.return_value = mock_invitation - - # Act - with app.test_request_context( - "/activate/check?workspace_id=workspace-123&email=invitee@example.com&token=valid_token" + def test_check_valid_invitation_token(self, app: Flask, invitation: dict[str, object]) -> None: + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + app.test_request_context( + "/activate/check?workspace_id=workspace-123&email=invitee@example.com&token=valid_token" + ), ): - api = ActivateCheckApi() - response = api.get() + response = ActivateCheckApi().get() - # Assert assert response["is_valid"] is True assert response["data"]["workspace_name"] == "Test Workspace" assert response["data"]["workspace_id"] == "workspace-123" assert response["data"]["email"] == "invitee@example.com" - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - def test_check_valid_invitation_token_includes_account_status( - self, mock_get_invitation: MagicMock, app: Flask, mock_invitation: MagicMock - ): - mock_account = MagicMock() - mock_account.status = AccountStatus.ACTIVE - mock_invitation["account"] = mock_account - mock_get_invitation.return_value = mock_invitation + def test_check_includes_persisted_account_status(self, app: Flask, invitation: dict[str, object]) -> None: + account = invitation["account"] + assert isinstance(account, Account) + account.status = AccountStatus.ACTIVE - with app.test_request_context("/activate/check?email=invitee@example.com&token=valid_token"): + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + app.test_request_context("/activate/check?email=invitee@example.com&token=valid_token"), + ): response = ActivateCheckApi().get() - assert response["is_valid"] is True assert response["data"]["account_status"] == AccountStatus.ACTIVE assert response["data"]["requires_setup"] is False - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - def test_check_invalid_invitation_token(self, mock_get_invitation, app: Flask): - """ - Test checking invalid invitation token. - - Verifies that: - - Invalid token returns is_valid as False - - No data is returned for invalid tokens - """ - # Arrange - mock_get_invitation.return_value = None - - # Act - with app.test_request_context( - "/activate/check?workspace_id=workspace-123&email=test@example.com&token=invalid_token" + def test_check_invalid_invitation_token(self, app: Flask) -> None: + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=None, + ), + app.test_request_context("/activate/check?email=test@example.com&token=invalid_token"), ): - api = ActivateCheckApi() - response = api.get() + assert ActivateCheckApi().get() == {"is_valid": False} - # Assert - assert response["is_valid"] is False - - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - def test_check_token_without_workspace_id( - self, mock_get_invitation: MagicMock, app: Flask, mock_invitation: MagicMock - ): - """ - Test checking token without workspace ID. - - Verifies that: - - Token can be checked without workspace_id parameter - - System handles None workspace_id gracefully - """ - # Arrange - mock_get_invitation.return_value = mock_invitation - - # Act - with app.test_request_context("/activate/check?email=invitee@example.com&token=valid_token"): - api = ActivateCheckApi() - response = api.get() - - # Assert - assert response["is_valid"] is True - mock_get_invitation.assert_called_once_with(None, "invitee@example.com", "valid_token", session=ANY) - - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - def test_check_token_without_email(self, mock_get_invitation: MagicMock, app: Flask, mock_invitation): - """ - Test checking token without email parameter. - - Verifies that: - - Token can be checked without email parameter - - System handles None email gracefully - """ - # Arrange - mock_get_invitation.return_value = mock_invitation - - # Act - with app.test_request_context("/activate/check?workspace_id=workspace-123&token=valid_token"): - api = ActivateCheckApi() - response = api.get() - - # Assert - assert response["is_valid"] is True - mock_get_invitation.assert_called_once_with("workspace-123", None, "valid_token", session=ANY) - - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - def test_check_token_normalizes_email_to_lowercase( - self, mock_get_invitation: MagicMock, app: Flask, mock_invitation: MagicMock - ): - """Ensure token validation uses lowercase emails.""" - mock_get_invitation.return_value = mock_invitation - - with app.test_request_context( - "/activate/check?workspace_id=workspace-123&email=Invitee@Example.com&token=valid_token" + @pytest.mark.parametrize( + ("query", "workspace_id", "email"), + [ + ("email=invitee@example.com&token=valid_token", None, "invitee@example.com"), + ("workspace_id=workspace-123&token=valid_token", "workspace-123", None), + ( + "workspace_id=workspace-123&email=Invitee@Example.com&token=valid_token", + "workspace-123", + "Invitee@Example.com", + ), + ], + ) + def test_check_forwards_optional_lookup_fields( + self, + app: Flask, + invitation: dict[str, object], + query: str, + workspace_id: str | None, + email: str | None, + ) -> None: + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ) as lookup, + app.test_request_context(f"/activate/check?{query}"), ): - api = ActivateCheckApi() - response = api.get() + assert ActivateCheckApi().get()["is_valid"] is True - assert response["is_valid"] is True - mock_get_invitation.assert_called_once_with("workspace-123", "Invitee@Example.com", "valid_token", session=ANY) + lookup.assert_called_once_with(workspace_id, email, "valid_token", session=ANY) + assert isinstance(lookup.call_args.kwargs["session"], Session) class TestActivateApi: - """Test cases for account activation endpoint.""" - - @pytest.fixture - def app(self): - """Create Flask test application.""" - app = Flask(__name__) - app.config["TESTING"] = True - return app - - @pytest.fixture - def mock_account(self): - """Create mock account object.""" - account = MagicMock() - account.id = "account-123" - account.email = "invitee@example.com" - account.status = AccountStatus.PENDING - return account - - @pytest.fixture - def mock_invitation(self, mock_account): - """Create mock invitation with account.""" - tenant = MagicMock() - tenant.id = "workspace-123" - tenant.name = "Test Workspace" - - return { - "data": {"email": "invitee@example.com"}, - "tenant": tenant, - "account": mock_account, - } - - @pytest.fixture(autouse=True) - def mock_switch_tenant(self): - 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, + sqlite_session: Session, app: Flask, - mock_invitation: MagicMock, - mock_account: MagicMock, - mock_switch_tenant: MagicMock, - ): + invitation: dict[str, object], + switch_tenant: Mock, + ) -> None: """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") + invited_account = invitation["account"] + assert isinstance(invited_account, Account) + invited_account.status = AccountStatus.ACTIVE + data = invitation["data"] + assert isinstance(data, dict) + data["requires_setup"] = False + sqlite_session.commit() + current_account = Mock(id="current-account-id") - with app.test_request_context( - "/activate", - method="POST", - json={ - "token": "valid_token", - }, + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + patch.object( + activate_module, + "current_account_with_tenant", + return_value=(current_account, "current-workspace-id"), + ), + patch.object(activate_module, "extract_access_token", return_value="access-token") as extract_access_token, + patch.object(activate_module.RegisterService, "revoke_token") as revoke_token, + patch.object(activate_module.TenantService, "create_tenant_member") as create_tenant_member, + pytest.raises(InvitationAccountMismatchError), ): - with pytest.raises(InvitationAccountMismatchError): - ActivateApi().post() + _post(app, {"token": "valid_token"}) - 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() + extract_access_token.assert_called_once() + revoke_token.assert_not_called() + create_tenant_member.assert_not_called() + switch_tenant.assert_not_called() + assert sqlite_session.scalar(select(func.count(TenantAccountJoin.id))) == 0 - @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") - def test_successful_account_activation( + def test_successful_account_activation_persists_membership( self, - mock_db: MagicMock, - mock_revoke_token: MagicMock, - mock_get_invitation: MagicMock, + sqlite_session: Session, app: Flask, - mock_invitation: MagicMock, - mock_account, - ): - """ - Test successful account activation. + invitation: dict[str, object], + ) -> None: + invited_account = invitation["account"] + invited_tenant = invitation["tenant"] + assert isinstance(invited_account, Account) + assert isinstance(invited_tenant, Tenant) + account_id = invited_account.id + tenant_id = invited_tenant.id - Verifies that: - - Account is activated with user preferences - - Account status is set to ACTIVE - - Invitation token is revoked - """ - # Arrange - mock_get_invitation.return_value = mock_invitation - - # Act - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "valid_token", - "name": "John Doe", - "interface_language": "en-US", - "timezone": "UTC", - }, + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + patch.object(activate_module.RegisterService, "revoke_token") as revoke_token, ): - api = ActivateApi() - response = api.post() + response = _post(app, _setup_payload()) - # Assert - assert response["result"] == "success" - assert mock_account.name == "John Doe" - assert mock_account.interface_language == "en-US" - assert mock_account.timezone == "UTC" - assert mock_account.status == AccountStatus.ACTIVE - assert mock_account.initialized_at is not None - mock_revoke_token.assert_called_once_with("workspace-123", "invitee@example.com", "valid_token") + sqlite_session.expire_all() + account = sqlite_session.get(Account, account_id) + assert account is not None + membership = sqlite_session.scalar( + select(TenantAccountJoin).where( + TenantAccountJoin.account_id == account_id, + TenantAccountJoin.tenant_id == tenant_id, + ) + ) + assert membership is not None + assert membership.role == TenantAccountRole.NORMAL + assert membership.current is True + assert membership.last_opened_at is not None + assert response == {"result": "success"} + assert account.name == "John Doe" + assert account.interface_language == "en-US" + assert account.timezone == "UTC" + assert account.interface_theme == "light" + assert account.status == AccountStatus.ACTIVE + assert account.initialized_at is not None + revoke_token.assert_called_once_with("workspace-123", "invitee@example.com", "valid_token") - @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.db") - def test_activation_rejects_missing_setup_fields_before_consuming_invitation( + def test_missing_setup_fields_does_not_consume_invitation_or_create_membership( self, - mock_db, - mock_revoke_token, - mock_get_invitation, - mock_create_tenant_member, + sqlite_session: Session, app: Flask, - mock_invitation, - mock_switch_tenant, - ): - mock_invitation["data"]["requires_setup"] = True - mock_get_invitation.return_value = mock_invitation - mock_db.session.scalar.return_value = None + invitation: dict[str, object], + switch_tenant: Mock, + ) -> None: + data = invitation["data"] + assert isinstance(data, dict) + data["requires_setup"] = True - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "valid_token", - }, + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + patch.object(activate_module.RegisterService, "revoke_token") as revoke_token, + pytest.raises(AlreadyActivateError), ): - with pytest.raises(AlreadyActivateError): - ActivateApi().post() + _post(app, _setup_payload(name=None, interface_language=None, timezone=None)) - mock_revoke_token.assert_not_called() - mock_create_tenant_member.assert_not_called() - mock_switch_tenant.assert_not_called() + assert sqlite_session.scalar(select(func.count(TenantAccountJoin.id))) == 0 + revoke_token.assert_not_called() + switch_tenant.assert_not_called() - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - def test_activation_with_invalid_token(self, mock_get_invitation, app: Flask): - """ - Test account activation with invalid token. - - Verifies that: - - AlreadyActivateError is raised for invalid tokens - - No account changes are made - """ - # Arrange - mock_get_invitation.return_value = None - - # Act & Assert - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "invalid_token", - "name": "John Doe", - "interface_language": "en-US", - "timezone": "UTC", - }, + def test_activation_with_invalid_token(self, app: Flask, switch_tenant: Mock) -> None: + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=None, + ), + pytest.raises(AlreadyActivateError), ): - api = ActivateApi() - with pytest.raises(AlreadyActivateError): - api.post() + _post(app, _setup_payload(token="invalid_token")) + switch_tenant.assert_not_called() - @patch("controllers.console.auth.activate.dify_config.BILLING_ENABLED", True) - @patch("controllers.console.auth.activate.BillingService.is_email_in_freeze") - @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") - def test_activation_rejects_account_in_billing_freeze( + def test_billing_freeze_leaves_persisted_account_pending( self, - mock_db, - mock_revoke_token, - mock_get_invitation, - mock_is_email_in_freeze, + sqlite_session: Session, app: Flask, - mock_invitation, - mock_account, - ): - """Frozen deleted-account emails cannot be reactivated through invitation links.""" - mock_account.email = "Invitee@Example.com" - mock_get_invitation.return_value = mock_invitation - mock_is_email_in_freeze.return_value = True + invitation: dict[str, object], + switch_tenant: Mock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + account = invitation["account"] + assert isinstance(account, Account) + account.email = "Invitee@Example.com" + sqlite_session.commit() + monkeypatch.setattr(activate_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "valid_token", - "name": "John Doe", - "interface_language": "en-US", - "timezone": "UTC", - }, + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + patch.object(activate_module.RegisterService, "revoke_token") as revoke_token, + patch.object(activate_module.BillingService, "is_email_in_freeze", return_value=True) as is_frozen, + pytest.raises(AccountInFreezeError), ): - api = ActivateApi() - with pytest.raises(AccountInFreezeError): - api.post() + _post(app, _setup_payload()) - mock_is_email_in_freeze.assert_called_once_with("Invitee@Example.com") - mock_revoke_token.assert_not_called() - mock_db.session.commit.assert_not_called() - assert mock_account.status == AccountStatus.PENDING - - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - @patch("controllers.console.auth.activate.RegisterService.revoke_token") - @patch("controllers.console.auth.activate.db") - def test_activation_sets_interface_theme( - self, - mock_db, - mock_revoke_token, - mock_get_invitation, - app: Flask, - mock_invitation, - mock_account, - ): - """ - Test that activation sets default interface theme. - - Verifies that: - - Interface theme is set to 'light' by default - """ - # Arrange - mock_get_invitation.return_value = mock_invitation - - # Act - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "valid_token", - "name": "John Doe", - "interface_language": "en-US", - "timezone": "UTC", - }, - ): - api = ActivateApi() - api.post() - - # Assert - assert mock_account.interface_theme == "light" + sqlite_session.refresh(account) + assert account.status == AccountStatus.PENDING + assert sqlite_session.scalar(select(func.count(TenantAccountJoin.id))) == 0 + is_frozen.assert_called_once_with("Invitee@Example.com") + revoke_token.assert_not_called() + switch_tenant.assert_not_called() @pytest.mark.parametrize( ("language", "timezone"), @@ -448,236 +316,136 @@ class TestActivateApi: ("es-ES", "Europe/Madrid"), ], ) - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - @patch("controllers.console.auth.activate.RegisterService.revoke_token") - @patch("controllers.console.auth.activate.db") def test_activation_with_different_locales( self, - mock_db, - mock_revoke_token, - mock_get_invitation, app: Flask, - mock_invitation, - mock_account, - language, - timezone, - ): - """ - Test account activation with various language and timezone combinations. - - Verifies that: - - Different languages are accepted - - Different timezones are accepted - - User preferences are properly stored - """ - # Arrange - mock_get_invitation.return_value = mock_invitation - - # Act - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "valid_token", - "name": "Test User", - "interface_language": language, - "timezone": timezone, - }, + invitation: dict[str, object], + switch_tenant: Mock, + language: str, + timezone: str, + ) -> None: + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + patch.object(activate_module.RegisterService, "revoke_token"), ): - api = ActivateApi() - response = api.post() + assert _post(app, _setup_payload(interface_language=language, timezone=timezone)) == {"result": "success"} - # Assert - assert response["result"] == "success" - assert mock_account.interface_language == language - assert mock_account.timezone == timezone + account = invitation["account"] + assert isinstance(account, Account) + assert account.interface_language == language + assert account.timezone == timezone - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - @patch("controllers.console.auth.activate.RegisterService.revoke_token") - @patch("controllers.console.auth.activate.db") - def test_activation_returns_success_response( + def test_activation_without_workspace_id_revokes_normalized_email( self, - mock_db: MagicMock, - mock_revoke_token: MagicMock, - mock_get_invitation: MagicMock, app: Flask, - mock_invitation: MagicMock, - ): - """ - Test that activation returns a success response without authentication tokens. - - Verifies that: - - Response contains a success result - - No token data is returned - """ - # Arrange - mock_get_invitation.return_value = mock_invitation - - # Act - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "valid_token", - "name": "John Doe", - "interface_language": "en-US", - "timezone": "UTC", - }, + invitation: dict[str, object], + switch_tenant: Mock, + ) -> None: + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ) as lookup, + patch.object(activate_module.RegisterService, "revoke_token") as revoke_token, ): - api = ActivateApi() - response = api.post() + response = _post( + app, + _setup_payload(workspace_id=None, email="Invitee@Example.com"), + ) - # Assert assert response == {"result": "success"} + lookup.assert_called_once_with(None, "Invitee@Example.com", "valid_token", session=ANY) + revoke_token.assert_called_once_with(None, "invitee@example.com", "valid_token") - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - @patch("controllers.console.auth.activate.RegisterService.revoke_token") - @patch("controllers.console.auth.activate.db") - def test_activation_without_workspace_id( + def test_existing_active_account_gets_tenant_scoped_admin_membership( self, - mock_db: MagicMock, - mock_revoke_token: MagicMock, - mock_get_invitation: MagicMock, + sqlite_session: Session, app: Flask, - mock_invitation: MagicMock, - ): - """ - Test account activation without workspace_id. - - Verifies that: - - Activation can proceed without workspace_id - - Token revocation handles None workspace_id - """ - # Arrange - mock_get_invitation.return_value = mock_invitation - - # Act - with app.test_request_context( - "/activate", - method="POST", - json={ - "email": "invitee@example.com", - "token": "valid_token", - "name": "John Doe", - "interface_language": "en-US", - "timezone": "UTC", - }, - ): - api = ActivateApi() - response = api.post() - - # Assert - assert response["result"] == "success" - mock_revoke_token.assert_called_once_with(None, "invitee@example.com", "valid_token") - - @patch("controllers.console.auth.activate.RegisterService.get_invitation_with_case_fallback") - @patch("controllers.console.auth.activate.RegisterService.revoke_token") - @patch("controllers.console.auth.activate.db") - def test_activation_normalizes_email_before_lookup( - self, - mock_db: MagicMock, - mock_revoke_token: MagicMock, - mock_get_invitation: MagicMock, - app: Flask, - mock_invitation: MagicMock, - mock_account: MagicMock, - ): - """Ensure uppercase emails are normalized before lookup and revocation.""" - mock_get_invitation.return_value = mock_invitation - - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "Invitee@Example.com", - "token": "valid_token", - "name": "John Doe", - "interface_language": "en-US", - "timezone": "UTC", - }, - ): - api = ActivateApi() - response = api.post() - - assert response["result"] == "success" - mock_get_invitation.assert_called_once_with("workspace-123", "Invitee@Example.com", "valid_token", session=ANY) - mock_revoke_token.assert_called_once_with("workspace-123", "invitee@example.com", "valid_token") - - @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.db") - def test_activation_for_existing_active_account_creates_membership_on_acceptance( - self, - mock_db: 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, - ): - mock_account.status = AccountStatus.ACTIVE - mock_invitation["data"]["role"] = "admin" - mock_invitation["data"]["requires_setup"] = False - mock_get_invitation.return_value = mock_invitation - mock_db.session.scalar.return_value = None - - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "valid_token", - }, - ): - response = ActivateApi().post() - - assert response["result"] == "success" - mock_create_tenant_member.assert_called_once_with( - mock_invitation["tenant"], mock_account, mock_db.session(), role=TenantAccountRole.ADMIN + invitation: dict[str, object], + switch_tenant: Mock, + ) -> None: + session = sqlite_session + account = invitation["account"] + tenant = invitation["tenant"] + data = invitation["data"] + assert isinstance(account, Account) + assert isinstance(tenant, Tenant) + assert isinstance(data, dict) + account.status = AccountStatus.ACTIVE + data.update({"role": "admin", "requires_setup": False}) + other_tenant = Tenant(name="Other Workspace") + other_tenant.id = "workspace-456" + session.add(other_tenant) + session.flush() + session.add( + TenantAccountJoin( + tenant_id=other_tenant.id, + account_id=account.id, + role=TenantAccountRole.NORMAL, + ) ) - mock_switch_tenant.assert_called_once_with(mock_account, mock_invitation["tenant"].id, session=ANY) - mock_revoke_token.assert_called_once_with("workspace-123", "invitee@example.com", "valid_token") + session.commit() - @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.db") - def test_activation_legacy_active_member_invitation_does_not_require_setup( - self, - mock_db: 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, - ): - mock_account.status = AccountStatus.ACTIVE - mock_get_invitation.return_value = mock_invitation - mock_db.session.scalar.return_value = "membership-id" - - with app.test_request_context( - "/activate", - method="POST", - json={ - "workspace_id": "workspace-123", - "email": "invitee@example.com", - "token": "valid_token", - }, + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + patch.object(activate_module.RegisterService, "revoke_token"), ): - response = ActivateApi().post() + assert _post( + app, + {"workspace_id": tenant.id, "email": account.email, "token": "valid_token"}, + ) == {"result": "success"} - assert response["result"] == "success" - mock_create_tenant_member.assert_not_called() - mock_switch_tenant.assert_called_once_with(mock_account, mock_invitation["tenant"].id, session=ANY) - mock_revoke_token.assert_called_once_with("workspace-123", "invitee@example.com", "valid_token") + memberships = session.scalars(select(TenantAccountJoin).where(TenantAccountJoin.account_id == account.id)).all() + assert {(row.tenant_id, row.role) for row in memberships} == { + (other_tenant.id, TenantAccountRole.NORMAL), + (tenant.id, TenantAccountRole.ADMIN), + } + + def test_existing_membership_is_not_duplicated( + self, + sqlite_session: Session, + app: Flask, + invitation: dict[str, object], + switch_tenant: Mock, + ) -> None: + session = sqlite_session + account = invitation["account"] + tenant = invitation["tenant"] + assert isinstance(account, Account) + assert isinstance(tenant, Tenant) + account.status = AccountStatus.ACTIVE + session.add( + TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + role=TenantAccountRole.EDITOR, + ) + ) + session.commit() + + with ( + patch.object( + activate_module.RegisterService, + "get_invitation_with_case_fallback", + return_value=invitation, + ), + patch.object(activate_module.RegisterService, "revoke_token"), + ): + assert _post( + app, + {"workspace_id": tenant.id, "email": account.email, "token": "valid_token"}, + ) == {"result": "success"} + + assert session.scalar(select(func.count(TenantAccountJoin.id))) == 1 + membership = session.scalar(select(TenantAccountJoin)) + assert membership is not None + assert membership.role == TenantAccountRole.EDITOR diff --git a/api/tests/unit_tests/controllers/console/auth/test_authentication_security.py b/api/tests/unit_tests/controllers/console/auth/test_authentication_security.py index 17bee94c520..dc4260e94b1 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_authentication_security.py +++ b/api/tests/unit_tests/controllers/console/auth/test_authentication_security.py @@ -10,6 +10,7 @@ from flask_restx import Api import services.errors.account from controllers.console.auth.error import AuthenticationFailedError from controllers.console.auth.login import LoginApi +from enums import DeploymentEdition def encode_password(password: str) -> str: @@ -33,7 +34,7 @@ class TestAuthenticationSecurity: @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.AccountService.authenticate") @patch("controllers.console.auth.login.AccountService.add_login_error_rate_limit") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") def test_login_invalid_email_with_registration_allowed( self, mock_get_invitation, mock_add_rate_limit, mock_authenticate, mock_is_rate_limit, mock_features, mock_db @@ -65,7 +66,7 @@ class TestAuthenticationSecurity: @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.AccountService.authenticate") @patch("controllers.console.auth.login.AccountService.add_login_error_rate_limit") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") def test_login_wrong_password_returns_error( self, mock_get_invitation, mock_add_rate_limit, mock_authenticate, mock_is_rate_limit, mock_db @@ -97,7 +98,7 @@ class TestAuthenticationSecurity: @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.AccountService.authenticate") @patch("controllers.console.auth.login.AccountService.add_login_error_rate_limit") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") def test_login_invalid_email_with_registration_disabled( self, mock_get_invitation, mock_add_rate_limit, mock_authenticate, mock_is_rate_limit, mock_features, mock_db diff --git a/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py b/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py index b231826aeac..549a78a6915 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py +++ b/api/tests/unit_tests/controllers/console/auth/test_data_source_bearer_auth.py @@ -3,25 +3,16 @@ from __future__ import annotations from datetime import UTC, datetime from inspect import unwrap from types import SimpleNamespace -from unittest.mock import ANY, PropertyMock, patch +from unittest.mock import ANY, patch -from controllers.console import console_ns from controllers.console.auth.data_source_bearer_auth import ( + ApiKeyAuthBindingPayload, ApiKeyAuthDataSource, ApiKeyAuthDataSourceBinding, ApiKeyAuthDataSourceBindingDelete, ) -def _payload_patch(payload: dict): - return patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ) - - def test_list_data_source_auth_uses_injected_tenant_id() -> None: api = ApiKeyAuthDataSource() method = unwrap(api.get) @@ -56,14 +47,14 @@ def test_create_data_source_auth_binding_uses_injected_tenant_id() -> None: "provider": "custom", "credentials": {"auth_type": "api_key", "config": {"api_key": "secret"}}, } + req_data = ApiKeyAuthBindingPayload.model_validate(payload) with ( - _payload_patch(payload), patch("controllers.console.auth.data_source_bearer_auth.db"), patch("controllers.console.auth.data_source_bearer_auth.ApiKeyAuthService.validate_api_key_auth_args"), patch("controllers.console.auth.data_source_bearer_auth.ApiKeyAuthService.create_provider_auth") as create_auth, ): - result, status = method(api, "tenant-1") + result, status = method(api, req_data, "tenant-1") create_auth.assert_called_once_with("tenant-1", payload, session=ANY) assert result == {"result": "success"} diff --git a/api/tests/unit_tests/controllers/console/auth/test_email_register.py b/api/tests/unit_tests/controllers/console/auth/test_email_register.py index 4c0b27d554f..6c945dc5f57 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_email_register.py +++ b/api/tests/unit_tests/controllers/console/auth/test_email_register.py @@ -11,8 +11,8 @@ from controllers.console.auth.email_register import ( EmailRegisterResetApi, EmailRegisterSendEmailApi, ) -from enums.deployment_edition import DeploymentEdition -from services.feature_service import SystemFeatureModel +from enums import DeploymentEdition +from services.entities.feature_entities import SystemFeatureModel class TestEmailRegisterSendEmailApi: @@ -41,8 +41,8 @@ class TestEmailRegisterSendEmailApi: is_allow_register=True, ) with ( - patch("controllers.console.auth.email_register.dify_config.BILLING_ENABLED", True), - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.auth.email_register.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags), ): with app.test_request_context( @@ -86,7 +86,7 @@ class TestEmailRegisterCheckApi: is_allow_register=True, ) with ( - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags), ): with app.test_request_context( @@ -138,7 +138,7 @@ class TestEmailRegisterResetApi: is_allow_register=True, ) with ( - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags), ): with app.test_request_context( @@ -190,7 +190,7 @@ class TestEmailRegisterResetApi: is_allow_register=True, ) with ( - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags), ): with app.test_request_context( @@ -247,7 +247,7 @@ class TestEmailRegisterResetApi: is_allow_register=True, ) with ( - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=feature_flags), ): with app.test_request_context( diff --git a/api/tests/unit_tests/controllers/console/auth/test_email_verification.py b/api/tests/unit_tests/controllers/console/auth/test_email_verification.py index eef39e8d208..4c8173ce80d 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_email_verification.py +++ b/api/tests/unit_tests/controllers/console/auth/test_email_verification.py @@ -15,8 +15,20 @@ import pytest from flask import Flask from pydantic import ValidationError -from controllers.console.auth.error import EmailCodeError, InvalidEmailError, InvalidTokenError -from controllers.console.auth.login import EmailCodeLoginApi, EmailCodeLoginPayload, EmailCodeLoginSendEmailApi +from controllers.console.auth.error import ( + EmailCodeError, + InvalidEmailError, + InvalidTokenError, + TurnstileServiceUnavailableError, + TurnstileVerificationFailedError, +) +from controllers.console.auth.login import ( + EmailCodeLoginApi, + EmailCodeLoginPayload, + EmailCodeLoginSendEmailApi, + EmailCodeSendPayload, + EmailPayload, +) from controllers.console.error import ( AccountInFreezeError, AccountNotFound, @@ -24,7 +36,9 @@ from controllers.console.error import ( NotAllowedCreateWorkspace, WorkspacesLimitExceeded, ) +from enums import DeploymentEdition from services.errors.account import AccountRegisterError +from services.turnstile_service import TurnstileChallengeRejectedError, TurnstileUpstreamError def encode_code(code: str) -> str: @@ -44,6 +58,11 @@ def test_email_code_login_payload_rejects_invalid_timezone(): ) +def test_turnstile_token_is_scoped_to_email_code_send_payload(): + assert "turnstile_token" in EmailCodeSendPayload.model_fields + assert "turnstile_token" not in EmailPayload.model_fields + + class TestEmailCodeLoginSendEmailApi: """Test cases for sending email verification codes.""" @@ -153,7 +172,8 @@ class TestEmailCodeLoginSendEmailApi: @patch("controllers.console.wraps.db") @patch("controllers.console.auth.login.AccountService.is_email_send_ip_limit") - def test_send_email_code_ip_rate_limited(self, mock_is_ip_limit, mock_db, app: Flask): + @patch("controllers.console.auth.login.TurnstileService.verify") + def test_send_email_code_ip_rate_limited(self, mock_verify, mock_is_ip_limit, mock_db, app: Flask): """ Test email code sending blocked by IP rate limit. @@ -165,10 +185,105 @@ class TestEmailCodeLoginSendEmailApi: mock_is_ip_limit.return_value = True # Act & Assert - with app.test_request_context("/email-code-login", method="POST", json={"email": "test@example.com"}): - api = EmailCodeLoginSendEmailApi() + with ( + patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + app.test_request_context("/email-code-login", method="POST", json={"email": "test@example.com"}), + ): with pytest.raises(EmailSendIpLimitError): - api.post() + EmailCodeLoginSendEmailApi().post() + + mock_verify.assert_not_called() + + @patch("controllers.console.wraps.db") + @patch("controllers.console.auth.login.AccountService.is_email_send_ip_limit", return_value=False) + @patch("controllers.console.auth.login.AccountService.get_user_through_email") + @patch("controllers.console.auth.login.AccountService.send_email_code_login_email", return_value="token") + @patch("controllers.console.auth.login.TurnstileService.verify") + def test_cloud_send_verifies_turnstile_before_sending_email( + self, + mock_verify, + mock_send_email, + mock_get_user, + mock_is_ip_limit, + mock_db, + app: Flask, + mock_account, + ): + mock_get_user.return_value = mock_account + + with ( + patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + app.test_request_context( + "/email-code-login", + method="POST", + json={"email": "test@example.com", "turnstile_token": "verified-token"}, + headers={"CF-Connecting-IP": "203.0.113.8"}, + ), + ): + response = EmailCodeLoginSendEmailApi().post() + + assert response["result"] == "success" + mock_verify.assert_called_once_with(token="verified-token", remote_ip="203.0.113.8") + mock_send_email.assert_called_once() + + @pytest.mark.parametrize( + ("service_error", "http_error"), + [ + (TurnstileChallengeRejectedError(), TurnstileVerificationFailedError), + (TurnstileUpstreamError(), TurnstileServiceUnavailableError), + ], + ) + @patch("controllers.console.wraps.db") + @patch("controllers.console.auth.login.AccountService.is_email_send_ip_limit", return_value=False) + @patch("controllers.console.auth.login.AccountService.get_user_through_email") + def test_cloud_send_maps_turnstile_errors_without_looking_up_account( + self, + mock_get_user, + mock_is_ip_limit, + mock_db, + app: Flask, + service_error: Exception, + http_error: type[Exception], + ): + with ( + patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + patch("controllers.console.auth.login.TurnstileService.verify", side_effect=service_error), + app.test_request_context( + "/email-code-login", + method="POST", + json={"email": "test@example.com", "turnstile_token": "challenge-token"}, + ), + pytest.raises(http_error), + ): + EmailCodeLoginSendEmailApi().post() + + mock_get_user.assert_not_called() + + @patch("controllers.console.wraps.db") + @patch("controllers.console.auth.login.AccountService.is_email_send_ip_limit", return_value=False) + @patch("controllers.console.auth.login.AccountService.get_user_through_email") + @patch("controllers.console.auth.login.AccountService.send_email_code_login_email", return_value="token") + @patch("controllers.console.auth.login.TurnstileService.verify") + def test_self_hosted_send_does_not_call_turnstile( + self, + mock_verify, + mock_send_email, + mock_get_user, + mock_is_ip_limit, + mock_db, + app: Flask, + mock_account, + ): + mock_get_user.return_value = mock_account + + with ( + patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), + app.test_request_context("/email-code-login", method="POST", json={"email": "test@example.com"}), + ): + response = EmailCodeLoginSendEmailApi().post() + + assert response["result"] == "success" + mock_verify.assert_not_called() @patch("controllers.console.wraps.db") @patch("controllers.console.auth.login.AccountService.is_email_send_ip_limit") diff --git a/api/tests/unit_tests/controllers/console/auth/test_forgot_password.py b/api/tests/unit_tests/controllers/console/auth/test_forgot_password.py index 48a6b98a1c3..438bd35169f 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_forgot_password.py +++ b/api/tests/unit_tests/controllers/console/auth/test_forgot_password.py @@ -13,10 +13,10 @@ from controllers.console.auth.forgot_password import ( ForgotPasswordResetApi, ForgotPasswordSendEmailApi, ) -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from models.account import Account from models.engine import db -from services.feature_service import SystemFeatureModel +from services.entities.feature_entities import SystemFeatureModel @pytest.fixture @@ -61,7 +61,7 @@ class TestForgotPasswordSendEmailApi: "controllers.console.auth.forgot_password.FeatureService.get_system_features", return_value=controller_features, ), - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features), ): with app.test_request_context( @@ -108,7 +108,7 @@ class TestForgotPasswordCheckApi: enable_email_password_login=True, ) with ( - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features), ): with app.test_request_context( @@ -154,7 +154,7 @@ class TestForgotPasswordResetApi: enable_email_password_login=True, ) with ( - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features), ): with database_app.test_request_context( diff --git a/api/tests/unit_tests/controllers/console/auth/test_login_logout.py b/api/tests/unit_tests/controllers/console/auth/test_login_logout.py index 51428d3d883..970acd52cfa 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_login_logout.py +++ b/api/tests/unit_tests/controllers/console/auth/test_login_logout.py @@ -29,6 +29,7 @@ from controllers.console.error import ( SeatsLimitExceeded, WorkspacesLimitExceeded, ) +from enums import DeploymentEdition from services.entities.auth_entities import LoginFailureReason from services.errors.account import AccountLoginError, AccountPasswordError, SeatsLimitExceededError @@ -86,7 +87,7 @@ class TestLoginApi: return token_pair @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") @patch("controllers.console.auth.login.AccountService.authenticate") @@ -137,7 +138,7 @@ class TestLoginApi: assert response.json["result"] == "success" @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") @patch("controllers.console.auth.login.AccountService.authenticate") @@ -190,7 +191,7 @@ class TestLoginApi: assert response.json["result"] == "success" @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") def test_login_fails_when_rate_limited( @@ -223,7 +224,7 @@ class TestLoginApi: assert warn_records[0].args[1] == LoginFailureReason.LOGIN_RATE_LIMITED @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", True) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) @patch("controllers.console.auth.login.BillingService.is_email_in_freeze") def test_login_fails_when_account_frozen( self, mock_is_frozen, mock_db, app: Flask, caplog: pytest.LogCaptureFixture @@ -254,7 +255,7 @@ class TestLoginApi: assert warn_records[0].args[1] == LoginFailureReason.ACCOUNT_IN_FREEZE @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") @patch("controllers.console.auth.login.AccountService.authenticate") @@ -301,7 +302,7 @@ class TestLoginApi: assert warn_records[0].args[1] == LoginFailureReason.INVALID_CREDENTIALS @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") @patch("controllers.console.auth.login.AccountService.authenticate") @@ -338,7 +339,7 @@ class TestLoginApi: assert warn_records[0].args[1] == LoginFailureReason.ACCOUNT_BANNED @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") @patch("controllers.console.auth.login.AccountService.authenticate") @@ -382,7 +383,7 @@ class TestLoginApi: login_api.post() @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") def test_login_invitation_email_mismatch(self, mock_get_invitation, mock_is_rate_limit, mock_db, app: Flask): @@ -412,7 +413,7 @@ class TestLoginApi: login_api.post() @patch("controllers.console.wraps.db") - @patch("controllers.console.auth.login.dify_config.BILLING_ENABLED", False) + @patch("controllers.console.auth.login.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) @patch("controllers.console.auth.login.AccountService.is_login_error_rate_limit") @patch("controllers.console.auth.login.RegisterService.get_invitation_with_case_fallback") @patch("controllers.console.auth.login.AccountService.authenticate") diff --git a/api/tests/unit_tests/controllers/console/auth/test_oauth_timezone.py b/api/tests/unit_tests/controllers/console/auth/test_oauth_timezone.py index 75545cc27e7..eaf339b62a6 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_oauth_timezone.py +++ b/api/tests/unit_tests/controllers/console/auth/test_oauth_timezone.py @@ -4,6 +4,7 @@ import pytest from flask import Flask from controllers.console.auth.oauth import OAuthLogin, _generate_account +from enums import DeploymentEdition from libs.oauth import OAuthUserInfo from services.errors.account import AccountRegisterError @@ -116,7 +117,7 @@ def test_generate_account_rejects_new_user_when_registration_disabled( app: Flask, ): mock_feature_service.get_system_features.return_value.is_allow_register = False - mock_config.BILLING_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY user_info = OAuthUserInfo(id="github-123", name="Test User", email="user@example.com") with app.test_request_context(headers={"Accept-Language": "en-US,en;q=0.9"}): diff --git a/api/tests/unit_tests/controllers/console/auth/test_password_reset.py b/api/tests/unit_tests/controllers/console/auth/test_password_reset.py index 6da8cb5a51d..c2375fab888 100644 --- a/api/tests/unit_tests/controllers/console/auth/test_password_reset.py +++ b/api/tests/unit_tests/controllers/console/auth/test_password_reset.py @@ -23,9 +23,9 @@ from controllers.console.auth.forgot_password import ( ForgotPasswordSendEmailApi, ) from controllers.console.error import AccountNotFound, EmailSendIpLimitError -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from models.account import Account, Tenant, TenantAccountJoin -from services.feature_service import SystemFeatureModel +from services.entities.feature_entities import SystemFeatureModel SQLITE_MODELS = (Account, Tenant, TenantAccountJoin) @@ -46,7 +46,7 @@ def _bind_database_session(session: Session) -> Generator[scoped_session[Session 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.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr( "controllers.console.wraps.FeatureService.get_system_features", lambda: SystemFeatureModel( diff --git a/api/tests/unit_tests/controllers/console/billing/test_billing.py b/api/tests/unit_tests/controllers/console/billing/test_billing.py index 52a85136723..e9ac3fd7f81 100644 --- a/api/tests/unit_tests/controllers/console/billing/test_billing.py +++ b/api/tests/unit_tests/controllers/console/billing/test_billing.py @@ -1,15 +1,24 @@ import base64 import json -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest from flask import Flask -from werkzeug.exceptions import BadRequest +from sqlalchemy.orm import Session +from werkzeug.exceptions import BadRequest, UnprocessableEntity +from controllers.console import wraps as console_wraps from controllers.console.billing.billing import PartnerTenants -from models.account import Account +from enums import DeploymentEdition +from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole +from models.model import DifySetup +@pytest.mark.parametrize( + "sqlite_session", + [(DifySetup, Account, Tenant, TenantAccountJoin)], + indirect=True, +) class TestPartnerTenants: """Unit tests for PartnerTenants controller.""" @@ -22,13 +31,27 @@ class TestPartnerTenants: return app @pytest.fixture - def mock_account(self): - """Create a mock account.""" - account = MagicMock(spec=Account) - account.id = "account-123" - account.email = "test@example.com" - account.current_tenant_id = "tenant-456" - account.is_authenticated = True + def mock_account(self, sqlite_session: Session): + """Persist an initialized account with an owner workspace membership.""" + tenant = Tenant(name="Billing Tenant") + account = Account(name="Billing User", email="test@example.com") + sqlite_session.add_all([tenant, account]) + sqlite_session.flush() + sqlite_session.add_all( + [ + TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + current=True, + role=TenantAccountRole.OWNER, + invited_by=None, + ), + DifySetup(version="test"), + ] + ) + sqlite_session.commit() + account._current_tenant = tenant + sqlite_session.expunge(account) return account @pytest.fixture @@ -38,16 +61,18 @@ class TestPartnerTenants: yield mock_service @pytest.fixture - def mock_decorators(self): - """Mock decorators to avoid database access.""" + def mock_decorators(self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + """Keep authentication mocked while the setup guard uses SQLite.""" + console_wraps._is_setup_completed.reset_success() + monkeypatch.setattr(console_wraps.db, "session", sqlite_session) with ( - patch("controllers.console.wraps.db") as mock_db, - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("libs.login.dify_config.LOGIN_DISABLED", False), patch("libs.login.check_csrf_token") as mock_csrf, ): mock_csrf.return_value = None - yield {"db": mock_db, "csrf": mock_csrf} + yield mock_csrf + console_wraps._is_setup_completed.reset_success() def test_put_success(self, app: Flask, mock_account, mock_billing_service, mock_decorators): """Test successful partner tenants bindings sync.""" @@ -66,7 +91,7 @@ class TestPartnerTenants: with ( patch( "controllers.console.wraps.current_account_with_tenant", - return_value=(mock_account, "tenant-456"), + return_value=(mock_account, mock_account.current_tenant_id), ), patch("libs.login._get_user", return_value=mock_account), ): @@ -93,7 +118,7 @@ class TestPartnerTenants: with ( patch( "controllers.console.wraps.current_account_with_tenant", - return_value=(mock_account, "tenant-456"), + return_value=(mock_account, mock_account.current_tenant_id), ), patch("libs.login._get_user", return_value=mock_account), ): @@ -105,7 +130,7 @@ class TestPartnerTenants: assert "Invalid partner_key" in str(exc_info.value) def test_put_missing_click_id(self, app: Flask, mock_account, mock_billing_service, mock_decorators): - """Test that missing click_id raises BadRequest.""" + """Test that missing click_id raises UnprocessableEntity (422).""" # Arrange partner_key_encoded = base64.b64encode(b"partner-key-123").decode("utf-8") @@ -117,15 +142,15 @@ class TestPartnerTenants: with ( patch( "controllers.console.wraps.current_account_with_tenant", - return_value=(mock_account, "tenant-456"), + return_value=(mock_account, mock_account.current_tenant_id), ), patch("libs.login._get_user", return_value=mock_account), ): resource = PartnerTenants() # Act & Assert - # Validation should raise BadRequest for missing required field - with pytest.raises(BadRequest): + # Validation should raise UnprocessableEntity (422) for missing required field + with pytest.raises(UnprocessableEntity): resource.put(partner_key_encoded) def test_put_billing_service_json_decode_error( @@ -159,7 +184,7 @@ class TestPartnerTenants: with ( patch( "controllers.console.wraps.current_account_with_tenant", - return_value=(mock_account, "tenant-456"), + return_value=(mock_account, mock_account.current_tenant_id), ), patch("libs.login._get_user", return_value=mock_account), ): @@ -190,7 +215,7 @@ class TestPartnerTenants: with ( patch( "controllers.console.wraps.current_account_with_tenant", - return_value=(mock_account, "tenant-456"), + return_value=(mock_account, mock_account.current_tenant_id), ), patch("libs.login._get_user", return_value=mock_account), ): @@ -216,7 +241,7 @@ class TestPartnerTenants: with ( patch( "controllers.console.wraps.current_account_with_tenant", - return_value=(mock_account, "tenant-456"), + return_value=(mock_account, mock_account.current_tenant_id), ), patch("libs.login._get_user", return_value=mock_account), ): @@ -242,7 +267,7 @@ class TestPartnerTenants: with ( patch( "controllers.console.wraps.current_account_with_tenant", - return_value=(mock_account, "tenant-456"), + return_value=(mock_account, mock_account.current_tenant_id), ), patch("libs.login._get_user", return_value=mock_account), ): diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py index 8f66ca5c993..753236f6dfc 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py @@ -14,9 +14,15 @@ from controllers.console.datasets.rag_pipeline.datasource_auth import ( DatasourceAuthListApi, DatasourceAuthOauthCustomClient, DatasourceAuthUpdateApi, + DatasourceCredentialDeletePayload, + DatasourceCredentialPayload, + DatasourceCredentialUpdatePayload, + DatasourceCustomClientPayload, + DatasourceDefaultPayload, DatasourceHardCodeAuthListApi, DatasourceOAuthCallback, DatasourcePluginOAuthAuthorizationUrl, + DatasourceUpdateNamePayload, DatasourceUpdateProviderNameApi, ) from core.plugin.impl.oauth import OAuthHandler @@ -405,7 +411,8 @@ class TestDatasourceAuth: return_value=None, ) as add_api_key_provider, ): - response, status = method(api, "tenant-1", _PROVIDER_ID) + req_data = DatasourceCredentialPayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", _PROVIDER_ID) assert response == _success_response() assert status == 200 @@ -431,7 +438,7 @@ class TestDatasourceAuth: ), ): with pytest.raises(ValueError): - method(api, "tenant-1", "notion") + method(api, DatasourceCredentialPayload.model_validate(payload), "tenant-1", "notion") def test_get_success(self, app: Flask): api = DatasourceAuth() @@ -462,7 +469,7 @@ class TestDatasourceAuth: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(ValueError): - method(api, "tenant-1", "notion") + method(api, DatasourceCredentialPayload.model_validate(payload), "tenant-1", "notion") def test_get_empty_list(self, app: Flask): api = DatasourceAuth() @@ -499,7 +506,8 @@ class TestDatasourceAuthDeleteApi: return_value=None, ) as remove_datasource_credentials, ): - response, status = method(api, "tenant-1", _PROVIDER_ID) + req_data = DatasourceCredentialDeletePayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", _PROVIDER_ID) assert response == _success_response() assert status == 200 @@ -522,7 +530,7 @@ class TestDatasourceAuthDeleteApi: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(ValueError): - method(api, "tenant-1", "notion") + method(api, DatasourceCredentialDeletePayload.model_validate(payload), "tenant-1", "notion") class TestDatasourceAuthUpdateApi: @@ -545,7 +553,8 @@ class TestDatasourceAuthUpdateApi: return_value=None, ) as update_datasource_credentials, ): - response, status = method(api, "tenant-1", _PROVIDER_ID) + req_data = DatasourceCredentialUpdatePayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", _PROVIDER_ID) assert response == _success_response() assert status == 201 @@ -573,7 +582,8 @@ class TestDatasourceAuthUpdateApi: return_value=None, ) as update_mock, ): - response, status = method(api, "tenant-1", "notion") + req_data = DatasourceCredentialUpdatePayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", "notion") assert response == _success_response() update_mock.assert_called_once() @@ -595,7 +605,8 @@ class TestDatasourceAuthUpdateApi: return_value=None, ), ): - response, status = method(api, "tenant-1", "notion") + req_data = DatasourceCredentialUpdatePayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", "notion") assert response == _success_response() assert status == 201 @@ -615,7 +626,8 @@ class TestDatasourceAuthUpdateApi: return_value=None, ) as update_mock, ): - response, status = method(api, "tenant-1", "notion") + req_data = DatasourceCredentialUpdatePayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", "notion") assert response == _success_response() update_mock.assert_called_once() @@ -716,7 +728,8 @@ class TestDatasourceAuthOauthCustomClient: return_value=None, ) as setup_custom_client, ): - response, status = method(api, "tenant-1", _PROVIDER_ID) + req_data = DatasourceCustomClientPayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", _PROVIDER_ID) assert response == _success_response() assert status == 200 @@ -760,7 +773,8 @@ class TestDatasourceAuthOauthCustomClient: return_value=None, ), ): - response, status = method(api, "tenant-1", "notion") + req_data = DatasourceCustomClientPayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", "notion") assert response == _success_response() assert status == 200 @@ -783,7 +797,8 @@ class TestDatasourceAuthOauthCustomClient: return_value=None, ) as setup_mock, ): - response, status = method(api, "tenant-1", "notion") + req_data = DatasourceCustomClientPayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", "notion") assert response == _success_response() setup_mock.assert_called_once() @@ -808,7 +823,8 @@ class TestDatasourceAuthDefaultApi: return_value=None, ) as set_default_datasource_provider, ): - response, status = method(api, "tenant-1", _PROVIDER_ID) + req_data = DatasourceDefaultPayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", _PROVIDER_ID) assert response == _success_response() assert status == 200 @@ -828,7 +844,7 @@ class TestDatasourceAuthDefaultApi: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(ValueError): - method(api, "tenant-1", "notion") + method(api, DatasourceDefaultPayload.model_validate(payload), "tenant-1", "notion") class TestDatasourceUpdateProviderNameApi: @@ -847,7 +863,8 @@ class TestDatasourceUpdateProviderNameApi: return_value=None, ) as update_datasource_provider_name, ): - response, status = method(api, "tenant-1", _PROVIDER_ID) + req_data = DatasourceUpdateNamePayload.model_validate(payload) + response, status = method(api, req_data, "tenant-1", _PROVIDER_ID) assert response == _success_response() assert status == 200 @@ -871,7 +888,7 @@ class TestDatasourceUpdateProviderNameApi: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(ValueError): - method(api, "tenant-1", "notion") + method(api, DatasourceUpdateNamePayload.model_validate(payload), "tenant-1", "notion") def test_update_name_missing_credential_id(self, app: Flask): api = DatasourceUpdateProviderNameApi() @@ -884,4 +901,4 @@ class TestDatasourceUpdateProviderNameApi: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(ValueError): - method(api, "tenant-1", "notion") + method(api, DatasourceUpdateNamePayload.model_validate(payload), "tenant-1", "notion") diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_content_preview.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_content_preview.py index eccf5ed56d1..b3aed7be87c 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_content_preview.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_content_preview.py @@ -7,6 +7,7 @@ from flask import Flask from controllers.console import console_ns from controllers.console.datasets.rag_pipeline.datasource_content_preview import ( DataSourceContentPreviewApi, + Parser, ) from models import Account from models.dataset import Pipeline @@ -32,7 +33,7 @@ class TestDataSourceContentPreviewApi: payload = self._valid_payload() - pipeline = MagicMock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") node_id = "node-1" account = make_account() @@ -49,7 +50,8 @@ class TestDataSourceContentPreviewApi: return_value=service_instance, ), ): - response, status = method(api, account, pipeline, node_id) + req_data = Parser.model_validate(payload) + response, status = method(api, req_data, account, pipeline, node_id) service_instance.run_datasource_node_preview.assert_called_once_with( pipeline=pipeline, @@ -72,7 +74,7 @@ class TestDataSourceContentPreviewApi: # datasource_type missing } - pipeline = MagicMock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") account = make_account() with ( @@ -80,7 +82,7 @@ class TestDataSourceContentPreviewApi: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(ValueError): - method(api, account, pipeline, "node-1") + method(api, Parser.model_validate(payload), account, pipeline, "node-1") def test_post_without_credential_id(self, app: Flask): api = DataSourceContentPreviewApi() @@ -92,7 +94,7 @@ class TestDataSourceContentPreviewApi: "credential_id": None, } - pipeline = MagicMock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") account = make_account() service_instance = MagicMock() @@ -106,7 +108,8 @@ class TestDataSourceContentPreviewApi: return_value=service_instance, ), ): - response, status = method(api, account, pipeline, "node-1") + req_data = Parser.model_validate(payload) + response, status = method(api, req_data, account, pipeline, "node-1") service_instance.run_datasource_node_preview.assert_called_once() assert status == 200 diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py index 860e7b6d286..91815612355 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline.py @@ -1,28 +1,32 @@ from __future__ import annotations from collections.abc import Iterator -from inspect import unwrap -from types import SimpleNamespace -from unittest.mock import PropertyMock, patch +from inspect import getclosurevars, unwrap +from unittest.mock import ANY, PropertyMock, patch import pytest from flask import Flask from sqlalchemy import Engine from sqlalchemy.orm import Session -from werkzeug.exceptions import NotFound +from werkzeug.exceptions import Forbidden, NotFound from controllers.console import console_ns from controllers.console.datasets.rag_pipeline import rag_pipeline as module from controllers.console.datasets.rag_pipeline.rag_pipeline import ( CustomizedPipelineTemplateApi, + CustomizedPipelineTemplatePayload, PipelineTemplateDetailApi, + PipelineTemplateDetailQuery, PipelineTemplateListApi, + PipelineTemplateListQuery, PublishCustomizedPipelineTemplateApi, ) -from models.account import Account -from models.dataset import PipelineCustomizedTemplate +from models.account import Account, TenantAccountRole +from models.dataset import Pipeline, PipelineCustomizedTemplate from models.engine import db from services.entities.knowledge_entities.rag_pipeline_entities import PipelineTemplateInfoEntity +from services.errors.account import NoPermissionError +from services.errors.rag_pipeline import RagPipelineResourceNotFoundError def _template_item() -> dict[str, object]: @@ -62,6 +66,12 @@ def _account() -> Account: return account +def _pipeline() -> Pipeline: + pipeline = Pipeline(tenant_id="tenant-1", name="Pipeline") + pipeline.id = "pipeline-1" + return pipeline + + @pytest.fixture def database_app() -> Iterator[Flask]: app = Flask(__name__) @@ -90,7 +100,7 @@ class TestPipelineTemplateListApi: app.test_request_context("/rag/pipeline/templates"), patch.object(module.RagPipelineService, "get_pipeline_templates", side_effect=get_pipeline_templates), ): - response, status = method(api, session, tenant_id) + response, status = method(api, PipelineTemplateListQuery(), session, tenant_id) assert status == 200 assert service_calls == [("built-in", "en-US", tenant_id)] @@ -120,7 +130,9 @@ class TestPipelineTemplateListApi: app.test_request_context("/rag/pipeline/templates?type=customized&language=ja-JP"), patch.object(module.RagPipelineService, "get_pipeline_templates", side_effect=get_pipeline_templates), ): - response, status = method(api, session, tenant_id) + response, status = method( + api, PipelineTemplateListQuery(type="customized", language="ja-JP"), session, tenant_id + ) assert status == 200 assert response == {"pipeline_templates": []} @@ -131,11 +143,13 @@ class TestPipelineTemplateDetailApi: def test_get_serializes_template_detail(self, app: Flask, sqlite_engine: Engine) -> None: api = PipelineTemplateDetailApi() method = unwrap(api.get) - service_calls: list[tuple[str, str]] = [] + service_calls: list[tuple[str, str, str]] = [] - def get_pipeline_template_detail(template_id: str, type: str, *, session) -> dict[str, object]: + def get_pipeline_template_detail( + template_id: str, current_tenant_id: str, type: str, *, session + ) -> dict[str, object]: del session - service_calls.append((template_id, type)) + service_calls.append((template_id, current_tenant_id, type)) return _template_detail() with ( @@ -147,18 +161,20 @@ class TestPipelineTemplateDetailApi: side_effect=get_pipeline_template_detail, ), ): - response, status = method(api, session, "template-1") + response, status = method( + api, PipelineTemplateDetailQuery(type="customized"), session, "tenant-1", "template-1" + ) assert status == 200 assert response == {**_template_detail(), "created_by": None} - assert service_calls == [("template-1", "customized")] + assert service_calls == [("template-1", "tenant-1", "customized")] def test_get_raises_not_found_without_custom_response_body(self, app: Flask, sqlite_engine: Engine) -> None: api = PipelineTemplateDetailApi() method = unwrap(api.get) - def get_pipeline_template_detail(template_id: str, type: str, *, session) -> None: - del template_id, type, session + def get_pipeline_template_detail(template_id: str, current_tenant_id: str, type: str, *, session) -> None: + del template_id, current_tenant_id, type, session with ( Session(sqlite_engine) as session, @@ -170,7 +186,7 @@ class TestPipelineTemplateDetailApi: ), pytest.raises(NotFound), ): - method(api, session, "missing") + method(api, PipelineTemplateDetailQuery(), session, "tenant-1", "missing") class TestCustomizedPipelineTemplateApi: @@ -198,7 +214,9 @@ class TestCustomizedPipelineTemplateApi: patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), patch.object(module.RagPipelineService, "update_customized_pipeline_template", side_effect=update_template), ): - response, status = method(api, tenant_id, account, "template-1") + response, status = method( + api, CustomizedPipelineTemplatePayload.model_validate(payload), tenant_id, account, "template-1" + ) assert (response, status) == ("", 204) assert len(service_calls) == 1 @@ -242,7 +260,9 @@ class TestCustomizedPipelineTemplateApi: patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), patch.object(module.RagPipelineService, "update_customized_pipeline_template", side_effect=update_template), ): - response, status = method(api, tenant_id, account, "template-1") + response, status = method( + api, CustomizedPipelineTemplatePayload.model_validate(payload), tenant_id, account, "template-1" + ) assert (response, status) == ("", 204) assert len(service_calls) == 1 @@ -277,9 +297,7 @@ class TestCustomizedPipelineTemplateApi: assert deleted_templates == [("template-1", tenant_id)] @pytest.mark.parametrize("sqlite_session", [(PipelineCustomizedTemplate,)], indirect=True) - def test_post_exports_yaml_from_orm_template( - self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session - ) -> None: + def test_post_exports_yaml_from_orm_template(self, app: Flask, sqlite_session: Session) -> None: api = CustomizedPipelineTemplateApi() method = unwrap(api.post) template = PipelineCustomizedTemplate( @@ -297,111 +315,201 @@ class TestCustomizedPipelineTemplateApi: template.id = "template-1" sqlite_session.add(template) sqlite_session.commit() - monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine)) with app.test_request_context("/rag/pipeline/customized/templates/template-1", method="POST"): - response, status = method(api, "template-1") + response, status = method( + api, + sqlite_session, + "00000000-0000-0000-0000-000000000001", + "template-1", + ) assert status == 200 assert response == {"data": "dsl: value"} @pytest.mark.parametrize("sqlite_session", [(PipelineCustomizedTemplate,)], indirect=True) - def test_post_raises_when_template_is_missing( - self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session - ) -> None: + def test_post_returns_not_found_for_other_tenant(self, app: Flask, sqlite_session: Session) -> None: api = CustomizedPipelineTemplateApi() method = unwrap(api.post) - assert sqlite_session.get(PipelineCustomizedTemplate, "missing") is None - monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine)) + template = PipelineCustomizedTemplate( + tenant_id="00000000-0000-0000-0000-000000000002", + name="Other tenant template", + description="Description", + chunk_structure="general", + icon={}, + position=1, + yaml_content="secret: value", + install_count=0, + language="en-US", + created_by="00000000-0000-0000-0000-000000000003", + ) + template.id = "template-1" + sqlite_session.add(template) + sqlite_session.commit() - with app.test_request_context("/rag/pipeline/customized/templates/missing", method="POST"): - with pytest.raises(ValueError, match="Customized pipeline template not found"): - method(api, "missing") + with ( + app.test_request_context("/rag/pipeline/customized/templates/template-1", method="POST"), + pytest.raises(NotFound, match="Customized pipeline template not found"), + ): + method( + api, + sqlite_session, + "00000000-0000-0000-0000-000000000001", + "template-1", + ) class TestPublishCustomizedPipelineTemplateApi: - def test_post_validates_payload_and_returns_empty_204(self, app: Flask) -> None: + def test_post_uses_pipeline_release_rbac_scene(self) -> None: + method = PublishCustomizedPipelineTemplateApi.post + while "scene" not in getclosurevars(method).nonlocals: + method = method.__wrapped__ + + assert getclosurevars(method).nonlocals["scene"] == module.RBACPermission.DATASET_PIPELINE_RELEASE + + def test_post_validates_payload_and_returns_empty_204(self) -> None: api = PublishCustomizedPipelineTemplateApi() method = unwrap(api.post) payload = _payload() account = _account() - tenant_id = "tenant-1" - service_calls: list[tuple[str, dict[str, object], Account, str]] = [] - - class Service: - def __init__(self, *args, **kwargs) -> None: - pass - - def publish_customized_pipeline_template( - self, - pipeline_id: str, - data: dict[str, object], - current_user: Account, - current_tenant_id: str, - *, - session, - ) -> None: - del session - service_calls.append((pipeline_id, data, current_user, current_tenant_id)) + pipeline = _pipeline() + dataset = object() with ( - app.test_request_context("/rag/pipelines/pipeline-1/customized/publish", method="POST", json=payload), - patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), - patch.object(module, "RagPipelineService", Service), + patch.object(module.dify_config, "RBAC_ENABLED", True), + patch.object(Pipeline, "retrieve_dataset", return_value=dataset), + patch.object(module.DatasetService, "check_dataset_permission") as legacy_acl, + patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish, ): - response, status = method(api, tenant_id, account, "pipeline-1") + response, status = method(api, CustomizedPipelineTemplatePayload.model_validate(payload), account, pipeline) assert (response, status) == ("", 204) - assert service_calls == [("pipeline-1", payload, account, tenant_id)] + publish.assert_called_once_with(pipeline, dataset, payload, account, session=ANY) + legacy_acl.assert_not_called() - def test_post_allows_missing_icon_info_for_publish_service_fallback(self, app: Flask) -> None: + @pytest.mark.parametrize( + ("payload", "expected_icon_info"), + [ + ( + {"name": "Published template", "description": "Description"}, + {"icon": "", "icon_background": None, "icon_type": None, "icon_url": None}, + ), + ({"name": "Published template", "description": "Description", "icon_info": {}}, {}), + ], + ) + def test_post_preserves_valid_icon_info( + self, + payload: dict[str, object], + expected_icon_info: dict[str, object | None], + ) -> None: api = PublishCustomizedPipelineTemplateApi() method = unwrap(api.post) - payload: dict[str, object] = { - "name": "Published template", - "description": "Description", - } account = _account() - tenant_id = "tenant-1" - service_calls: list[tuple[str, dict[str, object], Account, str]] = [] - - class Service: - def __init__(self, *args, **kwargs) -> None: - pass - - def publish_customized_pipeline_template( - self, - pipeline_id: str, - data: dict[str, object], - current_user: Account, - current_tenant_id: str, - *, - session, - ) -> None: - del session - service_calls.append((pipeline_id, data, current_user, current_tenant_id)) + pipeline = _pipeline() + dataset = object() with ( - app.test_request_context("/rag/pipelines/pipeline-1/customized/publish", method="POST", json=payload), - patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), - patch.object(module, "RagPipelineService", Service), + patch.object(module.dify_config, "RBAC_ENABLED", True), + patch.object(Pipeline, "retrieve_dataset", return_value=dataset), + patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish, ): - response, status = method(api, tenant_id, account, "pipeline-1") + response, status = method(api, CustomizedPipelineTemplatePayload.model_validate(payload), account, pipeline) assert (response, status) == ("", 204) - assert service_calls == [ - ( - "pipeline-1", - { - **payload, - "icon_info": { - "icon": "", - "icon_background": None, - "icon_type": None, - "icon_url": None, - }, - }, - account, - tenant_id, - ) - ] + publish.assert_called_once_with(pipeline, dataset, ANY, account, session=ANY) + assert publish.call_args.args[2]["icon_info"] == expected_icon_info + + def test_post_translates_missing_owned_resource_to_not_found(self) -> None: + api = PublishCustomizedPipelineTemplateApi() + method = unwrap(api.post) + payload = _payload() + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", True), + patch.object(Pipeline, "retrieve_dataset", return_value=object()), + patch.object( + module.RagPipelineService, + "publish_customized_pipeline_template", + side_effect=RagPipelineResourceNotFoundError("Workflow not found"), + ), + pytest.raises(NotFound, match="Workflow not found"), + ): + method(api, CustomizedPipelineTemplatePayload.model_validate(payload), _account(), _pipeline()) + + def test_post_allows_legacy_dataset_operator_after_dataset_acl(self) -> None: + api = PublishCustomizedPipelineTemplateApi() + method = unwrap(api.post) + account = _account() + account.role = TenantAccountRole.DATASET_OPERATOR + pipeline = _pipeline() + dataset = object() + payload = _payload() + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", False), + patch.object(Pipeline, "retrieve_dataset", return_value=dataset), + patch.object(module.DatasetService, "check_dataset_permission") as check_permission, + patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish, + ): + response = method(api, CustomizedPipelineTemplatePayload.model_validate(payload), account, pipeline) + + assert response == ("", 204) + assert check_permission.call_args.args[:2] == (dataset, account) + publish.assert_called_once_with(pipeline, dataset, payload, account, session=ANY) + + def test_post_rejects_legacy_non_editor_before_dataset_acl(self) -> None: + api = PublishCustomizedPipelineTemplateApi() + method = unwrap(api.post) + account = _account() + account.role = TenantAccountRole.NORMAL + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", False), + patch.object(Pipeline, "retrieve_dataset", return_value=object()), + patch.object(module.DatasetService, "check_dataset_permission") as check_permission, + patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish, + pytest.raises(Forbidden), + ): + method(api, CustomizedPipelineTemplatePayload.model_validate(_payload()), account, _pipeline()) + + check_permission.assert_not_called() + publish.assert_not_called() + + def test_post_rejects_legacy_dataset_acl_before_publish(self) -> None: + api = PublishCustomizedPipelineTemplateApi() + method = unwrap(api.post) + account = _account() + account.role = TenantAccountRole.EDITOR + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", False), + patch.object(Pipeline, "retrieve_dataset", return_value=object()), + patch.object( + module.DatasetService, + "check_dataset_permission", + side_effect=NoPermissionError("Dataset is private"), + ), + patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish, + pytest.raises(Forbidden, match="Dataset is private"), + ): + method(api, CustomizedPipelineTemplatePayload.model_validate(_payload()), account, _pipeline()) + + publish.assert_not_called() + + def test_post_rejects_missing_legacy_dataset_before_publish(self) -> None: + api = PublishCustomizedPipelineTemplateApi() + method = unwrap(api.post) + account = _account() + account.role = TenantAccountRole.EDITOR + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", False), + patch.object(Pipeline, "retrieve_dataset", return_value=None), + patch.object(module.DatasetService, "check_dataset_permission") as check_permission, + patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish, + pytest.raises(NotFound, match="Dataset not found"), + ): + method(api, CustomizedPipelineTemplatePayload.model_validate(_payload()), account, _pipeline()) + + check_permission.assert_not_called() + publish.assert_not_called() diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py index e9567f891af..3f8e1a2756c 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py @@ -15,6 +15,7 @@ from controllers.console.datasets.error import DatasetNameDuplicateError from controllers.console.datasets.rag_pipeline.rag_pipeline_datasets import ( CreateEmptyRagPipelineDatasetApi, CreateRagPipelineDatasetApi, + RagPipelineDatasetImportPayload, ) from services.entities.dsl_entities import ImportStatus @@ -50,7 +51,7 @@ class TestCreateRagPipelineDatasetApi: return_value=mock_service, ), ): - response, status = method(api, "tenant-1", user) + response, status = method(api, RagPipelineDatasetImportPayload.model_validate(payload), "tenant-1", user) assert status == 201 assert response == { @@ -75,7 +76,7 @@ class TestCreateRagPipelineDatasetApi: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(Forbidden): - method(api, "tenant-1", user) + method(api, RagPipelineDatasetImportPayload.model_validate(payload), "tenant-1", user) def test_post_dataset_name_duplicate(self, app: Flask) -> None: api = CreateRagPipelineDatasetApi() @@ -96,7 +97,7 @@ class TestCreateRagPipelineDatasetApi: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, "tenant-1", user) + method(api, RagPipelineDatasetImportPayload.model_validate(payload), "tenant-1", user) def test_post_invalid_payload(self, app: Flask) -> None: api = CreateRagPipelineDatasetApi() @@ -110,7 +111,7 @@ class TestCreateRagPipelineDatasetApi: patch.object(type(console_ns), "payload", payload), ): with pytest.raises(ValueError): - method(api, "tenant-1", user) + method(api, RagPipelineDatasetImportPayload.model_validate(payload), "tenant-1", user) class TestCreateEmptyRagPipelineDatasetApi: diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_draft_variable.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_draft_variable.py index 0d18120b715..a000a6a6ed5 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_draft_variable.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_draft_variable.py @@ -8,13 +8,16 @@ from controllers.common.errors import InvalidArgumentError, NotFoundError from controllers.console import console_ns from controllers.console.app.error import DraftWorkflowNotExist from controllers.console.datasets.rag_pipeline.rag_pipeline_draft_variable import ( + PaginationQuery, RagPipelineEnvironmentVariableCollectionApi, RagPipelineNodeVariableCollectionApi, RagPipelineSystemVariableCollectionApi, RagPipelineVariableApi, RagPipelineVariableCollectionApi, RagPipelineVariableResetApi, + WorkflowDraftVariablePatchPayload, ) +from core.workflow.llm_environment_variable import LLMEnvironmentVariable from core.workflow.variable_prefixes import SYSTEM_VARIABLE_NODE_ID from graphon.variables.types import SegmentType from models.account import Account, TenantAccountRole @@ -71,7 +74,7 @@ class TestRagPipelineVariableCollectionApi: return_value=draft_srv, ), ): - result = method(api, editor_user, pipeline) + result = method(api, PaginationQuery(page=1, limit=10), editor_user, pipeline) assert result is var_list draft_srv.list_variables_without_values.assert_called_once_with( @@ -99,7 +102,7 @@ class TestRagPipelineVariableCollectionApi: ), ): with pytest.raises(DraftWorkflowNotExist): - method(api, editor_user, pipeline) + method(api, PaginationQuery(), editor_user, pipeline) def test_delete_variables_success(self, app: Flask, fake_db, editor_user): api = RagPipelineVariableCollectionApi() @@ -197,7 +200,7 @@ class TestRagPipelineVariableApi: ), ): with pytest.raises(InvalidArgumentError): - method(api, editor_user, pipeline, "v1") + method(api, WorkflowDraftVariablePatchPayload.model_validate(payload), editor_user, pipeline, "v1") def test_delete_variable_success(self, app: Flask, fake_db, editor_user): api = RagPipelineVariableApi() @@ -315,3 +318,35 @@ class TestSystemAndEnvironmentVariablesApi: result = method(api, editor_user, pipeline) assert len(result["items"]) == 1 + assert result["items"][0]["value_type"] == "string" + + def test_environment_variables_preserve_number_subtype_and_llm_type(self, app: Flask, editor_user): + api = RagPipelineEnvironmentVariableCollectionApi() + method = unwrap(api.get) + number_var = MagicMock( + id="number", + name="NUMBER", + description="", + selector=["env", "NUMBER"], + value_type=MagicMock(value="integer"), + value=1, + ) + llm_var = LLMEnvironmentVariable( + id="llm", + name="MODEL", + value={"provider": "provider", "name": "model", "mode": "chat"}, + selector=["env", "MODEL"], + ) + rag_srv = MagicMock() + rag_srv.get_draft_workflow.return_value = MagicMock(environment_variables=[number_var, llm_var]) + + with ( + app.test_request_context("/"), + patch( + "controllers.console.datasets.rag_pipeline.rag_pipeline_draft_variable.RagPipelineService", + return_value=rag_srv, + ), + ): + result = method(api, editor_user, MagicMock(id="p1")) + + assert [item["value_type"] for item in result["items"]] == ["integer", "llm"] diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py similarity index 89% rename from api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py rename to api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py index b9a0029d131..ad6d974bdb8 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py @@ -1,7 +1,9 @@ -"""Testcontainers integration tests for rag_pipeline_import controller endpoints.""" +"""Unit tests for rag_pipeline_import controller endpoints.""" from __future__ import annotations +from collections.abc import Iterator +from inspect import unwrap from unittest.mock import MagicMock, patch import pytest @@ -9,23 +11,31 @@ from flask import Flask from controllers.console import console_ns from controllers.console.datasets.rag_pipeline.rag_pipeline_import import ( + IncludeSecretQuery, RagPipelineExportApi, RagPipelineImportApi, RagPipelineImportCheckDependenciesApi, RagPipelineImportConfirmApi, + RagPipelineImportPayload, ) from core.plugin.entities.plugin import PluginDependency, PluginDependencyType from models.dataset import Pipeline +from models.engine import db from services.entities.dsl_entities import CheckDependenciesResult, ImportStatus from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineImportInfo -from tests.test_containers_integration_tests.controllers.console.helpers import unwrap + + +@pytest.fixture +def app() -> Iterator[Flask]: + app = Flask(__name__) + app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:" + db.init_app(app) + + with app.app_context(): + yield app class TestRagPipelineImportApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def _payload(self, mode: str = "create") -> dict[str, str]: return { "mode": mode, @@ -59,7 +69,7 @@ class TestRagPipelineImportApi: return_value=service, ), ): - response, status = method(api, user) + response, status = method(api, RagPipelineImportPayload(mode="create"), user) assert status == 200 assert response == { @@ -97,7 +107,7 @@ class TestRagPipelineImportApi: return_value=service, ), ): - response, status = method(api, user) + response, status = method(api, RagPipelineImportPayload(mode="create"), user) assert status == 400 assert response["status"] == "failed" @@ -131,7 +141,7 @@ class TestRagPipelineImportApi: return_value=service, ), ): - response, status = method(api, user) + response, status = method(api, RagPipelineImportPayload(mode="create"), user) assert status == 202 assert response["status"] == "pending" @@ -140,10 +150,6 @@ class TestRagPipelineImportApi: class TestRagPipelineImportConfirmApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_confirm_success(self, app: Flask) -> None: api = RagPipelineImportConfirmApi() method = unwrap(api.post) @@ -205,15 +211,11 @@ class TestRagPipelineImportConfirmApi: class TestRagPipelineImportCheckDependenciesApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_get_success(self, app: Flask) -> None: api = RagPipelineImportCheckDependenciesApi() method = unwrap(api.get) - pipeline = MagicMock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") result = CheckDependenciesResult() service = MagicMock() @@ -235,7 +237,7 @@ class TestRagPipelineImportCheckDependenciesApi: api = RagPipelineImportCheckDependenciesApi() method = unwrap(api.get) - pipeline = MagicMock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") dependency = PluginDependency( type=PluginDependencyType.Marketplace, value=PluginDependency.Marketplace( @@ -272,15 +274,11 @@ class TestRagPipelineImportCheckDependenciesApi: class TestRagPipelineExportApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_get_with_include_secret(self, app: Flask) -> None: api = RagPipelineExportApi() method = unwrap(api.get) - pipeline = MagicMock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") service = MagicMock() service.export_rag_pipeline_dsl.return_value = "yaml: data" @@ -291,7 +289,7 @@ class TestRagPipelineExportApi: return_value=service, ), ): - response, status = method(api, pipeline) + response, status = method(api, IncludeSecretQuery(), pipeline) assert status == 200 assert response == {"data": "yaml: data"} diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py index c0f9a902a56..fdc9dfe54f7 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py @@ -1,40 +1,62 @@ -"""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. -""" +"""Unit coverage for RAG workflow controllers using real models and disposable SQLite state.""" from __future__ import annotations +import json +from collections.abc import Iterator from datetime import datetime from inspect import unwrap as unwrap_all -from types import SimpleNamespace -from unittest.mock import PropertyMock, patch +from unittest.mock import MagicMock, patch +from uuid import UUID import pytest from flask import Flask -from sqlalchemy.engine import Engine +from sqlalchemy import Engine from sqlalchemy.orm import Session +from werkzeug.exceptions import Forbidden, NotFound from controllers.console.datasets.rag_pipeline import rag_pipeline_workflow as module -from models.account import Account, TenantAccountRole -from models.dataset import Pipeline +from controllers.console.datasets.rag_pipeline.rag_pipeline_workflow import ( + DraftWorkflowRunPayload, + NodeIdQuery, + PublishedWorkflowRunPayload, + RagPipelineRecommendedPluginQuery, + WorkflowListQuery, + WorkflowUpdatePayload, +) +from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError +from models.account import Account, Tenant, TenantAccountRole +from models.dataset import Dataset, Pipeline +from models.engine import db +from models.enums import PermissionEnum +from models.tools import WorkflowToolProvider +from models.workflow import Workflow, WorkflowType +from services.errors.llm import InvokeRateLimitError +from services.errors.rag_pipeline import RagPipelineResourceNotFoundError +from services.rag_pipeline.rag_pipeline import RagPipelineService + +DEFAULT_WORKFLOW_TENANT_ID = "00000000-0000-0000-0000-000000000001" +DEFAULT_WORKFLOW_APP_ID = "00000000-0000-0000-0000-000000000002" +DEFAULT_WORKFLOW_CREATED_BY = "00000000-0000-0000-0000-000000000003" +DEFAULT_WORKFLOW_ID = "00000000-0000-0000-0000-000000000004" +DEFAULT_DATASET_ID = "44444444-4444-4444-4444-444444444444" -def _make_workflow(**overrides): - workflow = SimpleNamespace( - id="workflow-1", - graph_dict={"nodes": [], "edges": []}, - features_dict={"file_upload": {"enabled": False}}, - unique_hash="hash-1", - version="1", +def _make_workflow(**overrides: object) -> Workflow: + workflow = Workflow( + id=DEFAULT_WORKFLOW_ID, + tenant_id=DEFAULT_WORKFLOW_TENANT_ID, + app_id=DEFAULT_WORKFLOW_APP_ID, + type=WorkflowType.WORKFLOW, + version=Workflow.VERSION_DRAFT, marked_name="Release 1", marked_comment="Initial release", - created_by_account=SimpleNamespace(id="user-1", name="Alice", email="alice@example.com"), + graph=json.dumps({"nodes": [], "edges": []}), + features=json.dumps({"file_upload": {"enabled": False}}), + created_by=DEFAULT_WORKFLOW_CREATED_BY, created_at=datetime(2024, 1, 1, 12, 0, 0), - updated_by_account=None, + updated_by=None, updated_at=datetime(2024, 1, 1, 12, 1, 0), - tool_published=False, environment_variables=[], conversation_variables=[], rag_pipeline_variables=[], @@ -46,137 +68,148 @@ def _make_workflow(**overrides): def _account() -> Account: account = Account(name="Alice", email="alice@example.com") - account.id = "user-1" + account.id = DEFAULT_WORKFLOW_CREATED_BY account.role = TenantAccountRole.EDITOR + tenant = Tenant(name="Tenant") + tenant.id = DEFAULT_WORKFLOW_TENANT_ID + account._current_tenant = tenant return account def _pipeline() -> Pipeline: - pipeline = Pipeline(tenant_id="tenant-1", name="Pipeline", description="desc") - pipeline.id = "pipeline-1" + pipeline = Pipeline(tenant_id=DEFAULT_WORKFLOW_TENANT_ID, name="Pipeline", description="desc") + pipeline.id = DEFAULT_WORKFLOW_APP_ID return pipeline -def test_draft_rag_pipeline_workflow_get_serializes_response_model(monkeypatch: pytest.MonkeyPatch) -> None: - workflow = _make_workflow() - monkeypatch.setattr( - module, - "RagPipelineService", - lambda *_args, **_kwargs: SimpleNamespace(get_draft_workflow=lambda **_kwargs: workflow), +def _dataset(*, tenant_id: str = DEFAULT_WORKFLOW_TENANT_ID, maintainer: str = DEFAULT_WORKFLOW_CREATED_BY) -> Dataset: + return Dataset( + id=DEFAULT_DATASET_ID, + tenant_id=tenant_id, + name="Dataset", + created_by=maintainer, + maintainer=maintainer, + permission=PermissionEnum.ONLY_ME, + provider="vendor", ) + +def _persist_workflow(workflow: Workflow) -> None: + db.session.add(workflow) + db.session.commit() + db.session.expunge(workflow) + + +@pytest.fixture +def database_app() -> Iterator[Flask]: + app = Flask(__name__) + app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:" + db.init_app(app) + + with app.app_context(): + Account.__table__.create(db.engine) + WorkflowToolProvider.__table__.create(db.engine) + Workflow.__table__.create(db.engine) + db.session.add(_account()) + db.session.commit() + + try: + yield app + finally: + db.session.remove() + + +@pytest.mark.usefixtures("database_app") +def test_draft_rag_pipeline_workflow_get_serializes_response_model() -> None: + workflow = _make_workflow() + expected_hash = workflow.unique_hash + _persist_workflow(workflow) + api = module.DraftRagPipelineApi() handler = unwrap_all(api.get) response = handler(api, _pipeline()) - assert response["id"] == "workflow-1" + assert response["id"] == DEFAULT_WORKFLOW_ID assert response["graph"] == {"nodes": [], "edges": []} assert response["features"] == {"file_upload": {"enabled": False}} - assert response["hash"] == "hash-1" - assert response["created_by"] == {"id": "user-1", "name": "Alice", "email": "alice@example.com"} + assert response["hash"] == expected_hash + assert response["created_by"] == { + "id": DEFAULT_WORKFLOW_CREATED_BY, + "name": "Alice", + "email": "alice@example.com", + } assert response["updated_by"] is None assert response["created_at"] == int(datetime(2024, 1, 1, 12, 0, 0).timestamp()) assert response["updated_at"] == int(datetime(2024, 1, 1, 12, 1, 0).timestamp()) def test_published_rag_pipeline_workflows_serialize_items_before_session_closes( - app, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine + database_app: Flask, ) -> None: api = module.PublishedAllRagPipelineApi() handler = unwrap_all(api.get) - session_state: dict[str, Session] = {} + workflow = _make_workflow(version="1") + _persist_workflow(workflow) + pipeline = _pipeline() + pipeline.workflow_id = DEFAULT_WORKFLOW_ID - base_workflow = _make_workflow() + with database_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, WorkflowListQuery(page=1, limit=10, user_id="", named_only=False), _account(), pipeline=pipeline + ) - class _Workflow: - def __getattr__(self, name: str): - assert session_state["session"].in_transaction() is True - return getattr(base_workflow, name) - - 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=_get_all_published_workflow), - ) - - 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["items"][0]["id"] == DEFAULT_WORKFLOW_ID assert response["page"] == 1 assert response["limit"] == 10 assert response["has_more"] is False def test_rag_pipeline_workflow_patch_serializes_response_model( - app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine + database_app: Flask, ) -> None: workflow = _make_workflow(marked_name="Updated release") - captured_session: dict[str, Session] = {} - - def _update_workflow(**kwargs): - captured_session["session"] = kwargs["session"] - assert kwargs["session"].in_transaction() is True - return workflow - - monkeypatch.setattr( - module, - "RagPipelineService", - lambda *_args, **_kwargs: SimpleNamespace(update_workflow=_update_workflow), - ) + expected_hash = workflow.unique_hash + _persist_workflow(workflow) payload: dict[str, object] = {"marked_name": "Updated release"} api = module.RagPipelineByIdApi() handler = unwrap_all(api.patch) - with Session(sqlite_engine) as request_session: - monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine, session=lambda: request_session)) - with ( - app.test_request_context("/rag/pipelines/pipeline-1/workflows/workflow-1", method="PATCH", json=payload), - patch.object(type(module.console_ns), "payload", new_callable=PropertyMock, return_value=payload), - ): - response = handler( - api, - _account(), - pipeline=_pipeline(), - workflow_id="workflow-1", - ) + with database_app.test_request_context( + f"/rag/pipelines/{DEFAULT_WORKFLOW_APP_ID}/workflows/{DEFAULT_WORKFLOW_ID}", method="PATCH", json=payload + ): + response = handler( + api, + WorkflowUpdatePayload.model_validate(payload), + _account(), + pipeline=_pipeline(), + workflow_id=DEFAULT_WORKFLOW_ID, + ) - assert captured_session["session"].in_transaction() is False - assert response["id"] == "workflow-1" + assert response["id"] == DEFAULT_WORKFLOW_ID assert response["marked_name"] == "Updated release" - assert response["hash"] == "hash-1" + assert response["hash"] == expected_hash -def test_default_rag_pipeline_block_configs_serializes_root_response(monkeypatch: pytest.MonkeyPatch) -> None: +@pytest.mark.usefixtures("database_app") +def test_default_rag_pipeline_block_configs_serializes_root_response() -> None: block_configs = [{"type": "start", "config": {"title": "Start"}}] - monkeypatch.setattr( - module, - "RagPipelineService", - lambda *_args, **_kwargs: SimpleNamespace(get_default_block_configs=lambda: block_configs), - ) api = module.DefaultRagPipelineBlockConfigsApi() handler = unwrap_all(api.get) - response = handler(api, _pipeline()) + with patch.object(RagPipelineService, "get_default_block_configs", return_value=block_configs): + response = handler(api, _pipeline()) assert response == block_configs -def test_draft_rag_pipeline_second_step_parameters_serializes_variables(app, monkeypatch: pytest.MonkeyPatch) -> None: +def test_draft_rag_pipeline_second_step_parameters_serializes_variables(database_app: Flask) -> None: variables = [ { "belong_to_node_id": "shared", @@ -187,36 +220,225 @@ def test_draft_rag_pipeline_second_step_parameters_serializes_variables(app, mon "required": True, } ] - monkeypatch.setattr( - module, - "RagPipelineService", - lambda *_args, **_kwargs: SimpleNamespace(get_second_step_parameters=lambda **_kwargs: variables), - ) - api = module.DraftRagPipelineSecondStepApi() handler = unwrap_all(api.get) - with app.test_request_context("/?node_id=node-1"): - response = handler(api, _pipeline()) + with ( + database_app.test_request_context("/?node_id=node-1"), + patch.object(RagPipelineService, "get_second_step_parameters", return_value=variables), + ): + response = handler(api, NodeIdQuery(node_id="node-1"), _pipeline()) assert response["variables"] == variables -def test_rag_pipeline_recommended_plugins_serializes_known_envelope(app, monkeypatch: pytest.MonkeyPatch) -> None: +def test_rag_pipeline_recommended_plugins_serializes_known_envelope(database_app: Flask) -> None: recommended_plugins = { "installed_recommended_plugins": [{"name": "Dify Extractor", "meta": {"version": "1.0.0"}}], "uninstalled_recommended_plugins": [{"plugin_id": "langgenius/notion_datasource"}], } - monkeypatch.setattr( - module, - "RagPipelineService", - lambda *_args, **_kwargs: SimpleNamespace(get_recommended_plugins=lambda *_args: recommended_plugins), - ) - api = module.RagPipelineRecommendedPluginApi() handler = unwrap_all(api.get) - with app.test_request_context("/?type=tool"): - response = handler(api, "tenant-1", _account()) + with ( + database_app.test_request_context("/?type=tool"), + patch.object(RagPipelineService, "get_recommended_plugins", return_value=recommended_plugins), + ): + response = handler(api, RagPipelineRecommendedPluginQuery(type="tool"), DEFAULT_WORKFLOW_TENANT_ID, _account()) assert response == recommended_plugins + + +def test_rag_pipeline_transform_rejects_read_only_member(sqlite_engine: Engine) -> None: + account = _account() + account.role = TenantAccountRole.NORMAL + api = module.RagPipelineTransformApi() + handler = unwrap_all(api.post) + + with Session(sqlite_engine) as session: + session.add(_dataset()) + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", False), + pytest.raises(Forbidden), + ): + handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID)) + + +def test_rag_pipeline_transform_rejects_dataset_from_another_tenant_before_service_call( + sqlite_engine: Engine, +) -> None: + api = module.RagPipelineTransformApi() + handler = unwrap_all(api.post) + + with Session(sqlite_engine) as session: + session.add(_dataset(tenant_id="00000000-0000-0000-0000-000000000099")) + + with ( + patch.object(module.RagPipelineTransformService, "transform_dataset") as transform_dataset, + pytest.raises(NotFound), + ): + handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, _account(), UUID(DEFAULT_DATASET_ID)) + + transform_dataset.assert_not_called() + + +def test_rag_pipeline_transform_enforces_legacy_dataset_permission_before_service_call( + sqlite_engine: Engine, +) -> None: + api = module.RagPipelineTransformApi() + handler = unwrap_all(api.post) + + with Session(sqlite_engine) as session: + session.add(_dataset(maintainer="00000000-0000-0000-0000-000000000099")) + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", False), + patch.object(module.RagPipelineTransformService, "transform_dataset") as transform_dataset, + pytest.raises(Forbidden), + ): + handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, _account(), UUID(DEFAULT_DATASET_ID)) + + transform_dataset.assert_not_called() + + +def test_rag_pipeline_transform_passes_authorized_dataset_and_account_to_service( + sqlite_engine: Engine, +) -> None: + api = module.RagPipelineTransformApi() + handler = unwrap_all(api.post) + account = _account() + expected = {"pipeline_id": "pipeline-1", "dataset_id": DEFAULT_DATASET_ID, "status": "success"} + + with Session(sqlite_engine) as session: + dataset = _dataset() + session.add(dataset) + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", False), + patch.object(module.RagPipelineTransformService, "transform_dataset", return_value=expected) as transform, + ): + response = handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID)) + + transform.assert_called_once_with(dataset, account.id, session) + + assert response == expected + + +def test_rag_pipeline_transform_maps_missing_pipeline_to_not_found(sqlite_engine: Engine) -> None: + api = module.RagPipelineTransformApi() + handler = unwrap_all(api.post) + + with Session(sqlite_engine) as session: + session.add(_dataset()) + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", False), + patch.object( + module.RagPipelineTransformService, + "transform_dataset", + side_effect=RagPipelineResourceNotFoundError("Pipeline not found"), + ), + pytest.raises(NotFound, match="Pipeline not found"), + ): + handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, _account(), UUID(DEFAULT_DATASET_ID)) + + +def test_rag_pipeline_transform_skips_legacy_acl_when_rbac_is_enabled(sqlite_engine: Engine) -> None: + api = module.RagPipelineTransformApi() + handler = unwrap_all(api.post) + account = _account() + account.role = TenantAccountRole.NORMAL + expected = {"pipeline_id": "pipeline-1", "dataset_id": DEFAULT_DATASET_ID, "status": "success"} + + with Session(sqlite_engine) as session: + session.add(_dataset(maintainer="00000000-0000-0000-0000-000000000099")) + + with ( + patch.object(module.dify_config, "RBAC_ENABLED", True), + patch.object(module.RagPipelineTransformService, "transform_dataset", return_value=expected) as transform, + ): + response = handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID)) + + assert response == expected + transform.assert_called_once() + + +@pytest.mark.parametrize( + ("api_type", "payload"), + [ + ( + module.DraftRagPipelineRunApi, + {"inputs": {}, "datasource_type": "x", "datasource_info_list": [], "start_node_id": "node-1"}, + ), + ( + module.PublishedRagPipelineRunApi, + { + "inputs": {}, + "datasource_type": "x", + "datasource_info_list": [], + "start_node_id": "node-1", + "response_mode": "blocking", + }, + ), + ], +) +def test_rag_pipeline_run_uses_sqlite_session( + app: Flask, + sqlite_engine: Engine, + api_type: type, + payload: dict[str, object], +) -> None: + api = api_type() + handler = unwrap_all(api.post) + pipeline = _pipeline() + + with ( + Session(sqlite_engine) as session, + app.test_request_context("/", json=payload), + patch.object(module, "load_rag_pipeline", return_value=pipeline) as load_pipeline, + patch.object(module.PipelineGenerateService, "generate", return_value=MagicMock()) as generate, + patch.object(module.helper, "compact_generate_response", return_value={"ok": True}), + ): + req_data = ( + DraftWorkflowRunPayload.model_validate(payload) + if api_type is module.DraftRagPipelineRunApi + else PublishedWorkflowRunPayload.model_validate(payload) + ) + response = handler(api, req_data, session, _account(), pipeline.id) + + assert response == {"ok": True} + load_pipeline.assert_called_once_with(session, pipeline.id) + assert generate.call_args.kwargs["session"] is session + assert session.get_bind() is sqlite_engine + + +@pytest.mark.parametrize("api_type", [module.DraftRagPipelineRunApi, module.PublishedRagPipelineRunApi]) +def test_rag_pipeline_run_translates_rate_limit( + app: Flask, + sqlite_engine: Engine, + api_type: type, +) -> None: + payload = { + "inputs": {}, + "datasource_type": "x", + "datasource_info_list": [], + "start_node_id": "node-1", + } + api = api_type() + handler = unwrap_all(api.post) + pipeline = _pipeline() + req_data = ( + DraftWorkflowRunPayload.model_validate(payload) + if api_type is module.DraftRagPipelineRunApi + else PublishedWorkflowRunPayload.model_validate(payload) + ) + + with ( + Session(sqlite_engine) as session, + app.test_request_context("/", json=payload), + patch.object(module, "load_rag_pipeline", return_value=pipeline), + patch.object(module.PipelineGenerateService, "generate", side_effect=InvokeRateLimitError("limit")), + pytest.raises(InvokeRateLimitHttpError), + ): + handler(api, req_data, session, _account(), pipeline.id) diff --git a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow_apis.py similarity index 72% rename from api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py rename to api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow_apis.py index f8abd102143..caef20a3d69 100644 --- a/api/tests/test_containers_integration_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow_apis.py @@ -1,8 +1,10 @@ -"""Testcontainers integration tests for rag_pipeline_workflow controller endpoints.""" +"""Unit tests for rag_pipeline_workflow controller endpoints.""" from __future__ import annotations import json +from collections.abc import Iterator +from dataclasses import dataclass from datetime import datetime from inspect import unwrap from typing import TypedDict, Unpack @@ -11,19 +13,23 @@ from uuid import uuid4 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 BadRequest, Forbidden, HTTPException, NotFound +import models.workflow as workflow_models import services -from controllers.console import console_ns from controllers.console.app.error import DraftWorkflowNotExist, DraftWorkflowNotSync +from controllers.console.datasets.rag_pipeline import rag_pipeline_workflow as workflow_controller from controllers.console.datasets.rag_pipeline.rag_pipeline_workflow import ( + DatasourceVariablesPayload, + DefaultBlockConfigQuery, DefaultRagPipelineBlockConfigApi, DraftRagPipelineApi, - DraftRagPipelineRunApi, + NodeRunPayload, + NodeRunRequiredPayload, PublishedAllRagPipelineApi, PublishedRagPipelineApi, - PublishedRagPipelineRunApi, RagPipelineByIdApi, RagPipelineDatasourceVariableApi, RagPipelineDraftNodeRunApi, @@ -31,12 +37,13 @@ from controllers.console.datasets.rag_pipeline.rag_pipeline_workflow import ( RagPipelineDraftRunLoopNodeApi, RagPipelineDraftWorkflowRestoreApi, RagPipelineRecommendedPluginApi, + RagPipelineRecommendedPluginQuery, RagPipelineTaskStopApi, - RagPipelineTransformApi, RagPipelineWorkflowLastRunApi, RagPipelineWorkflowRunNodeExecutionListApi, + WorkflowListQuery, + WorkflowUpdatePayload, ) -from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError from graphon.enums import WorkflowNodeExecutionStatus from libs.datetime_utils import naive_utc_now from models.account import Account, TenantAccountRole @@ -44,7 +51,6 @@ from models.dataset import Pipeline from models.enums import CreatorUserRole from models.workflow import Workflow, WorkflowNodeExecutionModel, WorkflowNodeExecutionTriggeredFrom from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError -from services.errors.llm import InvokeRateLimitError DEFAULT_WORKFLOW_TENANT_ID = "00000000-0000-0000-0000-000000000001" DEFAULT_WORKFLOW_APP_ID = "00000000-0000-0000-0000-000000000002" @@ -52,6 +58,38 @@ DEFAULT_WORKFLOW_CREATED_BY = "00000000-0000-0000-0000-000000000003" type WorkflowVariablePayload = dict[str, object] +@dataclass(frozen=True) +class SQLiteDatabase: + """Expose the concrete SQLite engine and scoped session interface used by controller code.""" + + engine: Engine + session: scoped_session[Session] + + +@pytest.fixture(autouse=True) +def sqlite_database( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, +) -> Iterator[scoped_session[Session]]: + """Route controller transactions and model author lookups through SQLite.""" + + database_session = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False)) + database = SQLiteDatabase(engine=sqlite_engine, session=database_session) + monkeypatch.setattr(workflow_controller, "db", database) + monkeypatch.setattr(workflow_models, "db", database) + + with database_session() as session: + default_author = Account(name="Default Author", email="default-author@example.com") + default_author.id = DEFAULT_WORKFLOW_CREATED_BY + session.add(default_author) + session.commit() + + try: + yield database_session + finally: + database_session.remove() + + def empty_mapping() -> dict[str, object]: return {} @@ -206,18 +244,15 @@ def make_pipeline( @pytest.fixture -def workflow_author(db_session_with_containers: Session) -> Account: +def workflow_author(sqlite_database: scoped_session[Session]) -> Account: account = Account(name="Alice", email=f"alice-{uuid4()}@example.com") - db_session_with_containers.add(account) - db_session_with_containers.commit() + account.id = str(uuid4()) + sqlite_database.add(account) + sqlite_database.commit() return account class TestDraftWorkflowApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_get_draft_success(self, app: Flask, workflow_author: Account) -> None: api = DraftRagPipelineApi() method = unwrap(api.get) @@ -278,7 +313,6 @@ class TestDraftWorkflowApi: with ( app.test_request_context("/", json={"graph": empty_mapping(), "features": empty_mapping()}), - patch.object(type(console_ns), "payload", {"graph": empty_mapping(), "features": empty_mapping()}), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService", return_value=service, @@ -373,10 +407,6 @@ class TestDraftWorkflowApi: class TestDraftRunNodes: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_iteration_node_success(self, app: Flask) -> None: api = RagPipelineDraftRunIterationNodeApi() method = unwrap(api.post) @@ -386,7 +416,6 @@ class TestDraftRunNodes: with ( app.test_request_context("/", json={"inputs": empty_mapping()}), - patch.object(type(console_ns), "payload", {"inputs": empty_mapping()}), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate_single_iteration", return_value=MagicMock(), @@ -396,7 +425,7 @@ class TestDraftRunNodes: return_value={"ok": True}, ), ): - result = method(api, user, pipeline, "node") + result = method(api, NodeRunPayload(), user, pipeline, "node") assert result == {"ok": True} def test_iteration_node_conversation_not_exists(self, app: Flask) -> None: @@ -408,14 +437,13 @@ class TestDraftRunNodes: with ( app.test_request_context("/", json={"inputs": empty_mapping()}), - patch.object(type(console_ns), "payload", {"inputs": empty_mapping()}), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate_single_iteration", side_effect=services.errors.conversation.ConversationNotExistsError(), ), ): with pytest.raises(NotFound): - method(api, user, pipeline, "node") + method(api, NodeRunPayload(), user, pipeline, "node") def test_loop_node_success(self, app: Flask) -> None: api = RagPipelineDraftRunLoopNodeApi() @@ -426,7 +454,6 @@ class TestDraftRunNodes: with ( app.test_request_context("/", json={"inputs": empty_mapping()}), - patch.object(type(console_ns), "payload", {"inputs": empty_mapping()}), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate_single_loop", return_value=MagicMock(), @@ -436,89 +463,10 @@ class TestDraftRunNodes: return_value={"ok": True}, ), ): - assert method(api, user, pipeline, "node") == {"ok": True} - - -class TestPipelineRunApis: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_draft_run_success(self, app: Flask) -> None: - api = DraftRagPipelineRunApi() - method = unwrap(api.post) - - pipeline = make_pipeline() - user = make_account() - session = MagicMock(spec=Session) - - payload = { - "inputs": empty_mapping(), - "datasource_type": "x", - "datasource_info_list": empty_list(), - "start_node_id": "n", - } - - with ( - app.test_request_context("/", json=payload), - patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.load_rag_pipeline", - return_value=pipeline, - ) as load_pipeline, - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate", - return_value=MagicMock(), - ) as generate, - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.helper.compact_generate_response", - return_value={"ok": True}, - ), - ): - assert method(api, session, user, pipeline.id) == {"ok": True} - - load_pipeline.assert_called_once_with(session, pipeline.id) - assert generate.call_args.kwargs["session"] is session - - def test_draft_run_rate_limit(self, app: Flask) -> None: - api = DraftRagPipelineRunApi() - method = unwrap(api.post) - - pipeline = make_pipeline() - user = make_account() - session = MagicMock(spec=Session) - payload: dict[str, object] = { - "inputs": empty_mapping(), - "datasource_type": "x", - "datasource_info_list": empty_list(), - "start_node_id": "n", - } - - with ( - app.test_request_context("/", json=payload), - patch.object( - type(console_ns), - "payload", - payload, - ), - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.load_rag_pipeline", - return_value=pipeline, - ), - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate", - side_effect=InvokeRateLimitError("limit"), - ), - ): - with pytest.raises(InvokeRateLimitHttpError): - method(api, session, user, pipeline.id) + assert method(api, NodeRunPayload(), user, pipeline, "node") == {"ok": True} class TestDraftNodeRun: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_execution_not_found(self, app: Flask) -> None: api = RagPipelineDraftNodeRunApi() method = unwrap(api.post) @@ -531,22 +479,17 @@ class TestDraftNodeRun: with ( app.test_request_context("/", json={"inputs": empty_mapping()}), - patch.object(type(console_ns), "payload", {"inputs": empty_mapping()}), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService", return_value=service, ), ): with pytest.raises(ValueError): - method(api, user, pipeline, "node") + method(api, NodeRunRequiredPayload(inputs={}), user, pipeline, "node") class TestPublishedPipelineApis: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_publish_success(self, app: Flask, db_session_with_containers: Session) -> None: + def test_publish_success(self, app: Flask) -> None: api = PublishedRagPipelineApi() method = unwrap(api.post) @@ -557,10 +500,6 @@ class TestPublishedPipelineApis: description="test", created_by=str(uuid4()), ) - db_session_with_containers.add(pipeline) - db_session_with_containers.commit() - db_session_with_containers.expire_all() - user = make_account(id="u1") workflow = make_workflow(id=str(uuid4()), created_at=naive_utc_now()) @@ -582,10 +521,6 @@ class TestPublishedPipelineApis: class TestMiscApis: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_task_stop(self, app: Flask) -> None: api = RagPipelineTaskStopApi() method = unwrap(api.post) @@ -603,18 +538,6 @@ class TestMiscApis: stop_mock.assert_called_once() assert result["result"] == "success" - def test_transform_forbidden(self, app: Flask) -> None: - api = RagPipelineTransformApi() - method = unwrap(api.post) - - user = make_account(role=TenantAccountRole.NORMAL) - - with ( - app.test_request_context("/"), - ): - with pytest.raises(Forbidden): - method(api, MagicMock(spec=Session), user, "ds1") - def test_recommended_plugins(self, app: Flask) -> None: api = RagPipelineRecommendedPluginApi() method = unwrap(api.get) @@ -635,90 +558,12 @@ class TestMiscApis: return_value=service, ), ): - result = method(api, tenant_id, user) + result = method(api, RagPipelineRecommendedPluginQuery(type="all"), tenant_id, user) assert result == recommended_plugins service.get_recommended_plugins.assert_called_once_with("all", user, tenant_id) -class TestPublishedRagPipelineRunApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - - def test_published_run_success(self, app: Flask) -> None: - api = PublishedRagPipelineRunApi() - method = unwrap(api.post) - - pipeline = make_pipeline() - user = make_account() - session = MagicMock(spec=Session) - - payload = { - "inputs": empty_mapping(), - "datasource_type": "x", - "datasource_info_list": empty_list(), - "start_node_id": "n", - "response_mode": "blocking", - } - - with ( - app.test_request_context("/", json=payload), - patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.load_rag_pipeline", - return_value=pipeline, - ) as load_pipeline, - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate", - return_value=MagicMock(), - ) as generate, - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.helper.compact_generate_response", - return_value={"ok": True}, - ), - ): - result = method(api, session, user, pipeline.id) - assert result == {"ok": True} - - load_pipeline.assert_called_once_with(session, pipeline.id) - assert generate.call_args.kwargs["session"] is session - - def test_published_run_rate_limit(self, app: Flask) -> None: - api = PublishedRagPipelineRunApi() - method = unwrap(api.post) - - pipeline = make_pipeline() - user = make_account() - session = MagicMock(spec=Session) - - payload = { - "inputs": empty_mapping(), - "datasource_type": "x", - "datasource_info_list": empty_list(), - "start_node_id": "n", - } - - with ( - app.test_request_context("/", json=payload), - patch.object(type(console_ns), "payload", payload), - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.load_rag_pipeline", - return_value=pipeline, - ), - patch( - "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService.generate", - side_effect=InvokeRateLimitError("limit"), - ), - ): - with pytest.raises(InvokeRateLimitHttpError): - method(api, session, user, pipeline.id) - - class TestDefaultBlockConfigApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_get_block_config_success(self, app: Flask) -> None: api = DefaultRagPipelineBlockConfigApi() method = unwrap(api.get) @@ -735,7 +580,7 @@ class TestDefaultBlockConfigApi: return_value=service, ), ): - result = method(api, pipeline, "llm") + result = method(api, DefaultBlockConfigQuery(q="{}"), pipeline, "llm") assert result == {"k": "v"} def test_get_block_config_invalid_json(self, app: Flask) -> None: @@ -746,14 +591,10 @@ class TestDefaultBlockConfigApi: with app.test_request_context("/?q=bad-json"): with pytest.raises(ValueError): - method(api, pipeline, "llm") + method(api, DefaultBlockConfigQuery(q="bad-json"), pipeline, "llm") class TestPublishedAllRagPipelineApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_get_published_workflows_success(self, app: Flask) -> None: api = PublishedAllRagPipelineApi() method = unwrap(api.get) @@ -771,7 +612,7 @@ class TestPublishedAllRagPipelineApi: return_value=service, ), ): - result = method(api, user, pipeline) + result = method(api, WorkflowListQuery(), user, pipeline) assert result["items"][0]["id"] == "w1" assert result["items"][0]["graph"] == {"nodes": [], "edges": []} @@ -788,14 +629,10 @@ class TestPublishedAllRagPipelineApi: app.test_request_context("/?user_id=u2"), ): with pytest.raises(Forbidden): - method(api, user, pipeline) + method(api, WorkflowListQuery(user_id="u2"), user, pipeline) class TestRagPipelineByIdApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_patch_success(self, app: Flask) -> None: api = RagPipelineByIdApi() method = unwrap(api.patch) @@ -812,13 +649,12 @@ class TestRagPipelineByIdApi: with ( app.test_request_context("/", json=payload), - patch.object(type(console_ns), "payload", payload), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService", return_value=service, ), ): - result = method(api, user, pipeline, "w1") + result = method(api, WorkflowUpdatePayload.model_validate(payload), user, pipeline, "w1") assert result["id"] == "w1" assert result["marked_name"] == "test" @@ -831,11 +667,8 @@ class TestRagPipelineByIdApi: pipeline = make_pipeline() user = make_account() - with ( - app.test_request_context("/", json={}), - patch.object(type(console_ns), "payload", empty_mapping()), - ): - result, status = method(api, user, pipeline, "w1") + with app.test_request_context("/", json={}): + result, status = method(api, WorkflowUpdatePayload(), user, pipeline, "w1") assert status == 400 def test_delete_success(self, app: Flask) -> None: @@ -870,10 +703,6 @@ class TestRagPipelineByIdApi: class TestRagPipelineWorkflowLastRunApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_last_run_success(self, app: Flask) -> None: api = RagPipelineWorkflowLastRunApi() method = unwrap(api.get) @@ -919,10 +748,6 @@ class TestRagPipelineWorkflowLastRunApi: class TestRagPipelineWorkflowRunNodeExecutionListApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_get_node_executions_passes_current_user(self, app: Flask) -> None: api = RagPipelineWorkflowRunNodeExecutionListApi() method = unwrap(api.get) @@ -955,10 +780,6 @@ class TestRagPipelineWorkflowRunNodeExecutionListApi: class TestRagPipelineDatasourceVariableApi: - @pytest.fixture - def app(self, flask_app_with_containers: Flask) -> Flask: - return flask_app_with_containers - def test_set_datasource_variables_success(self, app: Flask) -> None: api = RagPipelineDatasourceVariableApi() method = unwrap(api.post) @@ -978,12 +799,11 @@ class TestRagPipelineDatasourceVariableApi: with ( app.test_request_context("/", json=payload), - patch.object(type(console_ns), "payload", payload), patch( "controllers.console.datasets.rag_pipeline.rag_pipeline_workflow.RagPipelineService", return_value=service, ), ): - result = method(api, user, pipeline) + result = method(api, DatasourceVariablesPayload.model_validate(payload), user, pipeline) assert result["node_id"] == "n1" assert result["process_data"] == {} diff --git a/api/tests/unit_tests/controllers/console/datasets/test_data_source.py b/api/tests/unit_tests/controllers/console/datasets/test_data_source.py index a6cb79417a7..c5339e166c2 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_data_source.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_data_source.py @@ -1,18 +1,22 @@ from __future__ import annotations import inspect -from collections.abc import Callable +from collections.abc import Callable, Iterator from datetime import UTC, datetime -from typing import cast +from typing import Literal, cast from unittest.mock import MagicMock, PropertyMock, patch -from uuid import uuid4 +from uuid import UUID import pytest from flask import Flask +from sqlalchemy import select +from sqlalchemy.orm import Session +from werkzeug.exceptions import NotFound from controllers.console.datasets import data_source as module -from controllers.console.datasets.data_source import DataSourceApi, DataSourceNotionListApi +from controllers.console.datasets.data_source import DataSourceApi, DataSourceNotionListApi, DataSourceNotionListQuery from models import Account, DataSourceOauthBinding +from models.engine import db ControllerMethod = Callable[..., tuple[dict[str, object], int]] @@ -22,10 +26,15 @@ def unwrap(func: object) -> ControllerMethod: @pytest.fixture -def flask_app() -> Flask: +def flask_app() -> Iterator[Flask]: app = Flask(__name__) app.config["TESTING"] = True - return app + app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:" + db.init_app(app) + + with app.app_context(): + DataSourceOauthBinding.__table__.create(db.engine) + yield app @pytest.fixture @@ -35,9 +44,13 @@ def current_user() -> Account: return account -def test_get_data_source_integrates_serializes_orm_binding(flask_app: Flask) -> None: +TENANT_ID = "11111111-1111-1111-1111-111111111111" +BINDING_ID = "22222222-2222-2222-2222-222222222222" + + +def _add_binding(session: Session, *, disabled: bool) -> DataSourceOauthBinding: binding = DataSourceOauthBinding( - tenant_id="tenant-1", + tenant_id=TENANT_ID, access_token="token", provider="notion", source_info={ @@ -55,24 +68,31 @@ def test_get_data_source_integrates_serializes_orm_binding(flask_app: Flask) -> } ], }, + disabled=disabled, ) - binding.id = "binding-1" + binding.id = BINDING_ID binding.created_at = datetime(2026, 5, 25, 1, 2, 3, tzinfo=UTC) - binding.disabled = False + session.add(binding) + session.commit() + return binding - with ( - flask_app.test_request_context("/"), - patch.object(module.db.session, "scalars", return_value=MagicMock(all=lambda: [binding])), - ): - response, status = unwrap(DataSourceApi().get)(DataSourceApi(), "tenant-1") + +def test_get_data_source_integrates_serializes_orm_binding( + flask_app: Flask, +) -> None: + binding = _add_binding(db.session, disabled=False) + expected_created_at = int(binding.created_at.timestamp()) + + with flask_app.test_request_context("/"): + response, status = unwrap(DataSourceApi().get)(DataSourceApi(), TENANT_ID) assert status == 200 assert response == { "data": [ { - "id": "binding-1", + "id": BINDING_ID, "provider": "notion", - "created_at": 1779670923, + "created_at": expected_created_at, "is_bound": True, "disabled": False, "source_info": { @@ -96,34 +116,75 @@ def test_get_data_source_integrates_serializes_orm_binding(flask_app: Flask) -> } -def test_get_data_source_integrates_preserves_empty_list_when_no_binding(flask_app: Flask) -> None: - with ( - flask_app.test_request_context("/"), - patch.object(module.db.session, "scalars", return_value=MagicMock(all=lambda: [])), - ): - response, status = unwrap(DataSourceApi().get)(DataSourceApi(), "tenant-1") +def test_get_data_source_integrates_preserves_empty_list_when_no_binding( + flask_app: Flask, +) -> None: + with flask_app.test_request_context("/"): + response, status = unwrap(DataSourceApi().get)(DataSourceApi(), TENANT_ID) assert status == 200 assert response == {"data": []} -def test_patch_data_source_binding_uses_injected_session(flask_app: Flask) -> None: - binding = MagicMock(disabled=True) - session = MagicMock() - session.scalar.return_value = binding +@pytest.mark.parametrize( + ("disabled", "action", "expected_disabled"), + [(True, "enable", False), (False, "disable", True)], +) +@pytest.mark.parametrize("sqlite_session", [(DataSourceOauthBinding,)], indirect=True) +def test_patch_data_source_binding_updates_state( + flask_app: Flask, + sqlite_session: Session, + disabled: bool, + action: Literal["enable", "disable"], + expected_disabled: bool, +) -> None: + _add_binding(sqlite_session, disabled=disabled) + sqlite_session.expunge_all() with flask_app.test_request_context("/"): - response, status = unwrap(DataSourceApi().patch)(DataSourceApi(), session, "tenant-1", uuid4(), "enable") + response, status = unwrap(DataSourceApi().patch)( + DataSourceApi(), sqlite_session, TENANT_ID, UUID(BINDING_ID), action + ) + sqlite_session.flush() + sqlite_session.expire_all() + binding = sqlite_session.scalar(select(DataSourceOauthBinding).where(DataSourceOauthBinding.id == BINDING_ID)) assert status == 200 assert response == {"result": "success"} - assert binding.disabled is False - session.scalar.assert_called_once() - session.add.assert_not_called() - session.commit.assert_not_called() + assert binding is not None + assert binding.disabled is expected_disabled -def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask, current_user: Account) -> None: +@pytest.mark.parametrize("sqlite_session", [(DataSourceOauthBinding,)], indirect=True) +def test_patch_data_source_binding_rejects_unknown_binding( + flask_app: Flask, + sqlite_session: Session, +) -> None: + with flask_app.test_request_context("/"), pytest.raises(NotFound, match="Data source binding not found"): + unwrap(DataSourceApi().patch)(DataSourceApi(), sqlite_session, TENANT_ID, UUID(BINDING_ID), "enable") + + +@pytest.mark.parametrize(("disabled", "action"), [(False, "enable"), (True, "disable")]) +@pytest.mark.parametrize("sqlite_session", [(DataSourceOauthBinding,)], indirect=True) +def test_patch_data_source_binding_rejects_current_state( + flask_app: Flask, + sqlite_session: Session, + disabled: bool, + action: Literal["enable", "disable"], +) -> None: + _add_binding(sqlite_session, disabled=disabled) + sqlite_session.expunge_all() + + with flask_app.test_request_context("/"), pytest.raises(ValueError): + unwrap(DataSourceApi().patch)(DataSourceApi(), sqlite_session, TENANT_ID, UUID(BINDING_ID), action) + + +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_notion_pre_import_pages_serializes_frontend_list_shape( + flask_app: Flask, + current_user: Account, + sqlite_session: Session, +) -> None: page = MagicMock( page_id="page-1", page_name="Page", @@ -145,8 +206,6 @@ def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask get_online_document_pages=MagicMock(return_value=iter([online_document_message])), datasource_provider_type=MagicMock(return_value="online_document"), ) - session = MagicMock() - with ( flask_app.test_request_context("/?credential_id=credential-1"), patch.object( @@ -158,7 +217,11 @@ def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask patch("core.datasource.datasource_manager.DatasourceManager.get_datasource_runtime", return_value=runtime), ): response, status = unwrap(DataSourceNotionListApi().get)( - DataSourceNotionListApi(), session, "tenant-1", current_user + DataSourceNotionListApi(), + DataSourceNotionListQuery(credential_id="credential-1"), + sqlite_session, + "tenant-1", + current_user, ) assert status == 200 @@ -183,3 +246,50 @@ def test_notion_pre_import_pages_serializes_frontend_list_shape(flask_app: Flask } runtime.get_online_document_pages.assert_called_once() assert runtime.get_online_document_pages.call_args.kwargs["datasource_parameters"] == {} + + +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_notion_pre_import_pages_rejects_missing_credential( + flask_app: Flask, + current_user: Account, + sqlite_session: Session, +) -> None: + with ( + flask_app.test_request_context("/?credential_id=credential-1"), + patch.object(module.DatasourceProviderService, "get_datasource_credentials", return_value=None), + pytest.raises(NotFound, match="Credential not found"), + ): + unwrap(DataSourceNotionListApi().get)( + DataSourceNotionListApi(), + DataSourceNotionListQuery(credential_id="credential-1"), + sqlite_session, + TENANT_ID, + current_user, + ) + + +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_notion_pre_import_pages_rejects_non_notion_dataset( + flask_app: Flask, + current_user: Account, + sqlite_session: Session, +) -> None: + dataset = MagicMock(data_source_type="other_type") + + with ( + flask_app.test_request_context("/?credential_id=credential-1&dataset_id=dataset-1"), + patch.object( + module.DatasourceProviderService, + "get_datasource_credentials", + return_value={"token": "token"}, + ), + patch.object(module.DatasetService, "get_dataset", return_value=dataset), + pytest.raises(ValueError, match="Dataset is not notion type"), + ): + unwrap(DataSourceNotionListApi().get)( + DataSourceNotionListApi(), + DataSourceNotionListQuery(credential_id="credential-1", dataset_id="dataset-1"), + sqlite_session, + TENANT_ID, + current_user, + ) diff --git a/api/tests/unit_tests/controllers/console/datasets/test_data_source_notion_apis.py b/api/tests/unit_tests/controllers/console/datasets/test_data_source_notion_apis.py new file mode 100644 index 00000000000..ca3d18973cf --- /dev/null +++ b/api/tests/unit_tests/controllers/console/datasets/test_data_source_notion_apis.py @@ -0,0 +1,173 @@ +"""Unit tests for controllers.console.datasets.data_source Notion endpoints.""" + +from __future__ import annotations + +import inspect +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask +from sqlalchemy.orm import Session +from werkzeug.exceptions import NotFound + +from controllers.console.datasets.data_source import ( + DataSourceNotionDatasetSyncApi, + DataSourceNotionDocumentSyncApi, + DataSourceNotionIndexingEstimateApi, + DataSourceNotionPreviewApi, + DataSourceNotionPreviewQuery, + NotionEstimatePayload, +) +from core.rag.index_processor.constant.index_type import IndexStructureType +from models import Account + + +@pytest.fixture +def current_user() -> Account: + account = Account(name="Test User", email="u1@example.com") + account.id = "u1" + return account + + +class TestDataSourceNotionPreviewApi: + def test_get_preview_success(self, app: Flask) -> None: + api = DataSourceNotionPreviewApi() + method = inspect.unwrap(api.get) + + extractor = MagicMock(extract=lambda: [MagicMock(page_content="hello")]) + + with ( + app.test_request_context("/?credential_id=c1"), + patch( + "controllers.console.datasets.data_source.DatasourceProviderService.get_datasource_credentials", + return_value={"integration_secret": "t"}, + ), + patch( + "controllers.console.datasets.data_source.NotionExtractor", + return_value=extractor, + ), + ): + response, status = method(api, DataSourceNotionPreviewQuery(credential_id="c1"), "tenant-1", "p1", "page") + + assert status == 200 + + +class TestDataSourceNotionIndexingEstimateApi: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_post_indexing_estimate_success(self, app: Flask, sqlite_session: Session) -> None: + api = DataSourceNotionIndexingEstimateApi() + method = inspect.unwrap(api.post) + + empty_rules: dict[str, object] = {} + payload: dict[str, object] = { + "notion_info_list": [ + { + "workspace_id": "w1", + "credential_id": "c1", + "pages": [{"page_id": "p1", "type": "page"}], + } + ], + "process_rule": {"rules": empty_rules}, + "doc_form": IndexStructureType.PARAGRAPH_INDEX, + "doc_language": "English", + } + + with ( + app.test_request_context("/", method="POST", json=payload, headers={"Content-Type": "application/json"}), + patch( + "controllers.console.datasets.data_source.DocumentService.estimate_args_validate", + ), + patch( + "controllers.console.datasets.data_source.IndexingRunner.indexing_estimate", + return_value=MagicMock(model_dump=lambda: {"total_pages": 1}), + ), + ): + response, status = method(api, NotionEstimatePayload.model_validate(payload), sqlite_session, "tenant-1") + + assert status == 200 + + +class TestDataSourceNotionDatasetSyncApi: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_get_success(self, app: Flask, sqlite_session: Session) -> None: + api = DataSourceNotionDatasetSyncApi() + method = inspect.unwrap(api.get) + + with ( + app.test_request_context("/"), + patch( + "controllers.console.datasets.data_source.DatasetService.get_dataset", + return_value=MagicMock(), + ), + patch( + "controllers.console.datasets.data_source.DocumentService.get_document_by_dataset_id", + return_value=[MagicMock(id="d1")], + ), + patch( + "controllers.console.datasets.data_source.document_indexing_sync_task.delay", + return_value=None, + ), + ): + response, status = method(api, sqlite_session, "ds-1") + + assert status == 200 + + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_get_dataset_not_found(self, app: Flask, sqlite_session: Session) -> None: + api = DataSourceNotionDatasetSyncApi() + method = inspect.unwrap(api.get) + + with ( + app.test_request_context("/"), + patch( + "controllers.console.datasets.data_source.DatasetService.get_dataset", + return_value=None, + ), + ): + with pytest.raises(NotFound): + method(api, sqlite_session, "ds-1") + + +class TestDataSourceNotionDocumentSyncApi: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_get_success(self, app: Flask, sqlite_session: Session) -> None: + api = DataSourceNotionDocumentSyncApi() + method = inspect.unwrap(api.get) + + with ( + app.test_request_context("/"), + patch( + "controllers.console.datasets.data_source.DatasetService.get_dataset", + return_value=MagicMock(), + ), + patch( + "controllers.console.datasets.data_source.DocumentService.get_document", + return_value=MagicMock(), + ), + patch( + "controllers.console.datasets.data_source.document_indexing_sync_task.delay", + return_value=None, + ), + ): + response, status = method(api, sqlite_session, "ds-1", "doc-1") + + assert status == 200 + + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_get_document_not_found(self, app: Flask, sqlite_session: Session) -> None: + api = DataSourceNotionDocumentSyncApi() + method = inspect.unwrap(api.get) + + with ( + app.test_request_context("/"), + patch( + "controllers.console.datasets.data_source.DatasetService.get_dataset", + return_value=MagicMock(), + ), + patch( + "controllers.console.datasets.data_source.DocumentService.get_document", + return_value=None, + ), + ): + with pytest.raises(NotFound): + method(api, sqlite_session, "ds-1", "doc-1") diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py index 541c4b3f392..1bbc48896cd 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py @@ -18,6 +18,7 @@ from controllers.console.datasets.datasets import ( DatasetApiDeleteApi, DatasetApiKeyApi, DatasetAutoDisableLogApi, + DatasetCreatePayload, DatasetEnableApiApi, DatasetErrorDocs, DatasetIndexingEstimateApi, @@ -28,18 +29,24 @@ from controllers.console.datasets.datasets import ( DatasetRelatedAppListApi, DatasetRetrievalSettingApi, DatasetRetrievalSettingMockApi, + DatasetUpdatePayload, DatasetUseCheckApi, + IndexingEstimatePayload, + _get_retrieval_methods_by_vector_type, ) from controllers.console.datasets.error import DatasetInUseError, DatasetNameDuplicateError, IndexingEstimateError from core.entities.knowledge_entities import IndexingEstimate from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.provider_manager import ProviderManager +from core.rag.datasource.vdb.vector_type import VectorType from core.rag.index_processor.constant.index_type import IndexStructureType +from core.rag.retrieval.retrieval_methods import RetrievalMethod from extensions.storage.storage_type import StorageType from models.account import Account, TenantAccountRole from models.dataset import Dataset, DatasetQuery, Document from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, IndexingStatus from models.model import ApiToken, App, AppMode, IconType, UploadFile +from services.dataset_ref_service import DatasetRef from services.dataset_service import DatasetPermissionService, DatasetService from services.enterprise import rbac_service as enterprise_rbac_service @@ -486,7 +493,7 @@ class TestDatasetListApiPost: patch.object(type(console_ns), "payload", payload), patch.object(DatasetService, "create_empty_dataset", return_value=dataset), ): - _, status = method(api, MagicMock(), "tenant-1", user) + _, status = method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", user) assert status == 201 def test_post_forbidden(self, app: Flask): @@ -496,7 +503,7 @@ class TestDatasetListApiPost: user = make_account(TenantAccountRole.NORMAL) with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", user) + method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", user) def test_post_duplicate_name(self, app: Flask): api = DatasetListApi() @@ -511,14 +518,14 @@ class TestDatasetListApiPost: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, MagicMock(), "tenant-1", user) + method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", user) def test_post_invalid_payload_missing_name(self, app: Flask): api = DatasetListApi() method = unwrap(api.post) with app.test_request_context("/datasets", json={}), patch.object(type(console_ns), "payload", {}): with pytest.raises(ValueError): - method(api, MagicMock(), "tenant-1", make_account()) + method(api, DatasetCreatePayload(), MagicMock(), "tenant-1", make_account()) def test_post_invalid_indexing_technique(self, app: Flask): api = DatasetListApi() @@ -526,7 +533,7 @@ class TestDatasetListApiPost: payload = {"name": "bad", "indexing_technique": "invalid-tech"} with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(ValueError, match="Invalid indexing technique"): - method(api, MagicMock(), "tenant-1", make_account()) + method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", make_account()) def test_post_invalid_provider(self, app: Flask): api = DatasetListApi() @@ -534,7 +541,7 @@ class TestDatasetListApiPost: payload = {"name": "bad", "provider": "unknown"} with app.test_request_context("/datasets", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(ValueError, match="Invalid provider"): - method(api, MagicMock(), "tenant-1", make_account()) + method(api, DatasetCreatePayload(**payload), MagicMock(), "tenant-1", make_account()) class TestDatasetApiGet: @@ -689,7 +696,7 @@ class TestDatasetApiPatch: patch.object(DatasetService, "update_dataset", return_value=dataset), patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=[]), ): - result, status = method(api, MagicMock(), tenant_id, user, dataset_id) + result, status = method(api, DatasetUpdatePayload(), MagicMock(), tenant_id, user, dataset_id) assert status == 200 assert result["partial_member_list"] == [] @@ -701,7 +708,7 @@ class TestDatasetApiPatch: patch.object(DatasetService, "get_dataset", return_value=None), ): with pytest.raises(NotFound, match="Dataset not found"): - method(api, MagicMock(), "tenant-1", make_account(), "missing") + method(api, DatasetUpdatePayload(), MagicMock(), "tenant-1", make_account(), "missing") def test_patch_permission_denied(self, app: Flask): api = DatasetApi() @@ -716,7 +723,7 @@ class TestDatasetApiPatch: patch.object(DatasetPermissionService, "check_permission", side_effect=Forbidden("no permission")), ): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant", make_account(), dataset_id) + method(api, DatasetUpdatePayload(), MagicMock(), "tenant", make_account(), dataset_id) def test_patch_partial_members_update(self, app: Flask): api = DatasetApi() @@ -733,7 +740,7 @@ class TestDatasetApiPatch: patch.object(DatasetPermissionService, "update_partial_member_list", return_value=None), patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=["u1", "u2"]), ): - result, _ = method(api, MagicMock(), "tenant", make_account(), dataset_id) + result, _ = method(api, DatasetUpdatePayload(), MagicMock(), "tenant", make_account(), dataset_id) assert result["partial_member_list"] == ["u1", "u2"] def test_patch_clear_partial_members(self, app: Flask): @@ -751,7 +758,7 @@ class TestDatasetApiPatch: patch.object(DatasetPermissionService, "clear_partial_member_list", return_value=None), patch.object(DatasetPermissionService, "get_dataset_partial_member_list", return_value=[]), ): - result, _ = method(api, MagicMock(), "tenant", make_account(), dataset_id) + result, _ = method(api, DatasetUpdatePayload(), MagicMock(), "tenant", make_account(), dataset_id) assert result["partial_member_list"] == [] @@ -805,29 +812,65 @@ class TestDatasetApiDelete: class TestDatasetUseCheckApi: - def test_get_use_check_true(self, app: Flask): + @pytest.mark.parametrize("is_using", [True, False]) + def test_get_use_check(self, app: Flask, is_using: bool): api = DatasetUseCheckApi() method = unwrap(api.get) dataset_id = "dataset-id" + dataset = make_dataset(id=dataset_id) + current_user = make_account() + session = MagicMock() with ( app.test_request_context(f"/datasets/{dataset_id}/use-check"), - patch.object(DatasetService, "dataset_use_check", return_value=True), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset, + patch.object(DatasetService, "check_dataset_permission") as check_permission, + patch.object(DatasetService, "dataset_use_check", return_value=is_using) as dataset_use_check, ): - result, status = method(api, MagicMock(), dataset_id) + result, status = method(api, session, "tenant-1", current_user, dataset_id) assert status == 200 - assert result == {"is_using": True} + assert result == {"is_using": is_using} + get_dataset.assert_called_once_with(dataset_id, "tenant-1", session=session) + check_permission.assert_called_once_with(dataset, current_user, session) + dataset_use_check.assert_called_once_with(DatasetRef("tenant-1", dataset_id), session) - def test_get_use_check_false(self, app: Flask): + def test_get_use_check_relies_on_rbac_in_rbac_mode(self, app: Flask): api = DatasetUseCheckApi() method = unwrap(api.get) - dataset_id = "dataset-id" + dataset = make_dataset(id="dataset-id") + session = MagicMock() with ( - app.test_request_context(f"/datasets/{dataset_id}/use-check"), + app.test_request_context("/datasets/dataset-id/use-check"), + patch("controllers.console.datasets.datasets.dify_config.RBAC_ENABLED", True), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission") as check_permission, patch.object(DatasetService, "dataset_use_check", return_value=False), ): - result, status = method(api, MagicMock(), dataset_id) + _, status = method(api, session, "tenant-1", make_account(), "dataset-id") + assert status == 200 - assert result == {"is_using": False} + check_permission.assert_not_called() + + +@pytest.mark.parametrize( + "api_cls", + [DatasetUseCheckApi, DatasetIndexingStatusApi, DatasetErrorDocs, DatasetAutoDisableLogApi], +) +def test_dataset_scoped_read_permission_denied(app: Flask, api_cls): + api = api_cls() + method = unwrap(api.get) + dataset = make_dataset(id="dataset-1") + session = MagicMock() + with ( + app.test_request_context("/"), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), + patch.object( + DatasetService, + "check_dataset_permission", + side_effect=services.errors.account.NoPermissionError("no permission"), + ), + ): + with pytest.raises(Forbidden, match="no permission"): + method(api, session, "tenant-1", make_account(), "dataset-1") class TestDatasetQueryApi: @@ -1011,7 +1054,12 @@ class TestDatasetIndexingEstimateApi: patch("controllers.console.datasets.datasets.DocumentService.estimate_args_validate", return_value=None), patch("controllers.console.datasets.datasets.IndexingRunner.indexing_estimate", return_value=mock_response), ): - response, status = method(api, session, "tenant-1") + response, status = method( + api, + IndexingEstimatePayload(**payload), + session, + "tenant-1", + ) assert status == 200 assert response == { "tokens": 0, @@ -1033,7 +1081,12 @@ class TestDatasetIndexingEstimateApi: patch("controllers.console.datasets.datasets.DocumentService.estimate_args_validate", return_value=None), ): with pytest.raises(NotFound): - method(api, session, "tenant-1") + method( + api, + IndexingEstimatePayload(**payload), + session, + "tenant-1", + ) def test_post_llm_bad_request_error(self, app: Flask): api = DatasetIndexingEstimateApi() @@ -1052,7 +1105,12 @@ class TestDatasetIndexingEstimateApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, session, "tenant-1") + method( + api, + IndexingEstimatePayload(**payload), + session, + "tenant-1", + ) def test_post_provider_token_not_init(self, app: Flask): api = DatasetIndexingEstimateApi() @@ -1071,7 +1129,12 @@ class TestDatasetIndexingEstimateApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, session, "tenant-1") + method( + api, + IndexingEstimatePayload(**payload), + session, + "tenant-1", + ) def test_post_generic_exception(self, app: Flask): api = DatasetIndexingEstimateApi() @@ -1089,7 +1152,12 @@ class TestDatasetIndexingEstimateApi: ), ): with pytest.raises(IndexingEstimateError): - method(api, session, "tenant-1") + method( + api, + IndexingEstimatePayload(**payload), + session, + "tenant-1", + ) class TestDatasetRelatedAppListApi: @@ -1210,6 +1278,8 @@ class TestDatasetIndexingStatusApi: def test_get_success_with_documents(self, app: Flask): api = DatasetIndexingStatusApi() method = unwrap(api.get) + dataset = make_dataset(id="dataset-1") + current_user = make_account() document = MagicMock() document.id = "doc-1" document.indexing_status = "completed" @@ -1224,28 +1294,43 @@ class TestDatasetIndexingStatusApi: session = MagicMock() session.scalars.return_value.all.return_value = [document] session.scalar.return_value = 3 - with app.test_request_context("/"): - response, status = method(api, session, "tenant-1", "dataset-1") + with ( + app.test_request_context("/"), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset, + patch.object(DatasetService, "check_dataset_permission") as check_permission, + ): + response, status = method(api, session, "tenant-1", current_user, "dataset-1") assert status == 200 assert "data" in response assert len(response["data"]) == 1 item = response["data"][0] assert item["completed_segments"] == 3 assert item["total_segments"] == 3 + get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session) + check_permission.assert_called_once_with(dataset, current_user, session) + assert {"dataset-1", "tenant-1"} <= set(session.scalars.call_args.args[0].compile().params.values()) + for segment_count_call in session.scalar.call_args_list: + assert {"dataset-1", "tenant-1", "doc-1"} <= set(segment_count_call.args[0].compile().params.values()) def test_get_success_no_documents(self, app: Flask): api = DatasetIndexingStatusApi() method = unwrap(api.get) + dataset = make_dataset(id="dataset-1") session = MagicMock() session.scalars.return_value.all.return_value = [] - with app.test_request_context("/"): - response, status = method(api, session, "tenant-1", "dataset-1") + with ( + app.test_request_context("/"), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), + ): + response, status = method(api, session, "tenant-1", make_account(), "dataset-1") assert status == 200 assert response == {"data": []} def test_segment_counts_different_values(self, app: Flask): api = DatasetIndexingStatusApi() method = unwrap(api.get) + dataset = make_dataset(id="dataset-1") document = MagicMock() document.id = "doc-1" document.indexing_status = "indexing" @@ -1260,8 +1345,12 @@ class TestDatasetIndexingStatusApi: session = MagicMock() session.scalars.return_value.all.return_value = [document] session.scalar.side_effect = [2, 5] - with app.test_request_context("/"): - response, status = method(api, session, "tenant-1", "dataset-1") + with ( + app.test_request_context("/"), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), + ): + response, status = method(api, session, "tenant-1", make_account(), "dataset-1") assert status == 200 item = response["data"][0] assert item["completed_segments"] == 2 @@ -1272,18 +1361,20 @@ class TestDatasetApiKeyApi: def test_get_api_keys_success(self, app: Flask): api = DatasetApiKeyApi() method = unwrap(api.get) - mock_key_1 = MagicMock(spec=ApiToken) - mock_key_1.id = "key-1" - mock_key_1.type = "dataset" - mock_key_1.token = "ds-abc" - mock_key_1.last_used_at = None - mock_key_1.created_at = None - mock_key_2 = MagicMock(spec=ApiToken) - mock_key_2.id = "key-2" - mock_key_2.type = "dataset" - mock_key_2.token = "ds-def" - mock_key_2.last_used_at = None - mock_key_2.created_at = None + mock_key_1 = ApiToken( + id="key-1", + type="dataset", + token="ds-abc", + last_used_at=None, + created_at=None, + ) + mock_key_2 = ApiToken( + id="key-2", + type="dataset", + token="ds-def", + last_used_at=None, + created_at=None, + ) session = MagicMock() session.scalars.return_value.all.return_value = [mock_key_1, mock_key_2] with app.test_request_context("/"): @@ -1355,27 +1446,47 @@ class TestDatasetApiDeleteApi: class TestDatasetEnableApiApi: - def test_enable_api(self, app: Flask): + @pytest.mark.parametrize(("status_value", "enabled"), [("enable", True), ("disable", False)]) + def test_update_api_status(self, app: Flask, status_value: str, enabled: bool): api = DatasetEnableApiApi() method = unwrap(api.post) + dataset = make_dataset(id="dataset-1") + current_user = make_account() + session = MagicMock() with ( app.test_request_context("/"), - patch("controllers.console.datasets.datasets.DatasetService.update_dataset_api_status", return_value=None), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset, + patch.object(DatasetService, "check_dataset_permission") as check_permission, + patch.object(DatasetService, "update_dataset_api_status") as update_status, ): - response, status = method(api, MagicMock(), "dataset-1", "enable") + response, status = method(api, session, "tenant-1", current_user, "dataset-1", status_value) assert status == 200 assert response["result"] == "success" + get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session) + check_permission.assert_called_once_with(dataset, current_user, session) + update_status.assert_called_once_with(dataset, enabled, current_user, session) - def test_disable_api(self, app: Flask): + def test_rejects_non_editor(self, app: Flask): api = DatasetEnableApiApi() method = unwrap(api.post) + dataset = make_dataset(id="dataset-1") + session = MagicMock() with ( app.test_request_context("/"), - patch("controllers.console.datasets.datasets.DatasetService.update_dataset_api_status", return_value=None), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), + patch.object(DatasetService, "update_dataset_api_status") as update_status, ): - response, status = method(api, MagicMock(), "dataset-1", "disable") - assert status == 200 - assert response["result"] == "success" + with pytest.raises(Forbidden): + method( + api, + session, + "tenant-1", + make_account(TenantAccountRole.NORMAL), + "dataset-1", + "enable", + ) + update_status.assert_not_called() class TestDatasetApiBaseUrlApi: @@ -1425,6 +1536,28 @@ class TestDatasetRetrievalSettingApi: response = method(api) assert "retrieval_method" in response + def test_tidb_vector_returns_semantic_only_when_fulltext_disabled(self): + with patch( + "controllers.console.datasets.datasets.dify_config.TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH", + False, + ): + response = _get_retrieval_methods_by_vector_type(VectorType.TIDB_VECTOR) + + assert response["retrieval_method"] == [RetrievalMethod.SEMANTIC_SEARCH.value] + + def test_tidb_vector_returns_full_methods_when_fulltext_enabled(self): + with patch( + "controllers.console.datasets.datasets.dify_config.TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH", + True, + ): + response = _get_retrieval_methods_by_vector_type(VectorType.TIDB_VECTOR) + + assert response["retrieval_method"] == [ + RetrievalMethod.SEMANTIC_SEARCH.value, + RetrievalMethod.FULL_TEXT_SEARCH.value, + RetrievalMethod.HYBRID_SEARCH.value, + ] + class TestDatasetRetrievalSettingMockApi: def test_get_success(self, app: Flask): @@ -1447,27 +1580,35 @@ class TestDatasetErrorDocs: method = unwrap(api.get) dataset = make_dataset(id="dataset-1") error_doc = make_document_status(id="error-doc", indexing_status=IndexingStatus.ERROR, error="failed") + current_user = make_account() + session = MagicMock() with ( app.test_request_context("/"), - patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset, + patch.object(DatasetService, "check_dataset_permission") as check_permission, patch( - "controllers.console.datasets.datasets.DocumentService.get_error_documents_by_dataset_id", + "controllers.console.datasets.datasets.DocumentService.get_error_documents_by_dataset_ref", return_value=[error_doc], - ), + ) as get_error_documents, ): - response, status = method(api, MagicMock(), "dataset-1") + response, status = method(api, session, "tenant-1", current_user, "dataset-1") assert status == 200 assert response["total"] == 1 + get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session) + check_permission.assert_called_once_with(dataset, current_user, session) + get_error_documents.assert_called_once_with(DatasetRef("tenant-1", "dataset-1"), session) def test_get_dataset_not_found(self, app: Flask): api = DatasetErrorDocs() method = unwrap(api.get) + session = MagicMock() with ( app.test_request_context("/"), - patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=None), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=None) as get_dataset, ): with pytest.raises(NotFound): - method(api, MagicMock(), "dataset-1") + method(api, session, "tenant-1", make_account(), "dataset-1") + get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session) class TestDatasetPermissionUserListApi: @@ -1511,23 +1652,29 @@ class TestDatasetAutoDisableLogApi: method = unwrap(api.get) dataset = make_dataset(id="dataset-1") logs = {"document_ids": ["doc-1"], "count": 1} + current_user = make_account() + session = MagicMock() with ( app.test_request_context("/"), - patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset), - patch( - "controllers.console.datasets.datasets.DatasetService.get_dataset_auto_disable_logs", return_value=logs - ), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset, + patch.object(DatasetService, "check_dataset_permission") as check_permission, + patch.object(DatasetService, "get_dataset_auto_disable_logs", return_value=logs) as get_logs, ): - response, status = method(api, MagicMock(), "dataset-1") + response, status = method(api, session, "tenant-1", current_user, "dataset-1") assert status == 200 assert response == logs + get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session) + check_permission.assert_called_once_with(dataset, current_user, session) + get_logs.assert_called_once_with(DatasetRef("tenant-1", "dataset-1"), session) def test_get_dataset_not_found(self, app: Flask): api = DatasetAutoDisableLogApi() method = unwrap(api.get) + session = MagicMock() with ( app.test_request_context("/"), - patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=None), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=None) as get_dataset, ): with pytest.raises(NotFound): - method(api, MagicMock(), "dataset-1") + method(api, session, "tenant-1", make_account(), "dataset-1") + get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session) diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py index 85ffd1b2d20..be4085dfb4a 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document.py @@ -1,10 +1,11 @@ import datetime from inspect import unwrap from types import SimpleNamespace -from unittest.mock import ANY, MagicMock, patch +from unittest.mock import MagicMock, patch import pytest from flask import Flask +from sqlalchemy import select from werkzeug.exceptions import Forbidden, NotFound import services @@ -21,13 +22,18 @@ from controllers.console.datasets.datasets_document import ( DocumentIndexingEstimateApi, DocumentIndexingStatusApi, DocumentMetadataApi, + DocumentMetadataUpdatePayload, + DocumentPauseApi, DocumentPipelineExecutionLogApi, DocumentProcessingApi, + DocumentRecoverApi, DocumentRenameApi, + DocumentResource, DocumentRetryApi, DocumentStatusApi, DocumentSummaryStatusApi, GetProcessRuleApi, + WebsiteDocumentSyncApi, ) from controllers.console.datasets.error import ( DocumentAlreadyFinishedError, @@ -38,9 +44,15 @@ from controllers.console.datasets.error import ( ) from core.entities.knowledge_entities import IndexingEstimate from core.rag.index_processor.constant.index_type import IndexStructureType -from models.dataset import Dataset +from models.dataset import Dataset, DatasetPermissionEnum from models.dataset import Document as DatasetDocument from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus +from services.dataset_ref_service import DatasetRef, DocumentRef +from services.enterprise.rbac_service import RBACResourceWhitelistScope, ReplaceMemberBindings +from services.vector_space_admission_service import ( + VECTOR_SPACE_ADMISSION_ERROR_CODE, + format_vector_space_admission_error, +) def make_serializable_document(**overrides): @@ -186,6 +198,12 @@ def tenant_ctx(): return (MagicMock(is_dataset_editor=True, id="u1"), "tenant-1") +@pytest.fixture(autouse=True) +def bypass_knowledge_rate_limit(): + with patch("controllers.console.datasets.datasets_document.check_knowledge_rate_limit") as check: + yield check + + @pytest.fixture def patch_tenant(tenant_ctx): return tenant_ctx @@ -217,6 +235,14 @@ def patch_dataset(dataset): yield +@pytest.fixture +def patch_scoped_dataset(dataset): + with patch( + "controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant", return_value=dataset + ): + yield + + @pytest.fixture def patch_permission(): with patch( @@ -476,6 +502,7 @@ class TestDatasetInitApi: with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), + patch("controllers.console.datasets.datasets_document.dify_config.RBAC_ENABLED", True), patch( "controllers.console.datasets.datasets_document.DocumentService.document_create_args_validate", return_value=None, @@ -484,6 +511,12 @@ class TestDatasetInitApi: "controllers.console.datasets.datasets_document.DocumentService.save_document_without_dataset_id", return_value=(created_dataset, [created_document], "batch-init"), ), + patch( + "controllers.console.datasets.datasets_document.enterprise_rbac_service.RBACService.DatasetAccess.replace_whitelist" + ) as replace_whitelist, + patch( + "controllers.console.datasets.datasets_document.initialize_created_app_rbac_access_task" + ) as initialize_rbac_task, ): response = method(api, session, tenant_id, user) assert response["dataset"]["id"] == "ds-1" @@ -491,6 +524,65 @@ class TestDatasetInitApi: assert response["documents"][0]["data_source_info"] == {} assert response["documents"][0]["doc_metadata"] == [] assert response["batch"] == "batch-init" + assert created_dataset.permission == DatasetPermissionEnum.ALL_TEAM + replace_whitelist.assert_called_once_with( + tenant_id, + user.id, + created_dataset.id, + ReplaceMemberBindings(scope=RBACResourceWhitelistScope.ALL), + ) + initialize_rbac_task.delay.assert_called_once_with(tenant_id, user.id, dataset_id=created_dataset.id) + + +class TestDocumentResource: + def test_get_document_resolves_owner_chain(self, dataset): + api = DocumentResource() + session = MagicMock() + user = MagicMock() + document = MagicMock() + + with ( + patch( + "controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant", + return_value=dataset, + ) as get_dataset, + patch( + "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission" + ) as check_permission, + patch( + "controllers.console.datasets.datasets_document.DatasetRefService.get_document_by_ref", + return_value=document, + ) as get_document, + ): + assert api.get_document(session, "ds-1", "doc-1", user, "tenant-1") is document + + get_dataset.assert_called_once_with("ds-1", "tenant-1", session=session) + check_permission.assert_called_once_with(dataset, user, session) + get_document.assert_called_once_with( + DocumentRef(dataset=DatasetRef(tenant_id="tenant-1", dataset_id="ds-1"), document_id="doc-1"), + session=session, + ) + + def test_get_document_relies_on_rbac_in_rbac_mode(self, dataset): + api = DocumentResource() + session = MagicMock() + with ( + patch("controllers.console.datasets.datasets_document.dify_config.RBAC_ENABLED", True), + patch( + "controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant", + return_value=dataset, + ), + patch( + "controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission" + ) as check_permission, + patch( + "controllers.console.datasets.datasets_document.DatasetRefService.get_document_by_ref", + return_value=MagicMock(), + ), + ): + api.get_document(session, "ds-1", "doc-1", MagicMock(), "tenant-1") + + check_permission.assert_not_called() class TestDocumentApi: @@ -627,15 +719,16 @@ class TestDocumentMetadataApi: payload = {"doc_type": "invoice", "doc_metadata": {"amount": 10, "invalid": "x"}} schema = {"amount": int} session = MagicMock() + req_data = DocumentMetadataUpdatePayload.model_validate(payload) with ( - app.test_request_context("/", json=payload), + app.test_request_context("/"), patch.object(api, "get_document", return_value=doc), patch( "controllers.console.datasets.datasets_document.DocumentService.DOCUMENT_METADATA_SCHEMA", {"invoice": schema}, ), ): - method(api, session, tenant_id, user, "ds-1", "doc-1") + method(api, req_data, session, tenant_id, user, "ds-1", "doc-1") assert doc.doc_metadata == {"amount": 10} def test_put_success(self, app: Flask, patch_tenant): @@ -645,32 +738,35 @@ class TestDocumentMetadataApi: document = MagicMock() payload = {"doc_type": "others", "doc_metadata": {"a": 1}} session = MagicMock() + req_data = DocumentMetadataUpdatePayload.model_validate(payload) with ( - app.test_request_context("/", json=payload), + app.test_request_context("/"), patch.object(api, "get_document", return_value=document), patch( "controllers.console.datasets.datasets_document.DocumentService.DOCUMENT_METADATA_SCHEMA", {"others": {}}, ), ): - response, status = method(api, session, tenant_id, user, "ds-1", "doc-1") + response, status = method(api, req_data, session, tenant_id, user, "ds-1", "doc-1") assert status == 200 def test_put_invalid_payload(self, app: Flask, patch_tenant): api = DocumentMetadataApi() method = unwrap(api.put) user, tenant_id = patch_tenant - with app.test_request_context("/", json={}), patch.object(api, "get_document", return_value=MagicMock()): + req_data = DocumentMetadataUpdatePayload.model_validate({}) + with app.test_request_context("/"), patch.object(api, "get_document", return_value=MagicMock()): with pytest.raises(ValueError): - method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") + method(api, req_data, MagicMock(), tenant_id, user, "ds-1", "doc-1") def test_put_invalid_doc_type(self, app: Flask, patch_tenant): api = DocumentMetadataApi() method = unwrap(api.put) user, tenant_id = patch_tenant payload = {"doc_type": "invalid", "doc_metadata": {}} + req_data = DocumentMetadataUpdatePayload.model_validate(payload) with ( - app.test_request_context("/", json=payload), + app.test_request_context("/"), patch.object(api, "get_document", return_value=MagicMock()), patch( "controllers.console.datasets.datasets_document.DocumentService.DOCUMENT_METADATA_SCHEMA", @@ -678,7 +774,7 @@ class TestDocumentMetadataApi: ), ): with pytest.raises(ValueError): - method(api, MagicMock(), tenant_id, user, "ds-1", "doc-1") + method(api, req_data, MagicMock(), tenant_id, user, "ds-1", "doc-1") class TestDocumentStatusApi: @@ -728,73 +824,219 @@ class TestDocumentStatusApi: class TestDocumentRetryApi: - def test_retry_archived_document_skipped(self, app: Flask, patch_tenant, patch_dataset): + def test_retry_archived_document_skipped(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission): api = DocumentRetryApi() method = unwrap(api.post) + user, tenant_id = patch_tenant payload = {"document_ids": ["doc-1"]} - doc = MagicMock(indexing_status="indexing") + doc = MagicMock(id="doc-1", indexing_status="indexing") + session = MagicMock() + session.scalars.return_value.all.return_value = [doc] with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=doc), patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=True), patch("controllers.console.datasets.datasets_document.DocumentService.retry_document") as retry_mock, ): - resp, status = method(api, MagicMock(), "ds-1") + resp, status = method(api, session, tenant_id, user, "ds-1") assert status == 204 - retry_mock.assert_called_once_with("ds-1", [], ANY) + retry_mock.assert_called_once_with("ds-1", [], session) - def test_retry_success(self, app: Flask, patch_tenant, patch_dataset): + def test_retry_success(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission): api = DocumentRetryApi() method = unwrap(api.post) + user, tenant_id = patch_tenant payload = {"document_ids": ["doc-1"]} - document = MagicMock(indexing_status=IndexingStatus.INDEXING, archived=False) + document = MagicMock(id="doc-1", indexing_status=IndexingStatus.INDEXING, archived=False) + session = MagicMock() + session.scalars.return_value.all.return_value = [document] with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document), patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False), patch( "controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None ) as retry_mock, ): - response, status = method(api, MagicMock(), "ds-1") + response, status = method(api, session, tenant_id, user, "ds-1") assert status == 204 - retry_mock.assert_called_once_with("ds-1", [document], ANY) + retry_mock.assert_called_once_with("ds-1", [document], session) - def test_retry_skips_completed_document(self, app: Flask, patch_tenant, patch_dataset): + def test_retry_loads_selected_documents_in_one_scoped_query( + self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission + ): api = DocumentRetryApi() method = unwrap(api.post) - payload = {"document_ids": ["doc-1"]} - document = MagicMock(indexing_status=IndexingStatus.COMPLETED, archived=False) + user, tenant_id = patch_tenant + payload = {"document_ids": ["doc-1", "doc-2"]} + first_document = MagicMock(id="doc-1", indexing_status=IndexingStatus.ERROR, archived=False) + second_document = MagicMock(id="doc-2", indexing_status=IndexingStatus.ERROR, archived=False) + session = MagicMock() + session.scalars.return_value.all.return_value = [first_document, second_document] + with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), - patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document), + patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False), patch( "controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None ) as retry_mock, ): - response, status = method(api, MagicMock(), "ds-1") + response, status = method(api, session, tenant_id, user, "ds-1") + assert status == 204 - retry_mock.assert_called_once_with("ds-1", [], ANY) + statement = session.scalars.call_args.args[0] + assert statement.compare( + select(DatasetDocument).where( + DatasetDocument.tenant_id == "tenant-1", + DatasetDocument.dataset_id == "ds-1", + DatasetDocument.id.in_(["doc-1", "doc-2"]), + ) + ) + retry_mock.assert_called_once_with("ds-1", [first_document, second_document], session) + + def test_retry_skips_completed_document(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission): + api = DocumentRetryApi() + method = unwrap(api.post) + user, tenant_id = patch_tenant + payload = {"document_ids": ["doc-1"]} + document = MagicMock(id="doc-1", indexing_status=IndexingStatus.COMPLETED, archived=False) + session = MagicMock() + session.scalars.return_value.all.return_value = [document] + with ( + app.test_request_context("/", json=payload), + patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None + ) as retry_mock, + ): + response, status = method(api, session, tenant_id, user, "ds-1") + assert status == 204 + retry_mock.assert_called_once_with("ds-1", [], session) + + def test_retry_foreign_dataset_has_no_side_effects(self, app: Flask, patch_tenant, bypass_knowledge_rate_limit): + api = DocumentRetryApi() + method = unwrap(api.post) + user, tenant_id = patch_tenant + session = MagicMock() + payload = {"document_ids": ["doc-1"]} + + with ( + app.test_request_context("/", json=payload), + patch.object(type(console_ns), "payload", payload), + patch( + "controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant", + return_value=None, + ), + patch("controllers.console.datasets.datasets_document.DocumentService.retry_document") as retry_document, + ): + with pytest.raises(NotFound): + method(api, session, tenant_id, user, "foreign-dataset") + + session.scalars.assert_not_called() + bypass_knowledge_rate_limit.assert_not_called() + retry_document.assert_not_called() + + +class TestDocumentPauseRecoverApi: + @pytest.mark.parametrize( + ("api_type", "service_method"), + [(DocumentPauseApi, "pause_document"), (DocumentRecoverApi, "recover_document")], + ) + def test_patch_uses_scoped_document( + self, app: Flask, patch_tenant, bypass_knowledge_rate_limit, api_type, service_method + ): + api = api_type() + method = unwrap(api.patch) + user, tenant_id = patch_tenant + session = MagicMock() + document = MagicMock() + + with ( + app.test_request_context("/"), + patch.object(api, "get_document", return_value=document) as get_document, + patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False), + patch( + f"controllers.console.datasets.datasets_document.DocumentService.{service_method}" + ) as process_document, + ): + response, status = method(api, session, tenant_id, user, "ds-1", "doc-1") + + assert (response, status) == ("", 204) + get_document.assert_called_once_with(session, "ds-1", "doc-1", user, tenant_id) + bypass_knowledge_rate_limit.assert_called_once_with() + process_document.assert_called_once_with(document, session) + + +class TestWebsiteDocumentSyncApi: + def test_get_uses_scoped_dataset_and_document(self, app: Flask, patch_tenant, dataset): + api = WebsiteDocumentSyncApi() + method = unwrap(api.get) + user, tenant_id = patch_tenant + session = MagicMock() + document = MagicMock(data_source_type="website_crawl") + + with ( + app.test_request_context("/"), + patch( + "controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant", + return_value=dataset, + ) as get_dataset, + patch.object(api, "get_document", return_value=document) as get_document, + patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False), + patch( + "controllers.console.datasets.datasets_document.DocumentService.sync_website_document" + ) as sync_document, + ): + response, status = method(api, session, tenant_id, user, "ds-1", "doc-1") + + assert status == 200 + assert response["result"] == "success" + get_dataset.assert_called_once_with("ds-1", tenant_id, session=session) + get_document.assert_called_once_with(session, dataset.id, "doc-1", user, tenant_id) + sync_document.assert_called_once_with(dataset, document, session) + + def test_get_rejects_non_editor_before_loading_document(self, app: Flask, dataset): + api = WebsiteDocumentSyncApi() + method = unwrap(api.get) + user = MagicMock(is_dataset_editor=False) + session = MagicMock() + + with ( + app.test_request_context("/"), + patch( + "controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant", + return_value=dataset, + ), + patch.object(api, "get_document") as get_document, + patch( + "controllers.console.datasets.datasets_document.DocumentService.sync_website_document" + ) as sync_document, + ): + with pytest.raises(Forbidden): + method(api, session, "tenant-1", user, "ds-1", "doc-1") + + get_document.assert_not_called() + sync_document.assert_not_called() class TestDocumentPipelineExecutionLogApi: - def test_get_log_success(self, app: Flask, patch_tenant, patch_dataset): + def test_get_log_success(self, app: Flask, patch_tenant): api = DocumentPipelineExecutionLogApi() method = unwrap(api.get) + user, tenant_id = patch_tenant log = MagicMock(datasource_info="{}", datasource_type="file", input_data={}, datasource_node_id="n1") + document = MagicMock(id="trusted-doc") session = MagicMock() session.scalar.return_value = log with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=MagicMock() - ), + patch.object(api, "get_document", return_value=document) as get_document, ): - response, status = method(api, session, "ds-1", "doc-1") + response, status = method(api, session, tenant_id, user, "ds-1", "doc-1") assert status == 200 + get_document.assert_called_once_with(session, "ds-1", "doc-1", user, tenant_id) + assert "trusted-doc" in session.scalar.call_args.args[0].compile().params.values() class TestDocumentGenerateSummaryApi: @@ -952,49 +1194,54 @@ class TestDocumentBatchDownloadZipApi: class TestDatasetDocumentListApiDelete: - def test_delete_success(self, app: Flask, patch_tenant, patch_dataset): + def test_delete_success(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission): """Test successful deletion of documents""" api = DatasetDocumentListApi() method = unwrap(api.delete) + user, tenant_id = patch_tenant + session = MagicMock() with ( app.test_request_context("/?document_id=doc-1&document_id=doc-2"), - patch( - "controllers.console.datasets.datasets_document.DatasetService.check_dataset_model_setting", - return_value=None, - ), patch("controllers.console.datasets.datasets_document.DocumentService.delete_documents", return_value=None), ): - response, status = method(api, MagicMock(), "ds-1") + response, status = method(api, session, tenant_id, user, "ds-1") assert status == 204 - def test_delete_indexing_error(self, app: Flask, patch_tenant, patch_dataset): + def test_delete_indexing_error(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission): """Test deletion with indexing error""" api = DatasetDocumentListApi() method = unwrap(api.delete) + user, tenant_id = patch_tenant with ( app.test_request_context("/?document_id=doc-1"), - patch( - "controllers.console.datasets.datasets_document.DatasetService.check_dataset_model_setting", - return_value=None, - ), patch( "controllers.console.datasets.datasets_document.DocumentService.delete_documents", side_effect=services.errors.document.DocumentIndexingError(), ), ): with pytest.raises(DocumentIndexingError): - method(api, MagicMock(), "ds-1") + method(api, MagicMock(), tenant_id, user, "ds-1") - def test_delete_dataset_not_found(self, app: Flask, patch_tenant): + def test_delete_dataset_not_found(self, app: Flask, patch_tenant, bypass_knowledge_rate_limit): """Test deletion when dataset not found""" api = DatasetDocumentListApi() method = unwrap(api.delete) + user, tenant_id = patch_tenant with ( app.test_request_context("/?document_id=doc-1"), - patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=None), + patch( + "controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant", + return_value=None, + ), + patch( + "controllers.console.datasets.datasets_document.DocumentService.delete_documents" + ) as delete_documents, ): with pytest.raises(NotFound): - method(api, MagicMock(), "ds-1") + method(api, MagicMock(), tenant_id, user, "foreign-dataset") + + bypass_knowledge_rate_limit.assert_not_called() + delete_documents.assert_not_called() class TestDocumentBatchIndexingEstimateApi: @@ -1080,9 +1327,10 @@ class TestDocumentBatchIndexingStatusApi: api = DocumentBatchIndexingStatusApi() method = unwrap(api.get) user, _ = patch_tenant + error = format_vector_space_admission_error(61, 50) document = MagicMock( id="doc-1", - indexing_status=IndexingStatus.COMPLETED, + indexing_status=IndexingStatus.ERROR, is_paused=False, processing_started_at=None, parsing_completed_at=None, @@ -1090,7 +1338,7 @@ class TestDocumentBatchIndexingStatusApi: splitting_completed_at=None, completed_at=None, paused_at=None, - error=None, + error=error, stopped_at=None, ) session = MagicMock() @@ -1101,14 +1349,17 @@ class TestDocumentBatchIndexingStatusApi: "data": [ { "id": "doc-1", - "indexing_status": "completed", + "indexing_status": "error", "processing_started_at": None, "parsing_completed_at": None, "cleaning_completed_at": None, "splitting_completed_at": None, "completed_at": None, "paused_at": None, - "error": None, + "error": error, + "error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE, + "estimated_vector_space_mb": 61, + "vector_space_limit_mb": 50, "stopped_at": None, "completed_segments": 2, "total_segments": 3, @@ -1291,22 +1542,6 @@ class TestDocumentPermissionCases: assert status == 200 assert response == {"tokens": 0, "total_price": 0, "currency": "USD", "total_segments": 0, "preview": []} - def test_document_tenant_mismatch(self, app: Flask): - api = DocumentApi() - method = unwrap(api.get) - user = MagicMock(is_dataset_editor=True) - document = MagicMock(tenant_id="other-tenant", dataset_process_rule=None) - with ( - app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=MagicMock() - ), - patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document), - patch("controllers.console.datasets.datasets_document.DatasetService.get_process_rules", return_value={}), - ): - with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") - def test_process_rule_get_by_document_success(self, app: Flask, patch_tenant): api = GetProcessRuleApi() method = unwrap(api.get) @@ -1386,15 +1621,16 @@ class TestDocumentListAdvancedCases: payload = {"doc_type": "contract", "doc_metadata": {"amount": 5000, "currency": "USD", "invalid_field": "x"}} schema = {"amount": int, "currency": str} session = MagicMock() + req_data = DocumentMetadataUpdatePayload.model_validate(payload) with ( - app.test_request_context("/", json=payload), + app.test_request_context("/"), patch.object(api, "get_document", return_value=doc), patch( "controllers.console.datasets.datasets_document.DocumentService.DOCUMENT_METADATA_SCHEMA", {"contract": schema}, ), ): - response, status = method(api, session, tenant_id, user, "ds-1", "doc-1") + response, status = method(api, req_data, session, tenant_id, user, "ds-1", "doc-1") assert status == 200 assert doc.doc_metadata == {"amount": 5000, "currency": "USD"} diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py index fa6d73078b7..4149a97eaba 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_document_download.py @@ -18,7 +18,7 @@ from zipfile import ZipFile import pytest from flask import Flask -from werkzeug.exceptions import Forbidden, NotFound +from werkzeug.exceptions import NotFound @pytest.fixture @@ -108,7 +108,11 @@ def _wire_common_success_mocks( import services.dataset_service as dataset_service_module # Return a dataset object and allow permission checks to pass. - monkeypatch.setattr(module.DatasetService, "get_dataset", lambda *_args, **_kwargs: SimpleNamespace(id="ds-1")) + monkeypatch.setattr( + module.DatasetService, + "get_dataset_for_tenant", + lambda *_args, **_kwargs: SimpleNamespace(id="ds-1", tenant_id="tenant-123"), + ) monkeypatch.setattr(module.DatasetService, "check_dataset_permission", lambda *_args, **_kwargs: None) # Return a document that will be validated inside DocumentResource.get_document. @@ -118,7 +122,11 @@ def _wire_common_success_mocks( data_source_type=data_source_type, upload_file_id=upload_file_id, ) - monkeypatch.setattr(module.DocumentService, "get_document", lambda *_args, **_kwargs: document) + monkeypatch.setattr( + module.DatasetRefService, + "get_document_by_ref", + lambda *_args, **_kwargs: document if document.tenant_id == "tenant-123" else None, + ) # Mock UploadFile lookup via FileService batch helper. upload_files_by_id: dict[str, object] = {} @@ -404,10 +412,10 @@ def test_document_download_rejects_when_upload_file_record_missing( method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") -def test_document_download_rejects_tenant_mismatch( +def test_document_download_rejects_document_owner_mismatch( app: Flask, datasets_document_module, monkeypatch: pytest.MonkeyPatch ) -> None: - """Ensure tenant mismatch is rejected by the shared `get_document()` permission check.""" + """Ensure an owner mismatch is rejected by the shared document resolver.""" _wire_common_success_mocks( module=datasets_document_module, @@ -422,5 +430,5 @@ def test_document_download_rejects_tenant_mismatch( with app.test_request_context("/datasets/ds-1/documents/doc-1/download", method="GET"): api = datasets_document_module.DocumentDownloadApi() method = unwrap(api.get) - with pytest.raises(Forbidden): + with pytest.raises(NotFound): method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1") diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py index 73c028da03b..bbcda2fb372 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py @@ -1,3 +1,4 @@ +from datetime import datetime from inspect import unwrap from types import SimpleNamespace from typing import Any, cast @@ -8,9 +9,11 @@ from flask import Flask from werkzeug.exceptions import Forbidden, NotFound import services +from controllers.common.controller_schemas import ChildChunkCreatePayload, ChildChunkUpdatePayload from controllers.console import console_ns from controllers.console.app.error import ProviderNotInitializeError from controllers.console.datasets.datasets_segments import ( + BatchImportPayload, ChildChunkAddApi, ChildChunkBatchUpdatePayload, ChildChunkUpdateApi, @@ -19,6 +22,8 @@ from controllers.console.datasets.datasets_segments import ( DatasetDocumentSegmentBatchImportApi, DatasetDocumentSegmentListApi, DatasetDocumentSegmentUpdateApi, + SegmentCreatePayload, + SegmentUpdatePayload, ) from controllers.console.datasets.error import ChildChunkDeleteIndexError, ChildChunkIndexingError, InvalidActionError from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError @@ -42,16 +47,17 @@ def _segment(): word_count=1, tokens=1, created_by="u1", + answer="a", + keywords=["test"], + index_node_id="n1", + index_node_hash="h", + status=SegmentStatus.COMPLETED, + updated_by="u1", ) + segment.id = "seg-1" - segment.answer = "a" - segment.keywords = ["test"] - segment.index_node_id = "n1" - segment.index_node_hash = "h" - segment.status = SegmentStatus.COMPLETED segment.created_at = naive_utc_now() segment.updated_at = naive_utc_now() - segment.updated_by = "u1" return segment @@ -65,9 +71,9 @@ def _child_chunk(): content="child", word_count=1, created_by="u1", + type=SegmentType.CUSTOMIZED, ) child_chunk.id = "cc-1" - child_chunk.type = SegmentType.CUSTOMIZED child_chunk.created_at = naive_utc_now() child_chunk.updated_at = naive_utc_now() return child_chunk @@ -353,7 +359,9 @@ class TestDatasetDocumentSegmentAddApi: return_value=None, ), ): - response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") + response, status = method( + api, SegmentCreatePayload(content="test content"), session, "tenant-1", user, "ds-1", "doc-1" + ) assert status == 200 assert response["data"]["id"] == "seg-1" @@ -377,7 +385,9 @@ class TestDatasetDocumentSegmentAddApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method( + api, SegmentCreatePayload(content="test content"), MagicMock(), "tenant-1", user, "ds-1", "doc-1" + ) def test_post_provider_token_not_init(self, app: Flask): api = DatasetDocumentSegmentAddApi() @@ -399,7 +409,9 @@ class TestDatasetDocumentSegmentAddApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method( + api, SegmentCreatePayload(content="test content"), MagicMock(), "tenant-1", user, "ds-1", "doc-1" + ) class TestDatasetDocumentSegmentUpdateApi: @@ -440,7 +452,9 @@ class TestDatasetDocumentSegmentUpdateApi: return_value=None, ), ): - response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1", "seg-1") + response, status = method( + api, SegmentUpdatePayload(content="test content"), session, "tenant-1", user, "ds-1", "doc-1", "seg-1" + ) assert status == 200 assert "data" in response @@ -466,7 +480,16 @@ class TestDatasetDocumentSegmentUpdateApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") + method( + api, + SegmentUpdatePayload(content="test content"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + "seg-1", + ) def test_patch_segment_not_found(self, app: Flask): api = DatasetDocumentSegmentUpdateApi() @@ -494,7 +517,16 @@ class TestDatasetDocumentSegmentUpdateApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") + method( + api, + SegmentUpdatePayload(content="test content"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + "seg-1", + ) def test_patch_llm_bad_request(self, app: Flask): api = DatasetDocumentSegmentUpdateApi() @@ -524,7 +556,16 @@ class TestDatasetDocumentSegmentUpdateApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") + method( + api, + SegmentUpdatePayload(content="test content"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + "seg-1", + ) class TestDatasetDocumentSegmentBatchImportApi: @@ -532,8 +573,19 @@ class TestDatasetDocumentSegmentBatchImportApi: api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - upload_file = MagicMock(spec=UploadFile) - upload_file.name = "test.csv" + upload_file = UploadFile( + tenant_id="tenant-id", + storage_type="opendal", + key="test-key", + name="test.csv", + size=0, + extension="txt", + mime_type="text/plain", + created_by_role="account", + created_by="account-id", + created_at=datetime.now(), + used=False, + ) user = MagicMock(id="u1") session = MagicMock() session.scalar.return_value = upload_file @@ -552,7 +604,9 @@ class TestDatasetDocumentSegmentBatchImportApi: return_value=None, ), ): - response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") + response, status = method( + api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + ) assert status == 200 assert response["job_status"] == "waiting" @@ -569,7 +623,15 @@ class TestDatasetDocumentSegmentBatchImportApi: patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method( + api, + BatchImportPayload(upload_file_id="test-file-id"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + ) def test_post_document_not_found(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() @@ -587,7 +649,15 @@ class TestDatasetDocumentSegmentBatchImportApi: patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=None), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method( + api, + BatchImportPayload(upload_file_id="test-file-id"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + ) def test_post_upload_file_not_found(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() @@ -607,7 +677,9 @@ class TestDatasetDocumentSegmentBatchImportApi: ), ): with pytest.raises(NotFound): - method(api, session, "tenant-1", user, "ds-1", "doc-1") + method( + api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + ) def test_post_invalid_file_type(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() @@ -629,7 +701,9 @@ class TestDatasetDocumentSegmentBatchImportApi: ), ): with pytest.raises(ValueError): - method(api, session, "tenant-1", user, "ds-1", "doc-1") + method( + api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + ) def test_post_async_task_failure(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() @@ -653,7 +727,9 @@ class TestDatasetDocumentSegmentBatchImportApi: "controllers.console.datasets.datasets_segments.redis_client.setnx", side_effect=Exception("redis down") ), ): - response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") + response, status = method( + api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + ) assert status == 500 assert "error" in response @@ -735,7 +811,9 @@ class TestChildChunkAddApi: return_value=child_chunk, ), ): - response, status = method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") + response, status = method( + api, ChildChunkCreatePayload(content="child"), MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1" + ) assert status == 200 assert response["data"]["id"] == "cc-1" @@ -766,7 +844,16 @@ class TestChildChunkAddApi: ), ): with pytest.raises(ChildChunkIndexingError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") + method( + api, + ChildChunkCreatePayload(content="child"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + "seg-1", + ) def test_post_permission_denied(self, app: Flask): api = ChildChunkAddApi() @@ -787,7 +874,16 @@ class TestChildChunkAddApi: ), ): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1") + method( + api, + ChildChunkCreatePayload(content="child"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + "seg-1", + ) class TestChildChunkUpdateApi: @@ -910,7 +1006,17 @@ class TestChildChunkUpdateApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + method( + api, + ChildChunkUpdatePayload(content="updated child"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + "seg-1", + "cc-1", + ) class TestSegmentListAdvancedCases: @@ -1035,7 +1141,9 @@ class TestSegmentOperationCases: ), ): with pytest.raises(ProviderTokenNotInitError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method( + api, SegmentCreatePayload(content="test content"), MagicMock(), "tenant-1", user, "ds-1", "doc-1" + ) def test_batch_import_with_document_not_found(self, app: Flask): """Test batch import with document not found""" @@ -1051,7 +1159,15 @@ class TestSegmentOperationCases: patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=None), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method( + api, + BatchImportPayload(upload_file_id="test-file-id"), + MagicMock(), + "tenant-1", + user, + "ds-1", + "doc-1", + ) def test_batch_import_with_invalid_file(self, app: Flask): """Test batch import with invalid file type""" @@ -1071,7 +1187,9 @@ class TestSegmentOperationCases: patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), ): with pytest.raises(NotFound): - method(api, session, "tenant-1", user, "ds-1", "doc-1") + method( + api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + ) def test_batch_import_with_async_task_failure(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() @@ -1079,8 +1197,20 @@ class TestSegmentOperationCases: user = MagicMock(is_dataset_editor=True) dataset = MagicMock() document = MagicMock() - upload_file = MagicMock(spec=UploadFile, extension="csv", id="file-1") - upload_file.name = "test.csv" + upload_file = UploadFile( + tenant_id="tenant-id", + storage_type="opendal", + key="test-key", + name="test.csv", + size=0, + extension="csv", + mime_type="text/csv", + created_by_role="account", + created_by="account-id", + created_at=datetime.now(), + used=False, + ) + upload_file.id = "file-1" payload = {"upload_file_id": "file-1"} session = MagicMock() session.scalar.return_value = upload_file @@ -1098,7 +1228,9 @@ class TestSegmentOperationCases: side_effect=Exception("Task failed"), ), ): - response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") + response, status = method( + api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + ) assert status == 500 assert "error" in response diff --git a/api/tests/unit_tests/controllers/console/datasets/test_external.py b/api/tests/unit_tests/controllers/console/datasets/test_external.py index 5addfd66e7f..c08b68c0fbc 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_external.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_external.py @@ -5,6 +5,7 @@ from unittest.mock import ANY, MagicMock, PropertyMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services @@ -12,14 +13,19 @@ from controllers.console import console_ns from controllers.console.datasets.error import DatasetNameDuplicateError from controllers.console.datasets.external import ( BedrockRetrievalApi, + BedrockRetrievalPayload, ExternalApiTemplateApi, ExternalApiTemplateListApi, + ExternalApiTemplateListQuery, ExternalApiUseCheckApi, ExternalDatasetCreateApi, + ExternalHitTestingPayload, + ExternalKnowledgeApiPayload, ExternalKnowledgeHitTestingApi, ) from models.account import Account, TenantAccountRole from services.dataset_service import DatasetService +from services.entities.external_knowledge_entities.external_knowledge_entities import ExternalDatasetCreatePayload from services.external_knowledge_service import ExternalDatasetService from services.hit_testing_service import HitTestingService from services.knowledge_service import ExternalDatasetTestService @@ -157,13 +163,21 @@ def _dataset_detail_object() -> SimpleNamespace: ) -class TestExternalApiTemplateListApi: +class _UsesSQLiteSession: + session: Session + + @pytest.fixture(autouse=True) + def _inject_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + + +class TestExternalApiTemplateListApi(_UsesSQLiteSession): def test_get_success(self, app: Flask): api = ExternalApiTemplateListApi() method = inspect.unwrap(api.get) api_item = _external_api_object("api-1") - session = MagicMock() + session = self.session with ( app.test_request_context("/?page=2&limit=1&keyword=vector"), @@ -173,7 +187,9 @@ class TestExternalApiTemplateListApi: return_value=([api_item], 3), ) as get_external_knowledge_apis, ): - resp, status = method(api, session, "tenant-1") + resp, status = method( + api, ExternalApiTemplateListQuery(page=2, limit=1, keyword="vector"), session, "tenant-1" + ) assert status == 200 assert resp == { @@ -200,7 +216,7 @@ class TestExternalApiTemplateListApi: }, } created = _external_api_object("api-created") - session = MagicMock() + session = self.session with ( app.test_request_context("/", json=payload), @@ -212,7 +228,9 @@ class TestExternalApiTemplateListApi: return_value=created, ) as create_external_knowledge_api, ): - resp, status = method(api, session, "tenant-1", current_user) + resp, status = method( + api, ExternalKnowledgeApiPayload.model_validate(payload), session, "tenant-1", current_user + ) assert status == 201 assert resp == _external_api_dict("api-created") @@ -238,7 +256,7 @@ class TestExternalApiTemplateListApi: patch.object(ExternalDatasetService, "validate_api_list"), ): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", current_user) + method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user) def test_post_duplicate_name(self, app: Flask, current_user: Account): api = ExternalApiTemplateListApi() @@ -257,15 +275,15 @@ class TestExternalApiTemplateListApi: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, MagicMock(), "tenant-1", current_user) + method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user) -class TestExternalApiTemplateApi: +class TestExternalApiTemplateApi(_UsesSQLiteSession): def test_get_success_returns_template_contract(self, app: Flask): api = ExternalApiTemplateApi() method = inspect.unwrap(api.get) template = _external_api_object("api-detail") - session = MagicMock() + session = self.session with ( app.test_request_context("/"), @@ -297,7 +315,7 @@ class TestExternalApiTemplateApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", "api-id") + method(api, self.session, "tenant-1", "api-id") def test_patch_success_uses_validated_payload_and_returns_template(self, app: Flask, current_user: Account): api = ExternalApiTemplateApi() @@ -312,7 +330,7 @@ class TestExternalApiTemplateApi: }, } updated = _external_api_object("api-updated") - session = MagicMock() + session = self.session with ( app.test_request_context("/", json=payload), @@ -324,7 +342,14 @@ class TestExternalApiTemplateApi: return_value=updated, ) as update_external_knowledge_api, ): - resp, status = method(api, session, "tenant-1", current_user, "api-updated") + resp, status = method( + api, + ExternalKnowledgeApiPayload.model_validate(payload), + session, + "tenant-1", + current_user, + "api-updated", + ) assert status == 200 assert resp == _external_api_dict("api-updated") @@ -346,15 +371,15 @@ class TestExternalApiTemplateApi: with app.test_request_context("/"): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", current_user, "api-id") + method(api, self.session, "tenant-1", current_user, "api-id") -class TestExternalApiUseCheckApi: +class TestExternalApiUseCheckApi(_UsesSQLiteSession): def test_get_scopes_usage_check_to_current_tenant(self, app: Flask): api = ExternalApiUseCheckApi() method = inspect.unwrap(api.get) - session = MagicMock() + session = self.session with ( app.test_request_context("/"), @@ -371,7 +396,7 @@ class TestExternalApiUseCheckApi: mock_use_check.assert_called_once_with("api-id", "tenant-1", session=ANY) -class TestExternalDatasetCreateApi: +class TestExternalDatasetCreateApi(_UsesSQLiteSession): def test_create_success(self, app: Flask, current_user: Account): api = ExternalDatasetCreateApi() method = inspect.unwrap(api.post) @@ -404,14 +429,16 @@ class TestExternalDatasetCreateApi: ) as dataset_response_source, ): session = MagicMock() - resp, status = method(api, session, "tenant-1", current_user) + resp, status = method( + api, ExternalDatasetCreatePayload.model_validate(payload), session, "tenant-1", current_user + ) assert status == 201 assert resp == _expected_dataset_detail_payload() create_external_dataset.assert_called_once_with( tenant_id="tenant-1", user_id="user-1", - args=payload, + args=ExternalDatasetCreatePayload.model_validate(payload), session=session, ) dataset_response_source.assert_called_once_with(dataset, session=session) @@ -432,10 +459,12 @@ class TestExternalDatasetCreateApi: patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), ): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", current_user) + method( + api, ExternalDatasetCreatePayload.model_validate(payload), self.session, "tenant-1", current_user + ) -class TestExternalKnowledgeHitTestingApi: +class TestExternalKnowledgeHitTestingApi(_UsesSQLiteSession): def test_hit_testing_dataset_not_found(self, app: Flask, current_user: Account): api = ExternalKnowledgeHitTestingApi() method = inspect.unwrap(api.post) @@ -449,7 +478,7 @@ class TestExternalKnowledgeHitTestingApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), current_user, "dataset-id") + method(api, ExternalHitTestingPayload(query="test"), self.session, current_user, "dataset-id") def test_hit_testing_success(self, app: Flask, current_user: Account): api = ExternalKnowledgeHitTestingApi() @@ -480,7 +509,7 @@ class TestExternalKnowledgeHitTestingApi: } ], } - session = MagicMock() + session = self.session with ( app.test_request_context("/", json=payload), @@ -495,7 +524,7 @@ class TestExternalKnowledgeHitTestingApi: ) as external_retrieve, patch("controllers.console.datasets.external.dump_response", side_effect=lambda _model, value: value), ): - resp = method(api, session, current_user, "dataset-id") + resp = method(api, ExternalHitTestingPayload.model_validate(payload), session, current_user, "dataset-id") assert resp == retrieve_response check_dataset_permission.assert_called_once_with(dataset, current_user, session) @@ -537,8 +566,6 @@ class TestBedrockRetrievalApi: ] } - session = MagicMock() - with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), @@ -548,7 +575,7 @@ class TestBedrockRetrievalApi: return_value=retrieval_response, ) as knowledge_retrieval, ): - resp, status = method() + resp, status = method(api, BedrockRetrievalPayload.model_validate(payload)) assert status == 200 assert resp == retrieval_response @@ -558,7 +585,7 @@ class TestBedrockRetrievalApi: assert knowledge_id == "knowledge-base-1" -class TestExternalApiTemplateListApiAdvanced: +class TestExternalApiTemplateListApiAdvanced(_UsesSQLiteSession): def test_post_duplicate_name_error(self, app: Flask, current_user: Account): api = ExternalApiTemplateListApi() method = inspect.unwrap(api.post) @@ -575,7 +602,7 @@ class TestExternalApiTemplateListApiAdvanced: ), ): with pytest.raises(DatasetNameDuplicateError): - method(api, MagicMock(), "tenant-1", current_user) + method(api, ExternalKnowledgeApiPayload.model_validate(payload), self.session, "tenant-1", current_user) def test_get_with_pagination(self, app: Flask): api = ExternalApiTemplateListApi() @@ -590,7 +617,7 @@ class TestExternalApiTemplateListApiAdvanced: return_value=(templates, 25), ) as get_external_knowledge_apis, ): - resp, status = method(api, MagicMock(), "tenant-1") + resp, status = method(api, ExternalApiTemplateListQuery(page=2, limit=3), self.session, "tenant-1") assert status == 200 assert resp == { @@ -603,7 +630,7 @@ class TestExternalApiTemplateListApiAdvanced: get_external_knowledge_apis.assert_called_once_with(2, 3, "tenant-1", None, session=ANY) -class TestExternalDatasetCreateApiAdvanced: +class TestExternalDatasetCreateApiAdvanced(_UsesSQLiteSession): def test_create_forbidden(self, app: Flask, current_user: Account): """Test creating external dataset without permission""" api = ExternalDatasetCreateApi() @@ -620,10 +647,12 @@ class TestExternalDatasetCreateApiAdvanced: with app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", current_user) + method( + api, ExternalDatasetCreatePayload.model_validate(payload), self.session, "tenant-1", current_user + ) -class TestExternalKnowledgeHitTestingApiAdvanced: +class TestExternalKnowledgeHitTestingApiAdvanced(_UsesSQLiteSession): def test_hit_testing_dataset_not_found(self, app: Flask, current_user: Account): """Test hit testing on non-existent dataset""" api = ExternalKnowledgeHitTestingApi() @@ -643,7 +672,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), current_user, "ds-1") + method(api, ExternalHitTestingPayload.model_validate(payload), self.session, current_user, "ds-1") def test_hit_testing_with_custom_retrieval_model(self, app: Flask, current_user: Account): api = ExternalKnowledgeHitTestingApi() @@ -655,7 +684,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced: "external_retrieval_model": {"type": "bm25"}, "metadata_filtering_conditions": {"status": "active"}, } - session = MagicMock() + session = self.session with ( app.test_request_context("/", json=payload), @@ -681,7 +710,7 @@ class TestExternalKnowledgeHitTestingApiAdvanced: }, ) as external_retrieve, ): - resp = method(api, session, current_user, "ds-1") + resp = method(api, ExternalHitTestingPayload.model_validate(payload), session, current_user, "ds-1") assert resp == { "query": {"content": "test query"}, @@ -726,4 +755,4 @@ class TestBedrockRetrievalApiAdvanced: ), ): with pytest.raises(ValueError): - method() + method(api, BedrockRetrievalPayload.model_validate(payload)) diff --git a/api/tests/unit_tests/controllers/console/datasets/test_external_dataset_payload.py b/api/tests/unit_tests/controllers/console/datasets/test_external_dataset_payload.py index 2ea1fcf5441..49c89b1adfc 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_external_dataset_payload.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_external_dataset_payload.py @@ -11,7 +11,7 @@ full Flask/RESTX request stack. import pytest from pydantic import ValidationError -from controllers.console.datasets.external import ExternalDatasetCreatePayload +from services.entities.external_knowledge_entities.external_knowledge_entities import ExternalDatasetCreatePayload def test_external_dataset_create_payload_allows_name_length_100() -> None: diff --git a/api/tests/unit_tests/controllers/console/datasets/test_metadata.py b/api/tests/unit_tests/controllers/console/datasets/test_metadata.py index 00f45d488a4..6cf5664bd96 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_metadata.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_metadata.py @@ -1,12 +1,13 @@ import uuid from inspect import unwrap -from unittest.mock import MagicMock, PropertyMock, patch +from unittest.mock import ANY, MagicMock, PropertyMock, patch import pytest from flask import Flask from pytest_mock import MockerFixture -from werkzeug.exceptions import NotFound +from werkzeug.exceptions import Forbidden, NotFound +from controllers.common.controller_schemas import MetadataUpdatePayload from controllers.console import console_ns from controllers.console.datasets.metadata import ( DatasetMetadataApi, @@ -18,12 +19,15 @@ from controllers.console.datasets.metadata import ( from models.account import Account from services.dataset_service import DatasetService from services.entities.knowledge_entities.knowledge_entities import MetadataArgs, MetadataOperationData +from services.errors.account import NoPermissionError +from services.errors.metadata import MetadataResourceNotFoundError from services.metadata_service import MetadataService @pytest.fixture def app(): app = Flask("test_dataset_metadata") + app.config["TESTING"] = True return app @@ -76,7 +80,9 @@ class TestDatasetMetadataCreateApi: MetadataService, "create_metadata", return_value={"id": "m1", "type": "string", "name": "author"} ), ): - result, status = method(api, MagicMock(), "tenant-1", current_user, dataset_id) + result, status = method( + api, MetadataArgs(type="string", name="author"), MagicMock(), "tenant-1", current_user, dataset_id + ) assert status == 201 assert result["type"] == "string" assert result["name"] == "author" @@ -92,16 +98,19 @@ class TestDatasetMetadataCreateApi: patch.object(DatasetService, "get_dataset", return_value=None), ): with pytest.raises(NotFound, match="Dataset not found"): - method(api, MagicMock(), "tenant-1", current_user, dataset_id) + method( + api, MetadataArgs(type="string", name="author"), MagicMock(), "tenant-1", current_user, dataset_id + ) class TestDatasetMetadataGetApi: - def test_get_metadata_success(self, app: Flask, dataset, dataset_id): + def test_get_metadata_success(self, app: Flask, current_user, dataset, dataset_id): api = DatasetMetadataCreateApi() method = unwrap(api.get) with ( app.test_request_context("/"), - patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset, + patch.object(DatasetService, "check_dataset_permission") as check_permission, patch.object( MetadataService, "get_dataset_metadatas", @@ -111,17 +120,63 @@ class TestDatasetMetadataGetApi: }, ), ): - result, status = method(api, MagicMock(), dataset_id) + session = MagicMock() + result, status = method(api, session, "tenant-1", current_user, dataset_id) assert status == 200 assert result["doc_metadata"] == [{"id": "m1", "name": "author", "type": "string", "count": 0}] assert result["built_in_field_enabled"] is False + get_dataset.assert_called_once_with(str(dataset_id), "tenant-1", session=session) + check_permission.assert_called_once_with(dataset, current_user, session) - def test_get_metadata_dataset_not_found(self, app: Flask, dataset_id): + def test_get_metadata_rejects_foreign_tenant_before_read(self, app: Flask, current_user, dataset_id): api = DatasetMetadataCreateApi() method = unwrap(api.get) - with app.test_request_context("/"), patch.object(DatasetService, "get_dataset", return_value=None): + session = MagicMock() + with ( + app.test_request_context("/"), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=None) as get_dataset, + patch.object(DatasetService, "check_dataset_permission") as check_permission, + patch.object(MetadataService, "get_dataset_metadatas") as get_metadata, + ): with pytest.raises(NotFound): - method(api, MagicMock(), dataset_id) + method(api, session, "tenant-1", current_user, dataset_id) + + get_dataset.assert_called_once_with(str(dataset_id), "tenant-1", session=session) + check_permission.assert_not_called() + get_metadata.assert_not_called() + + def test_get_metadata_relies_on_rbac_in_rbac_mode(self, app: Flask, current_user, dataset, dataset_id): + api = DatasetMetadataCreateApi() + method = unwrap(api.get) + with ( + app.test_request_context("/"), + patch("controllers.console.datasets.metadata.dify_config.RBAC_ENABLED", True), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission") as check_permission, + patch.object( + MetadataService, + "get_dataset_metadatas", + return_value={"doc_metadata": [], "built_in_field_enabled": False}, + ), + ): + _, status = method(api, MagicMock(), "tenant-1", current_user, dataset_id) + + assert status == 200 + check_permission.assert_not_called() + + def test_get_metadata_rejects_inaccessible_dataset(self, app: Flask, current_user, dataset, dataset_id): + api = DatasetMetadataCreateApi() + method = unwrap(api.get) + with ( + app.test_request_context("/"), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission", side_effect=NoPermissionError), + patch.object(MetadataService, "get_dataset_metadatas") as get_metadata, + ): + with pytest.raises(Forbidden): + method(api, MagicMock(), "tenant-1", current_user, dataset_id) + + get_metadata.assert_not_called() class TestDatasetMetadataApi: @@ -132,31 +187,43 @@ class TestDatasetMetadataApi: with ( app.test_request_context("/"), patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), - patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset, patch.object(DatasetService, "check_dataset_permission"), patch.object( MetadataService, "update_metadata_name", return_value={"id": "m1", "type": "string", "name": "updated-name"}, - ), + ) as update_metadata, ): - result, status = method(api, MagicMock(), "tenant-1", current_user, dataset_id, metadata_id) + result, status = method( + api, + MetadataUpdatePayload(name="updated-name"), + MagicMock(), + "tenant-1", + current_user, + dataset_id, + metadata_id, + ) assert status == 200 assert result["type"] == "string" assert result["name"] == "updated-name" + get_dataset.assert_called_once_with(str(dataset_id), "tenant-1", session=ANY) + update_metadata.assert_called_once_with(dataset, str(metadata_id), "updated-name", current_user, session=ANY) def test_delete_metadata_success(self, app: Flask, current_user, dataset, dataset_id, metadata_id): api = DatasetMetadataApi() method = unwrap(api.delete) with ( app.test_request_context("/"), - patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset, patch.object(DatasetService, "check_dataset_permission"), - patch.object(MetadataService, "delete_metadata"), + patch.object(MetadataService, "delete_metadata") as delete_metadata, ): - result, status = method(api, MagicMock(), current_user, dataset_id, metadata_id) + result, status = method(api, MagicMock(), "tenant-1", current_user, dataset_id, metadata_id) assert status == 204 assert result == "" + get_dataset.assert_called_once_with(str(dataset_id), "tenant-1", session=ANY) + delete_metadata.assert_called_once_with(dataset, str(metadata_id), ANY) class TestDatasetMetadataBuiltInFieldApi: @@ -199,11 +266,38 @@ class TestDocumentMetadataEditApi: with ( app.test_request_context("/"), patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload), - patch.object(DatasetService, "get_dataset", return_value=dataset), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), patch.object(DatasetService, "check_dataset_permission"), - patch.object(MetadataOperationData, "model_validate", return_value=MagicMock()), patch.object(MetadataService, "update_documents_metadata"), ): - result, status = method(api, MagicMock(), current_user, dataset_id) + result, status = method( + api, + MetadataOperationData( + operation_data=[{"document_id": "00000000-0000-0000-0000-000000000001", "metadata_list": []}] + ), + MagicMock(), + dataset.tenant_id, + current_user, + dataset_id, + ) assert status == 204 assert result == "" + + def test_update_document_metadata_translates_missing_resource(self, app: Flask, current_user, dataset, dataset_id): + api = DocumentMetadataEditApi() + method = unwrap(api.post) + request = MetadataOperationData(operation_data=[]) + with ( + app.test_request_context("/"), + patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset), + patch.object(DatasetService, "check_dataset_permission"), + patch.object( + MetadataService, + "update_documents_metadata", + side_effect=MetadataResourceNotFoundError("Metadata not found."), + ), + pytest.raises(NotFound) as exc_info, + ): + method(api, request, MagicMock(), dataset.tenant_id, current_user, dataset_id) + + assert exc_info.value.description == "Metadata not found." diff --git a/api/tests/unit_tests/controllers/console/datasets/test_website.py b/api/tests/unit_tests/controllers/console/datasets/test_website.py index 5c7b857c20e..b790059f0c7 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_website.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_website.py @@ -1,14 +1,15 @@ -from unittest.mock import Mock, PropertyMock, patch +from unittest.mock import Mock import pytest from flask import Flask from pytest_mock import MockerFixture -from controllers.console import console_ns from controllers.console.datasets.error import WebsiteCrawlError from controllers.console.datasets.website import ( WebsiteCrawlApi, + WebsiteCrawlPayload, WebsiteCrawlStatusApi, + WebsiteCrawlStatusQuery, ) from services.website_service import ( WebsiteCrawlApiRequest, @@ -58,16 +59,9 @@ class TestWebsiteCrawlApi: "url": "https://example.com", "options": {"depth": 1}, } + req_data = WebsiteCrawlPayload.model_validate(payload) - with ( - app.test_request_context("/", json=payload), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - ): + with app.test_request_context("/", json=payload): mock_request = Mock(spec=WebsiteCrawlApiRequest) mocker.patch.object( WebsiteCrawlApiRequest, @@ -81,7 +75,7 @@ class TestWebsiteCrawlApi: return_value={"job_id": "job-1"}, ) - result, status = method(api) + result, status = method(api, req_data) assert status == 200 assert result["job_id"] == "job-1" @@ -95,16 +89,9 @@ class TestWebsiteCrawlApi: "url": "bad-url", "options": {}, } + req_data = WebsiteCrawlPayload.model_validate(payload) - with ( - app.test_request_context("/", json=payload), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - ): + with app.test_request_context("/", json=payload): mocker.patch.object( WebsiteCrawlApiRequest, "from_args", @@ -112,7 +99,7 @@ class TestWebsiteCrawlApi: ) with pytest.raises(WebsiteCrawlError, match="invalid payload"): - method(api) + method(api, req_data) def test_crawl_service_error(self, app: Flask, mocker: MockerFixture): api = WebsiteCrawlApi() @@ -123,16 +110,9 @@ class TestWebsiteCrawlApi: "url": "https://example.com", "options": {}, } + req_data = WebsiteCrawlPayload.model_validate(payload) - with ( - app.test_request_context("/", json=payload), - patch.object( - type(console_ns), - "payload", - new_callable=PropertyMock, - return_value=payload, - ), - ): + with app.test_request_context("/", json=payload): mock_request = Mock(spec=WebsiteCrawlApiRequest) mocker.patch.object( WebsiteCrawlApiRequest, @@ -147,7 +127,7 @@ class TestWebsiteCrawlApi: ) with pytest.raises(WebsiteCrawlError, match="crawl failed"): - method(api) + method(api, req_data) class TestWebsiteCrawlStatusApi: @@ -156,14 +136,9 @@ class TestWebsiteCrawlStatusApi: method = unwrap(api.get) job_id = "job-123" - args = {"provider": "firecrawl"} + req_data = WebsiteCrawlStatusQuery.model_validate({"provider": "firecrawl"}) with app.test_request_context("/?provider=firecrawl"): - mocker.patch( - "controllers.console.datasets.website.request.args.to_dict", - return_value=args, - ) - mock_request = Mock(spec=WebsiteCrawlStatusApiRequest) mocker.patch.object( WebsiteCrawlStatusApiRequest, @@ -177,7 +152,7 @@ class TestWebsiteCrawlStatusApi: return_value={"status": "completed"}, ) - result, status = method(api, job_id) + result, status = method(api, req_data, job_id) assert status == 200 assert result["status"] == "completed" @@ -187,14 +162,9 @@ class TestWebsiteCrawlStatusApi: method = unwrap(api.get) job_id = "job-123" - args = {"provider": "firecrawl"} + req_data = WebsiteCrawlStatusQuery.model_validate({"provider": "firecrawl"}) with app.test_request_context("/?provider=firecrawl"): - mocker.patch( - "controllers.console.datasets.website.request.args.to_dict", - return_value=args, - ) - mocker.patch.object( WebsiteCrawlStatusApiRequest, "from_args", @@ -202,21 +172,16 @@ class TestWebsiteCrawlStatusApi: ) with pytest.raises(WebsiteCrawlError, match="invalid provider"): - method(api, job_id) + method(api, req_data, job_id) def test_get_status_service_error(self, app: Flask, mocker: MockerFixture): api = WebsiteCrawlStatusApi() method = unwrap(api.get) job_id = "job-123" - args = {"provider": "firecrawl"} + req_data = WebsiteCrawlStatusQuery.model_validate({"provider": "firecrawl"}) with app.test_request_context("/?provider=firecrawl"): - mocker.patch( - "controllers.console.datasets.website.request.args.to_dict", - return_value=args, - ) - mock_request = Mock(spec=WebsiteCrawlStatusApiRequest) mocker.patch.object( WebsiteCrawlStatusApiRequest, @@ -231,4 +196,4 @@ class TestWebsiteCrawlStatusApi: ) with pytest.raises(WebsiteCrawlError, match="status lookup failed"): - method(api, job_id) + method(api, req_data, job_id) diff --git a/api/tests/unit_tests/controllers/console/datasets/test_wraps.py b/api/tests/unit_tests/controllers/console/datasets/test_wraps.py index 80ee65a7094..0469aebda7a 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_wraps.py @@ -39,9 +39,11 @@ class TestGetRagPipeline: get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session_factory.return_value) def test_pipeline_found_and_injected(self, mocker: MockerFixture): - pipeline = Mock(spec=Pipeline) + pipeline = Pipeline( + tenant_id="tenant-1", + name="Test Pipeline", + ) pipeline.id = "pipeline-1" - pipeline.tenant_id = "tenant-1" @get_rag_pipeline def dummy_view(**kwargs): @@ -63,9 +65,8 @@ class TestGetRagPipeline: assert result is pipeline get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session_factory.return_value) - def test_load_rag_pipeline_uses_provided_session(self, mocker: MockerFixture): - pipeline = Mock(spec=Pipeline) - session = Mock(spec=Session) + def test_load_rag_pipeline_uses_provided_session(self, mocker: MockerFixture, sqlite_session: Session): + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") mocker.patch( "controllers.console.datasets.wraps.current_account_with_tenant", @@ -76,13 +77,13 @@ class TestGetRagPipeline: return_value=pipeline, ) - result = load_rag_pipeline(session, "pipeline-1") + result = load_rag_pipeline(sqlite_session, "pipeline-1") assert result is pipeline - get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=session) + get_pipeline_by_id.assert_called_once_with("pipeline-1", "tenant-1", session=sqlite_session) def test_pipeline_id_removed_from_kwargs(self, mocker: MockerFixture): - pipeline = Mock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") @get_rag_pipeline def dummy_view(**kwargs): @@ -105,7 +106,7 @@ class TestGetRagPipeline: assert result == "ok" def test_pipeline_id_cast_to_string(self, mocker: MockerFixture): - pipeline = Mock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline") @get_rag_pipeline def dummy_view(**kwargs): diff --git a/api/tests/unit_tests/controllers/console/explore/test_audio.py b/api/tests/unit_tests/controllers/console/explore/test_audio.py index b21e0e4b2c6..704b45698b6 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_audio.py +++ b/api/tests/unit_tests/controllers/console/explore/test_audio.py @@ -273,7 +273,10 @@ class TestChatTextApi: ), patch.object(audio_module.AudioService, "transcript_tts", transcript_tts), ): - resp = self.method(installed_app) + resp = self.method( + audio_module.TextToAudioPayload.model_validate({"message_id": "m1", "text": "hello", "voice": "v1"}), + installed_app, + ) assert resp == {"audio": "ok"} assert transcript_tts.call_args.kwargs["message_ref"] == MessageRef( @@ -295,7 +298,7 @@ class TestChatTextApi: ), ): with pytest.raises(ProviderNotInitializeError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_model_not_supported(self, app: Flask, installed_app): with ( @@ -310,7 +313,7 @@ class TestChatTextApi: ), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_invoke_error(self, app: Flask, installed_app): with ( @@ -325,7 +328,7 @@ class TestChatTextApi: ), ): with pytest.raises(CompletionRequestError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_unknown_exception(self, app: Flask, installed_app): with ( @@ -340,7 +343,7 @@ class TestChatTextApi: ), ): with pytest.raises(InternalServerError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_app_unavailable_tts(self, app: Flask, installed_app): with ( @@ -355,7 +358,7 @@ class TestChatTextApi: ), ): with pytest.raises(AppUnavailableError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_no_audio_uploaded_tts(self, app: Flask, installed_app): with ( @@ -370,7 +373,7 @@ class TestChatTextApi: ), ): with pytest.raises(NoAudioUploadedError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_audio_too_large_tts(self, app: Flask, installed_app): with ( @@ -385,7 +388,7 @@ class TestChatTextApi: ), ): with pytest.raises(AudioTooLargeError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_unsupported_audio_type_tts(self, app: Flask, installed_app): with ( @@ -400,7 +403,7 @@ class TestChatTextApi: ), ): with pytest.raises(audio_module.UnsupportedAudioTypeError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_provider_not_support_speech_to_text_tts(self, app: Flask, installed_app): with ( @@ -415,7 +418,7 @@ class TestChatTextApi: ), ): with pytest.raises(audio_module.ProviderNotSupportSpeechToTextError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) def test_quota_exceeded_tts(self, app: Flask, installed_app): with ( @@ -430,4 +433,4 @@ class TestChatTextApi: ), ): with pytest.raises(ProviderQuotaExceededError): - self.method(installed_app) + self.method(audio_module.TextToAudioPayload.model_validate({"text": "hi"}), installed_app) diff --git a/api/tests/unit_tests/controllers/console/explore/test_banner.py b/api/tests/unit_tests/controllers/console/explore/test_banner.py index 36b83151b72..0260cffde57 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_banner.py +++ b/api/tests/unit_tests/controllers/console/explore/test_banner.py @@ -1,43 +1,61 @@ -from collections.abc import Iterator from datetime import datetime -from inspect import unwrap +from typing import NamedTuple from uuid import uuid4 import pytest from flask import Flask -from sqlalchemy.engine import Engine -from sqlalchemy.orm import Session +from pydantic import ValidationError +from sqlalchemy.orm import Session, sessionmaker import controllers.console.explore.banner as banner_module -from models.base import TypeBase from models.enums import BannerStatus from models.model import ExporleBanner +from repositories.explore_banner_query_repository import ExploreBannerQueryRepository +from services.explore_banner_query_service import ExploreBannerQueryService, ExploreBannerRecord -@pytest.fixture -def banner_session(sqlite_engine: Engine) -> Iterator[Session]: - """Create the banner table without its PostgreSQL-cast server defaults.""" - table = TypeBase.metadata.tables[ExporleBanner.__tablename__] - status_default = table.c.status.server_default - language_default = table.c.language.server_default - table.c.status.server_default = None - table.c.language.server_default = None - try: - TypeBase.metadata.create_all(sqlite_engine, tables=[table]) - finally: - table.c.status.server_default = status_default - table.c.language.server_default = language_default +class FakeExploreBannerQuery: + def __init__(self, responses: dict[str, tuple[ExploreBannerRecord, ...]] | None = None) -> None: + self.responses = responses or {} + self.requested_languages: list[str] = [] - with Session(sqlite_engine, expire_on_commit=False) as session: - yield session + def list_enabled(self, language: str) -> tuple[ExploreBannerRecord, ...]: + self.requested_languages.append(language) + return self.responses.get(language, ()) -def _banner(*, text: str, language: str, link: str, created_at: datetime) -> ExporleBanner: +class _ApplicationServicesStub(NamedTuple): + explore_banner_queries: ExploreBannerQueryService + + +def _content( + title: str, + *, + category: str = "Featured", + description: str = "Banner description", +) -> dict[str, str]: + return { + "category": category, + "title": title, + "description": description, + "img-src": "https://example.com/banner.png", + } + + +def _banner( + *, + title: str, + language: str, + link: str, + created_at: datetime, + sort: int = 1, + status: BannerStatus = BannerStatus.ENABLED, +) -> ExporleBanner: banner = ExporleBanner( - content={"text": text}, + content=_content(title), link=link, - sort=1, - status=BannerStatus.ENABLED, + sort=sort, + status=status, language=language, ) banner.id = str(uuid4()) @@ -45,33 +63,150 @@ def _banner(*, text: str, language: str, link: str, created_at: datetime) -> Exp return banner +def _record( + *, + title: str = "hello", + category: str = "Featured", + description: str = "Banner description", +) -> ExploreBannerRecord: + return ExploreBannerRecord( + id="banner-1", + content=_content(title, category=category, description=description), + link="https://example.com", + sort=1, + status=BannerStatus.ENABLED.value, + created_at=datetime(2024, 1, 1), + ) + + +def _use_sqlite_banner_service( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], +) -> None: + service = ExploreBannerQueryService( + banners=ExploreBannerQueryRepository(sqlite_session_factory), + is_enabled=lambda: True, + ) + monkeypatch.setattr( + banner_module, + "application_services", + lambda: _ApplicationServicesStub(explore_banner_queries=service), + ) + + +class TestExploreBannerQueryService: + def test_returns_empty_without_querying_when_disabled(self) -> None: + banners = FakeExploreBannerQuery() + service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: False) + + assert service.list_for_language("fr-FR") == () + assert banners.requested_languages == [] + + def test_returns_requested_language(self) -> None: + record = _record() + banners = FakeExploreBannerQuery({"fr-FR": (record,)}) + service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True) + + assert service.list_for_language("fr-FR") == (record,) + assert banners.requested_languages == ["fr-FR"] + + def test_falls_back_to_en_us(self) -> None: + record = _record(title="fallback") + banners = FakeExploreBannerQuery({"en-US": (record,)}) + service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True) + + assert service.list_for_language("es-ES") == (record,) + assert banners.requested_languages == ["es-ES", "en-US"] + + def test_does_not_repeat_default_language_query(self) -> None: + banners = FakeExploreBannerQuery() + service = ExploreBannerQueryService(banners=banners, is_enabled=lambda: True) + + assert service.list_for_language("en-US") == () + assert banners.requested_languages == ["en-US"] + + +class TestExploreBannerQueryRepository: + def test_filters_language_and_status_and_orders_by_sort( + self, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: + created_at = datetime(2024, 1, 1) + second = _banner( + title="second", + language="fr-FR", + link="https://example.com/second", + created_at=created_at, + sort=2, + ) + first = _banner( + title="first", + language="fr-FR", + link="https://example.com/first", + created_at=created_at, + ) + disabled = _banner( + title="disabled", + language="fr-FR", + link="https://example.com/disabled", + created_at=created_at, + sort=0, + status=BannerStatus.DISABLED, + ) + english = _banner( + title="english", + language="en-US", + link="https://example.com/english", + created_at=created_at, + ) + sqlite_session.add_all([second, first, disabled, english]) + sqlite_session.commit() + + repository = ExploreBannerQueryRepository(sqlite_session_factory) + result = repository.list_enabled("fr-FR") + + assert [banner.id for banner in result] == [first.id, second.id] + assert result[0] == ExploreBannerRecord( + id=first.id, + content=_content("first"), + link="https://example.com/first", + sort=1, + status=BannerStatus.ENABLED.value, + created_at=created_at, + ) + + class TestBannerApi: - def test_get_banners_with_requested_language( + def test_get_serializes_requested_language( self, app: Flask, monkeypatch: pytest.MonkeyPatch, - banner_session: Session, - ): - api = banner_module.BannerApi() - method = unwrap(api.get) - + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: banner = _banner( - text="hello", + title="hello", language="fr-FR", link="https://example.com", created_at=datetime(2024, 1, 1), ) - banner_session.add(banner) - banner_session.commit() - monkeypatch.setattr(banner_module.db, "session", banner_session) + sqlite_session.add(banner) + sqlite_session.commit() + _use_sqlite_banner_service(monkeypatch, sqlite_session_factory) with app.test_request_context("/?language=fr-FR"): - result = method(api) + result = banner_module.BannerApi().get() assert result == [ { "id": banner.id, - "content": {"text": "hello"}, + "content": { + "category": "Featured", + "title": "hello", + "description": "Banner description", + "img-src": "https://example.com/banner.png", + }, "link": "https://example.com", "sort": 1, "status": "enabled", @@ -79,50 +214,75 @@ class TestBannerApi: } ] - def test_get_banners_fallback_to_en_us( + def test_get_uses_default_language( self, app: Flask, monkeypatch: pytest.MonkeyPatch, - banner_session: Session, - ): - api = banner_module.BannerApi() - method = unwrap(api.get) - + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: banner = _banner( - text="fallback", + title="default", language="en-US", - link="https://example.com/fallback", + link="https://example.com/default", created_at=datetime(2024, 1, 2), ) - banner_session.add(banner) - banner_session.commit() - monkeypatch.setattr(banner_module.db, "session", banner_session) + sqlite_session.add(banner) + sqlite_session.commit() + _use_sqlite_banner_service(monkeypatch, sqlite_session_factory) - with app.test_request_context("/?language=es-ES"): - result = method(api) + with app.test_request_context("/"): + result = banner_module.BannerApi().get() - assert result == [ - { - "id": banner.id, - "content": {"text": "fallback"}, - "link": "https://example.com/fallback", - "sort": 1, - "status": "enabled", - "created_at": "2024-01-02T00:00:00", - } - ] + assert result[0]["id"] == banner.id + assert result[0]["content"]["title"] == "default" - def test_get_banners_default_language_en_us( + def test_get_allows_empty_supporting_copy( self, app: Flask, monkeypatch: pytest.MonkeyPatch, - banner_session: Session, - ): - api = banner_module.BannerApi() - method = unwrap(api.get) - monkeypatch.setattr(banner_module.db, "session", banner_session) + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: + banner = _banner( + title="hello", + language="en-US", + link="https://example.com", + created_at=datetime(2024, 1, 3), + ) + banner.content["category"] = "" + banner.content["description"] = "" + sqlite_session.add(banner) + sqlite_session.commit() + _use_sqlite_banner_service(monkeypatch, sqlite_session_factory) with app.test_request_context("/"): - result = method(api) + result = banner_module.BannerApi().get() - assert result == [] + assert result[0]["content"] == { + "category": "", + "title": "hello", + "description": "", + "img-src": "https://example.com/banner.png", + } + + def test_get_rejects_invalid_content( + self, + app: Flask, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ) -> None: + banner = _banner( + title="invalid", + language="en-US", + link="https://example.com", + created_at=datetime(2024, 1, 4), + ) + banner.content = {"title": "invalid"} + sqlite_session.add(banner) + sqlite_session.commit() + _use_sqlite_banner_service(monkeypatch, sqlite_session_factory) + + with app.test_request_context("/"), pytest.raises(ValidationError): + banner_module.BannerApi().get() diff --git a/api/tests/unit_tests/controllers/console/explore/test_completion.py b/api/tests/unit_tests/controllers/console/explore/test_completion.py index 2e2e0f5e7e8..9ff013e27a2 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_completion.py +++ b/api/tests/unit_tests/controllers/console/explore/test_completion.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, PropertyMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError import controllers.console.explore.completion as completion_module @@ -57,7 +58,7 @@ def payload_patch(payload_data): class TestCompletionApi: - def test_post_success(self, app: Flask, completion_app, user, payload_patch): + def test_post_success(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -75,7 +76,13 @@ class TestCompletionApi: return_value=("ok", 200), ), ): - result = method(api, MagicMock(), user, completion_app) + result = method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) assert result == ("ok", 200) @@ -86,9 +93,15 @@ class TestCompletionApi: installed_app = _installed_app(AppMode.CHAT) with pytest.raises(NotCompletionAppError): - method(api, MagicMock(), user, installed_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate({"inputs": {}, "query": "hi"}), + MagicMock(), + user, + installed_app, + ) - def test_conversation_completed(self, app: Flask, completion_app, user, payload_patch): + def test_conversation_completed(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -102,9 +115,15 @@ class TestCompletionApi: ), ): with pytest.raises(ConversationCompletedError): - method(api, MagicMock(), user, completion_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) - def test_internal_error(self, app: Flask, completion_app, user, payload_patch): + def test_internal_error(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -118,9 +137,15 @@ class TestCompletionApi: ), ): with pytest.raises(InternalServerError): - method(api, MagicMock(), user, completion_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) - def test_conversation_not_exists(self, app: Flask, completion_app, user, payload_patch): + def test_conversation_not_exists(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -134,9 +159,15 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.NotFound): - method(api, MagicMock(), user, completion_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) - def test_app_unavailable(self, app: Flask, completion_app, user, payload_patch): + def test_app_unavailable(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -150,9 +181,15 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.AppUnavailableError): - method(api, MagicMock(), user, completion_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) - def test_provider_not_initialized(self, app: Flask, completion_app, user, payload_patch): + def test_provider_not_initialized(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -166,9 +203,15 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.ProviderNotInitializeError): - method(api, MagicMock(), user, completion_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) - def test_quota_exceeded(self, app: Flask, completion_app, user, payload_patch): + def test_quota_exceeded(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -182,9 +225,15 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.ProviderQuotaExceededError): - method(api, MagicMock(), user, completion_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) - def test_model_not_supported(self, app: Flask, completion_app, user, payload_patch): + def test_model_not_supported(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -198,9 +247,15 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.ProviderModelCurrentlyNotSupportError): - method(api, MagicMock(), user, completion_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) - def test_invoke_error(self, app: Flask, completion_app, user, payload_patch): + def test_invoke_error(self, app: Flask, completion_app, user, payload_patch, payload_data): api = completion_module.CompletionApi() method = unwrap(api.post) @@ -214,7 +269,13 @@ class TestCompletionApi: ), ): with pytest.raises(completion_module.CompletionRequestError): - method(api, MagicMock(), user, completion_app) + method( + api, + completion_module.CompletionMessageExplorePayload.model_validate(payload_data), + MagicMock(), + user, + completion_app, + ) class TestCompletionStopApi: @@ -239,7 +300,7 @@ class TestCompletionStopApi: class TestChatApi: - def test_post_success(self, app: Flask, chat_app, user, payload_patch): + def test_post_success(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -257,7 +318,9 @@ class TestChatApi: return_value=("ok", 200), ), ): - result = method(api, MagicMock(), user, chat_app) + result = method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) assert result == ("ok", 200) @@ -268,9 +331,15 @@ class TestChatApi: installed_app = _installed_app(AppMode.COMPLETION) with pytest.raises(NotChatAppError): - method(api, MagicMock(), user, installed_app) + method( + api, + completion_module.ChatMessagePayload.model_validate({"inputs": {}, "query": "hi"}), + MagicMock(), + user, + installed_app, + ) - def test_rate_limit_error(self, app: Flask, chat_app, user, payload_patch): + def test_rate_limit_error(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -284,9 +353,11 @@ class TestChatApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) - def test_conversation_completed_chat(self, app: Flask, chat_app, user, payload_patch): + def test_conversation_completed_chat(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -300,9 +371,11 @@ class TestChatApi: ), ): with pytest.raises(ConversationCompletedError): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) - def test_conversation_not_exists_chat(self, app: Flask, chat_app, user, payload_patch): + def test_conversation_not_exists_chat(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -316,9 +389,13 @@ class TestChatApi: ), ): with pytest.raises(completion_module.NotFound): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) - def test_invalid_conversation_id_fails_fast_as_not_found(self, app: Flask, chat_app, user) -> None: + def test_invalid_conversation_id_fails_fast_as_not_found( + self, app: Flask, chat_app, user, unbound_session: Session + ) -> None: # A nonexistent conversation_id must fail fast as 404, before the streaming # generator is created. Previously the lookup only ran inside the generator, # so an invalid id surfaced as a hang instead of a clean error. @@ -333,7 +410,7 @@ class TestChatApi: get_conversation_mock = MagicMock( side_effect=completion_module.services.errors.conversation.ConversationNotExistsError() ) - session = MagicMock() + session = unbound_session api = completion_module.ChatApi() method = unwrap(api.post) @@ -349,13 +426,21 @@ class TestChatApi: patch.object(completion_module.AppGenerateService, "generate", generate_mock), ): with pytest.raises(completion_module.NotFound): - method(api, session, user, chat_app) + method( + api, + completion_module.ChatMessagePayload.model_validate( + {"inputs": {}, "query": "hi", "conversation_id": conversation_id} + ), + session, + user, + chat_app, + ) # The lookup must run before generation, so the generator is never started. generate_mock.assert_not_called() assert get_conversation_mock.call_args.kwargs["session"] is session - def test_app_unavailable_chat(self, app: Flask, chat_app, user, payload_patch): + def test_app_unavailable_chat(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -369,9 +454,11 @@ class TestChatApi: ), ): with pytest.raises(completion_module.AppUnavailableError): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) - def test_provider_not_initialized_chat(self, app: Flask, chat_app, user, payload_patch): + def test_provider_not_initialized_chat(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -385,9 +472,11 @@ class TestChatApi: ), ): with pytest.raises(completion_module.ProviderNotInitializeError): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) - def test_quota_exceeded_chat(self, app: Flask, chat_app, user, payload_patch): + def test_quota_exceeded_chat(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -401,9 +490,11 @@ class TestChatApi: ), ): with pytest.raises(completion_module.ProviderQuotaExceededError): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) - def test_model_not_supported_chat(self, app: Flask, chat_app, user, payload_patch): + def test_model_not_supported_chat(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -417,9 +508,11 @@ class TestChatApi: ), ): with pytest.raises(completion_module.ProviderModelCurrentlyNotSupportError): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) - def test_invoke_error_chat(self, app: Flask, chat_app, user, payload_patch): + def test_invoke_error_chat(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -433,9 +526,11 @@ class TestChatApi: ), ): with pytest.raises(completion_module.CompletionRequestError): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) - def test_internal_error_chat(self, app: Flask, chat_app, user, payload_patch): + def test_internal_error_chat(self, app: Flask, chat_app, user, payload_patch, payload_data): api = completion_module.ChatApi() method = unwrap(api.post) @@ -449,7 +544,9 @@ class TestChatApi: ), ): with pytest.raises(InternalServerError): - method(api, MagicMock(), user, chat_app) + method( + api, completion_module.ChatMessagePayload.model_validate(payload_data), MagicMock(), user, chat_app + ) class TestChatStopApi: diff --git a/api/tests/unit_tests/controllers/console/explore/test_installed_app.py b/api/tests/unit_tests/controllers/console/explore/test_installed_app.py index c5780e46ede..2c1a1e488ca 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_installed_app.py +++ b/api/tests/unit_tests/controllers/console/explore/test_installed_app.py @@ -1,19 +1,36 @@ -from collections.abc import Callable -from contextlib import AbstractContextManager +from collections.abc import Callable, Generator +from contextlib import AbstractContextManager, contextmanager from datetime import datetime +from inspect import unwrap +from types import SimpleNamespace from unittest.mock import MagicMock, PropertyMock, patch import pytest from flask import Flask +from sqlalchemy import select +from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest, Forbidden, NotFound import controllers.console.explore.installed_app as module +import services.installed_app_service as service_module +from models.model import App, AppMode, AppModelConfig, IconType, InstalledApp, RecommendedApp +from models.workflow import Workflow, WorkflowKind, WorkflowType type Payload = dict[str, object] type PayloadPatch = Callable[[Payload], AbstractContextManager[object]] -from inspect import unwrap +def make_app_model(app_id: str) -> MagicMock: + app_model = MagicMock() + app_model.id = app_id + app_model.name = f"App {app_id}" + app_model.description = "Description" + app_model.mode = AppMode.CHAT + app_model.icon_type = IconType.EMOJI + app_model.icon = "robot" + app_model.icon_background = "#FFFFFF" + app_model.use_icon_as_answer_icon = False + return app_model @pytest.fixture @@ -33,7 +50,7 @@ def current_user(tenant_id: str) -> MagicMock: def installed_app() -> MagicMock: app = MagicMock() app.id = "ia1" - app.app = MagicMock(id="a1") + app.app = make_app_model("a1") app.app_owner_tenant_id = "t2" app.is_pinned = False app.last_used_at = datetime(2024, 1, 1) @@ -54,8 +71,30 @@ def payload_patch() -> PayloadPatch: class TestInstalledAppsListApi: + def test_list_query_defaults_to_20(self) -> None: + assert module.InstalledAppsListQuery().limit == 20 + + def test_response_schema_preserves_installed_app_domain_types(self) -> None: + app_schema = module.InstalledAppInfoResponse.model_json_schema(mode="serialization") + list_schema = module.InstalledAppListResponse.model_json_schema(mode="serialization") + + assert { + "id", + "name", + "description", + "mode", + "icon_type", + "icon", + "icon_background", + "use_icon_as_answer_icon", + "icon_url", + } <= set(app_schema["required"]) + assert set(app_schema["$defs"]["AppMode"]["enum"]) == {mode.value for mode in AppMode} + assert set(app_schema["$defs"]["IconType"]["enum"]) == {icon_type.value for icon_type in IconType} + assert "next_cursor" in list_schema["required"] + def test_published_app_filter_checks_publish_targets(self) -> None: - compiled_filter = str(module._published_app_filter().compile(compile_kwargs={"literal_binds": True})) + compiled_filter = str(service_module._published_app_filter().compile(compile_kwargs={"literal_binds": True})) assert "workflows" in compiled_filter assert "app_model_configs" in compiled_filter @@ -77,7 +116,7 @@ class TestInstalledAppsListApi: patch.object(module.db, "session", session), patch.object(module.TenantService, "get_user_role", return_value="owner"), patch.object( - module.FeatureService, + service_module.FeatureService, "get_system_features", return_value=MagicMock(webapp_auth=MagicMock(enabled=False)), ), @@ -87,27 +126,208 @@ class TestInstalledAppsListApi: assert "installed_apps" in result assert result["installed_apps"][0]["editable"] is True assert result["installed_apps"][0]["uninstallable"] is False + assert result["has_more"] is False + assert result["next_cursor"] is None + executed_stmt = session.execute.call_args.args[0] + assert 21 in executed_stmt.compile().params.values() def test_get_installed_apps_with_app_id_filter(self, app: Flask, current_user: MagicMock, tenant_id: str) -> None: api = module.InstalledAppsListApi() method = unwrap(api.get) session = MagicMock() - session.execute.return_value.all.return_value = [] + session.execute.return_value.all.return_value = list[tuple[InstalledApp, App]]() with ( app.test_request_context("/?app_id=a1"), patch.object(module.db, "session", session), patch.object(module.TenantService, "get_user_role", return_value="member"), patch.object( - module.FeatureService, + service_module.FeatureService, "get_system_features", return_value=MagicMock(webapp_auth=MagicMock(enabled=False)), ), ): result = method(api, tenant_id, current_user) - assert result == {"installed_apps": []} + assert result == {"installed_apps": [], "has_more": False, "next_cursor": None} + + def test_get_installed_apps_escapes_name_search(self, app: Flask, current_user: MagicMock, tenant_id: str) -> None: + api = module.InstalledAppsListApi() + method = unwrap(api.get) + session = MagicMock() + session.execute.return_value.all.return_value = list[tuple[InstalledApp, App]]() + + with ( + app.test_request_context("/?name=Sales%25_Q3"), + patch.object(module.db, "session", session), + patch.object(module.TenantService, "get_user_role", return_value="owner"), + patch.object( + service_module.FeatureService, + "get_system_features", + return_value=MagicMock(webapp_auth=MagicMock(enabled=False)), + ), + ): + method(api, tenant_id, current_user) + + executed_stmt = session.execute.call_args.args[0] + assert r"%Sales\%\_Q3%" in executed_stmt.compile().params.values() + + def test_get_installed_apps_returns_cursor_when_more_apps_exist( + self, app: Flask, current_user: MagicMock, tenant_id: str + ) -> None: + api = module.InstalledAppsListApi() + method = unwrap(api.get) + rows = [] + for index in range(3): + installed_app = MagicMock( + id=f"ia{index}", + app_owner_tenant_id="t2", + is_pinned=index == 0, + last_used_at=datetime(2024, 1, 3 - index), + ) + app_model = make_app_model(f"a{index}") + rows.append((installed_app, app_model)) + + session = MagicMock() + session.execute.return_value.all.return_value = rows + + with ( + app.test_request_context("/?limit=2"), + patch.object(module.db, "session", session), + patch.object(module.TenantService, "get_user_role", return_value="owner"), + patch.object( + service_module.FeatureService, + "get_system_features", + return_value=MagicMock(webapp_auth=MagicMock(enabled=False)), + ), + ): + result = method(api, tenant_id, current_user) + + assert [item["id"] for item in result["installed_apps"]] == ["ia0", "ia1"] + assert result["has_more"] is True + assert result["next_cursor"] + decoded_cursor = module._decode_installed_app_cursor(result["next_cursor"]) + assert decoded_cursor is not None + assert decoded_cursor.installed_app_id == "ia1" + + def test_get_installed_apps_filters_permissions_before_filling_page( + self, app: Flask, current_user: MagicMock, tenant_id: str + ) -> None: + api = module.InstalledAppsListApi() + method = unwrap(api.get) + rows = [] + for index in range(3): + installed_app = MagicMock( + id=f"ia{index}", + app_owner_tenant_id="t2", + is_pinned=False, + last_used_at=datetime(2024, 1, 3 - index), + ) + app_model = make_app_model(f"a{index}") + rows.append((installed_app, app_model)) + + session = MagicMock() + session.execute.return_value.all.return_value = rows + restricted = MagicMock(access_mode="restricted") + + with ( + app.test_request_context("/?limit=1"), + patch.object(module.db, "session", session), + patch.object(module.TenantService, "get_user_role", return_value="member"), + patch.object( + service_module.FeatureService, + "get_system_features", + return_value=MagicMock(webapp_auth=MagicMock(enabled=True)), + ), + patch.object( + service_module.EnterpriseService.WebAppAuth, + "batch_get_app_access_mode_by_id", + return_value={"a0": restricted, "a1": restricted, "a2": restricted}, + ), + patch.object( + service_module.EnterpriseService.WebAppAuth, + "batch_is_user_allowed_to_access_webapps", + return_value={"a0": False, "a1": True, "a2": True}, + ), + ): + result = method(api, tenant_id, current_user) + + assert [item["id"] for item in result["installed_apps"]] == ["ia1"] + assert result["has_more"] is True + + def test_get_installed_apps_scans_past_denied_candidate_batch( + self, app: Flask, current_user: MagicMock, tenant_id: str + ) -> None: + api = module.InstalledAppsListApi() + method = unwrap(api.get) + allowed_rows = [ + ( + MagicMock( + id=f"allowed-{index}", + app_owner_tenant_id="t2", + is_pinned=False, + last_used_at=datetime(2024, 1, 2) if index == 0 else datetime(2023, 12, 31), + ), + make_app_model(f"allowed-app-{index}"), + ) + for index in range(2) + ] + denied_rows = [ + ( + MagicMock( + id=f"denied-{index:03}", + app_owner_tenant_id="t2", + is_pinned=False, + last_used_at=datetime(2024, 1, 1), + ), + make_app_model(f"denied-app-{index:03}"), + ) + for index in range(1) + ] + first_batch = [allowed_rows[0], *denied_rows] + first_result = MagicMock() + first_result.all.return_value = first_batch + second_result = MagicMock() + second_result.all.return_value = [allowed_rows[1]] + session = MagicMock() + session.execute.side_effect = [first_result, second_result] + + with ( + app.test_request_context("/?limit=1"), + patch.object(module.db, "session", session), + patch.object(module.TenantService, "get_user_role", return_value="member"), + patch.object( + service_module.FeatureService, + "get_system_features", + return_value=MagicMock(webapp_auth=MagicMock(enabled=True)), + ), + patch.object( + service_module, + "_filter_rows_by_webapp_auth", + side_effect=[[allowed_rows[0]], [allowed_rows[1]]], + ), + ): + result = method(api, tenant_id, current_user) + + assert [item["id"] for item in result["installed_apps"]] == ["allowed-0"] + assert result["has_more"] is True + assert session.execute.call_count == 2 + second_stmt = session.execute.call_args_list[1].args[0] + assert "denied-000" in second_stmt.compile().params.values() + next_cursor = module._decode_installed_app_cursor(result["next_cursor"]) + assert next_cursor is not None + assert next_cursor.installed_app_id == "denied-000" + + def test_get_installed_apps_rejects_invalid_cursor( + self, app: Flask, current_user: MagicMock, tenant_id: str + ) -> None: + api = module.InstalledAppsListApi() + method = unwrap(api.get) + + with app.test_request_context("/?cursor=not-a-cursor"): + with pytest.raises(BadRequest, match="Invalid cursor"): + method(api, tenant_id, current_user) def test_get_installed_apps_with_webapp_auth_enabled( self, app: Flask, current_user: MagicMock, tenant_id: str, installed_app: MagicMock @@ -127,17 +347,17 @@ class TestInstalledAppsListApi: patch.object(module.db, "session", session), patch.object(module.TenantService, "get_user_role", return_value="owner"), patch.object( - module.FeatureService, + service_module.FeatureService, "get_system_features", return_value=MagicMock(webapp_auth=MagicMock(enabled=True)), ), patch.object( - module.EnterpriseService.WebAppAuth, + service_module.EnterpriseService.WebAppAuth, "batch_get_app_access_mode_by_id", return_value={"a1": mock_webapp_setting}, ), patch.object( - module.EnterpriseService.WebAppAuth, + service_module.EnterpriseService.WebAppAuth, "batch_is_user_allowed_to_access_webapps", return_value={"a1": True}, ), @@ -145,6 +365,8 @@ class TestInstalledAppsListApi: result = method(api, tenant_id, current_user) assert len(result["installed_apps"]) == 1 + executed_stmt = session.execute.call_args.args[0] + assert 40 in executed_stmt.compile().params.values() def test_get_installed_apps_with_webapp_auth_user_denied( self, app: Flask, current_user: MagicMock, tenant_id: str, installed_app: MagicMock @@ -164,17 +386,17 @@ class TestInstalledAppsListApi: patch.object(module.db, "session", session), patch.object(module.TenantService, "get_user_role", return_value="member"), patch.object( - module.FeatureService, + service_module.FeatureService, "get_system_features", return_value=MagicMock(webapp_auth=MagicMock(enabled=True)), ), patch.object( - module.EnterpriseService.WebAppAuth, + service_module.EnterpriseService.WebAppAuth, "batch_get_app_access_mode_by_id", return_value={"a1": mock_webapp_setting}, ), patch.object( - module.EnterpriseService.WebAppAuth, + service_module.EnterpriseService.WebAppAuth, "batch_is_user_allowed_to_access_webapps", return_value={"a1": False}, ), @@ -201,12 +423,12 @@ class TestInstalledAppsListApi: patch.object(module.db, "session", session), patch.object(module.TenantService, "get_user_role", return_value="owner"), patch.object( - module.FeatureService, + service_module.FeatureService, "get_system_features", return_value=MagicMock(webapp_auth=MagicMock(enabled=True)), ), patch.object( - module.EnterpriseService.WebAppAuth, + service_module.EnterpriseService.WebAppAuth, "batch_get_app_access_mode_by_id", return_value={"a1": mock_webapp_setting}, ), @@ -221,14 +443,14 @@ class TestInstalledAppsListApi: method = unwrap(api.get) session = MagicMock() - session.execute.return_value.all.return_value = [] + session.execute.return_value.all.return_value = list[tuple[InstalledApp, App]]() with ( app.test_request_context("/"), patch.object(module.db, "session", session), patch.object(module.TenantService, "get_user_role", return_value="owner"), patch.object( - module.FeatureService, + service_module.FeatureService, "get_system_features", return_value=MagicMock(webapp_auth=MagicMock(enabled=False)), ), @@ -244,14 +466,14 @@ class TestInstalledAppsListApi: method = unwrap(api.get) session = MagicMock() - session.execute.return_value.all.return_value = [] + session.execute.return_value.all.return_value = list[tuple[InstalledApp, App]]() with ( app.test_request_context("/"), patch.object(module.db, "session", session), patch.object(module.TenantService, "get_user_role", return_value="owner"), patch.object( - module.FeatureService, + service_module.FeatureService, "get_system_features", return_value=MagicMock(webapp_auth=MagicMock(enabled=False)), ), @@ -267,14 +489,14 @@ class TestInstalledAppsListApi: method = unwrap(api.get) session = MagicMock() - session.execute.return_value.all.return_value = [] + session.execute.return_value.all.return_value = list[tuple[InstalledApp, App]]() with ( app.test_request_context("/"), patch.object(module.db, "session", session), patch.object(module.TenantService, "get_user_role", return_value="owner"), patch.object( - module.FeatureService, + service_module.FeatureService, "get_system_features", return_value=MagicMock(webapp_auth=MagicMock(enabled=False)), ), @@ -326,7 +548,7 @@ class TestInstalledAppsCreateApi: payload_patch({"app_id": "a1"}), patch.object(module.db, "session", session), ): - result = method(api, tenant_id) + result = method(api, module.InstalledAppCreatePayload.model_validate({"app_id": "a1"}), tenant_id) assert result == {"message": "App installed successfully"} assert recommended.install_count == 1 @@ -344,7 +566,7 @@ class TestInstalledAppsCreateApi: patch.object(module.db, "session", session), ): with pytest.raises(NotFound): - method(api, tenant_id) + method(api, module.InstalledAppCreatePayload.model_validate({"app_id": "a1"}), tenant_id) def test_post_app_not_public(self, app: Flask, tenant_id: str, payload_patch: PayloadPatch) -> None: api = module.InstalledAppsListApi() @@ -365,10 +587,56 @@ class TestInstalledAppsCreateApi: patch.object(module.db, "session", session), ): with pytest.raises(Forbidden): - method(api, tenant_id) + method(api, module.InstalledAppCreatePayload.model_validate({"app_id": "a1"}), tenant_id) class TestInstalledAppApi: + def test_get_installed_app( + self, + app: Flask, + current_user: MagicMock, + tenant_id: str, + installed_app: MagicMock, + ) -> None: + api = module.InstalledAppApi() + method = unwrap(api.get) + app_model = installed_app.app + session = MagicMock() + session.scalar.return_value = app_model + + with ( + app.test_request_context("/"), + patch.object(module.db, "session", session), + patch.object(module.TenantService, "get_user_role", return_value="owner"), + ): + result = method(api, tenant_id, current_user, installed_app) + + assert result["id"] == installed_app.id + assert result["app"]["id"] == app_model.id + assert result["app"]["mode"] == AppMode.CHAT + assert result["app"]["icon_type"] == IconType.EMOJI + assert result["app"]["use_icon_as_answer_icon"] is False + assert result["editable"] is True + + def test_get_installed_app_rejects_unpublished_app( + self, + app: Flask, + current_user: MagicMock, + tenant_id: str, + installed_app: MagicMock, + ) -> None: + api = module.InstalledAppApi() + method = unwrap(api.get) + session = MagicMock() + session.scalar.return_value = None + + with ( + app.test_request_context("/"), + patch.object(module.db, "session", session), + ): + with pytest.raises(NotFound, match="Installed app not found"): + method(api, tenant_id, current_user, installed_app) + def test_delete_success(self, tenant_id: str, installed_app: MagicMock) -> None: api = module.InstalledAppApi() method = unwrap(api.delete) @@ -397,7 +665,7 @@ class TestInstalledAppApi: payload_patch({"is_pinned": True}), patch.object(module.db, "session"), ): - result = method(installed_app) + result = method(api, module.InstalledAppUpdatePayload.model_validate({"is_pinned": True}), installed_app) assert installed_app.is_pinned is True assert result["result"] == "success" @@ -407,6 +675,307 @@ class TestInstalledAppApi: method = unwrap(api.patch) with app.test_request_context("/", json={}), payload_patch({}), patch.object(module.db, "session"): - result = method(installed_app) + result = method(api, module.InstalledAppUpdatePayload.model_validate({}), installed_app) assert result["result"] == "success" + + +def _persist_app( + session: Session, + *, + app_id: str = "app-1", + tenant_id: str = "owner-tenant", + mode: AppMode = AppMode.CHAT, + public: bool = True, + published: bool = True, +) -> App: + app = App( + id=app_id, + tenant_id=tenant_id, + name=f"App {app_id}", + description="description", + mode=mode, + icon_type=None, + icon=None, + icon_background=None, + enable_site=True, + enable_api=True, + is_public=public, + max_active_requests=None, + ) + session.add(app) + session.flush() + if published and mode in {AppMode.WORKFLOW, AppMode.ADVANCED_CHAT}: + workflow = Workflow( + id=f"workflow-{app_id}", + tenant_id=tenant_id, + app_id=app_id, + type=WorkflowType.WORKFLOW, + kind=WorkflowKind.STANDARD, + version="1", + graph='{"nodes":[],"edges":[]}', + features="{}", + created_by="user-1", + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + session.add(workflow) + app.workflow_id = workflow.id + elif published: + model_config = AppModelConfig(app_id=app_id) + session.add(model_config) + session.flush() + app.app_model_config_id = model_config.id + session.commit() + return app + + +def _persist_installed_app( + session: Session, + app: App, + *, + tenant_id: str, + pinned: bool = False, +) -> InstalledApp: + installed = InstalledApp( + app_id=app.id, + tenant_id=tenant_id, + app_owner_tenant_id=app.tenant_id, + is_pinned=pinned, + last_used_at=datetime(2024, 1, 1), + ) + session.add(installed) + session.commit() + return installed + + +@contextmanager +def _sqlite_controller_context( + database: Session, *, role: str = "owner", auth_enabled: bool = False +) -> Generator[None]: + session_proxy = MagicMock(wraps=database) + session_proxy.return_value = database + with ( + patch.object(module.db, "session", session_proxy), + patch.object(module.TenantService, "get_user_role", return_value=role), + patch.object( + service_module.FeatureService, + "get_system_features", + return_value=MagicMock(webapp_auth=MagicMock(enabled=auth_enabled)), + ), + ): + yield + + +def test_sqlite_get_installed_apps_filters_tenant_publication_mode_and_app_id( + app: Flask, + current_user: MagicMock, + tenant_id: str, + sqlite_session: Session, +) -> None: + database = sqlite_session + chat = _persist_app(database, app_id="chat") + workflow = _persist_app(database, app_id="workflow", mode=AppMode.WORKFLOW) + unpublished = _persist_app(database, app_id="unpublished", published=False) + agent = _persist_app(database, app_id="agent", mode=AppMode.AGENT) + foreign = _persist_app(database, app_id="foreign") + for model, installed_tenant in ( + (chat, tenant_id), + (workflow, tenant_id), + (unpublished, tenant_id), + (agent, tenant_id), + (foreign, "other-tenant"), + ): + _persist_installed_app(database, model, tenant_id=installed_tenant) + + api = module.InstalledAppsListApi() + method = unwrap(api.get) + with app.test_request_context("/"), _sqlite_controller_context(database): + result = method(api, tenant_id, current_user) + + assert {item["app"]["id"] for item in result["installed_apps"]} == {"chat", "workflow"} + assert all(item["editable"] is True for item in result["installed_apps"]) + assert all(item["uninstallable"] is False for item in result["installed_apps"]) + + with app.test_request_context("/?app_id=workflow"), _sqlite_controller_context(database, role="member"): + filtered = method(api, tenant_id, current_user) + assert [item["app"]["id"] for item in filtered["installed_apps"]] == ["workflow"] + assert filtered["installed_apps"][0]["editable"] is False + + +def test_sqlite_get_installed_apps_applies_web_auth_permission_state( + app: Flask, + current_user: MagicMock, + tenant_id: str, + sqlite_session: Session, +) -> None: + database = sqlite_session + allowed = _persist_app(database, app_id="allowed") + denied = _persist_app(database, app_id="denied") + sso = _persist_app(database, app_id="sso") + for model in (allowed, denied, sso): + _persist_installed_app(database, model, tenant_id=tenant_id) + settings = { + "allowed": SimpleNamespace(access_mode="restricted"), + "denied": SimpleNamespace(access_mode="restricted"), + "sso": SimpleNamespace(access_mode="sso_verified"), + } + + api = module.InstalledAppsListApi() + method = unwrap(api.get) + with ( + app.test_request_context("/"), + _sqlite_controller_context(database, auth_enabled=True), + patch.object( + service_module.EnterpriseService.WebAppAuth, "batch_get_app_access_mode_by_id", return_value=settings + ), + patch.object( + service_module.EnterpriseService.WebAppAuth, + "batch_is_user_allowed_to_access_webapps", + return_value={"allowed": True, "denied": False}, + ), + ): + result = method(api, tenant_id, current_user) + + assert [item["app"]["id"] for item in result["installed_apps"]] == ["allowed"] + + +def test_sqlite_post_installs_public_recommended_app_and_is_idempotent( + app: Flask, + tenant_id: str, + payload_patch: PayloadPatch, + sqlite_session: Session, +) -> None: + database = sqlite_session + app_model = _persist_app(database, public=True) + recommended = RecommendedApp( + app_id=app_model.id, + description={"en-US": "recommended"}, + copyright="copyright", + privacy_policy="https://example.com/privacy", + category="productivity", + ) + database.add(recommended) + database.commit() + recommended_id = recommended.id + app_owner_tenant_id = app_model.tenant_id + api = module.InstalledAppsListApi() + method = unwrap(api.post) + + for _ in range(2): + with ( + app.test_request_context("/", json={"app_id": app_model.id}), + payload_patch({"app_id": app_model.id}), + patch.object(module.db, "session", database), + ): + assert method( + api, module.InstalledAppCreatePayload.model_validate({"app_id": app_model.id}), tenant_id + ) == {"message": "App installed successfully"} + + # End the request-scoped session so this assertion only observes committed data. + bind = database.get_bind() + database.close() + with Session(bind) as verification_session: + installed = verification_session.scalars(select(InstalledApp)).all() + assert len(installed) == 1 + assert installed[0].tenant_id == tenant_id + assert installed[0].app_owner_tenant_id == app_owner_tenant_id + persisted_recommendation = verification_session.get(RecommendedApp, recommended_id) + assert persisted_recommendation is not None + assert persisted_recommendation.install_count == 1 + + +def test_sqlite_post_enforces_recommendation_and_public_state( + app: Flask, + tenant_id: str, + payload_patch: PayloadPatch, + sqlite_session: Session, +) -> None: + database = sqlite_session + api = module.InstalledAppsListApi() + method = unwrap(api.post) + with ( + app.test_request_context("/", json={"app_id": "missing"}), + payload_patch({"app_id": "missing"}), + patch.object(module.db, "session", database), + pytest.raises(NotFound), + ): + method(api, module.InstalledAppCreatePayload.model_validate({"app_id": "missing"}), tenant_id) + + private_app = _persist_app(database, app_id="private", public=False) + database.add( + RecommendedApp( + app_id=private_app.id, + description={}, + copyright="copyright", + privacy_policy="privacy", + category="category", + ) + ) + database.commit() + with ( + app.test_request_context("/", json={"app_id": private_app.id}), + payload_patch({"app_id": private_app.id}), + patch.object(module.db, "session", database), + pytest.raises(Forbidden), + ): + method(api, module.InstalledAppCreatePayload.model_validate({"app_id": private_app.id}), tenant_id) + + +def test_sqlite_delete_removes_foreign_installed_app_and_rejects_owned_app( + tenant_id: str, + sqlite_session: Session, +) -> None: + database = sqlite_session + foreign_app = _persist_app(database, app_id="foreign") + installed = _persist_installed_app(database, foreign_app, tenant_id=tenant_id) + installed_id = installed.id + api = module.InstalledAppApi() + with patch.object(module.db, "session", database): + response, status = unwrap(api.delete)(api, tenant_id, installed) + assert (response, status) == ("", 204) + assert database.get(InstalledApp, installed_id) is None + + owned_app = _persist_app(database, app_id="owned", tenant_id=tenant_id) + owned_install = _persist_installed_app(database, owned_app, tenant_id=tenant_id) + with pytest.raises(BadRequest): + unwrap(api.delete)(api, tenant_id, owned_install) + assert database.get(InstalledApp, owned_install.id) is not None + + +def test_sqlite_patch_persists_pin_and_noop_payload( + app: Flask, + tenant_id: str, + payload_patch: PayloadPatch, + sqlite_session: Session, +) -> None: + database = sqlite_session + app_model = _persist_app(database) + installed = _persist_installed_app(database, app_model, tenant_id=tenant_id) + api = module.InstalledAppApi() + with ( + app.test_request_context("/", json={"is_pinned": True}), + payload_patch({"is_pinned": True}), + patch.object(module.db, "session", database), + ): + assert ( + unwrap(api.patch)(api, module.InstalledAppUpdatePayload.model_validate({"is_pinned": True}), installed)[ + "result" + ] + == "success" + ) + database.expire_all() + persisted_installed_app = database.get(InstalledApp, installed.id) + assert persisted_installed_app is not None + assert persisted_installed_app.is_pinned is True + + with ( + app.test_request_context("/", json={}), + payload_patch({}), + patch.object(module.db, "session", database), + ): + assert ( + unwrap(api.patch)(api, module.InstalledAppUpdatePayload.model_validate({}), installed)["result"] + == "success" + ) diff --git a/api/tests/unit_tests/controllers/console/explore/test_message.py b/api/tests/unit_tests/controllers/console/explore/test_message.py index a66869488b8..cb5a50f2346 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_message.py +++ b/api/tests/unit_tests/controllers/console/explore/test_message.py @@ -160,7 +160,9 @@ class TestMessageFeedbackApi: "create_feedback", ), ): - result = method(MagicMock(), installed_app, "mid") + result = method( + module.MessageFeedbackPayload.model_validate({"rating": "like"}), MagicMock(), installed_app, "mid" + ) assert result["result"] == "success" @@ -179,7 +181,7 @@ class TestMessageFeedbackApi: ), ): with pytest.raises(NotFound): - method(MagicMock(), installed_app, "mid") + method(module.MessageFeedbackPayload.model_validate({}), MagicMock(), installed_app, "mid") class TestMessageMoreLikeThisApi: diff --git a/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py b/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py index 4adeaaa90dd..d354e22541d 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py +++ b/api/tests/unit_tests/controllers/console/explore/test_recommended_app.py @@ -1,9 +1,12 @@ from inspect import unwrap from unittest.mock import ANY, patch +import pytest from flask import Flask +from pydantic import ValidationError import controllers.console.explore.recommended_app as module +from controllers.console.explore.recommended_app import RecommendedAppsQuery from models import Account from models.model import AppMode, IconType @@ -30,7 +33,7 @@ class TestRecommendedAppListApi: return_value=result_data, ) as service_mock, ): - result = method(api, make_account("fr-FR")) + result = method(api, RecommendedAppsQuery(language="en-US"), make_account("fr-FR")) service_mock.assert_called_once_with("en-US", session=ANY) assert result == result_data @@ -49,7 +52,7 @@ class TestRecommendedAppListApi: return_value=result_data, ) as service_mock, ): - result = method(api, make_account("fr-FR")) + result = method(api, RecommendedAppsQuery(), make_account("fr-FR")) service_mock.assert_called_once_with("fr-FR", session=ANY) assert result == result_data @@ -68,7 +71,7 @@ class TestRecommendedAppListApi: return_value=result_data, ) as service_mock, ): - result = method(api, make_account(None)) + result = method(api, RecommendedAppsQuery(), make_account(None)) service_mock.assert_called_once_with(module.languages[0], session=ANY) assert result == result_data @@ -89,7 +92,7 @@ class TestLearnDifyAppListApi: return_value=result_data, ) as service_mock, ): - result = method(api, make_account("fr-FR")) + result = method(api, RecommendedAppsQuery(language="en-US"), make_account("fr-FR")) service_mock.assert_called_once_with("en-US", session=ANY) assert result == result_data @@ -108,7 +111,7 @@ class TestLearnDifyAppListApi: return_value=result_data, ) as service_mock, ): - result = method(api, make_account("fr-FR")) + result = method(api, RecommendedAppsQuery(), make_account("fr-FR")) service_mock.assert_called_once_with("fr-FR", session=ANY) assert result == result_data @@ -119,7 +122,13 @@ class TestRecommendedAppApi: api = module.RecommendedAppApi() method = unwrap(api.get) - result_data = {"id": "app1"} + result_data = { + "id": "app1", + "name": "App", + "mode": "chat", + "export_data": "{}", + "can_trial": False, + } with ( app.test_request_context("/"), @@ -132,7 +141,7 @@ class TestRecommendedAppApi: result = method(api, "11111111-1111-1111-1111-111111111111") service_mock.assert_called_once_with("11111111-1111-1111-1111-111111111111", session=ANY) - assert result == result_data + assert result == {**result_data, "icon": None, "icon_background": None} class TestRecommendedAppResponseModels: @@ -198,6 +207,7 @@ class TestRecommendedAppResponseModels: "categories": ["Workflow"], "position": 1, "is_listed": True, + "can_trial": False, } ], } @@ -205,3 +215,7 @@ class TestRecommendedAppResponseModels: assert response["recommended_apps"][0]["app_id"] == "app-1" assert response["recommended_apps"][0]["categories"] == ["Workflow"] + + def test_recommended_app_response_requires_can_trial(self): + with pytest.raises(ValidationError): + module.RecommendedAppResponse.model_validate({"app_id": "app-1"}) diff --git a/api/tests/unit_tests/controllers/console/explore/test_saved_message.py b/api/tests/unit_tests/controllers/console/explore/test_saved_message.py index ae36d69f7dd..a685c6c8fbf 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_saved_message.py +++ b/api/tests/unit_tests/controllers/console/explore/test_saved_message.py @@ -100,7 +100,7 @@ class TestSavedMessageListApi: payload_patch(payload), patch.object(module.SavedMessageService, "save") as save_mock, ): - result = method(api, current_user, installed_app) + result = method(api, module.SavedMessageCreatePayload.model_validate(payload), current_user, installed_app) save_mock.assert_called_once() assert save_mock.call_args.args[1] is current_user @@ -124,7 +124,7 @@ class TestSavedMessageListApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), installed_app) + method(api, module.SavedMessageCreatePayload.model_validate(payload), MagicMock(), installed_app) class TestSavedMessageApi: diff --git a/api/tests/unit_tests/controllers/console/explore/test_trial.py b/api/tests/unit_tests/controllers/console/explore/test_trial.py index 426f83d5047..d8838a76251 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_trial.py +++ b/api/tests/unit_tests/controllers/console/explore/test_trial.py @@ -8,7 +8,9 @@ from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest -from flask import Flask +from flask import Flask, request +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, InternalServerError, NotFound import controllers.console.explore.trial as module @@ -26,17 +28,20 @@ from controllers.console.explore.error import ( NotCompletionAppError, NotWorkflowAppError, ) +from controllers.console.explore.trial import ChatRequest, CompletionRequest, TextToSpeechRequest, WorkflowRunRequest from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError from core.errors.error import ( ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError, ) +from core.helper import encrypter +from core.workflow.llm_environment_variable import LLMEnvironmentVariable from graphon.model_runtime.errors.invoke import InvokeError -from graphon.variables import StringVariable +from graphon.variables import SecretVariable, StringVariable from models import Account from models.account import TenantStatus -from models.model import AppMode +from models.model import AppMode, Site from services.app_ref_service import AppRef, MessageRef from services.errors.audio import SpeechToTextDisabledServiceError from services.errors.conversation import ConversationNotExistsError @@ -45,6 +50,16 @@ from services.errors.llm import InvokeRateLimitError unwrap: Any = inspect_unwrap +class _UsesSQLiteSession: + sqlite_session: Session + + @pytest.fixture(autouse=True) + def _provide_sqlite_session(self, sqlite_engine: Engine): + with Session(sqlite_engine, expire_on_commit=False) as session: + self.sqlite_session = session + yield + + @pytest.fixture def account() -> Account: acc = Account(name="User", email="user@example.com") @@ -58,6 +73,18 @@ def _file_data() -> Any: return file_data +def _persist_site(sqlite_session: Session, app_id: str) -> Site: + site = Site( + app_id=app_id, + title="Trial Site", + default_language="en-US", + customize_token_strategy="uuid", + ) + sqlite_session.add(site) + sqlite_session.commit() + return site + + @pytest.fixture def trial_app_chat() -> MagicMock: app = MagicMock() @@ -104,7 +131,7 @@ def test_trial_workflow_uses_trial_scoped_simple_account_model() -> None: assert module.simple_account_model.__schema__["properties"].keys() >= {"id", "name", "email"} -def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask): +def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask, unbound_session: Session): class DatasetListItem: id = "dataset-1" name = "Dataset" @@ -123,8 +150,6 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask): api = module.DatasetListApi() method = unwrap(api.get) app_model = SimpleNamespace(tenant_id="tenant-1") - session = MagicMock() - with ( app.test_request_context("/?page=1&limit=20&ids=dataset-1"), patch.object( @@ -133,9 +158,9 @@ def test_trial_dataset_list_preserves_slim_dataset_fields(app: Flask): return_value=([DatasetListItem()], 1), ) as get_datasets, ): - result = method(api, session, app_model) + result = method(api, unbound_session, app_model) - get_datasets.assert_called_once_with(["dataset-1"], "tenant-1", session=session) + get_datasets.assert_called_once_with(["dataset-1"], "tenant-1", session=unbound_session) assert result == { "data": [ { @@ -168,8 +193,9 @@ def test_trial_app_handlers_use_explicit_read_session(api_type: type) -> None: assert tuple(signature(api_type.get).parameters)[:3] == ("self", "session", "app_model") -def test_trial_app_detail_serializes_with_explicit_session(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - session = MagicMock() +def test_trial_app_detail_serializes_with_explicit_session( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: app_model = MagicMock() response_view = MagicMock() get_app = MagicMock(return_value=app_model) @@ -181,11 +207,11 @@ def test_trial_app_detail_serializes_with_explicit_session(app: Flask, monkeypat monkeypatch.setattr(module.TrialAppDetailResponse, "model_validate", MagicMock(return_value=validated)) with app.test_request_context("/"): - result = unwrap(module.AppApi.get)(module.AppApi(), session, app_model) + result = unwrap(module.AppApi.get)(module.AppApi(), unbound_session, app_model) assert result == {"id": "app-1"} - get_app.assert_called_once_with(app_model, session=session) - build_view.assert_called_once_with(app_model, session=session) + get_app.assert_called_once_with(app_model, session=unbound_session) + build_view.assert_called_once_with(app_model, session=unbound_session) module.TrialAppDetailResponse.model_validate.assert_called_once_with(response_view, from_attributes=True) @@ -227,14 +253,20 @@ class TestTrialAppRemoteFileUploadApi: upload.assert_called_once_with(current_user=account, resource_tenant_id="app-tenant-id") -class TestTrialAppWorkflowRunApi: +class TestTrialAppWorkflowRunApi(_UsesSQLiteSession): def test_not_workflow_app(self, app: Flask, account: Account) -> None: api = module.TrialAppWorkflowRunApi() method = unwrap(api.post) - with app.test_request_context("/"): + with app.test_request_context("/", json={"inputs": {}}): with pytest.raises(NotWorkflowAppError): - method(api, MagicMock(), account, MagicMock(mode=AppMode.CHAT)) + method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + MagicMock(mode=AppMode.CHAT), + ) def test_success(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -245,7 +277,13 @@ class TestTrialAppWorkflowRunApi: patch.object(module.AppGenerateService, "generate", return_value=MagicMock()), patch.object(module.RecommendedAppService, "add_trial_app_record"), ): - result = method(api, MagicMock(), account, trial_app_workflow) + result = method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_workflow, + ) assert result is not None @@ -262,7 +300,13 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, MagicMock(), account, trial_app_workflow) + method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_workflow, + ) def test_workflow_quota_exceeded(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -277,7 +321,13 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(ProviderQuotaExceededError): - method(api, MagicMock(), account, trial_app_workflow) + method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_workflow, + ) def test_workflow_model_not_support(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -292,7 +342,13 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(api, MagicMock(), account, trial_app_workflow) + method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_workflow, + ) def test_workflow_invoke_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -307,7 +363,13 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(CompletionRequestError): - method(api, MagicMock(), account, trial_app_workflow) + method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_workflow, + ) def test_workflow_rate_limit_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -322,7 +384,13 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, MagicMock(), account, trial_app_workflow) + method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_workflow, + ) def test_workflow_value_error(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -337,7 +405,13 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(ValueError): - method(api, MagicMock(), account, trial_app_workflow) + method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_workflow, + ) def test_workflow_generic_exception(self, app: Flask, trial_app_workflow: MagicMock, account: Account) -> None: api = module.TrialAppWorkflowRunApi() @@ -352,17 +426,29 @@ class TestTrialAppWorkflowRunApi: ), ): with pytest.raises(InternalServerError): - method(api, MagicMock(), account, trial_app_workflow) + method( + api, + WorkflowRunRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_workflow, + ) -class TestTrialChatApi: +class TestTrialChatApi(_UsesSQLiteSession): def test_not_chat_app(self, app: Flask, account: Account) -> None: api = module.TrialChatApi() method = unwrap(api.post) with app.test_request_context("/", json={"inputs": {}, "query": "hi"}): with pytest.raises(NotChatAppError): - method(api, MagicMock(), account, MagicMock(mode="completion")) + method( + api, + ChatRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + MagicMock(mode="completion"), + ) def test_success(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -373,7 +459,9 @@ class TestTrialChatApi: patch.object(module.AppGenerateService, "generate", return_value=MagicMock()), patch.object(module.RecommendedAppService, "add_trial_app_record"), ): - result = method(api, MagicMock(), account, trial_app_chat) + result = method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) assert result is not None @@ -390,7 +478,9 @@ class TestTrialChatApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_conversation_completed(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -405,7 +495,9 @@ class TestTrialChatApi: ), ): with pytest.raises(ConversationCompletedError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_app_config_broken(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -420,7 +512,9 @@ class TestTrialChatApi: ), ): with pytest.raises(AppUnavailableError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -435,7 +529,9 @@ class TestTrialChatApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -450,7 +546,9 @@ class TestTrialChatApi: ), ): with pytest.raises(ProviderQuotaExceededError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_model_not_support(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -465,7 +563,9 @@ class TestTrialChatApi: ), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -480,7 +580,9 @@ class TestTrialChatApi: ), ): with pytest.raises(CompletionRequestError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_rate_limit_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -495,7 +597,9 @@ class TestTrialChatApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_value_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -510,7 +614,9 @@ class TestTrialChatApi: ), ): with pytest.raises(ValueError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) def test_chat_generic_exception(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatApi() @@ -525,17 +631,25 @@ class TestTrialChatApi: ), ): with pytest.raises(InternalServerError): - method(api, MagicMock(), account, trial_app_chat) + method( + api, ChatRequest.model_validate(request.get_json()), self.sqlite_session, account, trial_app_chat + ) -class TestTrialCompletionApi: +class TestTrialCompletionApi(_UsesSQLiteSession): def test_not_completion_app(self, app: Flask, account: Account) -> None: api = module.TrialCompletionApi() method = unwrap(api.post) with app.test_request_context("/", json={"inputs": {}, "query": ""}): with pytest.raises(NotCompletionAppError): - method(api, MagicMock(), account, MagicMock(mode=AppMode.CHAT)) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + MagicMock(mode=AppMode.CHAT), + ) def test_success(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -546,7 +660,13 @@ class TestTrialCompletionApi: patch.object(module.AppGenerateService, "generate", return_value=MagicMock()), patch.object(module.RecommendedAppService, "add_trial_app_record"), ): - result = method(api, MagicMock(), account, trial_app_completion) + result = method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) assert result is not None @@ -563,7 +683,13 @@ class TestTrialCompletionApi: ), ): with pytest.raises(AppUnavailableError): - method(api, MagicMock(), account, trial_app_completion) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) def test_completion_provider_not_init(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -578,7 +704,13 @@ class TestTrialCompletionApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, MagicMock(), account, trial_app_completion) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) def test_completion_quota_exceeded(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -593,7 +725,13 @@ class TestTrialCompletionApi: ), ): with pytest.raises(ProviderQuotaExceededError): - method(api, MagicMock(), account, trial_app_completion) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) def test_completion_model_not_support(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -608,7 +746,13 @@ class TestTrialCompletionApi: ), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(api, MagicMock(), account, trial_app_completion) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) def test_completion_invoke_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -623,7 +767,13 @@ class TestTrialCompletionApi: ), ): with pytest.raises(CompletionRequestError): - method(api, MagicMock(), account, trial_app_completion) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) def test_completion_rate_limit_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -638,7 +788,13 @@ class TestTrialCompletionApi: ), ): with pytest.raises(InternalServerError): - method(api, MagicMock(), account, trial_app_completion) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) def test_completion_value_error(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -653,7 +809,13 @@ class TestTrialCompletionApi: ), ): with pytest.raises(ValueError): - method(api, MagicMock(), account, trial_app_completion) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) def test_completion_generic_exception(self, app: Flask, trial_app_completion: MagicMock, account: Account) -> None: api = module.TrialCompletionApi() @@ -668,7 +830,13 @@ class TestTrialCompletionApi: ), ): with pytest.raises(InternalServerError): - method(api, MagicMock(), account, trial_app_completion) + method( + api, + CompletionRequest.model_validate(request.get_json()), + self.sqlite_session, + account, + trial_app_completion, + ) class TestTrialMessageSuggestedQuestionApi: @@ -713,14 +881,14 @@ class TestTrialMessageSuggestedQuestionApi: class TestTrialAppParameterApi: - def test_app_unavailable(self) -> None: + def test_app_unavailable(self, unbound_session: Session) -> None: api = module.TrialAppParameterApi() method = unwrap(api.get) with pytest.raises(AppUnavailableError): - method(api, MagicMock(), None) + method(api, unbound_session, None) - def test_success_non_workflow(self, valid_parameters: dict[str, object]) -> None: + def test_success_non_workflow(self, valid_parameters: dict[str, object], unbound_session: Session) -> None: api = module.TrialAppParameterApi() method = unwrap(api.get) @@ -730,7 +898,6 @@ class TestTrialAppParameterApi: mode=AppMode.CHAT, app_model_config_with_session=MagicMock(return_value=app_model_config), ) - session = MagicMock() annotation_reply = {"enabled": False} with ( @@ -748,14 +915,14 @@ class TestTrialAppParameterApi: return_value=MagicMock(model_dump=lambda mode=None: {"ok": True}), ), ): - result = method(api, session, app_model) + result = method(api, unbound_session, app_model) assert result == {"ok": True} - app_model.app_model_config_with_session.assert_called_once_with(session=session) - load_annotation_reply.assert_called_once_with(session, "app-1") + app_model.app_model_config_with_session.assert_called_once_with(session=unbound_session) + load_annotation_reply.assert_called_once_with(unbound_session, "app-1") app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) - def test_success_workflow(self, valid_parameters: dict[str, object]) -> None: + def test_success_workflow(self, valid_parameters: dict[str, object], unbound_session: Session) -> None: api = module.TrialAppParameterApi() method = unwrap(api.get) @@ -765,8 +932,6 @@ class TestTrialAppParameterApi: mode=AppMode.WORKFLOW, workflow_with_session=MagicMock(return_value=workflow), ) - session = MagicMock() - with ( patch.object(module, "get_parameters_from_feature_dict", return_value=valid_parameters), patch.object( @@ -775,10 +940,10 @@ class TestTrialAppParameterApi: return_value=MagicMock(model_dump=lambda mode=None: {"ok": True}), ), ): - result = method(api, session, app_model) + result = method(api, unbound_session, app_model) assert result == {"ok": True} - app_model.workflow_with_session.assert_called_once_with(session=session) + app_model.workflow_with_session.assert_called_once_with(session=unbound_session) workflow.user_input_form.assert_called_once_with(to_old_structure=True) @@ -817,7 +982,11 @@ class TestTrialChatAudioApi: ), ): with pytest.raises(module.AppUnavailableError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_no_audio_uploaded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatAudioApi() @@ -836,7 +1005,11 @@ class TestTrialChatAudioApi: ), ): with pytest.raises(module.NoAudioUploadedError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_missing_file_field_returns_400(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: """A multipart POST with no `file` field must surface as 400, not 500. @@ -857,7 +1030,11 @@ class TestTrialChatAudioApi: patch.object(module.AudioService, "transcript_asr", side_effect=fake_asr), ): with pytest.raises(module.NoAudioUploadedError) as exc_info: - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) assert exc_info.value.code == 400 @@ -878,7 +1055,11 @@ class TestTrialChatAudioApi: ), ): with pytest.raises(module.AudioTooLargeError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_unsupported_audio_type(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatAudioApi() @@ -897,7 +1078,11 @@ class TestTrialChatAudioApi: ), ): with pytest.raises(module.UnsupportedAudioTypeError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_provider_not_support_tts(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatAudioApi() @@ -916,7 +1101,11 @@ class TestTrialChatAudioApi: ), ): with pytest.raises(module.ProviderNotSupportSpeechToTextError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_speech_to_text_disabled(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatAudioApi() @@ -934,7 +1123,11 @@ class TestTrialChatAudioApi: ), ): with pytest.raises(SpeechToTextDisabledError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatAudioApi() @@ -949,7 +1142,11 @@ class TestTrialChatAudioApi: patch.object(module.AudioService, "transcript_asr", side_effect=ProviderTokenNotInitError("test")), ): with pytest.raises(ProviderNotInitializeError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatAudioApi() @@ -964,7 +1161,11 @@ class TestTrialChatAudioApi: patch.object(module.AudioService, "transcript_asr", side_effect=QuotaExceededError()), ): with pytest.raises(ProviderQuotaExceededError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) class TestTrialChatTextApi: @@ -977,7 +1178,9 @@ class TestTrialChatTextApi: patch.object(module.AudioService, "transcript_tts", return_value={"audio": "base64_data"}), patch.object(module.RecommendedAppService, "add_trial_app_record"), ): - result = method(api, account, trial_app_chat) + result = method( + api, TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), account, trial_app_chat + ) assert result == {"audio": "base64_data"} @@ -992,7 +1195,9 @@ class TestTrialChatTextApi: patch.object(module.AudioService, "transcript_tts", transcript_tts), patch.object(module.RecommendedAppService, "add_trial_app_record"), ): - result = method(api, account, trial_app_chat) + result = method( + api, TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), account, trial_app_chat + ) assert result == {"audio": "base64_data"} assert transcript_tts.call_args.kwargs["message_ref"] == MessageRef( @@ -1014,7 +1219,12 @@ class TestTrialChatTextApi: ), ): with pytest.raises(module.AppUnavailableError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) def test_provider_not_support(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() @@ -1029,7 +1239,12 @@ class TestTrialChatTextApi: ), ): with pytest.raises(module.ProviderNotSupportSpeechToTextError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) def test_audio_too_large(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() @@ -1044,7 +1259,12 @@ class TestTrialChatTextApi: ), ): with pytest.raises(module.AudioTooLargeError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) def test_no_audio_uploaded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() @@ -1059,7 +1279,12 @@ class TestTrialChatTextApi: ), ): with pytest.raises(module.NoAudioUploadedError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) def test_provider_not_init(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() @@ -1070,7 +1295,12 @@ class TestTrialChatTextApi: patch.object(module.AudioService, "transcript_tts", side_effect=ProviderTokenNotInitError("test")), ): with pytest.raises(ProviderNotInitializeError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() @@ -1081,7 +1311,12 @@ class TestTrialChatTextApi: patch.object(module.AudioService, "transcript_tts", side_effect=QuotaExceededError()), ): with pytest.raises(ProviderQuotaExceededError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) def test_model_not_support(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() @@ -1092,7 +1327,12 @@ class TestTrialChatTextApi: patch.object(module.AudioService, "transcript_tts", side_effect=ModelCurrentlyNotSupportError()), ): with pytest.raises(ProviderModelCurrentlyNotSupportError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) def test_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() @@ -1103,21 +1343,24 @@ class TestTrialChatTextApi: patch.object(module.AudioService, "transcript_tts", side_effect=InvokeError("test error")), ): with pytest.raises(CompletionRequestError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) class TestTrialAppWorkflowTaskStopApi: def test_not_workflow_app(self, app: Flask, trial_app_chat: MagicMock) -> None: api = module.TrialAppWorkflowTaskStopApi() - method = unwrap(api.post) - with app.test_request_context("/"): + with app.test_request_context("/", json={"inputs": {}}): with pytest.raises(NotWorkflowAppError): - method(api, trial_app_chat, str(uuid4())) + api.post(trial_app_chat, str(uuid4())) def test_success(self, app: Flask, trial_app_workflow: MagicMock) -> None: api = module.TrialAppWorkflowTaskStopApi() - method = unwrap(api.post) task_id = str(uuid4()) with ( @@ -1125,7 +1368,7 @@ class TestTrialAppWorkflowTaskStopApi: patch.object(module.AppQueueManager, "set_stop_flag_no_user_check") as mock_set_flag, patch.object(module.GraphEngineManager, "send_stop_command") as mock_send_cmd, ): - result = method(api, trial_app_workflow, task_id) + result = api.post(trial_app_workflow, task_id) assert result == {"result": "success"} mock_set_flag.assert_called_once_with(task_id) @@ -1133,49 +1376,55 @@ class TestTrialAppWorkflowTaskStopApi: class TestTrialSitApi: - def test_no_site(self, app: Flask) -> None: + @pytest.mark.parametrize("sqlite_session", [(Site,)], indirect=True) + def test_no_site( + self, + app: Flask, + sqlite_session: Session, + ) -> None: api = module.TrialSitApi() method = unwrap(api.get) app_model = MagicMock() - app_model.id = "a1" - session = MagicMock() - session.scalar.return_value = None + app_model.id = str(uuid4()) with app.test_request_context("/"): with pytest.raises(Forbidden): - method(api, session, app_model) + method(api, sqlite_session, app_model) - session.scalar.assert_called_once() - - def test_archived_tenant(self, app: Flask) -> None: + @pytest.mark.parametrize("sqlite_session", [(Site,)], indirect=True) + def test_archived_tenant( + self, + app: Flask, + sqlite_session: Session, + ) -> None: api = module.TrialSitApi() method = unwrap(api.get) - site = MagicMock() - app_model = SimpleNamespace(id="a1", tenant_id="tenant-1") + app_model = SimpleNamespace(id=str(uuid4()), tenant_id="tenant-1") tenant = SimpleNamespace(status=TenantStatus.ARCHIVE) - session = MagicMock() - session.scalar.return_value = site + _persist_site(sqlite_session, app_model.id) with ( app.test_request_context("/"), patch.object(module.TenantService, "get_tenant_by_id", return_value=tenant) as get_tenant_by_id, ): with pytest.raises(Forbidden): - method(api, session, app_model) + method(api, sqlite_session, app_model) - session.scalar.assert_called_once() - get_tenant_by_id.assert_called_once_with("tenant-1", session=session) + get_tenant_by_id.assert_called_once_with("tenant-1", session=sqlite_session) - def test_success(self, app: Flask) -> None: + @pytest.mark.parametrize("sqlite_session", [(Site,)], indirect=True) + def test_success( + self, + app: Flask, + sqlite_session: Session, + ) -> None: api = module.TrialSitApi() method = unwrap(api.get) - site = MagicMock() - app_model = SimpleNamespace(id="a1", tenant_id="tenant-1") + app_model = SimpleNamespace(id=str(uuid4()), tenant_id="tenant-1") tenant = SimpleNamespace(status=TenantStatus.NORMAL) - session = MagicMock() - session.scalar.return_value = site + site = _persist_site(sqlite_session, app_model.id) with ( app.test_request_context("/"), @@ -1185,15 +1434,15 @@ class TestTrialSitApi: mock_validate_result = MagicMock() mock_validate_result.model_dump.return_value = {"name": "test", "icon": "icon"} mock_validate.return_value = mock_validate_result - result = method(api, session, app_model) + result = method(api, sqlite_session, app_model) assert result == {"name": "test", "icon": "icon"} - session.scalar.assert_called_once() - get_tenant_by_id.assert_called_once_with("tenant-1", session=session) + get_tenant_by_id.assert_called_once_with("tenant-1", session=sqlite_session) + mock_validate.assert_called_once_with(site) class TestAppWorkflowApi: - def test_uses_injected_session(self) -> None: + def test_uses_injected_session(self, unbound_session: Session) -> None: api = module.AppWorkflowApi() method = unwrap(api.get) created_by = SimpleNamespace(id="account-1", name="Creator", email="creator@example.com") @@ -1207,7 +1456,18 @@ class TestAppWorkflowApi: marked_comment="", created_at=datetime(2024, 1, 1, tzinfo=UTC), updated_at=datetime(2024, 1, 2, tzinfo=UTC), - environment_variables=[], + environment_variables=[ + SecretVariable( + id="env-secret", + name="api_key", + value="plaintext-secret", + ), + LLMEnvironmentVariable( + id="env-llm", + name="shared_model", + value={"provider": "provider", "name": "model", "mode": "chat"}, + ), + ], conversation_variables=[ StringVariable( id="conversation-variable-1", @@ -1225,9 +1485,7 @@ class TestAppWorkflowApi: workflow_id="workflow-1", workflow_with_session=MagicMock(return_value=workflow), ) - session = MagicMock() - - result = method(api, session, app_model) + result = method(api, unbound_session, app_model) assert result == { "id": "workflow-1", @@ -1242,7 +1500,24 @@ class TestAppWorkflowApi: "updated_by": None, "updated_at": 1704153600, "tool_published": True, - "environment_variables": [], + "environment_variables": [ + { + "value_type": "secret", + "value": encrypter.full_mask_token(), + "id": "env-secret", + "name": "api_key", + "description": "", + "selector": [], + }, + { + "value_type": "llm", + "value": {"provider": "provider", "name": "model", "mode": "chat"}, + "id": "env-llm", + "name": "shared_model", + "description": "", + "selector": [], + }, + ], "conversation_variables": [ { "id": "conversation-variable-1", @@ -1254,10 +1529,10 @@ class TestAppWorkflowApi: ], "rag_pipeline_variables": [], } - app_model.workflow_with_session.assert_called_once_with(session=session) - workflow.get_created_by_account.assert_called_once_with(session=session) - workflow.get_updated_by_account.assert_called_once_with(session=session) - workflow.get_tool_published.assert_called_once_with(session=session) + app_model.workflow_with_session.assert_called_once_with(session=unbound_session) + workflow.get_created_by_account.assert_called_once_with(session=unbound_session) + workflow.get_updated_by_account.assert_called_once_with(session=unbound_session) + workflow.get_tool_published.assert_called_once_with(session=unbound_session) class TestTrialChatAudioApiExceptionHandlers: @@ -1278,7 +1553,11 @@ class TestTrialChatAudioApiExceptionHandlers: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_quota_exceeded(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatAudioApi() @@ -1297,7 +1576,11 @@ class TestTrialChatAudioApiExceptionHandlers: ), ): with pytest.raises(ProviderQuotaExceededError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) def test_invoke_error(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatAudioApi() @@ -1316,7 +1599,11 @@ class TestTrialChatAudioApiExceptionHandlers: ), ): with pytest.raises(CompletionRequestError): - method(api, account, trial_app_chat) + method( + api, + account, + trial_app_chat, + ) class TestTrialChatTextApiExceptionHandlers: @@ -1333,7 +1620,12 @@ class TestTrialChatTextApiExceptionHandlers: ), ): with pytest.raises(module.AppUnavailableError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) def test_unsupported_audio_type(self, app: Flask, trial_app_chat: MagicMock, account: Account) -> None: api = module.TrialChatTextApi() @@ -1348,4 +1640,9 @@ class TestTrialChatTextApiExceptionHandlers: ), ): with pytest.raises(module.UnsupportedAudioTypeError): - method(api, account, trial_app_chat) + method( + api, + TextToSpeechRequest.model_validate(request.get_json(silent=True) or {}), + account, + trial_app_chat, + ) diff --git a/api/tests/unit_tests/controllers/console/explore/test_workflow.py b/api/tests/unit_tests/controllers/console/explore/test_workflow.py index 8bcce0c22b2..247669fe125 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_workflow.py +++ b/api/tests/unit_tests/controllers/console/explore/test_workflow.py @@ -5,6 +5,7 @@ import pytest from flask import Flask from werkzeug.exceptions import InternalServerError +from controllers.common.controller_schemas import WorkflowRunPayload from controllers.console.explore.error import NotWorkflowAppError from controllers.console.explore.workflow import ( InstalledAppWorkflowRunApi, @@ -62,11 +63,18 @@ class TestInstalledAppWorkflowRunApi: with app.test_request_context("/"): with pytest.raises(NotWorkflowAppError): - method(api, MagicMock(), MagicMock(), non_workflow_installed_app) + method( + api, + WorkflowRunPayload.model_validate({"inputs": {}}), + MagicMock(), + MagicMock(), + non_workflow_installed_app, + ) def test_success(self, app: Flask, installed_workflow_app, user, payload): api = InstalledAppWorkflowRunApi() method = unwrap(api.post) + req_data = WorkflowRunPayload.model_validate(payload) with ( app.test_request_context("/", json=payload), @@ -75,7 +83,7 @@ class TestInstalledAppWorkflowRunApi: return_value=MagicMock(), ) as generate_mock, ): - result = method(api, MagicMock(), user, installed_workflow_app) + result = method(api, req_data, MagicMock(), user, installed_workflow_app) generate_mock.assert_called_once() assert generate_mock.call_args.kwargs["user"] is user @@ -84,6 +92,7 @@ class TestInstalledAppWorkflowRunApi: def test_rate_limit_error(self, app: Flask, installed_workflow_app, user, payload): api = InstalledAppWorkflowRunApi() method = unwrap(api.post) + req_data = WorkflowRunPayload.model_validate(payload) with ( app.test_request_context("/", json=payload), @@ -93,11 +102,12 @@ class TestInstalledAppWorkflowRunApi: ), ): with pytest.raises(InvokeRateLimitHttpError): - method(api, MagicMock(), user, installed_workflow_app) + method(api, req_data, MagicMock(), user, installed_workflow_app) def test_unexpected_exception(self, app: Flask, installed_workflow_app, user, payload): api = InstalledAppWorkflowRunApi() method = unwrap(api.post) + req_data = WorkflowRunPayload.model_validate(payload) with ( app.test_request_context("/", json=payload), @@ -107,7 +117,7 @@ class TestInstalledAppWorkflowRunApi: ), ): with pytest.raises(InternalServerError): - method(api, MagicMock(), user, installed_workflow_app) + method(api, req_data, MagicMock(), user, installed_workflow_app) class TestInstalledAppWorkflowTaskStopApi: diff --git a/api/tests/unit_tests/controllers/console/explore/test_wraps.py b/api/tests/unit_tests/controllers/console/explore/test_wraps.py index a60c13315b8..f2eb8523bbf 100644 --- a/api/tests/unit_tests/controllers/console/explore/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/explore/test_wraps.py @@ -255,11 +255,9 @@ def test_trial_feature_enable_disabled(): def view(): return "ok" - features = MagicMock(enable_trial_app=False) - with patch( - "controllers.console.explore.wraps.FeatureService.get_system_features", - return_value=features, + "controllers.console.explore.wraps.RecommendedAppService.is_trial_app_enabled", + return_value=False, ): with pytest.raises(Forbidden): view() @@ -270,11 +268,9 @@ def test_trial_feature_enable_enabled(): def view(): return "ok" - features = MagicMock(enable_trial_app=True) - with patch( - "controllers.console.explore.wraps.FeatureService.get_system_features", - return_value=features, + "controllers.console.explore.wraps.RecommendedAppService.is_trial_app_enabled", + return_value=True, ): assert view() == "ok" @@ -285,5 +281,9 @@ def test_installed_app_resource_decorators(): def test_trial_app_resource_decorators(): - decorators = TrialAppResource.method_decorators - assert len(decorators) == 3 + assert TrialAppResource.method_decorators == [ + trial_app_required, + trial_feature_enable, + wraps_module.account_initialization_required, + wraps_module.login_required, + ] diff --git a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py index 8a542ef269e..b5b2d79411b 100644 --- a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py +++ b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow.py @@ -132,7 +132,14 @@ def test_draft_workflow_post_returns_400_for_invalid_graph(app: Flask, monkeypat method="POST", json={"graph": {"nodes": [], "edges": []}, "hash": "hash-1"}, ): - response, status_code = handler(api, user, snippet) + response, status_code = handler( + api, + snippet_workflow_module.SnippetDraftSyncPayload.model_validate( + {"graph": {"nodes": [], "edges": []}, "hash": "hash-1"} + ), + user, + snippet, + ) assert status_code == 400 assert response == {"message": "invalid graph"} @@ -244,7 +251,11 @@ def test_list_published_snippet_workflows_includes_input_fields( handler = unwrap(api.get) with app.test_request_context("/snippets/snippet-1/workflows?page=1&limit=20"): - response = handler(api, snippet=snippet) + response = handler( + api, + snippet_workflow_module.SnippetWorkflowListQuery.model_validate({"page": 1, "limit": 20}), + snippet=snippet, + ) assert response["items"][0]["input_fields"] == input_fields @@ -411,7 +422,15 @@ def test_update_published_snippet_workflow_returns_updated_workflow( method="PATCH", json={"marked_name": "v1", "marked_comment": "first version"}, ): - response = handler(api, user, snippet, workflow_id="workflow-1") + response = handler( + api, + snippet_workflow_module.WorkflowUpdatePayload.model_validate( + {"marked_name": "v1", "marked_comment": "first version"} + ), + user, + snippet, + workflow_id="workflow-1", + ) update_workflow.assert_called_once() update_call = update_workflow.call_args.kwargs @@ -432,7 +451,13 @@ def test_update_published_snippet_workflow_returns_400_when_no_fields(app: Flask handler = unwrap(api.patch) with app.test_request_context("/snippets/snippet-1/workflows/workflow-1", method="PATCH", json={}): - response, status_code = handler(api, _account("account-1"), _snippet(), workflow_id="workflow-1") + response, status_code = handler( + api, + snippet_workflow_module.WorkflowUpdatePayload(), + _account("account-1"), + _snippet(), + workflow_id="workflow-1", + ) assert status_code == 400 assert response == {"message": "No valid fields to update"} @@ -468,7 +493,13 @@ def test_update_published_snippet_workflow_raises_not_found( json={"marked_name": "v1"}, ): with pytest.raises(NotFound, match="Workflow not found"): - handler(api, user, snippet, workflow_id="missing-workflow") + handler( + api, + snippet_workflow_module.WorkflowUpdatePayload.model_validate({"marked_name": "v1"}), + user, + snippet, + workflow_id="missing-workflow", + ) sqlite_session.refresh(snippet) assert snippet.name == "Snippet" diff --git a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow_draft_variable.py b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow_draft_variable.py index 0a9382f6d59..258c86e53bd 100644 --- a/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow_draft_variable.py +++ b/api/tests/unit_tests/controllers/console/snippets/test_snippet_workflow_draft_variable.py @@ -11,6 +11,7 @@ from sqlalchemy.orm import Session, scoped_session, sessionmaker from controllers.console.snippets import snippet_workflow_draft_variable as module from graphon.variables import StringSegment +from graphon.variables.types import SegmentType from models.account import Account, AccountStatus from models.workflow import WorkflowDraftVariable, WorkflowDraftVariableFile from services.workflow_draft_variable_service import WorkflowDraftVariableList @@ -254,6 +255,7 @@ def test_variable_patch_returns_persisted_variable_without_committing_when_no_ch with app.test_request_context("/", method="PATCH", json={}): result = handler( api, + module.WorkflowDraftVariableUpdatePayload(), _make_account(), snippet=SimpleNamespace(id="snippet-1", tenant_id="tenant-1"), variable_id="var-1", @@ -331,7 +333,7 @@ def test_environment_variables_returns_workflow_environment_variables( name="API_KEY", description="secret", selector=["env", "API_KEY"], - value_type=SimpleNamespace(exposed_type=Mock(return_value=SimpleNamespace(value="secret"))), + value_type=SegmentType.SECRET, value="sk-test", ) monkeypatch.setattr( diff --git a/api/tests/unit_tests/controllers/console/tag/test_tags.py b/api/tests/unit_tests/controllers/console/tag/test_tags.py index 7d267b50153..aa04caed752 100644 --- a/api/tests/unit_tests/controllers/console/tag/test_tags.py +++ b/api/tests/unit_tests/controllers/console/tag/test_tags.py @@ -11,10 +11,15 @@ from werkzeug.exceptions import Forbidden import controllers.console.tag.tags as module from controllers.console import console_ns from controllers.console.tag.tags import ( + TagBasePayload, TagBindingCollectionApi, + TagBindingPayload, TagBindingRemoveApi, + TagBindingRemovePayload, TagListApi, + TagListQueryParam, TagUpdateDeleteApi, + TagUpdateRequestPayload, ) from models import Account from models.account import AccountStatus, TenantAccountRole @@ -130,7 +135,7 @@ class TestTagListApi: ], ), ): - result, status = method(api, "tenant-1") + result, status = method(api, TagListQueryParam(type="knowledge"), "tenant-1") assert status == 200 assert result == [{"id": "1", "name": "tag", "type": "knowledge", "binding_count": "1"}] @@ -153,7 +158,7 @@ class TestTagListApi: ], ) as get_tags_mock, ): - result, status = method(api, "tenant-1") + result, status = method(api, TagListQueryParam(type="snippet"), "tenant-1") get_tags_mock.assert_called_once() assert get_tags_mock.call_args.args == ("snippet", "tenant-1", None) @@ -161,35 +166,35 @@ class TestTagListApi: assert status == 200 assert result == [{"id": "1", "name": "snippet-tag", "type": "snippet", "binding_count": "1"}] - def test_post_success(self, app: Flask, admin_user, tag, payload_patch): + def test_post_success(self, app: Flask, admin_user, tag): api = TagListApi() method = unwrap(api.post) payload = {"name": "test-tag", "type": "knowledge"} + req_data = TagBasePayload.model_validate(payload) with app.test_request_context("/", json=payload): with ( - payload_patch(payload), patch( "controllers.console.tag.tags.TagService.save_tags", return_value=tag, ), ): - result, status = method(api, admin_user) + result, status = method(api, req_data, admin_user) assert status == 200 assert result["name"] == "test-tag" assert result["binding_count"] == "0" - def test_post_snippet_tag_checks_snippet_rbac_when_enabled(self, app: Flask, admin_user, tag, payload_patch): + def test_post_snippet_tag_checks_snippet_rbac_when_enabled(self, app: Flask, admin_user, tag): api = TagListApi() method = unwrap(api.post) payload = {"name": "snippet-tag", "type": "snippet"} + req_data = TagBasePayload.model_validate(payload) with app.test_request_context("/", json=payload): with ( - payload_patch(payload), patch("controllers.console.tag.tags.dify_config.RBAC_ENABLED", True), patch( "controllers.console.tag.tags.current_account_with_tenant", @@ -201,7 +206,7 @@ class TestTagListApi: return_value=tag, ), ): - method(api, admin_user) + method(api, req_data, admin_user) enforce_mock.assert_called_once_with( tenant_id="tenant-1", @@ -211,30 +216,25 @@ class TestTagListApi: resource_required=False, ) - def test_post_forbidden(self, app: Flask, readonly_user, payload_patch): + def test_post_forbidden(self, app: Flask, readonly_user): api = TagListApi() method = unwrap(api.post) - payload = {"name": "x"} - - with app.test_request_context("/", json=payload): - with ( - payload_patch(payload), - ): - with pytest.raises(Forbidden): - method(api, readonly_user) + with app.test_request_context("/"): + with pytest.raises(Forbidden): + method(api, TagBasePayload(name="test", type=TagType.KNOWLEDGE), readonly_user) class TestTagUpdateDeleteApi: - def test_patch_success(self, app: Flask, admin_user, tag, payload_patch, sqlite_engine: Engine): + def test_patch_success(self, app: Flask, admin_user, tag, sqlite_engine: Engine): api = TagUpdateDeleteApi() method = unwrap(api.patch) payload = {"name": "updated"} + req_data = TagUpdateRequestPayload.model_validate(payload) with app.test_request_context("/", json=payload): with ( - payload_patch(payload), patch( "controllers.console.tag.tags.TagService.update_tags", return_value=tag, @@ -244,7 +244,7 @@ class TestTagUpdateDeleteApi: return_value=3, ), ): - result, status = method(api, admin_user, "tag-1") + result, status = method(api, req_data, admin_user, "tag-1") assert status == 200 update_payload, tag_id, session = update_tags_mock.call_args.args @@ -253,18 +253,13 @@ class TestTagUpdateDeleteApi: _assert_sqlite_session(session, sqlite_engine) assert result["binding_count"] == "3" - def test_patch_forbidden(self, app: Flask, readonly_user, payload_patch): + def test_patch_forbidden(self, app: Flask, readonly_user): api = TagUpdateDeleteApi() method = unwrap(api.patch) - payload = {"name": "x"} - - with app.test_request_context("/", json=payload): - with ( - payload_patch(payload), - ): - with pytest.raises(Forbidden): - method(api, readonly_user, "tag-1") + with app.test_request_context("/"): + with pytest.raises(Forbidden): + method(api, TagUpdateRequestPayload(name="test"), readonly_user, "tag-1") def test_delete_success(self, app: Flask, admin_user, sqlite_engine: Engine): api = TagUpdateDeleteApi() @@ -383,7 +378,7 @@ class TestTagBindingCollectionApi: payload_patch(payload), patch("controllers.console.tag.tags.TagService.save_tag_binding") as save_mock, ): - result, status = method(api, admin_user) + result, status = method(api, TagBindingPayload.model_validate(payload), admin_user) save_mock.assert_called_once() assert status == 200 @@ -404,7 +399,7 @@ class TestTagBindingCollectionApi: payload_patch(payload), patch("controllers.console.tag.tags.TagService.save_tag_binding") as save_mock, ): - result, status = method(api, admin_user) + result, status = method(api, TagBindingPayload.model_validate(payload), admin_user) save_mock.assert_called_once() binding_payload = save_mock.call_args.args[0] @@ -422,7 +417,11 @@ class TestTagBindingCollectionApi: payload_patch({}), ): with pytest.raises(Forbidden): - method(api, readonly_user) + method( + api, + TagBindingPayload(tag_ids=["tag-1"], target_id="target-1", type=TagType.KNOWLEDGE), + readonly_user, + ) class TestTagBindingRemoveApi: @@ -441,7 +440,7 @@ class TestTagBindingRemoveApi: payload_patch(payload), patch("controllers.console.tag.tags.TagService.delete_tag_binding") as delete_mock, ): - result, status = method(api, admin_user) + result, status = method(api, TagBindingRemovePayload.model_validate(payload), admin_user) delete_mock.assert_called_once() delete_payload = delete_mock.call_args.args[0] @@ -458,7 +457,11 @@ class TestTagBindingRemoveApi: payload_patch({}), ): with pytest.raises(Forbidden): - method(api, readonly_user) + method( + api, + TagBindingRemovePayload(tag_ids=["tag-1"], target_id="target-1", type=TagType.KNOWLEDGE), + readonly_user, + ) class TestTagResponseModel: diff --git a/api/tests/unit_tests/controllers/console/test_apikey.py b/api/tests/unit_tests/controllers/console/test_apikey.py index 0435ad0996a..5693c102a5e 100644 --- a/api/tests/unit_tests/controllers/console/test_apikey.py +++ b/api/tests/unit_tests/controllers/console/test_apikey.py @@ -2,18 +2,31 @@ from __future__ import annotations import inspect from collections.abc import Callable -from types import SimpleNamespace from typing import cast from unittest.mock import MagicMock, patch +from uuid import UUID import pytest -from werkzeug.exceptions import Forbidden +from flask import Flask +from sqlalchemy import event, select +from sqlalchemy.orm import Session +from werkzeug.exceptions import BadRequest, Forbidden, NotFound -from controllers.console.apikey import BaseApiKeyListResource, BaseApiKeyResource +from configs import dify_config +from controllers.console.agent.roster import AgentApiKeyListApi +from controllers.console.apikey import ( + AppApiKeyListResource, + BaseApiKeyListResource, + BaseApiKeyResource, +) +from controllers.console.datasets.datasets import DatasetApiKeyApi +from core.rbac import RBACPermission, RBACResourceScope +from enums import DeploymentEdition from models import Account from models.account import AccountStatus, TenantAccountRole from models.enums import ApiTokenType -from models.model import ApiToken, App +from models.model import ApiToken, App, AppMode, IconType +from services.agent.errors import AgentAccessNotReadyError def _make_list_resource() -> BaseApiKeyListResource: @@ -44,80 +57,130 @@ def _make_account(role: TenantAccountRole) -> Account: return account -def test_list_api_keys_uses_injected_session_and_tenant_id() -> None: +def _persist_app(session: Session, *, mode: AppMode = AppMode.CHAT) -> App: + app = App( + id="app-1", + tenant_id="tenant-1", + name="API key app", + mode=mode, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#ffffff", + enable_site=False, + enable_api=True, + ) + session.add(app) + session.flush() + return app + + +def test_list_api_keys_uses_injected_session_and_tenant_id(sqlite_session: Session) -> None: resource = _make_list_resource() raw_get = cast( Callable[[BaseApiKeyListResource, object, str, str], dict[str, object]], inspect.unwrap(BaseApiKeyListResource.get), ) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() - api_key = SimpleNamespace( - id="key-1", + session = sqlite_session + _persist_app(session) + api_key = ApiToken( type=ApiTokenType.APP, token="app-token", - last_used_at=None, - created_at=None, + app_id="app-1", + tenant_id="tenant-1", ) - session.scalars.return_value.all.return_value = [api_key] + api_key.id = "key-1" + session.add(api_key) + session.add( + ApiToken( + type=ApiTokenType.APP, + token="foreign-app-token", + app_id="app-1", + tenant_id="tenant-2", + ) + ) + legacy_api_key = ApiToken(type=ApiTokenType.APP, token="legacy-app-token", app_id="app-1", tenant_id=None) + session.add(legacy_api_key) + session.commit() result = raw_get(resource, session, "app-1", "tenant-1") + data = cast(list[dict[str, object]], result["data"]) - session.execute.assert_called_once() - session.scalars.assert_called_once() - assert result == { - "data": [ - { - "id": "key-1", - "type": "app", - "token": "app-token", - "last_used_at": None, - "created_at": None, - } - ] - } + assert {item["token"] for item in data} == {"app-token", "legacy-app-token"} -def test_create_api_key_uses_injected_session_and_tenant_id() -> None: +def test_create_api_key_uses_injected_session_and_tenant_id(sqlite_session: Session) -> None: resource = _make_list_resource() raw_post = cast( Callable[[BaseApiKeyListResource, object, str, str], tuple[dict[str, object], int]], inspect.unwrap(BaseApiKeyListResource.post), ) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() - session.scalar.return_value = 0 - - def add_api_token(api_token: ApiToken) -> None: - api_token.id = "key-1" + session = sqlite_session + _persist_app(session) + session.add_all( + [ + ApiToken(type=ApiTokenType.APP, token=f"foreign-token-{index}", app_id="app-1", tenant_id="tenant-2") + for index in range(resource.max_keys) + ] + ) + session.commit() + commits: list[str] = [] + event.listen(session, "after_commit", lambda _session: commits.append("commit")) with patch( "controllers.console.apikey.ApiToken.generate_api_key", return_value="app-generated-token" ) as generate_api_key: - session.add.side_effect = add_api_token - result, status = raw_post(resource, session, "app-1", "tenant-1") assert status == 201 assert result["token"] == "app-generated-token" - api_token = session.add.call_args.args[0] + api_token = session.scalar(select(ApiToken).where(ApiToken.token == "app-generated-token")) + assert api_token is not None assert api_token.app_id == "app-1" assert api_token.tenant_id == "tenant-1" assert api_token.type == ApiTokenType.APP generate_api_key.assert_called_once_with("app-", 24, session=session) - session.execute.assert_called_once() - session.scalar.assert_called_once() - session.commit.assert_called_once() + assert commits == ["commit"] -def test_delete_api_key_rejects_non_admin_account() -> None: +def test_create_api_key_counts_legacy_tokens(sqlite_session: Session) -> None: + resource = _make_list_resource() + _persist_app(sqlite_session) + sqlite_session.add_all( + [ + ApiToken(type=ApiTokenType.APP, token=f"legacy-token-{index}", app_id="app-1", tenant_id=None) + for index in range(resource.max_keys) + ] + ) + sqlite_session.commit() + + with pytest.raises(BadRequest): + resource._create_api_key("app-1", "tenant-1", session=sqlite_session) + + +def test_create_agent_api_key_requires_published_access(sqlite_session: Session) -> None: + resource = _make_list_resource() + session = sqlite_session + app = _persist_app(session, mode=AppMode.AGENT) + + with patch( + "controllers.console.apikey.AppService.ensure_agent_app_access_ready", + side_effect=AgentAccessNotReadyError(), + ) as ensure_access_ready: + with pytest.raises(AgentAccessNotReadyError): + resource._create_api_key("app-1", "tenant-1", session=session) + + ensure_access_ready.assert_called_once_with(app, session=session) + assert session.scalar(select(ApiToken)) is None + + +def test_delete_api_key_rejects_non_admin_account(sqlite_session: Session) -> None: resource = _make_key_resource() raw_delete = cast( Callable[[BaseApiKeyResource, object, str, str, str, Account], tuple[str, int]], inspect.unwrap(BaseApiKeyResource.delete), ) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() + session = sqlite_session + _persist_app(session) with pytest.raises(Forbidden): raw_delete( @@ -129,20 +192,21 @@ def test_delete_api_key_rejects_non_admin_account() -> None: _make_account(TenantAccountRole.NORMAL), ) - session.execute.assert_called_once() - session.scalar.assert_not_called() - -def test_delete_api_key_uses_injected_session_user_and_tenant() -> None: +def test_delete_api_key_uses_injected_session_user_and_tenant(sqlite_session: Session) -> None: resource = _make_key_resource() raw_delete = cast( Callable[[BaseApiKeyResource, object, str, str, str, Account], tuple[str, int]], inspect.unwrap(BaseApiKeyResource.delete), ) - api_key = SimpleNamespace(token="app-token", type=ApiTokenType.APP) - session = MagicMock() - session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace() - session.scalar.return_value = api_key + session = sqlite_session + _persist_app(session) + api_key = ApiToken(type=ApiTokenType.APP, token="app-token", app_id="app-1", tenant_id=None) + api_key.id = "key-1" + session.add(api_key) + session.commit() + commits: list[str] = [] + event.listen(session, "after_commit", lambda _session: commits.append("commit")) with patch("controllers.console.apikey.ApiTokenCache.delete") as delete_cache: result, status = raw_delete( @@ -155,8 +219,105 @@ def test_delete_api_key_uses_injected_session_user_and_tenant() -> None: ) delete_cache.assert_called_once_with("app-token", ApiTokenType.APP) - assert session.execute.call_count == 2 - session.scalar.assert_called_once() - session.commit.assert_called_once() + assert session.get(ApiToken, "key-1") is None + assert commits == ["commit"] assert result == "" assert status == 204 + + +def test_delete_api_key_rejects_foreign_tenant_token(sqlite_session: Session) -> None: + resource = _make_key_resource() + session = sqlite_session + _persist_app(session) + api_key = ApiToken(type=ApiTokenType.APP, token="foreign-token", app_id="app-1", tenant_id="tenant-2") + api_key.id = "key-1" + session.add(api_key) + session.commit() + + with patch("controllers.console.apikey.ApiTokenCache.delete") as delete_cache: + with pytest.raises(NotFound): + resource._delete_api_key( + "app-1", + "key-1", + "tenant-1", + _make_account(TenantAccountRole.OWNER), + session=session, + ) + + delete_cache.assert_not_called() + assert session.get(ApiToken, "key-1") is api_key + + +def test_api_key_lists_require_matching_rbac_permission() -> None: + app = Flask(__name__) + account = _make_account(TenantAccountRole.OWNER) + api_id = UUID("00000000-0000-0000-0000-000000000001") + cases = [ + ( + lambda: AppApiKeyListResource().get(resource_id=api_id), + [(RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION, True)], + ), + ( + lambda: AgentApiKeyListApi().get(agent_id=api_id), + [ + (RBACResourceScope.WORKSPACE, RBACPermission.AGENT_MANAGE, False), + (RBACResourceScope.APP, RBACPermission.APP_RELEASE_AND_VERSION, True), + ], + ), + ( + lambda: DatasetApiKeyApi().get(), + [(RBACResourceScope.DATASET, RBACPermission.DATASET_API_KEY_MANAGE, False)], + ), + ] + + with ( + app.test_request_context("/"), + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + patch.object(dify_config, "LOGIN_DISABLED", True), + patch.object(dify_config, "RBAC_ENABLED", True), + patch("controllers.console.wraps.current_account_with_tenant", return_value=(account, "tenant-1")), + patch("controllers.common.wraps.current_account_with_tenant", return_value=(account, "tenant-1")), + patch.object(BaseApiKeyListResource, "_get_api_key_list") as get_api_key_list, + ): + for invoke, expected_gates in cases: + with patch( + "controllers.common.wraps.enforce_rbac_access", + side_effect=[None] * (len(expected_gates) - 1) + [Forbidden()], + ) as enforce_rbac_access: + with pytest.raises(Forbidden): + invoke() + + assert [ + (kwargs["resource_type"], kwargs["scene"], kwargs["resource_required"]) + for _, kwargs in enforce_rbac_access.call_args_list + ] == expected_gates + + get_api_key_list.assert_not_called() + + +def test_api_key_lists_reject_legacy_read_only_members() -> None: + app = Flask(__name__) + account = _make_account(TenantAccountRole.NORMAL) + api_id = UUID("00000000-0000-0000-0000-000000000001") + current_user = MagicMock() + current_user._get_current_object.return_value = account + current_user.has_edit_permission = False + + with ( + app.test_request_context("/"), + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + patch.object(dify_config, "LOGIN_DISABLED", True), + patch.object(dify_config, "RBAC_ENABLED", False), + patch("libs.login.current_user", current_user), + patch("controllers.console.wraps.current_account_with_tenant", return_value=(account, "tenant-1")), + patch.object(BaseApiKeyListResource, "_get_api_key_list") as get_api_key_list, + ): + for invoke in ( + lambda: AppApiKeyListResource().get(resource_id=api_id), + lambda: AgentApiKeyListApi().get(agent_id=api_id), + lambda: DatasetApiKeyApi().get(), + ): + with pytest.raises(Forbidden): + invoke() + + get_api_key_list.assert_not_called() diff --git a/api/tests/unit_tests/controllers/console/test_extension.py b/api/tests/unit_tests/controllers/console/test_extension.py index 8ea327dfdce..5d38d982e55 100644 --- a/api/tests/unit_tests/controllers/console/test_extension.py +++ b/api/tests/unit_tests/controllers/console/test_extension.py @@ -20,6 +20,7 @@ from controllers.console.extension import ( APIBasedExtensionDetailAPI, CodeBasedExtensionAPI, ) +from enums import DeploymentEdition if _NEEDS_METHOD_VIEW_CLEANUP: del builtins.__dict__["MethodView"] @@ -62,9 +63,9 @@ def _mock_console_guards(monkeypatch: pytest.MonkeyPatch) -> MagicMock: account.id = "account-123" account.is_authenticated = True - monkeypatch.setattr(wraps_module.dify_config, "EDITION", "CLOUD") + monkeypatch.setattr(wraps_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) + monkeypatch.setattr(wraps_module.dify_config, "INIT_PASSWORD", "") monkeypatch.setattr("libs.login.dify_config.LOGIN_DISABLED", True) - monkeypatch.delenv("INIT_PASSWORD", raising=False) monkeypatch.setattr(wraps_module, "current_account_with_tenant", lambda: (account, "tenant-123")) # The login_required decorator consults the shared LocalProxy in libs.login. diff --git a/api/tests/unit_tests/controllers/console/test_fastopenapi_init_validate.py b/api/tests/unit_tests/controllers/console/test_fastopenapi_init_validate.py new file mode 100644 index 00000000000..1d73ee30a68 --- /dev/null +++ b/api/tests/unit_tests/controllers/console/test_fastopenapi_init_validate.py @@ -0,0 +1,149 @@ +"""HTTP contract tests for the FastOpenAPI initialization routes.""" + +from types import SimpleNamespace +from unittest.mock import Mock, create_autospec + +import pytest + +from controllers.console import init_validate +from dify_app import DifyApp +from enums import DeploymentEdition +from extensions import ext_fastopenapi +from services.init_validation_service import ( + AlreadyInitializedError, + InitValidationService, + InvalidInitializationPasswordError, +) + + +@pytest.fixture +def init_validation(monkeypatch: pytest.MonkeyPatch) -> Mock: + service = create_autospec(InitValidationService, instance=True, spec_set=True) + services = SimpleNamespace(init_validation=service) + monkeypatch.setattr(init_validate, "application_services", lambda: services) + return service + + +@pytest.fixture +def app() -> DifyApp: + app = DifyApp(__name__) + app.config.update(TESTING=True, SECRET_KEY="test-secret") + ext_fastopenapi.init_app(app) + return app + + +@pytest.mark.parametrize( + ("validated", "expected_status"), + [ + pytest.param(True, "finished", id="finished"), + pytest.param(False, "not_started", id="not-started"), + ], +) +def test_get_init_status( + app: DifyApp, + init_validation: Mock, + validated: bool, + expected_status: str, +) -> None: + init_validation.is_validated.return_value = validated + + response = app.test_client().get("/console/api/init") + + assert response.status_code == 200 + assert response.get_json() == {"status": expected_status} + init_validation.is_validated.assert_called_once_with(session_validated=False) + + +def test_validate_init_password_success( + app: DifyApp, + init_validation: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + "controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ) + client = app.test_client() + + response = client.post("/console/api/init", json={"password": "expected"}) + + assert response.status_code == 201 + assert response.get_json() == {"result": "success"} + init_validation.validate_password.assert_called_once_with("expected") + with client.session_transaction() as browser_session: + assert browser_session["is_init_validated"] is True + + +def test_validate_init_password_rejects_a_mismatch( + app: DifyApp, + init_validation: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + "controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ) + init_validation.validate_password.side_effect = InvalidInitializationPasswordError + client = app.test_client() + + response = client.post("/console/api/init", json={"password": "wrong"}) + + assert response.status_code == 401 + with client.session_transaction() as browser_session: + assert browser_session["is_init_validated"] is False + + +def test_validate_init_password_rejects_an_initialized_installation( + app: DifyApp, + init_validation: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + "controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ) + init_validation.validate_password.side_effect = AlreadyInitializedError + + response = app.test_client().post("/console/api/init", json={"password": "expected"}) + + assert response.status_code == 403 + + +@pytest.mark.parametrize( + "payload", + [ + pytest.param({}, id="missing-password"), + pytest.param({"password": "x" * 31}, id="password-too-long"), + ], +) +def test_validate_init_password_rejects_an_invalid_payload( + app: DifyApp, + init_validation: Mock, + monkeypatch: pytest.MonkeyPatch, + payload: dict[str, str], +) -> None: + monkeypatch.setattr( + "controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ) + + response = app.test_client().post("/console/api/init", json=payload) + + assert response.status_code == 422 + init_validation.validate_password.assert_not_called() + + +def test_validate_init_password_is_not_available_in_cloud( + app: DifyApp, + init_validation: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + "controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.CLOUD, + ) + + response = app.test_client().post("/console/api/init", json={"password": "expected"}) + + assert response.status_code == 404 + init_validation.validate_password.assert_not_called() diff --git a/api/tests/unit_tests/controllers/console/test_fastopenapi_ping.py b/api/tests/unit_tests/controllers/console/test_fastopenapi_ping.py deleted file mode 100644 index fc04ca078b7..00000000000 --- a/api/tests/unit_tests/controllers/console/test_fastopenapi_ping.py +++ /dev/null @@ -1,27 +0,0 @@ -import builtins - -import pytest -from flask import Flask -from flask.views import MethodView - -from extensions import ext_fastopenapi - -if not hasattr(builtins, "MethodView"): - builtins.MethodView = MethodView # type: ignore[attr-defined] - - -@pytest.fixture -def app() -> Flask: - app = Flask(__name__) - app.config["TESTING"] = True - return app - - -def test_console_ping_fastopenapi_returns_pong(app: Flask): - ext_fastopenapi.init_app(app) - - client = app.test_client() - response = client.get("/console/api/ping") - - assert response.status_code == 200 - assert response.get_json() == {"result": "pong"} diff --git a/api/tests/unit_tests/controllers/console/test_fastopenapi_setup.py b/api/tests/unit_tests/controllers/console/test_fastopenapi_setup.py index 2b385304d32..8018ed3c115 100644 --- a/api/tests/unit_tests/controllers/console/test_fastopenapi_setup.py +++ b/api/tests/unit_tests/controllers/console/test_fastopenapi_setup.py @@ -1,40 +1,88 @@ import builtins -from unittest.mock import patch +from datetime import datetime +from types import SimpleNamespace +from unittest.mock import Mock, create_autospec import pytest -from flask import Flask from flask.views import MethodView +from controllers.console import setup as setup_controller +from controllers.console import wraps +from controllers.console.error import AlreadySetupError, NotInitValidateError +from dify_app import DifyApp +from enums import DeploymentEdition from extensions import ext_fastopenapi +from services.setup_service import ( + InitializationValidationRequiredError, + SetupAlreadyCompletedError, + SetupInput, + SetupService, + SetupStatus, +) if not hasattr(builtins, "MethodView"): builtins.MethodView = MethodView # type: ignore[attr-defined] @pytest.fixture -def app() -> Flask: - app = Flask(__name__) +def setup_service(monkeypatch: pytest.MonkeyPatch) -> Mock: + service = create_autospec(SetupService, instance=True, spec_set=True) + services = SimpleNamespace(setup=service) + monkeypatch.setattr(setup_controller, "application_services", lambda: services) + return service + + +@pytest.fixture +def app() -> DifyApp: + app = DifyApp(__name__) app.config["TESTING"] = True + ext_fastopenapi.init_app(app) return app -def test_console_setup_fastopenapi_get_not_started(app: Flask): - ext_fastopenapi.init_app(app) +def test_console_setup_fastopenapi_get_not_started(app: DifyApp, setup_service: Mock) -> None: + setup_service.get_status.return_value = SetupStatus(completed=False) - with ( - patch("controllers.console.setup.dify_config.EDITION", "SELF_HOSTED"), - patch("controllers.console.setup.get_setup_status", return_value=None), - ): - client = app.test_client() - response = client.get("/console/api/setup") + response = app.test_client().get("/console/api/setup") assert response.status_code == 200 assert response.get_json() == {"step": "not_started", "setup_at": None} -def test_console_setup_fastopenapi_post_success(app: Flask): - ext_fastopenapi.init_app(app) +def test_console_setup_fastopenapi_get_finished(app: DifyApp, setup_service: Mock) -> None: + setup_at = datetime(2026, 8, 6, 10, 30) + setup_service.get_status.return_value = SetupStatus(completed=True, setup_at=setup_at) + response = app.test_client().get("/console/api/setup") + + assert response.status_code == 200 + assert response.get_json() == {"step": "finished", "setup_at": "2026-08-06T10:30:00"} + + +def test_console_setup_fastopenapi_get_finished_without_setup_time(app: DifyApp, setup_service: Mock) -> None: + setup_service.get_status.return_value = SetupStatus(completed=True) + + response = app.test_client().get("/console/api/setup") + + assert response.status_code == 200 + assert response.get_json() == {"step": "finished", "setup_at": None} + + +@pytest.mark.parametrize( + "deployment_edition", + [DeploymentEdition.COMMUNITY, DeploymentEdition.ENTERPRISE], + ids=["community", "enterprise"], +) +def test_console_setup_fastopenapi_post_success( + app: DifyApp, + setup_service: Mock, + monkeypatch: pytest.MonkeyPatch, + deployment_edition: DeploymentEdition, +) -> None: + monkeypatch.setattr(wraps.dify_config, "DEPLOYMENT_EDITION", deployment_edition) + monkeypatch.setattr(setup_controller, "is_init_validated", lambda: True) + mark_setup_completed = Mock() + monkeypatch.setattr(setup_controller, "mark_setup_completed", mark_setup_completed) payload = { "email": "admin@example.com", "name": "Admin", @@ -42,17 +90,158 @@ def test_console_setup_fastopenapi_post_success(app: Flask): "language": "en-US", } - with ( - patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"), - patch("controllers.console.setup.get_setup_status", return_value=None), - patch("controllers.console.setup.TenantService.get_tenant_count", return_value=0), - patch("controllers.console.setup.get_init_validate_status", return_value=True), - patch("controllers.console.setup.RegisterService.setup"), - patch("controllers.console.setup.mark_setup_completed") as mark_setup_completed, - ): - client = app.test_client() - response = client.post("/console/api/setup", json=payload) + response = app.test_client().post( + "/console/api/setup", + json=payload, + headers={"CF-Connecting-IP": "203.0.113.7"}, + ) assert response.status_code == 201 assert response.get_json() == {"result": "success"} + setup_service.initialize.assert_called_once_with( + SetupInput( + email="admin@example.com", + name="Admin", + password="Passw0rd1", + ip_address="203.0.113.7", + language="en-US", + ), + initialization_validated=True, + ) mark_setup_completed.assert_called_once_with() + + +def test_console_setup_fastopenapi_post_rejects_cloud_edition( + app: DifyApp, + setup_service: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) + + response = app.test_client().post( + "/console/api/setup", + json={ + "email": "admin@example.com", + "name": "Admin", + "password": "Passw0rd1", + "language": "en-US", + }, + ) + + assert response.status_code == 404 + setup_service.initialize.assert_not_called() + + +@pytest.mark.parametrize( + "payload", + [ + pytest.param( + { + "email": "not-an-email", + "name": "Admin", + "password": "Passw0rd1", + "language": "en-US", + }, + id="invalid-email", + ), + pytest.param( + { + "email": "admin@example.com", + "name": "Admin", + "password": "short", + "language": "en-US", + }, + id="invalid-password", + ), + pytest.param( + { + "email": "admin@example.com", + "name": "a" * 31, + "password": "Passw0rd1", + "language": "en-US", + }, + id="name-too-long", + ), + ], +) +def test_console_setup_fastopenapi_post_rejects_invalid_payload_before_service_call( + app: DifyApp, + setup_service: Mock, + monkeypatch: pytest.MonkeyPatch, + payload: dict[str, str], +) -> None: + monkeypatch.setattr(wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + + response = app.test_client().post("/console/api/setup", json=payload) + + assert response.status_code == 422 + setup_service.initialize.assert_not_called() + + +@pytest.mark.parametrize( + ("service_error", "expected_controller_error"), + [ + pytest.param(SetupAlreadyCompletedError(), AlreadySetupError, id="already-setup"), + pytest.param( + InitializationValidationRequiredError(), + NotInitValidateError, + id="init-validation-required", + ), + ], +) +def test_console_setup_translates_service_errors_to_controller_errors( + app: DifyApp, + setup_service: Mock, + monkeypatch: pytest.MonkeyPatch, + service_error: Exception, + expected_controller_error: type[Exception], +) -> None: + monkeypatch.setattr(wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + monkeypatch.setattr(setup_controller, "is_init_validated", lambda: False) + mark_setup_completed = Mock() + monkeypatch.setattr(setup_controller, "mark_setup_completed", mark_setup_completed) + setup_service.initialize.side_effect = service_error + + payload = setup_controller.SetupRequestPayload.model_validate( + { + "email": "admin@example.com", + "name": "Admin", + "password": "Passw0rd1", + "language": "en-US", + } + ) + with app.test_request_context( + "/console/api/setup", + method="POST", + headers={"CF-Connecting-IP": "203.0.113.7"}, + ): + with pytest.raises(expected_controller_error) as raised: + setup_controller.setup_system(payload) + + assert type(raised.value) is expected_controller_error + mark_setup_completed.assert_not_called() + + +def test_console_setup_fastopenapi_does_not_mark_setup_completed_when_service_fails( + app: DifyApp, + setup_service: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + monkeypatch.setattr(setup_controller, "is_init_validated", lambda: True) + mark_setup_completed = Mock() + monkeypatch.setattr(setup_controller, "mark_setup_completed", mark_setup_completed) + setup_service.initialize.side_effect = RuntimeError("provision failed") + + response = app.test_client().post( + "/console/api/setup", + json={ + "email": "admin@example.com", + "name": "Admin", + "password": "Passw0rd1", + "language": "en-US", + }, + ) + + assert response.status_code == 500 + mark_setup_completed.assert_not_called() diff --git a/api/tests/unit_tests/controllers/console/test_fastopenapi_system.py b/api/tests/unit_tests/controllers/console/test_fastopenapi_system.py new file mode 100644 index 00000000000..08aa94e6ca3 --- /dev/null +++ b/api/tests/unit_tests/controllers/console/test_fastopenapi_system.py @@ -0,0 +1,41 @@ +import builtins +from unittest.mock import patch + +import pytest +from flask.views import MethodView + +from configs import dify_config +from dify_app import DifyApp +from extensions import ext_fastopenapi + +if not hasattr(builtins, "MethodView"): + builtins.MethodView = MethodView # type: ignore[attr-defined] + + +@pytest.fixture +def app() -> DifyApp: + app = DifyApp(__name__) + app.config["TESTING"] = True + return app + + +def test_console_ping_fastopenapi_returns_pong(app: DifyApp) -> None: + ext_fastopenapi.init_app(app) + + response = app.test_client().get("/console/api/ping") + + assert response.status_code == 200 + assert response.get_json() == {"result": "pong"} + + +def test_console_version_fastopenapi_returns_current_version(app: DifyApp) -> None: + ext_fastopenapi.init_app(app) + + with patch("controllers.console.system.dify_config.CHECK_UPDATE_URL", None): + response = app.test_client().get("/console/api/version", query_string={"current_version": "0.0.0"}) + + assert response.status_code == 200 + assert response.get_json() == { + "version": dify_config.project.version, + "release_notes": "", + } diff --git a/api/tests/unit_tests/controllers/console/test_fastopenapi_version.py b/api/tests/unit_tests/controllers/console/test_fastopenapi_version.py deleted file mode 100644 index c5b4e0dfcf4..00000000000 --- a/api/tests/unit_tests/controllers/console/test_fastopenapi_version.py +++ /dev/null @@ -1,35 +0,0 @@ -import builtins -from unittest.mock import patch - -import pytest -from flask import Flask -from flask.views import MethodView - -from configs import dify_config -from extensions import ext_fastopenapi - -if not hasattr(builtins, "MethodView"): - builtins.MethodView = MethodView # type: ignore[attr-defined] - - -@pytest.fixture -def app() -> Flask: - app = Flask(__name__) - app.config["TESTING"] = True - return app - - -def test_console_version_fastopenapi_returns_current_version(app: Flask): - ext_fastopenapi.init_app(app) - - with patch("controllers.console.version.dify_config.CHECK_UPDATE_URL", None): - client = app.test_client() - response = client.get("/console/api/version", query_string={"current_version": "0.0.0"}) - - assert response.status_code == 200 - data = response.get_json() - assert data["version"] == dify_config.project.version - assert data["release_date"] == "" - assert data["release_notes"] == "" - assert data["can_auto_update"] is False - assert "features" in data diff --git a/api/tests/unit_tests/controllers/console/test_feature.py b/api/tests/unit_tests/controllers/console/test_feature.py index 19e30a1b08d..f8e7d94b496 100644 --- a/api/tests/unit_tests/controllers/console/test_feature.py +++ b/api/tests/unit_tests/controllers/console/test_feature.py @@ -1,16 +1,51 @@ from inspect import unwrap +from unittest.mock import create_autospec from pytest_mock import MockerFixture -from enums.deployment_edition import DeploymentEdition -from services.feature_service import ( +from enums import DeploymentEdition +from extensions.ext_application_services import ApplicationServices +from machinery.context import RequestContext +from services.entities.feature_entities import ( FeatureModel, LicenseLimitationModel, LicenseModel, LicenseStatus, LimitationModel, SystemFeatureModel, + VectorSpaceLimitationModel, ) +from services.explore_banner_query_service import ExploreBannerQueryService +from services.feature_query_service import FeatureQueryService +from services.init_validation_service import InitValidationService +from services.schema_definition_service import SchemaDefinitionService +from services.setup_service import SetupService +from services.workspace_member_query_service import WorkspaceMemberQueryService +from services.workspace_query_service import WorkspaceQueryService + + +def _request_context() -> RequestContext: + return RequestContext( + request_id="request_123", + trace_id=None, + account_id="account_123", + active_workspace_id="tenant_123", + ) + + +def _install_application_services(mocker: MockerFixture): + feature_queries = create_autospec(FeatureQueryService, instance=True, spec_set=True) + services = ApplicationServices( + explore_banner_queries=create_autospec(ExploreBannerQueryService, instance=True, spec_set=True), + schema_definitions=create_autospec(SchemaDefinitionService, instance=True, spec_set=True), + setup=create_autospec(SetupService, instance=True, spec_set=True), + feature_queries=feature_queries, + init_validation=create_autospec(InitValidationService, instance=True, spec_set=True), + workspace_queries=create_autospec(WorkspaceQueryService, instance=True, spec_set=True), + workspace_member_queries=create_autospec(WorkspaceMemberQueryService, instance=True, spec_set=True), + ) + mocker.patch("controllers.console.feature.application_services", return_value=services) + return feature_queries class TestFeatureApi: @@ -21,47 +56,72 @@ class TestFeatureApi: knowledge_rate_limit=42, vector_space=LimitationModel(size=1, limit=2), ) - get_features = mocker.patch("controllers.console.feature.FeatureService.get_features") + feature_queries = _install_application_services(mocker) + get_features = feature_queries.get_features get_features.return_value = features api = FeatureApi() raw_get = unwrap(FeatureApi.get) - result = raw_get(api, "tenant_123") + request_context = _request_context() + result = raw_get(api, request_context) expected = features.model_dump() expected.pop("vector_space") assert result == expected - get_features.assert_called_once_with("tenant_123", exclude_vector_space=True) + get_features.assert_called_once_with(request_context) class TestFeatureVectorSpaceApi: def test_get_vector_space_success(self, mocker: MockerFixture): from controllers.console.feature import FeatureVectorSpaceApi - get_vector_space = mocker.patch("controllers.console.feature.FeatureService.get_vector_space") - get_vector_space.return_value = LimitationModel(size=5120, limit=20480) + feature_queries = _install_application_services(mocker) + get_vector_space = feature_queries.get_vector_space + get_vector_space.return_value = VectorSpaceLimitationModel(size=5120, limit=20480) api = FeatureVectorSpaceApi() raw_get = unwrap(FeatureVectorSpaceApi.get) - result = raw_get(api, "tenant_123") + request_context = _request_context() + result = raw_get(api, request_context) assert result == {"size": 5120, "limit": 20480} - get_vector_space.assert_called_once_with("tenant_123") + get_vector_space.assert_called_once_with(request_context) + + def test_get_vector_space_preserves_unknown_usage(self, mocker: MockerFixture): + from controllers.console.feature import FeatureVectorSpaceApi + + feature_queries = _install_application_services(mocker) + get_vector_space = feature_queries.get_vector_space + get_vector_space.return_value = VectorSpaceLimitationModel(size=0, limit=50, usage_unknown=True) + + request_context = _request_context() + result = unwrap(FeatureVectorSpaceApi.get)(FeatureVectorSpaceApi(), request_context) + + assert result == {"size": 0, "limit": 50, "usage_unknown": True} + get_vector_space.assert_called_once_with(request_context) + + def test_vector_space_response_schema_marks_usage_unknown_optional(self): + schema = VectorSpaceLimitationModel.model_json_schema(mode="serialization") + + assert schema["required"] == ["size", "limit"] + assert schema["properties"]["usage_unknown"]["type"] == "boolean" + assert "usage_unknown" not in schema["required"] class TestTrialModelsApi: def test_get_trial_models_success(self, mocker: MockerFixture): from controllers.console.feature import TrialModelsApi - get_trial_models = mocker.patch("controllers.console.feature.FeatureService.get_trial_models") + feature_queries = _install_application_services(mocker) + get_trial_models = feature_queries.get_trial_models get_trial_models.return_value = ["langgenius/openai/openai"] api = TrialModelsApi() raw_get = unwrap(TrialModelsApi.get) - result = raw_get(api) + result = raw_get(api, _request_context()) assert result == {"trial_models": ["langgenius/openai/openai"]} get_trial_models.assert_called_once_with() @@ -71,7 +131,8 @@ class TestAppDslVersionApi: def test_get_app_dsl_version_success(self, mocker: MockerFixture): from controllers.console.feature import AppDslVersionApi - get_app_dsl_version = mocker.patch("controllers.console.feature.FeatureService.get_app_dsl_version") + feature_queries = _install_application_services(mocker) + get_app_dsl_version = feature_queries.get_app_dsl_version get_app_dsl_version.return_value = "0.6.0" api = AppDslVersionApi() @@ -93,10 +154,9 @@ class TestSystemFeatureApi: is_allow_register=True, enable_learn_app=True, ) - get_system_features = mocker.patch( - "controllers.console.feature.FeatureService.get_system_features", - return_value=system_features, - ) + feature_queries = _install_application_services(mocker) + get_system_features = feature_queries.get_system_features + get_system_features.return_value = system_features api = SystemFeatureApi() result = api.get() @@ -105,6 +165,8 @@ class TestSystemFeatureApi: assert result["is_allow_register"] is True assert result["enable_learn_app"] is True assert result["license"] == {"status": LicenseStatus.NONE} + assert result["sso_enforced_for_signin_protocol"] is None + assert result["webapp_auth"]["sso_config"]["protocol"] is None get_system_features.assert_called_once_with() @@ -117,14 +179,13 @@ class TestSystemFeatureLicenseApi: expired_at="2025-12-31", seats=LicenseLimitationModel(enabled=True, limit=5, size=2), ) - get_license = mocker.patch( - "controllers.console.feature.FeatureService.get_license", - return_value=license_model, - ) + feature_queries = _install_application_services(mocker) + get_license = feature_queries.get_license + get_license.return_value = license_model api = SystemFeatureLicenseApi() raw_get = unwrap(SystemFeatureLicenseApi.get) - result = raw_get(api) + result = raw_get(api, _request_context()) assert result == license_model.model_dump() assert result["seats"] == {"enabled": True, "limit": 5, "size": 2} diff --git a/api/tests/unit_tests/controllers/console/test_files.py b/api/tests/unit_tests/controllers/console/test_files.py index c36b6395291..f894e04f481 100644 --- a/api/tests/unit_tests/controllers/console/test_files.py +++ b/api/tests/unit_tests/controllers/console/test_files.py @@ -5,6 +5,7 @@ import pytest from flask import Flask from werkzeug.exceptions import Forbidden +from configs import dify_config from constants import DOCUMENT_EXTENSIONS from controllers.common.errors import ( BlockedFileExtensionError, @@ -86,12 +87,21 @@ class TestFileApiGet: api = FileApi() get_method = unwrap(api.get) - with app.test_request_context(): - data, status = get_method(api) + with ( + app.test_request_context(), + patch( + "controllers.console.files.FeatureService.get_knowledge_file_size_limit", + return_value=50, + ) as get_knowledge_file_size_limit, + ): + data, status = get_method(api, "tenant-1") assert status == 200 assert "file_size_limit" in data + assert data["knowledge_file_size_limit"] == 50 assert "batch_count_limit" in data + get_knowledge_file_size_limit.assert_called_once_with("tenant-1") + assert data["skill_file_size_limit"] == dify_config.UPLOAD_SKILL_FILE_SIZE_LIMIT class TestFileApiPost: @@ -198,6 +208,33 @@ class TestFileApiPost: assert result is upload_file assert mock_file_service.upload_file.call_args.kwargs["tenant_id"] == "app-tenant-id" + def test_dataset_source_from_query_uses_knowledge_limit( + self, + app: Flask, + mock_account_context, + mock_file_service, + ): + upload_file = MagicMock() + mock_file_service.upload_file.return_value = upload_file + + with ( + app.test_request_context( + "/?source=datasets", + method="POST", + data={"file": (io.BytesIO(b"hello"), "test.txt")}, + ), + patch( + "controllers.console.files.FeatureService.get_knowledge_file_size_limit", + return_value=50, + ) as get_knowledge_file_size_limit, + ): + result = upload_file_from_request(current_user=mock_account_context) + + assert result is upload_file + assert mock_file_service.upload_file.call_args.kwargs["source"] == "datasets" + assert mock_file_service.upload_file.call_args.kwargs["default_file_size_limit"] == 50 + get_knowledge_file_size_limit.assert_called_once_with(mock_account_context.current_tenant_id) + def test_upload_with_invalid_source(self, app: Flask, mock_account_context, mock_file_service): """Test that invalid source parameter gets normalized to None""" api = FileApi() diff --git a/api/tests/unit_tests/controllers/console/test_human_input_form.py b/api/tests/unit_tests/controllers/console/test_human_input_form.py index a9e847d7d7c..e40d85c331a 100644 --- a/api/tests/unit_tests/controllers/console/test_human_input_form.py +++ b/api/tests/unit_tests/controllers/console/test_human_input_form.py @@ -4,7 +4,7 @@ import json from datetime import UTC, datetime from inspect import unwrap from types import SimpleNamespace -from unittest.mock import Mock +from unittest.mock import ANY, Mock import pytest from flask import Flask, Response @@ -18,6 +18,7 @@ from controllers.console.human_input_form import ( WorkflowResponseConverter, _jsonify_form_definition, ) +from core.workflow.human_input_policy import HumanInputSurface from models.account import AccountStatus from models.enums import CreatorUserRole from models.human_input import RecipientType @@ -344,3 +345,62 @@ def test_workflow_events_finished(app: Flask, monkeypatch: pytest.MonkeyPatch) - assert response.mimetype == "text/event-stream" assert "data" in response.get_data(as_text=True) + + +def test_workflow_events_snapshot_can_continue_across_pauses(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + workflow_run = SimpleNamespace( + id="run-1", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + tenant_id="t1", + app_id="app-1", + finished_at=None, + ) + app_model = SimpleNamespace(mode=AppMode.WORKFLOW) + + class _RepoStub: + def get_workflow_run_by_id_and_tenant_id(self, **_kwargs): + return workflow_run + + workflow_generator = Mock() + workflow_generator.convert_to_event_stream.return_value = iter(["data: snapshot\n\n"]) + snapshot_builder = Mock(return_value=["snapshot-events"]) + + monkeypatch.setattr( + DifyAPIRepositoryFactory, + "create_api_workflow_run_repository", + lambda *_args, **_kwargs: _RepoStub(), + ) + monkeypatch.setattr( + "controllers.console.human_input_form._retrieve_app_for_workflow_run", + lambda *_args, **_kwargs: app_model, + ) + monkeypatch.setattr( + "controllers.console.human_input_form.WorkflowAppGenerator", + lambda: workflow_generator, + ) + monkeypatch.setattr( + "controllers.console.human_input_form.build_workflow_event_stream", + snapshot_builder, + ) + monkeypatch.setattr("controllers.console.human_input_form.db", SimpleNamespace(engine=object())) + + api = ConsoleWorkflowEventsApi() + handler = unwrap(api.get) + + with app.test_request_context( + "/console/api/workflow/run-1/events?include_state_snapshot=true&continue_on_pause=true", + method="GET", + ): + response = handler(api, "t1", SimpleNamespace(id="user-1"), workflow_run_id="run-1") + + assert response.get_data(as_text=True) == "data: snapshot\n\n" + snapshot_builder.assert_called_once_with( + app_mode=AppMode.WORKFLOW, + workflow_run=workflow_run, + tenant_id="t1", + app_id="app-1", + session_maker=ANY, + human_input_surface=HumanInputSurface.CONSOLE, + close_on_pause=False, + ) diff --git a/api/tests/unit_tests/controllers/console/test_init_validate.py b/api/tests/unit_tests/controllers/console/test_init_validate.py index 377135e3f2f..80145f7cae6 100644 --- a/api/tests/unit_tests/controllers/console/test_init_validate.py +++ b/api/tests/unit_tests/controllers/console/test_init_validate.py @@ -1,33 +1,56 @@ -"""Initialization validation tests with real setup-state persistence in SQLite.""" - -from __future__ import annotations +"""Tests for the Flask adapter around initialization validation.""" from types import SimpleNamespace +from unittest.mock import Mock, create_autospec import pytest from flask import Flask -from sqlalchemy.orm import Session -from controllers.console import init_validate +from controllers.console import init_validate, wraps from controllers.console.error import AlreadySetupError, InitValidateFailedError -from models.model import DifySetup +from enums import DeploymentEdition +from services.init_validation_service import ( + AlreadyInitializedError, + InitValidationService, + InvalidInitializationPasswordError, +) -def test_get_init_status_finished(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(init_validate, "get_init_validate_status", lambda: True) - result = init_validate.get_init_status() - assert result.status == "finished" +@pytest.fixture +def init_validation(monkeypatch: pytest.MonkeyPatch) -> Mock: + service = create_autospec(InitValidationService, instance=True, spec_set=True) + application_services = SimpleNamespace(init_validation=service) + monkeypatch.setattr(init_validate, "application_services", lambda: application_services) + return service -def test_get_init_status_not_started(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(init_validate, "get_init_validate_status", lambda: False) - result = init_validate.get_init_status() - assert result.status == "not_started" +def test_get_init_status_finished(app: Flask, init_validation: Mock) -> None: + init_validation.is_validated.return_value = True + app.secret_key = "test-secret" + + with app.test_request_context("/console/api/init", method="GET"): + result = init_validate.get_init_status() + + assert result.status == "finished" -def test_validate_init_password_already_setup(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED") - monkeypatch.setattr(init_validate.TenantService, "get_tenant_count", lambda *, session: 1) +def test_get_init_status_not_started(app: Flask, init_validation: Mock) -> None: + init_validation.is_validated.return_value = False + app.secret_key = "test-secret" + + with app.test_request_context("/console/api/init", method="GET"): + result = init_validate.get_init_status() + + assert result.status == "not_started" + + +def test_validate_init_password_already_setup( + app: Flask, + init_validation: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + init_validation.validate_password.side_effect = AlreadyInitializedError app.secret_key = "test-secret" with app.test_request_context("/console/api/init", method="POST"): @@ -35,10 +58,13 @@ def test_validate_init_password_already_setup(app: Flask, monkeypatch: pytest.Mo init_validate.validate_init_password(init_validate.InitValidatePayload(password="pw")) -def test_validate_init_password_wrong_password(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED") - monkeypatch.setattr(init_validate.TenantService, "get_tenant_count", lambda *, session: 0) - monkeypatch.setenv("INIT_PASSWORD", "expected") +def test_validate_init_password_wrong_password( + app: Flask, + init_validation: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + init_validation.validate_password.side_effect = InvalidInitializationPasswordError app.secret_key = "test-secret" with app.test_request_context("/console/api/init", method="POST"): @@ -47,58 +73,42 @@ def test_validate_init_password_wrong_password(app: Flask, monkeypatch: pytest.M assert init_validate.session.get("is_init_validated") is False -def test_validate_init_password_success(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED") - monkeypatch.setattr(init_validate.TenantService, "get_tenant_count", lambda *, session: 0) - monkeypatch.setenv("INIT_PASSWORD", "expected") +def test_validate_init_password_success( + app: Flask, + init_validation: Mock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(wraps.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) app.secret_key = "test-secret" with app.test_request_context("/console/api/init", method="POST"): result = init_validate.validate_init_password(init_validate.InitValidatePayload(password="expected")) + assert result.result == "success" assert init_validate.session.get("is_init_validated") is True + init_validation.validate_password.assert_called_once_with("expected") -def test_get_init_validate_status_not_self_hosted(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(init_validate.dify_config, "EDITION", "CLOUD") - assert init_validate.get_init_validate_status() is True - - -def test_get_init_validate_status_validated_session(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED") - monkeypatch.setenv("INIT_PASSWORD", "expected") - app.secret_key = "test-secret" - - with app.test_request_context("/console/api/init", method="GET"): - init_validate.session["is_init_validated"] = True - assert init_validate.get_init_validate_status() is True - - -@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) -def test_get_init_validate_status_setup_exists( - app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +@pytest.mark.parametrize( + ("session_value", "expected"), + [ + pytest.param(None, False, id="missing"), + pytest.param(False, False, id="not-validated"), + pytest.param(True, True, id="validated"), + ], +) +def test_is_init_validated_passes_session_state( + app: Flask, + init_validation: Mock, + session_value: bool | None, + expected: bool, ) -> None: - monkeypatch.setattr(init_validate.dify_config, "EDITION", "SELF_HOSTED") - monkeypatch.setenv("INIT_PASSWORD", "expected") - monkeypatch.setattr(init_validate, "db", SimpleNamespace(engine=sqlite_session.get_bind())) - sqlite_session.add(DifySetup(version="test-version")) - sqlite_session.commit() + init_validation.is_validated.return_value = True app.secret_key = "test-secret" with app.test_request_context("/console/api/init", method="GET"): - init_validate.session.pop("is_init_validated", None) - assert init_validate.get_init_validate_status() is True + if session_value is not None: + init_validate.session["is_init_validated"] = session_value - -@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, "db", SimpleNamespace(engine=sqlite_session.get_bind())) - app.secret_key = "test-secret" - - with app.test_request_context("/console/api/init", method="GET"): - init_validate.session.pop("is_init_validated", None) - assert init_validate.get_init_validate_status() is False + assert init_validate.is_init_validated() is True + init_validation.is_validated.assert_called_once_with(session_validated=expected) diff --git a/api/tests/unit_tests/controllers/console/test_onboarding.py b/api/tests/unit_tests/controllers/console/test_onboarding.py index c2afd7f7422..8d613f7c202 100644 --- a/api/tests/unit_tests/controllers/console/test_onboarding.py +++ b/api/tests/unit_tests/controllers/console/test_onboarding.py @@ -2,13 +2,12 @@ from __future__ import annotations from datetime import UTC, datetime from inspect import unwrap -from unittest.mock import Mock, PropertyMock, patch +from unittest.mock import Mock import pytest from flask import Flask from pydantic import ValidationError -from controllers.console import console_ns from controllers.console.onboarding import ( StepByStepTourStateApi, StepByStepTourStatePatchPayload, @@ -69,13 +68,13 @@ def test_patch_step_by_step_tour_state_passes_action_payload( method = unwrap(api.patch) payload = {"action": "complete_task", "task_id": "studio"} + req_data = StepByStepTourStatePatchPayload.model_validate(payload) with app.test_request_context( "/console/api/onboarding/step-by-step-tour/state", method="PATCH", json=payload, ): - with patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload): - result = method(api, "workspace-1", _account()) + result = method(api, req_data, "workspace-1", _account()) assert result["completed_task_ids"] == ["home"] patch_state.assert_called_once() diff --git a/api/tests/unit_tests/controllers/console/test_spec.py b/api/tests/unit_tests/controllers/console/test_spec.py index 44fb8345928..9e6083e950f 100644 --- a/api/tests/unit_tests/controllers/console/test_spec.py +++ b/api/tests/unit_tests/controllers/console/test_spec.py @@ -1,15 +1,24 @@ from inspect import unwrap -from unittest.mock import patch - -import pytest +from types import SimpleNamespace +from unittest.mock import create_autospec, patch import controllers.console.spec as spec_module +from dify_app import DifyApp +from extensions import ext_login +from machinery.context import RequestContext +from services.schema_definition_service import SchemaDefinitionService class TestSpecSchemaDefinitionsApi: - def test_get_success(self): + def test_get_success(self) -> None: api = spec_module.SpecSchemaDefinitionsApi() method = unwrap(api.get) + request_context = RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id="account-1", + active_workspace_id="workspace-1", + ) schema_definitions = [ { @@ -23,34 +32,65 @@ class TestSpecSchemaDefinitionsApi: } ] + service = create_autospec(SchemaDefinitionService, instance=True, spec_set=True) + service.list.return_value = tuple(schema_definitions) + with patch.object( spec_module, - "SchemaManager", - ) as schema_manager_cls: - schema_manager_cls.return_value.get_all_schema_definitions.return_value = schema_definitions + "application_services", + return_value=SimpleNamespace(schema_definitions=service), + ): + resp, status = method(api, request_context) - resp, status = method(api) - - assert status == 200 + assert status == spec_module.HTTPStatus.OK assert resp == schema_definitions - assert spec_module.SchemaDefinitionsResponse.model_validate(resp).model_dump(mode="json") == schema_definitions + service.list.assert_called_once_with() - def test_get_documents_tight_response_model(self): + def test_get_documents_tight_response_model(self) -> None: response = spec_module.SpecSchemaDefinitionsApi.get.__apidoc__["responses"]["200"] assert response[1].name == spec_module.SchemaDefinitionsResponse.__name__ - def test_get_exception_returns_empty_list(self, caplog: pytest.LogCaptureFixture): + def test_get_returns_empty_list_from_service(self) -> None: api = spec_module.SpecSchemaDefinitionsApi() method = unwrap(api.get) + request_context = RequestContext( + request_id="request-1", + trace_id=None, + account_id="account-1", + active_workspace_id=None, + ) + service = create_autospec(SchemaDefinitionService, instance=True, spec_set=True) + service.list.return_value = () with patch.object( spec_module, - "SchemaManager", - side_effect=Exception("boom"), + "application_services", + return_value=SimpleNamespace(schema_definitions=service), ): - resp, status = method(api) + resp, status = method(api, request_context) - assert status == 200 + assert status == spec_module.HTTPStatus.OK assert resp == [] - assert "boom" in caplog.text + + def test_get_rejects_unauthenticated_request_before_service_call(self) -> None: + app = DifyApp(__name__) + app.config["TESTING"] = True + ext_login.init_app(app) + api = spec_module.SpecSchemaDefinitionsApi() + service = create_autospec(SchemaDefinitionService, instance=True, spec_set=True) + + with ( + app.test_request_context("/console/api/spec/schema-definitions"), + patch("controllers.console.wraps._is_setup_completed", return_value=True), + patch("libs.login._resolve_current_user", return_value=None), + patch.object( + spec_module, + "application_services", + return_value=SimpleNamespace(schema_definitions=service), + ), + ): + response = api.get() + + assert response.status_code == 401 + service.list.assert_not_called() diff --git a/api/tests/unit_tests/controllers/console/test_system.py b/api/tests/unit_tests/controllers/console/test_system.py new file mode 100644 index 00000000000..3c390003eb1 --- /dev/null +++ b/api/tests/unit_tests/controllers/console/test_system.py @@ -0,0 +1,134 @@ +import logging +from unittest.mock import MagicMock, patch + +import pytest + +import controllers.console.system as system_module + + +class TestHasNewVersion: + def test_has_new_version_true(self) -> None: + result = system_module._has_new_version( + latest_version="1.2.0", + current_version="1.1.0", + ) + assert result is True + + def test_has_new_version_false(self) -> None: + result = system_module._has_new_version( + latest_version="1.0.0", + current_version="1.1.0", + ) + assert result is False + + def test_has_new_version_invalid_version(self, caplog: pytest.LogCaptureFixture) -> None: + with caplog.at_level(logging.WARNING, logger="controllers.console.system"): + result = system_module._has_new_version( + latest_version="invalid", + current_version="1.0.0", + ) + + assert result is False + assert "Invalid version format" in caplog.text + + +class TestCheckVersionUpdate: + def test_no_check_update_url(self) -> None: + query = system_module.VersionQuery(current_version="1.0.0") + + with ( + patch.object( + system_module.dify_config, + "CHECK_UPDATE_URL", + "", + ), + patch.object( + system_module.dify_config.project, + "version", + "1.0.0", + ), + ): + result = system_module.check_version_update(query) + + assert result == system_module.VersionResponse(version="1.0.0", release_notes="") + + def test_http_error_fallback(self, caplog: pytest.LogCaptureFixture) -> None: + query = system_module.VersionQuery(current_version="1.0.0") + + with ( + patch.object( + system_module.dify_config, + "CHECK_UPDATE_URL", + "http://example.com", + ), + patch.object( + system_module.httpx, + "get", + side_effect=Exception("boom"), + ), + caplog.at_level(logging.WARNING, logger="controllers.console.system"), + ): + result = system_module.check_version_update(query) + + assert result.version == "1.0.0" + assert "Check update version error" in caplog.text + + def test_new_version_available(self) -> None: + query = system_module.VersionQuery(current_version="1.0.0") + + response = MagicMock() + response.json.return_value = { + "version": "1.2.0", + "releaseNotes": "New features", + } + + with ( + patch.object( + system_module.dify_config, + "CHECK_UPDATE_URL", + "http://example.com", + ), + patch.object( + system_module.httpx, + "get", + return_value=response, + ), + patch.object( + system_module.dify_config.project, + "version", + "1.0.0", + ), + ): + result = system_module.check_version_update(query) + + assert result.version == "1.2.0" + assert result.release_notes == "New features" + + def test_no_new_version(self) -> None: + query = system_module.VersionQuery(current_version="1.2.0") + + response = MagicMock() + response.json.return_value = { + "version": "1.1.0", + } + + with ( + patch.object( + system_module.dify_config, + "CHECK_UPDATE_URL", + "http://example.com", + ), + patch.object( + system_module.httpx, + "get", + return_value=response, + ), + patch.object( + system_module.dify_config.project, + "version", + "1.2.0", + ), + ): + result = system_module.check_version_update(query) + + assert result == system_module.VersionResponse(version="1.2.0", release_notes="") diff --git a/api/tests/unit_tests/controllers/console/test_version.py b/api/tests/unit_tests/controllers/console/test_version.py deleted file mode 100644 index c17217b0187..00000000000 --- a/api/tests/unit_tests/controllers/console/test_version.py +++ /dev/null @@ -1,162 +0,0 @@ -import logging -from unittest.mock import MagicMock, patch - -import pytest - -import controllers.console.version as version_module - - -class TestHasNewVersion: - def test_has_new_version_true(self): - result = version_module._has_new_version( - latest_version="1.2.0", - current_version="1.1.0", - ) - assert result is True - - def test_has_new_version_false(self): - result = version_module._has_new_version( - latest_version="1.0.0", - current_version="1.1.0", - ) - assert result is False - - def test_has_new_version_invalid_version(self, caplog: pytest.LogCaptureFixture): - with caplog.at_level(logging.WARNING, logger="controllers.console.version"): - result = version_module._has_new_version( - latest_version="invalid", - current_version="1.0.0", - ) - - assert result is False - assert "Invalid version format" in caplog.text - - -class TestCheckVersionUpdate: - def test_no_check_update_url(self): - query = version_module.VersionQuery(current_version="1.0.0") - - with ( - patch.object( - version_module.dify_config, - "CHECK_UPDATE_URL", - "", - ), - patch.object( - version_module.dify_config.project, - "version", - "1.0.0", - ), - patch.object( - version_module.dify_config, - "CAN_REPLACE_LOGO", - True, - ), - patch.object( - version_module.dify_config, - "MODEL_LB_ENABLED", - False, - ), - ): - result = version_module.check_version_update(query) - - assert result.version == "1.0.0" - assert result.can_auto_update is False - assert result.features.can_replace_logo is True - assert result.features.model_load_balancing_enabled is False - - def test_http_error_fallback(self, caplog: pytest.LogCaptureFixture): - query = version_module.VersionQuery(current_version="1.0.0") - - with ( - patch.object( - version_module.dify_config, - "CHECK_UPDATE_URL", - "http://example.com", - ), - patch.object( - version_module.httpx, - "get", - side_effect=Exception("boom"), - ), - caplog.at_level(logging.WARNING, logger="controllers.console.version"), - ): - result = version_module.check_version_update(query) - - assert result.version == "1.0.0" - assert "Check update version error" in caplog.text - - def test_new_version_available(self): - query = version_module.VersionQuery(current_version="1.0.0") - - response = MagicMock() - response.json.return_value = { - "version": "1.2.0", - "releaseDate": "2024-01-01", - "releaseNotes": "New features", - "canAutoUpdate": True, - } - - with ( - patch.object( - version_module.dify_config, - "CHECK_UPDATE_URL", - "http://example.com", - ), - patch.object( - version_module.httpx, - "get", - return_value=response, - ), - patch.object( - version_module.dify_config.project, - "version", - "1.0.0", - ), - patch.object( - version_module.dify_config, - "CAN_REPLACE_LOGO", - False, - ), - patch.object( - version_module.dify_config, - "MODEL_LB_ENABLED", - True, - ), - ): - result = version_module.check_version_update(query) - - assert result.version == "1.2.0" - assert result.release_date == "2024-01-01" - assert result.release_notes == "New features" - assert result.can_auto_update is True - - def test_no_new_version(self): - query = version_module.VersionQuery(current_version="1.2.0") - - response = MagicMock() - response.json.return_value = { - "version": "1.1.0", - } - - with ( - patch.object( - version_module.dify_config, - "CHECK_UPDATE_URL", - "http://example.com", - ), - patch.object( - version_module.httpx, - "get", - return_value=response, - ), - patch.object( - version_module.dify_config.project, - "version", - "1.2.0", - ), - ): - result = version_module.check_version_update(query) - - assert result.version == "1.2.0" - assert result.can_auto_update is False diff --git a/api/tests/unit_tests/controllers/console/test_workflow_run_archive.py b/api/tests/unit_tests/controllers/console/test_workflow_run_archive.py index bff41cb97b1..552a3570900 100644 --- a/api/tests/unit_tests/controllers/console/test_workflow_run_archive.py +++ b/api/tests/unit_tests/controllers/console/test_workflow_run_archive.py @@ -54,7 +54,7 @@ def test_current_owner_or_admin_ids_returns_current_ids_for_manager( ("method", "args"), [ (WorkflowRunArchivesApi.get, ()), - (WorkflowRunArchiveDownloadsApi.post, ()), + (WorkflowRunArchiveDownloadsApi.post, (None,)), (WorkflowRunArchiveDownloadApi.get, ("download-1",)), (WorkflowRunArchiveDownloadFileApi.get, ("download-1",)), ], @@ -91,7 +91,6 @@ def test_workflow_run_archive_endpoints_require_cloud_paid_plan(method) -> None: assert { "only_edition_cloud", - "cloud_edition_billing_enabled", "cloud_edition_billing_paid_plan_required", } <= decorator_names assert "rbac_permission_required" not in decorator_names diff --git a/api/tests/unit_tests/controllers/console/test_workspace_account.py b/api/tests/unit_tests/controllers/console/test_workspace_account.py index 40a3f06ad29..ce10cc9e070 100644 --- a/api/tests/unit_tests/controllers/console/test_workspace_account.py +++ b/api/tests/unit_tests/controllers/console/test_workspace_account.py @@ -8,12 +8,14 @@ import pytest from flask import Flask from sqlalchemy.orm import Session, scoped_session, sessionmaker +from controllers.console.error import EducationDiscountTemporarilyPausedError from controllers.console.workspace.account import ( AccountDeleteUpdateFeedbackApi, ChangeEmailCheckApi, ChangeEmailResetApi, ChangeEmailSendEmailApi, CheckEmailUnique, + EducationApi, ) from models import Account, AccountIntegrate, AccountStatus, Tenant, TenantAccountJoin from models.account import TenantAccountRole @@ -106,6 +108,27 @@ def _build_change_email_token( raise AssertionError(f"Unsupported phase for test helper: {phase}") +class TestEducationApi: + @patch("controllers.console.workspace.account.BillingService.EducationIdentity.activate") + def test_post_returns_temporarily_paused_error_without_activating_discount( + self, mock_activate: MagicMock, app: Flask + ): + account = _build_account("student@example.edu") + + with app.test_request_context("/account/education", method="POST", json={}): + api = EducationApi() + method = inspect.unwrap(api.post) + with pytest.raises(EducationDiscountTemporarilyPausedError) as exc_info: + method(api, account) + + assert exc_info.value.data == { + "code": "education_discount_temporarily_paused", + "message": "Education discount temporarily paused, while we upgrade our security measures.", + "status": 503, + } + mock_activate.assert_not_called() + + class TestChangeEmailSend: @patch("controllers.console.workspace.account.AccountService.send_change_email_email") @patch("controllers.console.workspace.account.AccountService.is_email_send_ip_limit", return_value=False) diff --git a/api/tests/unit_tests/controllers/console/test_workspace_members.py b/api/tests/unit_tests/controllers/console/test_workspace_members.py index deadff06c94..cb7df1644b5 100644 --- a/api/tests/unit_tests/controllers/console/test_workspace_members.py +++ b/api/tests/unit_tests/controllers/console/test_workspace_members.py @@ -7,6 +7,7 @@ from flask import Flask, g from controllers.console.workspace.error import InvalidMemberRoleError from controllers.console.workspace.members import MemberInviteEmailApi +from enums import DeploymentEdition from models.account import Account, TenantAccountRole @@ -53,8 +54,7 @@ class TestMemberInviteEmailApi: patch("controllers.console.workspace.members.dify_config.RBAC_ENABLED", False), patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "https://console.example.com"), patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), ): with app.test_request_context( "/workspaces/current/members/invite-email", @@ -81,14 +81,12 @@ class TestMemberInviteEmailApi: @patch("controllers.console.workspace.members.FeatureService.get_features") @patch("controllers.console.workspace.members.RegisterService.invite_new_member") - @patch("controllers.console.workspace.members.current_account_with_tenant") @patch("controllers.console.wraps.db") @patch("libs.login.check_csrf_token", return_value=None) def test_invite_rbac_enabled_accepts_rbac_role_id( self, mock_csrf, mock_db, - mock_current_account, mock_invite_member, mock_get_features, app, @@ -98,8 +96,6 @@ class TestMemberInviteEmailApi: mock_invite_member.return_value = "rbac-token" tenant = SimpleNamespace(id="tenant-1", name="Test Tenant") - inviter = SimpleNamespace(email="inviter@example.com", current_tenant=tenant, status="active") - mock_current_account.return_value = (inviter, tenant.id) with patch("controllers.console.workspace.members.dify_config") as mock_config: mock_config.RBAC_ENABLED = True @@ -121,14 +117,12 @@ class TestMemberInviteEmailApi: assert call_args.kwargs["role"] == "rbac-role-id-abc" @patch("controllers.console.workspace.members.FeatureService.get_features") - @patch("controllers.console.workspace.members.current_account_with_tenant") @patch("controllers.console.wraps.db") @patch("libs.login.check_csrf_token", return_value=None) def test_invite_rbac_disabled_rejects_invalid_role( self, mock_csrf, mock_db, - mock_current_account, mock_get_features, app, ): @@ -136,8 +130,6 @@ class TestMemberInviteEmailApi: mock_get_features.return_value = _build_feature_flags() tenant = SimpleNamespace(id="tenant-1", name="Test Tenant") - inviter = SimpleNamespace(email="inviter@example.com", current_tenant=tenant, status="active") - mock_current_account.return_value = (inviter, tenant.id) with patch("controllers.console.workspace.members.dify_config") as mock_config: mock_config.RBAC_ENABLED = False @@ -158,14 +150,12 @@ class TestMemberInviteEmailApi: assert exc_info.value.data == {"code": "invalid_role", "message": "Invalid role.", "status": 400} @patch("controllers.console.workspace.members.FeatureService.get_features") - @patch("controllers.console.workspace.members.current_account_with_tenant") @patch("controllers.console.wraps.db") @patch("libs.login.check_csrf_token", return_value=None) def test_invite_rbac_disabled_rejects_owner_role( self, mock_csrf, mock_db, - mock_current_account, mock_get_features, app, ): @@ -173,8 +163,6 @@ class TestMemberInviteEmailApi: mock_get_features.return_value = _build_feature_flags() tenant = SimpleNamespace(id="tenant-1", name="Test Tenant") - inviter = SimpleNamespace(email="inviter@example.com", current_tenant=tenant, status="active") - mock_current_account.return_value = (inviter, tenant.id) with patch("controllers.console.workspace.members.dify_config") as mock_config: mock_config.RBAC_ENABLED = False diff --git a/api/tests/unit_tests/controllers/console/test_wraps.py b/api/tests/unit_tests/controllers/console/test_wraps.py index 3c19bba2a93..a1432098416 100644 --- a/api/tests/unit_tests/controllers/console/test_wraps.py +++ b/api/tests/unit_tests/controllers/console/test_wraps.py @@ -11,6 +11,7 @@ from sqlalchemy.orm import Session from werkzeug.exceptions import HTTPException from controllers.common.wraps import _extract_resource_id +from controllers.console import flask_admission from controllers.console.error import NotInitValidateError, NotSetupError, UnauthorizedAndForceLogout from controllers.console.workspace.error import AccountNotInitializedError from controllers.console.wraps import ( @@ -18,7 +19,6 @@ from controllers.console.wraps import ( RBACResourceScope, _is_setup_completed, account_initialization_required, - cloud_edition_billing_enabled, cloud_edition_billing_paid_plan_required, cloud_edition_billing_rate_limit_check, cloud_edition_billing_resource_check, @@ -35,10 +35,13 @@ from controllers.console.wraps import ( with_current_user, with_current_user_id, ) +from enums import DeploymentEdition +from libs.login import AccountWithTenant +from machinery.context import RequestContext from models import Account from models.account import AccountStatus, TenantAccountRole -from models.dataset import RateLimitLog -from services.feature_service import LicenseStatus +from models.dataset import Dataset, RateLimitLog +from services.entities.feature_entities import LicenseStatus @pytest.fixture(autouse=True) @@ -124,6 +127,71 @@ class TestAccountInitialization: class TestCurrentContextInjection: """Test request context injection decorators.""" + def test_console_account_admission_injects_request_context(self): + current_user = make_account() + + with ( + patch( + "controllers.console.flask_admission.setup_required", side_effect=lambda view: view + ) as setup_required, + patch( + "controllers.console.flask_admission.login_required", side_effect=lambda view: view + ) as login_required, + patch( + "controllers.console.flask_admission.account_initialization_required", side_effect=lambda view: view + ) as account_initialization_required, + patch( + "controllers.console.flask_admission.current_account_with_tenant", + return_value=AccountWithTenant(account=current_user, tenant_id="tenant-123"), + ), + patch("controllers.console.flask_admission.get_request_id", return_value="request-1"), + patch("controllers.console.flask_admission.get_trace_id", return_value="trace-1"), + ): + + class Handler: + @flask_admission.console_account_admission() + def get(self, request_context: RequestContext): + return request_context + + with Flask(__name__).test_request_context(): + result = Handler().get() + + assert result == RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id=current_user.id, + active_workspace_id="tenant-123", + ) + setup_required.assert_called_once() + login_required.assert_called_once() + account_initialization_required.assert_called_once() + + def test_console_account_admission_preserves_route_kwarg_named_request_context(self): + current_user = make_account() + + with ( + patch("controllers.console.flask_admission.setup_required", side_effect=lambda view: view), + patch("controllers.console.flask_admission.login_required", side_effect=lambda view: view), + patch("controllers.console.flask_admission.account_initialization_required", side_effect=lambda view: view), + patch( + "controllers.console.flask_admission.current_account_with_tenant", + return_value=AccountWithTenant(account=current_user, tenant_id="tenant-123"), + ), + patch("controllers.console.flask_admission.get_request_id", return_value="request-1"), + patch("controllers.console.flask_admission.get_trace_id", return_value="trace-1"), + ): + + class Handler: + @flask_admission.console_account_admission() + def get(self, admission_context: RequestContext, request_context: str): + return admission_context, request_context + + with Flask(__name__).test_request_context(): + admission_context, route_value = Handler().get(request_context="route-value") + + assert admission_context.active_workspace_id == "tenant-123" + assert route_value == "route-value" + def test_with_current_tenant_id_injects_tenant_id(self): class Handler: @with_current_tenant_id @@ -328,6 +396,36 @@ class TestRbacPermissionRequired: request.view_args = {"resource_id": "dataset-1"} assert _extract_resource_id(RBACResourceScope.DATASET, "tenant-1") == "dataset-1" + def test_extract_resource_id_scopes_pipeline_resolution_to_the_calling_tenant(self, sqlite_session: Session): + app = Flask(__name__) + pipeline_id = "00000000-0000-0000-0000-000000000001" + current_tenant_id = "00000000-0000-0000-0000-000000000002" + foreign_dataset = Dataset( + id="00000000-0000-0000-0000-000000000003", + tenant_id="00000000-0000-0000-0000-000000000004", + name="Foreign decoy", + created_by="00000000-0000-0000-0000-000000000005", + pipeline_id=pipeline_id, + ) + current_dataset = Dataset( + id="00000000-0000-0000-0000-000000000006", + tenant_id=current_tenant_id, + name="Current tenant dataset", + created_by="00000000-0000-0000-0000-000000000007", + pipeline_id=pipeline_id, + ) + sqlite_session.add_all([foreign_dataset, current_dataset]) + + unscoped_dataset = sqlite_session.scalar(select(Dataset).where(Dataset.pipeline_id == pipeline_id)) + assert unscoped_dataset is foreign_dataset + + with ( + app.test_request_context("/rag/pipelines/pipeline-1"), + patch("controllers.common.wraps.db", SimpleNamespace(session=sqlite_session)), + ): + request.view_args = {"pipeline_id": pipeline_id} + assert _extract_resource_id(RBACResourceScope.DATASET, current_tenant_id) == current_dataset.id + def test_extract_resource_id_resolves_agent_to_its_authz_app(self): app = Flask(__name__) @@ -441,7 +539,7 @@ class TestEditionChecks: return "cloud_success" # Act - with patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"): + with patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD): result = cloud_view() # Assert @@ -458,13 +556,13 @@ class TestEditionChecks: # Act & Assert with app.test_request_context(): - with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + with patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY): with pytest.raises(HTTPException) as exc_info: cloud_view() assert exc_info.value.code == 404 - def test_only_edition_enterprise_allows_when_enabled(self): - """Test enterprise edition decorator allows when ENTERPRISE_ENABLED is True""" + def test_only_edition_enterprise_allows_enterprise_edition(self): + """Test enterprise edition decorator allows the ENTERPRISE edition.""" # Arrange @only_edition_enterprise @@ -472,14 +570,14 @@ class TestEditionChecks: return "enterprise_success" # Act - with patch("controllers.console.wraps.dify_config.ENTERPRISE_ENABLED", True): + with patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE): result = enterprise_view() # Assert assert result == "enterprise_success" def test_only_edition_self_hosted_allows_self_hosted(self): - """Test self-hosted edition decorator allows SELF_HOSTED edition""" + """Test self-hosted edition decorator allows the COMMUNITY edition.""" # Arrange @only_edition_self_hosted @@ -487,49 +585,13 @@ class TestEditionChecks: return "self_hosted_success" # Act - with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + with patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY): result = self_hosted_view() # Assert assert result == "self_hosted_success" -class TestBillingEnabled: - """Test billing enabled decorator.""" - - def test_should_allow_when_billing_config_enabled(self): - """Test billing decorator uses local config without loading tenant features.""" - - @cloud_edition_billing_enabled - def billing_view(): - return "billing_success" - - with patch("controllers.console.wraps.dify_config.BILLING_ENABLED", True): - with patch("controllers.console.wraps.FeatureService.get_features") as get_features: - result = billing_view() - - assert result == "billing_success" - get_features.assert_not_called() - - def test_should_reject_when_billing_config_disabled(self): - """Test billing decorator rejects when local billing config is disabled.""" - app = create_app_with_login() - - @cloud_edition_billing_enabled - def billing_view(): - return "billing_success" - - with app.test_request_context(): - with patch("controllers.console.wraps.dify_config.BILLING_ENABLED", False): - with patch("controllers.console.wraps.FeatureService.get_features") as get_features: - with pytest.raises(HTTPException) as exc_info: - billing_view() - - assert exc_info.value.code == 403 - assert "Billing feature is not enabled" in str(exc_info.value.description) - get_features.assert_not_called() - - class TestBillingPaidPlanRequired: @pytest.mark.parametrize("plan", ["professional", "team"]) def test_should_allow_paid_plan(self, plan: str): @@ -621,7 +683,7 @@ class TestBillingResourceLimits: "controllers.console.wraps.current_account_with_tenant", return_value=(MockUser("test_user"), "tenant123") ): with ( - patch("controllers.console.wraps.dify_config.BILLING_ENABLED", True), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch( "controllers.console.wraps.FeatureService.get_vector_space", return_value=mock_vector_space ) as get_vector_space, @@ -693,6 +755,17 @@ class TestBillingResourceLimits: result = upload_document() assert result == "document_uploaded" + # Test 3: Form source must enforce the same quota as query source + with app.test_request_context("/", method="POST", data={"source": "datasets"}): + with patch( + "controllers.console.wraps.current_account_with_tenant", + return_value=(MockUser("test_user"), "tenant123"), + ): + with patch("controllers.console.wraps.FeatureService.get_features", return_value=mock_features): + with pytest.raises(HTTPException) as exc_info: + upload_document() + assert exc_info.value.code == 403 + class TestRateLimiting: """Test rate limiting decorator""" @@ -773,8 +846,8 @@ class TestRateLimiting: class TestCloudUtmRecord: """Test cloud UTM recording decorator.""" - def test_should_record_utm_when_billing_config_enabled_and_cookie_exists(self): - """Test UTM recording uses billing config without loading tenant features.""" + def test_should_record_utm_for_cloud_edition_and_cookie(self): + """Test Cloud UTM recording without loading tenant features.""" app = create_app_with_login() @cloud_utm_record @@ -783,7 +856,7 @@ class TestCloudUtmRecord: with app.test_request_context("/", headers={"Cookie": "utm_info={}"}): with ( - patch("controllers.console.wraps.dify_config.BILLING_ENABLED", True), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.wraps.current_account_with_tenant", return_value=(MockUser("u1"), "t1")), patch("controllers.console.wraps.OperationService.record_utm") as record_utm, patch("controllers.console.wraps.FeatureService.get_features") as get_features, @@ -794,8 +867,8 @@ class TestCloudUtmRecord: record_utm.assert_called_once_with("t1", {}) get_features.assert_not_called() - def test_should_skip_utm_when_billing_config_disabled(self): - """Test UTM recording skips tenant feature loading when billing config is disabled.""" + def test_should_skip_utm_outside_cloud_edition(self): + """Test UTM recording skips tenant feature loading outside the Cloud edition.""" app = create_app_with_login() @cloud_utm_record @@ -804,7 +877,7 @@ class TestCloudUtmRecord: with app.test_request_context("/", headers={"Cookie": "utm_info={}"}): with ( - patch("controllers.console.wraps.dify_config.BILLING_ENABLED", False), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), patch("controllers.console.wraps.current_account_with_tenant") as current_account, patch("controllers.console.wraps.OperationService.record_utm") as record_utm, patch("controllers.console.wraps.FeatureService.get_features") as get_features, @@ -830,7 +903,7 @@ class TestSystemSetup: return "admin_success" # Act - with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + with patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY): result = admin_view() # Assert @@ -845,24 +918,25 @@ class TestSystemSetup: def admin_view(): return "admin_success" - with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + with patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY): assert admin_view() == "admin_success" assert admin_view() == "admin_success" assert mock_db.session.scalar.call_count == 1 @patch("controllers.console.wraps.db") - @patch("controllers.console.wraps.os.environ.get") - def test_should_not_cache_missing_setup(self, mock_environ_get, mock_db): + def test_should_not_cache_missing_setup(self, mock_db): """Test that first-time bootstrap completion can be observed later in the same process""" mock_db.session.scalar.side_effect = [None, MagicMock()] - mock_environ_get.return_value = None @setup_required def admin_view(): return "admin_success" - with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + with ( + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), + patch("controllers.console.wraps.dify_config.INIT_PASSWORD", ""), + ): with pytest.raises(NotSetupError): admin_view() assert admin_view() == "admin_success" @@ -870,36 +944,38 @@ class TestSystemSetup: assert mock_db.session.scalar.call_count == 2 @patch("controllers.console.wraps.db") - @patch("controllers.console.wraps.os.environ.get") - def test_should_raise_not_init_validate_error_with_init_password(self, mock_environ_get, mock_db: MagicMock): + def test_should_raise_not_init_validate_error_with_init_password(self, mock_db: MagicMock): """Test NotInitValidateError when INIT_PASSWORD is set but setup not complete""" # Arrange mock_db.session.scalar.return_value = None # No setup - mock_environ_get.return_value = "some_password" @setup_required def admin_view(): return "admin_success" # Act & Assert - with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + with ( + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), + patch("controllers.console.wraps.dify_config.INIT_PASSWORD", "some_password"), + ): with pytest.raises(NotInitValidateError): admin_view() @patch("controllers.console.wraps.db") - @patch("controllers.console.wraps.os.environ.get") - def test_should_raise_not_setup_error_without_init_password(self, mock_environ_get, mock_db: MagicMock): + def test_should_raise_not_setup_error_without_init_password(self, mock_db: MagicMock): """Test NotSetupError when no INIT_PASSWORD and setup not complete""" # Arrange mock_db.session.scalar.return_value = None # No setup - mock_environ_get.return_value = None # No INIT_PASSWORD @setup_required def admin_view(): return "admin_success" # Act & Assert - with patch("controllers.console.wraps.dify_config.EDITION", "SELF_HOSTED"): + with ( + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), + patch("controllers.console.wraps.dify_config.INIT_PASSWORD", ""), + ): with pytest.raises(NotSetupError): admin_view() diff --git a/api/tests/unit_tests/controllers/console/workspace/test_accounts.py b/api/tests/unit_tests/controllers/console/workspace/test_accounts.py index 5861fa87871..d2857182858 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_accounts.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_accounts.py @@ -16,6 +16,7 @@ from controllers.console.auth.error import ( from controllers.console.error import AccountInFreezeError from controllers.console.workspace.account import ( AccountAvatarApi, + AccountAvatarQuery, AccountDeleteApi, AccountDeleteVerifyApi, AccountInitApi, @@ -35,6 +36,7 @@ from controllers.console.workspace.error import ( CurrentPasswordIncorrectError, InvalidAccountDeletionCodeError, ) +from enums import DeploymentEdition from extensions.storage.storage_type import StorageType from models import Account, AccountIntegrate, InvitationCode, Tenant, TenantAccountJoin from models.account import AccountStatus, InvitationCodeStatus, TenantAccountRole @@ -118,7 +120,7 @@ class TestAccountInitApi: with ( app.test_request_context("/account/init", json=payload), - patch("controllers.console.workspace.account.dify_config.EDITION", "CLOUD"), + patch("controllers.console.workspace.account.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("controllers.console.workspace.account.db.session", sqlite_session), ): resp = method(api, account) @@ -216,7 +218,7 @@ class TestAccountAvatarApiGet: return_value="https://signed/example", ) as sign_mock, ): - result = method(api, tenant_id, user) + result = method(api, AccountAvatarQuery(avatar=file_id), user) assert result == {"avatar_url": "https://signed/example"} sign_mock.assert_called_once_with(upload_file_id=file_id) @@ -252,11 +254,11 @@ class TestAccountAvatarApiGet: patch("controllers.console.workspace.account.db.session", sqlite_session), patch( "controllers.console.workspace.account.file_helpers.get_signed_file_url", - return_value="https://signed/leak", + return_value="https://signed/example", ) as sign_mock, ): with pytest.raises(NotFound): - method(api, tenant_id, user) + method(api, AccountAvatarQuery(avatar=file_id), user) sign_mock.assert_not_called() @@ -265,12 +267,13 @@ class TestAccountAvatarApiGet: [(Account, Tenant, TenantAccountJoin, UploadFile)], indirect=True, ) - def test_get_avatar_not_found_when_upload_belongs_to_other_tenant(self, app: Flask, sqlite_session: Session): + def test_get_avatar_signed_url_when_upload_owned_by_current_account_in_other_tenant( + self, app: Flask, sqlite_session: Session + ): api = AccountAvatarApi() method = inspect.unwrap(api.get) - user, tenant = persist_account_with_tenant(sqlite_session, "acc-owner") - tenant_id = tenant.id + user, _ = persist_account_with_tenant(sqlite_session, "acc-owner") file_id = "550e8400-e29b-41d4-a716-446655440002" other_tenant = Tenant(name="tenant-other") @@ -285,20 +288,19 @@ class TestAccountAvatarApiGet: patch("controllers.console.workspace.account.db.session", sqlite_session), patch( "controllers.console.workspace.account.file_helpers.get_signed_file_url", - return_value="https://signed/leak", + return_value="https://signed/example", ) as sign_mock, ): - with pytest.raises(NotFound): - method(api, tenant_id, user) + result = method(api, AccountAvatarQuery(avatar=file_id), user) - sign_mock.assert_not_called() + assert result == {"avatar_url": "https://signed/example"} + sign_mock.assert_called_once_with(upload_file_id=file_id) def test_get_avatar_https_pass_through_without_signing(self, app: Flask): api = AccountAvatarApi() method = inspect.unwrap(api.get) user = make_account("acc-owner") - tenant_id = "tenant-1" external = "https://cdn.example/avatar.png" with ( @@ -308,7 +310,7 @@ class TestAccountAvatarApiGet: return_value="https://signed/should-not-use", ) as sign_mock, ): - result = method(api, tenant_id, user) + result = method(api, AccountAvatarQuery(avatar=external), user) assert result == {"avatar_url": external} sign_mock.assert_not_called() diff --git a/api/tests/unit_tests/controllers/console/workspace/test_endpoint.py b/api/tests/unit_tests/controllers/console/workspace/test_endpoint.py index 66e6b8fe35a..a34349710ea 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_endpoint.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_endpoint.py @@ -11,11 +11,17 @@ from controllers.console.workspace.endpoint import ( DeprecatedEndpointDeleteApi, DeprecatedEndpointUpdateApi, EndpointCollectionApi, + EndpointCreatePayload, EndpointDisableApi, EndpointEnableApi, + EndpointIdPayload, EndpointItemApi, EndpointListApi, + EndpointListForPluginQuery, EndpointListForSinglePluginApi, + EndpointListQuery, + EndpointUpdatePayload, + LegacyEndpointUpdatePayload, ) from core.entities.provider_entities import ProviderConfig, ProviderConfigType from core.plugin.entities.endpoint import EndpointEntityWithInstance, EndpointProviderDeclaration @@ -60,12 +66,13 @@ class TestEndpointCollectionApi: "name": "endpoint", "settings": {"a": 1}, } + req_data = EndpointCreatePayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.create_endpoint", return_value=True), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is True @@ -78,6 +85,7 @@ class TestEndpointCollectionApi: "name": "endpoint", "settings": {}, } + req_data = EndpointCreatePayload(**payload) with ( app.test_request_context("/", json=payload), @@ -87,7 +95,7 @@ class TestEndpointCollectionApi: ), ): with pytest.raises(ValueError): - method(api, "t1", "u1") + method(api, req_data, "t1", "u1") def test_create_validation_error(self, app: Flask): api = EndpointCollectionApi() @@ -103,7 +111,7 @@ class TestEndpointCollectionApi: app.test_request_context("/", json=payload), ): with pytest.raises(ValueError): - method(api, "t1", "u1") + method(api, EndpointCreatePayload(**payload), "t1", "u1") class TestDeprecatedEndpointCreateApi: @@ -116,12 +124,13 @@ class TestDeprecatedEndpointCreateApi: "name": "endpoint", "settings": {"a": 1}, } + req_data = EndpointCreatePayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.create_endpoint", return_value=True), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is True @@ -139,7 +148,7 @@ class TestEndpointListApi: return_value=[endpoint_entity], ), ): - result = method(api, "t1", "u1") + result = method(api, EndpointListQuery(page=1, page_size=10), "t1", "u1") endpoint = result["endpoints"][0] assert endpoint["id"] == "e1" @@ -173,7 +182,7 @@ class TestEndpointListApi: app.test_request_context("/?page=0&page_size=10"), ): with pytest.raises(ValueError): - method(api, "t1", "u1") + method(api, EndpointListQuery(page=0, page_size=10), "t1", "u1") class TestEndpointListForSinglePluginApi: @@ -188,7 +197,7 @@ class TestEndpointListForSinglePluginApi: return_value=[_endpoint_entity()], ), ): - result = method(api, "t1", "u1") + result = method(api, EndpointListForPluginQuery(page=1, page_size=10, plugin_id="p1"), "t1", "u1") assert result["endpoints"][0]["id"] == "e1" assert result["endpoints"][0]["settings"]["api_key"] == "pl********et" @@ -202,7 +211,7 @@ class TestEndpointListForSinglePluginApi: app.test_request_context("/?page=1&page_size=10"), ): with pytest.raises(ValueError): - method(api, "t1", "u1") + method(api, EndpointListForPluginQuery(page=1, page_size=10), "t1", "u1") class TestEndpointItemApi: @@ -242,6 +251,7 @@ class TestEndpointItemApi: "name": "new-name", "settings": {"x": 1}, } + req_data = EndpointUpdatePayload(**payload) with ( app.test_request_context("/", method="PATCH", json=payload), @@ -250,7 +260,7 @@ class TestEndpointItemApi: return_value=True, ) as mock_update, ): - result = method(api, "t1", "u1", "e1") + result = method(api, req_data, "t1", "u1", "e1") assert result["success"] is True mock_update.assert_called_once_with( @@ -271,7 +281,7 @@ class TestEndpointItemApi: app.test_request_context("/", method="PATCH", json=payload), ): with pytest.raises(ValueError): - method(api, "t1", "u1", "e1") + method(api, EndpointUpdatePayload(**payload), "t1", "u1", "e1") def test_update_service_failure(self, app: Flask): api = EndpointItemApi() @@ -281,12 +291,13 @@ class TestEndpointItemApi: "name": "n", "settings": {}, } + req_data = EndpointUpdatePayload(**payload) with ( app.test_request_context("/", method="PATCH", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.update_endpoint", return_value=False), ): - result = method(api, "t1", "u1", "e1") + result = method(api, req_data, "t1", "u1", "e1") assert result["success"] is False @@ -297,12 +308,13 @@ class TestDeprecatedEndpointDeleteApi: method = inspect.unwrap(api.post) payload = {"endpoint_id": "e1"} + req_data = EndpointIdPayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.delete_endpoint", return_value=True), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is True @@ -314,19 +326,20 @@ class TestDeprecatedEndpointDeleteApi: app.test_request_context("/", json={}), ): with pytest.raises(ValueError): - method(api, "t1", "u1") + method(api, EndpointIdPayload(), "t1", "u1") def test_delete_service_failure(self, app: Flask): api = DeprecatedEndpointDeleteApi() method = inspect.unwrap(api.post) payload = {"endpoint_id": "e1"} + req_data = EndpointIdPayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.delete_endpoint", return_value=False), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is False @@ -341,12 +354,13 @@ class TestDeprecatedEndpointUpdateApi: "name": "new-name", "settings": {"x": 1}, } + req_data = LegacyEndpointUpdatePayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.update_endpoint", return_value=True), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is True @@ -360,7 +374,7 @@ class TestDeprecatedEndpointUpdateApi: app.test_request_context("/", json=payload), ): with pytest.raises(ValueError): - method(api, "t1", "u1") + method(api, LegacyEndpointUpdatePayload(**payload), "t1", "u1") def test_update_service_failure(self, app: Flask): api = DeprecatedEndpointUpdateApi() @@ -371,12 +385,13 @@ class TestDeprecatedEndpointUpdateApi: "name": "n", "settings": {}, } + req_data = LegacyEndpointUpdatePayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.update_endpoint", return_value=False), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is False @@ -417,12 +432,13 @@ class TestEndpointEnableApi: method = inspect.unwrap(api.post) payload = {"endpoint_id": "e1"} + req_data = EndpointIdPayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.enable_endpoint", return_value=True), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is True @@ -434,19 +450,20 @@ class TestEndpointEnableApi: app.test_request_context("/", json={}), ): with pytest.raises(ValueError): - method(api, "t1", "u1") + method(api, EndpointIdPayload(), "t1", "u1") def test_enable_service_failure(self, app: Flask): api = EndpointEnableApi() method = inspect.unwrap(api.post) payload = {"endpoint_id": "e1"} + req_data = EndpointIdPayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.enable_endpoint", return_value=False), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is False @@ -457,12 +474,13 @@ class TestEndpointDisableApi: method = inspect.unwrap(api.post) payload = {"endpoint_id": "e1"} + req_data = EndpointIdPayload(**payload) with ( app.test_request_context("/", json=payload), patch("controllers.console.workspace.endpoint.EndpointService.disable_endpoint", return_value=True), ): - result = method(api, "t1", "u1") + result = method(api, req_data, "t1", "u1") assert result["success"] is True @@ -474,4 +492,4 @@ class TestEndpointDisableApi: app.test_request_context("/", json={}), ): with pytest.raises(ValueError): - method(api, "t1", "u1") + method(api, EndpointIdPayload(), "t1", "u1") diff --git a/api/tests/unit_tests/controllers/console/workspace/test_members.py b/api/tests/unit_tests/controllers/console/workspace/test_members.py index 58f8468429a..e7b60184275 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_members.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_members.py @@ -1,6 +1,9 @@ from contextlib import nullcontext +from datetime import datetime +from http import HTTPStatus from inspect import unwrap from types import SimpleNamespace +from typing import NamedTuple, override from unittest.mock import MagicMock, patch import pytest @@ -28,93 +31,95 @@ from controllers.console.workspace.members import ( SendOwnerTransferEmailApi, _count_new_member_invites, ) +from enums import DeploymentEdition from libs.external_api import ExternalApi +from machinery.context import RequestContext from services.errors.account import AccountAlreadyInTenantError, SeatsLimitExceededError +from services.workspace_member_query_service import ( + WorkspaceMemberQueryService, + WorkspaceMemberRole, + WorkspaceMemberSummary, +) + + +class _RecordingWorkspaceMemberQueryService(WorkspaceMemberQueryService): + def __init__(self, result: tuple[WorkspaceMemberSummary, ...]) -> None: + self._result = result + self.contexts: list[RequestContext] = [] + + @override + def list_current(self, context: RequestContext) -> tuple[WorkspaceMemberSummary, ...]: + self.contexts.append(context) + return self._result + + +class _ApplicationServicesStub(NamedTuple): + workspace_member_queries: WorkspaceMemberQueryService class TestMemberListApi: - def test_get_success(self, app: Flask): + def test_get_passes_context_and_serializes_application_result(self, app: Flask) -> None: api = MemberListApi() method = unwrap(api.get) - - tenant = MagicMock() - user = MagicMock(current_tenant=tenant) - member = MagicMock() - member.id = "m1" - member.name = "Member" - member.email = "member@test.com" - member.avatar = "avatar.png" - member.current_role = SimpleNamespace(value="admin") - member.status = SimpleNamespace(value="active") - members = [member] + request_context = RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id="actor-1", + active_workspace_id="workspace-1", + ) + timestamp = datetime(2026, 1, 1) + workspace_member_queries = _RecordingWorkspaceMemberQueryService( + ( + WorkspaceMemberSummary( + id="member-1", + name="Member", + email="member@example.com", + avatar=None, + last_login_at=None, + last_active_at=timestamp, + created_at=timestamp, + role="owner", + roles=( + WorkspaceMemberRole(id="workspace.owner", name="Owner"), + WorkspaceMemberRole(id="workspace.editor", name="Editor"), + ), + status="active", + ), + ) + ) + application_services_stub = _ApplicationServicesStub(workspace_member_queries=workspace_member_queries) with ( app.test_request_context("/"), - patch("controllers.console.workspace.members.TenantService.get_tenant_members", return_value=members), - ): - result, status = method(api, user) - - assert status == 200 - assert len(result["accounts"]) == 1 - assert result["accounts"][0]["role"] == "admin" - assert result["accounts"][0]["roles"] == [{"id": "admin", "name": "admin"}] - - def test_get_with_rbac_enabled_fetches_roles_in_batch(self, app): - api = MemberListApi() - method = unwrap(api.get) - - tenant = MagicMock(id="tenant-1") - user = MagicMock(id="acct-1", current_tenant=tenant) - member = SimpleNamespace( - id="m1", - name="Member", - email="member@test.com", - avatar=None, - last_login_at=1, - last_active_at=2, - created_at=3, - current_role=SimpleNamespace(value="editor"), - status=SimpleNamespace(value="active"), - ) - role_item = SimpleNamespace( - account_id="m1", - roles=[ - SimpleNamespace(id="workspace.owner", name="Owner"), - SimpleNamespace(id="workspace.editor", name="Editor"), - ], - ) - - with ( - app.test_request_context("/"), - patch("controllers.console.workspace.members.current_account_with_tenant", return_value=(user, "tenant-1")), - patch("controllers.console.workspace.members.dify_config.RBAC_ENABLED", True), - patch("controllers.console.workspace.members.TenantService.get_tenant_members", return_value=[member]), patch( - "controllers.console.workspace.members.enterprise_rbac_service.RBACService.MemberRoles.batch_get", - return_value=[role_item], - ) as mock_batch_get, + "controllers.console.workspace.members.application_services", + return_value=application_services_stub, + ), ): - result, status = method(api) + result, status = method(api, request_context=request_context) - assert status == 200 - assert result["accounts"][0]["role"] == "editor" - assert result["accounts"][0]["roles"] == [ - {"id": "workspace.owner", "name": "Owner"}, - {"id": "workspace.editor", "name": "Editor"}, - ] - mock_batch_get.assert_called_once_with("tenant-1", "acct-1", ["m1"]) - - def test_get_no_tenant(self, app: Flask): - api = MemberListApi() - method = unwrap(api.get) - - user = MagicMock(current_tenant=None) - - with ( - app.test_request_context("/"), - ): - with pytest.raises(ValueError): - method(api, user) + assert status == HTTPStatus.OK + assert result == { + "accounts": [ + { + "id": "member-1", + "name": "Member", + "email": "member@example.com", + "avatar": None, + "avatar_url": None, + "last_login_at": None, + "last_active_at": int(timestamp.timestamp()), + "created_at": int(timestamp.timestamp()), + "role": "owner", + "roles": [ + {"id": "workspace.owner", "name": "Owner"}, + {"id": "workspace.editor", "name": "Editor"}, + ], + "status": "active", + } + ] + } + assert workspace_member_queries.contexts == [request_context] class TestMemberInviteEmailApi: @@ -130,7 +135,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = False features.workspace_members.is_available.return_value = True @@ -148,8 +152,7 @@ class TestMemberInviteEmailApi: "controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token" ) as mock_invite, patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), ): result, status = method(api, user) @@ -167,7 +170,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = True features.workspace_members.is_available.return_value = False @@ -180,20 +182,18 @@ class TestMemberInviteEmailApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features), patch("controllers.console.workspace.members._count_new_member_invites", return_value=(1, 1)), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE), ): with pytest.raises(WorkspaceMembersLimitExceeded): method(api, user) - def test_invite_billing_limit_exceeded(self, app: Flask): + def test_invite_cloud_member_limit_exceeded(self, app: Flask): api = MemberInviteEmailApi() method = unwrap(api.post) tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = True features.members.size = 9 features.members.limit = 10 features.workspace_members.enabled = False @@ -208,8 +208,7 @@ class TestMemberInviteEmailApi: patch("controllers.console.workspace.members.FeatureService.get_features", return_value=features), patch("controllers.console.workspace.members._count_new_member_invites", return_value=(2, 2)), patch("controllers.console.workspace.members._count_current_members", return_value=9), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", True), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), ): with pytest.raises(WorkspaceMembersLimitExceeded): method(api, user) @@ -221,7 +220,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = False features.workspace_members.is_available.return_value = True @@ -239,8 +237,7 @@ class TestMemberInviteEmailApi: side_effect=AccountAlreadyInTenantError(), ), patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), ): result, status = method(api, user) @@ -289,7 +286,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = False features.workspace_members.is_available.return_value = True @@ -307,8 +303,7 @@ class TestMemberInviteEmailApi: side_effect=Exception("boom"), ), patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), ): result, _ = method(api, user) @@ -321,7 +316,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = False license_info = MagicMock() license_info.seats.is_available.return_value = False @@ -340,8 +334,7 @@ class TestMemberInviteEmailApi: return_value=license_info, ) as mock_get_license, patch("controllers.console.workspace.members.RegisterService.invite_new_member") as mock_invite, - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE), ): with pytest.raises(SeatsLimitExceeded): method(api, user) @@ -357,7 +350,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = False license_info = MagicMock() license_info.seats.is_available.return_value = False @@ -379,8 +371,7 @@ class TestMemberInviteEmailApi: "controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token" ) as mock_invite, patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE), ): result, status = method(api, user) @@ -397,7 +388,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = False license_info = MagicMock() license_info.seats.is_available.return_value = True @@ -419,8 +409,7 @@ class TestMemberInviteEmailApi: "controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token" ) as mock_invite, patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE), ): result, status = method(api, user) @@ -437,7 +426,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = False license_info = MagicMock() license_info.seats.is_available.return_value = False @@ -457,8 +445,7 @@ class TestMemberInviteEmailApi: ) as mock_get_license, patch("controllers.console.workspace.members.RegisterService.invite_new_member", return_value="token"), patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), ): result, status = method(api, user) @@ -474,7 +461,6 @@ class TestMemberInviteEmailApi: tenant = MagicMock(id="t1") user = MagicMock(current_tenant=tenant) features = MagicMock() - features.billing.enabled = False features.workspace_members.enabled = False license_info = MagicMock() license_info.seats.is_available.return_value = True @@ -497,8 +483,7 @@ class TestMemberInviteEmailApi: side_effect=SeatsLimitExceededError("licensed seats limit exceeded"), ), patch("controllers.console.workspace.members.dify_config.CONSOLE_WEB_URL", "http://x"), - patch("controllers.console.workspace.members.dify_config.ENTERPRISE_ENABLED", True), - patch("controllers.console.workspace.members.dify_config.BILLING_ENABLED", False), + patch("controllers.console.workspace.members.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE), ): result, status = method(api, user) diff --git a/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py b/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py index af37c20edaa..5fe0f47ecde 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_model_providers.py @@ -12,12 +12,15 @@ from configs import dify_config from controllers.console.workspace.model_providers import ( ModelProviderCredentialApi, ModelProviderCredentialSwitchApi, + ModelProviderCreditsApi, ModelProviderIconApi, ModelProviderListApi, ModelProviderPaymentCheckoutUrlApi, + ModelProviderSummaryListApi, ModelProviderValidateApi, PreferredProviderTypeUpdateApi, ) +from core.entities.provider_entities import CredentialConfiguration from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.provider_entities import ConfigurateMethod @@ -27,9 +30,14 @@ from models.provider import ProviderType from services.entities.model_provider_entities import ( CustomConfigurationResponse, CustomConfigurationStatus, + ModelProviderCustomConfigurationSummaryResponse, + ModelProviderPluginSummaryResponse, + ModelProviderSummaryResponse, + ModelProviderSystemConfigurationSummaryResponse, ProviderResponse, SystemConfigurationResponse, ) +from services.workspace_service import EffectiveCreditPool VALID_UUID = "123e4567-e89b-12d3-a456-426614174000" INVALID_UUID = "123" @@ -140,6 +148,136 @@ class TestModelProviderListApi: assert result == {"data": []} +class TestModelProviderSummaryListApi: + def test_get_success(self, app: Flask): + api = ModelProviderSummaryListApi() + method = unwrap(api.get) + provider = ModelProviderSummaryResponse( + tenant_id="tenant1", + provider="langgenius/openai/openai", + plugin_id="langgenius/openai", + label=I18nObject(en_US="OpenAI"), + description=I18nObject(en_US="OpenAI models"), + icon_small=I18nObject(en_US="icon.svg"), + icon_small_dark=None, + supported_model_types=[ModelType.LLM], + configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL], + preferred_provider_type=ProviderType.CUSTOM, + is_configured=True, + custom_configuration=ModelProviderCustomConfigurationSummaryResponse( + status=CustomConfigurationStatus.ACTIVE, + has_custom_models=True, + available_credentials=[ + CredentialConfiguration( + credential_id=VALID_UUID, + credential_name="production", + ), + CredentialConfiguration( + credential_id="223e4567-e89b-12d3-a456-426614174000", + credential_name="backup", + ), + ], + current_credential_id=VALID_UUID, + current_credential_name="production", + current_credential_usable=True, + ), + system_configuration=ModelProviderSystemConfigurationSummaryResponse(enabled=False), + ) + plugin = ModelProviderPluginSummaryResponse( + installation_id="installation-1", + plugin_id="langgenius/openai", + plugin_unique_identifier="langgenius/openai:1.0.0@checksum", + runtime_type="local", + source="marketplace", + version="1.0.0", + ) + + with ( + app.test_request_context("/"), + patch( + "controllers.console.workspace.model_providers.ModelProviderService.get_provider_summary_list", + return_value=([provider], {"langgenius/openai": plugin}), + ) as get_provider_summary_list, + ): + result = method(api, "tenant1") + + get_provider_summary_list.assert_called_once_with(tenant_id="tenant1") + assert result["data"][0]["provider"] == "langgenius/openai/openai" + assert "tenant_id" not in result["data"][0] + assert result["data"][0]["custom_configuration"] == { + "status": "active", + "has_custom_models": True, + "available_credentials": [ + { + "credential_id": VALID_UUID, + "credential_name": "production", + }, + { + "credential_id": "223e4567-e89b-12d3-a456-426614174000", + "credential_name": "backup", + }, + ], + "current_credential_id": VALID_UUID, + "current_credential_name": "production", + "current_credential_usable": True, + } + assert result["plugins"]["langgenius/openai"]["installation_id"] == "installation-1" + + +class TestModelProviderCreditsApi: + def test_get_success(self): + api = ModelProviderCreditsApi() + method = unwrap(api.get) + session = SimpleNamespace() + credit_pool = EffectiveCreditPool( + plan="team", + pool_type="paid", + quota_limit=-1, + quota_used=999, + next_credit_reset_date=1775001600, + ) + + with patch( + "controllers.console.workspace.model_providers.WorkspaceService.get_effective_credit_pool", + return_value=credit_pool, + ) as get_effective_credit_pool: + result = method(api, session, "tenant1") + + get_effective_credit_pool.assert_called_once_with("tenant1", session=session) + assert result == { + "pool_type": "paid", + "quota_limit": -1, + "quota_used": 999, + "remaining_credits": -1, + "is_unlimited": True, + "is_exhausted": False, + "exhausted_at": None, + "next_credit_reset_date": 1775001600, + } + + def test_get_without_effective_pool(self): + api = ModelProviderCreditsApi() + method = unwrap(api.get) + session = SimpleNamespace() + + with patch( + "controllers.console.workspace.model_providers.WorkspaceService.get_effective_credit_pool", + return_value=EffectiveCreditPool(), + ): + result = method(api, session, "tenant1") + + assert result == { + "pool_type": None, + "quota_limit": None, + "quota_used": None, + "remaining_credits": None, + "is_unlimited": False, + "is_exhausted": True, + "exhausted_at": None, + "next_credit_reset_date": None, + } + + class TestModelProviderCredentialApi: def test_get_success(self, app: Flask): api = ModelProviderCredentialApi() diff --git a/api/tests/unit_tests/controllers/console/workspace/test_models.py b/api/tests/unit_tests/controllers/console/workspace/test_models.py index 424c1ef1d75..dd0718f33bc 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_models.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_models.py @@ -15,6 +15,16 @@ from controllers.console.workspace.models import ( ModelProviderModelEnableApi, ModelProviderModelParameterRuleApi, ModelProviderModelValidateApi, + ParserCreateCredential, + ParserDeleteCredential, + ParserDeleteModels, + ParserGetCredentials, + ParserGetDefault, + ParserParameter, + ParserPostDefault, + ParserPostModels, + ParserSwitch, + ParserValidate, ) from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.errors.validate import CredentialsValidateFailedError @@ -43,7 +53,7 @@ class TestDefaultModelApi: }, } - result = method(api, "tenant1") + result = method(api, ParserGetDefault(model_type=ModelType.LLM), "tenant1") assert "data" in result @@ -65,7 +75,7 @@ class TestDefaultModelApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.models.ModelProviderService"), ): - result = method(api, "tenant1") + result = method(api, ParserPostDefault.model_validate(payload), "tenant1") assert result["result"] == "success" @@ -79,7 +89,7 @@ class TestDefaultModelApi: ): service.return_value.get_default_model_of_model_type.return_value = None - result = method(api, "t1") + result = method(api, ParserGetDefault(model_type=ModelType.LLM), "t1") assert "data" in result @@ -117,7 +127,7 @@ class TestModelProviderModelApi: patch("controllers.console.workspace.models.ModelProviderService"), patch("controllers.console.workspace.models.ModelLoadBalancingService"), ): - result, status = method(api, "tenant1", "openai") + result, status = method(api, ParserPostModels.model_validate(payload), "tenant1", "openai") assert status == 200 @@ -134,7 +144,7 @@ class TestModelProviderModelApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.models.ModelProviderService"), ): - result, status = method(api, "tenant1", "openai") + result, status = method(api, ParserDeleteModels.model_validate(payload), "tenant1", "openai") assert status == 204 @@ -177,7 +187,13 @@ class TestModelProviderModelCredentialApi: provider_service.return_value.provider_manager.get_provider_model_available_credentials.return_value = [] lb_service.return_value.get_load_balancing_configs.return_value = (False, []) - result = method(api, "tenant1", SimpleNamespace(id="u1"), "openai") + result = method( + api, + ParserGetCredentials(model="gpt-4", model_type=ModelType.LLM), + "tenant1", + SimpleNamespace(id="u1"), + "openai", + ) assert "credentials" in result @@ -195,7 +211,7 @@ class TestModelProviderModelCredentialApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.models.ModelProviderService"), ): - result, status = method(api, "tenant1", "openai") + result, status = method(api, ParserCreateCredential.model_validate(payload), "tenant1", "openai") assert status == 201 @@ -212,7 +228,13 @@ class TestModelProviderModelCredentialApi: service.return_value.provider_manager.get_provider_model_available_credentials.return_value = [] lb.return_value.get_load_balancing_configs.return_value = (False, []) - result = method(api, "t1", SimpleNamespace(id="u1"), "openai") + result = method( + api, + ParserGetCredentials(model="gpt", model_type=ModelType.LLM), + "t1", + SimpleNamespace(id="u1"), + "openai", + ) assert result["credentials"] == {} @@ -230,7 +252,7 @@ class TestModelProviderModelCredentialApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.models.ModelProviderService"), ): - result, status = method(api, "t1", "openai") + result, status = method(api, ParserDeleteCredential.model_validate(payload), "t1", "openai") assert status == 204 @@ -250,7 +272,7 @@ class TestModelProviderModelCredentialSwitchApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.models.ModelProviderService"), ): - result = method(api, "tenant1", "openai") + result = method(api, ParserSwitch.model_validate(payload), "tenant1", "openai") assert result["result"] == "success" @@ -269,7 +291,7 @@ class TestModelEnableDisableApis: app.test_request_context("/", json=payload), patch("controllers.console.workspace.models.ModelProviderService"), ): - result = method(api, "tenant1", "openai") + result = method(api, ParserDeleteModels.model_validate(payload), "tenant1", "openai") assert result["result"] == "success" @@ -286,7 +308,7 @@ class TestModelEnableDisableApis: app.test_request_context("/", json=payload), patch("controllers.console.workspace.models.ModelProviderService"), ): - result = method(api, "tenant1", "openai") + result = method(api, ParserDeleteModels.model_validate(payload), "tenant1", "openai") assert result["result"] == "success" @@ -306,7 +328,7 @@ class TestModelProviderModelValidateApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.models.ModelProviderService"), ): - result = method(api, "tenant1", "openai") + result = method(api, ParserValidate.model_validate(payload), "tenant1", "openai") assert result["result"] == "success" @@ -327,7 +349,7 @@ class TestModelProviderModelValidateApi: ): service_mock.return_value.validate_model_credentials.side_effect = CredentialsValidateFailedError("invalid") - result = method(api, "tenant1", "openai") + result = method(api, ParserValidate.model_validate(payload), "tenant1", "openai") assert result["result"] == "error" @@ -343,7 +365,7 @@ class TestParameterAndAvailableModels: ): service_mock.return_value.get_model_parameter_rules.return_value = [] - result = method(api, "tenant1", "openai") + result = method(api, ParserParameter(model="gpt-4"), "tenant1", "openai") assert "data" in result @@ -371,7 +393,7 @@ class TestParameterAndAvailableModels: ): service.return_value.get_model_parameter_rules.return_value = [] - result = method(api, "t1", "openai") + result = method(api, ParserParameter(model="gpt"), "t1", "openai") assert result["data"] == [] diff --git a/api/tests/unit_tests/controllers/console/workspace/test_plugin.py b/api/tests/unit_tests/controllers/console/workspace/test_plugin.py index 0cb336cfea8..7e81596d443 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_plugin.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_plugin.py @@ -10,6 +10,25 @@ from werkzeug.datastructures import FileStorage from werkzeug.exceptions import Forbidden from controllers.console.workspace.plugin import ( + ParserAsset, + ParserAutoUpgradeChange, + ParserAutoUpgradeFetch, + ParserDynamicOptions, + ParserDynamicOptionsWithCredentials, + ParserExcludePlugin, + ParserGithubInstall, + ParserGithubUpgrade, + ParserGithubUpload, + ParserIcon, + ParserLatest, + ParserList, + ParserMarketplaceUpgrade, + ParserPermissionChange, + ParserPluginIdentifierQuery, + ParserPluginIdentifiers, + ParserReadme, + ParserTasks, + ParserUninstall, PluginAssetApi, PluginAutoUpgradeExcludePluginApi, PluginCategoryListApi, @@ -28,6 +47,8 @@ from controllers.console.workspace.plugin import ( PluginFetchMarketplacePkgApi, PluginFetchPermissionApi, PluginIconApi, + PluginInstalledIdsApi, + PluginInstalledIdsQuery, PluginInstallFromGithubApi, PluginInstallFromMarketplaceApi, PluginInstallFromPkgApi, @@ -41,12 +62,16 @@ from controllers.console.workspace.plugin import ( PluginUploadFromBundleApi, PluginUploadFromGithubApi, PluginUploadFromPkgApi, + _list_hardcoded_builtin_tool_providers, ) from core.plugin.entities.parameters import PluginParameterOption -from core.plugin.entities.plugin import PluginDeclaration, PluginEntity, PluginInstallation +from core.plugin.entities.plugin import PluginCategory, PluginDeclaration, PluginEntity, PluginInstallation from core.plugin.entities.plugin_daemon import PluginInstallTask from core.plugin.impl.exc import PluginDaemonClientSideError from core.plugin.plugin_service import PluginService +from core.tools.entities.api_entities import ToolProviderApiEntity +from core.tools.entities.common_entities import I18nObject +from core.tools.entities.tool_entities import ToolProviderType from models.account import ( Account, TenantAccountRole, @@ -362,7 +387,7 @@ class TestPluginListLatestVersionsApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.plugin.PluginService.list_latest_versions", return_value=versions), ): - result = method(api) + result = method(api, ParserLatest.model_validate(payload)) assert result == { "versions": { @@ -384,7 +409,7 @@ class TestPluginListLatestVersionsApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api) + result = method(api, ParserLatest.model_validate(payload)) assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -430,7 +455,7 @@ class TestPluginListApi: return_value=plugins_with_total, ) as mock_list_with_total, ): - result = method(api, "t1", "u1") + result = method(api, ParserList(page=1, page_size=10), "t1", "u1") assert result == {"plugins": [_expected_plugin_entity_dump()], "total": 1} mock_list_with_total.assert_called_once_with("t1", "u1", 1, 10) @@ -445,7 +470,7 @@ class TestPluginCategoryListApi: mock_list = MagicMock(list=[plugin_item], has_more=True) with ( - app.test_request_context("/?page=2&page_size=10"), + app.test_request_context("/?page=2&page_size=10&query=weather&tags=search&tags=rag&language=zh_Hans"), patch( "controllers.console.workspace.plugin.PluginService.list_by_category", return_value=mock_list ) as list_mock, @@ -456,18 +481,75 @@ class TestPluginCategoryListApi: ): result = method(api, "t1", "tool") - list_mock.assert_called_once() - assert list_mock.call_args.args[0] == "t1" - assert list_mock.call_args.args[1] == "tool" - assert list_mock.call_args.args[2] == 2 - assert list_mock.call_args.args[3] == 10 + list_mock.assert_called_once_with( + "t1", + "tool", + 2, + 10, + query="weather", + tags=["search", "rag"], + language="zh_Hans", + ) assert result["plugins"][0]["id"] == "entity-1" assert result["plugins"][0]["plugin_unique_identifier"] == "test-author/test-plugin:1.0.0@checksum" assert result["builtin_tools"][0]["id"] == "builtin" assert result["builtin_tools"][0]["type"] == "builtin" assert result["has_more"] is True assert "total" not in result - builtin_mock.assert_called_once_with("t1") + builtin_mock.assert_called_once_with( + "t1", + query="weather", + tags=["search", "rag"], + language="zh_Hans", + ) + + def test_builtin_tool_providers_use_the_category_list_filters(self): + search_provider = ToolProviderApiEntity( + id="search-provider", + author="dify", + name="search-provider", + description=I18nObject(en_US="Search provider", zh_Hans="搜索工具"), + icon="icon.svg", + label=I18nObject(en_US="Search", zh_Hans="搜索"), + type=ToolProviderType.BUILT_IN, + labels=["search"], + ) + rag_provider = ToolProviderApiEntity( + id="rag-provider", + author="dify", + name="rag-provider", + description=I18nObject(en_US="RAG provider", zh_Hans="知识库工具"), + icon="icon.svg", + label=I18nObject(en_US="RAG", zh_Hans="知识库"), + type=ToolProviderType.BUILT_IN, + labels=["rag"], + ) + + with ( + patch("controllers.console.workspace.plugin.ToolManager.list_default_builtin_providers", return_value=[]), + patch( + "controllers.console.workspace.plugin.ToolManager.list_hardcoded_providers", + return_value=[MagicMock(), MagicMock()], + ), + patch("controllers.console.workspace.plugin.is_filtered", return_value=False), + patch( + "controllers.console.workspace.plugin.ToolTransformService.builtin_provider_to_user_provider", + side_effect=[search_provider, rag_provider], + ), + patch("controllers.console.workspace.plugin.ToolTransformService.repack_provider"), + patch( + "controllers.console.workspace.plugin.BuiltinToolProviderSort.sort", + side_effect=lambda providers: providers, + ), + ): + result = _list_hardcoded_builtin_tool_providers( + "t1", + query="搜索", + tags=["search", "weather"], + language="zh_Hans", + ) + + assert [provider["id"] for provider in result] == ["search-provider"] def test_non_tool_category_does_not_include_builtin_tools(self, app: Flask): api = PluginCategoryListApi() @@ -522,7 +604,7 @@ class TestPluginIconApi: app.test_request_context("/?tenant_id=t1&filename=a.png"), patch("controllers.console.workspace.plugin.PluginService.get_asset", return_value=(b"x", "image/png")), ): - response = method(api) + response = method(api, ParserIcon.model_validate({"tenant_id": "t1", "filename": "a.png"})) assert response.mimetype == "image/png" @@ -536,7 +618,9 @@ class TestPluginAssetApi: app.test_request_context("/?plugin_unique_identifier=p&file_name=a.bin"), patch("controllers.console.workspace.plugin.PluginService.extract_asset", return_value=b"x"), ): - response = method(api, "t1") + response = method( + api, ParserAsset.model_validate({"plugin_unique_identifier": "p", "file_name": "a.bin"}), "t1" + ) assert response.mimetype == "application/octet-stream" @@ -591,25 +675,34 @@ class TestPluginInstallFromPkgApi: "controllers.console.workspace.plugin.PluginService.install_from_local_pkg", return_value={"ok": True} ), ): - result = method(api, "t1") + result = method(api, ParserPluginIdentifiers.model_validate(payload), "t1") assert result["ok"] is True class TestPluginUninstallApi: - def test_uninstall(self, app: Flask): + @pytest.mark.parametrize("preserve_credentials", [False, True]) + def test_uninstall(self, app: Flask, preserve_credentials: bool): api = PluginUninstallApi() method = unwrap(api.post) - payload = {"plugin_installation_id": "x"} + payload = { + "plugin_installation_id": "x", + "preserve_credentials": preserve_credentials, + } with ( app.test_request_context("/", json=payload), - patch("controllers.console.workspace.plugin.PluginService.uninstall", return_value=True), + patch("controllers.console.workspace.plugin.PluginService.uninstall", return_value=True) as uninstall_mock, ): - result = method(api, "t1") + result = method(api, ParserUninstall.model_validate(payload), "t1") assert result["success"] is True + uninstall_mock.assert_called_once_with( + "t1", + "x", + preserve_credentials=preserve_credentials, + ) class TestPluginChangePermissionApi: @@ -628,7 +721,7 @@ class TestPluginChangePermissionApi: app.test_request_context("/", json=payload), ): with pytest.raises(Forbidden): - method(api, "t1", user) + method(api, ParserPermissionChange(), "t1", user) def test_change_permission_success(self, app: Flask): api = PluginChangePermissionApi() @@ -645,7 +738,7 @@ class TestPluginChangePermissionApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.plugin.PluginPermissionService.change_permission", return_value=True), ): - result = method(api, "t1", user) + result = method(api, ParserPermissionChange(), "t1", user) assert result["success"] is True @@ -676,7 +769,14 @@ class TestPluginFetchDynamicSelectOptionsApi: return_value=[_dynamic_option()], ), ): - result = method(api, "t1", user) + result = method( + api, + ParserDynamicOptions.model_validate( + {"plugin_id": "p", "provider": "x", "action": "y", "parameter": "z", "provider_type": "tool"} + ), + "t1", + user, + ) assert result == {"options": [_expected_dynamic_option_dump()]} @@ -690,7 +790,7 @@ class TestPluginReadmeApi: app.test_request_context("/?plugin_unique_identifier=p"), patch("controllers.console.workspace.plugin.PluginService.fetch_plugin_readme", return_value="readme"), ): - result = method(api, "t1") + result = method(api, ParserReadme.model_validate({"plugin_unique_identifier": "p"}), "t1") assert result["readme"] == "readme" @@ -709,7 +809,7 @@ class TestPluginListInstallationsFromIdsApi: return_value=[_plugin_installation()], ), ): - result = method(api, "t1") + result = method(api, ParserLatest.model_validate(payload), "t1") assert result == {"plugins": [_expected_plugin_installation_dump()]} @@ -726,10 +826,43 @@ class TestPluginListInstallationsFromIdsApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserLatest.model_validate(payload), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) +class TestPluginInstalledIdsApi: + def test_success(self, app: Flask): + api = PluginInstalledIdsApi() + method = unwrap(api.get) + + with ( + app.test_request_context("/?category=tool"), + patch( + "controllers.console.workspace.plugin.PluginService.list_installed_plugin_ids", + return_value=["langgenius/openai", "langgenius/anthropic"], + ) as list_installed_plugin_ids, + ): + result = method(api, PluginInstalledIdsQuery.model_validate({"category": "tool"}), "t1") + + assert result == {"plugin_ids": ["langgenius/openai", "langgenius/anthropic"]} + list_installed_plugin_ids.assert_called_once_with("t1", PluginCategory.Tool) + + def test_daemon_error(self, app: Flask): + api = PluginInstalledIdsApi() + method = unwrap(api.get) + + with ( + app.test_request_context("/?category=tool"), + patch( + "controllers.console.workspace.plugin.PluginService.list_installed_plugin_ids", + side_effect=PluginDaemonClientSideError("error"), + ), + ): + result = method(api, PluginInstalledIdsQuery.model_validate({"category": "tool"}), "t1") + + assert result == ({"code": "plugin_error", "message": "error"}, 400) + + class TestPluginUploadFromGithubApi: def test_success(self, app: Flask, user): api = PluginUploadFromGithubApi() @@ -743,7 +876,7 @@ class TestPluginUploadFromGithubApi: "controllers.console.workspace.plugin.PluginService.upload_pkg_from_github", return_value={"ok": True} ), ): - result = method(api, "t1") + result = method(api, ParserGithubUpload.model_validate(payload), "t1") assert result["ok"] is True @@ -760,7 +893,7 @@ class TestPluginUploadFromGithubApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserGithubUpload.model_validate(payload), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -829,7 +962,7 @@ class TestPluginInstallFromGithubApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.plugin.PluginService.install_from_github", return_value={"ok": True}), ): - result = method(api, "t1") + result = method(api, ParserGithubInstall.model_validate(payload), "t1") assert result["ok"] is True @@ -851,7 +984,7 @@ class TestPluginInstallFromGithubApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserGithubInstall.model_validate(payload), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -869,7 +1002,7 @@ class TestPluginInstallFromMarketplaceApi: return_value={"ok": True}, ), ): - result = method(api, "t1") + result = method(api, ParserPluginIdentifiers.model_validate(payload), "t1") assert result["ok"] is True @@ -886,7 +1019,7 @@ class TestPluginInstallFromMarketplaceApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserPluginIdentifiers.model_validate(payload), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -902,7 +1035,7 @@ class TestPluginFetchMarketplacePkgApi: return_value=_plugin_declaration(), ), ): - result = method(api, "t1") + result = method(api, ParserPluginIdentifierQuery.model_validate({"plugin_unique_identifier": "p"}), "t1") assert result == {"manifest": _expected_plugin_declaration_dump()} @@ -917,7 +1050,7 @@ class TestPluginFetchMarketplacePkgApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserPluginIdentifierQuery.model_validate({"plugin_unique_identifier": "p"}), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -932,7 +1065,7 @@ class TestPluginFetchManifestApi: app.test_request_context("/?plugin_unique_identifier=p"), patch("controllers.console.workspace.plugin.PluginService.fetch_plugin_manifest", return_value=manifest), ): - result = method(api, "t1") + result = method(api, ParserPluginIdentifierQuery.model_validate({"plugin_unique_identifier": "p"}), "t1") assert result == {"manifest": _expected_plugin_declaration_dump()} @@ -947,7 +1080,7 @@ class TestPluginFetchManifestApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserPluginIdentifierQuery.model_validate({"plugin_unique_identifier": "p"}), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -963,7 +1096,7 @@ class TestPluginFetchInstallTasksApi: return_value=[_plugin_task()], ), ): - result = method(api, "t1") + result = method(api, ParserTasks(), "t1") assert result == {"tasks": [_expected_plugin_task_dump()]} @@ -978,7 +1111,7 @@ class TestPluginFetchInstallTasksApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserTasks(), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -1113,7 +1246,7 @@ class TestPluginUpgradeFromMarketplaceApi: return_value={"ok": True}, ), ): - result = method(api, "t1") + result = method(api, ParserMarketplaceUpgrade.model_validate(payload), "t1") assert result["ok"] is True @@ -1133,7 +1266,7 @@ class TestPluginUpgradeFromMarketplaceApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserMarketplaceUpgrade.model_validate(payload), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -1157,7 +1290,7 @@ class TestPluginUpgradeFromGithubApi: return_value={"ok": True}, ), ): - result = method(api, "t1") + result = method(api, ParserGithubUpgrade.model_validate(payload), "t1") assert result["ok"] is True @@ -1180,7 +1313,7 @@ class TestPluginUpgradeFromGithubApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1") + result = method(api, ParserGithubUpgrade.model_validate(payload), "t1") assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -1205,7 +1338,7 @@ class TestPluginFetchDynamicSelectOptionsWithCredentialsApi: return_value=[_dynamic_option()], ), ): - result = method(api, "t1", user) + result = method(api, ParserDynamicOptionsWithCredentials.model_validate(payload), "t1", user) assert result == {"options": [_expected_dynamic_option_dump()]} @@ -1229,7 +1362,7 @@ class TestPluginFetchDynamicSelectOptionsWithCredentialsApi: side_effect=PluginDaemonClientSideError("error"), ), ): - result = method(api, "t1", user) + result = method(api, ParserDynamicOptionsWithCredentials.model_validate(payload), "t1", user) assert result == ({"code": "plugin_error", "message": "error"}, 400) @@ -1257,7 +1390,7 @@ class TestPluginChangeAutoUpgradeApi: "controllers.console.workspace.plugin.PluginAutoUpgradeService.change_strategy", return_value=True ) as change, ): - result = method(api, "t1", user) + result = method(api, ParserAutoUpgradeChange.model_validate(payload), "t1", user) assert result["success"] is True change.assert_called_once() @@ -1285,7 +1418,7 @@ class TestPluginChangeAutoUpgradeApi: "controllers.console.workspace.plugin.PluginAutoUpgradeService.change_strategy", return_value=True ) as change, ): - result = method(api, "t1", user) + result = method(api, ParserAutoUpgradeChange.model_validate(payload), "t1", user) assert result["success"] is True change.assert_called_once() @@ -1312,7 +1445,7 @@ class TestPluginChangeAutoUpgradeApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.plugin.PluginAutoUpgradeService.change_strategy", return_value=False), ): - result = method(api, "t1", user) + result = method(api, ParserAutoUpgradeChange.model_validate(payload), "t1", user) assert result["success"] is False @@ -1338,7 +1471,11 @@ class TestPluginFetchAutoUpgradeApi: return_value=auto_upgrade, ), ): - result = method(api, "t1") + result = method( + api, + ParserAutoUpgradeFetch.model_validate({"category": TenantPluginAutoUpgradeCategory.TOOL.value}), + "t1", + ) assert result["category"] == TenantPluginAutoUpgradeCategory.TOOL assert result["auto_upgrade"]["upgrade_time_of_day"] == 1 @@ -1358,7 +1495,11 @@ class TestPluginFetchAutoUpgradeApi: return_value=78300, ), ): - result = method(api, "t1") + result = method( + api, + ParserAutoUpgradeFetch.model_validate({"category": TenantPluginAutoUpgradeCategory.MODEL.value}), + "t1", + ) assert result == { "category": TenantPluginAutoUpgradeCategory.MODEL, @@ -1383,7 +1524,7 @@ class TestPluginAutoUpgradeExcludePluginApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.plugin.PluginAutoUpgradeService.exclude_plugin", return_value=True), ): - result = method(api, "t1") + result = method(api, ParserExcludePlugin.model_validate(payload), "t1") assert result["success"] is True @@ -1397,6 +1538,6 @@ class TestPluginAutoUpgradeExcludePluginApi: app.test_request_context("/", json=payload), patch("controllers.console.workspace.plugin.PluginAutoUpgradeService.exclude_plugin", return_value=False), ): - result = method(api, "t1") + result = method(api, ParserExcludePlugin.model_validate(payload), "t1") assert result["success"] is False diff --git a/api/tests/unit_tests/controllers/console/workspace/test_rbac.py b/api/tests/unit_tests/controllers/console/workspace/test_rbac.py index a118f803a91..5d6599e23a1 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_rbac.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_rbac.py @@ -24,7 +24,10 @@ from flask import Flask from pydantic import ValidationError from werkzeug.exceptions import Forbidden, NotFound +from configs import dify_config from controllers.console.workspace import rbac as rbac_mod +from controllers.console.workspace.rbac import _RolesListQuery +from enums import DeploymentEdition @pytest.fixture @@ -35,7 +38,8 @@ def app(): def _enabled(enabled: bool): - return patch("controllers.console.workspace.rbac.dify_config.ENTERPRISE_ENABLED", enabled) + deployment_edition = DeploymentEdition.ENTERPRISE if enabled else DeploymentEdition.COMMUNITY + return patch("controllers.console.workspace.rbac.dify_config.DEPLOYMENT_EDITION", deployment_edition) class TestCurrentIds: @@ -51,6 +55,27 @@ class TestCurrentIds: assert rbac_mod._current_ids() == ("tenant-1", "acct-1") +class TestMyPermissions: + def test_returns_app_deploy_permission(self, app): + permissions = rbac_mod.svc.MyPermissionsResponse( + app=rbac_mod.svc.ResourcePermissionSnapshot( + default_permission_keys=["app.acl.deploy"], + ) + ) + with ( + app.test_request_context("/workspaces/current/rbac/my-permissions"), + patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")), + patch( + "controllers.console.workspace.rbac.svc.RBACService.MyPermissions.get", + return_value=permissions, + ) as mock_get, + ): + response = inspect.unwrap(rbac_mod.RBACMyPermissionsApi.get)(rbac_mod.RBACMyPermissionsApi()) + + assert response["app"]["default_permission_keys"] == ["app.acl.deploy"] + mock_get.assert_called_once() + + class TestAccessMatrixAccountNames: def test_hydrates_missing_account_names(self): items = [ @@ -174,7 +199,24 @@ class TestPaginationMapping: patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")), patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list") as mock_list, ): - response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi()) + response = inspect.unwrap(rbac_mod.RBACRolesApi.get)( + rbac_mod.RBACRolesApi(), + _RolesListQuery.model_validate({"page": 1, "limit": 2, "include_owner": 1}), + ) + + owner_permission_keys = rbac_mod._LEGACY_ROLE_PERMISSION_KEYS["owner"] + valid_owner_permission_keys = [] + for permission_key in owner_permission_keys: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD and "billing" in permission_key: + continue + valid_owner_permission_keys.append(permission_key) + + admin_permission_keys = rbac_mod._LEGACY_ROLE_PERMISSION_KEYS["admin"] + valid_admin_permission_keys = [] + for permission_key in admin_permission_keys: + if dify_config.DEPLOYMENT_EDITION != DeploymentEdition.CLOUD and "billing" in permission_key: + continue + valid_admin_permission_keys.append(permission_key) assert response["data"] == [ { @@ -185,7 +227,7 @@ class TestPaginationMapping: "name": "owner", "description": "", "is_builtin": True, - "permission_keys": list(dict.fromkeys(rbac_mod._LEGACY_ROLE_PERMISSION_KEYS["owner"])), + "permission_keys": valid_owner_permission_keys, "role_tag": "owner", }, { @@ -196,7 +238,7 @@ class TestPaginationMapping: "name": "admin", "description": "", "is_builtin": True, - "permission_keys": list(dict.fromkeys(rbac_mod._LEGACY_ROLE_PERMISSION_KEYS["admin"])), + "permission_keys": valid_admin_permission_keys, "role_tag": "", }, ] @@ -215,7 +257,7 @@ class TestPaginationMapping: patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")), patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list"), ): - response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi()) + response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi(), _RolesListQuery()) names = [r["name"] for r in response["data"]] assert "owner" not in names @@ -227,7 +269,10 @@ class TestPaginationMapping: patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")), patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list"), ): - response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi()) + response = inspect.unwrap(rbac_mod.RBACRolesApi.get)( + rbac_mod.RBACRolesApi(), + _RolesListQuery.model_validate({"include_owner": 1}), + ) names = [r["name"] for r in response["data"]] assert "owner" in names @@ -239,7 +284,7 @@ class TestPaginationMapping: patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-1")), patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list"), ): - response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi()) + response = inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi(), _RolesListQuery()) names = [r["name"] for r in response["data"]] assert "owner" not in names @@ -252,7 +297,10 @@ class TestPaginationMapping: patch("controllers.console.workspace.rbac.svc.RBACService.Roles.list") as mock_list, patch("controllers.console.workspace.rbac._dump", return_value={}), ): - inspect.unwrap(rbac_mod.RBACRolesApi.get)(rbac_mod.RBACRolesApi()) + inspect.unwrap(rbac_mod.RBACRolesApi.get)( + rbac_mod.RBACRolesApi(), + _RolesListQuery.model_validate({"page": 2, "limit": 50, "reverse": True, "include_owner": 1}), + ) _, kwargs = mock_list.call_args options = kwargs["options"] @@ -281,6 +329,46 @@ class TestResourceAccessScopeBindings: mock_sync_task.delay.assert_called_once_with("tenant-1", "acct-actor", "app-1") + def test_dataset_whitelist_all_schedules_member_policy_sync(self, app): + # Widening a dataset to the whole workspace only records the scope; without granting the + # default policy to the current members nobody actually gains access. + with ( + app.test_request_context( + "/workspaces/current/rbac/datasets/dataset-1/whitelist", + method="PUT", + json={"scope": "all"}, + ), + patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-actor")), + patch("controllers.console.workspace.rbac.dify_config.RBAC_ENABLED", True), + patch( + "controllers.console.workspace.rbac.svc.RBACService.DatasetAccess.replace_whitelist", + return_value=rbac_mod.svc.ResourceWhitelist(), + ), + patch("controllers.console.workspace.rbac.initialize_created_app_rbac_access_task") as mock_sync_task, + ): + inspect.unwrap(rbac_mod.RBACDatasetWhitelistApi.put)(rbac_mod.RBACDatasetWhitelistApi(), "dataset-1") + + mock_sync_task.delay.assert_called_once_with("tenant-1", "acct-actor", dataset_id="dataset-1") + + def test_dataset_whitelist_specific_does_not_schedule_member_policy_sync(self, app): + with ( + app.test_request_context( + "/workspaces/current/rbac/datasets/dataset-1/whitelist", + method="PUT", + json={"scope": "specific"}, + ), + patch("controllers.console.workspace.rbac._current_ids", return_value=("tenant-1", "acct-actor")), + patch("controllers.console.workspace.rbac.dify_config.RBAC_ENABLED", True), + patch( + "controllers.console.workspace.rbac.svc.RBACService.DatasetAccess.replace_whitelist", + return_value=rbac_mod.svc.ResourceWhitelist(), + ), + patch("controllers.console.workspace.rbac.initialize_created_app_rbac_access_task") as mock_sync_task, + ): + inspect.unwrap(rbac_mod.RBACDatasetWhitelistApi.put)(rbac_mod.RBACDatasetWhitelistApi(), "dataset-1") + + mock_sync_task.delay.assert_not_called() + def test_app_user_access_policy_assignment_forwards_ids(self, app): with ( app.test_request_context( diff --git a/api/tests/unit_tests/controllers/console/workspace/test_snippets.py b/api/tests/unit_tests/controllers/console/workspace/test_snippets.py index ee569e97f34..211d9229a2a 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_snippets.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_snippets.py @@ -13,7 +13,7 @@ from services.snippet_dsl_service import ImportStatus, SnippetImportInfo @pytest.fixture(autouse=True) -def _patch_snippet_service_factory(monkeypatch): +def _patch_snippet_service_factory(monkeypatch: pytest.MonkeyPatch): def factory(): return snippets_module.SnippetService.__new__(snippets_module.SnippetService) @@ -154,19 +154,14 @@ def test_create_snippet_defaults_unknown_type_and_returns_created(app: Flask, mo snippet = _snippet() create_snippet = Mock(return_value=snippet) monkeypatch.setattr(snippets_module.SnippetService, "create_snippet", create_snippet) - monkeypatch.setattr( - snippets_module.CreateSnippetPayload, - "model_validate", - Mock( - return_value=SimpleNamespace( - name="Snippet", - type="unknown", - description="Description", - graph=None, - icon_info=None, - input_fields=[], - ) - ), + + req_data = SimpleNamespace( + name="Snippet", + type="unknown", + description="Description", + graph=None, + icon_info=None, + input_fields=[], ) api = snippets_module.CustomizedSnippetsApi() @@ -177,7 +172,7 @@ def test_create_snippet_defaults_unknown_type_and_returns_created(app: Flask, mo method="POST", json={"name": "Snippet", "type": "node", "description": "Description"}, ): - response, status_code = handler(api, "tenant-1", user) + response, status_code = handler(api, req_data, "tenant-1", user) assert status_code == 201 assert response["id"] == "snippet-1" @@ -190,6 +185,17 @@ def test_create_snippet_rejects_forbidden_nodes(app: Flask, monkeypatch: pytest. create_snippet = Mock() monkeypatch.setattr(snippets_module.SnippetService, "create_snippet", create_snippet) + req_data = snippets_module.CreateSnippetPayload( + name="snippet with invalid node", + type="node", + graph={ + "nodes": [ + {"id": "knowledge-1", "data": {"type": "knowledge-retrieval"}}, + ], + "edges": [], + }, + ) + api = snippets_module.CustomizedSnippetsApi() handler = unwrap(api.post) @@ -207,7 +213,7 @@ def test_create_snippet_rejects_forbidden_nodes(app: Flask, monkeypatch: pytest. }, }, ): - response, status_code = handler(api, "tenant-1", user) + response, status_code = handler(api, req_data, "tenant-1", user) assert status_code == 400 assert "knowledge-retrieval" in response["message"] @@ -245,6 +251,8 @@ def test_patch_snippet_returns_400_for_empty_payload(app: Flask, monkeypatch: py user = _account("user-1") monkeypatch.setattr(snippets_module.SnippetService, "get_snippet_by_id", Mock(return_value=snippet)) + req_data = snippets_module.UpdateSnippetPayload() + api = snippets_module.CustomizedSnippetDetailApi() handler = unwrap(api.patch) @@ -253,7 +261,7 @@ def test_patch_snippet_returns_400_for_empty_payload(app: Flask, monkeypatch: py method="PATCH", json={}, ): - response, status_code = handler(api, "tenant-1", user, snippet_id="snippet-1") + response, status_code = handler(api, req_data, "tenant-1", user, snippet_id="snippet-1") assert status_code == 400 assert response == {"message": "No valid fields to update"} @@ -275,6 +283,8 @@ def test_patch_snippet_updates_and_commits(app: Flask, monkeypatch: pytest.Monke monkeypatch.setattr(snippets_module, "Session", SessionContext) monkeypatch.setattr(snippets_module, "db", SimpleNamespace(engine=object())) + req_data = snippets_module.UpdateSnippetPayload(name="New", icon_info={"icon": "star"}) + api = snippets_module.CustomizedSnippetDetailApi() handler = unwrap(api.patch) @@ -283,7 +293,7 @@ def test_patch_snippet_updates_and_commits(app: Flask, monkeypatch: pytest.Monke method="PATCH", json={"name": "New", "icon_info": {"icon": "star"}}, ): - response, status_code = handler(api, "tenant-1", user, snippet_id="snippet-1") + response, status_code = handler(api, req_data, "tenant-1", user, snippet_id="snippet-1") assert status_code == 200 assert response["id"] == "snippet-1" @@ -378,6 +388,8 @@ def test_import_snippet_returns_202_for_pending_confirmation(app: Flask, monkeyp Mock(return_value=SimpleNamespace(import_snippet=import_snippet)), ) + req_data = snippets_module.SnippetImportPayload(mode="yaml-content", yaml_content="kind: snippet") + api = snippets_module.CustomizedSnippetImportApi() handler = unwrap(api.post) @@ -386,7 +398,7 @@ def test_import_snippet_returns_202_for_pending_confirmation(app: Flask, monkeyp method="POST", json={"mode": "yaml-content", "yaml_content": "kind: snippet"}, ): - response, status_code = handler(api, session, user) + response, status_code = handler(api, req_data, session, user) assert status_code == 202 assert response["status"] == ImportStatus.PENDING.value @@ -413,6 +425,8 @@ def test_import_snippet_returns_400_for_failed_import(app: Flask, monkeypatch: p Mock(return_value=SimpleNamespace(import_snippet=import_snippet)), ) + req_data = snippets_module.SnippetImportPayload(mode="yaml-content", yaml_content="kind: snippet") + api = snippets_module.CustomizedSnippetImportApi() handler = unwrap(api.post) @@ -421,7 +435,7 @@ def test_import_snippet_returns_400_for_failed_import(app: Flask, monkeypatch: p method="POST", json={"mode": "yaml-content", "yaml_content": "kind: snippet"}, ): - response, status_code = handler(api, session, user) + response, status_code = handler(api, req_data, session, user) assert status_code == 400 assert response["error"] == "Invalid DSL" diff --git a/api/tests/unit_tests/controllers/console/workspace/test_tool_provider_apis.py b/api/tests/unit_tests/controllers/console/workspace/test_tool_provider_apis.py index 696ae70080c..1f9f4d863b4 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_tool_provider_apis.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_tool_provider_apis.py @@ -10,6 +10,13 @@ from flask import Flask from werkzeug.exceptions import Forbidden from controllers.console.workspace.tool_providers import ( + ApiToolProviderAddPayload, + ApiToolProviderDeletePayload, + ApiToolProviderUpdatePayload, + BuiltinProviderDefaultCredentialPayload, + BuiltinToolAddPayload, + BuiltinToolCredentialDeletePayload, + BuiltinToolUpdatePayload, ToolApiListApi, ToolApiProviderAddApi, ToolApiProviderDeleteApi, @@ -32,6 +39,7 @@ from controllers.console.workspace.tool_providers import ( ToolLabelsApi, ToolOAuthCallback, ToolOAuthCustomClient, + ToolOAuthCustomClientPayload, ToolPluginOAuthApi, ToolProviderListApi, ToolWorkflowListApi, @@ -39,6 +47,9 @@ from controllers.console.workspace.tool_providers import ( ToolWorkflowProviderDeleteApi, ToolWorkflowProviderGetApi, ToolWorkflowProviderUpdateApi, + WorkflowToolCreatePayload, + WorkflowToolDeletePayload, + WorkflowToolUpdatePayload, is_valid_url, ) from core.tools.entities.api_entities import ToolProviderApiEntity as CoreToolProviderApiEntity @@ -276,7 +287,8 @@ class TestBuiltinProviderApis: return_value={"result": "success"}, ), ): - assert method(api, "t1", "provider")["result"] == "success" + req = BuiltinToolCredentialDeletePayload(credential_id="cid") + assert method(api, req, "t1", "provider")["result"] == "success" def test_add_invalid_type(self, app: Flask) -> None: api = ToolBuiltinProviderAddApi() @@ -286,7 +298,13 @@ class TestBuiltinProviderApis: app.test_request_context("/", json={"credentials": empty_mapping(), "type": "invalid"}), ): with pytest.raises(ValueError): - method(api, "t", make_account(), "provider") + method( + api, + BuiltinToolAddPayload(credentials=empty_mapping(), type="invalid"), + "t", + make_account(), + "provider", + ) def test_add_success(self, app: Flask) -> None: api = ToolBuiltinProviderAddApi() @@ -301,7 +319,7 @@ class TestBuiltinProviderApis: return_value={"result": "success"}, ), ): - assert method(api, "t", make_account(), "provider")["result"] == "success" + assert method(api, BuiltinToolAddPayload(**payload), "t", make_account(), "provider")["result"] == "success" def test_update(self, app: Flask) -> None: api = ToolBuiltinProviderUpdateApi() @@ -316,7 +334,8 @@ class TestBuiltinProviderApis: return_value={"result": "success"}, ), ): - assert method(api, "t", make_account(), "provider")["result"] == "success" + req = BuiltinToolUpdatePayload(**payload) + assert method(api, req, "t", make_account(), "provider")["result"] == "success" def test_get_credentials(self, app: Flask) -> None: api = ToolBuiltinProviderGetCredentialsApi() @@ -369,7 +388,8 @@ class TestBuiltinProviderApis: return_value={"result": "success"}, ), ): - assert method(api, "t", "provider")["result"] == "success" + req = BuiltinProviderDefaultCredentialPayload(id="c1") + assert method(api, req, "t", "provider")["result"] == "success" def test_get_credential_info(self, app: Flask) -> None: api = ToolBuiltinProviderGetCredentialInfoApi() @@ -418,7 +438,7 @@ class TestApiProviderApis: return_value={"result": "success"}, ) as create_api_tool_provider, ): - assert method(api, "t", make_account()) == {"result": "success"} + assert method(api, ApiToolProviderAddPayload(**payload), "t", make_account()) == {"result": "success"} create_api_tool_provider.assert_called_once() assert create_api_tool_provider.call_args.args[3] == emoji_icon() @@ -472,7 +492,7 @@ class TestApiProviderApis: return_value={"result": "success"}, ) as update_api_tool_provider, ): - assert method(api, "t", make_account()) == {"result": "success"} + assert method(api, ApiToolProviderUpdatePayload(**payload), "t", make_account()) == {"result": "success"} update_api_tool_provider.assert_called_once() assert update_api_tool_provider.call_args.args[4] == emoji_icon() @@ -488,7 +508,8 @@ class TestApiProviderApis: return_value={"result": "success"}, ), ): - assert method(api, "t", make_account())["result"] == "success" + req = ApiToolProviderDeletePayload(provider="p") + assert method(api, req, "t", make_account())["result"] == "success" def test_get(self, app: Flask) -> None: api = ToolApiProviderGetApi() @@ -525,7 +546,7 @@ class TestWorkflowApis: return_value={"result": "success"}, ) as create_workflow_tool, ): - assert method(api, "t", make_account()) == {"result": "success"} + assert method(api, WorkflowToolCreatePayload(**payload), "t", make_account()) == {"result": "success"} create_workflow_tool.assert_called_once() assert create_workflow_tool.call_args.kwargs["icon"] == emoji_icon() @@ -549,7 +570,7 @@ class TestWorkflowApis: return_value={"result": "success"}, ) as update_workflow_tool, ): - result = method(api, "t", make_account()) + result = method(api, WorkflowToolUpdatePayload(**payload), "t", make_account()) assert result == {"result": "success"} update_workflow_tool.assert_called_once() @@ -566,7 +587,8 @@ class TestWorkflowApis: return_value={"result": "success"}, ), ): - assert method(api, "t", make_account())["result"] == "success" + req = WorkflowToolDeletePayload(workflow_tool_id="123e4567-e89b-12d3-a456-426614174000") + assert method(api, req, "t", make_account())["result"] == "success" def test_get_error(self, app: Flask) -> None: api = ToolWorkflowProviderGetApi() @@ -671,7 +693,8 @@ class TestOAuthCustomClient: return_value={"result": "success"}, ), ): - assert method(api, "t", "provider") == {"result": "success"} + req = ToolOAuthCustomClientPayload(client_params={"a": 1}) + assert method(api, req, "t", "provider") == {"result": "success"} def test_get_custom_client(self, app: Flask) -> None: api = ToolOAuthCustomClient() diff --git a/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py b/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py index 845d596294e..5d5f47333c0 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_tool_providers.py @@ -18,6 +18,7 @@ from sqlalchemy.orm import Session, scoped_session, sessionmaker from core.tools.entities.api_entities import ToolProviderApiEntity as CoreToolProviderApiEntity from core.tools.entities.common_entities import I18nObject from core.tools.entities.tool_entities import ToolParameter +from enums import DeploymentEdition from models import Account, BuiltinToolProvider, Tenant, TenantAccountJoin from models.account import TenantAccountRole from models.credential_permission import CredentialPermission @@ -74,8 +75,8 @@ def controller_module(monkeypatch: pytest.MonkeyPatch): global _WRAPS_MODULE wraps_module = importlib.import_module("controllers.console.wraps") _WRAPS_MODULE = wraps_module - monkeypatch.setattr(module.dify_config, "EDITION", "CLOUD") - monkeypatch.setattr(wraps_module.dify_config, "EDITION", "CLOUD") + monkeypatch.setattr(module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) + monkeypatch.setattr(wraps_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) login_module = importlib.import_module("libs.login") monkeypatch.setattr(login_module, "check_csrf_token", lambda *args, **kwargs: None) @@ -720,7 +721,7 @@ def test_tool_labels_list(app: Flask, controller_module, monkeypatch: pytest.Mon def test_resolve_identity_mode_none_keeps_current_when_enterprise(controller_module, monkeypatch: pytest.MonkeyPatch): """None means 'leave unchanged' — fall back to the stored mode (update path).""" identity_mode = importlib.import_module("core.entities.mcp_provider").IdentityMode - monkeypatch.setattr(controller_module.dify_config, "ENTERPRISE_ENABLED", True) + monkeypatch.setattr(controller_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) resolved = controller_module._resolve_identity_mode(None, current=identity_mode.IDP_TOKEN) @@ -730,7 +731,7 @@ def test_resolve_identity_mode_none_keeps_current_when_enterprise(controller_mod def test_resolve_identity_mode_explicit_value_overrides_current(controller_module, monkeypatch: pytest.MonkeyPatch): """An explicit value wins over the stored mode.""" identity_mode = importlib.import_module("core.entities.mcp_provider").IdentityMode - monkeypatch.setattr(controller_module.dify_config, "ENTERPRISE_ENABLED", True) + monkeypatch.setattr(controller_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) resolved = controller_module._resolve_identity_mode(identity_mode.OFF, current=identity_mode.IDP_TOKEN) @@ -743,7 +744,7 @@ def test_resolve_identity_mode_coerces_non_off_to_off_when_not_enterprise( """Gate: a non-EE deployment must never persist a non-OFF mode — the runtime won't forward, so the stored row must not imply it does.""" identity_mode = importlib.import_module("core.entities.mcp_provider").IdentityMode - monkeypatch.setattr(controller_module.dify_config, "ENTERPRISE_ENABLED", False) + monkeypatch.setattr(controller_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) # Both an explicit idp_token request AND an inherited non-OFF current # must collapse to OFF. @@ -759,6 +760,6 @@ def test_resolve_identity_mode_off_is_passthrough_when_not_enterprise( ): """OFF is always fine — the gate only neutralizes non-OFF values.""" identity_mode = importlib.import_module("core.entities.mcp_provider").IdentityMode - monkeypatch.setattr(controller_module.dify_config, "ENTERPRISE_ENABLED", False) + monkeypatch.setattr(controller_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) assert controller_module._resolve_identity_mode(None, current=identity_mode.OFF) == identity_mode.OFF diff --git a/api/tests/unit_tests/controllers/console/workspace/test_trigger_provider_apis.py b/api/tests/unit_tests/controllers/console/workspace/test_trigger_provider_apis.py index 6eca14aa273..074f8d527c2 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_trigger_provider_apis.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_trigger_provider_apis.py @@ -4,7 +4,7 @@ from __future__ import annotations from datetime import datetime from inspect import unwrap -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask @@ -15,15 +15,19 @@ from controllers.console.workspace.trigger_providers import ( TriggerOAuthAuthorizeApi, TriggerOAuthCallbackApi, TriggerOAuthClientManageApi, + TriggerOAuthClientPayload, TriggerProviderIconApi, TriggerProviderInfoApi, TriggerProviderListApi, TriggerSubscriptionBuilderBuildApi, TriggerSubscriptionBuilderCreateApi, + TriggerSubscriptionBuilderCreatePayload, TriggerSubscriptionBuilderGetApi, TriggerSubscriptionBuilderLogsApi, TriggerSubscriptionBuilderUpdateApi, + TriggerSubscriptionBuilderUpdatePayload, TriggerSubscriptionBuilderVerifyApi, + TriggerSubscriptionBuilderVerifyPayload, TriggerSubscriptionListApi, TriggerSubscriptionUpdateApi, TriggerSubscriptionVerifyApi, @@ -163,7 +167,13 @@ class TestTriggerSubscriptionBuilderApis: return_value=subscription_builder(), ), ): - result = method(api, "t1", mock_user(), "github") + result = method( + api, + TriggerSubscriptionBuilderCreatePayload(credential_type="UNAUTHORIZED"), + "t1", + mock_user(), + "github", + ) assert result["subscription_builder"]["id"] == "b1" def test_get_builder(self, app: Flask) -> None: @@ -175,9 +185,15 @@ class TestTriggerSubscriptionBuilderApis: patch( "controllers.console.workspace.trigger_providers.TriggerSubscriptionBuilderService.get_subscription_builder_by_id", return_value=subscription_builder(), - ), + ) as mock_get_builder, ): - assert method(api, "github", "b1")["id"] == "b1" + assert method(api, "t1", mock_user(), "github", "b1")["id"] == "b1" + mock_get_builder.assert_called_once_with( + tenant_id="t1", + user_id="u1", + provider_id=ANY, + subscription_builder_id="b1", + ) def test_verify_builder(self, app: Flask) -> None: api = TriggerSubscriptionBuilderVerifyApi() @@ -190,7 +206,14 @@ class TestTriggerSubscriptionBuilderApis: return_value={"verified": True}, ), ): - assert method(api, "t1", mock_user(), "github", "b1") == {"verified": True} + assert method( + api, + TriggerSubscriptionBuilderVerifyPayload(credentials={"a": 1}), + "t1", + mock_user(), + "github", + "b1", + ) == {"verified": True} def test_verify_builder_error(self, app: Flask) -> None: api = TriggerSubscriptionBuilderVerifyApi() @@ -204,7 +227,14 @@ class TestTriggerSubscriptionBuilderApis: ), ): with pytest.raises(ValueError): - method(api, "t1", mock_user(), "github", "b1") + method( + api, + TriggerSubscriptionBuilderVerifyPayload(credentials={}), + "t1", + mock_user(), + "github", + "b1", + ) def test_update_builder(self, app: Flask) -> None: api = TriggerSubscriptionBuilderUpdateApi() @@ -215,9 +245,26 @@ class TestTriggerSubscriptionBuilderApis: patch( "controllers.console.workspace.trigger_providers.TriggerSubscriptionBuilderService.update_trigger_subscription_builder", return_value=subscription_builder(), - ), + ) as mock_update_builder, ): - assert method(api, "t1", "github", "b1")["id"] == "b1" + assert ( + method( + api, + TriggerSubscriptionBuilderUpdatePayload(name="n"), + "t1", + mock_user(), + "github", + "b1", + )["id"] + == "b1" + ) + mock_update_builder.assert_called_once_with( + tenant_id="t1", + user_id="u1", + provider_id=ANY, + subscription_builder_id="b1", + subscription_builder_updater=ANY, + ) def test_logs(self, app: Flask) -> None: api = TriggerSubscriptionBuilderLogsApi() @@ -228,10 +275,16 @@ class TestTriggerSubscriptionBuilderApis: patch( "controllers.console.workspace.trigger_providers.TriggerSubscriptionBuilderService.list_logs", return_value=[request_log()], - ), + ) as mock_list_logs, ): - result = method(api, "github", "b1") + result = method(api, "t1", mock_user(), "github", "b1") assert result["logs"][0]["id"] == "log1" + mock_list_logs.assert_called_once_with( + tenant_id="t1", + user_id="u1", + provider_id=ANY, + subscription_builder_id="b1", + ) def test_build(self, app: Flask) -> None: api = TriggerSubscriptionBuilderBuildApi() @@ -244,7 +297,14 @@ class TestTriggerSubscriptionBuilderApis: return_value=None, ), ): - assert method(api, "t1", mock_user(), "github", "b1") == {"result": "success"} + assert method( + api, + TriggerSubscriptionBuilderUpdatePayload(name="x"), + "t1", + mock_user(), + "github", + "b1", + ) == {"result": "success"} class TestTriggerSubscriptionCrud: @@ -264,7 +324,12 @@ class TestTriggerSubscriptionCrud: ), patch("controllers.console.workspace.trigger_providers.TriggerProviderService.update_trigger_subscription"), ): - assert method(api, "t1", "s1") == {"result": "success"} + assert method( + api, + TriggerSubscriptionBuilderUpdatePayload(name="x"), + "t1", + "s1", + ) == {"result": "success"} def test_update_not_found(self, app: Flask) -> None: api = TriggerSubscriptionUpdateApi() @@ -278,7 +343,7 @@ class TestTriggerSubscriptionCrud: ), ): with pytest.raises(NotFoundError): - method(api, "t1", "x") + method(api, TriggerSubscriptionBuilderUpdatePayload(name="x"), "t1", "x") def test_update_rebuild(self, app: Flask) -> None: api = TriggerSubscriptionUpdateApi() @@ -300,7 +365,12 @@ class TestTriggerSubscriptionCrud: "controllers.console.workspace.trigger_providers.TriggerProviderService.rebuild_trigger_subscription" ), ): - assert method(api, "t1", "s1") == {"result": "success"} + assert method( + api, + TriggerSubscriptionBuilderUpdatePayload(credentials={}), + "t1", + "s1", + ) == {"result": "success"} class TestTriggerOAuthApis: @@ -377,10 +447,17 @@ class TestTriggerOAuthApis: ), patch( "controllers.console.workspace.trigger_providers.TriggerSubscriptionBuilderService.update_trigger_subscription_builder" - ), + ) as mock_update_builder, ): resp = method(api, "github") assert resp.status_code == 302 + mock_update_builder.assert_called_once_with( + tenant_id="t1", + user_id="u1", + provider_id=ANY, + subscription_builder_id="b1", + subscription_builder_updater=ANY, + ) def test_oauth_callback_no_oauth_client(self, app: Flask) -> None: api = TriggerOAuthCallbackApi() @@ -473,7 +550,12 @@ class TestTriggerOAuthClientManageApi: return_value={"result": "success"}, ), ): - assert method(api, "t1", "github") == {"result": "success"} + assert method( + api, + TriggerOAuthClientPayload(enabled=True), + "t1", + "github", + ) == {"result": "success"} def test_delete_client(self, app: Flask) -> None: api = TriggerOAuthClientManageApi() @@ -500,7 +582,7 @@ class TestTriggerOAuthClientManageApi: ), ): with pytest.raises(BadRequest): - method(api, "t1", "github") + method(api, TriggerOAuthClientPayload(enabled=True), "t1", "github") class TestTriggerSubscriptionVerifyApi: @@ -515,7 +597,14 @@ class TestTriggerSubscriptionVerifyApi: return_value={"verified": True}, ), ): - assert method(api, "t1", mock_user(), "github", "s1") == {"verified": True} + assert method( + api, + TriggerSubscriptionBuilderVerifyPayload(credentials={}), + "t1", + mock_user(), + "github", + "s1", + ) == {"verified": True} @pytest.mark.parametrize("raised_exception", [ValueError("bad"), Exception("boom")]) def test_verify_errors(self, app: Flask, raised_exception: Exception) -> None: @@ -530,4 +619,11 @@ class TestTriggerSubscriptionVerifyApi: ), ): with pytest.raises(BadRequest): - method(api, "t1", mock_user(), "github", "s1") + method( + api, + TriggerSubscriptionBuilderVerifyPayload(credentials={}), + "t1", + mock_user(), + "github", + "s1", + ) diff --git a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py index 436d9469992..dac0b09241d 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_workspace.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_workspace.py @@ -1,16 +1,18 @@ import logging from collections.abc import Iterator +from datetime import timedelta from http import HTTPStatus from inspect import unwrap from io import BytesIO -from unittest.mock import ANY, MagicMock, patch +from types import SimpleNamespace +from unittest.mock import MagicMock, patch import pytest from flask import Flask from sqlalchemy import Engine, event from sqlalchemy.orm import Session, scoped_session, sessionmaker from werkzeug.datastructures import FileStorage -from werkzeug.exceptions import Unauthorized +from werkzeug.exceptions import NotFound import services from controllers.common.errors import ( @@ -20,11 +22,13 @@ from controllers.common.errors import ( TooManyFilesError, UnsupportedFileTypeError, ) +from controllers.console import console_ns from controllers.console.error import AccountNotLinkTenantError +from controllers.console.workspace.error import CurrentWorkspaceArchivedError from controllers.console.workspace.workspace import ( + CurrentWorkspaceSummaryApi, CustomConfigWorkspaceApi, SwitchWorkspaceApi, - TenantApi, TenantInfoResponse, TenantListApi, WebappLogoWorkspaceApi, @@ -34,15 +38,19 @@ from controllers.console.workspace.workspace import ( WorkspacePermissionApi, WorkspacePermissionResponse, ) -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from libs.datetime_utils import naive_utc_now -from models.account import Account, Tenant, TenantCustomConfigDict, TenantStatus +from machinery.context import RequestContext +from models.account import Account, Tenant, TenantAccountJoin, TenantCustomConfigDict, TenantStatus +from repositories.workspace_query_repository import WorkspaceQueryRepository +from services import workspace_plan_gateway +from services.workspace_query_service import WorkspaceQueryService, WorkspaceRecord @pytest.fixture def workspace_session(sqlite_engine: Engine) -> Iterator[scoped_session[Session]]: """Provide the callable scoped session expected by Flask-SQLAlchemy controllers.""" - Tenant.metadata.create_all(sqlite_engine, tables=[Tenant.__table__]) + Tenant.metadata.create_all(sqlite_engine, tables=[Tenant.__table__, TenantAccountJoin.__table__]) session = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False)) try: yield session @@ -50,6 +58,33 @@ def workspace_session(sqlite_engine: Engine) -> Iterator[scoped_session[Session] session.remove() +@pytest.fixture +def workspace_plan_dependencies(monkeypatch: pytest.MonkeyPatch) -> tuple[MagicMock, MagicMock]: + get_plan_bulk = MagicMock() + get_features = MagicMock() + monkeypatch.setattr(workspace_plan_gateway.BillingService, "get_plan_bulk", get_plan_bulk) + monkeypatch.setattr(workspace_plan_gateway.FeatureService, "get_features", get_features) + return get_plan_bulk, get_features + + +def configure_workspace_plans( + monkeypatch: pytest.MonkeyPatch, + *, + edition: DeploymentEdition = DeploymentEdition.CLOUD, +) -> None: + monkeypatch.setattr( + workspace_plan_gateway, + "dify_config", + SimpleNamespace( + DEPLOYMENT_EDITION=edition, + ), + ) + + +def features_with_plan(plan: str) -> SimpleNamespace: + return SimpleNamespace(billing=SimpleNamespace(subscription=SimpleNamespace(plan=plan))) + + def make_account(account_id: str = "u1") -> Account: account = Account(name="Test User", email=f"{account_id}@example.com") account.id = account_id @@ -71,12 +106,6 @@ def make_tenant( return tenant -def make_membership(*, last_opened_at=None) -> MagicMock: - membership = MagicMock() - membership.last_opened_at = last_opened_at - return membership - - def make_account_with_tenant(tenant: Tenant) -> Account: account = make_account() account._current_tenant = tenant @@ -84,188 +113,190 @@ def make_account_with_tenant(tenant: Tenant) -> Account: class TestTenantListApi: - def test_get_success_saas_path(self, app: Flask): + def test_get_passes_context_and_serializes_workspaces(self): api = TenantListApi() method = unwrap(api.get) - tenant1 = make_tenant("t1", name="Tenant 1") - tenant2 = make_tenant("t2", name="Tenant 2") + request_context = RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id="account-1", + active_workspace_id="workspace-1", + ) + created_at = naive_utc_now() last_opened_at = naive_utc_now() - user = make_account() - with ( - app.test_request_context("/workspaces"), - patch( - "controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", - return_value=[(tenant1, make_membership(last_opened_at=last_opened_at)), (tenant2, make_membership())], + workspaces = MagicMock() + workspaces.list_for_account.return_value = ( + WorkspaceRecord( + id="workspace-1", + name="Workspace 1", + status=TenantStatus.NORMAL.value, + created_at=created_at, + last_opened_at=last_opened_at, ), - patch("controllers.console.workspace.workspace.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", True), - patch("controllers.console.workspace.workspace.dify_config.EDITION", "CLOUD"), - patch( - "controllers.console.workspace.workspace.BillingService.get_plan_bulk", - return_value={ - "t1": {"plan": CloudPlan.TEAM, "expiration_date": 0}, - "t2": {"plan": CloudPlan.PROFESSIONAL, "expiration_date": 0}, + WorkspaceRecord( + id="workspace-2", + name=None, + status=TenantStatus.NORMAL.value, + created_at=created_at, + last_opened_at=None, + ), + ) + plans = MagicMock() + plans.resolve_many.return_value = {"workspace-1": CloudPlan.TEAM} + workspace_queries = WorkspaceQueryService(workspaces=workspaces, plans=plans) + application_services_mock = SimpleNamespace(workspace_queries=workspace_queries) + + with patch( + "controllers.console.workspace.workspace.application_services", return_value=application_services_mock + ): + result, status = method(api, request_context=request_context) + + assert status == HTTPStatus.OK + assert result == { + "workspaces": [ + { + "id": "workspace-1", + "name": "Workspace 1", + "plan": "team", + "status": "normal", + "created_at": int(created_at.timestamp()), + "last_opened_at": int(last_opened_at.timestamp()), + "current": True, }, - ) as get_plan_bulk_mock, - patch("controllers.console.workspace.workspace.FeatureService.get_features") as get_features_mock, - ): - result, status = method(api, MagicMock(), "t1", user) - assert status == HTTPStatus.OK - assert len(result["workspaces"]) == 2 - assert result["workspaces"][0]["current"] is True - assert result["workspaces"][0]["plan"] == CloudPlan.TEAM - assert result["workspaces"][0]["last_opened_at"] == int(last_opened_at.timestamp()) - assert result["workspaces"][1]["plan"] == CloudPlan.PROFESSIONAL - assert result["workspaces"][1]["last_opened_at"] is None - get_plan_bulk_mock.assert_called_once_with(["t1", "t2"]) - get_features_mock.assert_not_called() + { + "id": "workspace-2", + "name": None, + "plan": "sandbox", + "status": "normal", + "created_at": int(created_at.timestamp()), + "last_opened_at": None, + "current": False, + }, + ] + } + workspaces.list_for_account.assert_called_once_with("account-1") + plans.resolve_many.assert_called_once_with(["workspace-1", "workspace-2"]) - def test_get_saas_path_partial_fallback_does_not_gate_plan_on_billing_enabled(self, app: Flask): - """Bulk omits a tenant: resolve plan via subscription.plan only; billing.enabled is not used. - billing.enabled is mocked False to prove the endpoint does not gate on it for this path - (SaaS contract treats enabled as on; display follows subscription.plan). - """ - api = TenantListApi() - method = unwrap(api.get) - tenant1 = make_tenant("t1", name="Tenant 1") - tenant2 = make_tenant("t2", name="Tenant 2") - features_t2 = MagicMock() - features_t2.billing.enabled = False - features_t2.billing.subscription.plan = CloudPlan.PROFESSIONAL - user = make_account() - with ( - app.test_request_context("/workspaces"), - patch( - "controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", - return_value=[(tenant1, make_membership()), (tenant2, make_membership())], +class TestWorkspaceQueryRepository: + def test_list_for_account_filters_orders_and_maps(self, workspace_session: scoped_session[Session]): + now = naive_utc_now() + earlier = make_tenant("workspace-1") + earlier.created_at = now - timedelta(days=1) + later = make_tenant("workspace-2") + later.created_at = now + archived = make_tenant("workspace-3", status=TenantStatus.ARCHIVE) + other_account = make_tenant("workspace-4") + last_opened_at = now - timedelta(hours=1) + workspace_session.add_all( + [ + earlier, + later, + archived, + other_account, + TenantAccountJoin( + tenant_id=earlier.id, + account_id="account-1", + last_opened_at=last_opened_at, + ), + TenantAccountJoin(tenant_id=later.id, account_id="account-1"), + TenantAccountJoin(tenant_id=archived.id, account_id="account-1"), + TenantAccountJoin(tenant_id=other_account.id, account_id="account-2"), + ] + ) + workspace_session.commit() + + result = WorkspaceQueryRepository(workspace_session.session_factory).list_for_account("account-1") + + assert result == ( + WorkspaceRecord( + id=earlier.id, + name=earlier.name, + status=TenantStatus.NORMAL.value, + created_at=earlier.created_at, + last_opened_at=last_opened_at, ), - patch("controllers.console.workspace.workspace.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", True), - patch("controllers.console.workspace.workspace.dify_config.EDITION", "CLOUD"), - patch( - "controllers.console.workspace.workspace.BillingService.get_plan_bulk", - return_value={"t1": {"plan": CloudPlan.TEAM, "expiration_date": 0}}, - ) as get_plan_bulk_mock, - patch( - "controllers.console.workspace.workspace.FeatureService.get_features", return_value=features_t2 - ) as get_features_mock, - ): - result, status = method(api, MagicMock(), "t1", user) - assert status == HTTPStatus.OK - assert result["workspaces"][0]["plan"] == CloudPlan.TEAM - assert result["workspaces"][1]["plan"] == CloudPlan.PROFESSIONAL - get_plan_bulk_mock.assert_called_once_with(["t1", "t2"]) - get_features_mock.assert_called_once_with("t2", exclude_vector_space=True) - - def test_get_saas_path_falls_back_to_legacy_feature_path_on_bulk_error( - self, app: Flask, caplog: pytest.LogCaptureFixture - ): - """Test fallback to FeatureService when bulk billing returns empty result. - - BillingService.get_plan_bulk catches exceptions internally and returns empty dict, - so we simulate the real failure mode by returning empty dict for non-empty input. - """ - api = TenantListApi() - method = unwrap(api.get) - tenant1 = make_tenant("t1", name="Tenant 1") - tenant2 = make_tenant("t2", name="Tenant 2") - features = MagicMock() - features.billing.enabled = False - features.billing.subscription.plan = CloudPlan.TEAM - user = make_account() - with ( - app.test_request_context("/workspaces"), - caplog.at_level(logging.WARNING, logger="controllers.console.workspace.workspace"), - patch( - "controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", - return_value=[(tenant1, make_membership()), (tenant2, make_membership())], + WorkspaceRecord( + id=later.id, + name=later.name, + status=TenantStatus.NORMAL.value, + created_at=later.created_at, + last_opened_at=None, ), - patch("controllers.console.workspace.workspace.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", True), - patch("controllers.console.workspace.workspace.dify_config.EDITION", "CLOUD"), - patch( - "controllers.console.workspace.workspace.BillingService.get_plan_bulk", return_value={} - ) as get_plan_bulk_mock, - patch( - "controllers.console.workspace.workspace.FeatureService.get_features", return_value=features - ) as get_features_mock, - ): - result, status = method(api, MagicMock(), "t2", user) - assert status == HTTPStatus.OK - assert result["workspaces"][0]["plan"] == CloudPlan.TEAM - assert result["workspaces"][1]["plan"] == CloudPlan.TEAM - get_plan_bulk_mock.assert_called_once_with(["t1", "t2"]) - assert get_features_mock.call_count == 2 - assert "get_plan_bulk returned empty result, falling back to legacy feature path" in caplog.messages + ) - def test_get_billing_disabled_community_path(self, app: Flask): - api = TenantListApi() - method = unwrap(api.get) - tenant = make_tenant("t1", name="Tenant") - features = MagicMock() - features.billing.enabled = False - features.billing.subscription.plan = CloudPlan.SANDBOX - user = make_account() - with ( - app.test_request_context("/workspaces"), - patch( - "controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", - return_value=[(tenant, make_membership())], - ), - patch("controllers.console.workspace.workspace.dify_config.ENTERPRISE_ENABLED", False), - patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", False), - patch("controllers.console.workspace.workspace.dify_config.EDITION", "SELF_HOSTED"), - patch( - "controllers.console.workspace.workspace.FeatureService.get_features", return_value=features - ) as get_features_mock, - ): - result, status = method(api, MagicMock(), "t1", user) - assert status == HTTPStatus.OK - assert result["workspaces"][0]["plan"] == CloudPlan.SANDBOX - get_features_mock.assert_called_once_with("t1", exclude_vector_space=True) - def test_get_enterprise_only_skips_feature_service(self, app: Flask): - api = TenantListApi() - method = unwrap(api.get) - tenant1 = make_tenant("t1", name="Tenant 1") - tenant2 = make_tenant("t2", name="Tenant 2") - user = make_account() - with ( - app.test_request_context("/workspaces"), - patch( - "controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", - return_value=[(tenant1, make_membership()), (tenant2, make_membership())], - ), - patch("controllers.console.workspace.workspace.dify_config.ENTERPRISE_ENABLED", True), - patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", False), - patch("controllers.console.workspace.workspace.dify_config.EDITION", "SELF_HOSTED"), - patch("controllers.console.workspace.workspace.FeatureService.get_features") as get_features_mock, - ): - result, status = method(api, MagicMock(), "t2", user) - assert status == HTTPStatus.OK - assert result["workspaces"][0]["plan"] == CloudPlan.SANDBOX - assert result["workspaces"][1]["plan"] == CloudPlan.SANDBOX - assert result["workspaces"][0]["current"] is False - assert result["workspaces"][1]["current"] is True - get_features_mock.assert_not_called() +class TestDeploymentWorkspacePlanGateway: + def test_saas_uses_bulk_plans_and_feature_fallback( + self, + monkeypatch: pytest.MonkeyPatch, + workspace_plan_dependencies: tuple[MagicMock, MagicMock], + ) -> None: + configure_workspace_plans(monkeypatch) + get_plan_bulk, get_features = workspace_plan_dependencies + get_plan_bulk.return_value = {"workspace-1": {"plan": CloudPlan.TEAM, "expiration_date": 0}} + get_features.return_value = features_with_plan(CloudPlan.PROFESSIONAL) - def test_get_enterprise_only_with_empty_tenants(self, app: Flask): - api = TenantListApi() - method = unwrap(api.get) - user = make_account() - with ( - app.test_request_context("/workspaces"), - patch("controllers.console.workspace.workspace.TenantService.get_workspaces_for_account", return_value=[]), - patch("controllers.console.workspace.workspace.dify_config.ENTERPRISE_ENABLED", True), - patch("controllers.console.workspace.workspace.dify_config.BILLING_ENABLED", False), - patch("controllers.console.workspace.workspace.dify_config.EDITION", "SELF_HOSTED"), - patch("controllers.console.workspace.workspace.FeatureService.get_features") as get_features_mock, - ): - result, status = method(api, MagicMock(), None, user) - assert status == HTTPStatus.OK - assert result["workspaces"] == [] - get_features_mock.assert_not_called() + result = workspace_plan_gateway.DeploymentWorkspacePlanGateway().resolve_many(["workspace-1", "workspace-2"]) + + assert result == {"workspace-1": CloudPlan.TEAM, "workspace-2": CloudPlan.PROFESSIONAL} + get_plan_bulk.assert_called_once() + assert list(get_plan_bulk.call_args.args[0]) == ["workspace-1", "workspace-2"] + get_features.assert_called_once_with("workspace-2", exclude_vector_space=True) + + def test_saas_empty_bulk_result_falls_back_to_features( + self, + monkeypatch: pytest.MonkeyPatch, + workspace_plan_dependencies: tuple[MagicMock, MagicMock], + caplog: pytest.LogCaptureFixture, + ) -> None: + configure_workspace_plans(monkeypatch) + get_plan_bulk, get_features = workspace_plan_dependencies + get_plan_bulk.return_value = {} + get_features.return_value = features_with_plan(CloudPlan.TEAM) + + with caplog.at_level(logging.WARNING, logger=workspace_plan_gateway.__name__): + result = workspace_plan_gateway.DeploymentWorkspacePlanGateway().resolve_many( + ["workspace-1", "workspace-2"] + ) + + assert result == {"workspace-1": CloudPlan.TEAM, "workspace-2": CloudPlan.TEAM} + assert "get_plan_bulk returned empty result, falling back to FeatureService" in caplog.messages + + def test_non_saas_uses_features( + self, + monkeypatch: pytest.MonkeyPatch, + workspace_plan_dependencies: tuple[MagicMock, MagicMock], + ) -> None: + configure_workspace_plans( + monkeypatch, + edition=DeploymentEdition.COMMUNITY, + ) + get_plan_bulk, get_features = workspace_plan_dependencies + get_features.return_value = features_with_plan(CloudPlan.SANDBOX) + + result = workspace_plan_gateway.DeploymentWorkspacePlanGateway().resolve_many(["workspace-1"]) + + assert result == {"workspace-1": CloudPlan.SANDBOX} + get_plan_bulk.assert_not_called() + get_features.assert_called_once_with("workspace-1", exclude_vector_space=True) + + def test_enterprise_only_skips_external_lookups( + self, + monkeypatch: pytest.MonkeyPatch, + workspace_plan_dependencies: tuple[MagicMock, MagicMock], + ) -> None: + configure_workspace_plans( + monkeypatch, + edition=DeploymentEdition.ENTERPRISE, + ) + get_plan_bulk, get_features = workspace_plan_dependencies + + result = workspace_plan_gateway.DeploymentWorkspacePlanGateway().resolve_many(["workspace-1", "workspace-2"]) + + assert result == {"workspace-1": CloudPlan.SANDBOX, "workspace-2": CloudPlan.SANDBOX} + get_plan_bulk.assert_not_called() + get_features.assert_not_called() class TestWorkspaceListApi: @@ -297,66 +328,59 @@ class TestWorkspaceListApi: assert result["has_more"] is True -class TestTenantApi: - def test_post_active_tenant(self, app: Flask): - api = TenantApi() - method = unwrap(api.post) +def test_legacy_current_workspace_routes_are_not_registered(): + urls = {url for _resource, resource_urls, _route_doc, _kwargs in console_ns.resources for url in resource_urls} + + assert "/workspaces/current" not in urls + assert "/info" not in urls + + +class TestCurrentWorkspaceSummaryApi: + def test_get_summary(self, app: Flask): + api = CurrentWorkspaceSummaryApi() + method = unwrap(api.get) tenant = make_tenant() user = make_account_with_tenant(tenant) + session = MagicMock() + summary = { + "id": tenant.id, + "name": tenant.name, + "role": "owner", + "plan": CloudPlan.SANDBOX, + "credits": 180, + } + with ( - app.test_request_context("/workspaces/current"), + app.test_request_context("/workspaces/current/summary"), patch( - "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "t1"} - ), + "controllers.console.workspace.workspace.WorkspaceService.get_current_workspace_summary", + return_value=summary, + ) as get_summary, ): - result, status = method(api, MagicMock(), user) + result, status = method(api, session, user) + assert status == HTTPStatus.OK - assert result["id"] == "t1" + assert result == { + "id": tenant.id, + "name": tenant.name, + "role": "owner", + "plan": "sandbox", + "credits": 180, + } + get_summary.assert_called_once_with(tenant, user.id, session=session) - def test_post_archived_with_switch(self, app: Flask): - api = TenantApi() - method = unwrap(api.post) - archived = make_tenant(status=TenantStatus.ARCHIVE) - new_tenant = make_tenant("new") - user = make_account_with_tenant(archived) - with ( - app.test_request_context("/workspaces/current"), - patch("controllers.console.workspace.workspace.TenantService.get_join_tenants", return_value=[new_tenant]), - patch("controllers.console.workspace.workspace.TenantService.switch_tenant") as switch_tenant, - patch( - "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "new"} - ), - ): - result, status = method(api, MagicMock(), user) - assert result["id"] == "new" - switch_tenant.assert_called_once_with(user, new_tenant.id, session=ANY) + def test_get_archived_tenant_returns_conflict(self, app: Flask): + api = CurrentWorkspaceSummaryApi() + method = unwrap(api.get) + tenant = make_tenant(status=TenantStatus.ARCHIVE) + user = make_account_with_tenant(tenant) - def test_post_archived_no_tenant(self, app: Flask): - api = TenantApi() - method = unwrap(api.post) - user = make_account_with_tenant(make_tenant(status=TenantStatus.ARCHIVE)) - with ( - app.test_request_context("/workspaces/current"), - patch("controllers.console.workspace.workspace.TenantService.get_join_tenants", return_value=[]), - ): - with pytest.raises(Unauthorized): + with app.test_request_context("/workspaces/current/summary"): + with pytest.raises(CurrentWorkspaceArchivedError) as exc_info: method(api, MagicMock(), user) - def test_post_info_path(self, app: Flask, caplog: pytest.LogCaptureFixture): - api = TenantApi() - method = unwrap(api.post) - tenant = make_tenant() - user = make_account_with_tenant(tenant) - with ( - app.test_request_context("/info"), - caplog.at_level(logging.WARNING, logger="controllers.console.workspace.workspace"), - patch( - "controllers.console.workspace.workspace.WorkspaceService.get_tenant_info", return_value={"id": "t1"} - ), - ): - result, status = method(api, MagicMock(), user) - assert "Deprecated URL /info was used." in caplog.messages - assert status == HTTPStatus.OK + assert exc_info.value.code == HTTPStatus.CONFLICT + assert exc_info.value.error_code == "current_workspace_archived" class TestTenantInfoResponse: @@ -430,6 +454,43 @@ class TestSwitchWorkspaceApi: class TestCustomConfigWorkspaceApi: + def test_get_workspace_not_found(self, app: Flask, workspace_session: scoped_session[Session]): + api = CustomConfigWorkspaceApi() + method = unwrap(api.get) + + with app.test_request_context("/workspaces/custom-config"), pytest.raises(NotFound): + method(api, workspace_session, "missing") + + def test_get_defaults(self, app: Flask, workspace_session: scoped_session[Session]): + api = CustomConfigWorkspaceApi() + method = unwrap(api.get) + tenant = make_tenant(custom_config={}) + workspace_session.add(tenant) + workspace_session.commit() + + with app.test_request_context("/workspaces/custom-config"): + result = method(api, workspace_session, tenant.id) + + assert result == {"remove_webapp_brand": False, "replace_webapp_logo": None} + + def test_get_configured_brand(self, app: Flask, workspace_session: scoped_session[Session]): + api = CustomConfigWorkspaceApi() + method = unwrap(api.get) + tenant = make_tenant(custom_config={"remove_webapp_brand": True, "replace_webapp_logo": "logo-file-id"}) + workspace_session.add(tenant) + workspace_session.commit() + + with ( + app.test_request_context("/workspaces/custom-config"), + patch("controllers.console.workspace.workspace.dify_config.FILES_URL", "https://files.example.com"), + ): + result = method(api, workspace_session, tenant.id) + + assert result == { + "remove_webapp_brand": True, + "replace_webapp_logo": f"https://files.example.com/files/workspaces/{tenant.id}/webapp-logo", + } + def test_post_success(self, app: Flask, workspace_session: scoped_session[Session]): api = CustomConfigWorkspaceApi() method = unwrap(api.post) @@ -580,9 +641,8 @@ class TestWorkspaceInfoApi: ), ), ): - session = MagicMock() - session.get.return_value = tenant - session.commit.side_effect = lambda: events.append("commit") + session = workspace_session() + event.listen(session, "after_commit", lambda _session: events.append("commit")) result = method(api, session, "t1") assert result["result"] == "success" assert events == ["commit", "get_tenant_info"] diff --git a/api/tests/unit_tests/controllers/files/test_upload.py b/api/tests/unit_tests/controllers/files/test_upload.py index 4f6fcd6f3e0..5e0f5273253 100644 --- a/api/tests/unit_tests/controllers/files/test_upload.py +++ b/api/tests/unit_tests/controllers/files/test_upload.py @@ -1,13 +1,17 @@ import io import types +from contextlib import contextmanager from inspect import unwrap from unittest.mock import patch import pytest +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden import controllers.files.upload as module from core.workflow.file_reference import build_file_reference +from models import Account, TenantAccountJoin +from models.account import AccountStatus def fake_request(args: dict, file=None): @@ -17,6 +21,22 @@ def fake_request(args: dict, file=None): ) +def _persist_account_memberships(session: Session) -> None: + account = Account(name="Tenant member", email="member@example.com", status=AccountStatus.ACTIVE) + account.id = "account-1" + decoy = Account(name="Other tenant member", email="decoy@example.com", status=AccountStatus.ACTIVE) + decoy.id = "account-outside-tenant" + session.add_all( + [ + account, + decoy, + TenantAccountJoin(tenant_id="tenant-1", account_id=account.id), + TenantAccountJoin(tenant_id="tenant-other", account_id=decoy.id), + ] + ) + session.commit() + + class DummyUser: def __init__(self, user_id="user-1"): self.id = user_id @@ -33,6 +53,16 @@ class DummyFile: return self.stream.read() +class RecordingStream(io.BytesIO): + def __init__(self, content: bytes, events: list[str]): + super().__init__(content) + self.events = events + + def read(self, *args, **kwargs): + self.events.append("file-read") + return super().read(*args, **kwargs) + + class DummyToolFile: def __init__(self, name="test.txt", mimetype="text/plain"): self.id = "file-id" @@ -94,6 +124,138 @@ class TestPluginUploadFileApi: assert tool_file_manager_instance.create_file_by_raw.call_args.kwargs["conversation_id"] == "conversation-1" mock_tool_file_manager.sign_file.assert_called_once_with(tool_file_id="file-id", extension=".docx") + @patch.object(module, "get_user") + @patch.object(module, "ToolFileManager") + @pytest.mark.parametrize("sqlite_session", [(Account, TenantAccountJoin)], indirect=True) + def test_account_upload_preserves_signed_account_owner( + self, + mock_tool_file_manager, + mock_get_user, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ): + _persist_account_memberships(sqlite_session) + events: list[str] = [] + dummy_file = DummyFile(filename="report.pdf", mimetype="application/pdf", content=b"account-owned") + dummy_file.stream = RecordingStream(b"account-owned", events) + + @contextmanager + def membership_session(): + events.append("membership-session-enter") + try: + yield sqlite_session + finally: + events.append("membership-session-exit") + + monkeypatch.setattr(module.session_factory, "create_session", membership_session) + monkeypatch.setattr( + module, + "request", + fake_request( + { + "timestamp": "123", + "nonce": "abc", + "sign": "sig", + "tenant_id": "tenant-1", + "user_id": "account-1", + "user_from": "account", + }, + file=dummy_file, + ), + ) + tool_file_manager = mock_tool_file_manager.return_value + tool_file_manager.create_file_by_raw.side_effect = lambda **_kwargs: ( + events.append("storage-create-file") or DummyToolFile(name="report.pdf", mimetype="application/pdf") + ) + mock_tool_file_manager.sign_file.return_value = "signed-url" + + with patch.object( + module, + "verify_plugin_file_signature", + side_effect=lambda **_kwargs: events.append("signature-verify") or True, + ) as verify_signature: + api = module.PluginUploadFileApi() + result, status_code = unwrap(api.post)(api) + + assert status_code == 201 + assert result["reference"] == build_file_reference(record_id="file-id") + assert events == [ + "membership-session-enter", + "membership-session-exit", + "signature-verify", + "file-read", + "storage-create-file", + ] + mock_get_user.assert_not_called() + verify_signature.assert_called_once_with( + filename="report.pdf", + mimetype="application/pdf", + tenant_id="tenant-1", + user_id="account-1", + conversation_id=None, + user_from="account", + timestamp="123", + nonce="abc", + sign="sig", + ) + tool_file_manager.create_file_by_raw.assert_called_once_with( + user_id="account-1", + tenant_id="tenant-1", + file_binary=b"account-owned", + mimetype="application/pdf", + filename="report.pdf", + conversation_id=None, + ) + + @patch.object(module, "verify_plugin_file_signature") + @patch.object(module, "get_user") + @patch.object(module, "ToolFileManager") + @pytest.mark.parametrize("sqlite_session", [(Account, TenantAccountJoin)], indirect=True) + def test_account_upload_rejects_owner_outside_tenant( + self, + mock_tool_file_manager, + mock_get_user, + mock_verify_signature, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ): + _persist_account_memberships(sqlite_session) + events: list[str] = [] + + @contextmanager + def membership_session(): + events.append("membership-session-enter") + try: + yield sqlite_session + finally: + events.append("membership-session-exit") + + monkeypatch.setattr(module.session_factory, "create_session", membership_session) + monkeypatch.setattr( + module, + "request", + fake_request( + { + "timestamp": "123", + "nonce": "abc", + "sign": "sig", + "tenant_id": "tenant-1", + "user_id": "account-outside-tenant", + "user_from": "account", + }, + file=DummyFile(), + ), + ) + + api = module.PluginUploadFileApi() + with pytest.raises(Forbidden): + unwrap(api.post)(api) + + assert events == ["membership-session-enter", "membership-session-exit"] + mock_get_user.assert_not_called() + mock_verify_signature.assert_not_called() + mock_tool_file_manager.assert_not_called() + def test_missing_file(self): module.request = fake_request( { diff --git a/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py b/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py index c7f788dcb55..c11a6a97302 100644 --- a/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py +++ b/api/tests/unit_tests/controllers/inner_api/app/test_dsl.py @@ -6,20 +6,46 @@ in test_auth_wraps.py; handler tests use inspect.unwrap() to bypass them. """ import inspect -from unittest.mock import ANY, MagicMock, patch +from types import SimpleNamespace +from unittest.mock import MagicMock, patch +from uuid import uuid4 import pytest from flask import Flask from pydantic import ValidationError +from sqlalchemy import event, select +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker +from controllers.inner_api.app import dsl as dsl_module from controllers.inner_api.app.dsl import ( EnterpriseAppDSLExport, EnterpriseAppDSLImport, InnerAppDSLImportPayload, _get_active_account, ) +from models import Account, App from models.account import AccountStatus +from models.model import AppMode, IconType from services.app_dsl_service import Import, ImportStatus +from services.errors.app import IsDraftWorkflowError, WorkflowNotFoundError + + +def _persist_app(session: Session) -> App: + app = App( + id=str(uuid4()), + tenant_id=str(uuid4()), + name="DSL App", + mode=AppMode.WORKFLOW, + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#ffffff", + enable_site=False, + enable_api=False, + ) + session.add(app) + session.commit() + return app class TestInnerAppDSLImportPayload: @@ -61,32 +87,29 @@ class TestInnerAppDSLImportPayload: class TestGetActiveAccount: """Test the _get_active_account helper function.""" - @patch("controllers.inner_api.app.dsl.db") - def test_returns_active_account(self, mock_db): - mock_account = MagicMock() - mock_account.status = AccountStatus.ACTIVE - mock_db.session.scalar.return_value = mock_account + def test_returns_active_account(self, sqlite_session: Session): + account = Account(name="Active", email="user@example.com", status=AccountStatus.ACTIVE) + sqlite_session.add(account) + sqlite_session.commit() - result = _get_active_account("user@example.com") + with patch.object(dsl_module.db, "session", sqlite_session): + result = _get_active_account("user@example.com") - assert result is mock_account - mock_db.session.scalar.assert_called_once() + assert result is account - @patch("controllers.inner_api.app.dsl.db") - def test_returns_none_for_inactive_account(self, mock_db): - mock_account = MagicMock() - mock_account.status = AccountStatus.BANNED - mock_db.session.scalar.return_value = mock_account + def test_returns_none_for_inactive_account(self, sqlite_session: Session): + account = Account(name="Banned", email="banned@example.com", status=AccountStatus.BANNED) + sqlite_session.add(account) + sqlite_session.commit() - result = _get_active_account("banned@example.com") + with patch.object(dsl_module.db, "session", sqlite_session): + result = _get_active_account("banned@example.com") assert result is None - @patch("controllers.inner_api.app.dsl.db") - def test_returns_none_for_nonexistent_email(self, mock_db): - mock_db.session.scalar.return_value = None - - result = _get_active_account("missing@example.com") + def test_returns_none_for_nonexistent_email(self, sqlite_session: Session): + with patch.object(dsl_module.db, "session", sqlite_session): + result = _get_active_account("missing@example.com") assert result is None @@ -102,20 +125,29 @@ class TestEnterpriseAppDSLImport: return EnterpriseAppDSLImport() @pytest.fixture - def _mock_import_deps(self): - """Patch db, Session, and AppDslService for import handler tests.""" - mock_session = MagicMock() - mock_session.__enter__ = MagicMock(return_value=mock_session) - mock_session.__exit__ = MagicMock(return_value=False) + def _mock_import_deps(self, sqlite_engine: Engine): + """Bind the handler Session to SQLite and isolate the DSL service boundary.""" + self._transaction_events: list[str] = [] + + def on_commit(session: Session) -> None: + if session.get_bind() is sqlite_engine: + self._transaction_events.append("commit") + + def on_rollback(session: Session) -> None: + if session.get_bind() is sqlite_engine: + self._transaction_events.append("rollback") + + event.listen(Session, "after_commit", on_commit) + event.listen(Session, "after_rollback", on_rollback) with ( - patch("controllers.inner_api.app.dsl.db"), - patch("controllers.inner_api.app.dsl.Session", return_value=mock_session), + patch.object(dsl_module, "db", SimpleNamespace(engine=sqlite_engine)), patch("controllers.inner_api.app.dsl.AppDslService") as mock_dsl_cls, ): - self._mock_session = mock_session self._mock_dsl = MagicMock() mock_dsl_cls.return_value = self._mock_dsl yield + event.remove(Session, "after_commit", on_commit) + event.remove(Session, "after_rollback", on_rollback) def _make_import_result(self, status: ImportStatus, **kwargs) -> Import: result = Import( @@ -145,9 +177,9 @@ class TestEnterpriseAppDSLImport: body, status_code = result assert status_code == 200 assert body["status"] == "completed" - mock_account.set_tenant_id_with_session.assert_called_once_with("ws-123", session=self._mock_session) - self._mock_session.commit.assert_called_once_with() - self._mock_session.rollback.assert_not_called() + call_session = mock_account.set_tenant_id_with_session.call_args.kwargs["session"] + assert isinstance(call_session, Session) + assert self._transaction_events == ["commit"] @pytest.mark.usefixtures("_mock_import_deps") @patch("controllers.inner_api.app.dsl._get_active_account") @@ -163,13 +195,14 @@ class TestEnterpriseAppDSLImport: assert status_code == 202 assert body["status"] == "pending" - self._mock_session.commit.assert_called_once_with() - self._mock_session.rollback.assert_not_called() + assert self._transaction_events == ["commit"] @pytest.mark.usefixtures("_mock_import_deps") @patch("controllers.inner_api.app.dsl._get_active_account") def test_import_failed_returns_400(self, mock_get_account, api_instance, app: Flask): - mock_get_account.return_value = MagicMock() + mock_account = MagicMock() + mock_account.set_tenant_id_with_session.side_effect = lambda _tenant_id, *, session: session.execute(select(1)) + mock_get_account.return_value = mock_account self._mock_dsl.import_app.return_value = self._make_import_result(ImportStatus.FAILED) unwrapped = inspect.unwrap(api_instance.post) @@ -180,8 +213,7 @@ class TestEnterpriseAppDSLImport: assert status_code == 400 assert body["status"] == "failed" - self._mock_session.rollback.assert_called_once_with() - self._mock_session.commit.assert_not_called() + assert self._transaction_events == ["rollback"] @patch("controllers.inner_api.app.dsl._get_active_account") def test_import_account_not_found_returns_404(self, mock_get_account, api_instance, app: Flask): @@ -204,48 +236,232 @@ class TestEnterpriseAppDSLExport: Uses inspect.unwrap() to bypass auth/setup decorators. """ + def test_export_documents_query_parameters(self): + params = EnterpriseAppDSLExport.get.__apidoc__["params"] + + assert params["include_secret"]["in"] == "query" + assert params["include_secret"]["type"] == "boolean" + assert params["workflow_id"]["in"] == "query" + assert params["workflow_id"]["type"] == "string" + assert params["workflow_id"]["format"] == "uuid" + @pytest.fixture def api_instance(self): return EnterpriseAppDSLExport() + @pytest.fixture + def scoped_db(self, sqlite_session_factory: sessionmaker[Session]): + db_session = scoped_session(sqlite_session_factory) + with patch.object(dsl_module, "db", SimpleNamespace(session=db_session)): + yield db_session + db_session.remove() + @patch("controllers.inner_api.app.dsl.AppDslService") - @patch("controllers.inner_api.app.dsl.db") - def test_export_success_returns_200(self, mock_db, mock_dsl_cls, api_instance, app: Flask): - mock_app = MagicMock() - mock_db.session.get.return_value = mock_app + def test_export_success_returns_200( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) mock_dsl_cls.export_dsl.return_value = "version: 0.6.0\nkind: app\n" unwrapped = inspect.unwrap(api_instance.get) with app.test_request_context("?include_secret=false"): - result = unwrapped(api_instance, app_id="app-123") + result = unwrapped(api_instance, app_id=app_model.id) body, status_code = result assert status_code == 200 assert body["data"] == "version: 0.6.0\nkind: app\n" - mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=False) + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is False @patch("controllers.inner_api.app.dsl.AppDslService") - @patch("controllers.inner_api.app.dsl.db") - def test_export_with_secret(self, mock_db, mock_dsl_cls, api_instance, app: Flask): - mock_app = MagicMock() - mock_db.session.get.return_value = mock_app + def test_export_with_secret( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) mock_dsl_cls.export_dsl.return_value = "yaml-data" unwrapped = inspect.unwrap(api_instance.get) with app.test_request_context("?include_secret=true"): - result = unwrapped(api_instance, app_id="app-123") + result = unwrapped(api_instance, app_id=app_model.id) body, status_code = result assert status_code == 200 - mock_dsl_cls.export_dsl.assert_called_once_with(app_model=mock_app, session=ANY, include_secret=True) + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is True - @patch("controllers.inner_api.app.dsl.db") - def test_export_app_not_found_returns_404(self, mock_db, api_instance, app: Flask): - mock_db.session.get.return_value = None + @patch("controllers.inner_api.app.dsl.AppDslService") + def test_export_selected_workflow_forwards_canonical_uuid( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) + mock_dsl_cls.export_dsl.return_value = "yaml-data" + workflow_id = "F1FD7266-56FC-45C7-9D81-A72CD5A1B4F6" + unwrapped = inspect.unwrap(api_instance.get) + with app.test_request_context(f"?workflow_id={workflow_id}"): + body, status_code = unwrapped(api_instance, app_id=app_model.id) + + assert status_code == 200 + assert body["data"] == "yaml-data" + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is False + assert call_kwargs["workflow_id"] == "f1fd7266-56fc-45c7-9d81-a72cd5a1b4f6" + + @patch("controllers.inner_api.app.dsl.AppDslService") + def test_export_selected_workflow_with_secret( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) + mock_dsl_cls.export_dsl.return_value = "yaml-data" + workflow_id = "f1fd7266-56fc-45c7-9d81-a72cd5a1b4f6" + + unwrapped = inspect.unwrap(api_instance.get) + with app.test_request_context(f"?include_secret=true&workflow_id={workflow_id}"): + body, status_code = unwrapped(api_instance, app_id=app_model.id) + + assert status_code == 200 + assert body["data"] == "yaml-data" + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is True + assert call_kwargs["workflow_id"] == workflow_id + + @patch("controllers.inner_api.app.dsl.AppDslService") + def test_export_rejects_invalid_selected_workflow_id( + self, + mock_dsl_cls, + api_instance, + app: Flask, + scoped_db, + ): + assert scoped_db() is not None + unwrapped = inspect.unwrap(api_instance.get) + with app.test_request_context("?workflow_id=not-a-uuid"): + body, status_code = unwrapped(api_instance, app_id=str(uuid4())) + + assert status_code == 400 + assert body == { + "code": "invalid_workflow_id", + "message": "workflow_id must be a valid UUID", + "status": 400, + } + mock_dsl_cls.export_dsl.assert_not_called() + + @patch("controllers.inner_api.app.dsl.AppDslService") + def test_export_selected_missing_workflow_returns_404( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) + mock_dsl_cls.export_dsl.side_effect = WorkflowNotFoundError("selected workflow not found") + workflow_id = "f1fd7266-56fc-45c7-9d81-a72cd5a1b4f6" + + unwrapped = inspect.unwrap(api_instance.get) + with app.test_request_context(f"?workflow_id={workflow_id}"): + body, status_code = unwrapped(api_instance, app_id=app_model.id) + + assert status_code == 404 + assert body == { + "code": "workflow_version_not_found", + "message": "selected workflow not found", + "status": 404, + } + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is False + assert call_kwargs["workflow_id"] == workflow_id + + @patch("controllers.inner_api.app.dsl.AppDslService") + def test_export_selected_draft_workflow_returns_400( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) + mock_dsl_cls.export_dsl.side_effect = IsDraftWorkflowError("selected workflow is a draft") + workflow_id = "f1fd7266-56fc-45c7-9d81-a72cd5a1b4f6" + + unwrapped = inspect.unwrap(api_instance.get) + with app.test_request_context(f"?workflow_id={workflow_id}"): + body, status_code = unwrapped(api_instance, app_id=app_model.id) + + assert status_code == 400 + assert body == { + "code": "workflow_version_not_published", + "message": "selected workflow is a draft", + "status": 400, + } + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is False + assert call_kwargs["workflow_id"] == workflow_id + + @patch("controllers.inner_api.app.dsl.AppDslService") + def test_export_without_selected_workflow_preserves_workflow_error( + self, + mock_dsl_cls, + api_instance, + app: Flask, + sqlite_session: Session, + scoped_db, + ): + app_model = _persist_app(sqlite_session) + mock_dsl_cls.export_dsl.side_effect = WorkflowNotFoundError( + "Missing draft workflow configuration, please check." + ) + + unwrapped = inspect.unwrap(api_instance.get) + with app.test_request_context(): + with pytest.raises(WorkflowNotFoundError, match="Missing draft workflow configuration"): + unwrapped(api_instance, app_id=app_model.id) + + call_kwargs = mock_dsl_cls.export_dsl.call_args.kwargs + assert call_kwargs["app_model"].id == app_model.id + assert call_kwargs["session"] is scoped_db() + assert call_kwargs["include_secret"] is False + assert "workflow_id" not in call_kwargs + + def test_export_app_not_found_returns_404(self, api_instance, app: Flask, scoped_db): + assert scoped_db() is not None unwrapped = inspect.unwrap(api_instance.get) with app.test_request_context("?include_secret=false"): - result = unwrapped(api_instance, app_id="nonexistent") + result = unwrapped(api_instance, app_id=str(uuid4())) body, status_code = result assert status_code == 404 diff --git a/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_config.py b/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_config.py index eef6fdf31c0..79ecdc08cb9 100644 --- a/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_config.py +++ b/api/tests/unit_tests/controllers/inner_api/plugin/test_agent_config.py @@ -3,19 +3,17 @@ from __future__ import annotations import inspect -from types import SimpleNamespace from unittest.mock import patch from flask import Flask from controllers.inner_api.plugin.agent_config import ( + AgentConfigDownloadRequestApi, AgentConfigEnvApi, - AgentConfigFilePullApi, AgentConfigManifestApi, AgentConfigNoteApi, AgentConfigPushApi, AgentConfigSkillInspectApi, - AgentConfigSkillPullApi, ) from services.agent_config_service import AgentConfigServiceError @@ -108,26 +106,40 @@ def test_manifest_happy_path_calls_service() -> None: assert service.return_value.manifest.call_args.kwargs["config_version_kind"].value == "build_draft" -def test_skill_pull_returns_send_file_response() -> None: - raw = _raw(AgentConfigSkillPullApi.get) +def test_download_request_returns_origin_free_metadata() -> None: + raw = _raw(AgentConfigDownloadRequestApi.post) + payload = { + "tenant_id": "tenant-1", + "user_id": "user-1", + "config_version_id": "cfg-1", + "config_version_kind": "build_draft", + "config": {"kind": "skill", "name": "alpha"}, + } - with app.test_request_context( - "/?tenant_id=tenant-1&user_id=user-1&config_version_id=cfg-1&config_version_kind=build_draft" - ): + with app.test_request_context("/", method="POST", json=payload): with patch(f"{MODULE}.AgentConfigService") as service: - service.return_value.pull_skill.return_value = SimpleNamespace( - payload=b"zip-bytes", - mime_type="application/zip", - filename="alpha.zip", - ) - response = raw(AgentConfigSkillPullApi(), "agent-1", "alpha") + service.return_value.request_download.return_value.filename = "alpha.zip" + service.return_value.request_download.return_value.mime_type = "application/zip" + service.return_value.request_download.return_value.size = 123 + service.return_value.request_download.return_value.download_uri = "/files/tools/skill.zip?sign=1" + body = raw(AgentConfigDownloadRequestApi(), "agent-1") - response.direct_passthrough = False - assert response.status_code == 200 - assert response.mimetype == "application/zip" - assert response.get_data() == b"zip-bytes" - assert "filename=alpha.zip" in response.headers["Content-Disposition"] - assert service.return_value.pull_skill.call_args.kwargs["user_id"] == "user-1" + assert body == { + "filename": "alpha.zip", + "mime_type": "application/zip", + "size": 123, + "download_uri": "/files/tools/skill.zip?sign=1", + } + assert service.return_value.request_download.call_args.kwargs == { + "tenant_id": "tenant-1", + "agent_id": "agent-1", + "user_id": "user-1", + "config_version_id": "cfg-1", + "config_version_kind": service.return_value.request_download.call_args.kwargs["config_version_kind"], + "kind": "skill", + "name": "alpha", + } + assert service.return_value.request_download.call_args.kwargs["config_version_kind"].value == "build_draft" def test_skill_inspect_happy_path_returns_service_payload() -> None: @@ -144,28 +156,6 @@ def test_skill_inspect_happy_path_returns_service_payload() -> None: assert service.return_value.inspect_skill.call_args.kwargs["user_id"] == "user-1" -def test_file_pull_returns_send_file_response() -> None: - raw = _raw(AgentConfigFilePullApi.get) - - with app.test_request_context( - "/?tenant_id=tenant-1&user_id=user-1&config_version_id=cfg-1&config_version_kind=build_draft" - ): - with patch(f"{MODULE}.AgentConfigService") as service: - service.return_value.pull_file.return_value = SimpleNamespace( - payload=b"file-bytes", - mime_type="text/plain", - filename="guide.txt", - ) - response = raw(AgentConfigFilePullApi(), "agent-1", "guide.txt") - - response.direct_passthrough = False - assert response.status_code == 200 - assert response.mimetype == "text/plain" - assert response.get_data() == b"file-bytes" - assert "filename=guide.txt" in response.headers["Content-Disposition"] - assert service.return_value.pull_file.call_args.kwargs["user_id"] == "user-1" - - def test_push_happy_path_validates_body_and_preserves_execution_user() -> None: raw = _raw(AgentConfigPushApi.post) payload = { @@ -253,6 +243,22 @@ def test_manifest_invalid_query_returns_400() -> None: assert body["code"] == "invalid_request" +def test_download_request_rejects_invalid_config_name() -> None: + raw = _raw(AgentConfigDownloadRequestApi.post) + payload = { + "tenant_id": "tenant-1", + "config_version_id": "cfg-1", + "config_version_kind": "draft", + "config": {"kind": "file", "name": "../guide.txt"}, + } + + with app.test_request_context("/", method="POST", json=payload): + body, status = raw(AgentConfigDownloadRequestApi(), "agent-1") + + assert status == 400 + assert body["code"] == "invalid_request" + + def test_push_invalid_body_returns_400() -> None: raw = _raw(AgentConfigPushApi.post) diff --git a/api/tests/unit_tests/controllers/inner_api/plugin/test_plugin.py b/api/tests/unit_tests/controllers/inner_api/plugin/test_plugin.py index 249d8926124..c67e622ae51 100644 --- a/api/tests/unit_tests/controllers/inner_api/plugin/test_plugin.py +++ b/api/tests/unit_tests/controllers/inner_api/plugin/test_plugin.py @@ -263,11 +263,12 @@ class TestPluginUploadFileRequestApi: assert hasattr(api_instance, "post") assert callable(api_instance.post) - @patch("controllers.inner_api.plugin.plugin.get_signed_file_url_for_plugin") - def test_post_returns_signed_url(self, mock_get_url, api_instance, app: Flask): + @patch("controllers.inner_api.plugin.plugin.get_signed_file_uri_for_plugin") + def test_post_returns_signed_url(self, mock_get_uri, api_instance, app: Flask, monkeypatch: pytest.MonkeyPatch): """Test that post() generates a signed URL and returns it""" # Arrange - mock_get_url.return_value = "https://storage.example.com/signed-upload-url" + mock_get_uri.return_value = "/files/upload/for-plugin?sign=1" + monkeypatch.setattr(plugin_module.dify_config, "INTERNAL_FILES_URL", "http://api:5001") mock_tenant = MagicMock() mock_tenant.id = "tenant-id" mock_user = MagicMock() @@ -282,14 +283,14 @@ class TestPluginUploadFileRequestApi: result = raw_post(api_instance, user_model=mock_user, tenant_model=mock_tenant, payload=mock_payload) # Assert - mock_get_url.assert_called_once_with( + mock_get_uri.assert_called_once_with( filename="test.pdf", mimetype="application/pdf", tenant_id="tenant-id", user_id="user-id", conversation_id="conversation-id", ) - assert result["data"]["url"] == "https://storage.example.com/signed-upload-url" + assert result["data"]["url"] == "http://api:5001/files/upload/for-plugin?sign=1" class TestPluginDownloadFileRequestApi: @@ -304,6 +305,13 @@ class TestPluginDownloadFileRequestApi: assert callable(api_instance.post) @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) + @pytest.mark.parametrize( + ("for_external", "expected_url"), + [ + (True, "https://files.example.com/files/tools/report.pdf?sign=1"), + (False, "http://api:5001/files/tools/report.pdf?sign=1"), + ], + ) @patch("controllers.inner_api.plugin.plugin.FileRequestService") def test_post_returns_signed_download_url( self, @@ -312,6 +320,8 @@ class TestPluginDownloadFileRequestApi: app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session, + for_external: bool, + expected_url: str, ): tenant = Tenant( name="Plugin Tenant", @@ -324,18 +334,20 @@ class TestPluginDownloadFileRequestApi: sqlite_session.commit() monkeypatch.setattr(plugin_module.db, "session", sqlite_session) mock_service = mock_service_cls.return_value - mock_service.request_download_url.return_value = MagicMock( + mock_service.request_download.return_value = MagicMock( filename="report.pdf", mime_type="application/pdf", size=123, - download_url="https://files.example.com/download", + download_uri="/files/tools/report.pdf?sign=1", ) + monkeypatch.setattr(plugin_module.dify_config, "FILES_URL", "https://files.example.com") + monkeypatch.setattr(plugin_module.dify_config, "INTERNAL_FILES_URL", "http://api:5001") mock_payload = MagicMock() mock_payload.tenant_id = tenant.id mock_payload.user_id = "user-id" mock_payload.user_from = "account" mock_payload.invoke_from = "debugger" - mock_payload.for_external = False + mock_payload.for_external = for_external reference = build_file_reference(record_id="tool-file-1") mock_payload.file.model_dump.return_value = { "transfer_method": "tool_file", @@ -345,19 +357,18 @@ class TestPluginDownloadFileRequestApi: raw_post = _extract_raw_post(PluginDownloadFileRequestApi) result = raw_post(api_instance, payload=mock_payload) - mock_service.request_download_url.assert_called_once_with( + mock_service.request_download.assert_called_once_with( tenant_id=tenant.id, user_id="user-id", user_from="account", invoke_from="debugger", file_mapping={"transfer_method": "tool_file", "reference": reference}, - for_external=False, ) assert result["data"] == { "filename": "report.pdf", "mime_type": "application/pdf", "size": 123, - "download_url": "https://files.example.com/download", + "download_url": expected_url, } diff --git a/api/tests/unit_tests/controllers/inner_api/test_agent_files.py b/api/tests/unit_tests/controllers/inner_api/test_agent_files.py new file mode 100644 index 00000000000..d6a90456129 --- /dev/null +++ b/api/tests/unit_tests/controllers/inner_api/test_agent_files.py @@ -0,0 +1,125 @@ +import inspect +from collections.abc import Callable +from types import SimpleNamespace +from typing import cast +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask +from sqlalchemy.orm import Session + +from controllers.inner_api.agent.files import ( + AgentFileDownloadRequestApi, + AgentFileUploadRequestApi, +) +from core.workflow.file_reference import build_file_reference +from services.file_request_service import DownloadFileRequestResult + +MODULE = "controllers.inner_api.agent.files" + + +def _raw[R](method: Callable[..., R]) -> Callable[..., R]: + return cast(Callable[..., R], inspect.unwrap(method)) + + +def test_upload_request_returns_origin_free_uri(app: Flask, unbound_session: Session) -> None: + payload = { + "tenant_id": "tenant-1", + "user_id": "execution-user-1", + "filename": "report.pdf", + "mimetype": "application/pdf", + "conversation_id": "conversation-1", + } + tenant = SimpleNamespace(id="tenant-1") + user = SimpleNamespace(id="canonical-end-user-1") + session = unbound_session + with app.test_request_context("/", method="POST", json=payload): + with ( + patch(f"{MODULE}.TenantService") as tenant_service, + patch(f"{MODULE}.get_user", return_value=user), + patch(f"{MODULE}.get_signed_file_uri_for_plugin", return_value="/files/upload/for-plugin?sign=1") as sign, + ): + tenant_service.get_tenant_by_id.return_value = tenant + response = _raw(AgentFileUploadRequestApi.post)(AgentFileUploadRequestApi(), session) + + assert response == {"upload_uri": "/files/upload/for-plugin?sign=1"} + tenant_service.get_tenant_by_id.assert_called_once_with("tenant-1", session=session) + sign.assert_called_once_with( + filename="report.pdf", + mimetype="application/pdf", + tenant_id="tenant-1", + user_id="canonical-end-user-1", + conversation_id="conversation-1", + user_from=None, + ) + + +def test_download_request_returns_origin_free_uri_for_sandbox(app: Flask, unbound_session: Session) -> None: + reference = build_file_reference(record_id="tool-file-1") + payload = { + "tenant_id": "tenant-1", + "user_id": "user-1", + "user_from": "account", + "invoke_from": "debugger", + "file": {"transfer_method": "tool_file", "reference": reference}, + "for_frontend": False, + } + session = unbound_session + with app.test_request_context("/", method="POST", json=payload): + with ( + patch(f"{MODULE}.TenantService") as tenant_service, + patch(f"{MODULE}.FileRequestService") as service, + ): + tenant_service.get_tenant_by_id.return_value = MagicMock() + service.return_value.request_download.return_value = DownloadFileRequestResult( + filename="report.pdf", + mime_type="application/pdf", + size=123, + download_uri="/files/tools/tool-file-1.pdf?sign=1", + ) + response = _raw(AgentFileDownloadRequestApi.post)(AgentFileDownloadRequestApi(), session) + + assert response == { + "filename": "report.pdf", + "mime_type": "application/pdf", + "size": 123, + "download_uri": "/files/tools/tool-file-1.pdf?sign=1", + } + service.return_value.request_download.assert_called_once_with( + tenant_id="tenant-1", + user_id="user-1", + user_from="account", + invoke_from="debugger", + file_mapping={"transfer_method": "tool_file", "reference": reference}, + ) + + +def test_download_request_binds_frontend_url( + app: Flask, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: + reference = build_file_reference(record_id="tool-file-1") + payload = { + "tenant_id": "tenant-1", + "user_id": "user-1", + "user_from": "account", + "invoke_from": "debugger", + "file": {"transfer_method": "tool_file", "reference": reference}, + "for_frontend": True, + } + monkeypatch.setattr(f"{MODULE}.dify_config.FILES_URL", "https://files.example.com") + session = unbound_session + with app.test_request_context("/", method="POST", json=payload): + with ( + patch(f"{MODULE}.TenantService") as tenant_service, + patch(f"{MODULE}.FileRequestService") as service, + ): + tenant_service.get_tenant_by_id.return_value = MagicMock() + service.return_value.request_download.return_value = DownloadFileRequestResult( + filename="report.pdf", + mime_type="application/pdf", + size=123, + download_uri="/files/tools/tool-file-1.pdf?sign=1", + ) + response = _raw(AgentFileDownloadRequestApi.post)(AgentFileDownloadRequestApi(), session) + + assert response["download_uri"] == "https://files.example.com/files/tools/tool-file-1.pdf?sign=1" diff --git a/api/tests/unit_tests/controllers/inner_api/test_agent_tools.py b/api/tests/unit_tests/controllers/inner_api/test_agent_tools.py index ea53555807e..2a92ecbbc10 100644 --- a/api/tests/unit_tests/controllers/inner_api/test_agent_tools.py +++ b/api/tests/unit_tests/controllers/inner_api/test_agent_tools.py @@ -1,6 +1,6 @@ """Unit tests for the Agent tool inner API controller.""" -from collections.abc import Iterator +from collections.abc import Generator from contextlib import contextmanager from unittest.mock import patch @@ -46,7 +46,7 @@ def _payload() -> dict[str, object]: @contextmanager -def _agent_inner_auth() -> Iterator[None]: +def _agent_inner_auth() -> Generator[None]: with ( patch("configs.dify_config.PLUGIN_DAEMON_KEY", "plugin-daemon-key"), patch("configs.dify_config.INNER_API_KEY_FOR_PLUGIN", "inner-key"), diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_composition.py b/api/tests/unit_tests/controllers/openapi/auth/test_composition.py index 11cd1aa1380..41d4416efc5 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_composition.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_composition.py @@ -13,6 +13,7 @@ from controllers.openapi.auth.verify import ( check_workspace_role, ) from core.rbac import RBACPermission, RBACResourceScope +from enums import DeploymentEdition from libs.oauth_bearer import Scope, TokenType from models.account import TenantAccountRole from services.enterprise.enterprise_service import WebAppAccessMode @@ -70,12 +71,10 @@ def test_router_routes_contain_both_token_types(): assert TokenType.OAUTH_EXTERNAL_SSO in auth_router._routes -def test_external_sso_route_has_ee_required_edition(): +def test_external_sso_route_requires_enterprise_edition(): route = auth_router._routes[TokenType.OAUTH_EXTERNAL_SSO] assert isinstance(route, PipelineRoute) - from controllers.openapi.auth.data import Edition - - assert route.required_edition == frozenset({Edition.EE}) + assert route.required_edition == frozenset({DeploymentEdition.ENTERPRISE}) def test_account_route_has_no_required_edition(): @@ -147,7 +146,7 @@ def _selected_webapp_steps(*, scope, app_access_mode): """ from unittest.mock import MagicMock, patch - from controllers.openapi.auth.data import AuthData, Edition + from controllers.openapi.auth.data import AuthData ctx = RequestContext( token_type=TokenType.OAUTH_ACCOUNT, @@ -164,7 +163,10 @@ def _selected_webapp_steps(*, scope, app_access_mode): features.webapp_auth.enabled = True selected = [] with ( - patch("controllers.openapi.auth.conditions.current_edition", return_value=Edition.EE), + patch( + "controllers.openapi.auth.conditions.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.ENTERPRISE, + ), patch("controllers.openapi.auth.conditions.FeatureService.get_system_features", return_value=features), ): for step in account_pipeline._auth: diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_conditions.py b/api/tests/unit_tests/controllers/openapi/auth/test_conditions.py index fa882ca96c4..b9cd877f0bf 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_conditions.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_conditions.py @@ -1,9 +1,9 @@ from unittest.mock import MagicMock, patch from controllers.openapi.auth.conditions import ( - EDITION_CE, - EDITION_EE, - EDITION_SAAS, + EDITION_CLOUD, + EDITION_COMMUNITY, + EDITION_ENTERPRISE, HAS_ALLOWED_ROLES, HAS_RBAC, LOADED_APP_IS_PRIVATE, @@ -18,8 +18,9 @@ from controllers.openapi.auth.conditions import ( data_cond, request_cond, ) -from controllers.openapi.auth.data import AuthData, Edition, RBACRequirement, RequestContext +from controllers.openapi.auth.data import AuthData, RBACRequirement, RequestContext from core.rbac import RBACPermission, RBACResourceScope +from enums import DeploymentEdition from libs.oauth_bearer import Scope, TokenType from models.account import TenantAccountRole from services.enterprise.enterprise_service import WebAppAccessMode @@ -115,22 +116,31 @@ def test_path_has_app_id_false(): assert PATH_HAS_APP_ID(_ctx(path_params={})) is False -def test_edition_ce(): - with patch("controllers.openapi.auth.conditions.current_edition", return_value=Edition.CE): - assert EDITION_CE(_ctx()) is True - assert EDITION_EE(_ctx()) is False - assert EDITION_SAAS(_ctx()) is False +def test_edition_community(): + with patch( + "controllers.openapi.auth.conditions.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ): + assert EDITION_COMMUNITY(_ctx()) is True + assert EDITION_ENTERPRISE(_ctx()) is False + assert EDITION_CLOUD(_ctx()) is False -def test_edition_ee(): - with patch("controllers.openapi.auth.conditions.current_edition", return_value=Edition.EE): - assert EDITION_EE(_ctx()) is True - assert EDITION_CE(_ctx()) is False +def test_edition_enterprise(): + with patch( + "controllers.openapi.auth.conditions.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.ENTERPRISE, + ): + assert EDITION_ENTERPRISE(_ctx()) is True + assert EDITION_COMMUNITY(_ctx()) is False -def test_edition_saas(): - with patch("controllers.openapi.auth.conditions.current_edition", return_value=Edition.SAAS): - assert EDITION_SAAS(_ctx()) is True +def test_edition_cloud(): + with patch( + "controllers.openapi.auth.conditions.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.CLOUD, + ): + assert EDITION_CLOUD(_ctx()) is True def test_webapp_auth_enabled(): diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_data.py b/api/tests/unit_tests/controllers/openapi/auth/test_data.py index 7ee1a83ec00..e5171e64303 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_data.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_data.py @@ -1,41 +1,16 @@ import uuid -from unittest.mock import patch import pytest from pydantic import ValidationError from controllers.openapi.auth.data import ( AuthData, - Edition, ExternalIdentity, RequestContext, - current_edition, ) -from enums.deployment_edition import DeploymentEdition from libs.oauth_bearer import Scope, TokenType -def test_current_edition_saas(): - with patch("controllers.openapi.auth.data.dify_config") as cfg: - cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD - cfg.ENTERPRISE_ENABLED = True - assert current_edition() == Edition.SAAS - - -def test_current_edition_ee(): - with patch("controllers.openapi.auth.data.dify_config") as cfg: - cfg.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE - cfg.ENTERPRISE_ENABLED = True - assert current_edition() == Edition.EE - - -def test_current_edition_ce(): - with patch("controllers.openapi.auth.data.dify_config") as cfg: - cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY - cfg.ENTERPRISE_ENABLED = False - assert current_edition() == Edition.CE - - def test_external_identity_frozen(): ei = ExternalIdentity(email="a@b.com", issuer="idp") with pytest.raises(ValidationError): diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_pipeline.py b/api/tests/unit_tests/controllers/openapi/auth/test_pipeline.py index 37b400b92d0..a9f5df5aae6 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_pipeline.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_pipeline.py @@ -5,8 +5,9 @@ import pytest from flask import Flask from werkzeug.exceptions import Forbidden, NotFound, Unauthorized -from controllers.openapi.auth.data import AuthData, Edition +from controllers.openapi.auth.data import AuthData from controllers.openapi.auth.pipeline import AuthPipeline, PipelineRoute, PipelineRouter +from enums import DeploymentEdition from libs.oauth_bearer import Scope, TokenType @@ -75,9 +76,12 @@ def test_guard_edition_gate_returns_404(app): router = _make_router() with app.test_request_context("/test"): - with patch("controllers.openapi.auth.pipeline.current_edition", return_value=Edition.CE): + with patch( + "controllers.openapi.auth.pipeline.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ): - @router.guard(scope=Scope.FULL, edition=frozenset({Edition.EE})) + @router.guard(scope=Scope.FULL, edition=frozenset({DeploymentEdition.ENTERPRISE})) def view(*, auth_data): pass @@ -93,7 +97,10 @@ def test_guard_token_type_gate_returns_403(app): patch("controllers.openapi.auth.pipeline.extract_bearer", return_value="tok"), patch("controllers.openapi.auth.pipeline.get_authenticator") as mock_auth, patch("controllers.openapi.auth.pipeline.emit_wrong_surface"), - patch("controllers.openapi.auth.pipeline.current_edition", return_value=Edition.CE), + patch( + "controllers.openapi.auth.pipeline.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ), ): identity = _fake_identity() identity.token_type = TokenType.OAUTH_EXTERNAL_SSO @@ -114,7 +121,10 @@ def test_guard_unregistered_token_type_returns_403(app): with ( patch("controllers.openapi.auth.pipeline.extract_bearer", return_value="tok"), patch("controllers.openapi.auth.pipeline.get_authenticator") as mock_auth, - patch("controllers.openapi.auth.pipeline.current_edition", return_value=Edition.CE), + patch( + "controllers.openapi.auth.pipeline.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ), ): identity = _fake_identity() identity.token_type = TokenType.OAUTH_EXTERNAL_SSO @@ -196,14 +206,17 @@ def test_guard_resets_auth_ctx_on_exception(app): def test_router_rejects_token_type_on_wrong_edition(app): pipeline = AuthPipeline(prepare=[], auth=[]) - route = PipelineRoute(pipeline, required_edition=frozenset({Edition.EE})) + route = PipelineRoute(pipeline, required_edition=frozenset({DeploymentEdition.ENTERPRISE})) router = PipelineRouter({TokenType.OAUTH_EXTERNAL_SSO: route}) with app.test_request_context("/test", headers={"Authorization": "Bearer tok"}): with ( patch("controllers.openapi.auth.pipeline.extract_bearer", return_value="tok"), patch("controllers.openapi.auth.pipeline.get_authenticator") as mock_auth, - patch("controllers.openapi.auth.pipeline.current_edition", return_value=Edition.CE), + patch( + "controllers.openapi.auth.pipeline.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ), ): identity = _make_identity(token_type=TokenType.OAUTH_EXTERNAL_SSO) mock_auth.return_value.authenticate.return_value = identity diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py b/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py index 3b714a84da1..0fc691152f2 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_prepare.py @@ -4,6 +4,7 @@ from unittest.mock import MagicMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound, Unauthorized from controllers.openapi.auth.data import AuthData, ExternalIdentity @@ -150,10 +151,10 @@ def test_load_account_skips_when_already_set(): assert data.caller is existing_caller -def test_load_account_sets_current_tenant_when_tenant_present(): +def test_load_account_sets_current_tenant_when_tenant_present(sqlite_session: Session): account = MagicMock() tenant = MagicMock() - session = MagicMock() + session = sqlite_session data = _make_auth_data(account_id=uuid.uuid4(), tenant=tenant) with ( patch("controllers.openapi.auth.prepare.AccountService.get_account_by_id", return_value=account), diff --git a/api/tests/unit_tests/controllers/openapi/auth/test_verify.py b/api/tests/unit_tests/controllers/openapi/auth/test_verify.py index 012b9880a12..af6515ea7a4 100644 --- a/api/tests/unit_tests/controllers/openapi/auth/test_verify.py +++ b/api/tests/unit_tests/controllers/openapi/auth/test_verify.py @@ -1,5 +1,5 @@ import uuid -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest from flask import Flask @@ -61,7 +61,7 @@ def test_check_app_access_passes_when_tenant_none(): def test_check_app_access_passes_when_member(): - tenant = MagicMock(spec=Tenant) + tenant = Tenant(name="Test Tenant") tenant.id = "t1" data = _data(account_id=uuid.uuid4(), tenant=tenant) with patch("controllers.openapi.auth.verify.TenantService.account_belongs_to_tenant", return_value=True): @@ -69,7 +69,7 @@ def test_check_app_access_passes_when_member(): def test_check_app_access_raises_when_not_member(): - tenant = MagicMock(spec=Tenant) + tenant = Tenant(name="Test Tenant") tenant.id = "t1" data = _data(account_id=uuid.uuid4(), tenant=tenant) with patch("controllers.openapi.auth.verify.TenantService.account_belongs_to_tenant", return_value=False): @@ -113,7 +113,7 @@ def test_check_rbac_raises_when_context_missing(): def test_check_rbac_enforces_for_account_caller(): - tenant = MagicMock(spec=Tenant) + tenant = Tenant(name="Test Tenant") tenant.id = "t1" account_id = uuid.uuid4() data = _data( @@ -144,13 +144,13 @@ def test_check_acl_raises_when_app_or_mode_missing(): def test_check_acl_account_allowed_for_public(): - app = MagicMock(spec=App) + app = App() data = _data(token_type=TokenType.OAUTH_ACCOUNT, app=app, app_access_mode=WebAppAccessMode.PUBLIC) check_acl(data) def test_check_acl_external_sso_blocked_for_private(): - app = MagicMock(spec=App) + app = App() data = _data( token_type=TokenType.OAUTH_EXTERNAL_SSO, app=app, @@ -161,7 +161,7 @@ def test_check_acl_external_sso_blocked_for_private(): def test_check_acl_external_sso_allowed_for_sso_verified(): - app = MagicMock(spec=App) + app = App() data = _data( token_type=TokenType.OAUTH_EXTERNAL_SSO, app=app, @@ -176,8 +176,9 @@ def test_check_private_app_permission_raises_when_app_none(): def test_check_private_app_permission_raises_when_user_not_allowed(): - app = MagicMock(spec=App) - app.id = "app-1" + app = App( + id="app-1", + ) data = _data(account_id=uuid.uuid4(), app=app) target = "controllers.openapi.auth.verify.EnterpriseService.WebAppAuth.is_user_allowed_to_access_webapp" with patch(target, return_value=False): @@ -186,8 +187,9 @@ def test_check_private_app_permission_raises_when_user_not_allowed(): def test_check_private_app_permission_passes_when_allowed(): - app = MagicMock(spec=App) - app.id = "app-1" + app = App( + id="app-1", + ) data = _data(account_id=uuid.uuid4(), app=app) target = "controllers.openapi.auth.verify.EnterpriseService.WebAppAuth.is_user_allowed_to_access_webapp" with patch(target, return_value=True): @@ -208,7 +210,7 @@ def test_check_workspace_mismatch_passes_when_tenant_none(flask_app): def test_check_workspace_mismatch_passes_when_ids_match(flask_app): - tenant = MagicMock(spec=Tenant) + tenant = Tenant(name="Test Tenant") tid = uuid.uuid4() tenant.id = tid with flask_app.test_request_context(f"/test?workspace_id={tid}"): @@ -218,7 +220,7 @@ def test_check_workspace_mismatch_passes_when_ids_match(flask_app): def test_check_workspace_mismatch_raises_422_on_mismatch(flask_app): from werkzeug.exceptions import UnprocessableEntity - tenant = MagicMock(spec=Tenant) + tenant = Tenant(name="Test Tenant") tenant.id = uuid.uuid4() other_id = uuid.uuid4() with flask_app.test_request_context(f"/test?workspace_id={other_id}"): @@ -227,7 +229,7 @@ def test_check_workspace_mismatch_raises_422_on_mismatch(flask_app): def test_check_workspace_mismatch_passes_when_no_request_workspace_id(flask_app): - tenant = MagicMock(spec=Tenant) + tenant = Tenant(name="Test Tenant") tenant.id = uuid.uuid4() with flask_app.test_request_context("/test"): check_workspace_mismatch(_data(tenant=tenant, path_params={})) @@ -267,14 +269,16 @@ def test_check_workspace_role_passes_when_role_allowed(): def test_check_app_api_enabled_passes_when_enabled(): - app = MagicMock(spec=App) - app.enable_api = True + app = App( + enable_api=True, + ) check_app_api_enabled(_data(app=app)) def test_check_app_api_enabled_raises_forbidden_when_disabled(): - app = MagicMock(spec=App) - app.enable_api = False + app = App( + enable_api=False, + ) with pytest.raises(Forbidden, match="service_api_disabled"): check_app_api_enabled(_data(app=app)) diff --git a/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py b/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py index f5e9b16d29c..e852189b76b 100644 --- a/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py +++ b/api/tests/unit_tests/controllers/openapi/test_app_describe_builder.py @@ -1,6 +1,8 @@ from types import SimpleNamespace from unittest.mock import MagicMock +from sqlalchemy.orm import Session + from controllers.openapi._input_schema import EMPTY_INPUT_SCHEMA from controllers.openapi.apps import _EMPTY_PARAMETERS, build_app_describe_response from controllers.service_api.app.error import AppUnavailableError @@ -25,9 +27,9 @@ def _app() -> _FakeApp: ) -def test_fields_none_returns_all_blocks(monkeypatch): +def test_fields_none_returns_all_blocks(monkeypatch, unbound_session: Session): app = _app() - session = MagicMock() + session = unbound_session parameters_payload = MagicMock(return_value={"k": "v"}) input_schema = MagicMock(return_value={"s": 1}) monkeypatch.setattr("controllers.openapi.apps.parameters_payload", parameters_payload) @@ -41,8 +43,8 @@ def test_fields_none_returns_all_blocks(monkeypatch): input_schema.assert_called_once_with(app, session=session) -def test_fields_subset_limits_blocks(monkeypatch): - session = MagicMock() +def test_fields_subset_limits_blocks(monkeypatch, unbound_session: Session): + session = unbound_session monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={"k": "v"})) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={"s": 1})) resp = build_app_describe_response(_app(), ["info"], session=session) @@ -51,8 +53,8 @@ def test_fields_subset_limits_blocks(monkeypatch): assert resp.input_schema is None -def test_info_omits_author_and_tags(monkeypatch): - session = MagicMock() +def test_info_omits_author_and_tags(monkeypatch, unbound_session: Session): + session = unbound_session monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={})) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={})) resp = build_app_describe_response(_app(), ["info"], session=session) @@ -62,21 +64,21 @@ def test_info_omits_author_and_tags(monkeypatch): assert not hasattr(resp.info, "tags") -def test_parameters_fallback_on_app_unavailable(monkeypatch): +def test_parameters_fallback_on_app_unavailable(monkeypatch, unbound_session: Session): def _raise(app, *, session): raise AppUnavailableError() monkeypatch.setattr("controllers.openapi.apps.parameters_payload", _raise) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", MagicMock(return_value={"s": 1})) - resp = build_app_describe_response(_app(), ["parameters"], session=MagicMock()) + resp = build_app_describe_response(_app(), ["parameters"], session=unbound_session) assert resp.parameters == dict(_EMPTY_PARAMETERS) -def test_input_schema_fallback_on_app_unavailable(monkeypatch): +def test_input_schema_fallback_on_app_unavailable(monkeypatch, unbound_session: Session): def _raise(app, *, session): raise AppUnavailableError() monkeypatch.setattr("controllers.openapi.apps.parameters_payload", MagicMock(return_value={"k": "v"})) monkeypatch.setattr("controllers.openapi.apps.build_input_schema", _raise) - resp = build_app_describe_response(_app(), ["input_schema"], session=MagicMock()) + resp = build_app_describe_response(_app(), ["input_schema"], session=unbound_session) assert resp.input_schema == dict(EMPTY_INPUT_SCHEMA) diff --git a/api/tests/unit_tests/controllers/openapi/test_app_payloads.py b/api/tests/unit_tests/controllers/openapi/test_app_payloads.py index 2e9e7bc06a8..083a7bd980b 100644 --- a/api/tests/unit_tests/controllers/openapi/test_app_payloads.py +++ b/api/tests/unit_tests/controllers/openapi/test_app_payloads.py @@ -8,6 +8,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from controllers.openapi.apps import ( # pyright: ignore[reportPrivateUsage] _EMPTY_PARAMETERS, @@ -35,10 +36,10 @@ def _fake_app(**overrides): return SimpleNamespace(**base) -def test_parameters_payload_raises_app_unavailable_when_no_config(): +def test_parameters_payload_raises_app_unavailable_when_no_config(unbound_session: Session): app = _fake_app(mode="chat") app.app_model_config_with_session = MagicMock(return_value=None) - session = MagicMock() + session = unbound_session with pytest.raises(AppUnavailableError): parameters_payload(app, session=session) diff --git a/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py b/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py index 49f7bea5cd3..038a17ad1df 100644 --- a/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py +++ b/api/tests/unit_tests/controllers/openapi/test_apps_permitted_external_query.py @@ -14,6 +14,7 @@ from unittest.mock import MagicMock, patch import pytest from pydantic import ValidationError +from sqlalchemy.orm import Session from controllers.openapi.apps_permitted_external import ( PermittedExternalAppDescribeApi, @@ -69,10 +70,10 @@ def test_query_accepts_valid_mode(): assert q.mode.value == "chat" -def test_describe_forwards_request_session_to_response_builder(): +def test_describe_forwards_request_session_to_response_builder(unbound_session: Session): api = PermittedExternalAppDescribeApi() method = inspect.unwrap(api.get) - session = MagicMock() + session = unbound_session app = MagicMock() auth_data = SimpleNamespace(app=app) query = SimpleNamespace(fields={"info"}) diff --git a/api/tests/unit_tests/controllers/openapi/test_device_sso.py b/api/tests/unit_tests/controllers/openapi/test_device_sso.py index 655a5940f47..38c5249bbc3 100644 --- a/api/tests/unit_tests/controllers/openapi/test_device_sso.py +++ b/api/tests/unit_tests/controllers/openapi/test_device_sso.py @@ -143,7 +143,7 @@ def test_device_error_redirect_drops_malformed_user_code(): def _ee_features(): - from services.feature_service import LicenseStatus + from services.entities.feature_entities import LicenseStatus m = MagicMock() m.license.status = LicenseStatus.ACTIVE diff --git a/api/tests/unit_tests/controllers/openapi/test_input_schema.py b/api/tests/unit_tests/controllers/openapi/test_input_schema.py index 133072ad33e..bb042cad089 100644 --- a/api/tests/unit_tests/controllers/openapi/test_input_schema.py +++ b/api/tests/unit_tests/controllers/openapi/test_input_schema.py @@ -3,6 +3,7 @@ from __future__ import annotations import pytest +from sqlalchemy.orm import Session from controllers.openapi._input_schema import _form_to_jsonschema @@ -97,6 +98,7 @@ from models.model import AppMode def _stub_app(mode: AppMode, *, form: list[dict] | None = None, has_workflow: bool | None = None): """Returns a MagicMock whose explicit config getters are wired up.""" app = MagicMock() + app.id = "00000000-0000-0000-0000-000000000001" app.mode = mode if mode in (AppMode.WORKFLOW, AppMode.ADVANCED_CHAT): if has_workflow is False: @@ -111,20 +113,15 @@ def _stub_app(mode: AppMode, *, form: list[dict] | None = None, has_workflow: bo app.app_model_config_with_session.return_value = None else: app_model_config = MagicMock() + app_model_config.app_id = app.id app_model_config.to_dict.return_value = {"user_input_form": form or []} app.app_model_config_with_session.return_value = app_model_config return app -def _session() -> MagicMock: - session = MagicMock() - session.scalar.return_value = None - return session - - -def test_chat_mode_includes_query() -> None: +def test_chat_mode_includes_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.CHAT, form=[{"text-input": {"variable": "x", "label": "X", "required": True}}]) - session = _session() + session = sqlite_session schema = build_input_schema(app, session=session) assert schema["$schema"] == "https://json-schema.org/draft/2020-12/schema" assert "query" in schema["properties"] @@ -137,33 +134,33 @@ def test_chat_mode_includes_query() -> None: app.app_model_config_with_session.return_value.to_dict.assert_called_once_with(annotation_reply={"enabled": False}) -def test_agent_chat_mode_includes_query() -> None: +def test_agent_chat_mode_includes_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.AGENT_CHAT, form=[]) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert "query" in schema["properties"] -def test_advanced_chat_mode_includes_query() -> None: +def test_advanced_chat_mode_includes_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.ADVANCED_CHAT, form=[]) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert "query" in schema["properties"] -def test_workflow_mode_omits_query() -> None: +def test_workflow_mode_omits_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.WORKFLOW, form=[]) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert "query" not in schema["properties"] assert schema["required"] == ["inputs"] -def test_completion_mode_omits_query() -> None: +def test_completion_mode_omits_query(sqlite_session: Session) -> None: app = _stub_app(AppMode.COMPLETION, form=[]) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert "query" not in schema["properties"] assert schema["required"] == ["inputs"] -def test_inputs_required_driven_by_form() -> None: +def test_inputs_required_driven_by_form(sqlite_session: Session) -> None: app = _stub_app( AppMode.CHAT, form=[ @@ -171,20 +168,20 @@ def test_inputs_required_driven_by_form() -> None: {"text-input": {"variable": "context", "label": "Context", "required": False}}, ], ) - schema = build_input_schema(app, session=_session()) + schema = build_input_schema(app, session=sqlite_session) assert schema["properties"]["inputs"]["required"] == ["industry"] -def test_misconfigured_chat_raises_app_unavailable() -> None: +def test_misconfigured_chat_raises_app_unavailable(sqlite_session: Session) -> None: app = _stub_app(AppMode.CHAT, has_workflow=False) with pytest.raises(AppUnavailableError): - build_input_schema(app, session=_session()) + build_input_schema(app, session=sqlite_session) -def test_misconfigured_workflow_raises_app_unavailable() -> None: +def test_misconfigured_workflow_raises_app_unavailable(sqlite_session: Session) -> None: app = _stub_app(AppMode.WORKFLOW, has_workflow=False) with pytest.raises(AppUnavailableError): - build_input_schema(app, session=_session()) + build_input_schema(app, session=sqlite_session) def test_empty_input_schema_sentinel_shape() -> None: diff --git a/api/tests/unit_tests/controllers/openapi/test_meta_version.py b/api/tests/unit_tests/controllers/openapi/test_meta_version.py index 57de0c517f4..3da3c4fca21 100644 --- a/api/tests/unit_tests/controllers/openapi/test_meta_version.py +++ b/api/tests/unit_tests/controllers/openapi/test_meta_version.py @@ -4,6 +4,8 @@ from __future__ import annotations import pytest +from enums import DeploymentEdition + def test_version_endpoint_returns_200_without_auth(openapi_app): client = openapi_app.test_client() @@ -15,7 +17,7 @@ def test_version_endpoint_returns_200_without_auth(openapi_app): assert "version" in payload assert "edition" in payload assert isinstance(payload["version"], str) - assert payload["edition"] in ("SELF_HOSTED", "CLOUD") + assert payload["edition"] in {edition.value for edition in DeploymentEdition} def test_version_endpoint_ignores_bearer_header(openapi_app): @@ -35,7 +37,7 @@ def test_version_endpoint_ignores_bearer_header(openapi_app): def test_version_endpoint_reflects_edition_config(openapi_app, monkeypatch: pytest.MonkeyPatch): from configs import dify_config - monkeypatch.setattr(dify_config, "EDITION", "CLOUD") + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) client = openapi_app.test_client() response = client.get("/openapi/v1/_version") @@ -44,13 +46,13 @@ def test_version_endpoint_reflects_edition_config(openapi_app, monkeypatch: pyte assert response.get_json()["edition"] == "CLOUD" -def test_version_endpoint_falls_back_to_self_hosted_on_unexpected_edition(openapi_app, monkeypatch: pytest.MonkeyPatch): +def test_version_endpoint_reflects_enterprise_edition(openapi_app, monkeypatch: pytest.MonkeyPatch): from configs import dify_config - monkeypatch.setattr(dify_config, "EDITION", "EXPERIMENTAL") + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) client = openapi_app.test_client() response = client.get("/openapi/v1/_version") assert response.status_code == 200 - assert response.get_json()["edition"] == "SELF_HOSTED" + assert response.get_json()["edition"] == "ENTERPRISE" diff --git a/api/tests/unit_tests/controllers/openapi/test_oauth_sso_claims.py b/api/tests/unit_tests/controllers/openapi/test_oauth_sso_claims.py index 9551a37142f..2d58c9499cc 100644 --- a/api/tests/unit_tests/controllers/openapi/test_oauth_sso_claims.py +++ b/api/tests/unit_tests/controllers/openapi/test_oauth_sso_claims.py @@ -17,7 +17,7 @@ def app() -> Flask: def _ee_features(): - from services.feature_service import LicenseStatus + from services.entities.feature_entities import LicenseStatus m = MagicMock() m.license.status = LicenseStatus.ACTIVE diff --git a/api/tests/unit_tests/controllers/openapi/test_oauth_sso_host_header.py b/api/tests/unit_tests/controllers/openapi/test_oauth_sso_host_header.py index 2eff2d25414..95f31482e9f 100644 --- a/api/tests/unit_tests/controllers/openapi/test_oauth_sso_host_header.py +++ b/api/tests/unit_tests/controllers/openapi/test_oauth_sso_host_header.py @@ -18,7 +18,7 @@ def app() -> Flask: def _ee_features(): - from services.feature_service import LicenseStatus + from services.entities.feature_entities import LicenseStatus m = MagicMock() m.license.status = LicenseStatus.ACTIVE diff --git a/api/tests/unit_tests/controllers/service_api/app/test_annotation.py b/api/tests/unit_tests/controllers/service_api/app/test_annotation.py index 6d60af54e2e..866018bd49d 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_annotation.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_annotation.py @@ -110,25 +110,28 @@ class TestAppModelPatterns: def test_app_model_has_required_fields(self): """Test App model has required fields for annotation operations.""" - app = Mock(spec=App) - app.id = str(uuid.uuid4()) - app.status = "normal" - app.enable_api = True + app = App( + id=str(uuid.uuid4()), + status="normal", + enable_api=True, + ) assert app.id is not None assert app.status == "normal" assert app.enable_api def test_app_model_disabled_api(self): """Test app with disabled API access.""" - app = Mock(spec=App) - app.enable_api = False + app = App( + enable_api=False, + ) assert not app.enable_api def test_app_model_archived_status(self): """Test app with archived status.""" - app = Mock(spec=App) - app.status = "archived" + app = App( + status="archived", + ) assert app.status == "archived" diff --git a/api/tests/unit_tests/controllers/service_api/app/test_app.py b/api/tests/unit_tests/controllers/service_api/app/test_app.py index c48cb343950..25d31328e4a 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_app.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_app.py @@ -1,606 +1,423 @@ -""" -Unit tests for Service API App controllers +"""SQLite-backed tests for Service API application controllers. + +The authentication decorator resolves the app, tenant, and tenant owner before +the controller runs. Controller/model code then reads configuration, workflow, +tags, and author information through two additional references to the database +extension. Tests bind all of those references to one explicit scoped SQLite +session and persist visibility and cross-tenant decoys instead of fabricating ORM +lookup results. """ -import uuid -from unittest.mock import ANY, Mock, patch +import json +from collections.abc import Iterator +from dataclasses import dataclass +from unittest.mock import Mock +from uuid import uuid4 import pytest from flask import Flask +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker +from werkzeug.exceptions import Forbidden, Unauthorized +from controllers.service_api.app import app as app_controller from controllers.service_api.app.app import AppInfoApi, AppMetaApi, AppParameterApi from controllers.service_api.app.error import AgentNotPublishedError, AppUnavailableError from core.app.apps.agent_app.errors import AgentAppNotPublishedError -from models.account import TenantStatus -from models.model import App, AppMode -from tests.unit_tests.conftest import setup_mock_tenant_owner_execute_result +from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus +from models.base import TypeBase +from models.enums import EndUserType +from models.model import ( + App, + AppAnnotationSetting, + AppMode, + AppModelConfig, + CustomizeTokenStrategy, + EndUser, + Site, + Tag, + TagBinding, + TagType, +) +from models.workflow import Workflow, WorkflowType -def _configure_current_app_mock(mock_current_app): - mock_current_app.login_manager = Mock() - mock_current_app._get_current_object = Mock(return_value=Mock()) +@dataclass(frozen=True) +class _DatabaseBinding: + engine: Engine + session: scoped_session[Session] -class TestAppParameterApi: - """Test suite for AppParameterApi""" +@dataclass(frozen=True) +class _Token: + app_id: str + tenant_id: str - @pytest.fixture - def app(self): - """Create Flask test application.""" - app = Flask(__name__) - app.config["TESTING"] = True - return app - @pytest.fixture - def mock_app_model(self): - """Create a mock App model.""" - app = Mock(spec=App) - app.id = str(uuid.uuid4()) - app.tenant_id = str(uuid.uuid4()) - app.mode = AppMode.CHAT - app.status = "normal" - app.enable_api = True - app.app_model_config_with_session.return_value = None - app.workflow_with_session.return_value = None - return app +@dataclass(frozen=True) +class AppDatabase: + """Persisted application graph used by the decorated controller methods.""" - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - def test_get_parameters_for_chat_app( - self, mock_db, mock_validate_token, mock_current_app, mock_user_logged_in, app: Flask, mock_app_model - ): - """Test retrieving parameters for a chat app.""" - # Arrange - _configure_current_app_mock(mock_current_app) + session_maker: sessionmaker[Session] + registry: scoped_session[Session] + tenant_id: str + app_id: str + owner_id: str + config_id: str + workflow_id: str - mock_config = Mock() - mock_config.id = str(uuid.uuid4()) - mock_config.to_dict.return_value = { - "user_input_form": [{"type": "text", "label": "Name", "variable": "name", "required": True}], - "suggested_questions": [], - } - mock_app_model.app_model_config = mock_config - mock_app_model.app_model_config_with_session.return_value = mock_config - mock_app_model.workflow = None + def update_app(self, **values: object) -> None: + with self.session_maker.begin() as session: + app = session.get_one(App, self.app_id) + for key, value in values.items(): + setattr(app, key, value) - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app_model.id - mock_api_token.tenant_id = mock_app_model.tenant_id - mock_validate_token.return_value = mock_api_token + def update_tenant(self, **values: object) -> None: + with self.session_maker.begin() as session: + tenant = session.get_one(Tenant, self.tenant_id) + for key, value in values.items(): + setattr(tenant, key, value) - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL + def delete_row(self, model: type[object], object_id: str) -> None: + table = model.__table__ # type: ignore[attr-defined] + with self.session_maker.begin() as session: + session.execute(table.delete().where(table.c.id == object_id)) - # Mock DB queries for app and tenant - mock_db.session.get.side_effect = [ - mock_app_model, - mock_tenant, - ] + def replace_tags(self, *names: str) -> None: + """Replace the visible app's tenant-owned tag bindings with persisted tags.""" + with self.session_maker.begin() as session: + session.execute( + TagBinding.__table__.delete().where( + TagBinding.tenant_id == self.tenant_id, + TagBinding.target_id == self.app_id, + ) + ) + tags = [ + Tag(tenant_id=self.tenant_id, type=TagType.APP, name=name, created_by=self.owner_id) for name in names + ] + session.add_all(tags) + session.flush() + session.add_all( + TagBinding( + tenant_id=self.tenant_id, + tag_id=tag.id, + target_id=self.app_id, + created_by=self.owner_id, + ) + for tag in tags + ) - # Mock tenant owner info for login - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - # Act - with ( - app.test_request_context("/parameters", method="GET", headers={"Authorization": "Bearer test_token"}), - patch( - "controllers.service_api.app.app.load_annotation_reply_config", - return_value={"enabled": False}, - ), - ): - api = AppParameterApi() - response = api.get() +@pytest.fixture +def flask_app() -> Flask: + app = Flask(__name__) + app.config["TESTING"] = True + return app - # Assert - assert "opening_statement" in response - assert "suggested_questions" in response - assert "user_input_form" in response - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - def test_get_parameters_for_workflow_app( - self, mock_db, mock_validate_token, mock_current_app, mock_user_logged_in, app: Flask, mock_app_model - ): - """Test retrieving parameters for a workflow app.""" - # Arrange - _configure_current_app_mock(mock_current_app) +@pytest.fixture +def app_db(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[AppDatabase]: + """Create the minimal controller/model schema and bind every DB reference explicitly.""" - mock_app_model.mode = AppMode.WORKFLOW - mock_workflow = Mock() - mock_workflow.features_dict = {"suggested_questions": []} - mock_workflow.user_input_form.return_value = [{"type": "text", "label": "Input", "variable": "input"}] - mock_app_model.workflow = mock_workflow - mock_app_model.workflow_with_session.return_value = mock_workflow - mock_app_model.app_model_config = None + tables = [ + Tenant.__table__, + Account.__table__, + TenantAccountJoin.__table__, + App.__table__, + AppModelConfig.__table__, + AppAnnotationSetting.__table__, + Workflow.__table__, + Tag.__table__, + TagBinding.__table__, + Site.__table__, + EndUser.__table__, + ] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + registry = scoped_session(maker) + binding = _DatabaseBinding(engine=sqlite_engine, session=registry) + monkeypatch.setattr("controllers.service_api.wraps.db", binding) + monkeypatch.setattr(app_controller, "db", binding) + monkeypatch.setattr("models.model.db", binding) + monkeypatch.setattr("models.account.db", binding) - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app_model.id - mock_api_token.tenant_id = mock_app_model.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - - mock_db.session.get.side_effect = [ - mock_app_model, - mock_tenant, - ] - - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - - # Act - with app.test_request_context("/parameters", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppParameterApi() - response = api.get() - - # Assert - assert "user_input_form" in response - assert "opening_statement" in response - - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - @patch("controllers.service_api.app.app._get_agent_app_feature_dict_and_user_input_form") - def test_get_parameters_for_agent_app( - self, - mock_get_agent_parameters, - mock_db, - mock_validate_token, - mock_current_app, - mock_user_logged_in, - app: Flask, - mock_app_model, - ): - """Test retrieving parameters for an Agent App from Agent Soul app variables.""" - _configure_current_app_mock(mock_current_app) - - mock_app_model.mode = AppMode.AGENT - mock_app_model.app_model_config = None - mock_app_model.workflow = None - mock_get_agent_parameters.return_value = ( - {"opening_statement": "Hi from Agent"}, - [{"text-input": {"label": "topic", "variable": "topic", "required": True}}], + tenant_id = str(uuid4()) + app_id = str(uuid4()) + owner_id = str(uuid4()) + config_id = str(uuid4()) + workflow_id = str(uuid4()) + other_tenant_id = str(uuid4()) + with maker.begin() as session: + tenant = Tenant(name="Visible tenant") + tenant.id = tenant_id + other_tenant = Tenant(name="Other tenant") + other_tenant.id = other_tenant_id + owner = Account(name="Test Author", email="owner@example.com") + owner.id = owner_id + app = App( + id=app_id, + tenant_id=tenant_id, + name="Test App", + description="A test application", + mode=AppMode.CHAT, + icon_type=None, + icon=None, + icon_background=None, + app_model_config_id=config_id, + workflow_id=workflow_id, + enable_site=True, + enable_api=True, + max_active_requests=None, + created_by=owner_id, + ) + config = AppModelConfig( + app_id=app_id, + opening_statement="Hello", + suggested_questions=json.dumps(["Question?"]), + user_input_form=json.dumps([{"text-input": {"label": "Name", "variable": "name", "required": True}}]), + ) + config.id = config_id + workflow = Workflow.new( + tenant_id=tenant_id, + app_id=app_id, + type=WorkflowType.WORKFLOW.value, + version="1", + graph=json.dumps({"nodes": [{"id": "start", "data": {"type": "start", "variables": []}}]}), + features=json.dumps({"suggested_questions": []}), + created_by=owner_id, + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + workflow.id = workflow_id + target_tag = Tag(tenant_id=tenant_id, type=TagType.APP, name="test-tag", created_by=owner_id) + other_tag = Tag(tenant_id=other_tenant_id, type=TagType.APP, name="foreign-tag", created_by=owner_id) + session.add_all([tenant, other_tenant, owner, app, config, workflow, target_tag, other_tag]) + session.flush() + session.add_all( + [ + TenantAccountJoin( + tenant_id=tenant_id, + account_id=owner_id, + current=True, + role=TenantAccountRole.OWNER, + ), + TagBinding( + tenant_id=tenant_id, + tag_id=target_tag.id, + target_id=app_id, + created_by=owner_id, + ), + # Same target ID but another tenant: App.tags must exclude it. + TagBinding( + tenant_id=other_tenant_id, + tag_id=other_tag.id, + target_id=app_id, + created_by=owner_id, + ), + Site( + app_id=app_id, + title="Published site", + icon_type=None, + icon=None, + icon_background=None, + description="Site decoy", + default_language="en-US", + customize_token_strategy=CustomizeTokenStrategy.MUST, + code="visible-site", + ), + EndUser( + tenant_id=other_tenant_id, + app_id=app_id, + type=EndUserType.BROWSER, + name="Cross-tenant visitor", + is_anonymous=False, + session_id="visitor-session", + ), + ] ) - mock_api_token = Mock() - mock_api_token.app_id = mock_app_model.id - mock_api_token.tenant_id = mock_app_model.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - mock_db.session.get.side_effect = [mock_app_model, mock_tenant] - - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - - with app.test_request_context("/parameters", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppParameterApi() - response = api.get() - - assert response["opening_statement"] == "Hi from Agent" - assert response["user_input_form"] == [ - {"text-input": {"label": "topic", "variable": "topic", "required": True}} - ] - mock_get_agent_parameters.assert_called_once_with(mock_app_model, session=ANY) - - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - @patch( - "controllers.service_api.app.app.get_published_agent_app_feature_dict_and_user_input_form", - side_effect=AgentAppNotPublishedError("Agent has not been published"), + database = AppDatabase( + session_maker=maker, + registry=registry, + tenant_id=tenant_id, + app_id=app_id, + owner_id=owner_id, + config_id=config_id, + workflow_id=workflow_id, ) - def test_get_parameters_for_unpublished_agent_app_raises_friendly_error( - self, - mock_get_agent_parameters, - mock_db, - mock_validate_token, - mock_current_app, - mock_user_logged_in, - app: Flask, - mock_app_model, - ): - _configure_current_app_mock(mock_current_app) - - mock_app_model.mode = AppMode.AGENT - mock_api_token = Mock() - mock_api_token.app_id = mock_app_model.id - mock_api_token.tenant_id = mock_app_model.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - mock_db.session.get.side_effect = [mock_app_model, mock_tenant] - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, Mock(current_tenant=mock_tenant)) - - with app.test_request_context("/parameters", method="GET", headers={"Authorization": "Bearer test_token"}): - with pytest.raises(AgentNotPublishedError): - AppParameterApi().get() - - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - def test_get_parameters_raises_error_when_chat_config_missing( - self, mock_db, mock_validate_token, mock_current_app, mock_user_logged_in, app: Flask, mock_app_model - ): - """Test that AppUnavailableError is raised when chat app has no config.""" - # Arrange - _configure_current_app_mock(mock_current_app) - - mock_app_model.app_model_config = None - mock_app_model.workflow = None - - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app_model.id - mock_api_token.tenant_id = mock_app_model.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - - mock_db.session.get.side_effect = [ - mock_app_model, - mock_tenant, - ] - - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - - # Act & Assert - with app.test_request_context("/parameters", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppParameterApi() - with pytest.raises(AppUnavailableError): - api.get() - - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - def test_get_parameters_raises_error_when_workflow_missing( - self, mock_db, mock_validate_token, mock_current_app, mock_user_logged_in, app: Flask, mock_app_model - ): - """Test that AppUnavailableError is raised when workflow app has no workflow.""" - # Arrange - _configure_current_app_mock(mock_current_app) - - mock_app_model.mode = AppMode.WORKFLOW - mock_app_model.workflow = None - mock_app_model.app_model_config = None - - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app_model.id - mock_api_token.tenant_id = mock_app_model.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - - mock_db.session.get.side_effect = [ - mock_app_model, - mock_tenant, - ] - - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - - # Act & Assert - with app.test_request_context("/parameters", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppParameterApi() - with pytest.raises(AppUnavailableError): - api.get() + try: + yield database + finally: + registry.remove() -class TestAppMetaApi: - """Test suite for AppMetaApi""" +@pytest.fixture +def authenticated_controller(app_db: AppDatabase, monkeypatch: pytest.MonkeyPatch) -> Iterator[AppDatabase]: + """Patch only token validation and Flask login signaling around real ORM auth.""" - @pytest.fixture - def app(self): - """Create Flask test application.""" - app = Flask(__name__) - app.config["TESTING"] = True - return app - - @pytest.fixture - def mock_app_model(self): - """Create a mock App model.""" - app = Mock(spec=App) - app.id = str(uuid.uuid4()) - app.status = "normal" - app.enable_api = True - return app - - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - @patch("controllers.service_api.app.app.AppService") - def test_get_app_meta( - self, - mock_app_service, - mock_db, - mock_validate_token, - mock_current_app, - mock_user_logged_in, - app: Flask, - mock_app_model, - ): - """Test retrieving app metadata via AppService.""" - # Arrange - _configure_current_app_mock(mock_current_app) - - mock_service_instance = Mock() - mock_service_instance.get_app_meta.return_value = { - "tool_icons": {}, - "AgentIcons": {}, - } - mock_app_service.return_value = mock_service_instance - - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app_model.id - mock_api_token.tenant_id = mock_app_model.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - - mock_db.session.get.side_effect = [ - mock_app_model, - mock_tenant, - ] - - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - - # Act - with app.test_request_context("/meta", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppMetaApi() - response = api.get() - - # Assert - mock_service_instance.get_app_meta.assert_called_once_with(mock_app_model, session=ANY) - assert response == {"tool_icons": {}, "AgentIcons": {}} - - -class TestAppInfoApi: - """Test suite for AppInfoApi""" - - @pytest.fixture - def app(self): - """Create Flask test application.""" - app = Flask(__name__) - app.config["TESTING"] = True - return app - - @pytest.fixture - def mock_app_model(self): - """Create a mock App model with all required attributes.""" - app = Mock(spec=App) - app.id = str(uuid.uuid4()) - app.tenant_id = str(uuid.uuid4()) - app.name = "Test App" - app.description = "A test application" - app.mode = AppMode.CHAT - app.author_name = "Test Author" - app.status = "normal" - app.enable_api = True - - # Mock tags relationship - mock_tag = Mock() - mock_tag.name = "test-tag" - app.tags = [mock_tag] - - return app - - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - def test_get_app_info( - self, mock_db, mock_validate_token, mock_current_app, mock_user_logged_in, app: Flask, mock_app_model - ): - """Test retrieving basic app information.""" - _configure_current_app_mock(mock_current_app) - - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app_model.id - mock_api_token.tenant_id = mock_app_model.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - - mock_db.session.get.side_effect = [ - mock_app_model, - mock_tenant, - ] - - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - - # Act - with app.test_request_context("/info", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppInfoApi() - response = api.get() - - # Assert - assert response["name"] == "Test App" - assert response["description"] == "A test application" - assert response["tags"] == ["test-tag"] - assert response["mode"] == AppMode.CHAT - assert response["author_name"] == "Test Author" - - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - def test_get_app_info_with_multiple_tags( - self, mock_db, mock_validate_token, mock_current_app, mock_user_logged_in, app - ): - """Test retrieving app info with multiple tags.""" - # Arrange - _configure_current_app_mock(mock_current_app) - - mock_app = Mock(spec=App) - mock_app.id = str(uuid.uuid4()) - mock_app.tenant_id = str(uuid.uuid4()) - mock_app.name = "Multi Tag App" - mock_app.description = "App with multiple tags" - mock_app.mode = AppMode.WORKFLOW - mock_app.author_name = "Author" - mock_app.status = "normal" - mock_app.enable_api = True - - tag1, tag2, tag3 = Mock(), Mock(), Mock() - tag1.name = "tag-one" - tag2.name = "tag-two" - tag3.name = "tag-three" - mock_app.tags = [tag1, tag2, tag3] - - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app.id - mock_api_token.tenant_id = mock_app.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - - mock_db.session.get.side_effect = [ - mock_app, - mock_tenant, - ] - - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - - # Act - with app.test_request_context("/info", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppInfoApi() - response = api.get() - - # Assert - assert response["tags"] == ["tag-one", "tag-two", "tag-three"] - - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - def test_get_app_info_with_no_tags( - self, mock_db, mock_validate_token, mock_current_app, mock_user_logged_in, app: Flask - ): - """Test retrieving app info when app has no tags.""" - # Arrange - _configure_current_app_mock(mock_current_app) - - mock_app = Mock(spec=App) - mock_app.id = str(uuid.uuid4()) - mock_app.tenant_id = str(uuid.uuid4()) - mock_app.name = "No Tags App" - mock_app.description = "App without tags" - mock_app.mode = AppMode.COMPLETION - mock_app.author_name = "Author" - mock_app.tags = [] - mock_app.status = "normal" - mock_app.enable_api = True - - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app.id - mock_api_token.tenant_id = mock_app.tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - - mock_db.session.get.side_effect = [ - mock_app, - mock_tenant, - ] - - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) - - # Act - with app.test_request_context("/info", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppInfoApi() - response = api.get() - - # Assert - assert response["tags"] == [] - - @pytest.mark.parametrize( - "app_mode", - [AppMode.CHAT, AppMode.COMPLETION, AppMode.WORKFLOW, AppMode.ADVANCED_CHAT], + monkeypatch.setattr( + "controllers.service_api.wraps.validate_and_get_api_token", + Mock(return_value=_Token(app_id=app_db.app_id, tenant_id=app_db.tenant_id)), ) - @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.current_app") - @patch("controllers.service_api.wraps.validate_and_get_api_token") - @patch("controllers.service_api.wraps.db") - def test_get_app_info_returns_correct_mode( - self, mock_db, mock_validate_token, mock_current_app, mock_user_logged_in, app: Flask, app_mode - ): - """Test that all app modes are correctly returned.""" - # Arrange - _configure_current_app_mock(mock_current_app) + current_app = Mock() + current_app.login_manager = Mock() + current_app._get_current_object.return_value = Mock() + monkeypatch.setattr("controllers.service_api.wraps.current_app", current_app) + monkeypatch.setattr("controllers.service_api.wraps.user_logged_in", Mock()) + return app_db - mock_app = Mock(spec=App) - mock_app.id = str(uuid.uuid4()) - mock_app.tenant_id = str(uuid.uuid4()) - mock_app.name = "Test" - mock_app.description = "Test" - mock_app.mode = app_mode - mock_app.author_name = "Test" - mock_app.tags = [] - mock_app.status = "normal" - mock_app.enable_api = True - # Mock authentication - mock_api_token = Mock() - mock_api_token.app_id = mock_app.id - mock_api_token.tenant_id = mock_app.tenant_id - mock_validate_token.return_value = mock_api_token +@pytest.mark.usefixtures("authenticated_controller") +def test_get_parameters_for_persisted_chat_config(flask_app: Flask) -> None: + with flask_app.test_request_context("/parameters", headers={"Authorization": "Bearer token"}): + response = AppParameterApi().get() - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL + assert response["opening_statement"] == "Hello" + assert response["suggested_questions"] == ["Question?"] + assert response["user_input_form"] == [{"text-input": {"label": "Name", "variable": "name", "required": True}}] - mock_db.session.get.side_effect = [ - mock_app, - mock_tenant, - ] - mock_account = Mock() - mock_account.current_tenant = mock_tenant - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) +def test_get_parameters_for_persisted_workflow(flask_app: Flask, authenticated_controller: AppDatabase) -> None: + authenticated_controller.update_app(mode=AppMode.WORKFLOW, app_model_config_id=None) - # Act - with app.test_request_context("/info", method="GET", headers={"Authorization": "Bearer test_token"}): - api = AppInfoApi() - response = api.get() + with flask_app.test_request_context("/parameters", headers={"Authorization": "Bearer token"}): + response = AppParameterApi().get() - # Assert - assert response["mode"] == app_mode + assert response["user_input_form"] == [] + assert response["suggested_questions"] == [] + + +def test_get_parameters_for_agent_uses_persisted_app( + flask_app: Flask, authenticated_controller: AppDatabase, monkeypatch: pytest.MonkeyPatch +) -> None: + authenticated_controller.update_app(mode=AppMode.AGENT, app_model_config_id=None, workflow_id=None) + user_input_form = [{"text-input": {"label": "Topic", "variable": "topic", "required": True}}] + agent_parameters = Mock( + return_value=( + {"opening_statement": "Hi from Agent"}, + user_input_form, + ) + ) + monkeypatch.setattr(app_controller, "_get_agent_app_feature_dict_and_user_input_form", agent_parameters) + + with flask_app.test_request_context("/parameters", headers={"Authorization": "Bearer token"}): + response = AppParameterApi().get() + + assert response["opening_statement"] == "Hi from Agent" + assert response["user_input_form"] == user_input_form + agent_parameters.assert_called_once() + (app_model,) = agent_parameters.call_args.args + assert app_model.id == authenticated_controller.app_id + assert agent_parameters.call_args.kwargs["session"] is authenticated_controller.registry() + + +def test_unpublished_agent_raises_friendly_error( + flask_app: Flask, authenticated_controller: AppDatabase, monkeypatch: pytest.MonkeyPatch +) -> None: + authenticated_controller.update_app(mode=AppMode.AGENT) + monkeypatch.setattr( + app_controller, + "get_published_agent_app_feature_dict_and_user_input_form", + Mock(side_effect=AgentAppNotPublishedError("not published")), + ) + + with flask_app.test_request_context("/parameters", headers={"Authorization": "Bearer token"}): + with pytest.raises(AgentNotPublishedError): + AppParameterApi().get() + + +@pytest.mark.parametrize( + ("mode", "field"), + [(AppMode.CHAT, "app_model_config_id"), (AppMode.WORKFLOW, "workflow_id")], +) +def test_parameters_reject_missing_persisted_configuration( + flask_app: Flask, + authenticated_controller: AppDatabase, + mode: AppMode, + field: str, +) -> None: + authenticated_controller.update_app(mode=mode, **{field: None}) + + with flask_app.test_request_context("/parameters", headers={"Authorization": "Bearer token"}): + with pytest.raises(AppUnavailableError): + AppParameterApi().get() + + +def test_get_meta_passes_real_session_and_app( + flask_app: Flask, authenticated_controller: AppDatabase, monkeypatch: pytest.MonkeyPatch +) -> None: + service = Mock() + service.get_app_meta.return_value = {"tool_icons": {}} + monkeypatch.setattr(app_controller, "AppService", Mock(return_value=service)) + + with flask_app.test_request_context("/meta", headers={"Authorization": "Bearer token"}): + response = AppMetaApi().get() + + service.get_app_meta.assert_called_once() + (app_model,) = service.get_app_meta.call_args.args + assert app_model.id == authenticated_controller.app_id + assert service.get_app_meta.call_args.kwargs["session"] is authenticated_controller.registry() + assert response == {"tool_icons": {}} + + +@pytest.mark.parametrize("mode", [AppMode.CHAT, AppMode.COMPLETION, AppMode.WORKFLOW, AppMode.ADVANCED_CHAT]) +def test_get_info_reads_author_and_tenant_scoped_tags( + flask_app: Flask, + authenticated_controller: AppDatabase, + mode: AppMode, +) -> None: + authenticated_controller.update_app(mode=mode) + + with flask_app.test_request_context("/info", headers={"Authorization": "Bearer token"}): + response = AppInfoApi().get() + + assert response == { + "name": "Test App", + "description": "A test application", + "tags": ["test-tag"], + "mode": mode, + "author_name": "Test Author", + } + + +@pytest.mark.parametrize( + "tag_names", + [(), ("tag-one", "tag-two", "tag-three")], + ids=["zero-tags", "multiple-tags"], +) +def test_get_info_handles_zero_or_multiple_tags( + flask_app: Flask, + authenticated_controller: AppDatabase, + tag_names: tuple[str, ...], +) -> None: + authenticated_controller.replace_tags(*tag_names) + + with flask_app.test_request_context("/info", headers={"Authorization": "Bearer token"}): + response = AppInfoApi().get() + + assert len(response["tags"]) == len(tag_names) + assert set(response["tags"]) == set(tag_names) + + +@pytest.mark.parametrize("state", ["missing", "disabled", "archived", "ownerless"]) +def test_authentication_rejects_empty_or_invisible_database_state( + flask_app: Flask, + authenticated_controller: AppDatabase, + state: str, +) -> None: + expected_error: type[Exception] = Forbidden + if state == "missing": + authenticated_controller.delete_row(App, authenticated_controller.app_id) + elif state == "disabled": + authenticated_controller.update_app(enable_api=False) + elif state == "archived": + authenticated_controller.update_tenant(status=TenantStatus.ARCHIVE) + else: + with authenticated_controller.session_maker.begin() as session: + session.execute(TenantAccountJoin.__table__.delete()) + expected_error = Unauthorized + + with flask_app.test_request_context("/info", headers={"Authorization": "Bearer token"}): + with pytest.raises(expected_error): + AppInfoApi().get() diff --git a/api/tests/unit_tests/controllers/service_api/app/test_audio.py b/api/tests/unit_tests/controllers/service_api/app/test_audio.py index 2019891d281..4fb56e6ac8b 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_audio.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_audio.py @@ -143,11 +143,12 @@ class TestAudioServiceMockedBehavior: @pytest.fixture def mock_app(self): - """Create mock app model.""" + """Create an app model.""" from models.model import App - app = Mock(spec=App) - app.id = str(uuid.uuid4()) + app = App( + id=str(uuid.uuid4()), + ) return app @pytest.fixture @@ -159,7 +160,7 @@ class TestAudioServiceMockedBehavior: return mock @patch.object(AudioService, "transcript_asr") - def test_transcript_asr_returns_response(self, mock_asr, mock_app, mock_file): + def test_transcript_asr_returns_response(self, mock_asr, mock_app, mock_file, sqlite_session: Session): """Test ASR transcription returns response dict.""" mock_response = {"text": "Transcribed text"} mock_asr.return_value = mock_response @@ -167,14 +168,13 @@ class TestAudioServiceMockedBehavior: result = AudioService.transcript_asr( app_model=mock_app, file=mock_file, - session=Mock(), + session=sqlite_session, end_user="user_123", ) assert result["text"] == "Transcribed text" @patch.object(AudioService, "transcript_tts") - @pytest.mark.parametrize("sqlite_session", [()], indirect=True) def test_transcript_tts_returns_response(self, mock_tts, mock_app, sqlite_session: Session): """Test TTS transcription returns response.""" mock_response = {"audio": "base64_audio_data"} diff --git a/api/tests/unit_tests/controllers/service_api/app/test_completion.py b/api/tests/unit_tests/controllers/service_api/app/test_completion.py index de3ffa82018..39c986e3ca2 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_completion.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_completion.py @@ -42,7 +42,7 @@ from controllers.service_api.app.error import ( ) from core.app.apps.agent_app.errors import AgentAppNotPublishedError from core.errors.error import QuotaExceededError -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from graphon.model_runtime.errors.invoke import InvokeError from models.base import TypeBase from models.enums import ConversationFromSource, EndUserType @@ -556,7 +556,7 @@ class TestChatApiController: self, app: Flask, monkeypatch: pytest.MonkeyPatch, orm_session: Session ) -> None: completion_module = sys.modules["controllers.service_api.app.completion"] - monkeypatch.setattr(completion_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(completion_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) billing_get_info = Mock(return_value={"enabled": True, "subscription": {"plan": CloudPlan.SANDBOX}}) generate = Mock() @@ -582,12 +582,13 @@ class TestChatApiController: assert exc_info.value.error_code == "workflow_version_execution_not_allowed" @pytest.mark.parametrize( - ("billing_config_enabled", "billing_enabled", "plan", "workflow_id"), + ("deployment_edition", "billing_enabled", "plan", "workflow_id"), [ - (False, True, CloudPlan.SANDBOX, str(uuid.uuid4())), - (True, False, CloudPlan.SANDBOX, str(uuid.uuid4())), - (True, True, CloudPlan.PROFESSIONAL, str(uuid.uuid4())), - (True, True, CloudPlan.SANDBOX, None), + (DeploymentEdition.COMMUNITY, True, CloudPlan.SANDBOX, str(uuid.uuid4())), + (DeploymentEdition.ENTERPRISE, True, CloudPlan.SANDBOX, str(uuid.uuid4())), + (DeploymentEdition.CLOUD, False, CloudPlan.SANDBOX, str(uuid.uuid4())), + (DeploymentEdition.CLOUD, True, CloudPlan.PROFESSIONAL, str(uuid.uuid4())), + (DeploymentEdition.CLOUD, True, CloudPlan.SANDBOX, None), ], ) def test_allows_default_or_entitled_workflow_version_execution( @@ -595,13 +596,13 @@ class TestChatApiController: app: Flask, monkeypatch: pytest.MonkeyPatch, orm_session: Session, - billing_config_enabled: bool, + deployment_edition: DeploymentEdition, billing_enabled: bool, plan: CloudPlan, workflow_id: str | None, ) -> None: completion_module = sys.modules["controllers.service_api.app.completion"] - monkeypatch.setattr(completion_module.dify_config, "BILLING_ENABLED", billing_config_enabled) + monkeypatch.setattr(completion_module.dify_config, "DEPLOYMENT_EDITION", deployment_edition) billing_get_info = Mock(return_value={"enabled": billing_enabled, "subscription": {"plan": plan}}) generate = Mock(return_value={"result": "ok"}) @@ -621,7 +622,7 @@ class TestChatApiController: assert response == {"result": "ok"} generate.assert_called_once() - if billing_config_enabled and workflow_id: + if deployment_edition == DeploymentEdition.CLOUD and workflow_id: billing_get_info.assert_called_once_with(app_model.tenant_id, exclude_vector_space=True) else: billing_get_info.assert_not_called() diff --git a/api/tests/unit_tests/controllers/service_api/app/test_conversation.py b/api/tests/unit_tests/controllers/service_api/app/test_conversation.py index 09f0d430eff..ac79b565f7a 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_conversation.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_conversation.py @@ -56,8 +56,9 @@ from services.errors.conversation import ( def _end_user(user_id: str = "end-user-1") -> EndUser: - end_user = EndUser() - end_user.id = user_id + end_user = EndUser( + id=user_id, + ) return end_user @@ -375,8 +376,9 @@ class TestConversationAppModeValidation: Verifies that CHAT, AGENT_CHAT, AGENT, and ADVANCED_CHAT modes pass validation without raising NotChatAppError. """ - app = Mock(spec=App) - app.mode = mode + app = App( + mode=mode, + ) # Validation should pass without raising for chat modes app_mode = AppMode.value_of(app.mode) @@ -388,8 +390,9 @@ class TestConversationAppModeValidation: Verifies that calling a conversation endpoint with a COMPLETION mode app raises NotChatAppError. """ - app = Mock(spec=App) - app.mode = AppMode.COMPLETION + app = App( + mode=AppMode.COMPLETION, + ) app_mode = AppMode.value_of(app.mode) assert app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT} @@ -402,8 +405,9 @@ class TestConversationAppModeValidation: Verifies that calling a conversation endpoint with a WORKFLOW mode app raises NotChatAppError. """ - app = Mock(spec=App) - app.mode = AppMode.WORKFLOW + app = App( + mode=AppMode.WORKFLOW, + ) app_mode = AppMode.value_of(app.mode) assert app_mode not in {AppMode.CHAT, AppMode.AGENT_CHAT, AppMode.ADVANCED_CHAT, AppMode.AGENT} @@ -478,8 +482,8 @@ class TestConversationService: mock_pagination.return_value = mock_result result = ConversationService.pagination_by_last_id( - app_model=Mock(spec=App), - user=Mock(spec=EndUser), + app_model=App(), + user=EndUser(), last_id=None, limit=20, invoke_from=Mock(), @@ -490,7 +494,6 @@ class TestConversationService: assert hasattr(result, "limit") assert hasattr(result, "has_more") - @pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True) def test_rename_returns_conversation(self, sqlite_session: Session): """Test rename returns updated conversation.""" conversation_id = "00000000-0000-0000-0000-000000000001" @@ -498,8 +501,9 @@ class TestConversationService: sqlite_session.add(conversation) sqlite_session.commit() - app_model = App() - app_model.id = "app-1" + app_model = App( + id="app-1", + ) end_user = _end_user() result = ConversationService.rename( @@ -537,7 +541,6 @@ class TestConversationApiController: with pytest.raises(NotChatAppError): handler(api, app_model=app_model, end_user=end_user) - @pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True) def test_list_last_not_found( self, app: Flask, diff --git a/api/tests/unit_tests/controllers/service_api/app/test_file.py b/api/tests/unit_tests/controllers/service_api/app/test_file.py index 184723a06ae..2a714600be9 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_file.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_file.py @@ -210,9 +210,10 @@ from inspect import unwrap def mock_app_model(): from models import App - app = Mock(spec=App) - app.id = str(uuid.uuid4()) - app.tenant_id = str(uuid.uuid4()) + app = App( + id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + ) return app @@ -220,8 +221,9 @@ def mock_app_model(): def mock_end_user(): from models import EndUser - user = Mock(spec=EndUser) - user.id = str(uuid.uuid4()) + user = EndUser( + id=str(uuid.uuid4()), + ) return user diff --git a/api/tests/unit_tests/controllers/service_api/app/test_file_preview.py b/api/tests/unit_tests/controllers/service_api/app/test_file_preview.py index 14a00a9af82..33b6e83094e 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_file_preview.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_file_preview.py @@ -1,17 +1,31 @@ -""" -Unit tests for Service API File Preview endpoint +"""Unit tests for the Service API file-preview endpoint. + +Ownership checks run against persisted message, file, app, and upload rows so the +tests exercise the same SQLAlchemy statements and tenant boundary as production. +Storage remains mocked because it is the external I/O boundary of the endpoint. """ import logging -import uuid +from collections.abc import Iterator +from dataclasses import dataclass +from datetime import datetime +from decimal import Decimal from typing import Protocol, cast from unittest.mock import Mock, patch +from uuid import uuid4 import pytest +from sqlalchemy import event +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from controllers.service_api.app.error import FileAccessDeniedError, FileNotFoundError from controllers.service_api.app.file_preview import FilePreviewApi -from models.model import App, EndUser, Message, MessageFile, UploadFile +from extensions.storage.storage_type import StorageType +from graphon.file import FileTransferMethod, FileType +from models.base import TypeBase +from models.enums import ConversationFromSource, CreatorUserRole +from models.model import App, AppMode, Message, MessageFile, UploadFile class _FilePreviewLogRecord(Protocol): @@ -20,367 +34,252 @@ class _FilePreviewLogRecord(Protocol): error: str +@dataclass(frozen=True) +class _Database: + """Expose the real test session through the interface used by the controller.""" + + session: Session + + +@dataclass(frozen=True) +class _PreviewRecords: + app: App + message: Message + message_file: MessageFile + upload_file: UploadFile + + +@pytest.fixture +def database(sqlite_engine: Engine) -> Iterator[_Database]: + """Create only the tables required by file ownership validation.""" + + models = (App, Message, MessageFile, UploadFile) + tables = [TypeBase.metadata.tables[model.__tablename__] for model in models] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + with Session(sqlite_engine, expire_on_commit=False) as session: + yield _Database(session) + + +@pytest.fixture +def file_preview_api() -> FilePreviewApi: + """Create the resource instance under test.""" + + return FilePreviewApi() + + +def _upload_file(*, tenant_id: str, file_id: str | None = None) -> UploadFile: + upload_file = UploadFile( + tenant_id=tenant_id, + storage_type=StorageType.LOCAL, + key="storage/key/test_file.jpg", + name="test_file.jpg", + size=1024, + extension="jpg", + mime_type="image/jpeg", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + created_at=datetime(2026, 1, 1), + used=True, + ) + if file_id is not None: + upload_file.id = file_id + return upload_file + + +def _persist_preview_records( + session: Session, + *, + app_id: str | None = None, + app_tenant_id: str | None = None, + upload_tenant_id: str | None = None, +) -> _PreviewRecords: + app_id = app_id or str(uuid4()) + app_tenant_id = app_tenant_id or str(uuid4()) + upload_file = _upload_file(tenant_id=upload_tenant_id or app_tenant_id) + app = App( + id=app_id, + tenant_id=app_tenant_id, + name="Preview app", + description="", + mode=AppMode.CHAT, + icon_type=None, + icon="", + icon_background=None, + enable_site=True, + enable_api=True, + ) + message = Message( + id=str(uuid4()), + app_id=app_id, + conversation_id=str(uuid4()), + _inputs={}, + query="preview", + message={}, + message_unit_price=Decimal(0), + answer="answer", + answer_unit_price=Decimal(0), + currency="USD", + from_source=ConversationFromSource.API, + ) + message_file = MessageFile( + message_id=message.id, + type=FileType.IMAGE, + transfer_method=FileTransferMethod.LOCAL_FILE, + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + upload_file_id=upload_file.id, + ) + session.add_all([app, message, message_file, upload_file]) + session.commit() + return _PreviewRecords(app=app, message=message, message_file=message_file, upload_file=upload_file) + + class TestFilePreviewApi: - """Test suite for FilePreviewApi""" + """Exercise ownership validation and response construction.""" - @pytest.fixture - def file_preview_api(self): - """Create FilePreviewApi instance for testing""" - return FilePreviewApi() + def test_validate_file_ownership_success(self, file_preview_api: FilePreviewApi, database: _Database): + records = _persist_preview_records(database.session) - @pytest.fixture - def mock_app(self): - """Mock App model""" - app = Mock(spec=App) - app.id = str(uuid.uuid4()) - app.tenant_id = str(uuid.uuid4()) - return app + with patch("controllers.service_api.app.file_preview.db", database): + message_file, upload_file = file_preview_api._validate_file_ownership( + records.upload_file.id, records.app.id + ) - @pytest.fixture - def mock_end_user(self): - """Mock EndUser model""" - end_user = Mock(spec=EndUser) - end_user.id = str(uuid.uuid4()) - return end_user + assert message_file.id == records.message_file.id + assert upload_file.id == records.upload_file.id + assert upload_file.tenant_id == records.app.tenant_id - @pytest.fixture - def mock_upload_file(self): - """Mock UploadFile model""" - upload_file = Mock(spec=UploadFile) - upload_file.id = str(uuid.uuid4()) - upload_file.name = "test_file.jpg" - upload_file.extension = "jpg" - upload_file.mime_type = "image/jpeg" - upload_file.size = 1024 - upload_file.key = "storage/key/test_file.jpg" - upload_file.tenant_id = str(uuid.uuid4()) - return upload_file + def test_validate_file_ownership_file_not_found(self, file_preview_api: FilePreviewApi, database: _Database): + with patch("controllers.service_api.app.file_preview.db", database): + with pytest.raises(FileNotFoundError, match="File not found in message context"): + file_preview_api._validate_file_ownership(str(uuid4()), str(uuid4())) - @pytest.fixture - def mock_message_file(self): - """Mock MessageFile model""" - message_file = Mock(spec=MessageFile) - message_file.id = str(uuid.uuid4()) - message_file.upload_file_id = str(uuid.uuid4()) - message_file.message_id = str(uuid.uuid4()) - return message_file + def test_validate_file_ownership_access_denied(self, file_preview_api: FilePreviewApi, database: _Database): + records = _persist_preview_records(database.session) - @pytest.fixture - def mock_message(self): - """Mock Message model""" - message = Mock(spec=Message) - message.id = str(uuid.uuid4()) - message.app_id = str(uuid.uuid4()) - return message + with patch("controllers.service_api.app.file_preview.db", database): + with pytest.raises(FileAccessDeniedError, match="not owned by requesting app"): + file_preview_api._validate_file_ownership(records.upload_file.id, str(uuid4())) - def test_validate_file_ownership_success( - self, file_preview_api: FilePreviewApi, mock_app, mock_upload_file, mock_message_file, mock_message - ): - """Test successful file ownership validation""" - file_id = str(uuid.uuid4()) - app_id = mock_app.id + def test_validate_file_ownership_upload_file_not_found(self, file_preview_api: FilePreviewApi, database: _Database): + records = _persist_preview_records(database.session) + database.session.delete(records.upload_file) + database.session.commit() - # Set up the mocks - mock_upload_file.tenant_id = mock_app.tenant_id - mock_message.app_id = app_id - mock_message_file.upload_file_id = file_id - mock_message_file.message_id = mock_message.id + with patch("controllers.service_api.app.file_preview.db", database): + with pytest.raises(FileNotFoundError, match="Upload file record not found"): + file_preview_api._validate_file_ownership(records.upload_file.id, records.app.id) - with patch("controllers.service_api.app.file_preview.db") as mock_db: - # Mock scalar() for MessageFile and Message queries - mock_db.session.scalar.side_effect = [ - mock_message_file, # MessageFile query - mock_message, # Message query - ] - # Mock get() for UploadFile and App PK lookups - mock_db.session.get.side_effect = [ - mock_upload_file, # UploadFile query - mock_app, # App query for tenant validation - ] + def test_validate_file_ownership_tenant_mismatch(self, file_preview_api: FilePreviewApi, database: _Database): + records = _persist_preview_records(database.session, upload_tenant_id=str(uuid4())) - # Execute the method - result_message_file, result_upload_file = file_preview_api._validate_file_ownership(file_id, app_id) - - # Assertions - assert result_message_file == mock_message_file - assert result_upload_file == mock_upload_file - - def test_validate_file_ownership_file_not_found(self, file_preview_api: FilePreviewApi): - """Test file ownership validation when MessageFile not found""" - file_id = str(uuid.uuid4()) - app_id = str(uuid.uuid4()) - - with patch("controllers.service_api.app.file_preview.db") as mock_db: - # Mock MessageFile not found via scalar() - mock_db.session.scalar.return_value = None - - # Execute and assert exception - with pytest.raises(FileNotFoundError) as exc_info: - file_preview_api._validate_file_ownership(file_id, app_id) - - assert "File not found in message context" in str(exc_info.value) - - def test_validate_file_ownership_access_denied(self, file_preview_api: FilePreviewApi, mock_message_file): - """Test file ownership validation when Message not owned by app""" - file_id = str(uuid.uuid4()) - app_id = str(uuid.uuid4()) - - with patch("controllers.service_api.app.file_preview.db") as mock_db: - # Mock MessageFile found but Message not owned by app via scalar() - mock_db.session.scalar.side_effect = [ - mock_message_file, # MessageFile query - found - None, # Message query - not found (access denied) - ] - - # Execute and assert exception - with pytest.raises(FileAccessDeniedError) as exc_info: - file_preview_api._validate_file_ownership(file_id, app_id) - - assert "not owned by requesting app" in str(exc_info.value) - - def test_validate_file_ownership_upload_file_not_found( - self, file_preview_api: FilePreviewApi, mock_message_file, mock_message - ): - """Test file ownership validation when UploadFile not found""" - file_id = str(uuid.uuid4()) - app_id = str(uuid.uuid4()) - - with patch("controllers.service_api.app.file_preview.db") as mock_db: - # Mock scalar() for MessageFile and Message - mock_db.session.scalar.side_effect = [ - mock_message_file, # MessageFile query - found - mock_message, # Message query - found - ] - # Mock get() for UploadFile - not found - mock_db.session.get.return_value = None - - # Execute and assert exception - with pytest.raises(FileNotFoundError) as exc_info: - file_preview_api._validate_file_ownership(file_id, app_id) - - assert "Upload file record not found" in str(exc_info.value) - - def test_validate_file_ownership_tenant_mismatch( - self, file_preview_api: FilePreviewApi, mock_app, mock_upload_file, mock_message_file, mock_message - ): - """Test file ownership validation with tenant mismatch""" - file_id = str(uuid.uuid4()) - app_id = mock_app.id - - # Set up tenant mismatch - mock_upload_file.tenant_id = "different_tenant_id" - mock_app.tenant_id = "app_tenant_id" - mock_message.app_id = app_id - mock_message_file.upload_file_id = file_id - mock_message_file.message_id = mock_message.id - - with patch("controllers.service_api.app.file_preview.db") as mock_db: - # Mock scalar() for MessageFile and Message queries - mock_db.session.scalar.side_effect = [ - mock_message_file, # MessageFile query - mock_message, # Message query - ] - # Mock get() for UploadFile and App PK lookups - mock_db.session.get.side_effect = [ - mock_upload_file, # UploadFile query - mock_app, # App query for tenant validation - ] - - # Execute and assert exception - with pytest.raises(FileAccessDeniedError) as exc_info: - file_preview_api._validate_file_ownership(file_id, app_id) - - assert "tenant mismatch" in str(exc_info.value) + with patch("controllers.service_api.app.file_preview.db", database): + with pytest.raises(FileAccessDeniedError, match="tenant mismatch"): + file_preview_api._validate_file_ownership(records.upload_file.id, records.app.id) def test_validate_file_ownership_invalid_input(self, file_preview_api: FilePreviewApi): - """Test file ownership validation with invalid input""" - - # Test with empty file_id - with pytest.raises(FileAccessDeniedError) as exc_info: + with pytest.raises(FileAccessDeniedError, match="Invalid file or app identifier"): file_preview_api._validate_file_ownership("", "app_id") - assert "Invalid file or app identifier" in str(exc_info.value) - # Test with empty app_id - with pytest.raises(FileAccessDeniedError) as exc_info: + with pytest.raises(FileAccessDeniedError, match="Invalid file or app identifier"): file_preview_api._validate_file_ownership("file_id", "") - assert "Invalid file or app identifier" in str(exc_info.value) - def test_build_file_response_basic(self, file_preview_api: FilePreviewApi, mock_upload_file): - """Test basic file response building""" - mock_generator = Mock() + @pytest.mark.parametrize( + ("as_attachment", "mime_type", "name", "extension", "size"), + [ + (False, "image/jpeg", "test_file.jpg", "jpg", 1024), + (True, "image/jpeg", "test_file.jpg", "jpg", 1024), + (False, "text/html", "unsafe.html", "html", 1024), + (False, "video/mp4", "test_file.mp4", "mp4", 1024), + (False, "image/jpeg", "test_file.jpg", "jpg", 0), + ], + ) + def test_build_file_response( + self, + file_preview_api: FilePreviewApi, + as_attachment: bool, + mime_type: str, + name: str, + extension: str, + size: int, + ): + upload_file = _upload_file(tenant_id=str(uuid4())) + upload_file.mime_type = mime_type + upload_file.name = name + upload_file.extension = extension + upload_file.size = size - response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False) + response = file_preview_api._build_file_response(Mock(), upload_file, as_attachment) - # Check response properties - assert response.mimetype == mock_upload_file.mime_type assert response.direct_passthrough is True - assert response.headers["Content-Length"] == str(mock_upload_file.size) assert "Cache-Control" in response.headers - - def test_build_file_response_as_attachment(self, file_preview_api: FilePreviewApi, mock_upload_file): - """Test file response building with attachment flag""" - mock_generator = Mock() - - response = file_preview_api._build_file_response(mock_generator, mock_upload_file, True) - - # Check attachment-specific headers - assert "attachment" in response.headers["Content-Disposition"] - assert mock_upload_file.name in response.headers["Content-Disposition"] - assert response.headers["Content-Type"] == "application/octet-stream" - - def test_build_file_response_html_forces_attachment(self, file_preview_api: FilePreviewApi, mock_upload_file): - """Test HTML files are forced to download""" - mock_generator = Mock() - mock_upload_file.mime_type = "text/html" - mock_upload_file.name = "unsafe.html" - mock_upload_file.extension = "html" - - response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False) - - assert "attachment" in response.headers["Content-Disposition"] - assert response.headers["Content-Type"] == "application/octet-stream" - assert response.headers["X-Content-Type-Options"] == "nosniff" - - def test_build_file_response_audio_video(self, file_preview_api: FilePreviewApi, mock_upload_file): - """Test file response building for audio/video files""" - mock_generator = Mock() - mock_upload_file.mime_type = "video/mp4" - - response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False) - - # Check Range support for media files - assert response.headers["Accept-Ranges"] == "bytes" - - def test_build_file_response_no_size(self, file_preview_api: FilePreviewApi, mock_upload_file): - """Test file response building when size is unknown""" - mock_generator = Mock() - mock_upload_file.size = 0 # Unknown size - - response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False) - - # Content-Length should not be set when size is unknown - assert "Content-Length" not in response.headers + assert ("Content-Length" in response.headers) is bool(size) + if as_attachment or mime_type == "text/html": + assert "attachment" in response.headers["Content-Disposition"] + assert response.headers["Content-Type"] == "application/octet-stream" + else: + assert response.mimetype == mime_type + if mime_type == "text/html": + assert response.headers["X-Content-Type-Options"] == "nosniff" + if mime_type.startswith("video/"): + assert response.headers["Accept-Ranges"] == "bytes" @patch("controllers.service_api.app.file_preview.storage") - def test_get_method_integration( - self, - mock_storage, - file_preview_api: FilePreviewApi, - mock_app, - mock_end_user, - mock_upload_file, - mock_message_file, - mock_message, + def test_components_use_validated_file( + self, mock_storage: Mock, file_preview_api: FilePreviewApi, database: _Database ): - """Test the full GET method integration (without decorator)""" - file_id = str(uuid.uuid4()) - app_id = mock_app.id + records = _persist_preview_records(database.session) + generator = Mock() - # Set up mocks - mock_upload_file.tenant_id = mock_app.tenant_id - mock_message.app_id = app_id - mock_message_file.upload_file_id = file_id - mock_message_file.message_id = mock_message.id + with patch("controllers.service_api.app.file_preview.db", database): + message_file, upload_file = file_preview_api._validate_file_ownership( + records.upload_file.id, records.app.id + ) + response = file_preview_api._build_file_response(generator, upload_file, False) - mock_generator = Mock() - mock_storage.load.return_value = mock_generator - - with patch("controllers.service_api.app.file_preview.db") as mock_db: - # Mock scalar() for MessageFile and Message queries - mock_db.session.scalar.side_effect = [ - mock_message_file, # MessageFile query - mock_message, # Message query - ] - # Mock get() for UploadFile and App PK lookups - mock_db.session.get.side_effect = [ - mock_upload_file, # UploadFile query - mock_app, # App query for tenant validation - ] - - # Test the core logic directly without Flask decorators - # Validate file ownership - result_message_file, result_upload_file = file_preview_api._validate_file_ownership(file_id, app_id) - assert result_message_file == mock_message_file - assert result_upload_file == mock_upload_file - - # Test file response building - response = file_preview_api._build_file_response(mock_generator, mock_upload_file, False) - assert response is not None - - # Verify storage was called correctly - mock_storage.load.assert_not_called() # Since we're testing components separately + assert message_file.id == records.message_file.id + assert response.mimetype == "image/jpeg" + mock_storage.load.assert_not_called() @patch("controllers.service_api.app.file_preview.storage") - def test_storage_error_handling( - self, - mock_storage, - file_preview_api: FilePreviewApi, - mock_app, - mock_upload_file, - mock_message_file, - mock_message, + def test_storage_error_remains_external( + self, mock_storage: Mock, file_preview_api: FilePreviewApi, database: _Database ): - """Test storage error handling in the core logic""" - file_id = str(uuid.uuid4()) - app_id = mock_app.id + records = _persist_preview_records(database.session) + mock_storage.load.side_effect = OSError("Storage error") - # Set up mocks - mock_upload_file.tenant_id = mock_app.tenant_id - mock_message.app_id = app_id - mock_message_file.upload_file_id = file_id - mock_message_file.message_id = mock_message.id + with patch("controllers.service_api.app.file_preview.db", database): + _, upload_file = file_preview_api._validate_file_ownership(records.upload_file.id, records.app.id) - # Mock storage error - mock_storage.load.side_effect = Exception("Storage error") - - with patch("controllers.service_api.app.file_preview.db") as mock_db: - # Mock scalar() for MessageFile and Message queries - mock_db.session.scalar.side_effect = [ - mock_message_file, # MessageFile query - mock_message, # Message query - ] - # Mock get() for UploadFile and App PK lookups - mock_db.session.get.side_effect = [ - mock_upload_file, # UploadFile query - mock_app, # App query for tenant validation - ] - - # First validate file ownership works - result_message_file, result_upload_file = file_preview_api._validate_file_ownership(file_id, app_id) - assert result_message_file == mock_message_file - assert result_upload_file == mock_upload_file - - # Test storage error handling - with pytest.raises(Exception) as exc_info: - mock_storage.load(mock_upload_file.key, stream=True) - - assert "Storage error" in str(exc_info.value) + with pytest.raises(OSError, match="Storage error"): + mock_storage.load(upload_file.key, stream=True) def test_validate_file_ownership_unexpected_error_logging( - self, file_preview_api: FilePreviewApi, caplog: pytest.LogCaptureFixture + self, + file_preview_api: FilePreviewApi, + database: _Database, + sqlite_engine: Engine, + caplog: pytest.LogCaptureFixture, ): - """Test that unexpected errors are logged properly""" - file_id = str(uuid.uuid4()) - app_id = str(uuid.uuid4()) + file_id = str(uuid4()) + app_id = str(uuid4()) - with patch("controllers.service_api.app.file_preview.db") as mock_db: - # Mock database scalar to raise unexpected exception - mock_db.session.scalar.side_effect = Exception("Unexpected database error") + def fail_statement(*_args: object) -> None: + raise RuntimeError("Unexpected database error") - # Execute and assert exception - with caplog.at_level(logging.ERROR, logger="controllers.service_api.app.file_preview"): - with pytest.raises(FileAccessDeniedError) as exc_info: - file_preview_api._validate_file_ownership(file_id, app_id) + event.listen(sqlite_engine, "before_cursor_execute", fail_statement) + try: + with patch("controllers.service_api.app.file_preview.db", database): + with caplog.at_level(logging.ERROR, logger="controllers.service_api.app.file_preview"): + with pytest.raises(FileAccessDeniedError, match="File access validation failed"): + file_preview_api._validate_file_ownership(file_id, app_id) + finally: + event.remove(sqlite_engine, "before_cursor_execute", fail_statement) - # Verify error message - assert "File access validation failed" in str(exc_info.value) - - # Verify logging was called with the structured context fields. The ``extra`` keys - # are attached to the LogRecord as attributes, so they are not in ``caplog.text``. - assert len(caplog.records) == 1 - log_record = caplog.records[0] - assert log_record.getMessage() == "Unexpected error during file ownership validation" - record = cast(_FilePreviewLogRecord, log_record) - assert record.file_id == file_id - assert record.app_id == app_id - assert record.error == "Unexpected database error" + assert len(caplog.records) == 1 + log_record = caplog.records[0] + assert log_record.getMessage() == "Unexpected error during file ownership validation" + record = cast(_FilePreviewLogRecord, log_record) + assert record.file_id == file_id + assert record.app_id == app_id + assert record.error == "Unexpected database error" diff --git a/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py b/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py index c5072ec2e70..57c5c038424 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_hitl_service_api.py @@ -36,11 +36,12 @@ from core.workflow.nodes.human_input.entities import ParagraphInputConfig, UserA 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 enums import DeploymentEdition from graphon.entities import WorkflowStartReason from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus from graphon.runtime import GraphRuntimeState, VariablePool from models.account import Account -from models.enums import CreatorUserRole +from models.enums import CreatorUserRole, MessageStatus from models.human_input import HumanInputForm from models.model import AppMode from models.workflow import WorkflowRun @@ -117,10 +118,8 @@ def _build_service_api_pause_converter() -> WorkflowResponseConverter: workflow_id="workflow-id", workflow_execution_id="run-id", ) - user = MagicMock(spec=Account) + user = Account(name="Tester", email="tester@example.com") user.id = "account-id" - user.name = "Tester" - user.email = "tester@example.com" return WorkflowResponseConverter( application_generate_entity=application_generate_entity, user=user, @@ -271,8 +270,7 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext: workflow_execution_id="run-1", ) runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0) - runtime_state.register_paused_node("node-1") - runtime_state.outputs = {"result": "value"} + runtime_state.set_output("result", "value") wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity) return WorkflowResumptionContext( generate_entity=wrapper, @@ -381,7 +379,7 @@ class TestHitlServiceApi: monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, ) -> None: - monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(ags_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) monkeypatch.setattr(ags_module, "RateLimit", _DummyRateLimit) workflow = MagicMock() @@ -456,7 +454,6 @@ class TestHitlServiceApi: def test_advanced_chat_blocking_pipeline_pause_payload_contract(self) -> None: from core.app.app_config.entities import AppAdditionalFeatures from core.app.apps.advanced_chat.generate_task_pipeline import AdvancedChatAppGenerateTaskPipeline - from models.enums import MessageStatus from models.model import EndUser app_config = WorkflowUIBasedAppConfig( @@ -618,7 +615,6 @@ class TestHitlServiceApi: assert response.data.paused_nodes == ["node-1"] assert response.data.reasons == [{"TYPE": "human_input_required", "form_id": "form-1", "expiration_time": 1}] - @pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) def test_service_api_pause_event_serializes_hitl_reason( self, monkeypatch: pytest.MonkeyPatch, @@ -633,7 +629,7 @@ class TestHitlServiceApi: reason=WorkflowStartReason.INITIAL, ) - expiration_time = datetime(2024, 1, 1) + expiration_time = datetime(2024, 1, 1, tzinfo=UTC) _persist_human_input_form(sqlite_session, expiration_time=expiration_time) monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=sqlite_engine)) @@ -696,7 +692,6 @@ class TestHitlServiceApi: assert hi_resp.data.expiration_time == int(expiration_time.timestamp()) # Snapshot payload contract - @pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) def test_snapshot_events_include_pause_payload_contract( self, monkeypatch: pytest.MonkeyPatch, @@ -706,7 +701,7 @@ class TestHitlServiceApi: 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) + expiration_time = datetime(2024, 1, 1, tzinfo=UTC) _persist_human_input_form(sqlite_session, expiration_time=expiration_time) monkeypatch.setattr( "services.workflow_event_snapshot_service.load_form_dispositions_by_form_id", diff --git a/api/tests/unit_tests/controllers/service_api/app/test_message.py b/api/tests/unit_tests/controllers/service_api/app/test_message.py index 40186036ed7..2920e6565ad 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_message.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_message.py @@ -15,12 +15,15 @@ Focus on: """ import uuid +from collections.abc import Iterator from inspect import unwrap from types import SimpleNamespace from unittest.mock import Mock, patch import pytest from flask import Flask +from sqlalchemy import Engine +from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest, InternalServerError, NotFound from controllers.service_api.app.error import NotChatAppError @@ -44,6 +47,14 @@ from services.errors.message import ( from services.message_service import MessageService +@pytest.fixture +def orm_session(sqlite_engine: Engine) -> Iterator[Session]: + """Provide a real caller-owned session for MessageService interface tests.""" + + with Session(sqlite_engine, expire_on_commit=False) as session: + yield session + + class TestMessageListQuery: """Test suite for MessageListQuery Pydantic model.""" @@ -253,7 +264,7 @@ class TestMessageService: assert callable(MessageService.get_suggested_questions_after_answer) @patch.object(MessageService, "pagination_by_first_id") - def test_pagination_by_first_id_returns_pagination_result(self, mock_pagination): + def test_pagination_by_first_id_returns_pagination_result(self, mock_pagination, orm_session: Session): """Test pagination_by_first_id returns expected format.""" mock_result = Mock() mock_result.data = [] @@ -262,12 +273,12 @@ class TestMessageService: mock_pagination.return_value = mock_result result = MessageService.pagination_by_first_id( - app_model=Mock(spec=App), - user=Mock(spec=EndUser), + app_model=App(), + user=EndUser(), conversation_id=str(uuid.uuid4()), first_id=None, limit=20, - session=Mock(), + session=orm_session, ) assert hasattr(result, "data") @@ -275,7 +286,7 @@ class TestMessageService: assert hasattr(result, "has_more") @patch.object(MessageService, "pagination_by_first_id") - def test_pagination_raises_conversation_not_exists_error(self, mock_pagination): + def test_pagination_raises_conversation_not_exists_error(self, mock_pagination, orm_session: Session): """Test pagination raises ConversationNotExistsError.""" import services.errors.conversation @@ -283,62 +294,62 @@ class TestMessageService: with pytest.raises(services.errors.conversation.ConversationNotExistsError): MessageService.pagination_by_first_id( - app_model=Mock(spec=App), - user=Mock(spec=EndUser), + app_model=App(), + user=EndUser(), conversation_id="invalid_id", first_id=None, limit=20, - session=Mock(), + session=orm_session, ) @patch.object(MessageService, "pagination_by_first_id") - def test_pagination_raises_first_message_not_exists_error(self, mock_pagination): + def test_pagination_raises_first_message_not_exists_error(self, mock_pagination, orm_session: Session): """Test pagination raises FirstMessageNotExistsError.""" mock_pagination.side_effect = FirstMessageNotExistsError() with pytest.raises(FirstMessageNotExistsError): MessageService.pagination_by_first_id( - app_model=Mock(spec=App), - user=Mock(spec=EndUser), + app_model=App(), + user=EndUser(), conversation_id=str(uuid.uuid4()), first_id="invalid_first_id", limit=20, - session=Mock(), + session=orm_session, ) @patch.object(MessageService, "create_feedback") - def test_create_feedback_with_rating_and_content(self, mock_create_feedback): + def test_create_feedback_with_rating_and_content(self, mock_create_feedback, orm_session: Session): """Test create_feedback with rating and content.""" mock_create_feedback.return_value = None MessageService.create_feedback( - app_model=Mock(spec=App), + app_model=App(), message_id=str(uuid.uuid4()), - user=Mock(spec=EndUser), + user=EndUser(), rating=FeedbackRating.LIKE, content="Great response!", - session=Mock(), + session=orm_session, ) mock_create_feedback.assert_called_once() @patch.object(MessageService, "create_feedback") - def test_create_feedback_raises_message_not_exists_error(self, mock_create_feedback): + def test_create_feedback_raises_message_not_exists_error(self, mock_create_feedback, orm_session: Session): """Test create_feedback raises MessageNotExistsError.""" mock_create_feedback.side_effect = MessageNotExistsError() with pytest.raises(MessageNotExistsError): MessageService.create_feedback( - app_model=Mock(spec=App), + app_model=App(), message_id="invalid_message_id", - user=Mock(spec=EndUser), + user=EndUser(), rating=FeedbackRating.LIKE, content=None, - session=Mock(), + session=orm_session, ) @patch.object(MessageService, "get_all_messages_feedbacks") - def test_get_all_messages_feedbacks_returns_list(self, mock_get_feedbacks): + def test_get_all_messages_feedbacks_returns_list(self, mock_get_feedbacks, orm_session: Session): """Test get_all_messages_feedbacks returns list of feedbacks.""" mock_feedbacks = [ {"message_id": str(uuid.uuid4()), "rating": "like"}, @@ -346,54 +357,54 @@ class TestMessageService: ] mock_get_feedbacks.return_value = mock_feedbacks - result = MessageService.get_all_messages_feedbacks(app_model=Mock(spec=App), page=1, limit=20, session=Mock()) + result = MessageService.get_all_messages_feedbacks(app_model=App(), page=1, limit=20, session=orm_session) assert len(result) == 2 assert result[0]["rating"] == "like" @patch.object(MessageService, "get_suggested_questions_after_answer") - def test_get_suggested_questions_returns_questions_list(self, mock_get_questions): + def test_get_suggested_questions_returns_questions_list(self, mock_get_questions, orm_session: Session): """Test get_suggested_questions_after_answer returns list of questions.""" mock_questions = ["What about this aspect?", "Can you elaborate on that?", "How does this relate to...?"] mock_get_questions.return_value = mock_questions result = MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), - user=Mock(spec=EndUser), + app_model=App(), + user=EndUser(), message_id=str(uuid.uuid4()), invoke_from=Mock(), - session=Mock(), + session=orm_session, ) assert len(result) == 3 assert isinstance(result[0], str) @patch.object(MessageService, "get_suggested_questions_after_answer") - def test_get_suggested_questions_raises_disabled_error(self, mock_get_questions): + def test_get_suggested_questions_raises_disabled_error(self, mock_get_questions, orm_session: Session): """Test get_suggested_questions_after_answer raises SuggestedQuestionsAfterAnswerDisabledError.""" mock_get_questions.side_effect = SuggestedQuestionsAfterAnswerDisabledError() with pytest.raises(SuggestedQuestionsAfterAnswerDisabledError): MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), - user=Mock(spec=EndUser), + app_model=App(), + user=EndUser(), message_id=str(uuid.uuid4()), invoke_from=Mock(), - session=Mock(), + session=orm_session, ) @patch.object(MessageService, "get_suggested_questions_after_answer") - def test_get_suggested_questions_raises_message_not_exists_error(self, mock_get_questions): + def test_get_suggested_questions_raises_message_not_exists_error(self, mock_get_questions, orm_session: Session): """Test get_suggested_questions_after_answer raises MessageNotExistsError.""" mock_get_questions.side_effect = MessageNotExistsError() with pytest.raises(MessageNotExistsError): MessageService.get_suggested_questions_after_answer( - app_model=Mock(spec=App), - user=Mock(spec=EndUser), + app_model=App(), + user=EndUser(), message_id="invalid_message_id", invoke_from=Mock(), - session=Mock(), + session=orm_session, ) diff --git a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py index 7975a935f93..c1cf3539477 100644 --- a/api/tests/unit_tests/controllers/service_api/app/test_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/app/test_workflow.py @@ -42,7 +42,7 @@ from controllers.service_api.app.workflow import ( ) from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError from core.app.entities.app_invoke_entities import InvokeFrom -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from graphon.enums import WorkflowExecutionStatus from models import Account from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom @@ -582,10 +582,10 @@ class TestWorkflowRunApi: 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 + 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", True) + monkeypatch.setattr(workflow_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) billing_get_info = Mock(return_value={"enabled": True, "subscription": {"plan": CloudPlan.SANDBOX}}) generate = Mock(return_value={"result": "ok"}) @@ -598,7 +598,7 @@ class TestWorkflowRunApi: with app.test_request_context("/workflows/run", method="POST", json={"inputs": {}}): response = handler( api, - session=Mock(), + session=sqlite_session, app_model=_make_app_model(), end_user=_make_end_user(), ) @@ -609,9 +609,11 @@ class TestWorkflowRunApi: class TestWorkflowRunByIdApi: - def test_rejects_sandbox_plan_with_upgrade_error(self, app: Flask, monkeypatch: pytest.MonkeyPatch) -> None: + def test_rejects_sandbox_plan_with_upgrade_error( + 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", True) + monkeypatch.setattr(workflow_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) billing_get_info = Mock(return_value={"enabled": True, "subscription": {"plan": CloudPlan.SANDBOX}}) generate = Mock() @@ -626,7 +628,7 @@ class TestWorkflowRunByIdApi: with pytest.raises(WorkflowVersionExecutionNotAllowedError) as exc_info: handler( api, - session=Mock(), + session=sqlite_session, app_model=app_model, end_user=_make_end_user(), workflow_id="w1", @@ -646,23 +648,25 @@ class TestWorkflowRunByIdApi: } @pytest.mark.parametrize( - ("billing_config_enabled", "billing_enabled", "plan"), + ("deployment_edition", "billing_enabled", "plan"), [ - (False, True, CloudPlan.SANDBOX), - (True, False, CloudPlan.SANDBOX), - (True, True, CloudPlan.PROFESSIONAL), + (DeploymentEdition.COMMUNITY, True, CloudPlan.SANDBOX), + (DeploymentEdition.ENTERPRISE, True, CloudPlan.SANDBOX), + (DeploymentEdition.CLOUD, False, CloudPlan.SANDBOX), + (DeploymentEdition.CLOUD, True, CloudPlan.PROFESSIONAL), ], ) def test_allows_execution_outside_enabled_sandbox_plan( self, app: Flask, monkeypatch: pytest.MonkeyPatch, - billing_config_enabled: bool, + deployment_edition: DeploymentEdition, billing_enabled: bool, plan: CloudPlan, + sqlite_session: Session, ) -> None: workflow_module = sys.modules["controllers.service_api.app.workflow"] - monkeypatch.setattr(workflow_module.dify_config, "BILLING_ENABLED", billing_config_enabled) + monkeypatch.setattr(workflow_module.dify_config, "DEPLOYMENT_EDITION", deployment_edition) billing_get_info = Mock(return_value={"enabled": billing_enabled, "subscription": {"plan": plan}}) generate = Mock(return_value={"result": "ok"}) @@ -676,7 +680,7 @@ class TestWorkflowRunByIdApi: with app.test_request_context("/workflows/w1/run", method="POST", json={"inputs": {}}): response = handler( api, - session=Mock(), + session=sqlite_session, app_model=app_model, end_user=_make_end_user(), workflow_id="w1", @@ -684,7 +688,7 @@ class TestWorkflowRunByIdApi: assert response.get_json() == {"result": "ok"} generate.assert_called_once() - if billing_config_enabled: + if deployment_edition == DeploymentEdition.CLOUD: billing_get_info.assert_called_once_with(app_model.tenant_id, exclude_vector_space=True) else: billing_get_info.assert_not_called() @@ -692,7 +696,7 @@ class TestWorkflowRunByIdApi: @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(workflow_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) monkeypatch.setattr( AppGenerateService, "generate", @@ -711,7 +715,7 @@ class TestWorkflowRunByIdApi: @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(workflow_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) monkeypatch.setattr( AppGenerateService, "generate", diff --git a/api/tests/unit_tests/controllers/service_api/conftest.py b/api/tests/unit_tests/controllers/service_api/conftest.py index bede4d75850..126ac8e8763 100644 --- a/api/tests/unit_tests/controllers/service_api/conftest.py +++ b/api/tests/unit_tests/controllers/service_api/conftest.py @@ -83,25 +83,27 @@ def mock_app_id(): @pytest.fixture def mock_end_user(mock_tenant_id): """Create a mock EndUser model with required attributes.""" - user = Mock(spec=EndUser) - user.id = str(uuid.uuid4()) - user.external_user_id = f"external_{uuid.uuid4().hex[:8]}" - user.tenant_id = mock_tenant_id + user = EndUser( + id=str(uuid.uuid4()), + external_user_id=f"external_{uuid.uuid4().hex[:8]}", + tenant_id=mock_tenant_id, + ) return user @pytest.fixture def mock_app_model(mock_app_id, mock_tenant_id): - """Create a mock App model with all required attributes for API testing.""" - app = Mock(spec=App) - app.id = mock_app_id - app.tenant_id = mock_tenant_id - app.name = "Test App" - app.description = "A test application" - app.mode = AppMode.CHAT + """Create an App model with all required attributes for API testing.""" + app = App( + id=mock_app_id, + tenant_id=mock_tenant_id, + name="Test App", + description="A test application", + mode=AppMode.CHAT, + status="normal", + enable_api=True, + ) app.author_name = "Test Author" - app.status = "normal" - app.enable_api = True app.tags = [] # Mock workflow for workflow apps @@ -113,7 +115,7 @@ def mock_app_model(mock_app_id, mock_tenant_id): @pytest.fixture def mock_tenant(mock_tenant_id): - """Create a mock Tenant model.""" + """Create a Tenant model.""" tenant = Mock() tenant.id = mock_tenant_id tenant.status = TenantStatus.NORMAL @@ -122,7 +124,7 @@ def mock_tenant(mock_tenant_id): @pytest.fixture def mock_account(): - """Create a mock Account model.""" + """Create an Account model.""" account = Mock() account.id = str(uuid.uuid4()) return account @@ -151,50 +153,55 @@ def mock_dataset_api_token(mock_tenant_id): @pytest.fixture def mock_dataset(): - """Create a mock Dataset model.""" + """Create a Dataset model.""" from models.dataset import Dataset - dataset = Mock(spec=Dataset) - dataset.id = str(uuid.uuid4()) - dataset.tenant_id = str(uuid.uuid4()) - dataset.name = "Test Dataset" - dataset.indexing_technique = "economy" - dataset.embedding_model = None - dataset.embedding_model_provider = None + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + name="Test Dataset", + indexing_technique="economy", + embedding_model=None, + embedding_model_provider=None, + ) return dataset @pytest.fixture def mock_document(): - """Create a mock Document model.""" + """Create a Document model.""" from models.dataset import Document - document = Mock(spec=Document) - document.id = str(uuid.uuid4()) - document.dataset_id = str(uuid.uuid4()) - document.tenant_id = str(uuid.uuid4()) - document.name = "test_document.txt" - document.indexing_status = "completed" - document.enabled = True - document.doc_form = IndexStructureType.PARAGRAPH_INDEX + document = Document( + id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + name="test_document.txt", + indexing_status="completed", + enabled=True, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) return document @pytest.fixture def mock_segment(): - """Create a mock DocumentSegment model.""" + """Create a DocumentSegment model.""" from models.dataset import DocumentSegment - segment = Mock(spec=DocumentSegment) + segment = DocumentSegment( + tenant_id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + document_id=str(uuid.uuid4()), + position=1, + content="Test segment content", + word_count=3, + tokens=0, + created_by="account-id", + enabled=True, + status="completed", + ) segment.id = str(uuid.uuid4()) - segment.document_id = str(uuid.uuid4()) - segment.dataset_id = str(uuid.uuid4()) - segment.tenant_id = str(uuid.uuid4()) - segment.content = "Test segment content" - segment.word_count = 3 - segment.position = 1 - segment.enabled = True - segment.status = "completed" return segment @@ -203,9 +210,15 @@ def mock_child_chunk(): """Create a mock ChildChunk model.""" from models.dataset import ChildChunk - child_chunk = Mock(spec=ChildChunk) + child_chunk = ChildChunk( + tenant_id=str(uuid.uuid4()), + dataset_id="dataset-id", + document_id="document-id", + segment_id=str(uuid.uuid4()), + position=1, + content="Test child chunk content", + word_count=0, + created_by="account-id", + ) child_chunk.id = str(uuid.uuid4()) - child_chunk.segment_id = str(uuid.uuid4()) - child_chunk.tenant_id = str(uuid.uuid4()) - child_chunk.content = "Test child chunk content" return child_chunk diff --git a/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py b/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py index 406037e268d..d7c524e6532 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py @@ -24,10 +24,18 @@ from unittest.mock import Mock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.datastructures import FileStorage from werkzeug.exceptions import Forbidden, NotFound -from controllers.common.errors import FilenameNotExistsError, NoFileUploadedError, TooManyFilesError +from controllers.common.errors import ( + FilenameNotExistsError, + NoFileUploadedError, + TooManyFilesError, +) +from controllers.common.errors import ( + FileTooLargeError as FileTooLargeHTTPError, +) from controllers.service_api.dataset.error import PipelineRunError from controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow import ( DatasourceNodeRunApi, @@ -38,7 +46,9 @@ from controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow import ( ) from core.app.entities.app_invoke_entities import InvokeFrom from models.account import Account -from services.errors.file import FileTooLargeError, UnsupportedFileTypeError +from models.dataset import Dataset +from services.errors.file import FileTooLargeError as FileTooLargeServiceError +from services.errors.file import UnsupportedFileTypeError from services.rag_pipeline.entity.pipeline_service_api_entities import ( DatasourceNodeRunApiEntity, PipelineRunApiEntity, @@ -46,6 +56,20 @@ from services.rag_pipeline.entity.pipeline_service_api_entities import ( from services.rag_pipeline.rag_pipeline import RagPipelineService +def _persist_dataset(session: Session, *, tenant_id: str, dataset_id: str) -> Dataset: + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Pipeline dataset", + created_by="account-1", + data_source_type=None, + indexing_technique=None, + ) + session.add(dataset) + session.commit() + return dataset + + class TestDatasourceNodeRunPayload: """Test suite for DatasourceNodeRunPayload Pydantic model.""" @@ -127,7 +151,7 @@ class TestFileUploadErrors: def test_file_too_large_error(self): """Test FileTooLargeError can be raised.""" - error = FileTooLargeError("File exceeds size limit") + error = FileTooLargeServiceError("File exceeds size limit") assert error is not None def test_unsupported_file_type_error(self): @@ -470,7 +494,7 @@ class TestDatasourceNodeRunApiPost: @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.PipelineGenerator") @patch( "controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", - new_callable=lambda: Mock(spec=Account), + new_callable=lambda: Account(name="Test Account", email="test@example.com"), ) @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db") @@ -546,17 +570,18 @@ class TestPipelineRunApiPost: @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.PipelineGenerateService") @patch( "controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", - new_callable=lambda: Mock(spec=Account), + new_callable=lambda: Account(name="Test Account", email="test@example.com"), ) @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.RagPipelineService") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns") - def test_post_success_streaming(self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app): + def test_post_success_streaming( + self, mock_ns, mock_svc_cls, mock_current_user, mock_gen_svc, mock_helper, app, sqlite_session: Session + ): """Test successful pipeline run with streaming response.""" tenant_id = str(uuid.uuid4()) dataset_id = str(uuid.uuid4()) - session = Mock() - session.scalar.return_value = Mock() + _persist_dataset(sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id) mock_ns.payload = { "inputs": {"key": "val"}, @@ -577,33 +602,31 @@ class TestPipelineRunApiPost: with app.test_request_context("/datasets/test/pipeline/run", method="POST"): api = PipelineRunApi() - response = api.post.__wrapped__(api, session, tenant_id=tenant_id, dataset_id=dataset_id) + response = api.post.__wrapped__(api, sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id) assert response == {"result": "ok"} - mock_svc_cls.assert_called_once_with(session) + mock_svc_cls.assert_called_once_with(sqlite_session) mock_gen_svc.generate.assert_called_once() - def test_post_not_found(self, app: Flask): + def test_post_not_found(self, app: Flask, sqlite_session: Session): """Test NotFound when dataset check fails.""" - session = Mock() - session.scalar.return_value = None - with app.test_request_context("/datasets/test/pipeline/run", method="POST"): api = PipelineRunApi() with pytest.raises(NotFound): api.post.__wrapped__( api, - session, + sqlite_session, tenant_id=str(uuid.uuid4()), dataset_id=str(uuid.uuid4()), ) @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user", new="not_account") @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.service_api_ns") - def test_post_forbidden_non_account_user(self, mock_ns, app: Flask): + def test_post_forbidden_non_account_user(self, mock_ns, app: Flask, sqlite_session: Session): """Test Forbidden when current_user is not an Account.""" - session = Mock() - session.scalar.return_value = Mock() + tenant_id = str(uuid.uuid4()) + dataset_id = str(uuid.uuid4()) + _persist_dataset(sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id) mock_ns.payload = { "inputs": {}, "datasource_type": "online_document", @@ -618,9 +641,9 @@ class TestPipelineRunApiPost: with pytest.raises(Forbidden): api.post.__wrapped__( api, - session, - tenant_id=str(uuid.uuid4()), - dataset_id=str(uuid.uuid4()), + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, ) @@ -666,6 +689,38 @@ class TestFileUploadApiPost: assert response["name"] == "doc.pdf" assert response["extension"] == "pdf" + @patch( + "controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.FeatureService" + ".get_knowledge_file_size_limit", + return_value=15, + ) + @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.FileService") + @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.current_user") + @patch("controllers.service_api.dataset.rag_pipeline.rag_pipeline_workflow.db") + def test_upload_file_too_large_returns_http_413( + self, mock_db, mock_current_user, mock_file_svc_cls, mock_get_limit, app: Flask + ): + mock_current_user.__bool__ = Mock(return_value=True) + mock_file_svc_cls.return_value.upload_file.side_effect = FileTooLargeServiceError() + file_data = FileStorage( + stream=io.BytesIO(b"oversized content"), + filename="doc.pdf", + content_type="application/pdf", + ) + + with app.test_request_context( + "/datasets/pipeline/file-upload", + method="POST", + content_type="multipart/form-data", + data={"file": file_data}, + ): + with pytest.raises(FileTooLargeHTTPError) as exc_info: + KnowledgebasePipelineFileUploadApi().post(tenant_id="tenant-1") + + assert exc_info.value.code == 413 + assert exc_info.value.error_code == "file_too_large" + mock_get_limit.assert_called_once_with("tenant-1") + def test_upload_no_file(self, app: Flask): """Test error when no file is uploaded.""" with app.test_request_context( diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_apis.py b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_apis.py new file mode 100644 index 00000000000..8d306a85dcf --- /dev/null +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_apis.py @@ -0,0 +1,738 @@ +"""Unit tests for Service API dataset controller behavior. + +Service boundaries stay mocked, while ORM collaborators are concrete model instances +persisted in one in-memory SQLite session. The controller's ``db.session`` and the +session passed to unwrapped ``@with_session`` endpoints both use that same session, +so model properties and service call contracts exercise real SQLAlchemy behavior. +""" + +import uuid +from datetime import UTC, datetime +from inspect import unwrap +from typing import cast +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask +from sqlalchemy.orm import Session, scoped_session, sessionmaker +from werkzeug.exceptions import Forbidden, NotFound + +import services +from controllers.service_api.dataset.error import DatasetInUseError, DatasetNameDuplicateError, InvalidActionError +from extensions.ext_database import db +from models.account import Account, Tenant, TenantAccountRole +from models.dataset import AppDatasetJoin, Dataset, DatasetMetadata, Document +from models.enums import PermissionEnum +from models.model import App, Tag, TagBinding + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- +DATASET_MODEL_TABLES = ( + Account, + Tenant, + Dataset, + Document, + App, + AppDatasetJoin, + DatasetMetadata, + Tag, + TagBinding, +) +pytestmark = pytest.mark.parametrize("sqlite_session", [DATASET_MODEL_TABLES], indirect=True) + + +@pytest.fixture(autouse=True) +def controller_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Route controller and model database access through the test's SQLite session.""" + + # Flask-SQLAlchemy exposes a callable registry that also proxies Session methods. + # Seed that registry with this fixture's Session so both access styles share one transaction. + existing_session_factory = cast(sessionmaker[Session], lambda: sqlite_session) + session_registry = scoped_session(existing_session_factory) + monkeypatch.setattr(db, "session", session_registry) + return sqlite_session + + +@pytest.fixture +def tenant(controller_session: Session) -> Tenant: + tenant = Tenant(name="Dataset API Tenant") + controller_session.add(tenant) + controller_session.flush() + return tenant + + +@pytest.fixture +def account(controller_session: Session, tenant: Tenant, monkeypatch: pytest.MonkeyPatch) -> Account: + account = Account(name="Dataset API User", email=f"dataset-api-{uuid.uuid4()}@example.com") + account.role = TenantAccountRole.OWNER + account._current_tenant = tenant + controller_session.add(account) + controller_session.flush() + + # Inject the concrete account at the controller boundary without relying on Flask-Login globals. + from controllers.service_api.dataset import dataset as dataset_module + + monkeypatch.setattr(dataset_module, "current_user", account) + return account + + +def make_dataset( + session: Session, + tenant: Tenant, + account: Account, + **overrides: object, +) -> Dataset: + """Create and flush a real dataset so its database-backed properties can be serialized.""" + + base: dict[str, object] = { + "id": str(uuid.uuid4()), + "tenant_id": tenant.id, + "name": "Dataset", + "description": "desc", + "provider": "vendor", + "permission": PermissionEnum.ONLY_ME, + "data_source_type": None, + "indexing_technique": "economy", + "created_by": account.id, + "created_at": datetime(2024, 1, 1, 12, 0, 0, tzinfo=UTC), + "updated_by": None, + "updated_at": datetime(2024, 1, 1, 12, 0, 0, tzinfo=UTC), + "embedding_model": None, + "embedding_model_provider": None, + "retrieval_model": None, + "summary_index_setting": None, + "built_in_field_enabled": False, + "pipeline_id": None, + "runtime_mode": "general", + "chunk_structure": None, + "icon_info": None, + "enable_api": False, + "is_multimodal": False, + } + base.update(overrides) + dataset = Dataset(**base) + session.add(dataset) + session.flush() + return dataset + + +@pytest.fixture +def dataset(controller_session: Session, tenant: Tenant, account: Account) -> Dataset: + return make_dataset(controller_session, tenant, account) + + +DATASET_DETAIL_KEYS = { + "id", + "name", + "description", + "provider", + "permission", + "data_source_type", + "indexing_technique", + "app_count", + "document_count", + "word_count", + "created_by", + "author_name", + "created_at", + "updated_by", + "updated_at", + "embedding_model", + "embedding_model_provider", + "embedding_available", + "retrieval_model_dict", + "summary_index_setting", + "tags", + "doc_form", + "external_knowledge_info", + "external_retrieval_model", + "doc_metadata", + "built_in_field_enabled", + "pipeline_id", + "runtime_mode", + "chunk_structure", + "icon_info", + "is_published", + "total_documents", + "total_available_documents", + "enable_api", + "is_multimodal", + "maintainer", +} + + +def assert_dataset_detail_shape(response: dict[str, object], *, with_partial_members: bool = False) -> None: + expected_keys = set(DATASET_DETAIL_KEYS) + if with_partial_members: + expected_keys.add("partial_member_list") + assert set(response) == expected_keys + assert isinstance(response["created_at"], int) + assert isinstance(response["updated_at"], int) + retrieval_model = response["retrieval_model_dict"] + assert isinstance(retrieval_model, dict) + assert set(retrieval_model) == { + "search_method", + "reranking_enable", + "reranking_mode", + "reranking_model", + "weights", + "top_k", + "score_threshold_enabled", + "score_threshold", + } + external_retrieval_model = response["external_retrieval_model"] + if external_retrieval_model is not None: + assert isinstance(external_retrieval_model, dict) + assert set(external_retrieval_model) == { + "top_k", + "score_threshold", + "score_threshold_enabled", + } + if not with_partial_members: + assert "partial_member_list" not in response + + +# --------------------------------------------------------------------------- +# API endpoint tests — DatasetListApi +# --------------------------------------------------------------------------- + + +class TestDatasetListApiGet: + """Test suite for DatasetListApi.get() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_list_datasets_success( + self, + mock_dataset_svc: MagicMock, + mock_provider_mgr: MagicMock, + app: Flask, + account: Account, + tenant: Tenant, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetListApi + + mock_dataset_svc.get_datasets.return_value = ([make_dataset(controller_session, tenant, account)], 1) + mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]() + + with app.test_request_context("/datasets?page=1&limit=20", method="GET"): + api = DatasetListApi() + response, status = unwrap(api.get)(api, controller_session, tenant_id=tenant.id) + + assert status == 200 + assert set(response) == {"data", "has_more", "limit", "total", "page"} + assert response["has_more"] is False + assert response["limit"] == 20 + assert response["total"] == 1 + assert response["page"] == 1 + assert len(response["data"]) == 1 + assert_dataset_detail_shape(response["data"][0]) + + @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_list_datasets_preserves_repeated_tag_ids( + self, + mock_dataset_svc: MagicMock, + mock_provider_mgr: MagicMock, + app: Flask, + account: Account, + tenant: Tenant, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetListApi + + mock_dataset_svc.get_datasets.return_value = ([make_dataset(controller_session, tenant, account)], 1) + mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]() + + with app.test_request_context("/datasets?tag_ids=tag-a&tag_ids=tag-b", method="GET"): + api = DatasetListApi() + response, status = unwrap(api.get)(api, controller_session, tenant_id=tenant.id) + page, limit, session, tenant_id, user, keyword, tag_ids, include_all = ( + mock_dataset_svc.get_datasets.call_args.args + ) + assert user is account + + assert status == 200 + assert response["total"] == 1 + assert (page, limit, session, tenant_id, keyword, tag_ids, include_all) == ( + 1, + 20, + controller_session, + tenant.id, + None, + ["tag-a", "tag-b"], + False, + ) + + +class TestDatasetListApiPost: + """Test suite for DatasetListApi.post() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_create_dataset_success( + self, + mock_dataset_svc: MagicMock, + app: Flask, + account: Account, + tenant: Tenant, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetListApi + + mock_dataset_svc.create_empty_dataset.return_value = make_dataset( + controller_session, tenant, account, name="New Dataset" + ) + + with app.test_request_context( + "/datasets", + method="POST", + json={"name": "New Dataset"}, + ): + api = DatasetListApi() + response, status = unwrap(api.post)(api, controller_session, tenant_id=tenant.id) + + assert status == 200 + assert_dataset_detail_shape(response) + assert response["name"] == "New Dataset" + mock_dataset_svc.create_empty_dataset.assert_called_once() + + @pytest.mark.usefixtures("account") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_create_dataset_duplicate_name( + self, + mock_dataset_svc: MagicMock, + app: Flask, + tenant: Tenant, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetListApi + + mock_dataset_svc.create_empty_dataset.side_effect = services.errors.dataset.DatasetNameDuplicateError() + + with app.test_request_context( + "/datasets", + method="POST", + json={"name": "Existing Dataset"}, + ): + api = DatasetListApi() + with pytest.raises(DatasetNameDuplicateError): + unwrap(api.post)(api, controller_session, tenant_id=tenant.id) + + +# --------------------------------------------------------------------------- +# API endpoint tests — DatasetApi +# --------------------------------------------------------------------------- + + +class TestDatasetApiGet: + """Test suite for DatasetApi.get() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_get_dataset_success( + self, + mock_dataset_svc: MagicMock, + mock_provider_mgr: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.check_dataset_permission.return_value = None + mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]() + + with app.test_request_context( + f"/datasets/{dataset.id}", + method="GET", + ): + api = DatasetApi() + response, status = unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id) + + assert status == 200 + assert_dataset_detail_shape(response) + assert response["embedding_available"] is True + assert response["retrieval_model_dict"]["search_method"] == "keyword_search" + + @patch("controllers.service_api.dataset.dataset.DatasetPermissionService") + @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_get_dataset_partial_members_shape( + self, + mock_dataset_svc: MagicMock, + mock_provider_mgr: MagicMock, + mock_perm_svc: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + dataset.permission = PermissionEnum.PARTIAL_TEAM + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.check_dataset_permission.return_value = None + mock_perm_svc.get_dataset_partial_member_list.return_value = ["user-1", "user-2"] + mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]() + + with app.test_request_context( + f"/datasets/{dataset.id}", + method="GET", + ): + api = DatasetApi() + response, status = unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id) + + assert status == 200 + assert_dataset_detail_shape(response, with_partial_members=True) + assert response["partial_member_list"] == ["user-1", "user-2"] + + @patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_get_dataset_uses_default_external_retrieval_model( + self, + mock_dataset_svc: MagicMock, + mock_provider_mgr: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + dataset.retrieval_model = None + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.check_dataset_permission.return_value = None + mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]() + + with app.test_request_context(f"/datasets/{dataset.id}", method="GET"): + api = DatasetApi() + response, status = unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id) + + assert status == 200 + assert_dataset_detail_shape(response) + assert response["external_retrieval_model"] == { + "top_k": 2, + "score_threshold": 0.0, + "score_threshold_enabled": None, + } + + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_get_dataset_not_found( + self, + mock_dataset_svc: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + mock_dataset_svc.get_dataset.return_value = None + + with app.test_request_context( + f"/datasets/{dataset.id}", + method="GET", + ): + api = DatasetApi() + with pytest.raises(NotFound): + unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id) + + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_get_dataset_no_permission( + self, + mock_dataset_svc: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.check_dataset_permission.side_effect = services.errors.account.NoPermissionError() + + with app.test_request_context( + f"/datasets/{dataset.id}", + method="GET", + ): + api = DatasetApi() + with pytest.raises(Forbidden): + unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id) + + +class TestDatasetApiPatch: + """Test suite for DatasetApi.patch() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.DatasetPermissionService") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_patch_dataset_success_shape( + self, + mock_dataset_svc: MagicMock, + mock_perm_svc: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + dataset.name = "Updated Dataset" + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.update_dataset.return_value = dataset + mock_perm_svc.check_permission.return_value = None + mock_perm_svc.get_dataset_partial_member_list.return_value = ["user-1"] + + payload = { + "name": "Updated Dataset", + "permission": "partial_members", + "partial_member_list": [{"user_id": "user-1", "role": "editor"}], + } + with app.test_request_context( + f"/datasets/{dataset.id}", + method="PATCH", + json=payload, + ): + api = DatasetApi() + response, status = unwrap(api.patch)( + api, + controller_session, + _=dataset.tenant_id, + dataset_id=dataset.id, + ) + + assert status == 200 + assert_dataset_detail_shape(response, with_partial_members=True) + assert response["name"] == "Updated Dataset" + assert response["partial_member_list"] == ["user-1"] + mock_dataset_svc.update_dataset.assert_called_once() + _, update_data, _ = mock_dataset_svc.update_dataset.call_args.args + session = mock_dataset_svc.update_dataset.call_args.kwargs["session"] + assert session is controller_session + assert update_data["name"] == "Updated Dataset" + assert update_data["permission"] == "partial_members" + mock_perm_svc.update_partial_member_list.assert_called_once_with( + dataset.tenant_id, + dataset.id, + [{"user_id": "user-1", "role": "editor"}], + controller_session, + ) + + +class TestDatasetApiDelete: + """Test suite for DatasetApi.delete() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.DatasetPermissionService") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_delete_dataset_success( + self, + mock_dataset_svc: MagicMock, + mock_perm_svc: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + mock_dataset_svc.delete_dataset.return_value = True + + with app.test_request_context( + f"/datasets/{dataset.id}", + method="DELETE", + ): + api = DatasetApi() + result = unwrap(api.delete)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id) + + assert result == ("", 204) + mock_perm_svc.clear_partial_member_list.assert_called_once_with(dataset.id, controller_session) + + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_delete_dataset_not_found( + self, + mock_dataset_svc: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + mock_dataset_svc.delete_dataset.return_value = False + + with app.test_request_context( + f"/datasets/{dataset.id}", + method="DELETE", + ): + api = DatasetApi() + with pytest.raises(NotFound): + unwrap(api.delete)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id) + + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_delete_dataset_in_use( + self, + mock_dataset_svc: MagicMock, + app: Flask, + dataset: Dataset, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetApi + + mock_dataset_svc.delete_dataset.side_effect = services.errors.dataset.DatasetInUseError() + + with app.test_request_context( + f"/datasets/{dataset.id}", + method="DELETE", + ): + api = DatasetApi() + with pytest.raises(DatasetInUseError): + unwrap(api.delete)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id) + + +# --------------------------------------------------------------------------- +# API endpoint tests — DocumentStatusApi +# --------------------------------------------------------------------------- + + +class TestDocumentStatusApiPatch: + """Test suite for DocumentStatusApi.patch() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.DocumentService") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_batch_update_status_success( + self, + mock_dataset_svc: MagicMock, + mock_doc_svc: MagicMock, + app: Flask, + tenant: Tenant, + dataset: Dataset, + ) -> None: + from controllers.service_api.dataset.dataset import DocumentStatusApi + + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.check_dataset_permission.return_value = None + mock_dataset_svc.check_dataset_model_setting.return_value = None + mock_doc_svc.batch_update_document_status.return_value = None + + with app.test_request_context( + f"/datasets/{dataset.id}/documents/status/enable", + method="PATCH", + json={"document_ids": ["doc-1", "doc-2"]}, + ): + api = DocumentStatusApi() + response, status = api.patch( + tenant_id=tenant.id, + dataset_id=dataset.id, + action="enable", + ) + + assert status == 200 + assert response["result"] == "success" + + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_batch_update_status_dataset_not_found( + self, + mock_dataset_svc: MagicMock, + app: Flask, + tenant: Tenant, + dataset: Dataset, + ) -> None: + from controllers.service_api.dataset.dataset import DocumentStatusApi + + mock_dataset_svc.get_dataset.return_value = None + + with app.test_request_context( + f"/datasets/{dataset.id}/documents/status/enable", + method="PATCH", + json={"document_ids": ["doc-1"]}, + ): + api = DocumentStatusApi() + with pytest.raises(NotFound): + api.patch( + tenant_id=tenant.id, + dataset_id=dataset.id, + action="enable", + ) + + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_batch_update_status_permission_error( + self, + mock_dataset_svc: MagicMock, + app: Flask, + tenant: Tenant, + dataset: Dataset, + ) -> None: + from controllers.service_api.dataset.dataset import DocumentStatusApi + + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.check_dataset_permission.side_effect = services.errors.account.NoPermissionError( + "No permission" + ) + + with app.test_request_context( + f"/datasets/{dataset.id}/documents/status/enable", + method="PATCH", + json={"document_ids": ["doc-1"]}, + ): + api = DocumentStatusApi() + with pytest.raises(Forbidden): + api.patch( + tenant_id=tenant.id, + dataset_id=dataset.id, + action="enable", + ) + + @patch("controllers.service_api.dataset.dataset.DocumentService") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_batch_update_status_indexing_error( + self, + mock_dataset_svc: MagicMock, + mock_doc_svc: MagicMock, + app: Flask, + tenant: Tenant, + dataset: Dataset, + ) -> None: + from controllers.service_api.dataset.dataset import DocumentStatusApi + + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.check_dataset_permission.return_value = None + mock_dataset_svc.check_dataset_model_setting.return_value = None + mock_doc_svc.batch_update_document_status.side_effect = services.errors.document.DocumentIndexingError() + + with app.test_request_context( + f"/datasets/{dataset.id}/documents/status/enable", + method="PATCH", + json={"document_ids": ["doc-1"]}, + ): + api = DocumentStatusApi() + with pytest.raises(InvalidActionError): + api.patch( + tenant_id=tenant.id, + dataset_id=dataset.id, + action="enable", + ) + + @patch("controllers.service_api.dataset.dataset.DocumentService") + @patch("controllers.service_api.dataset.dataset.DatasetService") + def test_batch_update_status_value_error( + self, + mock_dataset_svc: MagicMock, + mock_doc_svc: MagicMock, + app: Flask, + tenant: Tenant, + dataset: Dataset, + ) -> None: + from controllers.service_api.dataset.dataset import DocumentStatusApi + + mock_dataset_svc.get_dataset.return_value = dataset + mock_dataset_svc.check_dataset_permission.return_value = None + mock_dataset_svc.check_dataset_model_setting.return_value = None + mock_doc_svc.batch_update_document_status.side_effect = ValueError("Invalid action") + + with app.test_request_context( + f"/datasets/{dataset.id}/documents/status/enable", + method="PATCH", + json={"document_ids": ["doc-1"]}, + ): + api = DocumentStatusApi() + with pytest.raises(InvalidActionError): + api.patch( + tenant_id=tenant.id, + dataset_id=dataset.id, + action="enable", + ) diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_payloads.py b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_payloads.py new file mode 100644 index 00000000000..d6755e957a6 --- /dev/null +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_payloads.py @@ -0,0 +1,207 @@ +"""Unit tests for Service API dataset request payloads.""" + +from typing import Literal + +import pytest + +from controllers.service_api.dataset.dataset import ( + DatasetCreatePayload, + DatasetListQuery, + DatasetUpdatePayload, + TagBindingPayload, + TagCreatePayload, + TagDeletePayload, + TagUnbindingPayload, + TagUpdatePayload, +) +from models.dataset import DatasetPermissionEnum + + +class TestDatasetCreatePayload: + """Test suite for DatasetCreatePayload Pydantic model.""" + + def test_payload_with_required_name(self) -> None: + payload = DatasetCreatePayload(name="Test Dataset") + assert payload.name == "Test Dataset" + assert payload.description == "" + assert payload.permission == DatasetPermissionEnum.ONLY_ME + + def test_payload_with_all_fields(self) -> None: + payload = DatasetCreatePayload( + name="Full Dataset", + description="A comprehensive dataset description", + indexing_technique="high_quality", + permission=DatasetPermissionEnum.ALL_TEAM, + provider="vendor", + embedding_model="text-embedding-ada-002", + embedding_model_provider="openai", + ) + assert payload.name == "Full Dataset" + assert payload.description == "A comprehensive dataset description" + assert payload.indexing_technique == "high_quality" + assert payload.permission == DatasetPermissionEnum.ALL_TEAM + assert payload.provider == "vendor" + assert payload.embedding_model == "text-embedding-ada-002" + assert payload.embedding_model_provider == "openai" + + def test_payload_name_length_validation_min(self) -> None: + with pytest.raises(ValueError): + DatasetCreatePayload(name="") + + def test_payload_name_length_validation_max(self) -> None: + with pytest.raises(ValueError): + DatasetCreatePayload(name="A" * 41) + + def test_payload_description_max_length(self) -> None: + with pytest.raises(ValueError): + DatasetCreatePayload(name="Dataset", description="A" * 401) + + @pytest.mark.parametrize("technique", ["high_quality", "economy"]) + def test_payload_valid_indexing_techniques(self, technique: Literal["high_quality", "economy"]) -> None: + payload = DatasetCreatePayload(name="Dataset", indexing_technique=technique) + assert payload.indexing_technique == technique + + def test_payload_with_external_knowledge_settings(self) -> None: + payload = DatasetCreatePayload( + name="External Dataset", external_knowledge_api_id="api_123", external_knowledge_id="knowledge_456" + ) + assert payload.external_knowledge_api_id == "api_123" + assert payload.external_knowledge_id == "knowledge_456" + + +class TestDatasetUpdatePayload: + """Test suite for DatasetUpdatePayload Pydantic model.""" + + def test_payload_all_optional(self) -> None: + payload = DatasetUpdatePayload() + assert payload.name is None + assert payload.description is None + assert payload.permission is None + + def test_payload_with_partial_update(self) -> None: + payload = DatasetUpdatePayload(name="Updated Name", description="Updated description") + assert payload.name == "Updated Name" + assert payload.description == "Updated description" + + def test_payload_with_permission_change(self) -> None: + payload = DatasetUpdatePayload( + permission=DatasetPermissionEnum.PARTIAL_TEAM, + partial_member_list=[{"user_id": "user_123", "role": "editor"}], + ) + assert payload.permission == DatasetPermissionEnum.PARTIAL_TEAM + assert payload.partial_member_list is not None + assert len(payload.partial_member_list) == 1 + + def test_payload_name_length_validation(self) -> None: + with pytest.raises(ValueError): + DatasetUpdatePayload(name="") + with pytest.raises(ValueError): + DatasetUpdatePayload(name="A" * 41) + + +class TestDatasetListQuery: + """Test suite for DatasetListQuery Pydantic model.""" + + def test_query_with_defaults(self) -> None: + query = DatasetListQuery() + assert query.page == 1 + assert query.limit == 20 + assert query.keyword is None + assert query.include_all is False + assert query.tag_ids == [] + + def test_query_with_all_filters(self) -> None: + query = DatasetListQuery( + page=3, limit=50, keyword="machine learning", include_all=True, tag_ids=["tag1", "tag2", "tag3"] + ) + assert query.page == 3 + assert query.limit == 50 + assert query.keyword == "machine learning" + assert query.include_all is True + assert len(query.tag_ids) == 3 + + def test_query_with_tag_filter(self) -> None: + query = DatasetListQuery(tag_ids=["tag_abc", "tag_def"]) + assert query.tag_ids == ["tag_abc", "tag_def"] + + +class TestTagCreatePayload: + """Test suite for TagCreatePayload Pydantic model.""" + + def test_payload_with_name(self) -> None: + payload = TagCreatePayload(name="New Tag") + assert payload.name == "New Tag" + + def test_payload_name_length_min(self) -> None: + with pytest.raises(ValueError): + TagCreatePayload(name="") + + def test_payload_name_length_max(self) -> None: + with pytest.raises(ValueError): + TagCreatePayload(name="A" * 51) + + def test_payload_with_unicode_name(self) -> None: + payload = TagCreatePayload(name="标签 🏷️ Тег") + assert payload.name == "标签 🏷️ Тег" + + +class TestTagUpdatePayload: + """Test suite for TagUpdatePayload Pydantic model.""" + + def test_payload_with_name_and_id(self) -> None: + payload = TagUpdatePayload(name="Updated Tag", tag_id="tag_123") + assert payload.name == "Updated Tag" + assert payload.tag_id == "tag_123" + + def test_payload_requires_tag_id(self) -> None: + with pytest.raises(ValueError): + TagUpdatePayload.model_validate({"name": "Updated Tag"}) + + +class TestTagDeletePayload: + """Test suite for TagDeletePayload Pydantic model.""" + + def test_payload_with_tag_id(self) -> None: + payload = TagDeletePayload(tag_id="tag_to_delete") + assert payload.tag_id == "tag_to_delete" + + def test_payload_requires_tag_id(self) -> None: + with pytest.raises(ValueError): + TagDeletePayload.model_validate({}) + + +class TestTagBindingPayload: + """Test suite for TagBindingPayload Pydantic model.""" + + def test_payload_with_valid_data(self) -> None: + payload = TagBindingPayload(tag_ids=["tag1", "tag2"], target_id="dataset_123") + assert len(payload.tag_ids) == 2 + assert payload.target_id == "dataset_123" + + def test_payload_rejects_empty_tag_ids(self) -> None: + with pytest.raises(ValueError) as exc_info: + TagBindingPayload(tag_ids=[], target_id="dataset_123") + assert "Tag IDs is required" in str(exc_info.value) + + def test_payload_single_tag_id(self) -> None: + payload = TagBindingPayload(tag_ids=["single_tag"], target_id="dataset_456") + assert payload.tag_ids == ["single_tag"] + + +class TestTagUnbindingPayload: + """Test suite for TagUnbindingPayload Pydantic model.""" + + def test_payload_with_valid_data(self) -> None: + payload = TagUnbindingPayload(tag_ids=["tag_123"], target_id="dataset_456") + assert payload.tag_ids == ["tag_123"] + assert payload.target_id == "dataset_456" + + def test_payload_normalizes_legacy_tag_id(self) -> None: + payload = TagUnbindingPayload(tag_id="tag_123", target_id="dataset_456") + assert payload.tag_ids == ["tag_123"] + assert payload.target_id == "dataset_456" + + def test_payload_rejects_empty_tag_ids(self) -> None: + with pytest.raises(ValueError) as exc_info: + TagUnbindingPayload(tag_ids=[], target_id="dataset_456") + assert "Tag IDs is required" in str(exc_info.value) diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py index c7bac28b694..9fae95edcfc 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_segment.py @@ -100,9 +100,9 @@ def _child_chunk() -> ChildChunk: content="child chunk content", word_count=3, created_by="account-1", + type=SegmentType.CUSTOMIZED, ) child_chunk.id = "child-1" - child_chunk.type = SegmentType.CUSTOMIZED child_chunk.created_at = naive_utc_now() child_chunk.updated_at = naive_utc_now() return child_chunk @@ -358,34 +358,64 @@ class TestSegmentServiceMockedBehavior: @pytest.fixture def mock_dataset(self): """Create mock dataset.""" - dataset = Mock(spec=Dataset) - dataset.id = str(uuid.uuid4()) - dataset.tenant_id = str(uuid.uuid4()) + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + ) return dataset @pytest.fixture def mock_document(self): """Create mock document.""" - document = Mock(spec=Document) - document.id = str(uuid.uuid4()) - document.dataset_id = str(uuid.uuid4()) - document.indexing_status = "completed" - document.enabled = True + document = Document( + id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + indexing_status="completed", + enabled=True, + ) return document @pytest.fixture def mock_segment(self): """Create mock segment.""" - segment = Mock(spec=DocumentSegment) + segment = DocumentSegment( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id=str(uuid.uuid4()), + position=1, + content="Test content", + word_count=0, + tokens=0, + created_by="account-id", + ) segment.id = str(uuid.uuid4()) - segment.document_id = str(uuid.uuid4()) - segment.content = "Test content" return segment @patch.object(SegmentService, "multi_create_segment") def test_create_segments_returns_list(self, mock_create, mock_dataset, mock_document): """Test segment creation returns list of segments.""" - mock_segments = [Mock(spec=DocumentSegment), Mock(spec=DocumentSegment)] + mock_segments = [ + DocumentSegment( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + position=1, + content="", + word_count=0, + tokens=0, + created_by="account-id", + ), + DocumentSegment( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + position=1, + content="", + word_count=0, + tokens=0, + created_by="account-id", + ), + ] mock_create.return_value = mock_segments session = Mock() @@ -459,17 +489,33 @@ class TestChildChunkServiceMockedBehavior: @pytest.fixture def mock_segment(self): """Create mock segment.""" - segment = Mock(spec=DocumentSegment) + segment = DocumentSegment( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + position=1, + content="", + word_count=0, + tokens=0, + created_by="account-id", + ) segment.id = str(uuid.uuid4()) return segment @pytest.fixture def mock_child_chunk(self): """Create mock child chunk.""" - chunk = Mock(spec=ChildChunk) + chunk = ChildChunk( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + segment_id=str(uuid.uuid4()), + position=1, + content="Child chunk content", + word_count=0, + created_by="account-id", + ) chunk.id = str(uuid.uuid4()) - chunk.segment_id = str(uuid.uuid4()) - chunk.content = "Child chunk content" return chunk @patch.object(SegmentService, "create_child_chunk") @@ -480,8 +526,8 @@ class TestChildChunkServiceMockedBehavior: result = SegmentService.create_child_chunk( content="New chunk content", segment=mock_segment, - document=Mock(spec=Document), - dataset=Mock(spec=Dataset), + document=Document(), + dataset=Dataset(), session=Mock(), ) @@ -524,16 +570,33 @@ class TestChildChunkServiceMockedBehavior: @patch.object(SegmentService, "update_child_chunk") def test_update_child_chunk_returns_updated_chunk(self, mock_update, mock_child_chunk): """Test update_child_chunk returns updated chunk.""" - updated_chunk = Mock(spec=ChildChunk) - updated_chunk.content = "Updated content" + updated_chunk = ChildChunk( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + segment_id="segment-id", + position=1, + content="Updated content", + word_count=0, + created_by="account-id", + ) mock_update.return_value = updated_chunk result = SegmentService.update_child_chunk( content="Updated content", child_chunk=mock_child_chunk, - segment=Mock(spec=DocumentSegment), - document=Mock(spec=Document), - dataset=Mock(spec=Dataset), + segment=DocumentSegment( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + position=1, + content="", + word_count=0, + tokens=0, + created_by="account-id", + ), + document=Document(), + dataset=Dataset(), session=Mock(), ) @@ -545,26 +608,30 @@ class TestDocumentValidation: def test_document_indexing_status_completed_is_valid(self): """Test that completed indexing status is valid.""" - document = Mock(spec=Document) - document.indexing_status = "completed" + document = Document( + indexing_status="completed", + ) assert document.indexing_status == "completed" def test_document_indexing_status_indexing_is_invalid(self): """Test that indexing status is invalid for segment operations.""" - document = Mock(spec=Document) - document.indexing_status = "indexing" + document = Document( + indexing_status="indexing", + ) assert document.indexing_status != "completed" def test_document_enabled_true_is_valid(self): """Test that enabled=True is valid.""" - document = Mock(spec=Document) - document.enabled = True + document = Document( + enabled=True, + ) assert document.enabled def test_document_enabled_false_is_invalid(self): """Test that enabled=False is invalid for segment operations.""" - document = Mock(spec=Document) - document.enabled = False + document = Document( + enabled=False, + ) assert not document.enabled @@ -573,10 +640,11 @@ class TestDatasetModels: def test_dataset_has_required_fields(self): """Test Dataset model has required fields.""" - dataset = Mock(spec=Dataset) - dataset.id = str(uuid.uuid4()) - dataset.tenant_id = str(uuid.uuid4()) - dataset.indexing_technique = "economy" + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + indexing_technique="economy", + ) assert dataset.id is not None assert dataset.tenant_id is not None @@ -584,11 +652,17 @@ class TestDatasetModels: def test_document_segment_has_required_fields(self): """Test DocumentSegment model has required fields.""" - segment = Mock(spec=DocumentSegment) + segment = DocumentSegment( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id=str(uuid.uuid4()), + position=1, + content="Test content", + word_count=0, + tokens=0, + created_by="account-id", + ) segment.id = str(uuid.uuid4()) - segment.document_id = str(uuid.uuid4()) - segment.content = "Test content" - segment.position = 1 assert segment.id is not None assert segment.document_id is not None @@ -596,10 +670,17 @@ class TestDatasetModels: def test_child_chunk_has_required_fields(self): """Test ChildChunk model has required fields.""" - chunk = Mock(spec=ChildChunk) + chunk = ChildChunk( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + segment_id=str(uuid.uuid4()), + position=1, + content="Chunk content", + word_count=0, + created_by="account-id", + ) chunk.id = str(uuid.uuid4()) - chunk.segment_id = str(uuid.uuid4()) - chunk.content = "Chunk content" assert chunk.id is not None assert chunk.segment_id is not None @@ -787,8 +868,9 @@ class TestSegmentIndexingRequirements: @pytest.mark.parametrize("technique", ["high_quality", "economy"]) def test_indexing_technique_values(self, technique): """Test valid indexing technique values.""" - dataset = Mock(spec=Dataset) - dataset.indexing_technique = technique + dataset = Dataset( + indexing_technique=technique, + ) assert dataset.indexing_technique in ["high_quality", "economy"] @pytest.mark.parametrize( @@ -803,8 +885,9 @@ class TestSegmentIndexingRequirements: ) def test_valid_indexing_statuses(self, status): """Test valid document indexing statuses.""" - document = Mock(spec=Document) - document.indexing_status = status + document = Document( + indexing_status=status, + ) assert document.indexing_status in { IndexingStatus.WAITING, IndexingStatus.PARSING, @@ -815,9 +898,10 @@ class TestSegmentIndexingRequirements: def test_completed_status_required_for_segments(self): """Test that completed status is required for segment operations.""" - document = Mock(spec=Document) - document.indexing_status = "completed" - document.enabled = True + document = Document( + indexing_status="completed", + enabled=True, + ) # Both conditions must be true assert document.indexing_status == "completed" diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_tag_apis.py b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_tag_apis.py new file mode 100644 index 00000000000..420b1fa1bf7 --- /dev/null +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_dataset_tag_apis.py @@ -0,0 +1,380 @@ +"""Unit tests for Service API dataset tag controller behavior. + +Service boundaries stay mocked, while users, tenants, and tags are real ORM objects +persisted in SQLite. Controller database calls share that SQLite session so assertions +cover the concrete objects and session passed across the controller boundary. +""" + +import uuid +from inspect import unwrap +from typing import cast +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask +from sqlalchemy.orm import Session, scoped_session, sessionmaker +from werkzeug.exceptions import Forbidden + +from extensions.ext_database import db +from models.account import Account, Tenant, TenantAccountRole +from models.enums import TagType +from models.model import Tag + +TAG_MODEL_TABLES = (Account, Tenant, Tag) +pytestmark = pytest.mark.parametrize("sqlite_session", [TAG_MODEL_TABLES], indirect=True) + + +@pytest.fixture(autouse=True) +def controller_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Route controller database access through the test's SQLite session.""" + + # Flask-SQLAlchemy exposes a callable registry that also proxies Session methods. + # Seed that registry with this fixture's Session so both access styles share one transaction. + existing_session_factory = cast(sessionmaker[Session], lambda: sqlite_session) + session_registry = scoped_session(existing_session_factory) + monkeypatch.setattr(db, "session", session_registry) + return sqlite_session + + +@pytest.fixture +def tenant(controller_session: Session) -> Tenant: + tenant = Tenant(name="Dataset Tag API Tenant") + controller_session.add(tenant) + controller_session.flush() + return tenant + + +@pytest.fixture +def account(controller_session: Session, tenant: Tenant, monkeypatch: pytest.MonkeyPatch) -> Account: + account = Account(name="Dataset Tag API User", email=f"dataset-tag-api-{uuid.uuid4()}@example.com") + account.role = TenantAccountRole.OWNER + account._current_tenant = tenant + controller_session.add(account) + controller_session.flush() + + # Inject the concrete account at the controller boundary without relying on Flask-Login globals. + from controllers.service_api.dataset import dataset as dataset_module + + monkeypatch.setattr(dataset_module, "current_user", account) + return account + + +def make_tag( + session: Session, + tenant: Tenant, + account: Account, + *, + id: str, + name: str, + binding_count: int | None = None, +) -> Tag: + """Create and flush a real tag, optionally adding the aggregate count returned by TagService.""" + + tag = Tag(tenant_id=tenant.id, type=TagType.KNOWLEDGE, name=name, created_by=account.id) + tag.id = id + session.add(tag) + session.flush() + if binding_count is not None: + tag.__dict__["binding_count"] = binding_count + return tag + + +class TestDatasetTagsApiGet: + """Test suite for DatasetTagsApi.get() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.TagService") + def test_list_tags_success( + self, + mock_tag_svc: MagicMock, + app: Flask, + account: Account, + tenant: Tenant, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetTagsApi + + tag = make_tag(controller_session, tenant, account, id="tag-1", name="Test Tag", binding_count=0) + mock_tag_svc.get_tags.return_value = [tag] + + with app.test_request_context("/datasets/tags", method="GET"): + api = DatasetTagsApi() + response, status = unwrap(api.get)(api, controller_session, _=None) + + assert status == 200 + assert response == [{"id": "tag-1", "name": "Test Tag", "type": "knowledge", "binding_count": "0"}] + mock_tag_svc.get_tags.assert_called_once_with("knowledge", tenant.id, session=controller_session) + + +class TestDatasetTagsApiPost: + """Test suite for DatasetTagsApi.post() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.TagService") + def test_create_tag_success( + self, + mock_tag_svc: MagicMock, + app: Flask, + account: Account, + tenant: Tenant, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetTagsApi + + tag = make_tag(controller_session, tenant, account, id="tag-new", name="New Tag") + mock_tag_svc.save_tags.return_value = tag + + with app.test_request_context( + "/datasets/tags", + method="POST", + json={"name": "New Tag"}, + ): + api = DatasetTagsApi() + response, status = unwrap(api.post)(api, controller_session, _=None) + + assert status == 200 + assert response == {"id": "tag-new", "name": "New Tag", "type": "knowledge", "binding_count": "0"} + mock_tag_svc.save_tags.assert_called_once() + + def test_create_tag_forbidden(self, app: Flask, account: Account) -> None: + from controllers.service_api.dataset.dataset import DatasetTagsApi + + account.role = TenantAccountRole.NORMAL + + with app.test_request_context( + "/datasets/tags", + method="POST", + json={"name": "New Tag"}, + ): + api = DatasetTagsApi() + with pytest.raises(Forbidden): + api.post(_=None) + + +class TestDatasetTagsApiPatch: + """Test suite for DatasetTagsApi.patch() endpoint.""" + + @patch("controllers.service_api.dataset.dataset.TagService") + @patch("controllers.service_api.dataset.dataset.service_api_ns") + def test_update_tag_success( + self, + mock_service_api_ns: MagicMock, + mock_tag_svc: MagicMock, + app: Flask, + account: Account, + tenant: Tenant, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetTagsApi + + tag = make_tag(controller_session, tenant, account, id="tag-1", name="Updated Tag") + mock_tag_svc.update_tags.return_value = tag + mock_tag_svc.get_tag_binding_count.return_value = 5 + mock_service_api_ns.payload = {"name": "Updated Tag", "tag_id": "tag-1"} + + with app.test_request_context( + "/datasets/tags", + method="PATCH", + json={"name": "Updated Tag", "tag_id": "tag-1"}, + ): + api = DatasetTagsApi() + response, status = unwrap(api.patch)(api, controller_session, _=None) + + assert status == 200 + assert response == {"id": "tag-1", "name": "Updated Tag", "type": "knowledge", "binding_count": "5"} + mock_tag_svc.update_tags.assert_called_once() + update_payload, tag_id, session = mock_tag_svc.update_tags.call_args.args + assert update_payload.name == "Updated Tag" + assert tag_id == "tag-1" + assert session is controller_session + + def test_update_tag_forbidden(self, app: Flask, account: Account) -> None: + from controllers.service_api.dataset.dataset import DatasetTagsApi + + account.role = TenantAccountRole.NORMAL + + with app.test_request_context( + "/datasets/tags", + method="PATCH", + json={"name": "Updated Tag", "tag_id": "tag-1"}, + ): + api = DatasetTagsApi() + with pytest.raises(Forbidden): + api.patch(_=None) + + +class TestDatasetTagsApiDelete: + """Test suite for DatasetTagsApi.delete() endpoint.""" + + @pytest.mark.usefixtures("account") + @patch("controllers.service_api.dataset.dataset.TagService") + @patch("controllers.service_api.dataset.dataset.service_api_ns") + def test_delete_tag_success( + self, + mock_service_api_ns: MagicMock, + mock_tag_svc: MagicMock, + app: Flask, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetTagsApi + + mock_tag_svc.delete_tag.return_value = None + mock_service_api_ns.payload = {"tag_id": "tag-1"} + + with app.test_request_context( + "/datasets/tags", + method="DELETE", + json={"tag_id": "tag-1"}, + ): + api = DatasetTagsApi() + result = unwrap(api.delete)(api, controller_session, _=None) + + assert result == ("", 204) + mock_tag_svc.delete_tag.assert_called_once_with("tag-1", controller_session, tag_type=TagType.KNOWLEDGE) + + +class TestDatasetTagsBindingStatusApi: + """Test suite for DatasetTagsBindingStatusApi endpoints.""" + + @patch("controllers.service_api.dataset.dataset.TagService") + def test_get_dataset_tags_binding_status( + self, + mock_tag_svc: MagicMock, + app: Flask, + account: Account, + tenant: Tenant, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetTagsBindingStatusApi + + tag = make_tag(controller_session, tenant, account, id="tag_1", name="Test Tag") + mock_tag_svc.get_tags_by_target_id.return_value = [tag] + + with app.test_request_context("/", method="GET"): + api = DatasetTagsBindingStatusApi() + response, status_code = unwrap(api.get)(api, controller_session, tenant.id, dataset_id="dataset_123") + + assert status_code == 200 + assert response["data"] == [{"id": "tag_1", "name": "Test Tag"}] + assert response["total"] == 1 + mock_tag_svc.get_tags_by_target_id.assert_called_once_with( + "knowledge", tenant.id, "dataset_123", controller_session + ) + + +class TestDatasetTagBindingApiPost: + """Test suite for DatasetTagBindingApi.post() endpoint.""" + + @pytest.mark.usefixtures("account") + @patch("controllers.service_api.dataset.dataset.TagService") + def test_bind_tags_success( + self, + mock_tag_svc: MagicMock, + app: Flask, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetTagBindingApi + + mock_tag_svc.save_tag_binding.return_value = None + + with app.test_request_context( + "/datasets/tags/binding", + method="POST", + json={"tag_ids": ["tag-1"], "target_id": "ds-1"}, + ): + api = DatasetTagBindingApi() + result = unwrap(api.post)(api, controller_session, _=None) + + assert result == ("", 204) + from services.tag_service import TagBindingCreatePayload + + mock_tag_svc.save_tag_binding.assert_called_once_with( + TagBindingCreatePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE), + controller_session, + ) + + def test_bind_tags_forbidden(self, app: Flask, account: Account) -> None: + from controllers.service_api.dataset.dataset import DatasetTagBindingApi + + account.role = TenantAccountRole.NORMAL + + with app.test_request_context( + "/datasets/tags/binding", + method="POST", + json={"tag_ids": ["tag-1"], "target_id": "ds-1"}, + ): + api = DatasetTagBindingApi() + with pytest.raises(Forbidden): + api.post(_=None) + + +class TestDatasetTagUnbindingApiPost: + """Test suite for DatasetTagUnbindingApi.post() endpoint.""" + + @pytest.mark.usefixtures("account") + @patch("controllers.service_api.dataset.dataset.TagService") + def test_unbind_tag_success( + self, + mock_tag_svc: MagicMock, + app: Flask, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi + + mock_tag_svc.delete_tag_binding.return_value = None + + with app.test_request_context( + "/datasets/tags/unbinding", + method="POST", + json={"tag_ids": ["tag-1"], "target_id": "ds-1"}, + ): + api = DatasetTagUnbindingApi() + result = unwrap(api.post)(api, controller_session, _=None) + + assert result == ("", 204) + from services.tag_service import TagBindingDeletePayload + + mock_tag_svc.delete_tag_binding.assert_called_once_with( + TagBindingDeletePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE), + controller_session, + ) + + @pytest.mark.usefixtures("account") + @patch("controllers.service_api.dataset.dataset.TagService") + def test_unbind_legacy_tag_id_success( + self, + mock_tag_svc: MagicMock, + app: Flask, + controller_session: Session, + ) -> None: + from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi + + mock_tag_svc.delete_tag_binding.return_value = None + + with app.test_request_context( + "/datasets/tags/unbinding", + method="POST", + json={"tag_id": "tag-1", "target_id": "ds-1"}, + ): + api = DatasetTagUnbindingApi() + result = unwrap(api.post)(api, controller_session, _=None) + + assert result == ("", 204) + from services.tag_service import TagBindingDeletePayload + + mock_tag_svc.delete_tag_binding.assert_called_once_with( + TagBindingDeletePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE), + controller_session, + ) + + def test_unbind_tag_forbidden(self, app: Flask, account: Account) -> None: + from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi + + account.role = TenantAccountRole.NORMAL + + with app.test_request_context( + "/datasets/tags/unbinding", + method="POST", + json={"tag_ids": ["tag-1"], "target_id": "ds-1"}, + ): + api = DatasetTagUnbindingApi() + with pytest.raises(Forbidden): + api.post(_=None) diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py index 9713caabeb5..19d37038421 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_document.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_document.py @@ -20,12 +20,14 @@ import json import uuid from dataclasses import dataclass from datetime import UTC, datetime -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound +from controllers.common.errors import FileTooLargeError as FileTooLargeHTTPError from controllers.service_api.dataset.document import ( DeprecatedDocumentAddByTextApi, DeprecatedDocumentUpdateByFileApi, @@ -43,10 +45,12 @@ from controllers.service_api.dataset.document import ( ) from controllers.service_api.dataset.error import ArchivedDocumentImmutableError from core.rag.index_processor.constant.index_type import IndexStructureType -from models.dataset import Dataset, Document -from models.enums import DataSourceType, DocumentCreatedFrom, DocumentDocType, IndexingStatus +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, DocumentDocType, IndexingStatus, SegmentStatus +from services.dataset_ref_service import DatasetRef from services.dataset_service import DocumentService from services.entities.knowledge_entities.knowledge_entities import ProcessRule, RetrievalModel +from services.errors.file import FileTooLargeError as FileTooLargeServiceError def _document_data_source_info() -> dict[str, str]: @@ -174,6 +178,26 @@ def _expected_document_response(document: Document) -> dict[str, object]: } +def _persist_segments(session: Session, document: Document, count: int = 5) -> None: + session.add_all( + [ + DocumentSegment( + tenant_id=document.tenant_id, + dataset_id=document.dataset_id, + document_id=document.id, + position=index, + content=f"segment {index}", + word_count=20, + tokens=5, + created_by=document.created_by, + status=SegmentStatus.COMPLETED, + ) + for index in range(1, count + 1) + ] + ) + session.flush() + + class TestDocumentTextCreatePayload: """Test suite for DocumentTextCreatePayload Pydantic model.""" @@ -368,10 +392,10 @@ class TestDocumentService: assert result.indexing_status == "completed" @patch.object(DocumentService, "delete_document") - def test_delete_document_called(self, mock_delete): + def test_delete_document_called(self, mock_delete, sqlite_session: Session): """Test delete_document is called with document.""" document = make_serializable_document() - session = Mock() + session = sqlite_session DocumentService.delete_document(document=document, session=session) mock_delete.assert_called_once_with(document=document, session=session) @@ -548,24 +572,28 @@ class TestDocumentDisplayStatusLogic: class TestDocumentServiceBatchMethods: """Test DocumentService batch operations.""" - def test_get_documents_by_ids(self): + def test_get_documents_by_ids(self, sqlite_session: Session): """Test batch retrieval of documents by IDs.""" dataset_id = str(uuid.uuid4()) doc_ids = [str(uuid.uuid4()), str(uuid.uuid4())] - session = Mock() - mock_result = Mock() - mock_result.all.return_value = [Mock(id=doc_ids[0]), Mock(id=doc_ids[1])] - session.scalars.return_value = mock_result + session = sqlite_session + session.add_all( + [ + make_serializable_document(id=document_id, tenant_id="tenant-id", dataset_id=dataset_id) + for document_id in doc_ids + ] + ) + session.flush() - documents = DocumentService.get_documents_by_ids(dataset_id, doc_ids, session) + documents = DocumentService.get_documents_by_ids(DatasetRef("tenant-id", dataset_id), doc_ids, session) assert len(documents) == 2 - session.scalars.assert_called_once() + assert {document.id for document in documents} == set(doc_ids) - def test_get_documents_by_ids_empty(self): + def test_get_documents_by_ids_empty(self, sqlite_session: Session): """Test batch retrieval with empty list returns empty.""" - assert DocumentService.get_documents_by_ids("ds_id", [], Mock()) == [] + assert DocumentService.get_documents_by_ids(DatasetRef("tenant-id", "ds_id"), [], sqlite_session) == [] class TestDocumentServiceFileOperations: @@ -573,7 +601,7 @@ class TestDocumentServiceFileOperations: @patch("services.dataset_service.file_helpers.get_signed_file_url") @patch("services.dataset_service.DocumentService._get_upload_file_for_upload_file_document") - def test_get_document_download_url(self, mock_get_file, mock_signed_url): + def test_get_document_download_url(self, mock_get_file, mock_signed_url, sqlite_session: Session): """Test generation of download URL.""" mock_doc = Mock() mock_file = Mock() @@ -581,7 +609,7 @@ class TestDocumentServiceFileOperations: mock_get_file.return_value = mock_file mock_signed_url.return_value = "https://example.com/download" - session = Mock() + session = sqlite_session url = DocumentService.get_document_download_url(mock_doc, session) assert url == "https://example.com/download" @@ -595,7 +623,7 @@ class TestDocumentServiceSaveValidation: @patch("services.dataset_service.DatasetService.check_doc_form") @patch("services.dataset_service.FeatureService.get_features") @patch("services.dataset_service.current_user") - def test_save_document_validates_doc_form(self, mock_user, mock_features, mock_check_form): + def test_save_document_validates_doc_form(self, mock_user, mock_features, mock_check_form, sqlite_session: Session): """Test that doc_form is validated during save.""" mock_user.current_tenant_id = "tenant_id" dataset = Mock() @@ -608,7 +636,7 @@ class TestDocumentServiceSaveValidation: pass mock_check_form.side_effect = TestStopError() - session = Mock() + session = sqlite_session # Skip actual logic by mocking dependent calls or raising error to stop early with pytest.raises(TestStopError): @@ -661,6 +689,7 @@ class TestDocumentApiGet: app: Flask, mock_tenant: str, mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test successful document retrieval with metadata='all'.""" # Arrange @@ -668,9 +697,10 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail + mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset mock_dataset_svc.get_process_rules.return_value = {"mode": "automatic", "rules": {}} - session = MagicMock() - session.scalar.side_effect = [5, 0] + session = sqlite_session + _persist_segments(session, mock_doc_detail) # Act with app.test_request_context( @@ -724,15 +754,16 @@ class TestDocumentApiGet: assert response["summary_index_status"] is None @patch("controllers.service_api.dataset.document.DocumentService") - def test_get_document_not_found(self, mock_doc_svc: Mock, app: Flask, mock_tenant: str) -> None: + def test_get_document_not_found( + self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, sqlite_session: Session + ) -> None: """Test 404 when document is not found.""" # Arrange dataset_id = str(uuid.uuid4()) mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant) mock_doc_svc.get_document.return_value = None - session = MagicMock() - session.scalar.return_value = mock_dataset + session = sqlite_session # Act & Assert with app.test_request_context( @@ -752,7 +783,12 @@ class TestDocumentApiGet: @patch("controllers.service_api.dataset.document.DocumentService") def test_get_document_forbidden_wrong_tenant( - self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document + self, + mock_doc_svc: Mock, + app: Flask, + mock_tenant: str, + mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test 403 when document tenant doesn't match request tenant.""" # Arrange @@ -761,8 +797,9 @@ class TestDocumentApiGet: mock_doc_detail.tenant_id = "different-tenant-id" mock_doc_svc.get_document.return_value = mock_doc_detail - session = MagicMock() - session.scalar.return_value = mock_dataset + session = sqlite_session + session.add(mock_dataset) + session.flush() # Act & Assert with app.test_request_context( @@ -782,7 +819,12 @@ class TestDocumentApiGet: @patch("controllers.service_api.dataset.document.DocumentService") def test_get_document_metadata_only( - self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document + self, + mock_doc_svc: Mock, + app: Flask, + mock_tenant: str, + mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test document retrieval with metadata='only'.""" # Arrange @@ -790,8 +832,9 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail - session = MagicMock() - session.scalar.return_value = mock_dataset + session = sqlite_session + session.add(mock_dataset) + session.flush() # Act with app.test_request_context( @@ -825,6 +868,7 @@ class TestDocumentApiGet: app: Flask, mock_tenant: str, mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test document retrieval with metadata='without'.""" # Arrange @@ -832,9 +876,10 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail + mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset mock_dataset_svc.get_process_rules.return_value = {"mode": "automatic", "rules": {}} - session = MagicMock() - session.scalar.side_effect = [5, 0] + session = sqlite_session + _persist_segments(session, mock_doc_detail) # Act with app.test_request_context( @@ -893,7 +938,12 @@ class TestDocumentApiGet: @patch("controllers.service_api.dataset.document.DocumentService") def test_get_document_invalid_metadata_value( - self, mock_doc_svc: Mock, app: Flask, mock_tenant: str, mock_doc_detail: Document + self, + mock_doc_svc: Mock, + app: Flask, + mock_tenant: str, + mock_doc_detail: Document, + sqlite_session: Session, ) -> None: """Test error when metadata parameter has invalid value.""" # Arrange @@ -901,8 +951,9 @@ class TestDocumentApiGet: mock_dataset = make_dataset(id=dataset_id, tenant_id=mock_tenant, summary_index_setting=None) mock_doc_svc.get_document.return_value = mock_doc_detail - session = MagicMock() - session.scalar.return_value = mock_dataset + session = sqlite_session + session.add(mock_dataset) + session.flush() # Act & Assert with app.test_request_context( @@ -1155,6 +1206,9 @@ class TestDocumentIndexingStatusApi: "completed_at": 1609459204, "paused_at": None, "error": None, + "error_code": None, + "estimated_vector_space_mb": None, + "vector_space_limit_mb": None, "stopped_at": None, "completed_segments": 5, "total_segments": 5, @@ -1459,6 +1513,7 @@ class TestDocumentUpdateByTextApiPost: app: Flask, mock_tenant, mock_dataset, + sqlite_session: Session, ): """Test successful document update by text.""" _setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant) @@ -1593,6 +1648,52 @@ class TestDocumentAddByFileApiPost: 200, ) + @patch( + "controllers.service_api.dataset.document.FeatureService.get_knowledge_file_size_limit", + return_value=15, + ) + @patch("controllers.service_api.dataset.document.FileService") + @patch("controllers.service_api.dataset.document.current_user") + @patch("controllers.service_api.dataset.document.db") + def test_add_by_file_too_large_returns_http_413( + self, + mock_db, + mock_current_user, + mock_file_svc_cls, + mock_get_limit, + app: Flask, + mock_tenant, + mock_dataset, + ): + mock_dataset.provider = "vendor" + mock_dataset.indexing_technique = "economy" + mock_dataset.chunk_structure = None + mock_db.session.scalar.return_value = mock_dataset + mock_current_user.__bool__ = Mock(return_value=True) + mock_file_svc_cls.return_value.upload_file.side_effect = FileTooLargeServiceError() + + from io import BytesIO + + data = { + "file": (BytesIO(b"oversized content"), "test.pdf", "application/pdf"), + "data": json.dumps({"process_rule": {"mode": "automatic", "rules": None}}), + } + with app.test_request_context( + f"/datasets/{mock_dataset.id}/document/create-by-file", + method="POST", + content_type="multipart/form-data", + data=data, + ): + api = DocumentAddByFileApi() + with pytest.raises(FileTooLargeHTTPError) as exc_info: + _unwrap_non_wrapped_controller(type(api).post)( + api, mock_db.session, tenant_id=mock_tenant, dataset_id=mock_dataset.id + ) + + assert exc_info.value.code == 413 + assert exc_info.value.error_code == "file_too_large" + mock_get_limit.assert_called_once_with(mock_tenant) + @patch("controllers.service_api.dataset.document.db") @patch("controllers.service_api.wraps.FeatureService") @patch("controllers.service_api.wraps.validate_and_get_api_token") @@ -1748,6 +1849,7 @@ class TestDocumentUpdateByFileApiPatch: app: Flask, mock_tenant, mock_dataset, + sqlite_session: Session, ): """Test legacy POST aliases still dispatch while marked deprecated.""" _setup_billing_mocks(mock_validate_token, mock_feature_svc, mock_tenant) @@ -1757,8 +1859,7 @@ class TestDocumentUpdateByFileApiPatch: ) doc_id = str(uuid.uuid4()) - session = MagicMock() - session.scalar.return_value = 0 + session = sqlite_session with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/{doc_id}/{route_name}", method="POST", diff --git a/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py b/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py index c30762911e2..3df48fdbfe5 100644 --- a/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py +++ b/api/tests/unit_tests/controllers/service_api/dataset/test_metadata.py @@ -17,10 +17,11 @@ Decorator strategy: import uuid from inspect import unwrap -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import ANY, MagicMock, Mock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound from controllers.service_api.dataset.metadata import ( @@ -30,6 +31,7 @@ from controllers.service_api.dataset.metadata import ( DatasetMetadataServiceApi, DocumentMetadataEditServiceApi, ) +from services.errors.metadata import MetadataResourceNotFoundError @pytest.fixture @@ -56,7 +58,15 @@ def mock_dataset(): # --------------------------------------------------------------------------- -class TestDatasetMetadataCreatePost: +class _UsesSQLiteSession: + session: Session + + @pytest.fixture(autouse=True) + def _inject_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + + +class TestDatasetMetadataCreatePost(_UsesSQLiteSession): """Tests for DatasetMetadataCreateServiceApi.post(). ``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check`` @@ -64,7 +74,7 @@ class TestDatasetMetadataCreatePost: """ @staticmethod - def _call_post(api, session: MagicMock, **kwargs): + def _call_post(api, session: Session, **kwargs): return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -91,7 +101,7 @@ class TestDatasetMetadataCreatePost: json={"type": "string", "name": "Author"}, ): api = DatasetMetadataCreateServiceApi() - session = MagicMock() + session = self.session response, status = self._call_post( api, session, @@ -120,7 +130,7 @@ class TestDatasetMetadataCreatePost: json={"type": "string", "name": "Author"}, ): api = DatasetMetadataCreateServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_post( api, @@ -130,7 +140,7 @@ class TestDatasetMetadataCreatePost: ) -class TestDatasetMetadataCreateGet: +class TestDatasetMetadataCreateGet(_UsesSQLiteSession): """Tests for DatasetMetadataCreateServiceApi.get().""" @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -191,14 +201,14 @@ class TestDatasetMetadataCreateGet: # --------------------------------------------------------------------------- -class TestDatasetMetadataServiceApiPatch: +class TestDatasetMetadataServiceApiPatch(_UsesSQLiteSession): """Tests for DatasetMetadataServiceApi.patch(). ``patch`` is wrapped by ``@cloud_edition_billing_rate_limit_check``. """ @staticmethod - def _call_patch(api, session: MagicMock, **kwargs): + def _call_patch(api, session: Session, **kwargs): return unwrap(api.patch)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -215,7 +225,7 @@ class TestDatasetMetadataServiceApiPatch: ): """Test successful metadata name update.""" metadata_id = str(uuid.uuid4()) - mock_dataset_svc.get_dataset.return_value = mock_dataset + mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset mock_dataset_svc.check_dataset_permission.return_value = None mock_meta_svc.update_metadata_name.return_value = {"id": metadata_id, "type": "string", "name": "New Name"} @@ -225,7 +235,7 @@ class TestDatasetMetadataServiceApiPatch: json={"name": "New Name"}, ): api = DatasetMetadataServiceApi() - session = MagicMock() + session = self.session response, status = self._call_patch( api, session, @@ -236,7 +246,12 @@ class TestDatasetMetadataServiceApiPatch: assert status == 200 assert response == {"id": metadata_id, "type": "string", "name": "New Name"} - mock_meta_svc.update_metadata_name.assert_called_once() + mock_dataset_svc.get_dataset_for_tenant.assert_called_once_with( + str(mock_dataset.id), mock_tenant.id, session=session + ) + mock_meta_svc.update_metadata_name.assert_called_once_with( + mock_dataset, metadata_id, "New Name", mock_current_user, session=session + ) @patch("controllers.service_api.dataset.metadata.DatasetService") def test_update_metadata_dataset_not_found( @@ -248,7 +263,7 @@ class TestDatasetMetadataServiceApiPatch: ): """Test 404 when dataset not found.""" metadata_id = str(uuid.uuid4()) - mock_dataset_svc.get_dataset.return_value = None + mock_dataset_svc.get_dataset_for_tenant.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/metadata/{metadata_id}", @@ -256,7 +271,7 @@ class TestDatasetMetadataServiceApiPatch: json={"name": "x"}, ): api = DatasetMetadataServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_patch( api, @@ -267,14 +282,14 @@ class TestDatasetMetadataServiceApiPatch: ) -class TestDatasetMetadataServiceApiDelete: +class TestDatasetMetadataServiceApiDelete(_UsesSQLiteSession): """Tests for DatasetMetadataServiceApi.delete(). ``delete`` is wrapped by ``@cloud_edition_billing_rate_limit_check``. """ @staticmethod - def _call_delete(api, session: MagicMock, **kwargs): + def _call_delete(api, session: Session, **kwargs): return unwrap(api.delete)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -291,7 +306,7 @@ class TestDatasetMetadataServiceApiDelete: ): """Test successful metadata deletion.""" metadata_id = str(uuid.uuid4()) - mock_dataset_svc.get_dataset.return_value = mock_dataset + mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset mock_dataset_svc.check_dataset_permission.return_value = None mock_meta_svc.delete_metadata.return_value = None @@ -300,7 +315,7 @@ class TestDatasetMetadataServiceApiDelete: method="DELETE", ): api = DatasetMetadataServiceApi() - session = MagicMock() + session = self.session response = self._call_delete( api, session, @@ -310,7 +325,10 @@ class TestDatasetMetadataServiceApiDelete: ) assert response == ("", 204) - mock_meta_svc.delete_metadata.assert_called_once() + mock_dataset_svc.get_dataset_for_tenant.assert_called_once_with( + str(mock_dataset.id), mock_tenant.id, session=session + ) + mock_meta_svc.delete_metadata.assert_called_once_with(mock_dataset, metadata_id, session) @patch("controllers.service_api.dataset.metadata.DatasetService") def test_delete_metadata_dataset_not_found( @@ -322,14 +340,14 @@ class TestDatasetMetadataServiceApiDelete: ): """Test 404 when dataset not found.""" metadata_id = str(uuid.uuid4()) - mock_dataset_svc.get_dataset.return_value = None + mock_dataset_svc.get_dataset_for_tenant.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/metadata/{metadata_id}", method="DELETE", ): api = DatasetMetadataServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_delete( api, @@ -345,7 +363,7 @@ class TestDatasetMetadataServiceApiDelete: # --------------------------------------------------------------------------- -class TestDatasetMetadataBuiltInFieldGet: +class TestDatasetMetadataBuiltInFieldGet(_UsesSQLiteSession): """Tests for DatasetMetadataBuiltInFieldServiceApi.get().""" @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -380,14 +398,14 @@ class TestDatasetMetadataBuiltInFieldGet: # --------------------------------------------------------------------------- -class TestDatasetMetadataBuiltInFieldAction: +class TestDatasetMetadataBuiltInFieldAction(_UsesSQLiteSession): """Tests for DatasetMetadataBuiltInFieldActionServiceApi.post(). ``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check``. """ @staticmethod - def _call_post(api, session: MagicMock, **kwargs): + def _call_post(api, session: Session, **kwargs): return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -411,7 +429,7 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() - session = MagicMock() + session = self.session response, status = self._call_post( api, session, @@ -445,7 +463,7 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() - session = MagicMock() + session = self.session response, status = self._call_post( api, session, @@ -473,7 +491,7 @@ class TestDatasetMetadataBuiltInFieldAction: method="POST", ): api = DatasetMetadataBuiltInFieldActionServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_post( api, @@ -489,14 +507,14 @@ class TestDatasetMetadataBuiltInFieldAction: # --------------------------------------------------------------------------- -class TestDocumentMetadataEditPost: +class TestDocumentMetadataEditPost(_UsesSQLiteSession): """Tests for DocumentMetadataEditServiceApi.post(). ``post`` is wrapped by ``@cloud_edition_billing_rate_limit_check``. """ @staticmethod - def _call_post(api, session: MagicMock, **kwargs): + def _call_post(api, session: Session, **kwargs): return unwrap(api.post)(api, session, **kwargs) @patch("controllers.service_api.dataset.metadata.MetadataService") @@ -512,7 +530,7 @@ class TestDocumentMetadataEditPost: mock_dataset, ): """Test successful documents metadata update.""" - mock_dataset_svc.get_dataset.return_value = mock_dataset + mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset mock_dataset_svc.check_dataset_permission.return_value = None mock_meta_svc.update_documents_metadata.return_value = None @@ -522,7 +540,7 @@ class TestDocumentMetadataEditPost: json={"operation_data": []}, ): api = DocumentMetadataEditServiceApi() - session = MagicMock() + session = self.session response, status = self._call_post( api, session, @@ -532,6 +550,12 @@ class TestDocumentMetadataEditPost: assert status == 200 assert response["result"] == "success" + mock_meta_svc.update_documents_metadata.assert_called_once_with( + mock_dataset, + ANY, + mock_current_user, + session=session, + ) @patch("controllers.service_api.dataset.metadata.DatasetService") def test_update_documents_metadata_dataset_not_found( @@ -542,7 +566,7 @@ class TestDocumentMetadataEditPost: mock_dataset, ): """Test 404 when dataset not found.""" - mock_dataset_svc.get_dataset.return_value = None + mock_dataset_svc.get_dataset_for_tenant.return_value = None with app.test_request_context( f"/datasets/{mock_dataset.id}/documents/metadata", @@ -550,7 +574,7 @@ class TestDocumentMetadataEditPost: json={"operation_data": []}, ): api = DocumentMetadataEditServiceApi() - session = MagicMock() + session = self.session with pytest.raises(NotFound): self._call_post( api, @@ -558,3 +582,34 @@ class TestDocumentMetadataEditPost: tenant_id=mock_tenant.id, dataset_id=mock_dataset.id, ) + + @patch("controllers.service_api.dataset.metadata.MetadataService") + @patch("controllers.service_api.dataset.metadata.DatasetService") + @patch("controllers.service_api.dataset.metadata.current_user") + def test_update_documents_metadata_translates_missing_resource( + self, + mock_current_user, + mock_dataset_svc, + mock_meta_svc, + app: Flask, + mock_tenant, + mock_dataset, + ): + mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset + mock_meta_svc.update_documents_metadata.side_effect = MetadataResourceNotFoundError("Document not found.") + + with app.test_request_context( + f"/datasets/{mock_dataset.id}/documents/metadata", + method="POST", + json={"operation_data": []}, + ): + api = DocumentMetadataEditServiceApi() + with pytest.raises(NotFound) as exc_info: + self._call_post( + api, + MagicMock(), + tenant_id=mock_tenant.id, + dataset_id=mock_dataset.id, + ) + + assert exc_info.value.description == "Document not found." diff --git a/api/tests/unit_tests/controllers/service_api/end_user/test_end_user.py b/api/tests/unit_tests/controllers/service_api/end_user/test_end_user.py index 687c1a67b1a..30449f4b6b3 100644 --- a/api/tests/unit_tests/controllers/service_api/end_user/test_end_user.py +++ b/api/tests/unit_tests/controllers/service_api/end_user/test_end_user.py @@ -1,5 +1,4 @@ from datetime import UTC, datetime -from unittest.mock import Mock from uuid import UUID, uuid4 import pytest @@ -18,25 +17,27 @@ class TestEndUserApi: @pytest.fixture def app_model(self) -> App: - app = Mock(spec=App) - app.id = str(uuid4()) - app.tenant_id = str(uuid4()) + app = App( + id=str(uuid4()), + tenant_id=str(uuid4()), + ) return app def test_get_end_user_returns_all_attributes( self, mocker: MockerFixture, resource: EndUserApi, app_model: App ) -> None: - end_user = Mock(spec=EndUser) - end_user.id = str(uuid4()) - end_user.tenant_id = app_model.tenant_id - end_user.app_id = app_model.id - end_user.type = EndUserType.SERVICE_API - end_user.external_user_id = "external-123" - end_user.name = "Alice" - end_user._is_anonymous = True - end_user.session_id = "session-xyz" - end_user.created_at = datetime(2024, 1, 1, tzinfo=UTC) - end_user.updated_at = datetime(2024, 1, 2, tzinfo=UTC) + end_user = EndUser( + id=str(uuid4()), + tenant_id=app_model.tenant_id, + app_id=app_model.id, + type=EndUserType.SERVICE_API, + external_user_id="external-123", + name="Alice", + _is_anonymous=True, + session_id="session-xyz", + created_at=datetime(2024, 1, 1, tzinfo=UTC), + updated_at=datetime(2024, 1, 2, tzinfo=UTC), + ) get_end_user_by_id = mocker.patch( "controllers.service_api.end_user.end_user.EndUserService.get_end_user_by_id", return_value=end_user diff --git a/api/tests/unit_tests/controllers/service_api/test_wraps.py b/api/tests/unit_tests/controllers/service_api/test_wraps.py index 5857e5c639a..a6d502c6613 100644 --- a/api/tests/unit_tests/controllers/service_api/test_wraps.py +++ b/api/tests/unit_tests/controllers/service_api/test_wraps.py @@ -3,11 +3,14 @@ Unit tests for Service API wraps (authentication decorators) """ import uuid +from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask -from werkzeug.exceptions import Forbidden, NotFound, Unauthorized +from sqlalchemy import select +from sqlalchemy.orm import Session +from werkzeug.exceptions import Forbidden, NotFound, ServiceUnavailable, Unauthorized from controllers.service_api.wraps import ( DatasetApiResource, @@ -20,13 +23,12 @@ from controllers.service_api.wraps import ( validate_app_token, validate_dataset_token, ) -from enums.cloud_plan import CloudPlan -from models.account import TenantStatus -from models.model import ApiToken -from tests.unit_tests.conftest import ( - setup_mock_dataset_owner_execute_result, - setup_mock_tenant_owner_execute_result, -) +from enums import CloudPlan, DeploymentEdition +from models import Account, Tenant, TenantAccountJoin +from models.account import TenantAccountRole +from models.dataset import Dataset, RateLimitLog +from models.enums import ApiTokenType +from models.model import ApiToken, App, AppMode, IconType def _configure_current_app_mock(mock_current_app): @@ -34,6 +36,51 @@ def _configure_current_app_mock(mock_current_app): mock_current_app._get_current_object = Mock(return_value=Mock()) +def _session_proxy(session: Session) -> MagicMock: + """Emulate Flask-SQLAlchemy's callable scoped-session proxy around a test session.""" + proxy = MagicMock(wraps=session) + proxy.return_value = session + return proxy + + +def _api_token(*, tenant_id: str, app_id: str | None = None, token_type: ApiTokenType) -> ApiToken: + return ApiToken( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + app_id=app_id, + type=token_type, + token="test_token", + ) + + +def _persist_workspace(session: Session) -> tuple[Tenant, Account, TenantAccountJoin]: + tenant = Tenant(name="Workspace") + account = Account(name="Owner", email=f"owner-{uuid.uuid4()}@example.com") + membership = TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + current=True, + role=TenantAccountRole.OWNER, + ) + session.add_all([tenant, account, membership]) + session.commit() + return tenant, account, membership + + +def _app_model(*, tenant_id: str, enable_api: bool = True) -> App: + return App( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Service API App", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#FFFFFF", + enable_site=False, + enable_api=enable_api, + ) + + class TestValidateAndGetApiToken: """Test suite for validate_and_get_api_token function""" @@ -70,21 +117,24 @@ class TestValidateAndGetApiToken: def test_valid_token_returns_api_token(self, mock_fetch_token, mock_cache_cls, mock_record_usage, app: Flask): """Test that valid token returns the ApiToken object.""" # Arrange - mock_api_token = Mock(spec=ApiToken) - mock_api_token.token = "valid_token_123" - mock_api_token.type = "app" + api_token = _api_token( + tenant_id=str(uuid.uuid4()), + app_id=str(uuid.uuid4()), + token_type=ApiTokenType.APP, + ) + api_token.token = "valid_token_123" mock_cache_instance = Mock() mock_cache_instance.get.return_value = None # Cache miss mock_cache_cls.get = mock_cache_instance.get - mock_fetch_token.return_value = mock_api_token + mock_fetch_token.return_value = api_token # Act with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer valid_token_123"}): result = validate_and_get_api_token("app") # Assert - assert result == mock_api_token + assert result == api_token @patch("controllers.service_api.wraps.record_token_usage") @patch("controllers.service_api.wraps.ApiTokenCache") @@ -117,116 +167,124 @@ class TestValidateAppToken: return app @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") @patch("controllers.service_api.wraps.current_app") + @pytest.mark.parametrize( + "sqlite_session", + [(App, ApiToken, Tenant, Account, TenantAccountJoin)], + indirect=True, + ) def test_valid_app_token_allows_access( - self, mock_current_app, mock_validate_token, mock_db, mock_user_logged_in, app + self, + mock_current_app, + mock_validate_token, + mock_user_logged_in, + app: Flask, + sqlite_session: Session, ): """Test that valid app token allows access to decorated view.""" # Arrange _configure_current_app_mock(mock_current_app) - mock_api_token = Mock() - mock_api_token.app_id = str(uuid.uuid4()) - mock_api_token.tenant_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_app = Mock() - mock_app.id = mock_api_token.app_id - mock_app.status = "normal" - mock_app.enable_api = True - mock_app.tenant_id = mock_api_token.tenant_id - - mock_tenant = Mock() - mock_tenant.status = TenantStatus.NORMAL - mock_tenant.id = mock_api_token.tenant_id - - mock_account = Mock() - mock_account.id = str(uuid.uuid4()) - - # Use side_effect to return app first, then tenant via session.get() - mock_db.session.get.side_effect = [mock_app, mock_tenant] - - # Mock the tenant owner execute result (execute(select(...)).one_or_none()) - setup_mock_tenant_owner_execute_result(mock_db, mock_tenant, mock_account) + tenant, account, _ = _persist_workspace(sqlite_session) + app_model = _app_model(tenant_id=tenant.id) + api_token = _api_token(tenant_id=tenant.id, app_id=app_model.id, token_type=ApiTokenType.APP) + sqlite_session.add_all([app_model, api_token]) + sqlite_session.commit() + mock_validate_token.return_value = api_token @validate_app_token def protected_view(app_model): return {"success": True, "app_id": app_model.id} # Act - with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}): + with ( + app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}), + patch("controllers.service_api.wraps.db.session", _session_proxy(sqlite_session)), + ): result = protected_view() # Assert assert result["success"] is True - assert result["app_id"] == mock_app.id + assert result["app_id"] == app_model.id + assert account.current_tenant_id == tenant.id - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") - def test_app_not_found_raises_forbidden(self, mock_validate_token, mock_db, app: Flask): + @pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) + def test_app_not_found_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session): """Test that Forbidden is raised when app no longer exists.""" # Arrange - mock_api_token = Mock() - mock_api_token.app_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_db.session.get.return_value = None + api_token = _api_token( + tenant_id=str(uuid.uuid4()), + app_id=str(uuid.uuid4()), + token_type=ApiTokenType.APP, + ) + mock_validate_token.return_value = api_token @validate_app_token def protected_view(**kwargs): return {"success": True} # Act & Assert - with app.test_request_context("/", method="GET"): + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.db.session", sqlite_session), + ): with pytest.raises(Forbidden) as exc_info: protected_view() assert "no longer exists" in str(exc_info.value) - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") - def test_app_status_abnormal_raises_forbidden(self, mock_validate_token, mock_db, app: Flask): + @pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) + def test_app_status_abnormal_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session): """Test that Forbidden is raised when app status is abnormal.""" # Arrange - mock_api_token = Mock() - mock_api_token.app_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_app = Mock() - mock_app.status = "abnormal" - mock_db.session.get.return_value = mock_app + app_model = _app_model(tenant_id=str(uuid.uuid4())) + sqlite_session.add(app_model) + sqlite_session.commit() + app_model.status = "abnormal" + mock_validate_token.return_value = _api_token( + tenant_id=app_model.tenant_id, + app_id=app_model.id, + token_type=ApiTokenType.APP, + ) @validate_app_token def protected_view(**kwargs): return {"success": True} # Act & Assert - with app.test_request_context("/", method="GET"): + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.db.session", sqlite_session), + ): with pytest.raises(Forbidden) as exc_info: protected_view() assert "status is abnormal" in str(exc_info.value) - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") - def test_app_api_disabled_raises_forbidden(self, mock_validate_token, mock_db, app: Flask): + @pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) + def test_app_api_disabled_raises_forbidden(self, mock_validate_token, app: Flask, sqlite_session: Session): """Test that Forbidden is raised when app API is disabled.""" # Arrange - mock_api_token = Mock() - mock_api_token.app_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_app = Mock() - mock_app.status = "normal" - mock_app.enable_api = False - mock_db.session.get.return_value = mock_app + app_model = _app_model(tenant_id=str(uuid.uuid4()), enable_api=False) + sqlite_session.add(app_model) + sqlite_session.commit() + mock_validate_token.return_value = _api_token( + tenant_id=app_model.tenant_id, + app_id=app_model.id, + token_type=ApiTokenType.APP, + ) @validate_app_token def protected_view(**kwargs): return {"success": True} # Act & Assert - with app.test_request_context("/", method="GET"): + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.db.session", sqlite_session), + ): with pytest.raises(Forbidden) as exc_info: protected_view() assert "API service has been disabled" in str(exc_info.value) @@ -280,6 +338,7 @@ class TestCloudEditionBillingResourceCheck: mock_vector_space = Mock() mock_vector_space.limit = 10 mock_vector_space.size = 5 + mock_vector_space.usage_unknown = False mock_get_vector_space.return_value = mock_vector_space @cloud_edition_billing_resource_check("vector_space", "dataset") @@ -289,7 +348,7 @@ class TestCloudEditionBillingResourceCheck: # Act with ( app.test_request_context("/", method="GET"), - patch("controllers.service_api.wraps.dify_config.BILLING_ENABLED", True), + patch("controllers.service_api.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), ): result = add_segment() @@ -298,6 +357,64 @@ class TestCloudEditionBillingResourceCheck: mock_get_vector_space.assert_called_once_with("tenant123") mock_get_features.assert_not_called() + @patch("controllers.service_api.wraps.validate_and_get_api_token") + @patch("controllers.service_api.wraps.FeatureService.get_features") + @patch("controllers.service_api.wraps.FeatureService.get_vector_space") + def test_rejects_sandbox_when_vector_space_usage_is_unknown( + self, mock_get_vector_space, mock_get_features, mock_validate_token, app: Flask + ): + mock_validate_token.return_value = Mock(tenant_id="tenant123") + mock_get_vector_space.return_value = Mock(size=0, limit=50, usage_unknown=True) + mock_get_features.return_value = SimpleNamespace( + billing=SimpleNamespace( + enabled=True, + subscription=SimpleNamespace(plan=CloudPlan.SANDBOX), + ) + ) + + @cloud_edition_billing_resource_check("vector_space", "dataset") + def upload_document(): + return "document_uploaded" + + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + pytest.raises(ServiceUnavailable) as exc_info, + ): + upload_document() + + assert "Please try again later" in str(exc_info.value) + mock_get_features.assert_called_once_with("tenant123", exclude_vector_space=True) + + @patch("controllers.service_api.wraps.validate_and_get_api_token") + @patch("controllers.service_api.wraps.FeatureService.get_features") + @patch("controllers.service_api.wraps.FeatureService.get_vector_space") + @pytest.mark.parametrize("plan", [CloudPlan.PROFESSIONAL, CloudPlan.TEAM]) + def test_allows_paid_plan_when_vector_space_usage_is_unknown( + self, mock_get_vector_space, mock_get_features, mock_validate_token, app: Flask, plan: CloudPlan + ): + mock_validate_token.return_value = Mock(tenant_id="tenant123") + mock_get_vector_space.return_value = Mock(size=0, limit=50, usage_unknown=True) + mock_get_features.return_value = SimpleNamespace( + billing=SimpleNamespace( + enabled=True, + subscription=SimpleNamespace(plan=plan), + ) + ) + + @cloud_edition_billing_resource_check("vector_space", "dataset") + def upload_document(): + return "document_uploaded" + + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + ): + result = upload_document() + + assert result == "document_uploaded" + mock_get_features.assert_called_once_with("tenant123", exclude_vector_space=True) + @patch("controllers.service_api.wraps.validate_and_get_api_token") @patch("controllers.service_api.wraps.FeatureService.get_features") def test_loads_features_when_checking_non_vector_space_limit( @@ -468,26 +585,35 @@ class TestCloudEditionBillingRateLimitCheck: @patch("controllers.service_api.wraps.validate_and_get_api_token") @patch("controllers.service_api.wraps.FeatureService.get_knowledge_rate_limit") - @patch("controllers.service_api.wraps.db") - @patch("controllers.service_api.wraps.sessionmaker") + @pytest.mark.parametrize("sqlite_session", [(RateLimitLog,)], indirect=True) def test_rejects_over_rate_limit( - self, mock_sessionmaker, mock_db, mock_get_rate_limit, mock_validate_token, app: Flask + self, + mock_get_rate_limit, + mock_validate_token, + app: Flask, + sqlite_session: Session, ): """Test that Forbidden is raised when over rate limit.""" # Arrange - mock_validate_token.return_value = Mock(tenant_id="tenant123") + tenant_id = str(uuid.uuid4()) + mock_validate_token.return_value = _api_token( + tenant_id=tenant_id, + token_type=ApiTokenType.DATASET, + ) mock_rate_limit = Mock() mock_rate_limit.enabled = True mock_rate_limit.limit = 10 mock_rate_limit.subscription_plan = "pro" mock_get_rate_limit.return_value = mock_rate_limit - rate_limit_log_session = MagicMock() - session_factory = MagicMock() - session_factory.begin.return_value.__enter__.return_value = rate_limit_log_session - mock_sessionmaker.return_value = session_factory - with patch("controllers.service_api.wraps.redis_client") as mock_redis: + with ( + patch("controllers.service_api.wraps.redis_client") as mock_redis, + patch( + "controllers.service_api.wraps.db", + SimpleNamespace(engine=sqlite_session.get_bind()), + ), + ): mock_redis.zcard.return_value = 15 # Over limit @cloud_edition_billing_rate_limit_check("knowledge", "dataset") @@ -499,9 +625,12 @@ class TestCloudEditionBillingRateLimitCheck: with pytest.raises(Forbidden) as exc_info: knowledge_request() assert "rate limit" in str(exc_info.value) - mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) - rate_limit_log_session.add.assert_called_once() - mock_db.session.commit.assert_not_called() + + persisted_logs = sqlite_session.scalars(select(RateLimitLog)).all() + assert len(persisted_logs) == 1 + assert persisted_logs[0].tenant_id == tenant_id + assert persisted_logs[0].subscription_plan == "pro" + assert persisted_logs[0].operation == "knowledge" class TestValidateDatasetToken: @@ -515,65 +644,62 @@ class TestValidateDatasetToken: return app @patch("controllers.service_api.wraps.user_logged_in") - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") @patch("controllers.service_api.wraps.current_app") - def test_valid_dataset_token(self, mock_current_app, mock_validate_token, mock_db, mock_user_logged_in, app: Flask): + @pytest.mark.parametrize( + "sqlite_session", + [(Tenant, Account, TenantAccountJoin)], + indirect=True, + ) + def test_valid_dataset_token( + self, + mock_current_app, + mock_validate_token, + mock_user_logged_in, + app: Flask, + sqlite_session: Session, + ): """Test that valid dataset token allows access.""" # Arrange _configure_current_app_mock(mock_current_app) - tenant_id = str(uuid.uuid4()) - mock_api_token = Mock() - mock_api_token.tenant_id = tenant_id - mock_validate_token.return_value = mock_api_token - - mock_tenant = Mock() - mock_tenant.id = tenant_id - mock_tenant.status = TenantStatus.NORMAL - - mock_ta = Mock() - mock_ta.account_id = str(uuid.uuid4()) - - mock_account = Mock() - mock_account.id = mock_ta.account_id - mock_account.current_tenant = mock_tenant - - # Mock the tenant account join query (execute(select(...)).one_or_none()) - setup_mock_dataset_owner_execute_result(mock_db, mock_tenant, mock_ta) - - # Mock the account lookup via session.get() - mock_db.session.get.return_value = mock_account + tenant, account, _ = _persist_workspace(sqlite_session) + api_token = _api_token(tenant_id=tenant.id, token_type=ApiTokenType.DATASET) + mock_validate_token.return_value = api_token @validate_dataset_token def protected_view(tenant_id): return {"success": True, "tenant_id": tenant_id} # Act - with app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}): + with ( + app.test_request_context("/", method="GET", headers={"Authorization": "Bearer test_token"}), + patch("controllers.service_api.wraps.db.session", _session_proxy(sqlite_session)), + ): result = protected_view() # Assert assert result["success"] is True - assert result["tenant_id"] == tenant_id + assert result["tenant_id"] == tenant.id + assert account.current_tenant_id == tenant.id - @patch("controllers.service_api.wraps.db") @patch("controllers.service_api.wraps.validate_and_get_api_token") - def test_dataset_not_found_raises_not_found(self, mock_validate_token, mock_db, app: Flask): + @pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) + def test_dataset_not_found_raises_not_found(self, mock_validate_token, app: Flask, sqlite_session: Session): """Test that NotFound is raised when dataset doesn't exist.""" # Arrange - mock_api_token = Mock() - mock_api_token.tenant_id = str(uuid.uuid4()) - mock_validate_token.return_value = mock_api_token - - mock_db.session.scalar.return_value = None + api_token = _api_token(tenant_id=str(uuid.uuid4()), token_type=ApiTokenType.DATASET) + mock_validate_token.return_value = api_token @validate_dataset_token def protected_view(dataset_id=None, **kwargs): return {"success": True} # Act & Assert - with app.test_request_context("/", method="GET"): + with ( + app.test_request_context("/", method="GET"), + patch("controllers.service_api.wraps.db.session", sqlite_session), + ): with pytest.raises(NotFound) as exc_info: protected_view(dataset_id=str(uuid.uuid4())) assert "Dataset not found" in str(exc_info.value) diff --git a/api/tests/unit_tests/controllers/test_compare_versions.py b/api/tests/unit_tests/controllers/test_compare_versions.py index 9db57a84460..f4c296898aa 100644 --- a/api/tests/unit_tests/controllers/test_compare_versions.py +++ b/api/tests/unit_tests/controllers/test_compare_versions.py @@ -1,6 +1,6 @@ import pytest -from controllers.console.version import _has_new_version +from controllers.console.system import _has_new_version @pytest.mark.parametrize( @@ -20,5 +20,5 @@ from controllers.console.version import _has_new_version ("1.0.0", "1.0.0-dev", True), ], ) -def test_has_new_version(latest_version, current_version, expected): +def test_has_new_version(latest_version: str, current_version: str, expected: bool) -> None: assert _has_new_version(latest_version=latest_version, current_version=current_version) == expected diff --git a/api/tests/unit_tests/controllers/test_rbac_route_resource_contract.py b/api/tests/unit_tests/controllers/test_rbac_route_resource_contract.py new file mode 100644 index 00000000000..23cf107813d --- /dev/null +++ b/api/tests/unit_tests/controllers/test_rbac_route_resource_contract.py @@ -0,0 +1,102 @@ +"""Guard against resource-scoped RBAC gates mounted on routes that carry no resource id. + +``rbac_permission_required`` defaults to ``resource_required=True``, which makes +``_extract_resource_id`` raise ``ValueError`` when the matched path holds none of the +accepted identifiers. The request then fails with a 400 before the view ever runs, so the +endpoint is unreachable for every tenant with ``RBAC_ENABLED``. Creation endpoints and +other workspace-level actions must opt out with ``resource_required=False``. +""" + +import ast +from pathlib import Path + +CONTROLLERS_DIR = Path(__file__).resolve().parents[3] / "controllers" + +# Mirrors the lookup order in controllers/common/wraps.py::_extract_resource_id. +ACCEPTED_PATH_ARGS = { + "APP": ("app_id", "agent_id", "resource_id"), + "DATASET": ("dataset_id", "pipeline_id", "resource_id"), +} + +# Known violations tracked separately: DatasetDocumentSegmentBatchImportApi binds one class +# to both the dataset-scoped import route and the job-scoped status route, so every method +# on it is reachable at a path carrying only a job id. Its permission points are genuinely +# per-dataset, so it needs the route split rather than resource_required=False. Remove +# these entries with that fix. +KNOWN_VIOLATIONS = { + ("console/datasets/datasets_segments.py", "DatasetDocumentSegmentBatchImportApi", "post"), + ("console/datasets/datasets_segments.py", "DatasetDocumentSegmentBatchImportApi", "get"), +} + + +def _decorator_name(node: ast.Call) -> str: + func = node.func + parts = [] + while isinstance(func, ast.Attribute): + parts.append(func.attr) + func = func.value + if isinstance(func, ast.Name): + parts.append(func.id) + return ".".join(reversed(parts)) + + +def _attribute_name(node: ast.expr | None) -> str | None: + return node.attr if isinstance(node, ast.Attribute) else None + + +def _route_paths(class_node: ast.ClassDef) -> list[str]: + paths: list[str] = [] + for decorator in class_node.decorator_list: + if isinstance(decorator, ast.Call) and _decorator_name(decorator).endswith(".route"): + paths.extend( + arg.value for arg in decorator.args if isinstance(arg, ast.Constant) and isinstance(arg.value, str) + ) + return paths + + +def _path_args(route: str) -> set[str]: + args = set() + args.update(segment.split(">")[0].split(":")[-1] for segment in route.split("<")[1:]) + return args + + +def _resource_scoped_gates(method: ast.FunctionDef | ast.AsyncFunctionDef) -> list[str]: + scopes = [] + for decorator in method.decorator_list: + if not (isinstance(decorator, ast.Call) and _decorator_name(decorator).endswith("rbac_permission_required")): + continue + keywords = {keyword.arg: keyword.value for keyword in decorator.keywords} + resource_required = keywords.get("resource_required") + if isinstance(resource_required, ast.Constant) and resource_required.value is False: + continue + scope = _attribute_name(decorator.args[0] if decorator.args else keywords.get("resource_type")) + if scope in ACCEPTED_PATH_ARGS: + scopes.append(scope) + return scopes + + +def test_resource_scoped_rbac_gates_have_a_resource_id_in_the_route() -> None: + violations = [] + + for path in sorted(CONTROLLERS_DIR.rglob("*.py")): + tree = ast.parse(path.read_text(encoding="utf-8")) + for class_node in (node for node in ast.walk(tree) if isinstance(node, ast.ClassDef)): + routes = _route_paths(class_node) + if not routes: + continue + methods = (node for node in class_node.body if isinstance(node, ast.FunctionDef | ast.AsyncFunctionDef)) + for method in methods: + for scope in _resource_scoped_gates(method): + accepted = set(ACCEPTED_PATH_ARGS[scope]) + unscoped = [route for route in routes if not _path_args(route) & accepted] + if not unscoped: + continue + key = (path.relative_to(CONTROLLERS_DIR).as_posix(), class_node.name, method.name) + if key in KNOWN_VIOLATIONS: + continue + violations.append(f"{key[0]}::{key[1]}.{key[2]} scope={scope} routes={unscoped}") + + assert not violations, ( + "resource-scoped rbac_permission_required on routes without a resource id; " + "pass resource_required=False for workspace-level actions:\n" + "\n".join(violations) + ) diff --git a/api/tests/unit_tests/controllers/test_swagger.py b/api/tests/unit_tests/controllers/test_swagger.py index 149af9ff76a..fbca5ee600c 100644 --- a/api/tests/unit_tests/controllers/test_swagger.py +++ b/api/tests/unit_tests/controllers/test_swagger.py @@ -574,7 +574,7 @@ 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): +def test_console_agent_debug_conversation_refresh_has_no_body(monkeypatch: pytest.MonkeyPatch): from configs import dify_config from controllers.console import bp as console_bp @@ -586,13 +586,9 @@ def test_console_agent_debug_conversation_refresh_body_is_optional(monkeypatch: 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"] + assert "requestBody" not in operation + assert "AgentDebugConversationRefreshPayload" not in payload["components"]["schemas"] def test_console_member_invite_documents_bad_request_response(monkeypatch: pytest.MonkeyPatch): @@ -634,6 +630,11 @@ def test_console_plugin_category_list_exported_schema_uses_typed_items(tmp_path) console_openapi_path = next(path for path in written_paths if path.name == "console-openapi.json") payload = json.loads(console_openapi_path.read_text(encoding="utf-8")) operation = payload["paths"]["/workspaces/current/plugin/{category}/list"]["get"] + parameters = {parameter["name"]: parameter for parameter in operation["parameters"]} + assert parameters["query"]["in"] == "query" + assert parameters["tags"]["in"] == "query" + assert parameters["tags"]["schema"]["type"] == "array" + assert parameters["language"]["in"] == "query" response_ref = operation["responses"]["200"]["content"]["application/json"]["schema"]["$ref"].removeprefix( "#/components/schemas/" ) @@ -661,3 +662,114 @@ def test_console_plugin_category_list_exported_schema_uses_typed_items(tmp_path) builtin_tool_schema = schemas["PluginCategoryBuiltinToolProviderResponse"] for field in ("plugin_unique_identifier", "team_credentials", "type", "tools"): assert field in builtin_tool_schema["properties"] + + +def test_console_installed_plugin_ids_exported_schema_is_lightweight(tmp_path): + from dev.generate_swagger_specs import generate_specs + + written_paths = generate_specs(tmp_path) + console_openapi_path = next(path for path in written_paths if path.name == "console-openapi.json") + payload = json.loads(console_openapi_path.read_text(encoding="utf-8")) + operation = payload["paths"]["/workspaces/current/plugin/installed-ids"]["get"] + parameters = {parameter["name"]: parameter for parameter in operation["parameters"]} + assert parameters["category"]["in"] == "query" + assert parameters["category"]["required"] is True + assert parameters["category"]["schema"]["enum"] == [ + "agent-strategy", + "datasource", + "extension", + "model", + "tool", + "trigger", + ] + response_ref = operation["responses"]["200"]["content"]["application/json"]["schema"]["$ref"].removeprefix( + "#/components/schemas/" + ) + response_schema = payload["components"]["schemas"][response_ref] + + assert response_schema["required"] == ["plugin_ids"] + assert response_schema["properties"] == { + "plugin_ids": { + "items": {"type": "string"}, + "title": "Plugin Ids", + "type": "array", + } + } + + +def test_console_model_provider_summary_exported_schema_is_lightweight(tmp_path): + from dev.generate_swagger_specs import generate_specs + + written_paths = generate_specs(tmp_path) + console_openapi_path = next(path for path in written_paths if path.name == "console-openapi.json") + payload = json.loads(console_openapi_path.read_text(encoding="utf-8")) + operation = payload["paths"]["/workspaces/current/model-providers/summary"]["get"] + assert operation.get("parameters", []) == [] + + response_ref = operation["responses"]["200"]["content"]["application/json"]["schema"]["$ref"].removeprefix( + "#/components/schemas/" + ) + response_schema = payload["components"]["schemas"][response_ref] + assert response_schema["required"] == ["data", "plugins"] + assert response_schema["properties"]["data"]["items"]["$ref"] == ( + "#/components/schemas/ModelProviderSummaryResponse" + ) + assert response_schema["properties"]["plugins"]["additionalProperties"]["$ref"] == ( + "#/components/schemas/ModelProviderPluginSummaryResponse" + ) + + provider_properties = payload["components"]["schemas"]["ModelProviderSummaryResponse"]["properties"] + assert set(provider_properties) == { + "configurate_methods", + "custom_configuration", + "description", + "icon_small", + "icon_small_dark", + "is_configured", + "label", + "plugin_id", + "preferred_provider_type", + "provider", + "supported_model_types", + "system_configuration", + } + assert "provider_credential_schema" not in provider_properties + assert "model_credential_schema" not in provider_properties + + custom_configuration_schema = payload["components"]["schemas"]["ModelProviderCustomConfigurationSummaryResponse"] + custom_configuration_properties = custom_configuration_schema["properties"] + assert set(custom_configuration_schema["required"]) == { + "available_credentials", + "current_credential_usable", + "has_custom_models", + "status", + } + assert set(custom_configuration_properties) == { + "available_credentials", + "current_credential_id", + "current_credential_name", + "current_credential_usable", + "has_custom_models", + "status", + } + assert custom_configuration_properties["available_credentials"]["items"]["$ref"] == ( + "#/components/schemas/CredentialConfiguration" + ) + assert "has_credentials" not in custom_configuration_properties + + credential_properties = payload["components"]["schemas"]["CredentialConfiguration"]["properties"] + assert set(credential_properties) == { + "credential_id", + "credential_name", + } + assert "encrypted_config" not in credential_properties + + plugin_properties = payload["components"]["schemas"]["ModelProviderPluginSummaryResponse"]["properties"] + assert set(plugin_properties) == { + "installation_id", + "plugin_id", + "plugin_unique_identifier", + "runtime_type", + "source", + "version", + } diff --git a/api/tests/unit_tests/controllers/web/conftest.py b/api/tests/unit_tests/controllers/web/conftest.py index b7f3244c6cd..a7198334a07 100644 --- a/api/tests/unit_tests/controllers/web/conftest.py +++ b/api/tests/unit_tests/controllers/web/conftest.py @@ -2,9 +2,6 @@ from __future__ import annotations -from types import SimpleNamespace -from typing import Any - import pytest from flask import Flask @@ -15,69 +12,3 @@ def app() -> Flask: flask_app = Flask(__name__) flask_app.config["TESTING"] = True return flask_app - - -class FakeSession: - """Stand-in for db.session that returns pre-seeded objects by model class name.""" - - def __init__(self, mapping: dict[str, Any] | None = None): - self._mapping: dict[str, Any] = mapping or {} - - def get(self, model: type, _ident: object) -> Any: - return self._mapping.get(model.__name__) - - def scalar(self, stmt: Any) -> Any: - try: - model = stmt.column_descriptions[0]["entity"] - except (AttributeError, IndexError, KeyError, TypeError): - return None - return self._mapping.get(model.__name__) - - -class FakeDB: - """Minimal db stub exposing engine and session.""" - - def __init__(self, session: FakeSession | None = None): - self.session = session or FakeSession() - self.engine = object() - - -def make_app_model( - *, - app_id: str = "app-1", - tenant_id: str = "tenant-1", - mode: str = "chat", - enable_site: bool = True, - status: str = "normal", -) -> SimpleNamespace: - """Build a fake App model with common defaults.""" - tenant = SimpleNamespace( - id=tenant_id, - status="normal", - plan="basic", - custom_config_dict={}, - ) - return SimpleNamespace( - id=app_id, - tenant_id=tenant_id, - tenant=tenant, - mode=mode, - enable_site=enable_site, - status=status, - workflow=None, - app_model_config=None, - ) - - -def make_end_user( - *, - user_id: str = "end-user-1", - session_id: str = "session-1", - external_user_id: str = "ext-user-1", -) -> SimpleNamespace: - """Build a fake EndUser model with common defaults.""" - return SimpleNamespace( - id=user_id, - session_id=session_id, - external_user_id=external_user_id, - ) diff --git a/api/tests/unit_tests/controllers/web/test_completion.py b/api/tests/unit_tests/controllers/web/test_completion.py index e3bbbe2c87c..e76a4d7d6a4 100644 --- a/api/tests/unit_tests/controllers/web/test_completion.py +++ b/api/tests/unit_tests/controllers/web/test_completion.py @@ -9,6 +9,7 @@ from unittest.mock import MagicMock, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from controllers.web.completion import ChatApi, ChatStopApi, CompletionApi, CompletionStopApi from controllers.web.error import ( @@ -168,9 +169,10 @@ class TestChatApi: mock_get_conversation: MagicMock, mock_generate: MagicMock, app: Flask, + unbound_session: Session, ) -> None: mock_ns.payload = {"inputs": {}, "query": "hi", "conversation_id": str(uuid.uuid4())} - session = MagicMock() + session = unbound_session with app.test_request_context("/chat-messages", method="POST"): unwrap(ChatApi.post)(ChatApi(), session, _chat_app(), _end_user()) diff --git a/api/tests/unit_tests/controllers/web/test_feature.py b/api/tests/unit_tests/controllers/web/test_feature.py index c4701119eff..68b8e72be71 100644 --- a/api/tests/unit_tests/controllers/web/test_feature.py +++ b/api/tests/unit_tests/controllers/web/test_feature.py @@ -2,32 +2,48 @@ from __future__ import annotations -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, create_autospec +import pytest from flask import Flask +from pytest_mock import MockerFixture from controllers.web.feature import SystemFeatureApi -from enums.deployment_edition import DeploymentEdition -from services.feature_service import SystemFeatureModel +from enums import DeploymentEdition +from services.entities.feature_entities import SystemFeatureModel +from services.feature_query_service import FeatureQueryService + + +def _install_feature_queries(mocker: MockerFixture) -> MagicMock: + feature_queries = create_autospec(FeatureQueryService, instance=True, spec_set=True) + application_services = mocker.patch("controllers.web.feature.application_services") + application_services.return_value.feature_queries = feature_queries + return feature_queries class TestSystemFeatureApi: - @patch("controllers.web.feature.FeatureService.get_system_features") - def test_returns_system_features(self, mock_features: MagicMock, app: Flask) -> None: - system_features = SystemFeatureModel(deployment_edition=DeploymentEdition.COMMUNITY) - mock_features.return_value = system_features + @pytest.mark.parametrize("deployment_edition", list(DeploymentEdition)) + def test_returns_system_features( + self, + deployment_edition: DeploymentEdition, + app: Flask, + mocker: MockerFixture, + ) -> None: + system_features = SystemFeatureModel(deployment_edition=deployment_edition) + feature_queries = _install_feature_queries(mocker) + feature_queries.get_system_features.return_value = system_features with app.test_request_context("/system-features"): result = SystemFeatureApi().get() - assert result == system_features.model_dump() - mock_features.assert_called_once() + assert result == system_features.model_dump(mode="json") + assert result["deployment_edition"] == deployment_edition.value + assert result["sso_enforced_for_signin_protocol"] is None + assert result["webapp_auth"]["sso_config"]["protocol"] is None + feature_queries.get_system_features.assert_called_once_with() - @patch("controllers.web.feature.FeatureService.get_system_features") - def test_unauthenticated_access(self, mock_features: MagicMock, app: Flask) -> None: + def test_unauthenticated_access(self) -> None: """SystemFeatureApi is unauthenticated by design — no WebApiResource decorator.""" - mock_features.return_value = SystemFeatureModel(deployment_edition=DeploymentEdition.COMMUNITY) - # Verify it's a bare Resource, not WebApiResource from flask_restx import Resource diff --git a/api/tests/unit_tests/controllers/web/test_human_input_form.py b/api/tests/unit_tests/controllers/web/test_human_input_form.py index 3408e3049d1..985c1e41ba9 100644 --- a/api/tests/unit_tests/controllers/web/test_human_input_form.py +++ b/api/tests/unit_tests/controllers/web/test_human_input_form.py @@ -22,7 +22,7 @@ from models import Tenant from models.enums import CustomizeTokenStrategy from models.human_input import RecipientType from models.model import App, AppMode, IconType, Site -from services.feature_service import FeatureModel +from services.entities.feature_entities import FeatureModel from services.human_input_service import FormExpiredError HumanInputFormApi = human_input_module.HumanInputFormApi @@ -163,6 +163,7 @@ def test_get_form_includes_site(monkeypatch: pytest.MonkeyPatch, app: Flask, dat assert body["expiration_time"] == int(expiration_time.timestamp()) assert body["site"] == { "app_id": app_model.id, + "mode": "chat", "end_user_id": None, "enable_site": True, "site": { @@ -383,6 +384,7 @@ def test_get_form_allows_backstage_token(monkeypatch: pytest.MonkeyPatch, app: F assert body["expiration_time"] == int(expiration_time.timestamp()) assert body["site"] == { "app_id": app_model.id, + "mode": "chat", "end_user_id": None, "enable_site": True, "site": { diff --git a/api/tests/unit_tests/controllers/web/test_site.py b/api/tests/unit_tests/controllers/web/test_site.py index 1c2a403994f..011d6b6a51e 100644 --- a/api/tests/unit_tests/controllers/web/test_site.py +++ b/api/tests/unit_tests/controllers/web/test_site.py @@ -2,17 +2,54 @@ from unittest.mock import MagicMock, patch from configs import dify_config from controllers.web import site as site_module +from enums import DeploymentEdition from extensions.storage.storage_type import StorageType -from models.model import IconType, Site +from models.model import AppMode, IconType, Site +from services.entities.feature_entities import FeatureModel + + +def test_app_site_api_returns_legacy_agent_compatible_mode() -> None: + app_model = MagicMock() + app_model.id = "app-id" + app_model.tenant_id = "tenant-id" + app_model.tenant = MagicMock(id="tenant-id", status="normal") + app_model.mode_compatible_with_agent_with_session.return_value = AppMode.AGENT_CHAT + end_user = MagicMock(id="end-user-id") + site = Site() + response = MagicMock() + response.model_dump.return_value = {"mode": AppMode.AGENT_CHAT} + + with ( + patch.object(site_module, "db") as mock_db, + patch.object(site_module.FeatureService, "get_features", return_value=FeatureModel(can_replace_logo=False)), + patch.object(site_module, "_build_site_icon_url", return_value=None), + patch.object(site_module.WebAppSiteResponse, "from_app_site", return_value=response) as mock_from_app_site, + ): + mock_db.session.scalar.return_value = site + result = site_module.AppSiteApi().get(app_model, end_user) + + assert result["mode"] == AppMode.AGENT_CHAT + app_model.mode_compatible_with_agent_with_session.assert_called_once_with(session=mock_db.session()) + mock_from_app_site.assert_called_once_with( + tenant=app_model.tenant, + app_model=app_model, + mode=AppMode.AGENT_CHAT, + site=site, + end_user_id=end_user.id, + features=FeatureModel(can_replace_logo=False), + can_replace_logo=False, + icon_url=None, + ) 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" + site = Site( + icon_type=IconType.IMAGE, + icon="11111111-1111-4111-8111-111111111111", + ) with ( - patch.object(dify_config, "EDITION", "CLOUD"), + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.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, @@ -34,12 +71,13 @@ def test_build_site_icon_url_uses_s3_presigned_url() -> None: 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" + site = Site( + icon_type=IconType.IMAGE, + icon="11111111-1111-4111-8111-111111111111", + ) with ( - patch.object(dify_config, "EDITION", "SELF_HOSTED"), + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), 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"), @@ -51,12 +89,13 @@ def test_build_site_icon_url_keeps_preview_url_for_self_hosted_s3() -> None: 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" + site = Site( + icon_type=IconType.IMAGE, + icon="11111111-1111-4111-8111-111111111111", + ) with ( - patch.object(dify_config, "EDITION", "CLOUD"), + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.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"), diff --git a/api/tests/unit_tests/controllers/web/test_web_forgot_password.py b/api/tests/unit_tests/controllers/web/test_web_forgot_password.py index 0d6fa621276..f2588ec3450 100644 --- a/api/tests/unit_tests/controllers/web/test_web_forgot_password.py +++ b/api/tests/unit_tests/controllers/web/test_web_forgot_password.py @@ -14,10 +14,10 @@ from controllers.web.forgot_password import ( ForgotPasswordResetApi, ForgotPasswordSendEmailApi, ) -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from models.account import Account from models.engine import db -from services.feature_service import SystemFeatureModel +from services.entities.feature_entities import SystemFeatureModel @pytest.fixture @@ -39,8 +39,7 @@ def _patch_wraps(): ) with ( patch("controllers.console.wraps.db") as mock_db, - patch("controllers.console.wraps.dify_config.ENTERPRISE_ENABLED", True), - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features), ): yield diff --git a/api/tests/unit_tests/controllers/web/test_web_login.py b/api/tests/unit_tests/controllers/web/test_web_login.py index 2ccb6b832fc..b4dae61bd91 100644 --- a/api/tests/unit_tests/controllers/web/test_web_login.py +++ b/api/tests/unit_tests/controllers/web/test_web_login.py @@ -6,13 +6,19 @@ from unittest.mock import ANY, MagicMock, patch import pytest from flask import Flask from jwt import InvalidTokenError +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker from werkzeug.exceptions import Unauthorized import services.errors.account +from controllers.console import wraps as console_wraps from controllers.web.login import EmailCodeLoginApi, EmailCodeLoginSendEmailApi, LoginApi, LoginStatusApi, LogoutApi -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition +from models.model import DifySetup from services.entities.auth_entities import LoginFailureReason +pytestmark = pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) + def encode_code(code: str) -> str: return base64.b64encode(code.encode("utf-8")).decode() @@ -33,17 +39,27 @@ def app(): @pytest.fixture(autouse=True) -def _patch_wraps(): +def _patch_wraps( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, +): wraps_features = SimpleNamespace(enable_email_password_login=True) - console_dify = SimpleNamespace(ENTERPRISE_ENABLED=True, DEPLOYMENT_EDITION=DeploymentEdition.CLOUD) - web_dify = SimpleNamespace(ENTERPRISE_ENABLED=True) + console_dify = SimpleNamespace(DEPLOYMENT_EDITION=DeploymentEdition.ENTERPRISE) + web_dify = SimpleNamespace(DEPLOYMENT_EDITION=DeploymentEdition.ENTERPRISE) + sqlite_session.add(DifySetup(version="test")) + sqlite_session.commit() + console_wraps._is_setup_completed.reset_success() + session_registry = scoped_session(sessionmaker(bind=sqlite_engine, expire_on_commit=False)) + monkeypatch.setattr(console_wraps.db, "session", session_registry) with ( - patch("controllers.console.wraps.db") as mock_db, patch("controllers.console.wraps.dify_config", console_dify), patch("controllers.console.wraps.FeatureService.get_system_features", return_value=wraps_features), patch("controllers.web.login.dify_config", web_dify), ): yield + session_registry.remove() + console_wraps._is_setup_completed.reset_success() class TestEmailCodeLoginSendEmailApi: diff --git a/api/tests/unit_tests/controllers/web/test_web_passport.py b/api/tests/unit_tests/controllers/web/test_web_passport.py index 4e1a24a4da2..82ec7f1bd44 100644 --- a/api/tests/unit_tests/controllers/web/test_web_passport.py +++ b/api/tests/unit_tests/controllers/web/test_web_passport.py @@ -2,11 +2,14 @@ from __future__ import annotations +import uuid from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest from flask import Flask +from sqlalchemy import Engine, select +from sqlalchemy.orm import Session from werkzeug.exceptions import NotFound, Unauthorized from controllers.web.error import WebAppAuthRequiredError @@ -16,9 +19,62 @@ from controllers.web.passport import ( exchange_token_for_existing_web_user, generate_session_id, ) +from models.base import TypeBase +from models.enums import CustomizeTokenStrategy, EndUserType +from models.model import App, AppMode, EndUser, IconType, Site from services.webapp_auth_service import WebAppAuthType +@pytest.fixture +def database_session(sqlite_engine: Engine): + models = (App, Site, EndUser) + tables = [model.metadata.tables[model.__tablename__] for model in models] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + with Session(sqlite_engine, expire_on_commit=False) as session: + with patch("controllers.web.passport.db.session", session): + yield session + + +def _persist_webapp( + session: Session, + *, + app_code: str = "code1", + enable_site: bool = True, +) -> tuple[App, Site]: + app_model = App( + id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + name="Web App", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#FFFFFF", + enable_site=enable_site, + enable_api=False, + ) + site = Site( + app_id=app_model.id, + title="Web App Site", + default_language="en-US", + customize_token_strategy=CustomizeTokenStrategy.UUID, + code=app_code, + ) + session.add_all([app_model, site]) + session.commit() + return app_model, site + + +def _end_user(app_model: App, *, session_id: str) -> EndUser: + return EndUser( + id=str(uuid.uuid4()), + tenant_id=app_model.tenant_id, + app_id=app_model.id, + type=EndUserType.BROWSER, + name="Web User", + session_id=session_id, + ) + + # --------------------------------------------------------------------------- # decode_enterprise_webapp_user_id # --------------------------------------------------------------------------- @@ -55,56 +111,48 @@ class TestDecodeEnterpriseWebappUserId: # generate_session_id # --------------------------------------------------------------------------- class TestGenerateSessionId: - @patch("controllers.web.passport.db") - def test_returns_unique_session_id(self, mock_db: MagicMock) -> None: - mock_db.session.scalar.return_value = 0 + def test_returns_unique_session_id(self, database_session: Session) -> None: sid = generate_session_id() assert isinstance(sid, str) assert len(sid) == 36 # UUID format - @patch("controllers.web.passport.db") - def test_retries_on_collision(self, mock_db: MagicMock) -> None: - # First call returns count=1 (collision), second returns 0 - mock_db.session.scalar.side_effect = [1, 0] - sid = generate_session_id() - assert isinstance(sid, str) - assert mock_db.session.scalar.call_count == 2 + def test_retries_on_collision(self, database_session: Session) -> None: + app_model, _ = _persist_webapp(database_session) + collision_id = str(uuid.uuid4()) + generated_id = str(uuid.uuid4()) + database_session.add(_end_user(app_model, session_id=collision_id)) + database_session.commit() + + with patch( + "controllers.web.passport.uuid.uuid4", + side_effect=[uuid.UUID(collision_id), uuid.UUID(generated_id)], + ): + sid = generate_session_id() + + assert sid == generated_id # --------------------------------------------------------------------------- # exchange_token_for_existing_web_user # --------------------------------------------------------------------------- class TestExchangeTokenForExistingWebUser: - @patch("controllers.web.passport.PassportService") - @patch("controllers.web.passport.db") - def test_external_auth_type_mismatch_raises(self, mock_db: MagicMock, mock_passport_cls: MagicMock) -> None: - site = SimpleNamespace(code="code1", app_id="app-1") - app_model = SimpleNamespace(id="app-1", status="normal", enable_site=True, tenant_id="t1") - mock_db.session.scalar.side_effect = [site, app_model] - + def test_external_auth_type_mismatch_raises(self, database_session: Session) -> None: + _persist_webapp(database_session) decoded = {"user_id": "u1", "auth_type": "internal"} # mismatch: expected "external" with pytest.raises(WebAppAuthRequiredError, match="external"): exchange_token_for_existing_web_user( app_code="code1", enterprise_user_decoded=decoded, auth_type=WebAppAuthType.EXTERNAL ) - @patch("controllers.web.passport.PassportService") - @patch("controllers.web.passport.db") - def test_internal_auth_type_mismatch_raises(self, mock_db: MagicMock, mock_passport_cls: MagicMock) -> None: - site = SimpleNamespace(code="code1", app_id="app-1") - app_model = SimpleNamespace(id="app-1", status="normal", enable_site=True, tenant_id="t1") - mock_db.session.scalar.side_effect = [site, app_model] - + def test_internal_auth_type_mismatch_raises(self, database_session: Session) -> None: + _persist_webapp(database_session) decoded = {"user_id": "u1", "auth_type": "external"} # mismatch: expected "internal" with pytest.raises(WebAppAuthRequiredError, match="internal"): exchange_token_for_existing_web_user( app_code="code1", enterprise_user_decoded=decoded, auth_type=WebAppAuthType.INTERNAL ) - @patch("controllers.web.passport.PassportService") - @patch("controllers.web.passport.db") - def test_site_not_found_raises(self, mock_db: MagicMock, mock_passport_cls: MagicMock) -> None: - mock_db.session.scalar.return_value = None + def test_site_not_found_raises(self, database_session: Session) -> None: decoded = {"user_id": "u1", "auth_type": "external"} with pytest.raises(NotFound): exchange_token_for_existing_web_user( @@ -125,69 +173,68 @@ class TestPassportResource: @patch("controllers.web.passport.PassportService") @patch("controllers.web.passport.generate_session_id", return_value="new-sess-id") - @patch("controllers.web.passport.db") @patch("controllers.web.passport.FeatureService.get_system_features") def test_creates_new_end_user_when_no_user_id( self, mock_features: MagicMock, - mock_db: MagicMock, mock_gen_session: MagicMock, mock_passport_cls: MagicMock, app: Flask, + database_session: Session, + sqlite_engine: Engine, ) -> None: mock_features.return_value = SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)) - site = SimpleNamespace(app_id="app-1", code="code1") - app_model = SimpleNamespace(id="app-1", status="normal", enable_site=True, tenant_id="t1") - mock_db.session.scalar.side_effect = [site, app_model] + app_model, _ = _persist_webapp(database_session) mock_passport_cls.return_value.issue.return_value = "issued-token" with app.test_request_context("/passport", headers={"X-App-Code": "code1"}): response = PassportResource().get() assert response["access_token"] == "issued-token" - mock_db.session.add.assert_called_once() - mock_db.session.commit.assert_called_once() + database_session.close() + with Session(sqlite_engine) as verification_session: + end_users = verification_session.scalars(select(EndUser)).all() + assert len(end_users) == 1 + assert end_users[0].session_id == "new-sess-id" + assert end_users[0].app_id == app_model.id + assert end_users[0].tenant_id == app_model.tenant_id @patch("controllers.web.passport.PassportService") - @patch("controllers.web.passport.db") @patch("controllers.web.passport.FeatureService.get_system_features") def test_reuses_existing_end_user_when_user_id_provided( self, mock_features: MagicMock, - mock_db: MagicMock, mock_passport_cls: MagicMock, app: Flask, + database_session: Session, ) -> None: mock_features.return_value = SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)) - site = SimpleNamespace(app_id="app-1", code="code1") - app_model = SimpleNamespace(id="app-1", status="normal", enable_site=True, tenant_id="t1") - existing_user = SimpleNamespace(id="eu-1", session_id="sess-existing") - mock_db.session.scalar.side_effect = [site, app_model, existing_user] + app_model, _ = _persist_webapp(database_session) + existing_user = _end_user(app_model, session_id="sess-existing") + database_session.add(existing_user) + database_session.commit() mock_passport_cls.return_value.issue.return_value = "reused-token" with app.test_request_context("/passport?user_id=sess-existing", headers={"X-App-Code": "code1"}): response = PassportResource().get() assert response["access_token"] == "reused-token" - # Should not create a new end user - mock_db.session.add.assert_not_called() + end_users = database_session.scalars(select(EndUser)).all() + assert [end_user.id for end_user in end_users] == [existing_user.id] - @patch("controllers.web.passport.db") @patch("controllers.web.passport.FeatureService.get_system_features") - def test_site_not_found_raises(self, mock_features: MagicMock, mock_db: MagicMock, app: Flask) -> None: + def test_site_not_found_raises(self, mock_features: MagicMock, app: Flask, database_session: Session) -> None: mock_features.return_value = SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)) - mock_db.session.scalar.return_value = None with app.test_request_context("/passport", headers={"X-App-Code": "code1"}): with pytest.raises(NotFound): PassportResource().get() - @patch("controllers.web.passport.db") @patch("controllers.web.passport.FeatureService.get_system_features") - def test_disabled_app_raises_not_found(self, mock_features: MagicMock, mock_db: MagicMock, app: Flask) -> None: + def test_disabled_app_raises_not_found( + self, mock_features: MagicMock, app: Flask, database_session: Session + ) -> None: mock_features.return_value = SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)) - site = SimpleNamespace(app_id="app-1", code="code1") - disabled_app = SimpleNamespace(id="app-1", status="normal", enable_site=False) - mock_db.session.scalar.side_effect = [site, disabled_app] + _persist_webapp(database_session, enable_site=False) with app.test_request_context("/passport", headers={"X-App-Code": "code1"}): with pytest.raises(NotFound): PassportResource().get() diff --git a/api/tests/unit_tests/controllers/web/test_workflow_events.py b/api/tests/unit_tests/controllers/web/test_workflow_events.py index ab2ca7dde67..d9eb5fc7e73 100644 --- a/api/tests/unit_tests/controllers/web/test_workflow_events.py +++ b/api/tests/unit_tests/controllers/web/test_workflow_events.py @@ -3,7 +3,7 @@ from __future__ import annotations from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import ANY, MagicMock, Mock, patch import pytest from flask import Flask @@ -11,6 +11,7 @@ from flask import Flask from controllers.common.errors import NotFoundError from controllers.web.workflow_events import WorkflowEventsApi from models.enums import CreatorUserRole +from models.model import AppMode def _workflow_app() -> SimpleNamespace: @@ -125,3 +126,39 @@ class TestWorkflowEventsApi: response = WorkflowEventsApi().get(_workflow_app(), _end_user(), "run-1") assert response.mimetype == "text/event-stream" + + @patch("controllers.web.workflow_events.DifyAPIRepositoryFactory") + @patch("controllers.web.workflow_events.db") + def test_snapshot_stream_can_continue_across_pauses( + self, mock_db: MagicMock, mock_factory: MagicMock, app: Flask, monkeypatch: pytest.MonkeyPatch + ) -> None: + mock_db.engine = "engine" + run = SimpleNamespace( + id="run-1", + app_id="app-1", + created_by_role=CreatorUserRole.END_USER, + created_by="eu-1", + finished_at=None, + ) + mock_repo = MagicMock() + mock_repo.get_workflow_run_by_id_and_tenant_id.return_value = run + mock_factory.create_api_workflow_run_repository.return_value = mock_repo + + workflow_generator = Mock() + workflow_generator.convert_to_event_stream.return_value = iter(["data: snapshot\n\n"]) + snapshot_builder = Mock(return_value=["snapshot-events"]) + monkeypatch.setattr("controllers.web.workflow_events.WorkflowAppGenerator", lambda: workflow_generator) + monkeypatch.setattr("controllers.web.workflow_events.build_workflow_event_stream", snapshot_builder) + + with app.test_request_context("/workflow/run-1/events?include_state_snapshot=true&continue_on_pause=true"): + response = WorkflowEventsApi().get(_workflow_app(), _end_user(), "run-1") + + assert response.get_data(as_text=True) == "data: snapshot\n\n" + snapshot_builder.assert_called_once_with( + app_mode=AppMode.WORKFLOW, + workflow_run=run, + tenant_id="tenant-1", + app_id="app-1", + session_maker=ANY, + close_on_pause=False, + ) diff --git a/api/tests/unit_tests/core/agent/strategy/test_plugin.py b/api/tests/unit_tests/core/agent/strategy/test_plugin.py index 15c441e5876..55d4ffcb263 100644 --- a/api/tests/unit_tests/core/agent/strategy/test_plugin.py +++ b/api/tests/unit_tests/core/agent/strategy/test_plugin.py @@ -5,7 +5,9 @@ from unittest.mock import MagicMock import pytest from pytest_mock import MockerFixture +from core.agent.plugin_entities import AgentStrategyParameter from core.agent.strategy.plugin import PluginAgentStrategy +from core.tools.entities.common_entities import I18nObject # ============================================================ # Fixtures @@ -103,6 +105,17 @@ class TestInitializeParameters: mock_declaration.parameters[0].init_frontend_parameter.assert_called_once_with("value1") mock_declaration.parameters[1].init_frontend_parameter.assert_called_once_with(None) + def test_initialize_parameters_allows_empty_tools_selection(self, strategy: PluginAgentStrategy) -> None: + tools_parameter = AgentStrategyParameter( + name="tools", + label=I18nObject(en_US="Tools"), + required=True, + type=AgentStrategyParameter.AgentStrategyParameterType.TOOLS_SELECTOR, + ) + strategy.declaration.parameters = [tools_parameter] + + assert strategy.initialize_parameters({"tools": []}) == {"tools": []} + @pytest.mark.parametrize( "input_params", [ diff --git a/api/tests/unit_tests/core/agent/test_cot_agent_runner.py b/api/tests/unit_tests/core/agent/test_cot_agent_runner.py index d3cdc6ab292..4f3c7bbc3b5 100644 --- a/api/tests/unit_tests/core/agent/test_cot_agent_runner.py +++ b/api/tests/unit_tests/core/agent/test_cot_agent_runner.py @@ -344,9 +344,20 @@ class TestRun: message = MagicMock() message.id = "msg-id" events: list[str] = [] - session = MagicMock() - session.commit.side_effect = lambda: events.append("commit") - session.close.side_effect = lambda: events.append("close") + session = runner.session + original_close = session.close + original_commit = session.commit + + def close_session() -> None: + events.append("close") + original_close() + + def commit_session() -> None: + events.append("commit") + original_commit() + + mocker.patch.object(session, "close", side_effect=close_session) + mocker.patch.object(session, "commit", side_effect=commit_session) def provider_chunks(): events.append("first-chunk") diff --git a/api/tests/unit_tests/core/agent/test_fc_agent_runner.py b/api/tests/unit_tests/core/agent/test_fc_agent_runner.py index 9ce87271bf7..ea8c3b87c42 100644 --- a/api/tests/unit_tests/core/agent/test_fc_agent_runner.py +++ b/api/tests/unit_tests/core/agent/test_fc_agent_runner.py @@ -1,5 +1,6 @@ import json from collections.abc import Iterator +from datetime import UTC, datetime from typing import Any from unittest.mock import MagicMock @@ -16,9 +17,12 @@ from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.model_runtime.entities.message_entities import ( DocumentPromptMessageContent, ImagePromptMessageContent, + PromptMessageContentType, TextPromptMessageContent, UserPromptMessage, ) +from models.enums import CreatorUserRole +from models.model import StorageType, UploadFile # ============================== # Dummy Helper Classes @@ -133,6 +137,7 @@ def runner(mocker: MockerFixture, sqlite_engine: Engine) -> Iterator[FunctionCal runner.history_prompt_messages = [] runner._current_thoughts = [] runner.files = [] + runner.vision_enabled = False runner.agent_callback = MagicMock() runner.session = Session(sqlite_engine) @@ -290,6 +295,96 @@ class TestClearUserPromptImageMessages: assert result[0].content == "hello\n[image]\n[file]" + def test_keeps_knowledge_retrieval_image_message(self, runner: FunctionCallAgentRunner): + text = TextPromptMessageContent(data="query") + image = ImagePromptMessageContent(format="url", mime_type="image/png") + user_msg = UserPromptMessage(name="knowledge_retrieval", content=[image, text]) + + result = runner._clear_user_prompt_image_messages([user_msg]) + + assert result[0].content == [image, text] + + +# ============================== +# Dataset Tool Image Content +# ============================== + + +class TestBuildDatasetToolImageContents: + def test_returns_empty_when_vision_disabled(self, runner: FunctionCallAgentRunner): + tool = MagicMock() + tool.__class__.__name__ = "DatasetRetrieverTool" + response = "![image](http://localhost:5001/files/890985e9-c2f1-484e-bc7b-62010a337e6d/file-preview)" + + assert runner._build_dataset_tool_image_contents(runner.session, response, tool) == [] + + def test_builds_image_contents_from_dataset_tool_preview_links( + self, runner: FunctionCallAgentRunner, mocker: MockerFixture + ): + from core.tools.utils.dataset_retriever_tool import DatasetRetrieverTool + + runner.vision_enabled = True + image_content = ImagePromptMessageContent(format="url", mime_type="image/png") + to_prompt_content = mocker.patch( + "core.agent.fc_agent_runner.file_manager.to_prompt_message_content", + return_value=image_content, + ) + grant_access = mocker.patch("core.agent.fc_agent_runner.grant_upload_file_access") + sign_preview = mocker.patch( + "core.agent.fc_agent_runner.sign_upload_file_preview_url", + return_value="http://localhost:5001/files/file-id/file-preview?sign=1", + ) + build_reference = mocker.patch("core.agent.fc_agent_runner.build_file_reference", return_value="file-ref") + + upload_file = UploadFile( + tenant_id="00000000-0000-0000-0000-000000000001", + storage_type=StorageType.LOCAL, + key="image_files/chart.png", + name="chart.png", + size=123, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="00000000-0000-0000-0000-000000000002", + created_at=datetime.now(UTC), + used=True, + ) + upload_file.id = "890985e9-c2f1-484e-bc7b-62010a337e6d" + non_image_file = UploadFile( + tenant_id=upload_file.tenant_id, + storage_type=StorageType.LOCAL, + key="files/report.pdf", + name="report.pdf", + size=10, + extension="pdf", + mime_type="application/pdf", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=upload_file.created_by, + created_at=datetime.now(UTC), + used=True, + ) + non_image_file.id = "11111111-1111-1111-1111-111111111111" + session = runner.session + session.add_all([upload_file, non_image_file]) + session.commit() + + response = ( + "![image](http://localhost:5001/files/890985e9-c2f1-484e-bc7b-62010a337e6d/file-preview?sign=1)\n" + "duplicate ![image](http://localhost:5001/files/890985e9-c2f1-484e-bc7b-62010a337e6d/file-preview)\n" + "file ![file](http://localhost:5001/files/11111111-1111-1111-1111-111111111111/file-preview)" + ) + + tool = MagicMock(spec=DatasetRetrieverTool) + contents = runner._build_dataset_tool_image_contents(session, response, tool) + + assert contents == [image_content] + assert contents[0].type == PromptMessageContentType.IMAGE + grant_access.assert_called_once() + assert list(grant_access.call_args.args[0]) == ["890985e9-c2f1-484e-bc7b-62010a337e6d"] + sign_preview.assert_called_once_with(upload_file.id, upload_file.extension) + build_reference.assert_called_once_with(record_id=str(upload_file.id)) + to_prompt_content.assert_called_once() + # ============================== # Run Method Tests @@ -314,13 +409,24 @@ class TestRunMethod: queue_calls = runner.queue_manager.publish.call_args_list assert any(call.args and call.args[0].__class__.__name__ == "QueueMessageEndEvent" for call in queue_calls) - def test_run_streaming_branch(self, runner: FunctionCallAgentRunner): + def test_run_streaming_branch(self, runner: FunctionCallAgentRunner, mocker: MockerFixture): message = MagicMock(id="m1") runner.stream_tool_call = True events: list[str] = [] - session = MagicMock() - session.commit.side_effect = lambda: events.append("commit") - session.close.side_effect = lambda: events.append("close") + session = runner.session + original_commit = session.commit + original_close = session.close + + def commit_session() -> None: + events.append("commit") + original_commit() + + def close_session() -> None: + events.append("close") + original_close() + + mocker.patch.object(session, "commit", side_effect=commit_session) + mocker.patch.object(session, "close", side_effect=close_session) content = [TextPromptMessageContent(data="hi")] chunk = DummyChunk(message=DummyMessage(content=content), usage=build_usage()) diff --git a/api/tests/unit_tests/core/agent/test_publish_visibility.py b/api/tests/unit_tests/core/agent/test_publish_visibility.py index 7e6b0964c83..042c88f3d17 100644 --- a/api/tests/unit_tests/core/agent/test_publish_visibility.py +++ b/api/tests/unit_tests/core/agent/test_publish_visibility.py @@ -71,6 +71,7 @@ def _add_agent( tenant_id="tenant-1", agent_id=agent_id, version=1, + home_snapshot_id=f"home-{agent_id}", config_snapshot=_agent_soul(), ) ) @@ -245,6 +246,7 @@ def test_workflow_callable_filter_distinguishes_never_published_from_dirty_draft tenant_id="tenant-1", agent_id=stale_published_snapshot.id, version=2, + home_snapshot_id="home-stale-published-old", config_snapshot=_agent_soul(), ) ) diff --git a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py index f347a5fae7e..b5396e3acd5 100644 --- a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_generator.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock import pytest from pydantic import BaseModel, ValidationError +from sqlalchemy.orm import Session from constants import UUID_NIL from core.app.app_config.entities import AppAdditionalFeatures, WorkflowUIBasedAppConfig @@ -25,7 +26,7 @@ from models.model import AppMode class TestAdvancedChatAppGeneratorValidation: - def test_generate_requires_query(self): + def test_generate_requires_query(self, unbound_session: Session): generator = AdvancedChatAppGenerator() with pytest.raises(ValueError, match="query is required"): @@ -37,10 +38,10 @@ class TestAdvancedChatAppGeneratorValidation: invoke_from=InvokeFrom.WEB_APP, workflow_run_id="run-id", streaming=False, - session=MagicMock(), + session=unbound_session, ) - def test_generate_requires_string_query(self): + def test_generate_requires_string_query(self, unbound_session: Session): generator = AdvancedChatAppGenerator() with pytest.raises(ValueError, match="query must be a string"): @@ -52,10 +53,10 @@ class TestAdvancedChatAppGeneratorValidation: invoke_from=InvokeFrom.WEB_APP, workflow_run_id="run-id", streaming=False, - session=MagicMock(), + session=unbound_session, ) - def test_single_iteration_generate_validates_args(self): + def test_single_iteration_generate_validates_args(self, unbound_session: Session): generator = AdvancedChatAppGenerator() with pytest.raises(ValueError, match="node_id is required"): @@ -66,7 +67,7 @@ class TestAdvancedChatAppGeneratorValidation: user=SimpleNamespace(), args={"inputs": {}}, streaming=False, - session=MagicMock(), + session=unbound_session, ) with pytest.raises(ValueError, match="inputs is required"): @@ -77,10 +78,10 @@ class TestAdvancedChatAppGeneratorValidation: user=SimpleNamespace(), args={}, streaming=False, - session=MagicMock(), + session=unbound_session, ) - def test_single_loop_generate_validates_args(self): + def test_single_loop_generate_validates_args(self, unbound_session: Session): generator = AdvancedChatAppGenerator() with pytest.raises(ValueError, match="node_id is required"): @@ -91,7 +92,7 @@ class TestAdvancedChatAppGeneratorValidation: user=SimpleNamespace(), args=SimpleNamespace(inputs={}), streaming=False, - session=MagicMock(), + session=unbound_session, ) with pytest.raises(ValueError, match="inputs is required"): @@ -102,7 +103,7 @@ class TestAdvancedChatAppGeneratorValidation: user=SimpleNamespace(), args=SimpleNamespace(inputs=None), streaming=False, - session=MagicMock(), + session=unbound_session, ) @@ -118,7 +119,7 @@ class TestAdvancedChatAppGeneratorInternals: workflow_id="workflow-id", ) - def test_generate_loads_conversation_and_files(self, monkeypatch: pytest.MonkeyPatch): + def test_generate_loads_conversation_and_files(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): generator = AdvancedChatAppGenerator() app_config = self._build_app_config() @@ -126,7 +127,7 @@ class TestAdvancedChatAppGeneratorInternals: built_files: list[object] = [] build_files_called = {"called": False} captured: dict[str, object] = {} - session = MagicMock() + session = unbound_session get_conversation = MagicMock(return_value=conversation) monkeypatch.setattr( @@ -207,7 +208,7 @@ class TestAdvancedChatAppGeneratorInternals: assert build_files_called["called"] is True assert get_conversation.call_args.kwargs["session"] is session - def test_resume_delegates_to_generate(self, monkeypatch: pytest.MonkeyPatch): + def test_resume_delegates_to_generate(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): generator = AdvancedChatAppGenerator() existing_trace_manager = SimpleNamespace(app_id="existing-app", user_id="existing-user") application_generate_entity = AdvancedChatAppGenerateEntity.model_construct( @@ -241,7 +242,7 @@ class TestAdvancedChatAppGeneratorInternals: user=SimpleNamespace(id="end-user-id", session_id="session-id"), conversation=SimpleNamespace(id="conversation-id"), message=SimpleNamespace(id="message-id"), - session=MagicMock(), + session=unbound_session, application_generate_entity=application_generate_entity, workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), @@ -254,7 +255,9 @@ class TestAdvancedChatAppGeneratorInternals: assert captured_entity.trace_manager is existing_trace_manager assert captured_graph_runtime_state is not None - def test_single_iteration_generate_builds_debug_task(self, monkeypatch: pytest.MonkeyPatch): + def test_single_iteration_generate_builds_debug_task( + self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): generator = AdvancedChatAppGenerator() app_config = self._build_app_config() captured: dict[str, object] = {} @@ -262,7 +265,7 @@ class TestAdvancedChatAppGeneratorInternals: draft_sessions: list[object] = [] var_loader = SimpleNamespace(loader="draft") workflow = SimpleNamespace(id="workflow-id") - session = MagicMock() + session = unbound_session monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.AdvancedChatAppConfigManager.get_app_config", @@ -318,7 +321,7 @@ class TestAdvancedChatAppGeneratorInternals: assert captured["application_generate_entity"].single_iteration_run.node_id == "node-1" assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1" - def test_single_loop_generate_builds_debug_task(self, monkeypatch: pytest.MonkeyPatch): + def test_single_loop_generate_builds_debug_task(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): generator = AdvancedChatAppGenerator() app_config = self._build_app_config() captured: dict[str, object] = {} @@ -326,7 +329,7 @@ class TestAdvancedChatAppGeneratorInternals: draft_sessions: list[object] = [] var_loader = SimpleNamespace(loader="draft") workflow = SimpleNamespace(id="workflow-id") - session = MagicMock() + session = unbound_session monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.AdvancedChatAppConfigManager.get_app_config", @@ -442,6 +445,13 @@ class TestAdvancedChatAppGeneratorInternals: def start(self): thread_data["started"] = True + def join(self, timeout): + thread_data["joined"] = True + thread_data["join_timeout"] = timeout + + def is_alive(self): + return False + monkeypatch.setattr("core.app.apps.advanced_chat.app_generator.threading.Thread", _Thread) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.db", SimpleNamespace(engine=object(), session=db_session) @@ -475,6 +485,8 @@ class TestAdvancedChatAppGeneratorInternals: assert response["response"] == {"raw": True} assert thread_data["started"] is True + assert thread_data["joined"] is True + assert thread_data["join_timeout"] == 300 assert "pause-layer" in thread_data["kwargs"]["graph_engine_layers"] assert generator._dialogue_count == 3 assert init_records.call_args.kwargs["session"] is db_session @@ -542,6 +554,13 @@ class TestAdvancedChatAppGeneratorInternals: def start(self): thread_data["started"] = True + def join(self, timeout): + thread_data["joined"] = True + thread_data["join_timeout"] = timeout + + def is_alive(self): + return False + monkeypatch.setattr("core.app.apps.advanced_chat.app_generator.threading.Thread", _Thread) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.db", SimpleNamespace(engine=object(), session=db_session) @@ -574,6 +593,8 @@ class TestAdvancedChatAppGeneratorInternals: init_records.assert_not_called() get_thread_messages_length.assert_called_once_with(conversation.id, session=db_session) assert thread_data["started"] is True + assert thread_data["joined"] is True + assert thread_data["join_timeout"] == 300 db_session.commit.assert_not_called() db_session.refresh.assert_not_called() db_session.close.assert_called_once() @@ -730,11 +751,13 @@ class TestAdvancedChatAppGeneratorInternals: monkeypatch.setattr("core.app.apps.advanced_chat.app_generator.preserve_flask_contexts", _fake_context) + workflow = SimpleNamespace(id="workflow-id", tenant_id="tenant", app_id="app") + class _Session: def __init__(self, *args, **kwargs): self.scalar = MagicMock( side_effect=[ - SimpleNamespace(id="workflow-id", tenant_id="tenant", app_id="app"), + workflow, SimpleNamespace(id="app"), ] ) @@ -754,6 +777,8 @@ class TestAdvancedChatAppGeneratorInternals: monkeypatch.setattr("core.app.apps.advanced_chat.app_generator.Session", _Session) monkeypatch.setattr("core.app.apps.advanced_chat.app_generator.AdvancedChatAppRunner", _Runner) + restore_workflow_run_graph = MagicMock() + monkeypatch.setattr(generator, "_restore_workflow_run_graph", restore_workflow_run_graph) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.db", SimpleNamespace(engine=object(), session=SimpleNamespace(close=lambda: None)), @@ -770,10 +795,12 @@ class TestAdvancedChatAppGeneratorInternals: workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), graph_engine_layers=(), - graph_runtime_state=None, + graph_runtime_state=SimpleNamespace(), ) queue_manager.publish_error.assert_not_called() + assert restore_workflow_run_graph.call_args.kwargs["workflow"] is workflow + assert restore_workflow_run_graph.call_args.kwargs["workflow_run_id"] == "run-id" def test_generate_worker_handles_validation_error(self, monkeypatch: pytest.MonkeyPatch): generator = AdvancedChatAppGenerator() @@ -1132,7 +1159,7 @@ class TestAdvancedChatAppGeneratorInternals: assert queue_manager.publish_error.called - def test_generate_debugger_enables_retrieve_source(self, monkeypatch: pytest.MonkeyPatch): + def test_generate_debugger_enables_retrieve_source(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): generator = AdvancedChatAppGenerator() app_config = WorkflowUIBasedAppConfig( @@ -1205,14 +1232,16 @@ class TestAdvancedChatAppGeneratorInternals: invoke_from=InvokeFrom.DEBUGGER, workflow_run_id="run-id", streaming=False, - session=MagicMock(), + session=unbound_session, ) assert result == {"ok": True} assert app_config.additional_features.show_retrieve_source is True assert captured["application_generate_entity"].query == "hello" - def test_generate_service_api_sets_parent_message_id(self, monkeypatch: pytest.MonkeyPatch): + def test_generate_service_api_sets_parent_message_id( + self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): generator = AdvancedChatAppGenerator() app_config = WorkflowUIBasedAppConfig( @@ -1285,7 +1314,7 @@ class TestAdvancedChatAppGeneratorInternals: invoke_from=InvokeFrom.SERVICE_API, workflow_run_id="run-id", streaming=False, - session=MagicMock(), + session=unbound_session, ) assert captured["application_generate_entity"].parent_message_id == UUID_NIL @@ -1303,7 +1332,9 @@ class TestAdvancedChatAppGeneratorResume: workflow_id="workflow-id", ) - def test_resume_restores_trace_manager_when_missing(self, monkeypatch: pytest.MonkeyPatch): + def test_resume_restores_trace_manager_when_missing( + self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session + ): generator = AdvancedChatAppGenerator() application_generate_entity = AdvancedChatAppGenerateEntity.model_construct( task_id="task", @@ -1349,7 +1380,7 @@ class TestAdvancedChatAppGeneratorResume: user=SimpleNamespace(id="end-user-id", session_id="session-id"), conversation=SimpleNamespace(id="conversation-id"), message=SimpleNamespace(id="message-id"), - session=MagicMock(), + session=unbound_session, application_generate_entity=application_generate_entity, workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), @@ -1363,7 +1394,7 @@ class TestAdvancedChatAppGeneratorResume: assert trace_manager.app_id == "app-id" assert trace_manager.user_id == "session-id" - def test_resume_preserves_existing_trace_manager(self, monkeypatch: pytest.MonkeyPatch): + def test_resume_preserves_existing_trace_manager(self, monkeypatch: pytest.MonkeyPatch, unbound_session: Session): generator = AdvancedChatAppGenerator() existing_trace_manager = SimpleNamespace(app_id="existing-app", user_id="existing-user") application_generate_entity = AdvancedChatAppGenerateEntity.model_construct( @@ -1397,7 +1428,7 @@ class TestAdvancedChatAppGeneratorResume: user=SimpleNamespace(id="end-user-id", session_id="session-id"), conversation=SimpleNamespace(id="conversation-id"), message=SimpleNamespace(id="message-id"), - session=MagicMock(), + session=unbound_session, application_generate_entity=application_generate_entity, workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), diff --git a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_input_moderation.py b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_input_moderation.py index 5a7b58e581a..7811ab78739 100644 --- a/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_input_moderation.py +++ b/api/tests/unit_tests/core/app/apps/advanced_chat/test_app_runner_input_moderation.py @@ -1,13 +1,18 @@ +from contextlib import contextmanager from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest +from sqlalchemy import event +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, sessionmaker import core.app.apps.advanced_chat.app_runner as module from core.app.apps.advanced_chat.app_runner import AdvancedChatAppRunner from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom from core.app.entities.queue_entities import QueueAnnotationReplyEvent, QueueStopEvent from core.moderation.base import ModerationError +from models.model import App, AppMode, IconType MINIMAL_GRAPH = { "nodes": [ @@ -24,10 +29,11 @@ MINIMAL_GRAPH = { @pytest.fixture -def build_runner(): +def build_runner(sqlite_session: Session): """Construct a minimal AdvancedChatAppRunner with heavy dependencies mocked.""" app_id = str(uuid4()) workflow_id = str(uuid4()) + tenant_id = str(uuid4()) # Mocks for constructor args mock_queue_manager = MagicMock() @@ -41,7 +47,7 @@ def build_runner(): mock_workflow = MagicMock() mock_workflow.id = workflow_id - mock_workflow.tenant_id = str(uuid4()) + mock_workflow.tenant_id = tenant_id mock_workflow.app_id = app_id mock_workflow.type = "chat" mock_workflow.graph_dict = MINIMAL_GRAPH @@ -50,7 +56,22 @@ def build_runner(): 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_config.tenant_id = tenant_id + + sqlite_session.add( + App( + id=app_id, + tenant_id=tenant_id, + name="Advanced chat app", + mode=AppMode.ADVANCED_CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#ffffff", + enable_site=False, + enable_api=False, + ) + ) + sqlite_session.commit() gen = MagicMock(spec=AdvancedChatAppGenerateEntity) gen.app_config = mock_app_config @@ -86,24 +107,8 @@ def build_runner(): def _patch_common_run_deps(runner: AdvancedChatAppRunner): """Context manager that patches common heavy deps used by run().""" - # create_session() returns a context manager whose body yields a session that - # supports both scalar() (app record lookup) and begin()/scalars().all() - # (conversation variable initialization). - mock_session = MagicMock() - mock_session.scalar.return_value = MagicMock() - mock_session.scalars.return_value.all.return_value = [] - - 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.multiple( "core.app.apps.advanced_chat.app_runner", - create_session=MagicMock(return_value=session_context), - select=MagicMock(), - session_factory=MagicMock(get_session_maker=MagicMock(return_value=MagicMock())), RedisChannel=MagicMock(), redis_client=MagicMock(), WorkflowEntry=MagicMock(**{"return_value.run.return_value": iter([])}), @@ -192,15 +197,15 @@ def test_run_returns_early_when_direct_output_via_handle_input_moderation(build_ mock_init_graph.assert_not_called() -def test_run_publishes_annotation_after_commit(build_runner): +def test_run_publishes_annotation_after_commit(build_runner, sqlite_engine: Engine): runner = build_runner events: list[str] = [] - session = MagicMock() - session.scalar.return_value = MagicMock() - session.commit.side_effect = lambda: events.append("commit") - session_context = MagicMock() - session_context.__enter__.return_value = session - session_context.__exit__.return_value = False + + def record_commit(session: Session) -> None: + if session.get_bind() is sqlite_engine: + events.append("commit") + + event.listen(Session, "after_commit", record_commit) annotation_reply = MagicMock(id="annotation-1", content="annotated answer") def publish(event): @@ -209,7 +214,6 @@ def test_run_publishes_annotation_after_commit(build_runner): with ( _patch_common_run_deps(runner), - patch.object(module, "create_session", return_value=session_context), patch.object( runner, "handle_input_moderation", @@ -220,19 +224,20 @@ def test_run_publishes_annotation_after_commit(build_runner): patch.object(runner, "_complete_with_stream_output"), ): runner.run() + event.remove(Session, "after_commit", record_commit) assert events == ["commit", "publish"] -def test_run_closes_scoped_session_before_workflow_run(build_runner): +def test_run_closes_scoped_session_before_workflow_run(build_runner, sqlite_session_factory: sessionmaker[Session]): runner = build_runner events = [] - mock_session = MagicMock() - mock_session.scalar.return_value = MagicMock() - session_context = MagicMock() - session_context.__enter__.return_value = mock_session - session_context.__exit__.side_effect = lambda exc_type, exc, tb: events.append("close") or False + @contextmanager + def observed_session(): + with sqlite_session_factory() as session: + yield session + events.append("close") workflow_entry = MagicMock() @@ -243,8 +248,7 @@ def test_run_closes_scoped_session_before_workflow_run(build_runner): workflow_entry.run.side_effect = run_workflow with ( - patch.object(module, "create_session", return_value=session_context), - patch.object(module, "session_factory", MagicMock(get_session_maker=MagicMock(return_value=MagicMock()))), + patch.object(module, "create_session", observed_session), patch.object(module, "RedisChannel"), patch.object(module, "redis_client"), patch.object(module, "WorkflowEntry", return_value=workflow_entry), diff --git a/api/tests/unit_tests/core/app/apps/advanced_chat/test_generate_task_pipeline_core.py b/api/tests/unit_tests/core/app/apps/advanced_chat/test_generate_task_pipeline_core.py index 1f975c74fb1..f9995597279 100644 --- a/api/tests/unit_tests/core/app/apps/advanced_chat/test_generate_task_pipeline_core.py +++ b/api/tests/unit_tests/core/app/apps/advanced_chat/test_generate_task_pipeline_core.py @@ -4,6 +4,7 @@ from contextlib import contextmanager from types import SimpleNamespace import pytest +from sqlalchemy.orm import Session from core.app.app_config.entities import AppAdditionalFeatures, WorkflowUIBasedAppConfig from core.app.apps.advanced_chat.generate_task_pipeline import ( @@ -54,10 +55,12 @@ from core.workflow.nodes.human_input.entities import UserActionConfig from core.workflow.nodes.human_input.pause_reason import DifyHITLEventType from core.workflow.system_variables import build_system_variables from graphon.enums import BuiltinNodeTypes +from graphon.file import FileTransferMethod, FileType +from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.runtime import GraphRuntimeState, VariablePool from libs.datetime_utils import naive_utc_now from models.enums import MessageStatus -from models.model import AppMode, EndUser +from models.model import AppMode, EndUser, Message, MessageFile from tests.workflow_test_utils import build_test_variable_pool @@ -110,6 +113,40 @@ def _make_pipeline(): return pipeline +def _persist_message(session: Session, *, message_id: str = "message-id") -> Message: + message = Message( + app_id="app", + model_provider="provider", + model_id="model", + override_model_configs=None, + conversation_id="conv-id", + inputs={}, + query="hello", + message="", + message_tokens=0, + message_unit_price=0, + message_price_unit=0, + 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="end-user", + from_account_id=None, + app_mode=AppMode.ADVANCED_CHAT, + status=MessageStatus.PAUSED, + ) + message.id = message_id + session.add(message) + session.commit() + return message + + class TestAdvancedChatGenerateTaskPipeline: def test_ensure_workflow_initialized_raises(self): pipeline = _make_pipeline() @@ -138,7 +175,7 @@ class TestAdvancedChatGenerateTaskPipeline: variables=build_system_variables(workflow_execution_id="run-id"), ), start_at=0.0, - total_tokens=7, + llm_usage=LLMUsage.empty_usage().model_copy(update={"total_tokens": 7}), node_run_steps=3, ) @@ -255,53 +292,29 @@ class TestAdvancedChatGenerateTaskPipeline: pipeline._base_task_pipeline.handle_error = lambda **kwargs: ValueError("boom") pipeline._base_task_pipeline.error_to_stream_response = lambda err: err - @contextmanager - def _fake_session(): - yield SimpleNamespace() - - pipeline._database_session = _fake_session - responses = list(pipeline._handle_error_event(QueueErrorEvent(error=ValueError("boom")))) assert isinstance(responses[0], ValueError) - def test_handle_workflow_started_event_sets_run_id(self, monkeypatch: pytest.MonkeyPatch): + def test_handle_workflow_started_event_sets_run_id(self, sqlite_session: Session): pipeline = _make_pipeline() + message = _persist_message(sqlite_session) + other_message = _persist_message(sqlite_session, message_id="other-message-id") pipeline._graph_runtime_state = GraphRuntimeState( variable_pool=build_test_variable_pool(variables=build_system_variables(workflow_execution_id="run-id")), start_at=0.0, ) pipeline._workflow_response_converter.workflow_start_to_stream_response = lambda **kwargs: "started" - # Track database operations for verification - executed_statements = [] - - @contextmanager - def _fake_session(): - sess = SimpleNamespace() - - def _execute(stmt): - executed_statements.append(stmt) - return SimpleNamespace() - - sess.execute = _execute - yield sess - - monkeypatch.setattr(pipeline, "_database_session", _fake_session) - monkeypatch.setattr(pipeline, "_get_message", lambda **kwargs: SimpleNamespace()) - responses = list(pipeline._handle_workflow_started_event(QueueWorkflowStartedEvent())) assert pipeline._workflow_run_id == "run-id" assert responses == ["started"] - # Verify database operation was executed - assert len(executed_statements) == 1 - # Verify the UPDATE statement targets the correct message and sets workflow_run_id - update_stmt = executed_statements[0] - stmt_str = str(update_stmt) - assert "UPDATE messages" in stmt_str - assert "WHERE messages.id" in stmt_str + sqlite_session.refresh(message) + sqlite_session.refresh(other_message) + assert message.workflow_run_id == "run-id" + assert other_message.workflow_run_id is None def test_message_end_to_stream_response_strips_annotation_reply(self): pipeline = _make_pipeline() @@ -464,7 +477,7 @@ class TestAdvancedChatGenerateTaskPipeline: assert list(pipeline._handle_loop_next_event(loop_next)) == ["loop_next"] assert list(pipeline._handle_loop_completed_event(loop_done)) == ["loop_done"] - def test_workflow_finish_handlers(self, monkeypatch: pytest.MonkeyPatch): + def test_workflow_finish_handlers(self): pipeline = _make_pipeline() pipeline._workflow_run_id = "run-id" pipeline._graph_runtime_state = GraphRuntimeState( @@ -482,12 +495,6 @@ class TestAdvancedChatGenerateTaskPipeline: pipeline._base_task_pipeline.error_to_stream_response = lambda err: err pipeline._get_message = lambda **kwargs: SimpleNamespace(id="message-id") - @contextmanager - def _fake_session(): - yield SimpleNamespace(scalar=lambda *args, **kwargs: None) - - monkeypatch.setattr(pipeline, "_database_session", _fake_session) - succeeded_responses = list(pipeline._handle_workflow_succeeded_event(QueueWorkflowSucceededEvent(outputs={}))) assert len(succeeded_responses) == 2 assert isinstance(succeeded_responses[0], MessageEndStreamResponse) @@ -644,12 +651,12 @@ class TestAdvancedChatGenerateTaskPipeline: assert list(pipeline._handle_human_input_form_timeout_event(timeout_event)) == ["timeout"] assert persisted == ["saved"] - def test_save_message_preserves_full_answer_and_sets_usage(self): + def test_save_message_preserves_full_answer_and_sets_usage(self, sqlite_session: Session): pipeline = _make_pipeline() pipeline._recorded_files = [ { - "type": "image", - "transfer_method": "remote", + "type": FileType.IMAGE, + "transfer_method": FileTransferMethod.REMOTE_URL, "remote_url": "http://example.com/file.png", "related_id": "file-id", } @@ -660,32 +667,7 @@ class TestAdvancedChatGenerateTaskPipeline: pipeline._task_state.first_token_time = pipeline._base_task_pipeline.start_at + 0.1 pipeline._task_state.last_token_time = pipeline._base_task_pipeline.start_at + 0.2 - message = SimpleNamespace( - id="message-id", - status=MessageStatus.PAUSED, - answer="", - updated_at=None, - provider_response_latency=None, - message_tokens=None, - message_unit_price=None, - message_price_unit=None, - answer_tokens=None, - answer_unit_price=None, - answer_price_unit=None, - total_price=None, - currency=None, - message_metadata=None, - invoke_from=InvokeFrom.WEB_APP, - from_account_id=None, - from_end_user_id="end-user", - ) - - class _Session: - def scalar(self, *args, **kwargs): - return message - - def add_all(self, items): - self.items = items + message = _persist_message(sqlite_session) graph_runtime_state = GraphRuntimeState( variable_pool=VariablePool.from_bootstrap( @@ -694,13 +676,15 @@ class TestAdvancedChatGenerateTaskPipeline: start_at=0.0, ) - pipeline._save_message(session=_Session(), graph_runtime_state=graph_runtime_state) + pipeline._save_message(session=sqlite_session, graph_runtime_state=graph_runtime_state) + sqlite_session.commit() assert message.status == MessageStatus.NORMAL assert message.answer == "![img](http://example.com/file.png) hello ![inline](http://llm.com/img.jpg)" assert message.message_metadata + assert sqlite_session.query(MessageFile).filter_by(message_id=message.id).count() == 1 - def test_handle_stop_event_saves_message_for_moderation(self, monkeypatch: pytest.MonkeyPatch): + def test_handle_stop_event_saves_message_for_moderation(self): pipeline = _make_pipeline() pipeline._message_end_to_stream_response = lambda: "end" saved: list[str] = [] @@ -710,18 +694,12 @@ class TestAdvancedChatGenerateTaskPipeline: pipeline._save_message = _save_message - @contextmanager - def _fake_session(): - yield SimpleNamespace() - - monkeypatch.setattr(pipeline, "_database_session", _fake_session) - responses = list(pipeline._handle_stop_event(QueueStopEvent(stopped_by=QueueStopEvent.StopBy.INPUT_MODERATION))) assert responses == ["end"] assert saved == ["saved"] - def test_handle_message_end_event_applies_output_moderation(self, monkeypatch: pytest.MonkeyPatch): + def test_handle_message_end_event_applies_output_moderation(self): pipeline = _make_pipeline() pipeline._graph_runtime_state = GraphRuntimeState( variable_pool=VariablePool.from_bootstrap( @@ -740,12 +718,6 @@ class TestAdvancedChatGenerateTaskPipeline: pipeline._save_message = _save_message - @contextmanager - def _fake_session(): - yield SimpleNamespace() - - monkeypatch.setattr(pipeline, "_database_session", _fake_session) - responses = list(pipeline._handle_advanced_chat_message_end_event(QueueAdvancedChatMessageEndEvent())) assert responses == ["replace", "end"] diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_app_generator.py b/api/tests/unit_tests/core/app/apps/agent_app/test_app_generator.py index 12c429e85b9..d346dfe9601 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_app_generator.py @@ -22,11 +22,10 @@ from core.app.apps.agent_app.app_generator import ( AgentAppGeneratorError, ) from core.app.apps.exc import GenerateTaskStoppedError -from core.app.entities.app_invoke_entities import AGENT_RUNTIME_EXIT_INTENT_ARG, InvokeFrom, UserFrom +from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom from core.app.entities.queue_entities import QueueAnnotationReplyEvent from core.workflow.file_reference import build_file_reference from models import Account, AppModelConfig -from models.agent import AgentConfigDraftType MODULE = "core.app.apps.agent_app.app_generator" @@ -79,14 +78,18 @@ class TestGenerateGuards: class TestGenerateSuccess: - def test_runtime_session_snapshot_id_preserves_snapshot_for_debugger_and_web_app(self): + def test_session_scope_config_version_id_preserves_draft_or_snapshot_id(self): assert ( - AgentAppGenerator._runtime_session_snapshot_id(invoke_from=InvokeFrom.DEBUGGER, snapshot_id="snap-1") - == "snap-1" + AgentAppGenerator._session_scope_config_version_id( + invoke_from=InvokeFrom.DEBUGGER, config_version_id="draft-1" + ) + == "draft-1" ) assert ( - AgentAppGenerator._runtime_session_snapshot_id(invoke_from=InvokeFrom.WEB_APP, snapshot_id="snap-1") - == "snap-1" + AgentAppGenerator._session_scope_config_version_id( + invoke_from=InvokeFrom.WEB_APP, config_version_id="snapshot-1" + ) + == "snapshot-1" ) def test_generate_orchestrates_and_starts_worker(self, generator, mocker: MockerFixture): @@ -145,84 +148,11 @@ class TestGenerateSuccess: draft_type=None, user=user, session=session, + conversation=None, ) session.get.assert_called_once_with(AppModelConfig, "config-1") assert generate_entity.call_args.kwargs["prompt_file_mappings"] == file_mappings - assert generate_entity.call_args.kwargs["agent_runtime_exit_intent"] == "suspend" - - def test_generate_uses_delete_exit_intent_from_internal_arg(self, generator, mocker: MockerFixture): - app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") - user = DummyAccount("user") - - generator._resolve_agent = mocker.MagicMock( - return_value=(mocker.MagicMock(id="agent1"), "snap1", "snapshot", mocker.MagicMock()) - ) - generator._prepare_user_inputs = mocker.MagicMock(return_value={}) - generator._init_generate_records = mocker.MagicMock( - return_value=(mocker.MagicMock(id="conv", mode="agent"), mocker.MagicMock(id="msg")) - ) - generator._handle_response = mocker.MagicMock(return_value="raw-response") - - mocker.patch( - f"{MODULE}.AgentAppConfigManager.get_app_config", - return_value=mocker.MagicMock(variables=[], tenant_id="tenant", app_id="app1"), - ) - mocker.patch(f"{MODULE}.ModelConfigConverter.convert", return_value=mocker.MagicMock(model="gpt-4o-mini")) - mocker.patch(f"{MODULE}.TraceQueueManager", return_value=mocker.MagicMock()) - generate_entity = mocker.patch( - f"{MODULE}.AgentAppGenerateEntity", return_value=mocker.MagicMock(task_id="t", user_id="user") - ) - mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock()) - mocker.patch(f"{MODULE}.threading.Thread", return_value=mocker.MagicMock()) - mocker.patch(f"{MODULE}.AgentAppGenerateResponseConverter.convert", return_value={"result": "ok"}) - - generator.generate( - app_model=app_model, - user=user, - args={"query": "hello", "inputs": {}, AGENT_RUNTIME_EXIT_INTENT_ARG: "delete"}, - invoke_from=InvokeFrom.DEBUGGER, - session=mocker.MagicMock(), - streaming=True, - ) - - assert generate_entity.call_args.kwargs["agent_runtime_exit_intent"] == "delete" - - def test_generate_falls_back_to_suspend_for_invalid_internal_exit_intent(self, generator, mocker: MockerFixture): - app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") - user = DummyAccount("user") - - generator._resolve_agent = mocker.MagicMock( - return_value=(mocker.MagicMock(id="agent1"), "snap1", "snapshot", mocker.MagicMock()) - ) - generator._prepare_user_inputs = mocker.MagicMock(return_value={}) - generator._init_generate_records = mocker.MagicMock( - return_value=(mocker.MagicMock(id="conv", mode="agent"), mocker.MagicMock(id="msg")) - ) - generator._handle_response = mocker.MagicMock(return_value="raw-response") - - mocker.patch( - f"{MODULE}.AgentAppConfigManager.get_app_config", - return_value=mocker.MagicMock(variables=[], tenant_id="tenant", app_id="app1"), - ) - mocker.patch(f"{MODULE}.ModelConfigConverter.convert", return_value=mocker.MagicMock(model="gpt-4o-mini")) - mocker.patch(f"{MODULE}.TraceQueueManager", return_value=mocker.MagicMock()) - generate_entity = mocker.patch( - f"{MODULE}.AgentAppGenerateEntity", return_value=mocker.MagicMock(task_id="t", user_id="user") - ) - mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock()) - mocker.patch(f"{MODULE}.threading.Thread", return_value=mocker.MagicMock()) - mocker.patch(f"{MODULE}.AgentAppGenerateResponseConverter.convert", return_value={"result": "ok"}) - - generator.generate( - app_model=app_model, - user=user, - args={"query": "hello", "inputs": {}, AGENT_RUNTIME_EXIT_INTENT_ARG: "bogus"}, - invoke_from=InvokeFrom.DEBUGGER, - session=mocker.MagicMock(), - streaming=True, - ) - - assert generate_entity.call_args.kwargs["agent_runtime_exit_intent"] == "suspend" + assert "agent_runtime_exit_intent" not in generate_entity.call_args.kwargs def test_generate_loads_existing_conversation(self, generator: AgentAppGenerator, mocker: MockerFixture): app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") @@ -263,6 +193,7 @@ class TestGenerateSuccess: user=user, session=session, ) + assert generator._resolve_agent.call_args.kwargs["conversation"].id == "conv" assert generator._init_generate_records.call_args.kwargs["session"] is session def test_generate_does_not_include_trace_session_id_in_extras( @@ -326,8 +257,10 @@ class TestGenerateWorker: generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock(id="conv")) generator._get_message = mocker.MagicMock(return_value=mocker.MagicMock(id="msg")) generator._run_input_guards = mocker.MagicMock(return_value=(handled, guard_query, None)) + resolved_agent = mocker.MagicMock(id="a") + resolved_config = mocker.MagicMock(id="s", home_snapshot_id="home-1") generator._resolve_agent_by_id = mocker.MagicMock( - return_value=(mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + return_value=(resolved_agent, resolved_config, mocker.MagicMock()) ) session = mocker.MagicMock() session.get.return_value = mocker.MagicMock(id="app1") @@ -345,7 +278,7 @@ class TestGenerateWorker: mocker.patch(f"{MODULE}.AgentAppRuntimeRequestBuilder", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.create_agent_backend_run_client", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.AgentBackendRunEventAdapter", return_value=mocker.MagicMock()) - mocker.patch(f"{MODULE}.AgentAppRuntimeSessionStore", return_value=mocker.MagicMock()) + mocker.patch(f"{MODULE}.AgentAppWorkspaceStore", return_value=mocker.MagicMock()) runner = mocker.MagicMock() if run_side_effect is not None: runner.run.side_effect = run_side_effect @@ -360,9 +293,8 @@ class TestGenerateWorker: *, is_resume=False, query="query", - runtime_session_snapshot_id="s", + session_scope_config_version_id="s", prompt_file_mappings=(), - agent_runtime_exit_intent="suspend", ): generator._generate_worker( flask_app=mocker.MagicMock(), @@ -370,8 +302,7 @@ class TestGenerateWorker: application_generate_entity=mocker.MagicMock( agent_id="a", agent_config_snapshot_id="s", - agent_runtime_session_snapshot_id=runtime_session_snapshot_id, - agent_runtime_exit_intent=agent_runtime_exit_intent, + agent_session_scope_config_version_id=session_scope_config_version_id, model_conf=mocker.MagicMock(model="m"), query=query, prompt_file_mappings=prompt_file_mappings, @@ -389,25 +320,19 @@ class TestGenerateWorker: self._call(generator, mocker, queue_manager) runner.run.assert_called_once() assert generator._resolve_agent_by_id.call_args.kwargs["session"] is resolver_session + assert runner.run.call_args.kwargs["home_snapshot_id"] == "home-1" + assert "home_snapshot_ref" not in runner.run.call_args.kwargs queue_manager.publish_error.assert_not_called() - def test_worker_passes_runtime_session_scope_to_runner(self, generator, mocker: MockerFixture): + def test_worker_passes_session_scope_config_version_to_runner(self, generator, mocker: MockerFixture): runner, _ = self._wire(generator, mocker) queue_manager = mocker.MagicMock() - self._call(generator, mocker, queue_manager, runtime_session_snapshot_id=None) + self._call(generator, mocker, queue_manager, session_scope_config_version_id=None) assert runner.run.call_args.kwargs["agent_config_snapshot_id"] == "s" assert runner.run.call_args.kwargs["session_scope_snapshot_id"] is None - def test_worker_forwards_runtime_exit_intent_to_runner(self, generator, mocker: MockerFixture): - runner, _ = self._wire(generator, mocker) - queue_manager = mocker.MagicMock() - - self._call(generator, mocker, queue_manager, agent_runtime_exit_intent="delete") - - assert runner.run.call_args.kwargs["agent_runtime_exit_intent"] == "delete" - def test_worker_appends_prompt_files_to_backend_query(self, generator, mocker: MockerFixture): runner, _ = self._wire(generator, mocker, guard_query="你看得见这张图片吗") queue_manager = mocker.MagicMock() @@ -526,7 +451,7 @@ class TestResumeAfterFormSubmission: mocker.patch(f"{MODULE}.TraceQueueManager", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.MessageBasedAppQueueManager", return_value=mocker.MagicMock()) mocker.patch(f"{MODULE}.threading.Thread", return_value=mocker.MagicMock()) - mocker.patch(f"{MODULE}.AgentAppRuntimeSessionStore") + generator._resolve_resume_draft = mocker.MagicMock(return_value=(None, None)) return ( mocker.patch( f"{MODULE}.AgentAppGenerateEntity", return_value=mocker.MagicMock(task_id="t", user_id="user") @@ -547,6 +472,7 @@ class TestResumeAfterFormSubmission: app_model=app_model, user=user, conversation_id="conv", + form_id="form-1", invoke_from=InvokeFrom.WEB_APP, session=session, ) @@ -573,6 +499,7 @@ class TestResumeAfterFormSubmission: app_model=mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent"), user=DummyAccount("user"), conversation_id="conv", + form_id="form-1", invoke_from=InvokeFrom.WEB_APP, session=session, ) @@ -584,25 +511,24 @@ class TestResumeAfterFormSubmission: self._wire(generator, mocker) conversation = mocker.MagicMock(id="conv", invoke_from=InvokeFrom.DEBUGGER) mocker.patch(f"{MODULE}.ConversationService.get_conversation", return_value=conversation) - session_store = mocker.patch(f"{MODULE}.AgentAppRuntimeSessionStore") - session_store.return_value.load_active_session_for_conversation.return_value = mocker.MagicMock( - scope=mocker.MagicMock(agent_config_snapshot_id="draft-build-1") - ) - draft_row = mocker.MagicMock(draft_type=AgentConfigDraftType.DEBUG_BUILD, account_id="user") - account_user = mocker.MagicMock(spec=Account) + generator._resolve_resume_draft.return_value = ("debug_build", "draft-build-1") + account_user = Account(name="Test Account", email="test@example.com") account_user.id = "user" app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent") app_model.app_model_config_id = "config-1" session = mocker.MagicMock() - session.scalar.side_effect = [draft_row, mocker.MagicMock(query="original question")] + session.scalar.return_value = mocker.MagicMock(query="original question") generator.resume_after_form_submission( app_model=app_model, user=account_user, conversation_id="conv", + form_id="form-1", invoke_from=InvokeFrom.DEBUGGER, session=session, ) assert generator._resolve_agent.call_args.kwargs["draft_type"] == "debug_build" + assert generator._resolve_agent.call_args.kwargs["draft_id"] == "draft-build-1" assert generator._resolve_agent.call_args.kwargs["session"] is session + assert generator._resolve_agent.call_args.kwargs["conversation"] is conversation diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_app_runner.py b/api/tests/unit_tests/core/app/apps/agent_app/test_app_runner.py index f623dd767d2..26e44377e3a 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_app_runner.py @@ -7,7 +7,6 @@ from __future__ import annotations from collections.abc import Callable, Iterator from datetime import UTC, datetime from decimal import Decimal -from types import SimpleNamespace from typing import Any, override from unittest.mock import MagicMock @@ -20,10 +19,12 @@ from dify_agent.protocol import ( CancelRunResponse, PydanticAIStreamRunEvent, RunEvent, + RunFailedEvent, + RunFailedEventData, + RunFailureType, RunStartedEvent, RunSucceededEvent, RunSucceededEventData, - RuntimeLayerSpec, ) from pydantic_ai.messages import ( FunctionToolCallEvent, @@ -34,9 +35,10 @@ from pydantic_ai.messages import ( ToolCallPart, ToolReturnPart, ) +from sqlalchemy import event, select +from sqlalchemy.orm import Session from clients.agent_backend import ( - AgentBackendError, AgentBackendRunEventAdapter, AgentBackendRunFailedError, AgentBackendRunFailedInternalEvent, @@ -49,7 +51,7 @@ from core.app.apps.agent_app.app_runner import AgentAppRunner from core.app.apps.agent_app.runtime_request_builder import AgentAppRuntimeRequestBuilder from core.app.apps.agent_app.session_store import AgentAppSessionScope, StoredAgentAppSession from core.app.apps.exc import GenerateTaskStoppedError -from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom +from core.app.entities.app_invoke_entities import DifyRunContext, InvokeFrom, UserFrom from core.app.entities.queue_entities import ( QueueAgentMessageEvent, QueueAgentThoughtEvent, @@ -57,20 +59,33 @@ from core.app.entities.queue_entities import ( QueueMessageEndEvent, ) from core.workflow.nodes.agent_v2.ask_human_resume import AskHumanResumeOutcome +from core.workflow.nodes.agent_v2.dify_tools_builder import WorkflowAgentToolLayers +from graphon.model_runtime.entities.llm_entities import LLMResult from graphon.model_runtime.errors.invoke import InvokeRateLimitError from models.agent_config_entities import AgentSoulConfig from models.model import MessageAgentThought +@pytest.fixture(autouse=True) +def bind_agent_db(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: + """Bind the runner's ORM writes to the shared SQLite session.""" + monkeypatch.setattr(app_runner_module.db, "session", sqlite_session) + + +def _thought_rows(session: Session) -> list[MessageAgentThought]: + session.expire_all() + return list(session.scalars(select(MessageAgentThought).order_by(MessageAgentThought.position)).all()) + + class _FakeCredentialsProvider: def fetch(self, provider_name: str, model_name: str) -> dict[str, Any]: return {"openai_api_key": "sk-test"} class _NoToolsBuilder: - def build_layers(self, **kwargs): + def build_layers(self, **kwargs: Any) -> WorkflowAgentToolLayers: del kwargs - return SimpleNamespace(plugin_tools=None, core_tools=None, exposed_tool_names=lambda: []) + return WorkflowAgentToolLayers() class _FakeQueueManager: @@ -105,6 +120,30 @@ class _RecordingFakeAgentBackendRunClient(FakeAgentBackendRunClient): return super().cancel_run(run_id, request=request) +class _RunLimitBindingLostFakeAgentBackendRunClient(FakeAgentBackendRunClient): + @override + def stream_events( + self, + run_id: str, + *, + after: str | None = None, + should_stop: Callable[[], bool] | None = None, + ) -> Iterator[RunEvent]: + del after, should_stop + created_at = datetime(2026, 1, 1, tzinfo=UTC) + yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) + yield RunFailedEvent( + id="2-0", + run_id=run_id, + created_at=created_at, + data=RunFailedEventData( + error="run limit reached", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + reason="binding_lost", + ), + ) + + class _StreamingFakeAgentBackendRunClient(FakeAgentBackendRunClient): @override def stream_events( @@ -155,43 +194,6 @@ class _StreamingFakeAgentBackendRunClient(FakeAgentBackendRunClient): ) -class _StreamingRecordingFakeAgentBackendRunClient(_RecordingFakeAgentBackendRunClient): - @override - def stream_events( - self, - run_id: str, - *, - after: str | None = None, - should_stop: Callable[[], bool] | None = None, - ) -> Iterator[RunEvent]: - del after, should_stop - created_at = datetime(2026, 1, 1, tzinfo=UTC) - yield RunStartedEvent(id="1-0", run_id=run_id, created_at=created_at) - yield PydanticAIStreamRunEvent( - id="2-0", - run_id=run_id, - created_at=created_at, - data=PartDeltaEvent(index=0, delta=TextPartDelta(content_delta="hello ")), - agent_message_delta="hello ", - ) - yield PydanticAIStreamRunEvent( - id="3-0", - run_id=run_id, - created_at=created_at, - data=PartDeltaEvent(index=0, delta=TextPartDelta(content_delta="agent")), - agent_message_delta="agent", - ) - yield RunSucceededEvent( - id="4-0", - run_id=run_id, - created_at=created_at, - data=RunSucceededEventData( - output={"text": "hello agent"}, - session_snapshot=CompositorSessionSnapshot(layers=[]), - ), - ) - - class _StreamingStopAfterFirstDeltaFakeAgentBackendRunClient(_RecordingFakeAgentBackendRunClient): def __init__(self, *, queue_manager: _FakeQueueManager, **kwargs: Any) -> None: super().__init__(**kwargs) @@ -393,86 +395,53 @@ class _ProcessStreamingFakeAgentBackendRunClient(FakeAgentBackendRunClient): ) -class _FakeDbSession: - def __init__(self) -> None: - self.rows: dict[str, MessageAgentThought] = {} - self.rollback_count = 0 - - def add(self, row: MessageAgentThought) -> None: - self.rows[str(row.id)] = row - - def commit(self) -> None: - pass - - def get(self, _model: type[MessageAgentThought], row_id: str) -> MessageAgentThought | None: - return self.rows.get(row_id) - - def delete(self, row: MessageAgentThought) -> None: - self.rows.pop(str(row.id), None) - - def rollback(self) -> None: - self.rollback_count += 1 - - class _FakeSessionStore: def __init__( self, loaded: CompositorSessionSnapshot | None = None, loaded_session: StoredAgentAppSession | None = None, - listed_sessions: list[StoredAgentAppSession] | None = None, + binding_id: str = "binding-1", + workspace_id: str = "workspace-1", + backend_binding_ref: str = "backend-binding-1", ) -> None: self.loaded = loaded self._loaded_session = loaded_session - self._listed_sessions = list(listed_sessions or []) - self.loaded_scopes: list[AgentAppSessionScope] = [] + self.binding_id = binding_id + self.workspace_id = workspace_id + self.backend_binding_ref = backend_binding_ref + self.resolved_scopes: list[AgentAppSessionScope] = [] self.saved: list[ tuple[ AgentAppSessionScope, str, CompositorSessionSnapshot | None, - list[RuntimeLayerSpec], str | None, str | None, ] ] = [] - self.cleaned: list[tuple[AgentAppSessionScope, str | None]] = [] - def load_active_snapshot(self, scope: AgentAppSessionScope) -> CompositorSessionSnapshot | None: - self.loaded_scopes.append(scope) - return self.loaded - - def load_active_session(self, scope: AgentAppSessionScope) -> StoredAgentAppSession | None: - self.loaded_scopes.append(scope) + def load_or_create(self, scope: AgentAppSessionScope) -> StoredAgentAppSession: + self.resolved_scopes.append(scope) if self._loaded_session is not None: return self._loaded_session - if self.loaded is None: - return None - return StoredAgentAppSession(scope=scope, session_snapshot=self.loaded, backend_run_id=None) - - def list_active_sessions_for_conversation( - self, *, tenant_id: str, app_id: str, conversation_id: str - ) -> list[StoredAgentAppSession]: - assert tenant_id == "tenant-1" - assert app_id == "app-1" - assert conversation_id == "conv-1" - return list(self._listed_sessions) + return StoredAgentAppSession( + scope=scope, + binding_id=self.binding_id, + workspace_id=self.workspace_id, + backend_binding_ref=self.backend_binding_ref, + session_snapshot=self.loaded, + ) def save_active_snapshot( self, *, - scope, - backend_run_id, - snapshot, - runtime_layer_specs, - pending_form_id=None, - pending_tool_call_id=None, + scope: AgentAppSessionScope, + binding_id: str, + snapshot: CompositorSessionSnapshot | None, + pending_form_id: str | None = None, + pending_tool_call_id: str | None = None, ) -> None: - self.saved.append( - (scope, backend_run_id, snapshot, list(runtime_layer_specs), pending_form_id, pending_tool_call_id) - ) - - def mark_cleaned(self, *, scope: AgentAppSessionScope, backend_run_id: str | None = None) -> None: - self.cleaned.append((scope, backend_run_id)) + self.saved.append((scope, binding_id, snapshot, pending_form_id, pending_tool_call_id)) class _MonotonicClock: @@ -501,8 +470,8 @@ def _soul() -> AgentSoulConfig: ) -def _dify_ctx() -> Any: - return SimpleNamespace( +def _dify_ctx() -> DifyRunContext: + return DifyRunContext( tenant_id="tenant-1", app_id="app-1", user_id="user-1", @@ -529,18 +498,18 @@ def _runner( ) -def _run(runner: AgentAppRunner, qm: _FakeQueueManager, *, agent_runtime_exit_intent: str = "suspend") -> None: +def _run(runner: AgentAppRunner, qm: _FakeQueueManager) -> None: runner.run( dify_context=_dify_ctx(), agent_id="agent-1", agent_config_snapshot_id="snap-1", agent_soul=_soul(), + home_snapshot_id="home-1", conversation_id="conv-1", query="hello", message_id="msg-1", model_name="gpt-4o-mini", queue_manager=qm, # type: ignore[arg-type] - agent_runtime_exit_intent=agent_runtime_exit_intent, # type: ignore[arg-type] ) @@ -548,9 +517,14 @@ def _message_end(qm: _FakeQueueManager) -> QueueMessageEndEvent: return next(e for e in qm.events if isinstance(e, QueueMessageEndEvent)) -def _saved_user_query(qm: _FakeQueueManager) -> str: +def _llm_result(qm: _FakeQueueManager) -> LLMResult: llm_result = _message_end(qm).llm_result assert llm_result is not None + return llm_result + + +def _saved_user_query(qm: _FakeQueueManager) -> str: + llm_result = _llm_result(qm) prompt_messages = llm_result.prompt_messages assert len(prompt_messages) == 1 content = prompt_messages[0].content @@ -558,7 +532,7 @@ def _saved_user_query(qm: _FakeQueueManager) -> str: return content -def test_successful_turn_publishes_chunk_and_message_end_and_saves_session(): +def test_successful_turn_publishes_chunk_and_message_end_and_saves_session() -> None: client = FakeAgentBackendRunClient() # SUCCESS: output {"text": "hello agent"} store = _FakeSessionStore() qm = _FakeQueueManager() @@ -573,144 +547,37 @@ def test_successful_turn_publishes_chunk_and_message_end_and_saves_session(): assert len(chunk_events) == 1 assert len(end_events) == 1 assert chunk_events[0].chunk.delta.message.content == "hello agent" - assert end_events[0].llm_result.message.content == "hello agent" - assert end_events[0].llm_result.model == "gpt-4o-mini" + assert _llm_result(qm).message.content == "hello agent" + assert _llm_result(qm).model == "gpt-4o-mini" assert _saved_user_query(qm) == "hello" # The conversation session snapshot is persisted for multi-turn continuity. assert store.saved - saved_scope, saved_run_id, saved_snapshot, saved_specs, pending_form_id, pending_tool_call_id = store.saved[0] + saved_scope, saved_binding_id, saved_snapshot, pending_form_id, pending_tool_call_id = store.saved[0] assert saved_scope.conversation_id == "conv-1" assert saved_scope.agent_config_snapshot_id == "snap-1" - assert saved_run_id == "fake-run-1" + assert saved_binding_id == "binding-1" assert saved_snapshot is not None - assert saved_specs # A successful turn carries no ask_human pause correlation. assert pending_form_id is None assert pending_tool_call_id is None - assert store.cleaned == [] -def test_successful_turn_enqueues_cleanup_for_superseded_sessions_after_saving_snapshot(monkeypatch): - superseded = StoredAgentAppSession( - scope=AgentAppSessionScope( - tenant_id="tenant-1", - app_id="app-1", - conversation_id="conv-1", - agent_id="agent-2", - agent_config_snapshot_id="snap-2", - ), - session_snapshot=CompositorSessionSnapshot(layers=[]), - backend_run_id="run-old", - runtime_layer_specs=[RuntimeLayerSpec(name="history", type="pydantic_ai.history")], - ) - current_scope_session = StoredAgentAppSession( - scope=AgentAppSessionScope( - tenant_id="tenant-1", - app_id="app-1", - conversation_id="conv-1", - agent_id="agent-1", - agent_config_snapshot_id="snap-1", - ), - session_snapshot=CompositorSessionSnapshot(layers=[]), - backend_run_id="run-current", - runtime_layer_specs=[RuntimeLayerSpec(name="history", type="pydantic_ai.history")], - ) - store = _FakeSessionStore(listed_sessions=[current_scope_session, superseded]) +def test_turn_uses_resolved_backend_binding_before_backend_invocation() -> None: client = FakeAgentBackendRunClient() - qm = _FakeQueueManager() - cleanup_delay = MagicMock() - monkeypatch.setattr(app_runner_module.cleanup_conversation_agent_runtime_session, "delay", cleanup_delay) + store = _FakeSessionStore(binding_id="binding-2", backend_binding_ref="backend-binding-2") - _run(_runner(client, store), qm) - - assert store.saved - cleanup_delay.assert_called_once() - payload = cleanup_delay.call_args.args[0] - assert payload["metadata"]["conversation_id"] == "conv-1" - assert payload["metadata"]["agent_id"] == "agent-2" - assert payload["metadata"]["previous_agent_backend_run_id"] == "run-old" - assert payload["idempotency_key"] == "tenant-1:app-1:conv-1:agent-2:snap-2:superseded-session-cleanup:run-old" - - -def test_superseded_session_cleanup_enqueue_failure_does_not_fail_turn(monkeypatch): - superseded = StoredAgentAppSession( - scope=AgentAppSessionScope( - tenant_id="tenant-1", - app_id="app-1", - conversation_id="conv-1", - agent_id="agent-2", - agent_config_snapshot_id="snap-2", - ), - session_snapshot=CompositorSessionSnapshot(layers=[]), - backend_run_id="run-old", - runtime_layer_specs=[RuntimeLayerSpec(name="history", type="pydantic_ai.history")], - ) - store = _FakeSessionStore(listed_sessions=[superseded]) - client = FakeAgentBackendRunClient() - qm = _FakeQueueManager() - cleanup_delay = MagicMock(side_effect=RuntimeError("queue down")) - monkeypatch.setattr(app_runner_module.cleanup_conversation_agent_runtime_session, "delay", cleanup_delay) - - _run(_runner(client, store), qm) - - cleanup_delay.assert_called_once() - end_events = [e for e in qm.events if isinstance(e, QueueMessageEndEvent)] - assert len(end_events) == 1 - assert end_events[0].llm_result.message.content == "hello agent" - - -def test_delete_on_exit_turn_marks_session_cleaned_without_saving_snapshot(): - client = _StreamingRecordingFakeAgentBackendRunClient() - store = _FakeSessionStore() - qm = _FakeQueueManager() - - _run(_runner(client, store), qm, agent_runtime_exit_intent="delete") + _run(_runner(client, store), _FakeQueueManager()) assert client.request is not None - assert client.request.on_exit.default.value == "delete" - assert store.saved == [] - assert len(store.cleaned) == 1 - cleaned_scope, cleaned_run_id = store.cleaned[0] - assert cleaned_scope.conversation_id == "conv-1" - assert cleaned_scope.agent_config_snapshot_id == "snap-1" - assert cleaned_run_id == "fake-run-1" - end_events = [e for e in qm.events if isinstance(e, QueueMessageEndEvent)] - assert len(end_events) == 1 - assert end_events[0].llm_result.message.content == "hello agent" + layers = {layer["name"]: layer for layer in client.request.model_dump(mode="json")["composition"]["layers"]} + assert layers["runtime"]["config"]["backend_binding_ref"] == "backend-binding-2" + assert store.saved[0][1] == "binding-2" + assert len(store.resolved_scopes) == 1 -def test_delete_on_exit_turn_swallows_cleanup_failure_after_success(): - client = _StreamingRecordingFakeAgentBackendRunClient() - store = _FakeSessionStore() - store.mark_cleaned = MagicMock(side_effect=RuntimeError("cleanup failed")) # type: ignore[method-assign] - qm = _FakeQueueManager() - - _run(_runner(client, store), qm, agent_runtime_exit_intent="delete") - - assert store.saved == [] - store.mark_cleaned.assert_called_once() - end_events = [e for e in qm.events if isinstance(e, QueueMessageEndEvent)] - assert len(end_events) == 1 - - -def test_delete_on_exit_turn_marks_session_cleaned_when_publish_fails(): - client = _StreamingRecordingFakeAgentBackendRunClient() - store = _FakeSessionStore() - store.mark_cleaned = MagicMock(side_effect=RuntimeError("cleanup failed")) # type: ignore[method-assign] - qm = _FakeQueueManager() - runner = _runner(client, store) - runner._publish_terminal_answer = MagicMock(side_effect=RuntimeError("publish failed")) - - with pytest.raises(RuntimeError, match="publish failed"): - _run(runner, qm, agent_runtime_exit_intent="delete") - - assert store.saved == [] - store.mark_cleaned.assert_called_once() - - -def test_successful_turn_routes_stream_text_to_agent_message_and_uses_terminal_output(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_successful_turn_routes_stream_text_to_agent_message_and_uses_terminal_output( + sqlite_session: Session, +) -> None: client = _StreamingFakeAgentBackendRunClient() store = _FakeSessionStore() qm = _FakeQueueManager() @@ -723,22 +590,21 @@ def test_successful_turn_routes_stream_text_to_agent_message_and_uses_terminal_o assert [event.chunk.delta.message.content for event in chunk_events] == ["hello agent"] assert [event.chunk.delta.message.content for event in agent_message_events] == ["hello ", "agent"] assert len(end_events) == 1 - assert end_events[0].llm_result.message.content == "hello agent" - assert end_events[0].llm_result.usage.prompt_tokens == 3 - assert end_events[0].llm_result.usage.completion_tokens == 5 - assert end_events[0].llm_result.usage.total_tokens == 8 - assert end_events[0].llm_result.usage.prompt_price == Decimal("0.000015") - assert end_events[0].llm_result.usage.completion_price == Decimal("0.000150") - assert end_events[0].llm_result.usage.total_price == Decimal("0.000165") - assert end_events[0].llm_result.usage.currency == "USD" - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + llm_result = _llm_result(qm) + assert llm_result.message.content == "hello agent" + assert llm_result.usage.prompt_tokens == 3 + assert llm_result.usage.completion_tokens == 5 + assert llm_result.usage.total_tokens == 8 + assert llm_result.usage.prompt_price == Decimal("0.000015") + assert llm_result.usage.completion_price == Decimal("0.000150") + assert llm_result.usage.total_price == Decimal("0.000165") + assert llm_result.usage.currency == "USD" + rows = _thought_rows(sqlite_session) assert rows == [] assert store.saved -def test_successful_turn_routes_single_agent_message_delta(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_successful_turn_routes_single_agent_message_delta(sqlite_session: Session) -> None: client = _StreamingSingleAgentMessageDeltaFakeAgentBackendRunClient() store = _FakeSessionStore() qm = _FakeQueueManager() @@ -751,12 +617,42 @@ def test_successful_turn_routes_single_agent_message_delta(monkeypatch): assert [event.chunk.delta.message.content for event in chunk_events] == ["hello agent"] assert [event.chunk.delta.message.content for event in agent_message_events] == ["hello"] assert len(end_events) == 1 - assert end_events[0].llm_result.message.content == "hello agent" - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + assert _llm_result(qm).message.content == "hello agent" + rows = _thought_rows(sqlite_session) assert rows == [] -def test_successful_turn_with_null_terminal_output_publishes_empty_answer_not_literal_null(): +def test_thought_commit_failure_rolls_back_and_turn_continues(sqlite_session: Session) -> None: + rollback_events: list[Session] = [] + should_fail = True + + def fail_first_commit(_session: Session) -> None: + nonlocal should_fail + if should_fail: + should_fail = False + raise RuntimeError("forced thought commit failure") + + def record_rollback(session: Session) -> None: + rollback_events.append(session) + + event.listen(sqlite_session, "before_commit", fail_first_commit) + event.listen(sqlite_session, "after_rollback", record_rollback) + try: + client = _StreamingSingleAgentMessageDeltaFakeAgentBackendRunClient() + store = _FakeSessionStore() + qm = _FakeQueueManager() + + _run(_runner(client, store), qm) + finally: + event.remove(sqlite_session, "before_commit", fail_first_commit) + event.remove(sqlite_session, "after_rollback", record_rollback) + + assert rollback_events == [sqlite_session] + assert _thought_rows(sqlite_session) == [] + assert _llm_result(qm).message.content == "hello agent" + + +def test_successful_turn_with_null_terminal_output_publishes_empty_answer_not_literal_null() -> None: client = _NullOutputFakeAgentBackendRunClient() store = _FakeSessionStore() qm = _FakeQueueManager() @@ -769,12 +665,12 @@ def test_successful_turn_with_null_terminal_output_publishes_empty_answer_not_li assert chunk_events == [] assert agent_message_events == [] assert len(end_events) == 1 - assert end_events[0].llm_result.message.content == "" + assert _llm_result(qm).message.content == "" -def test_successful_turn_with_stream_text_and_null_terminal_output_keeps_empty_message(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_successful_turn_with_stream_text_and_null_terminal_output_keeps_empty_message( + sqlite_session: Session, +) -> None: client = _StreamingTextNullOutputFakeAgentBackendRunClient() store = _FakeSessionStore() qm = _FakeQueueManager() @@ -787,15 +683,13 @@ def test_successful_turn_with_stream_text_and_null_terminal_output_keeps_empty_m assert chunk_events == [] assert [event.chunk.delta.message.content for event in agent_message_events] == ["streamed answer"] assert len(end_events) == 1 - assert end_events[0].llm_result.message.content == "" - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + assert _llm_result(qm).message.content == "" + rows = _thought_rows(sqlite_session) assert len(rows) == 1 assert rows[0].answer == "streamed answer" -def test_successful_turn_routes_agent_answer_to_agent_message(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_successful_turn_routes_agent_answer_to_agent_message(sqlite_session: Session) -> None: client = _AgentAnswerStreamingFakeAgentBackendRunClient() store = _FakeSessionStore() qm = _FakeQueueManager() @@ -808,21 +702,21 @@ def test_successful_turn_routes_agent_answer_to_agent_message(monkeypatch): assert [event.chunk.delta.message.content for event in agent_message_events] == ["hello ", "agent"] end_events = [e for e in qm.events if isinstance(e, QueueMessageEndEvent)] assert len(end_events) == 1 - assert end_events[0].llm_result.message.content == "final answer" + assert _llm_result(qm).message.content == "final answer" thought_events = [e for e in qm.events if isinstance(e, QueueAgentThoughtEvent)] assert len(thought_events) == 2 - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert len(rows) == 1 assert rows[0].answer == "hello agent" assert rows[0].thought == "" assert rows[0].tool == "" -def test_agent_message_deltas_are_debounced_to_agent_message(monkeypatch): +def test_agent_message_deltas_are_debounced_to_agent_message( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: monkeypatch.setattr(app_runner_module.time, "monotonic", _MonotonicClock(0.0, 0.2)) - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) client = _StreamingFakeAgentBackendRunClient() store = _FakeSessionStore() qm = _FakeQueueManager() @@ -833,13 +727,13 @@ def test_agent_message_deltas_are_debounced_to_agent_message(monkeypatch): agent_message_events = [e for e in qm.events if isinstance(e, QueueAgentMessageEvent)] assert [event.chunk.delta.message.content for event in chunk_events] == ["hello agent"] assert [event.chunk.delta.message.content for event in agent_message_events] == ["hello agent"] - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert rows == [] -def test_successful_turn_persists_thinking_and_tool_process_events(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_successful_turn_persists_thinking_and_tool_process_events( + sqlite_session: Session, +) -> None: client = _ProcessStreamingFakeAgentBackendRunClient() store = _FakeSessionStore() qm = _FakeQueueManager() @@ -853,7 +747,7 @@ def test_successful_turn_persists_thinking_and_tool_process_events(monkeypatch): thought_events = [e for e in qm.events if isinstance(e, QueueAgentThoughtEvent)] assert len(thought_events) >= 3 - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert rows[0].thought == "I need to inspect the file." assert rows[0].tool == "" assert rows[1].tool == "bash" @@ -862,9 +756,9 @@ def test_successful_turn_persists_thinking_and_tool_process_events(monkeypatch): assert len(rows) == 2 -def test_streaming_turn_cancels_after_persisting_seen_agent_answer(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_streaming_turn_cancels_after_persisting_seen_agent_answer( + sqlite_session: Session, +) -> None: store = _FakeSessionStore() qm = _FakeQueueManager() client = _StreamingStopAfterFirstDeltaFakeAgentBackendRunClient(queue_manager=qm) @@ -876,15 +770,15 @@ def test_streaming_turn_cancels_after_persisting_seen_agent_answer(monkeypatch): agent_message_events = [e for e in qm.events if isinstance(e, QueueAgentMessageEvent)] assert chunk_events == [] assert [event.chunk.delta.message.content for event in agent_message_events] == ["hello "] - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert len(rows) == 1 assert rows[0].answer == "hello " assert client.cancelled_run_ids == ["fake-run-1"] -def test_tool_result_without_identity_does_not_attach_to_previous_tool(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_tool_result_without_identity_does_not_attach_to_previous_tool( + sqlite_session: Session, +) -> None: qm = _FakeQueueManager() recorder = app_runner_module._AgentProcessRecorder( dify_context=_dify_ctx(), @@ -916,7 +810,7 @@ def test_tool_result_without_identity_does_not_attach_to_previous_tool(monkeypat ) ) - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert len(rows) == 2 assert rows[0].tool == "shell_run" assert rows[0].tool_input == '{"script": "npx skills find browser"}' @@ -926,9 +820,7 @@ def test_tool_result_without_identity_does_not_attach_to_previous_tool(monkeypat assert rows[1].observation == "Knowledge base search results: browser skill" -def test_answer_suffix_trim_keeps_non_terminal_prefix(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_answer_suffix_trim_keeps_non_terminal_prefix(sqlite_session: Session) -> None: qm = _FakeQueueManager() recorder = app_runner_module._AgentProcessRecorder( dify_context=_dify_ctx(), @@ -939,14 +831,12 @@ def test_answer_suffix_trim_keeps_non_terminal_prefix(monkeypatch): recorder.append_answer_text("intermediate final answer") recorder.trim_answer_suffix("final answer") - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert len(rows) == 1 assert rows[0].answer == "intermediate " -def test_tool_call_part_binds_late_call_id_to_delta_row(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_tool_call_part_binds_late_call_id_to_delta_row(sqlite_session: Session) -> None: qm = _FakeQueueManager() recorder = app_runner_module._AgentProcessRecorder( dify_context=_dify_ctx(), @@ -998,16 +888,14 @@ def test_tool_call_part_binds_late_call_id_to_delta_row(monkeypatch): ) ) - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert len(rows) == 1 assert rows[0].tool == "knowledge_base_search" assert rows[0].tool_input == '{"query": "browser"}' assert rows[0].observation == "Knowledge base search results: browser skill" -def test_thinking_after_tool_starts_new_snapshot_row(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_thinking_after_tool_starts_new_snapshot_row(sqlite_session: Session) -> None: qm = _FakeQueueManager() recorder = app_runner_module._AgentProcessRecorder( dify_context=_dify_ctx(), @@ -1056,16 +944,16 @@ def test_thinking_after_tool_starts_new_snapshot_row(monkeypatch): ) ) - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert [row.thought for row in rows] == ["The first thought.", "", "The next thought."] assert rows[0].id != rows[2].id assert rows[1].tool == "shell_run" assert rows[1].tool_input == '{"cmd": "date"}' -def test_tool_result_without_call_id_matches_unique_open_tool_name(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_tool_result_without_call_id_matches_unique_open_tool_name( + sqlite_session: Session, +) -> None: qm = _FakeQueueManager() recorder = app_runner_module._AgentProcessRecorder( dify_context=_dify_ctx(), @@ -1100,16 +988,16 @@ def test_tool_result_without_call_id_matches_unique_open_tool_name(monkeypatch): ) ) - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert len(rows) == 1 assert rows[0].tool == "knowledge_base_search" assert rows[0].tool_input == '{"query": "browser"}' assert rows[0].observation == "Knowledge base search results: browser skill" -def test_repeated_tool_calls_without_call_id_or_index_create_distinct_rows(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_repeated_tool_calls_without_call_id_or_index_create_distinct_rows( + sqlite_session: Session, +) -> None: qm = _FakeQueueManager() recorder = app_runner_module._AgentProcessRecorder( dify_context=_dify_ctx(), @@ -1170,7 +1058,7 @@ def test_repeated_tool_calls_without_call_id_or_index_create_distinct_rows(monke ) ) - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert len(rows) == 2 assert rows[0].tool == "shell_run" assert rows[0].tool_input == '{"script": "lookup find"}' @@ -1180,9 +1068,9 @@ def test_repeated_tool_calls_without_call_id_or_index_create_distinct_rows(monke assert rows[1].observation == "out output" -def test_repeated_tool_calls_with_placeholder_call_id_and_reused_index_create_distinct_rows(monkeypatch): - fake_session = _FakeDbSession() - monkeypatch.setattr(app_runner_module.db, "session", fake_session) +def test_repeated_tool_calls_with_placeholder_call_id_and_reused_index_create_distinct_rows( + sqlite_session: Session, +) -> None: qm = _FakeQueueManager() recorder = app_runner_module._AgentProcessRecorder( dify_context=_dify_ctx(), @@ -1221,7 +1109,7 @@ def test_repeated_tool_calls_with_placeholder_call_id_and_reused_index_create_di ) ) - rows = sorted(fake_session.rows.values(), key=lambda row: row.position) + rows = _thought_rows(sqlite_session) assert len(rows) == 2 assert rows[0].tool == "shell_run" assert rows[0].tool_input == '{"script": "lookup find"}' @@ -1231,7 +1119,7 @@ def test_repeated_tool_calls_with_placeholder_call_id_and_reused_index_create_di assert rows[1].observation == "out output" -def test_prior_session_snapshot_is_threaded_into_request(): +def test_prior_session_snapshot_is_threaded_into_request() -> None: prior = CompositorSessionSnapshot(layers=[]) client = FakeAgentBackendRunClient() store = _FakeSessionStore(loaded=prior) @@ -1243,7 +1131,7 @@ def test_prior_session_snapshot_is_threaded_into_request(): assert client.request.session_snapshot is prior -def test_debug_session_scope_can_reuse_conversation_across_config_snapshots(): +def test_debug_session_scope_can_reuse_conversation_across_config_snapshots() -> None: prior = CompositorSessionSnapshot(layers=[]) client = FakeAgentBackendRunClient() store = _FakeSessionStore(loaded=prior) @@ -1254,6 +1142,7 @@ def test_debug_session_scope_can_reuse_conversation_across_config_snapshots(): agent_id="agent-1", agent_config_snapshot_id="snap-new", agent_soul=_soul(), + home_snapshot_id="home-1", conversation_id="conv-1", query="hello", message_id="msg-1", @@ -1264,11 +1153,11 @@ def test_debug_session_scope_can_reuse_conversation_across_config_snapshots(): assert client.request is not None assert client.request.session_snapshot is prior - assert store.loaded_scopes[0].agent_config_snapshot_id is None - assert store.saved[0][0].agent_config_snapshot_id is None + assert store.resolved_scopes[0].agent_config_snapshot_id == "snap-new" + assert store.saved[0][0].agent_config_snapshot_id == "snap-new" -def test_failed_run_raises_agent_backend_error(): +def test_failed_run_raises_agent_backend_error() -> None: client = FakeAgentBackendRunClient(scenario=FakeAgentBackendScenario.FAILED) store = _FakeSessionStore() qm = _FakeQueueManager() @@ -1280,7 +1169,18 @@ def test_failed_run_raises_agent_backend_error(): assert store.saved == [] -def test_agent_backend_failure_to_exception_maps_rate_limit_reason(): +def test_failed_run_prefers_run_failure_type_over_binding_lost_reason() -> None: + client = _RunLimitBindingLostFakeAgentBackendRunClient() + store = _FakeSessionStore() + + with pytest.raises(AgentBackendRunFailedError) as raised: + _run(_runner(client, store), _FakeQueueManager()) + + assert raised.value.error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + assert raised.value.reason == "binding_lost" + + +def test_agent_backend_failure_to_exception_maps_rate_limit_reason() -> None: err = app_runner_module._agent_backend_failure_to_exception( AgentBackendRunFailedInternalEvent( run_id="run-1", @@ -1293,7 +1193,7 @@ def test_agent_backend_failure_to_exception_maps_rate_limit_reason(): assert str(err) == "quota exceeded" -def test_agent_backend_failure_to_exception_preserves_unknown_reason_context(): +def test_agent_backend_failure_to_exception_preserves_unknown_reason_context() -> None: err = app_runner_module._agent_backend_failure_to_exception( AgentBackendRunFailedInternalEvent( run_id="run-1", @@ -1315,7 +1215,27 @@ def test_agent_backend_failure_to_exception_preserves_unknown_reason_context(): assert str(err) == "Knowledge retrieval failed (agent_run_id=run-1)" -def test_stopped_task_cancels_agent_backend_run_and_skips_session_save(): +def test_agent_backend_failure_to_exception_prefers_run_failure_type_over_known_reason() -> None: + err = app_runner_module._agent_backend_failure_to_exception( + AgentBackendRunFailedInternalEvent( + run_id="run-1", + error="run limit reached", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + reason="InvokeRateLimitError", + ) + ) + + assert isinstance(err, AgentBackendRunFailedError) + assert err.error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + assert err.reason == "InvokeRateLimitError" + assert err.detail == { + "error": "run limit reached", + "reason": "InvokeRateLimitError", + "source_event_id": None, + } + + +def test_stopped_task_cancels_agent_backend_run_and_skips_session_save() -> None: client = _RecordingFakeAgentBackendRunClient() store = _FakeSessionStore() qm = _StoppedQueueManager() @@ -1327,14 +1247,14 @@ def test_stopped_task_cancels_agent_backend_run_and_skips_session_save(): assert store.saved == [] -def test_terminal_output_to_answer_handles_plain_string_and_dict(): +def test_terminal_output_to_answer_handles_plain_string_and_dict() -> None: assert AgentAppRunner._terminal_output_to_answer(None) == "" assert AgentAppRunner._terminal_output_to_answer("plain text") == "plain text" assert AgentAppRunner._terminal_output_to_answer({"text": "hi"}) == "hi" assert AgentAppRunner._terminal_output_to_answer({"a": 1}) == '{"a": 1}' -def test_ask_human_pauses_turn_creates_form_and_persists_correlation(): +def test_ask_human_pauses_turn_creates_form_and_persists_correlation() -> None: # ENG-635/637: the PAUSED scenario emits a dify.ask_human deferred call, so # the chat turn ends by creating a conversation-owned HITL form + saving the # pause correlation, instead of crashing. Stub the form repo (DB-free). @@ -1358,27 +1278,11 @@ def test_ask_human_pauses_turn_creates_form_and_persists_correlation(): assert _saved_user_query(qm) == "hello" # The pause correlation is persisted so a form submission can resume the run. assert store.saved - assert store.saved[0][4] == "form-1" - assert store.saved[0][5] == "fake-ask-human-1" + assert store.saved[0][3] == "form-1" + assert store.saved[0][4] == "fake-ask-human-1" -def test_delete_on_exit_deferred_tool_marks_session_cleaned_and_raises_error(): - client = FakeAgentBackendRunClient(scenario=FakeAgentBackendScenario.PAUSED) - store = _FakeSessionStore() - store.mark_cleaned = MagicMock(side_effect=RuntimeError("cleanup failed")) # type: ignore[method-assign] - qm = _FakeQueueManager() - runner = _runner(client, store) - runner._pause_for_ask_human = MagicMock() - - with pytest.raises(AgentBackendError, match="finalization cannot pause for human input"): - _run(runner, qm, agent_runtime_exit_intent="delete") - - runner._pause_for_ask_human.assert_not_called() - assert store.saved == [] - store.mark_cleaned.assert_called_once() - - -def test_submitted_form_resumes_turn_with_deferred_tool_results(monkeypatch): +def test_submitted_form_resumes_turn_with_deferred_tool_results(monkeypatch: pytest.MonkeyPatch) -> None: # ENG-638: a turn that runs while a pending form is answered threads the # human's reply into the request as deferred_tool_results. snapshot = CompositorSessionSnapshot(layers=[]) @@ -1389,9 +1293,12 @@ def test_submitted_form_resumes_turn_with_deferred_tool_results(monkeypatch): conversation_id="conv-1", agent_id="agent-1", agent_config_snapshot_id="snap-1", + home_snapshot_id="home-1", ), + binding_id="binding-1", + workspace_id="workspace-1", + backend_binding_ref="backend-binding-1", session_snapshot=snapshot, - backend_run_id="run-0", pending_form_id="form-1", pending_tool_call_id="call-1", ) diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_input_guards.py b/api/tests/unit_tests/core/app/apps/agent_app/test_input_guards.py index ec2706a2068..dc14fb9895c 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_input_guards.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_input_guards.py @@ -10,9 +10,9 @@ from __future__ import annotations from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session import core.app.features.annotation_reply.annotation_reply as annotation_mod import core.moderation.input_moderation as input_moderation_mod @@ -75,13 +75,13 @@ def _saved_user_query(events: list[Any]) -> str: class TestRunInputGuards: - def test_no_guards_passes_through(self, monkeypatch: pytest.MonkeyPatch): + def test_no_guards_passes_through(self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): _patch_moderation(monkeypatch, returns=(False, {}, "hello")) _patch_annotation(monkeypatch, reply=None) qm = _FakeQueueManager() handled, query, annotation_reply = AgentAppGenerator()._run_input_guards( - session=MagicMock(), + session=sqlite_session, application_generate_entity=_make_entity("hello"), app_model=SimpleNamespace(id="app-1"), message=SimpleNamespace(id="msg-1"), @@ -93,13 +93,13 @@ class TestRunInputGuards: assert annotation_reply is None assert qm.events == [] - def test_moderation_override_sanitizes_query(self, monkeypatch: pytest.MonkeyPatch): + def test_moderation_override_sanitizes_query(self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): _patch_moderation(monkeypatch, returns=(True, {}, "[redacted]")) _patch_annotation(monkeypatch, reply=None) qm = _FakeQueueManager() handled, query, annotation_reply = AgentAppGenerator()._run_input_guards( - session=MagicMock(), + session=sqlite_session, application_generate_entity=_make_entity("leak my secret"), app_model=SimpleNamespace(id="app-1"), message=SimpleNamespace(id="msg-1"), @@ -111,13 +111,13 @@ class TestRunInputGuards: assert annotation_reply is None assert qm.events == [] - def test_moderation_block_short_circuits(self, monkeypatch: pytest.MonkeyPatch): + def test_moderation_block_short_circuits(self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): _patch_moderation(monkeypatch, raises=ModerationError("blocked preset answer")) _patch_annotation(monkeypatch, reply=None) qm = _FakeQueueManager() handled, _, annotation_reply = AgentAppGenerator()._run_input_guards( - session=MagicMock(), + session=sqlite_session, application_generate_entity=_make_entity("forbidden"), app_model=SimpleNamespace(id="app-1"), message=SimpleNamespace(id="msg-1"), @@ -130,13 +130,13 @@ class TestRunInputGuards: assert _answer_text(qm.events) == "blocked preset answer" assert _saved_user_query(qm.events) == "forbidden" - def test_annotation_hit_short_circuits(self, monkeypatch: pytest.MonkeyPatch): + def test_annotation_hit_short_circuits(self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): _patch_moderation(monkeypatch, returns=(False, {}, "what is your name")) _patch_annotation(monkeypatch, reply=SimpleNamespace(id="anno-1", content="I am the annotated Iris.")) qm = _FakeQueueManager() handled, _, annotation_reply = AgentAppGenerator()._run_input_guards( - session=MagicMock(), + session=sqlite_session, application_generate_entity=_make_entity("what is your name"), app_model=SimpleNamespace(id="app-1"), message=SimpleNamespace(id="msg-1"), diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py b/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py index f65770665e4..37ea3338000 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_resolve_agent.py @@ -9,13 +9,19 @@ from __future__ import annotations from types import SimpleNamespace from typing import Any +from unittest.mock import MagicMock import pytest +from core.app.apps.agent_app import app_generator 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 AgentConfigDraft, AgentConfigDraftType, AgentScope, AgentSource +from models.agent import AgentConfigDraft, AgentConfigDraftType, AgentConfigVersionKind, AgentScope, AgentSource from models.agent_config_entities import AgentSoulConfig +from services.agent.workspace_service import ( + AgentWorkspaceBindingGenerationMismatchError, + AgentWorkspaceService, +) _SOUL_DICT = { "model": { @@ -34,8 +40,10 @@ class _FakeScalarSession: self._values = list(values) self.added: list[Any] = [] self.flush_count = 0 + self.scalar_statements: list[Any] = [] - def scalar(self, _stmt: Any) -> Any: + def scalar(self, stmt: Any) -> Any: + self.scalar_statements.append(stmt) return self._values.pop(0) if self._values else None def add(self, value: Any) -> None: @@ -46,7 +54,7 @@ class _FakeScalarSession: def _snapshot() -> SimpleNamespace: - return SimpleNamespace(id="snap-1", config_snapshot_dict=_SOUL_DICT) + return SimpleNamespace(id="snap-1", home_snapshot_id="home-1", config_snapshot_dict=_SOUL_DICT) class TestResolveAgentById: @@ -102,6 +110,7 @@ class TestResolveDebugDraft: tenant_id="t1", agent=agent, draft_type=None, + draft_id=None, account_id=None, session=session, ) @@ -127,10 +136,12 @@ class TestResolveDebugDraft: account_id=None, draft_owner_key="", base_snapshot_id="snap-1", + home_snapshot_id="home-1", config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "old"}}), ) active_snapshot = SimpleNamespace( id="snap-2", + home_snapshot_id="home-2", config_snapshot_dict={"prompt": {"system_prompt": "new"}}, ) session = _FakeScalarSession([draft, active_snapshot]) @@ -146,6 +157,7 @@ class TestResolveDebugDraft: assert resolved is draft assert resolved.id == "draft-1" assert resolved.base_snapshot_id == "snap-2" + assert resolved.home_snapshot_id == "home-2" assert resolved.config_snapshot_dict["prompt"]["system_prompt"] == "new" assert session.flush_count == 1 @@ -165,6 +177,7 @@ class TestResolveDebugDraft: account_id="account-1", draft_owner_key="account-1", base_snapshot_id="snap-1", + home_snapshot_id="home-build", config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "build edit"}}), ) session = _FakeScalarSession([draft]) @@ -182,8 +195,63 @@ class TestResolveDebugDraft: assert resolved.config_snapshot_dict["prompt"]["system_prompt"] == "build edit" assert session.flush_count == 0 + def test_build_draft_uses_exact_draft_id(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="exact-build-draft", + 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", + home_snapshot_id="home-build", + config_snapshot=AgentSoulConfig(), + ) + session = _FakeScalarSession([draft]) + statements: list[Any] = [] + scalar = session.scalar + + def capture_scalar(statement: Any) -> Any: + statements.append(statement) + return scalar(statement) + + session.scalar = capture_scalar # type: ignore[method-assign] + + resolved = AgentAppGenerator._resolve_debug_draft( + tenant_id="t1", + agent=agent, + draft_type=AgentConfigDraftType.DEBUG_BUILD.value, + draft_id="exact-build-draft", + account_id="account-1", + session=session, + ) + + assert resolved is draft + assert "agent_config_drafts.id =" in str(statements[0]) + assert "exact-build-draft" in statements[0].compile().params.values() + class TestResolveAgent: + @pytest.fixture(autouse=True) + def _publish_visibility(self, monkeypatch: pytest.MonkeyPatch) -> None: + def is_publish_visible(*, agent: SimpleNamespace, **_kwargs: object) -> bool: + if "publish_visible" in vars(agent): + return bool(agent.publish_visible) + return bool(agent.active_config_is_published) + + monkeypatch.setattr( + app_generator, + "agent_has_workflow_callable_active_snapshot", + is_publish_visible, + ) + def test_success_chains_to_resolve_by_id(self): bound_agent = SimpleNamespace( id="agent-1", @@ -216,6 +284,7 @@ class TestResolveAgent: source=AgentSource.AGENT_APP, active_config_snapshot_id="snap-1", active_config_is_published=False, + publish_visible=True, ) inner_agent = SimpleNamespace(id="agent-1") snapshot = _snapshot() @@ -235,6 +304,126 @@ class TestResolveAgent: assert config_version_kind == "snapshot" assert soul.prompt.system_prompt == "You are Iris." + def test_existing_conversation_resolves_binding_snapshot_instead_of_latest_active_snapshot( + self, monkeypatch: pytest.MonkeyPatch + ): + bound_agent = SimpleNamespace( + id="agent-1", + source=AgentSource.AGENT_APP, + active_config_snapshot_id="snap-2", + active_config_is_published=True, + ) + inner_agent = SimpleNamespace(id="agent-1") + pinned_snapshot = SimpleNamespace( + id="snap-1", + home_snapshot_id="home-1", + config_snapshot_dict=_SOUL_DICT, + ) + conversation = SimpleNamespace(id="conversation-1", agent_workspace_binding_id="binding-1") + binding = SimpleNamespace( + id="binding-1", + agent_id="agent-1", + base_home_snapshot_id="home-1", + agent_config_version_id="snap-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + get_active_binding = MagicMock(return_value=binding) + validate_generation = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active_binding) + monkeypatch.setattr(AgentWorkspaceService, "validate_binding_generation", validate_generation) + session = _FakeScalarSession([bound_agent, inner_agent, pinned_snapshot]) + app_model = SimpleNamespace(id="app-1", tenant_id="t1") + + _, config_id, config_version_kind, soul = AgentAppGenerator()._resolve_agent( + app_model, + invoke_from=InvokeFrom.WEB_APP, + draft_type=None, + user=SimpleNamespace(id="user-1"), + session=session, + conversation=conversation, + ) # type: ignore[arg-type] + + assert config_id == "snap-1" + assert config_version_kind == "snapshot" + assert soul.prompt.system_prompt == "You are Iris." + assert get_active_binding.call_args.kwargs["binding_id"] == "binding-1" + validate_generation.assert_called_once_with( + binding, + base_home_snapshot_id="home-1", + agent_config_version_id="snap-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + + def test_existing_conversation_rejects_unavailable_binding(self, monkeypatch: pytest.MonkeyPatch): + bound_agent = SimpleNamespace( + id="agent-1", + source=AgentSource.AGENT_APP, + active_config_snapshot_id="snap-active", + active_config_is_published=True, + ) + conversation = SimpleNamespace(id="conversation-1", agent_workspace_binding_id="binding-missing") + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=None)) + + with pytest.raises(AgentAppGeneratorError, match="Conversation participant Binding is unavailable"): + AgentAppGenerator()._resolve_agent( + SimpleNamespace(id="app-1", tenant_id="t1"), + invoke_from=InvokeFrom.WEB_APP, + draft_type=None, + user=SimpleNamespace(id="user-1"), + session=_FakeScalarSession([bound_agent]), + conversation=conversation, + ) # type: ignore[arg-type] + + @pytest.mark.parametrize( + ("binding_home_id", "binding_version_kind", "snapshot_home_id"), + [ + ("home-binding", AgentConfigVersionKind.SNAPSHOT, "home-other"), + ("home-pinned", AgentConfigVersionKind.DRAFT, "home-pinned"), + ], + ) + def test_existing_conversation_generation_mismatch_does_not_fallback_to_active_snapshot( + self, + monkeypatch: pytest.MonkeyPatch, + binding_home_id: str, + binding_version_kind: AgentConfigVersionKind, + snapshot_home_id: str, + ): + bound_agent = SimpleNamespace( + id="agent-1", + source=AgentSource.AGENT_APP, + active_config_snapshot_id="snap-active", + active_config_is_published=True, + ) + conversation = SimpleNamespace(id="conversation-1", agent_workspace_binding_id="binding-1") + binding = SimpleNamespace( + id="binding-1", + agent_id="agent-1", + base_home_snapshot_id=binding_home_id, + agent_config_version_id="snap-pinned", + agent_config_version_kind=binding_version_kind, + ) + pinned_snapshot = SimpleNamespace( + id="snap-pinned", + home_snapshot_id=snapshot_home_id, + config_snapshot_dict=_SOUL_DICT, + ) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=binding)) + session = _FakeScalarSession([bound_agent, SimpleNamespace(id="agent-1"), pinned_snapshot]) + + with pytest.raises(AgentWorkspaceBindingGenerationMismatchError): + AgentAppGenerator()._resolve_agent( + SimpleNamespace(id="app-1", tenant_id="t1"), + invoke_from=InvokeFrom.WEB_APP, + draft_type=None, + user=SimpleNamespace(id="user-1"), + session=session, + conversation=conversation, + ) # type: ignore[arg-type] + + snapshot_query_params = session.scalar_statements[-1].compile().params.values() + assert "snap-pinned" in snapshot_query_params + assert "snap-active" not in snapshot_query_params + def test_unpublished_imported_agent_is_not_available_to_public_runtime(self): bound_agent = SimpleNamespace( id="agent-1", @@ -254,6 +443,24 @@ class TestResolveAgent: session=session, ) # type: ignore[arg-type] + def test_never_published_agent_app_is_not_available_to_public_runtime(self): + bound_agent = SimpleNamespace( + id="agent-1", + source=AgentSource.AGENT_APP, + active_config_snapshot_id="snap-1", + active_config_is_published=False, + publish_visible=False, + ) + + with pytest.raises(AgentAppNotPublishedError, match="not been published"): + AgentAppGenerator()._resolve_agent( + SimpleNamespace(id="app-1", tenant_id="t1"), + invoke_from=InvokeFrom.WEB_APP, + draft_type=None, + user=SimpleNamespace(id="user-1"), + session=_FakeScalarSession([bound_agent]), + ) # type: ignore[arg-type] + def test_unpublished_imported_agent_remains_available_to_debugger(self): bound_agent = SimpleNamespace( id="agent-1", diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_runtime_request_builder.py b/api/tests/unit_tests/core/app/apps/agent_app/test_runtime_request_builder.py index b3ea05b8a87..a027fd31699 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_runtime_request_builder.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_runtime_request_builder.py @@ -44,6 +44,7 @@ class TestBuildForAgentApp: AgentBackendAgentAppRunInput( model=AgentBackendModelConfig(plugin_id="langgenius/openai", model_provider="openai", model="gpt-test"), execution_context=_exec_ctx(), + backend_binding_ref="binding-ref-1", user_prompt="hello", agent_soul_prompt="You are Iris.", ) @@ -65,6 +66,7 @@ class TestBuildForAgentApp: AgentBackendAgentAppRunInput( model=AgentBackendModelConfig(plugin_id="p/q", model_provider="openai", model="m"), execution_context=_exec_ctx(), + backend_binding_ref="binding-ref-1", user_prompt=" ", ) @@ -73,6 +75,7 @@ class TestBuildForAgentApp: AgentBackendAgentAppRunInput( model=AgentBackendModelConfig(plugin_id="langgenius/openai", model_provider="openai", model="gpt-test"), execution_context=_exec_ctx(), + backend_binding_ref="binding-ref-1", user_prompt="hi", ) ) @@ -142,7 +145,6 @@ def _ctx( *, query: str = "hello", agent_config_version_kind: str = "snapshot", - suspend_on_exit: bool = True, ) -> AgentAppRuntimeBuildContext: dify_context = SimpleNamespace( tenant_id="tenant-1", @@ -159,8 +161,9 @@ def _ctx( conversation_id="conv-1", user_query=query, idempotency_key="msg-1", + binding_id="binding-1", + backend_binding_ref="binding-ref-1", agent_config_version_kind=agent_config_version_kind, # type: ignore[arg-type] - suspend_on_exit=suspend_on_exit, ) @@ -191,6 +194,7 @@ class TestAgentAppRuntimeRequestBuilder: "agent_soul_prompt", "agent_app_user_prompt", "execution_context", + "runtime", DIFY_SHELL_LAYER_ID, DIFY_CONFIG_LAYER_ID, "history", @@ -243,16 +247,6 @@ class TestAgentAppRuntimeRequestBuilder: assert execution_context.config.agent_config_version_kind == "draft" assert config_layer.config.config_version.kind == "draft" - def test_build_uses_delete_on_exit_when_requested(self): - builder = AgentAppRuntimeRequestBuilder( - credentials_provider=_FakeCredentialsProvider(), - dify_tools_builder=_NoToolsBuilder(), # type: ignore[arg-type] - ) - - result = builder.build(_ctx(_soul_with_model(), suspend_on_exit=False)) - - assert result.request.on_exit.default.value == "delete" - def test_build_includes_plugin_tools_layer_returned_by_injected_builder_for_draft(self): soul = _soul_with_model() soul.tools.dify_tools = [ @@ -410,7 +404,7 @@ class TestAgentAppRuntimeRequestBuilder: ] assert shell_config["cli_tools"][0]["install_commands"] == ["apt-get install -y ripgrep"] assert shell_config["env"][0] == {"name": "PROJECT_NAME", "value": "demo"} - assert shell_config["sandbox"] == {"provider": "independent", "config": {"cpu": 2}} + assert "sandbox" not in shell_config assert result.metadata["agent_tools"] == { "dify_tool_count": 0, "dify_tool_names": [], @@ -461,7 +455,7 @@ class TestAgentAppConfigLayer: assert config.config.mentioned_file_names == [] # shell enters first; config uses that shell to materialize mentioned targets. names = [layer.name for layer in result.request.composition.layers] - assert names.index(DIFY_SHELL_LAYER_ID) == names.index("execution_context") + 1 + assert names.index(DIFY_SHELL_LAYER_ID) == names.index("execution_context") + 2 assert names.index(DIFY_CONFIG_LAYER_ID) == names.index(DIFY_SHELL_LAYER_ID) + 1 def test_config_layer_present_when_agent_soul_has_no_config_assets(self, monkeypatch: pytest.MonkeyPatch): @@ -484,7 +478,10 @@ class TestAgentAppConfigLayer: "mentioned_skill_names": [], "mentioned_file_names": [], } - assert layers[DIFY_SHELL_LAYER_ID].deps == {"execution_context": "execution_context"} + assert layers[DIFY_SHELL_LAYER_ID].deps == { + "execution_context": "execution_context", + "runtime": "runtime", + } assert layers[DIFY_SHELL_LAYER_ID].config.agent_stub_drive_ref is None def test_config_layer_for_build_draft_marks_config_writable(self): diff --git a/api/tests/unit_tests/core/app/apps/agent_app/test_session_store.py b/api/tests/unit_tests/core/app/apps/agent_app/test_session_store.py index fc9bf1b48e6..34109d343a3 100644 --- a/api/tests/unit_tests/core/app/apps/agent_app/test_session_store.py +++ b/api/tests/unit_tests/core/app/apps/agent_app/test_session_store.py @@ -1,348 +1,126 @@ -"""Unit tests for the conversation-keyed Agent App session store. - -Exercises the real ORM round-trip against the project's in-memory SQLite engine -(per-test create/drop of the unified ``agent_runtime_sessions`` table), so the -conversation owner path is verified without Postgres. -""" - -from __future__ import annotations - -from collections.abc import Generator +from types import SimpleNamespace +from unittest.mock import MagicMock import pytest from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot -from agenton.layers.base import LifecycleState -from dify_agent.protocol import RuntimeLayerSpec -from sqlalchemy import delete -from core.app.apps.agent_app.session_store import AgentAppRuntimeSessionStore, AgentAppSessionScope -from core.db.session_factory import session_factory -from models.agent import AgentRuntimeSession, AgentRuntimeSessionOwnerType, AgentRuntimeSessionStatus +from core.app.apps.agent_app.session_store import AgentAppSessionScope, AgentAppWorkspaceStore +from models.agent import ( + AgentConfigVersionKind, + AgentWorkspaceOwnerType, +) +from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService def _scope( - conversation_id: str = "conv-1", agent_id: str = "agent-1", agent_config_snapshot_id: str | None = "snap-1" + *, + kind: AgentConfigVersionKind = AgentConfigVersionKind.SNAPSHOT, + build_draft_id: str | None = None, + home_snapshot_id: str | None = "home-1", ) -> AgentAppSessionScope: return AgentAppSessionScope( tenant_id="tenant-1", app_id="app-1", - conversation_id=conversation_id, - agent_id=agent_id, - agent_config_snapshot_id=agent_config_snapshot_id, + conversation_id="conversation-1", + agent_id="agent-1", + agent_config_snapshot_id="config-1", + home_snapshot_id=home_snapshot_id, + agent_config_version_kind=kind, + build_draft_id=build_draft_id, ) -def _snapshot(messages: int = 1) -> CompositorSessionSnapshot: - return CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot( - name="history", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={"messages": [{"role": "user", "content": f"m{i}"} for i in range(messages)]}, - ) - ] +def _binding() -> SimpleNamespace: + return SimpleNamespace( + id="binding-1", + workspace_id="workspace-1", + backend_binding_ref="backend-binding-1", + agent_id="agent-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + base_home_snapshot_id="home-1", + session_snapshot=None, + pending_form_id=None, + pending_tool_call_id=None, ) -def _runtime_layer_specs() -> list[RuntimeLayerSpec]: - return [ - RuntimeLayerSpec(name="execution_context", type="dify.execution_context", config={"tenant_id": "tenant-1"}), - RuntimeLayerSpec(name="history", type="pydantic_ai.history"), - ] +def test_scope_selects_conversation_or_build_draft_workspace_owner() -> None: + assert _scope().workspace_owner.owner_type is AgentWorkspaceOwnerType.CONVERSATION + build_owner = _scope( + kind=AgentConfigVersionKind.BUILD_DRAFT, + build_draft_id="build-draft-1", + ).workspace_owner + assert build_owner.owner_type is AgentWorkspaceOwnerType.BUILD_DRAFT + assert build_owner.owner_id == "build-draft-1" -@pytest.fixture(autouse=True) -def _create_table() -> Generator[None, None, None]: - engine = session_factory.get_session_maker().kw["bind"] - AgentRuntimeSession.__table__.create(bind=engine, checkfirst=True) - yield - with session_factory.create_session() as session: - session.execute(delete(AgentRuntimeSession)) - session.commit() - AgentRuntimeSession.__table__.drop(bind=engine, checkfirst=True) +@pytest.mark.parametrize("home_snapshot_id", ["home-1", None]) +def test_load_or_create_persists_new_binding_on_caller(monkeypatch, home_snapshot_id: str | None) -> None: + caller = SimpleNamespace(agent_workspace_binding_id=None) + context = MagicMock() + session = context.__enter__.return_value + create = MagicMock(return_value=_binding()) + store = AgentAppWorkspaceStore() + monkeypatch.setattr("core.app.apps.agent_app.session_store.session_factory.create_session", lambda: context) + monkeypatch.setattr(store, "_load_caller", MagicMock(return_value=caller)) + monkeypatch.setattr(AgentWorkspaceService, "create_binding", create) + + stored = store.load_or_create(_scope(home_snapshot_id=home_snapshot_id)) + + assert stored.binding_id == "binding-1" + assert stored.workspace_id == "workspace-1" + assert stored.backend_binding_ref == "backend-binding-1" + assert caller.agent_workspace_binding_id == "binding-1" + assert create.call_args.kwargs["session"] is session + assert create.call_args.kwargs["base_home_snapshot_id"] == home_snapshot_id + session.commit.assert_called_once_with() -def test_load_returns_none_when_no_row(): - assert AgentAppRuntimeSessionStore().load_active_snapshot(_scope()) is None +def test_load_or_create_uses_exact_caller_binding(monkeypatch: pytest.MonkeyPatch) -> None: + caller = SimpleNamespace(agent_workspace_binding_id="binding-1") + context = MagicMock() + context.__enter__.return_value = MagicMock() + get_binding = MagicMock(return_value=_binding()) + create = MagicMock() + store = AgentAppWorkspaceStore() + monkeypatch.setattr("core.app.apps.agent_app.session_store.session_factory.create_session", lambda: context) + monkeypatch.setattr(store, "_load_caller", MagicMock(return_value=caller)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_binding) + monkeypatch.setattr(AgentWorkspaceService, "validate_binding_generation", MagicMock()) + monkeypatch.setattr(AgentWorkspaceService, "create_binding", create) + + stored = store.load_or_create(_scope()) + + assert stored.binding_id == "binding-1" + assert get_binding.call_args.kwargs["binding_id"] == "binding-1" + create.assert_not_called() -def test_save_creates_conversation_owned_row_and_round_trips(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(messages=2), - runtime_layer_specs=_runtime_layer_specs(), - ) +def test_normal_conversation_pointer_does_not_create_replacement_binding(monkeypatch: pytest.MonkeyPatch) -> None: + caller = SimpleNamespace(agent_workspace_binding_id="unavailable-binding") + context = MagicMock() + get_binding = MagicMock(return_value=None) + create = MagicMock() + store = AgentAppWorkspaceStore() + monkeypatch.setattr("core.app.apps.agent_app.session_store.session_factory.create_session", lambda: context) + monkeypatch.setattr(store, "_load_caller", MagicMock(return_value=caller)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_binding) + monkeypatch.setattr(AgentWorkspaceService, "create_binding", create) - loaded = store.load_active_snapshot(_scope()) - assert loaded is not None - assert loaded.layers[0].runtime_state["messages"] == [ - {"role": "user", "content": "m0"}, - {"role": "user", "content": "m1"}, - ] - with session_factory.create_session() as session: - row = session.query(AgentRuntimeSession).one() - assert row.owner_type == AgentRuntimeSessionOwnerType.CONVERSATION - assert row.conversation_id == "conv-1" - assert row.agent_config_snapshot_id == "snap-1" - assert row.workflow_run_id is None # conversation owner leaves workflow cols NULL - assert row.backend_run_id == "run-1" - assert "execution_context" in row.composition_layer_specs - assert "history" in row.composition_layer_specs + with pytest.raises(AgentWorkspaceNotFoundError, match="Caller participant Binding is unavailable"): + store.load_or_create(_scope()) + + assert get_binding.call_args.kwargs["binding_id"] == "unavailable-binding" + create.assert_not_called() -def test_save_is_noop_when_snapshot_missing(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-x", - snapshot=None, - runtime_layer_specs=_runtime_layer_specs(), - ) - with session_factory.create_session() as session: - assert session.query(AgentRuntimeSession).count() == 0 +def test_save_snapshot_targets_binding(monkeypatch: pytest.MonkeyPatch) -> None: + save = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "save_binding_session_snapshot", save) + snapshot = CompositorSessionSnapshot(layers=[]) + AgentAppWorkspaceStore().save_active_snapshot(scope=_scope(), binding_id="binding-1", snapshot=snapshot) -def test_second_turn_updates_same_conversation_row(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(messages=1), - runtime_layer_specs=_runtime_layer_specs(), - ) - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-2", - snapshot=_snapshot(messages=3), - runtime_layer_specs=_runtime_layer_specs(), - ) - with session_factory.create_session() as session: - rows = session.query(AgentRuntimeSession).all() - assert len(rows) == 1 - assert rows[0].backend_run_id == "run-2" - - -def test_debug_scope_with_null_snapshot_id_updates_same_conversation_row(): - store = AgentAppRuntimeSessionStore() - scope = _scope(agent_config_snapshot_id=None) - store.save_active_snapshot( - scope=scope, - backend_run_id="run-1", - snapshot=_snapshot(messages=1), - runtime_layer_specs=_runtime_layer_specs(), - ) - store.save_active_snapshot( - scope=scope, - backend_run_id="run-2", - snapshot=_snapshot(messages=3), - runtime_layer_specs=_runtime_layer_specs(), - ) - - loaded = store.load_active_snapshot(scope) - - assert loaded is not None - assert loaded.layers[0].runtime_state["messages"] == [ - {"role": "user", "content": "m0"}, - {"role": "user", "content": "m1"}, - {"role": "user", "content": "m2"}, - ] - with session_factory.create_session() as session: - row = session.query(AgentRuntimeSession).one() - assert row.agent_config_snapshot_id is None - assert row.backend_run_id == "run-2" - - -def test_mark_cleaned_then_load_returns_none_and_save_resurrects(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(), - runtime_layer_specs=_runtime_layer_specs(), - ) - store.mark_cleaned(scope=_scope(), backend_run_id="cleanup-1") - assert store.load_active_snapshot(_scope()) is None - # Re-entry revives the row. - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-2", - snapshot=_snapshot(messages=2), - runtime_layer_specs=_runtime_layer_specs(), - ) - with session_factory.create_session() as session: - row = session.query(AgentRuntimeSession).one() - assert row.status == AgentRuntimeSessionStatus.ACTIVE - assert row.cleaned_at is None - assert row.backend_run_id == "run-2" - - -def test_distinct_conversations_do_not_collide(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(conversation_id="conv-A"), - backend_run_id="a", - snapshot=_snapshot(), - runtime_layer_specs=_runtime_layer_specs(), - ) - store.save_active_snapshot( - scope=_scope(conversation_id="conv-B"), - backend_run_id="b", - snapshot=_snapshot(), - runtime_layer_specs=_runtime_layer_specs(), - ) - assert store.load_active_snapshot(_scope(conversation_id="conv-A")) is not None - assert store.load_active_snapshot(_scope(conversation_id="conv-B")) is not None - with session_factory.create_session() as session: - assert session.query(AgentRuntimeSession).count() == 2 - - -def test_distinct_agent_config_snapshots_keep_only_latest_active_session(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(agent_config_snapshot_id="snap-1"), - backend_run_id="a", - snapshot=_snapshot(), - runtime_layer_specs=_runtime_layer_specs(), - ) - store.save_active_snapshot( - scope=_scope(agent_config_snapshot_id="snap-2"), - backend_run_id="b", - snapshot=_snapshot(messages=2), - runtime_layer_specs=_runtime_layer_specs(), - ) - - assert store.load_active_snapshot(_scope(agent_config_snapshot_id="snap-1")) is None - assert store.load_active_snapshot(_scope(agent_config_snapshot_id="snap-2")) is not None - with session_factory.create_session() as session: - rows = session.query(AgentRuntimeSession).order_by(AgentRuntimeSession.backend_run_id).all() - assert len(rows) == 2 - assert [row.agent_config_snapshot_id for row in rows] == ["snap-1", "snap-2"] - assert [row.status for row in rows] == [AgentRuntimeSessionStatus.CLEANED, AgentRuntimeSessionStatus.ACTIVE] - - -def test_load_active_session_for_conversation_resolves_without_agent_or_config_scope(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(messages=2), - runtime_layer_specs=_runtime_layer_specs(), - ) - - loaded = store.load_active_session_for_conversation(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1") - assert loaded is not None - assert loaded.session_snapshot.layers[0].runtime_state["messages"] == [ - {"role": "user", "content": "m0"}, - {"role": "user", "content": "m1"}, - ] - assert [spec.name for spec in loaded.runtime_layer_specs] == ["execution_context", "history"] - - -def test_load_active_session_for_conversation_uses_latest_active_snapshot_after_config_change(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(agent_config_snapshot_id="snap-1"), - backend_run_id="a", - snapshot=_snapshot(), - runtime_layer_specs=_runtime_layer_specs(), - ) - store.save_active_snapshot( - scope=_scope(agent_config_snapshot_id="snap-2"), - backend_run_id="b", - snapshot=_snapshot(messages=3), - runtime_layer_specs=_runtime_layer_specs(), - ) - - loaded = store.load_active_session_for_conversation(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1") - - assert loaded is not None - assert loaded.session_snapshot.layers[0].runtime_state["messages"] == [ - {"role": "user", "content": "m0"}, - {"role": "user", "content": "m1"}, - {"role": "user", "content": "m2"}, - ] - - -def test_load_active_session_for_conversation_returns_none_when_cleaned_or_absent(): - store = AgentAppRuntimeSessionStore() - assert ( - store.load_active_session_for_conversation(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1") - is None - ) - - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(), - runtime_layer_specs=_runtime_layer_specs(), - ) - store.mark_cleaned(scope=_scope(), backend_run_id="cleanup-1") - assert ( - store.load_active_session_for_conversation(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1") - is None - ) - - -def test_load_active_session_for_conversation_isolates_other_conversations(): - store = AgentAppRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(conversation_id="conv-A"), - backend_run_id="a", - snapshot=_snapshot(), - runtime_layer_specs=_runtime_layer_specs(), - ) - - assert ( - store.load_active_session_for_conversation(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-B") - is None - ) - assert ( - store.load_active_session_for_conversation(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-A") - is not None - ) - - -def test_list_active_sessions_for_conversation_returns_all_active_rows(): - store = AgentAppRuntimeSessionStore() - with session_factory.create_session() as session: - session.add( - AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id="agent-1", - agent_config_snapshot_id="snap-1", - conversation_id="conv-1", - backend_run_id="run-1", - session_snapshot=_snapshot(messages=1).model_dump_json(), - composition_layer_specs='[{"name":"execution_context","type":"dify.execution_context","deps":{},"metadata":{},"config":{"tenant_id":"tenant-1"}},{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]', - status=AgentRuntimeSessionStatus.ACTIVE, - ) - ) - session.add( - AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id="agent-2", - agent_config_snapshot_id="snap-2", - conversation_id="conv-1", - backend_run_id="run-2", - session_snapshot=_snapshot(messages=2).model_dump_json(), - composition_layer_specs='[{"name":"execution_context","type":"dify.execution_context","deps":{},"metadata":{},"config":{"tenant_id":"tenant-1"}},{"name":"history","type":"pydantic_ai.history","deps":{},"metadata":{},"config":null}]', - status=AgentRuntimeSessionStatus.ACTIVE, - ) - ) - session.commit() - - loaded = store.list_active_sessions_for_conversation( - tenant_id="tenant-1", - app_id="app-1", - conversation_id="conv-1", - ) - - assert [session.scope.agent_id for session in loaded] == ["agent-2", "agent-1"] - assert all(session.scope.conversation_id == "conv-1" for session in loaded) + assert save.call_args.kwargs["binding_id"] == "binding-1" + assert save.call_args.kwargs["session_snapshot"] == snapshot.model_dump_json() diff --git a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_generator.py b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_generator.py index 0dadae0064b..29bfe2b4bb3 100644 --- a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_generator.py @@ -5,6 +5,7 @@ import logging import pytest from pydantic import ValidationError from pytest_mock import MockerFixture +from sqlalchemy.orm import Session, sessionmaker from core.app.apps.agent_chat.app_generator import AgentChatAppGenerator from core.app.apps.exc import GenerateTaskStoppedError @@ -30,12 +31,12 @@ def generator(mocker: MockerFixture): class TestAgentChatAppGeneratorGenerate: - def test_generate_rejects_blocking_mode(self, generator, mocker: MockerFixture): + def test_generate_rejects_blocking_mode(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model = mocker.MagicMock() user = DummyAccount("user") with pytest.raises(ValueError): generator.generate( - session=mocker.MagicMock(), + session=sqlite_session, app_model=app_model, user=user, args={}, @@ -43,44 +44,45 @@ class TestAgentChatAppGeneratorGenerate: streaming=False, ) - def test_generate_requires_query(self, generator, mocker: MockerFixture): + def test_generate_requires_query(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model = mocker.MagicMock() user = DummyAccount("user") with pytest.raises(ValueError): generator.generate( - session=mocker.MagicMock(), + session=sqlite_session, app_model=app_model, user=user, args={"inputs": {}}, invoke_from=mocker.MagicMock(), ) - def test_generate_rejects_non_string_query(self, generator, mocker: MockerFixture): + def test_generate_rejects_non_string_query(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model = mocker.MagicMock() user = DummyAccount("user") with pytest.raises(ValueError): generator.generate( - session=mocker.MagicMock(), + session=sqlite_session, app_model=app_model, user=user, args={"query": 123, "inputs": {}}, invoke_from=mocker.MagicMock(), ) - def test_generate_override_requires_debugger(self, generator, mocker: MockerFixture): + def test_generate_override_requires_debugger(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model = mocker.MagicMock() user = DummyAccount("user") + generator._get_app_model_config = mocker.MagicMock(return_value=mocker.MagicMock()) with pytest.raises(ValueError): generator.generate( - session=mocker.MagicMock(), + session=sqlite_session, app_model=app_model, user=user, args={"query": "hi", "inputs": {}, "model_config": {"model": {"provider": "p"}}}, invoke_from=InvokeFrom.WEB_APP, ) - def test_generate_success_with_debugger_override(self, generator, mocker: MockerFixture): + def test_generate_success_with_debugger_override(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat") app_model_config = mocker.MagicMock(id="cfg1") app_model_config.to_dict.return_value = {"model": {"provider": "p"}} @@ -154,7 +156,7 @@ class TestAgentChatAppGeneratorGenerate: "files": [{"id": "f1"}], "trace_session_id": "session-1", } - session = mocker.MagicMock() + session = sqlite_session result = generator.generate( session=session, @@ -173,7 +175,7 @@ class TestAgentChatAppGeneratorGenerate: inspect.signature(worker_call.kwargs["target"]).bind(**worker_call.kwargs["kwargs"]) thread_obj.start.assert_called_once() - def test_generate_without_file_config(self, generator, mocker: MockerFixture): + def test_generate_without_file_config(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model = mocker.MagicMock(id="app1", tenant_id="tenant", mode="agent-chat") app_model_config = mocker.MagicMock(id="cfg1", app_id="app1") app_model_config.to_dict.return_value = {"model": {"provider": "p"}} @@ -235,7 +237,7 @@ class TestAgentChatAppGeneratorGenerate: ) args = {"query": "hello", "inputs": {"name": "world"}} - session = mocker.MagicMock() + session = sqlite_session result = generator.generate( session=session, @@ -261,7 +263,9 @@ class TestAgentChatAppGeneratorWorker: mocker.patch("core.app.apps.agent_chat.app_generator.preserve_flask_contexts", ctx_manager) - def test_generate_worker_handles_generate_task_stopped(self, generator, mocker: MockerFixture): + def test_generate_worker_handles_generate_task_stopped( + self, generator, mocker: MockerFixture, sqlite_session_factory: sessionmaker[Session] + ): queue_manager = mocker.MagicMock() generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock()) generator._get_message = mocker.MagicMock(return_value=mocker.MagicMock()) @@ -269,10 +273,10 @@ class TestAgentChatAppGeneratorWorker: runner = mocker.MagicMock() runner.run.side_effect = GenerateTaskStoppedError() mocker.patch("core.app.apps.agent_chat.app_generator.AgentChatAppRunner", return_value=runner) - session_cm = mocker.MagicMock() - session_cm.__enter__.return_value = mocker.MagicMock() - create_session = mocker.patch("core.app.apps.agent_chat.app_generator.session_factory.create_session") - create_session.return_value = session_cm + mocker.patch( + "core.app.apps.agent_chat.app_generator.session_factory.create_session", + side_effect=sqlite_session_factory, + ) generator._generate_worker( flask_app=mocker.MagicMock(), @@ -294,7 +298,9 @@ class TestAgentChatAppGeneratorWorker: Exception("bad"), ], ) - def test_generate_worker_publishes_errors(self, generator, mocker: MockerFixture, error): + def test_generate_worker_publishes_errors( + self, generator, mocker: MockerFixture, error, sqlite_session_factory: sessionmaker[Session] + ): queue_manager = mocker.MagicMock() generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock()) generator._get_message = mocker.MagicMock(return_value=mocker.MagicMock()) @@ -302,10 +308,10 @@ class TestAgentChatAppGeneratorWorker: runner = mocker.MagicMock() runner.run.side_effect = error mocker.patch("core.app.apps.agent_chat.app_generator.AgentChatAppRunner", return_value=runner) - session_cm = mocker.MagicMock() - session_cm.__enter__.return_value = mocker.MagicMock() - create_session = mocker.patch("core.app.apps.agent_chat.app_generator.session_factory.create_session") - create_session.return_value = session_cm + mocker.patch( + "core.app.apps.agent_chat.app_generator.session_factory.create_session", + side_effect=sqlite_session_factory, + ) generator._generate_worker( flask_app=mocker.MagicMock(), @@ -319,7 +325,11 @@ class TestAgentChatAppGeneratorWorker: assert queue_manager.publish_error.called def test_generate_worker_logs_value_error_when_debug( - self, generator, mocker: MockerFixture, caplog: pytest.LogCaptureFixture + self, + generator, + mocker: MockerFixture, + caplog: pytest.LogCaptureFixture, + sqlite_session_factory: sessionmaker[Session], ): queue_manager = mocker.MagicMock() generator._get_conversation = mocker.MagicMock(return_value=mocker.MagicMock()) @@ -328,10 +338,10 @@ class TestAgentChatAppGeneratorWorker: runner = mocker.MagicMock() runner.run.side_effect = ValueError("bad") mocker.patch("core.app.apps.agent_chat.app_generator.AgentChatAppRunner", return_value=runner) - session_cm = mocker.MagicMock() - session_cm.__enter__.return_value = mocker.MagicMock() - create_session = mocker.patch("core.app.apps.agent_chat.app_generator.session_factory.create_session") - create_session.return_value = session_cm + mocker.patch( + "core.app.apps.agent_chat.app_generator.session_factory.create_session", + side_effect=sqlite_session_factory, + ) mocker.patch("core.app.apps.agent_chat.app_generator.dify_config", new=mocker.MagicMock(DEBUG=True)) diff --git a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_runner.py b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_runner.py index f4caca4e1d7..1bbfb9e5f1f 100644 --- a/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/agent_chat/test_agent_chat_app_runner.py @@ -1,41 +1,100 @@ +from datetime import datetime + import pytest from pytest_mock import MockerFixture +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker from core.agent.entities import AgentEntity from core.app.apps.agent_chat.app_runner import AgentChatAppRunner +from core.app.entities.app_invoke_entities import InvokeFrom from core.moderation.base import ModerationError from graphon.model_runtime.entities.llm_entities import LLMMode from graphon.model_runtime.entities.model_entities import ModelFeature, ModelPropertyKey +from models.enums import ConversationFromSource +from models.model import App, AppMode, Conversation, Message @pytest.fixture -def runner(): +def runner(sqlite_session: Session): + app = App( + id="app1", + tenant_id="tenant", + name="Agent chat app", + description="", + mode=AppMode.AGENT_CHAT, + enable_site=False, + enable_api=False, + ) + conversation = Conversation( + id="conv", + app_id=app.id, + app_model_config_id=None, + model_provider=None, + override_model_configs=None, + model_id=None, + mode=AppMode.AGENT_CHAT, + name="Conversation", + inputs={}, + introduction="", + system_instruction="", + system_instruction_tokens=0, + status="normal", + invoke_from=InvokeFrom.SERVICE_API, + from_source=ConversationFromSource.API, + from_end_user_id=None, + from_account_id="user", + ) + message = Message( + id="msg", + app_id=app.id, + conversation_id=conversation.id, + inputs={}, + query="q", + message={}, + message_tokens=0, + message_unit_price=0, + message_price_unit=0, + answer="", + answer_tokens=0, + answer_unit_price=0, + answer_price_unit=0, + provider_response_latency=0, + total_price=0, + currency="USD", + invoke_from=InvokeFrom.SERVICE_API, + from_source=ConversationFromSource.API, + from_end_user_id=None, + from_account_id="user", + app_mode=AppMode.AGENT_CHAT, + created_at=datetime(2025, 1, 1), + ) + sqlite_session.add_all([app, conversation, message]) + sqlite_session.commit() return AgentChatAppRunner() -def patch_create_session(mocker: MockerFixture, *, return_value=None, side_effect=None): - session = mocker.MagicMock() - if side_effect is not None: - session.scalar.side_effect = side_effect - else: - session.scalar.return_value = return_value - session_context = mocker.MagicMock() - session_context.__enter__.return_value = session - mocker.patch("core.app.apps.agent_chat.app_runner.create_session", return_value=session_context) - return session +@pytest.fixture(autouse=True) +def _patch_create_session(mocker: MockerFixture, sqlite_session_factory: sessionmaker[Session]) -> None: + mocker.patch("core.app.apps.agent_chat.app_runner.create_session", side_effect=sqlite_session_factory) class TestAgentChatAppRunnerRun: - def test_run_app_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture): + def test_run_app_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session): app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", agent=mocker.MagicMock()) generate_entity = mocker.MagicMock(app_config=app_config, inputs={}, query="q", files=[], stream=True) - patch_create_session(mocker, return_value=None) + app = sqlite_session.get(App, "app1") + assert app is not None + sqlite_session.delete(app) + sqlite_session.commit() with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), sqlite_session) - def test_run_moderation_error_direct_output(self, runner: AgentChatAppRunner, mocker: MockerFixture): + def test_run_moderation_error_direct_output( + self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session + ): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config.agent = mocker.MagicMock() @@ -49,16 +108,17 @@ class TestAgentChatAppRunnerRun: conversation_id=None, ) - patch_create_session(mocker, return_value=app_record) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", side_effect=ModerationError("bad")) mocker.patch.object(runner, "direct_output") - runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), sqlite_session) runner.direct_output.assert_called_once() - def test_run_annotation_reply_short_circuits(self, runner: AgentChatAppRunner, mocker: MockerFixture): + def test_run_annotation_reply_short_circuits( + self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session + ): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config.agent = mocker.MagicMock() @@ -74,7 +134,6 @@ class TestAgentChatAppRunnerRun: invoke_from=mocker.MagicMock(), ) - patch_create_session(mocker, return_value=app_record) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) annotation = mocker.MagicMock(id="anno", content="answer") @@ -82,14 +141,16 @@ class TestAgentChatAppRunnerRun: mocker.patch.object(runner, "direct_output") queue_manager = mocker.MagicMock() - write_session = mocker.MagicMock() + write_session = sqlite_session runner.run(generate_entity, queue_manager, mocker.MagicMock(), mocker.MagicMock(), write_session) queue_manager.publish.assert_called_once() assert annotation_query.call_args.kwargs["session"] is write_session runner.direct_output.assert_called_once() - def test_run_hosting_moderation_short_circuits(self, runner: AgentChatAppRunner, mocker: MockerFixture): + def test_run_hosting_moderation_short_circuits( + self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session + ): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config.agent = mocker.MagicMock() @@ -105,15 +166,14 @@ class TestAgentChatAppRunnerRun: user_id="user", ) - patch_create_session(mocker, return_value=app_record) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) mocker.patch.object(runner, "check_hosting_moderation", return_value=True) - runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), sqlite_session) - def test_run_model_schema_missing(self, runner: AgentChatAppRunner, mocker: MockerFixture): + def test_run_model_schema_missing(self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) @@ -135,7 +195,6 @@ class TestAgentChatAppRunnerRun: user_id="user", ) - patch_create_session(mocker, return_value=app_record) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) @@ -146,7 +205,7 @@ class TestAgentChatAppRunnerRun: mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), mocker.MagicMock(), mocker.MagicMock(), sqlite_session) @pytest.mark.parametrize( ("mode", "expected_runner"), @@ -155,109 +214,120 @@ class TestAgentChatAppRunnerRun: (LLMMode.COMPLETION, "CotCompletionAgentRunner"), ], ) - def test_run_chain_of_thought_modes(self, runner: AgentChatAppRunner, mocker: MockerFixture, mode, expected_runner): - app_record = mocker.MagicMock(id="app1", tenant_id="tenant") - app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) - app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) - - generate_entity = mocker.MagicMock( - app_config=app_config, - inputs={}, - query="q", - files=[], - stream=True, - model_conf=mocker.MagicMock( - provider_model_bundle=mocker.MagicMock(), - model="m", - provider="p", - credentials={"k": "v"}, - ), - conversation_id="conv", - invoke_from=mocker.MagicMock(), - user_id="user", - ) - - patch_create_session(mocker, return_value=app_record) - mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) - mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) - mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) - mocker.patch.object(runner, "check_hosting_moderation", return_value=False) - - model_schema = mocker.MagicMock() - model_schema.features = [] - model_schema.model_properties = {ModelPropertyKey.MODE: mode} - - llm_instance = mocker.MagicMock() - llm_instance.model_type_instance.get_model_schema.return_value = model_schema - mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) - - conversation = mocker.MagicMock(id="conv") - message = mocker.MagicMock(id="msg") - patch_create_session(mocker, side_effect=[app_record, conversation, message]) - - runner_cls = mocker.MagicMock() - mocker.patch(f"core.app.apps.agent_chat.app_runner.{expected_runner}", runner_cls) - - runner_instance = mocker.MagicMock() - runner_cls.return_value = runner_instance - events: list[str] = [] - runner_instance.run.side_effect = lambda **_kwargs: events.append("agent-run") or [] - mocker.patch.object(runner, "_handle_invoke_result") - session = mocker.MagicMock() - session.commit.side_effect = lambda: events.append("commit") - session.close.side_effect = lambda: events.append("close") - - runner.run(generate_entity, mocker.MagicMock(), conversation, message, session) - - assert events == ["commit", "close", "commit", "close", "agent-run"] - runner_instance.run.assert_called_once() - runner._handle_invoke_result.assert_called_once() - - def test_run_invalid_llm_mode_raises(self, runner: AgentChatAppRunner, mocker: MockerFixture): - app_record = mocker.MagicMock(id="app1", tenant_id="tenant") - app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) - app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) - - generate_entity = mocker.MagicMock( - app_config=app_config, - inputs={}, - query="q", - files=[], - stream=True, - model_conf=mocker.MagicMock( - provider_model_bundle=mocker.MagicMock(), - model="m", - provider="p", - credentials={"k": "v"}, - ), - conversation_id="conv", - invoke_from=mocker.MagicMock(), - user_id="user", - ) - - patch_create_session(mocker, return_value=app_record) - mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) - mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) - mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) - mocker.patch.object(runner, "check_hosting_moderation", return_value=False) - - model_schema = mocker.MagicMock() - model_schema.features = [] - model_schema.model_properties = {ModelPropertyKey.MODE: "invalid"} - - llm_instance = mocker.MagicMock() - llm_instance.model_type_instance.get_model_schema.return_value = model_schema - mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) - - conversation = mocker.MagicMock(id="conv") - message = mocker.MagicMock(id="msg") - patch_create_session(mocker, side_effect=[app_record, conversation, message]) - - with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), conversation, message, mocker.MagicMock()) - - def test_run_function_calling_strategy_selected_by_features( - self, runner: AgentChatAppRunner, mocker: MockerFixture + def test_run_chain_of_thought_modes( + self, + runner: AgentChatAppRunner, + mocker: MockerFixture, + mode, + expected_runner, + sqlite_session: Session, + ): + app_record = mocker.MagicMock(id="app1", tenant_id="tenant") + app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) + app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) + + generate_entity = mocker.MagicMock( + app_config=app_config, + inputs={}, + query="q", + files=[], + stream=True, + model_conf=mocker.MagicMock( + provider_model_bundle=mocker.MagicMock(), + model="m", + provider="p", + credentials={"k": "v"}, + ), + conversation_id="conv", + invoke_from=mocker.MagicMock(), + user_id="user", + ) + + mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) + mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) + mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) + mocker.patch.object(runner, "check_hosting_moderation", return_value=False) + + model_schema = mocker.MagicMock() + model_schema.features = [] + model_schema.model_properties = {ModelPropertyKey.MODE: mode} + + llm_instance = mocker.MagicMock() + llm_instance.model_type_instance.get_model_schema.return_value = model_schema + mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) + + conversation = mocker.MagicMock(id="conv") + message = mocker.MagicMock(id="msg") + + runner_cls = mocker.MagicMock() + mocker.patch(f"core.app.apps.agent_chat.app_runner.{expected_runner}", runner_cls) + + runner_instance = mocker.MagicMock() + runner_cls.return_value = runner_instance + events: list[str] = [] + runner_instance.run.side_effect = lambda **_kwargs: events.append("agent-run") or [] + mocker.patch.object(runner, "_handle_invoke_result") + session = sqlite_session + event.listen(session, "after_commit", lambda _session: events.append("commit")) + original_close = session.close + + def close_session() -> None: + events.append("close") + original_close() + + session.close = close_session + + runner.run(generate_entity, mocker.MagicMock(), conversation, message, session) + + assert events == ["commit", "close", "commit", "close", "agent-run"] + runner_instance.run.assert_called_once() + runner._handle_invoke_result.assert_called_once() + + def test_run_invalid_llm_mode_raises( + self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session + ): + app_record = mocker.MagicMock(id="app1", tenant_id="tenant") + app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) + app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.CHAIN_OF_THOUGHT) + + generate_entity = mocker.MagicMock( + app_config=app_config, + inputs={}, + query="q", + files=[], + stream=True, + model_conf=mocker.MagicMock( + provider_model_bundle=mocker.MagicMock(), + model="m", + provider="p", + credentials={"k": "v"}, + ), + conversation_id="conv", + invoke_from=mocker.MagicMock(), + user_id="user", + ) + + mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) + mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) + mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) + mocker.patch.object(runner, "check_hosting_moderation", return_value=False) + + model_schema = mocker.MagicMock() + model_schema.features = [] + model_schema.model_properties = {ModelPropertyKey.MODE: "invalid"} + + llm_instance = mocker.MagicMock() + llm_instance.model_type_instance.get_model_schema.return_value = model_schema + mocker.patch("core.app.apps.agent_chat.app_runner.ModelInstance", return_value=llm_instance) + + conversation = mocker.MagicMock(id="conv") + message = mocker.MagicMock(id="msg") + + with pytest.raises(ValueError): + runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session) + + def test_run_function_calling_strategy_selected_by_features( + self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session ): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) @@ -280,7 +350,6 @@ class TestAgentChatAppRunnerRun: user_id="user", ) - patch_create_session(mocker, return_value=app_record) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) @@ -296,7 +365,6 @@ class TestAgentChatAppRunnerRun: conversation = mocker.MagicMock(id="conv") message = mocker.MagicMock(id="msg") - patch_create_session(mocker, side_effect=[app_record, conversation, message]) runner_cls = mocker.MagicMock() mocker.patch("core.app.apps.agent_chat.app_runner.FunctionCallAgentRunner", runner_cls) @@ -306,12 +374,14 @@ class TestAgentChatAppRunnerRun: runner_instance.run.return_value = [] mocker.patch.object(runner, "_handle_invoke_result") - runner.run(generate_entity, mocker.MagicMock(), conversation, message, mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session) assert app_config.agent.strategy == AgentEntity.Strategy.FUNCTION_CALLING runner_instance.run.assert_called_once() - def test_run_conversation_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture): + def test_run_conversation_not_found( + self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session + ): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.FUNCTION_CALLING) @@ -333,7 +403,10 @@ class TestAgentChatAppRunnerRun: user_id="user", ) - patch_create_session(mocker, side_effect=[app_record, None]) + conversation_record = sqlite_session.get(Conversation, "conv") + assert conversation_record is not None + sqlite_session.delete(conversation_record) + sqlite_session.commit() mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) @@ -345,10 +418,10 @@ class TestAgentChatAppRunnerRun: mocker.MagicMock(), mocker.MagicMock(id="conv"), mocker.MagicMock(id="msg"), - mocker.MagicMock(), + sqlite_session, ) - def test_run_message_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture): + def test_run_message_not_found(self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config.agent = AgentEntity(provider="p", model="m", strategy=AgentEntity.Strategy.FUNCTION_CALLING) @@ -370,7 +443,10 @@ class TestAgentChatAppRunnerRun: user_id="user", ) - patch_create_session(mocker, side_effect=[app_record, mocker.MagicMock(id="conv"), None]) + message_record = sqlite_session.get(Message, "msg") + assert message_record is not None + sqlite_session.delete(message_record) + sqlite_session.commit() mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) @@ -382,10 +458,12 @@ class TestAgentChatAppRunnerRun: mocker.MagicMock(), mocker.MagicMock(id="conv"), mocker.MagicMock(id="msg"), - mocker.MagicMock(), + sqlite_session, ) - def test_run_invalid_agent_strategy_raises(self, runner: AgentChatAppRunner, mocker: MockerFixture): + def test_run_invalid_agent_strategy_raises( + self, runner: AgentChatAppRunner, mocker: MockerFixture, sqlite_session: Session + ): app_record = mocker.MagicMock(id="app1", tenant_id="tenant") app_config = mocker.MagicMock(app_id="app1", tenant_id="tenant", prompt_template=mocker.MagicMock()) app_config.agent = mocker.MagicMock(strategy="invalid", provider="p", model="m") @@ -407,7 +485,6 @@ class TestAgentChatAppRunnerRun: user_id="user", ) - patch_create_session(mocker, return_value=app_record) mocker.patch.object(runner, "organize_prompt_messages", return_value=([], None)) mocker.patch.object(runner, "moderation_for_inputs", return_value=(None, {}, "q")) mocker.patch.object(runner, "query_app_annotations_to_reply", return_value=None) @@ -423,7 +500,6 @@ class TestAgentChatAppRunnerRun: conversation = mocker.MagicMock(id="conv") message = mocker.MagicMock(id="msg") - patch_create_session(mocker, side_effect=[app_record, conversation, message]) with pytest.raises(ValueError): - runner.run(generate_entity, mocker.MagicMock(), conversation, message, mocker.MagicMock()) + runner.run(generate_entity, mocker.MagicMock(), conversation, message, sqlite_session) diff --git a/api/tests/unit_tests/core/app/apps/chat/test_app_config_manager.py b/api/tests/unit_tests/core/app/apps/chat/test_app_config_manager.py index 24ec3116b32..6bb0c0c977c 100644 --- a/api/tests/unit_tests/core/app/apps/chat/test_app_config_manager.py +++ b/api/tests/unit_tests/core/app/apps/chat/test_app_config_manager.py @@ -1,6 +1,8 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch +from sqlalchemy.orm import Session + from core.app.app_config.entities import EasyUIBasedAppModelConfigFrom, ModelConfigEntity, PromptTemplateEntity from core.app.apps.chat.app_config_manager import ChatAppConfigManager from models.model import AppMode @@ -76,7 +78,7 @@ class TestChatAppConfigManager: app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) - def test_config_validate_filters_related_keys(self): + def test_config_validate_filters_related_keys(self, unbound_session: Session): config = {"extra": 1} def _add_key(key, value): @@ -133,7 +135,7 @@ class TestChatAppConfigManager: side_effect=_add_key("sensitive_word_avoidance", 11), ), ): - filtered = ChatAppConfigManager.config_validate(session=MagicMock(), tenant_id="t1", config=config) + filtered = ChatAppConfigManager.config_validate(session=unbound_session, tenant_id="t1", config=config) assert filtered["model"] == 1 assert filtered["inputs"] == 2 diff --git a/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py b/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py index bdb706bde09..905e711414a 100644 --- a/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py +++ b/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py @@ -3,6 +3,7 @@ from types import SimpleNamespace from unittest.mock import ANY, MagicMock, Mock, patch import pytest +from sqlalchemy.orm import Session, sessionmaker from core.app.apps.chat.app_generator import ChatAppGenerator from core.app.apps.chat.app_runner import ChatAppRunner @@ -11,7 +12,7 @@ from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.queue_entities import QueueAnnotationReplyEvent from core.moderation.base import ModerationError from graphon.model_runtime.errors.invoke import InvokeAuthorizationError -from models.model import AppMode +from models.model import App, AppMode, IconType class DummyGenerateEntity: @@ -31,24 +32,39 @@ class DummyQueueManager: @contextmanager -def patched_create_session(*, return_value=None, side_effect=None): - session = MagicMock() - if side_effect is not None: - session.scalar.side_effect = side_effect - else: - session.scalar.return_value = return_value - session_context = MagicMock() - session_context.__enter__.return_value = session - with patch("core.app.apps.chat.app_runner.create_session", return_value=session_context): - yield session +def patched_create_session(session_factory: sessionmaker[Session]): + @contextmanager + def create_session(): + with session_factory() as session: + yield session + + with patch("core.app.apps.chat.app_runner.create_session", create_session): + yield + + +def _persist_app(session: Session) -> App: + app = App( + id="app-1", + tenant_id="tenant-1", + name="Chat app", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#ffffff", + enable_site=False, + enable_api=True, + ) + session.add(app) + session.commit() + return app class TestChatAppGenerator: - def test_generate_requires_query(self): + def test_generate_requires_query(self, unbound_session: Session): generator = ChatAppGenerator() with pytest.raises(ValueError): generator.generate( - session=MagicMock(), + session=unbound_session, app_model=SimpleNamespace(), user=SimpleNamespace(), args={"inputs": {}}, @@ -56,11 +72,11 @@ class TestChatAppGenerator: streaming=False, ) - def test_generate_rejects_non_string_query(self): + def test_generate_rejects_non_string_query(self, unbound_session: Session): generator = ChatAppGenerator() with pytest.raises(ValueError): generator.generate( - session=MagicMock(), + session=unbound_session, app_model=SimpleNamespace(), user=SimpleNamespace(), args={"query": 1, "inputs": {}}, @@ -68,11 +84,10 @@ class TestChatAppGenerator: streaming=False, ) - def test_generate_debugger_overrides_model_config(self): + def test_generate_debugger_overrides_model_config(self, unbound_session: Session): generator = ChatAppGenerator() app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") user = SimpleNamespace(id="user-1", session_id="session-1") - session = MagicMock() args = { "query": "hi", "inputs": {}, @@ -116,19 +131,20 @@ class TestChatAppGenerator: patch("core.app.apps.chat.app_generator.threading.Thread") as mock_thread, ): mock_thread.return_value.start.return_value = None - result = generator.generate(app_model, user, args, InvokeFrom.DEBUGGER, streaming=False, session=session) + result = generator.generate( + app_model, user, args, InvokeFrom.DEBUGGER, streaming=False, session=unbound_session + ) assert result == {"ok": True} - assert get_conversation.call_args.kwargs["session"] is session + assert get_conversation.call_args.kwargs["session"] is unbound_session assert generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1" - def test_generate_uses_session_for_annotation_reply(self): + def test_generate_uses_session_for_annotation_reply(self, unbound_session: Session): generator = ChatAppGenerator() app_model = SimpleNamespace(id="app-1", tenant_id="tenant-1") app_model_config = MagicMock(id="config-1", app_id="app-1") annotation_reply = {"enabled": False} user = SimpleNamespace(id="user-1", session_id="session-1") - session = MagicMock() with ( patch.object(ChatAppGenerator, "_get_app_model_config", return_value=app_model_config), @@ -148,14 +164,14 @@ class TestChatAppGenerator: user, {"query": "hi", "inputs": {}}, InvokeFrom.WEB_APP, - session=session, + session=unbound_session, ) - load_annotation_reply_config.assert_called_once_with(session, "app-1") + load_annotation_reply_config.assert_called_once_with(unbound_session, "app-1") app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) assert get_app_config.call_args.kwargs["annotation_reply"] is annotation_reply - def test_generate_rejects_model_config_override_for_non_debugger(self): + def test_generate_rejects_model_config_override_for_non_debugger(self, unbound_session: Session): generator = ChatAppGenerator() with pytest.raises(ValueError): with ( @@ -164,7 +180,7 @@ class TestChatAppGenerator: ), ): generator.generate( - session=MagicMock(), + session=unbound_session, app_model=SimpleNamespace(tenant_id="t1", id="a1", mode=AppMode.CHAT.value), user=SimpleNamespace(id="u1", session_id="s1"), args={"query": "hi", "inputs": {}, "model_config": {"foo": "bar"}}, @@ -172,7 +188,7 @@ class TestChatAppGenerator: streaming=False, ) - def test_generate_worker_handles_exceptions(self): + def test_generate_worker_handles_exceptions(self, unbound_session_factory: sessionmaker[Session]): generator = ChatAppGenerator() queue_manager = DummyQueueManager() entity = DummyGenerateEntity(task_id="t1", user_id="u1") @@ -181,10 +197,12 @@ class TestChatAppGenerator: patch.object(ChatAppGenerator, "_get_conversation", return_value=SimpleNamespace()), patch.object(ChatAppGenerator, "_get_message", return_value=SimpleNamespace()), patch("core.app.apps.chat.app_generator.ChatAppRunner.run", side_effect=InvokeAuthorizationError()), - patch("core.app.apps.chat.app_generator.session_factory.create_session") as create_session, + patch( + "core.app.apps.chat.app_generator.session_factory", + SimpleNamespace(create_session=unbound_session_factory), + ), patch("core.app.apps.chat.app_generator.db.session.close"), ): - create_session.return_value.__enter__.return_value = MagicMock() generator._generate_worker( flask_app=Mock(app_context=Mock(return_value=Mock(__enter__=Mock(), __exit__=Mock()))), application_generate_entity=entity, @@ -199,10 +217,12 @@ class TestChatAppGenerator: patch.object(ChatAppGenerator, "_get_conversation", return_value=SimpleNamespace()), patch.object(ChatAppGenerator, "_get_message", return_value=SimpleNamespace()), patch("core.app.apps.chat.app_generator.ChatAppRunner.run", side_effect=GenerateTaskStoppedError()), - patch("core.app.apps.chat.app_generator.session_factory.create_session") as create_session, + patch( + "core.app.apps.chat.app_generator.session_factory", + SimpleNamespace(create_session=unbound_session_factory), + ), patch("core.app.apps.chat.app_generator.db.session.close"), ): - create_session.return_value.__enter__.return_value = MagicMock() generator._generate_worker( flask_app=Mock(app_context=Mock(return_value=Mock(__enter__=Mock(), __exit__=Mock()))), application_generate_entity=entity, @@ -213,7 +233,7 @@ class TestChatAppGenerator: class TestChatAppRunner: - def test_run_raises_when_app_missing(self): + def test_run_raises_when_app_missing(self, sqlite_session_factory: sessionmaker[Session], unbound_session: Session): runner = ChatAppRunner() app_config = SimpleNamespace( app_id="app-1", tenant_id="tenant-1", prompt_template=None, external_data_variables=[] @@ -231,13 +251,20 @@ class TestChatAppRunner: invoke_from=InvokeFrom.SERVICE_API, ) - with patched_create_session(return_value=None): + with patched_create_session(sqlite_session_factory): with pytest.raises(ValueError): runner.run( - app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1"), MagicMock() + app_generate_entity, + DummyQueueManager(), + SimpleNamespace(), + SimpleNamespace(id="m1"), + unbound_session, ) - def test_run_moderation_error_direct_output(self): + def test_run_moderation_error_direct_output( + self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] + ): + _persist_app(sqlite_session) runner = ChatAppRunner() app_config = SimpleNamespace( app_id="app-1", @@ -261,18 +288,25 @@ class TestChatAppRunner: ) with ( - patched_create_session(return_value=SimpleNamespace(id="app-1", tenant_id="tenant-1")), + patched_create_session(sqlite_session_factory), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), patch.object(ChatAppRunner, "moderation_for_inputs", side_effect=ModerationError("blocked")), patch.object(ChatAppRunner, "direct_output") as mock_direct, ): runner.run( - app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1"), MagicMock() + app_generate_entity, + DummyQueueManager(), + SimpleNamespace(), + SimpleNamespace(id="m1"), + sqlite_session, ) mock_direct.assert_called_once() - def test_run_annotation_reply_short_circuits(self): + def test_run_annotation_reply_short_circuits( + self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] + ): + _persist_app(sqlite_session) runner = ChatAppRunner() app_config = SimpleNamespace( app_id="app-1", @@ -298,21 +332,23 @@ class TestChatAppRunner: annotation = SimpleNamespace(id="ann-1", content="answer") with ( - patched_create_session(return_value=SimpleNamespace(id="app-1", tenant_id="tenant-1")), + patched_create_session(sqlite_session_factory), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")), patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=annotation) as annotation_query, patch.object(ChatAppRunner, "direct_output") as mock_direct, ): queue_manager = DummyQueueManager() - write_session = MagicMock() - runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1"), write_session) + runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1"), sqlite_session) assert any(isinstance(item[0], QueueAnnotationReplyEvent) for item in queue_manager.published) - assert annotation_query.call_args.kwargs["session"] is write_session + assert annotation_query.call_args.kwargs["session"] is sqlite_session mock_direct.assert_called_once() - def test_run_returns_when_hosting_moderation_blocks(self): + def test_run_returns_when_hosting_moderation_blocks( + self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] + ): + _persist_app(sqlite_session) runner = ChatAppRunner() app_config = SimpleNamespace( app_id="app-1", @@ -336,17 +372,24 @@ class TestChatAppRunner: ) with ( - patched_create_session(return_value=SimpleNamespace(id="app-1", tenant_id="tenant-1")), + patched_create_session(sqlite_session_factory), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")), patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=None), patch.object(ChatAppRunner, "check_hosting_moderation", return_value=True), ): runner.run( - app_generate_entity, DummyQueueManager(), SimpleNamespace(), SimpleNamespace(id="m1"), MagicMock() + app_generate_entity, + DummyQueueManager(), + SimpleNamespace(), + SimpleNamespace(id="m1"), + sqlite_session, ) - def test_run_closes_explicit_session_before_stream_consumption(self): + def test_run_closes_explicit_session_before_stream_consumption( + self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] + ): + _persist_app(sqlite_session) runner = ChatAppRunner() app_config = SimpleNamespace( app_id="app-1", @@ -372,9 +415,8 @@ class TestChatAppRunner: events = [] queue_manager = DummyQueueManager() model_instance = MagicMock() - session = MagicMock() - session.commit.side_effect = lambda: events.append("commit") - session.close.side_effect = lambda: events.append("close") + original_commit = sqlite_session.commit + original_close = sqlite_session.close def invoke_stream(): events.append("first-chunk") @@ -385,7 +427,7 @@ class TestChatAppRunner: return invoke_stream() with ( - patched_create_session(return_value=SimpleNamespace(id="app-1", tenant_id="tenant-1")), + patched_create_session(sqlite_session_factory), patch.object(ChatAppRunner, "organize_prompt_messages", return_value=([], [])), patch.object(ChatAppRunner, "moderation_for_inputs", return_value=(None, {}, "hi")), patch.object(ChatAppRunner, "query_app_annotations_to_reply", return_value=None), @@ -397,9 +439,11 @@ class TestChatAppRunner: side_effect=lambda invoke_result, **kwargs: list(invoke_result), ) as mock_handle, patch("core.app.apps.chat.app_runner.ModelInstance", return_value=model_instance), + patch.object(sqlite_session, "commit", side_effect=lambda: (events.append("commit"), original_commit())[1]), + patch.object(sqlite_session, "close", side_effect=lambda: (events.append("close"), original_close())[1]), ): model_instance.invoke_llm.side_effect = invoke_llm - runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1"), session) + runner.run(app_generate_entity, queue_manager, SimpleNamespace(), SimpleNamespace(id="m1"), sqlite_session) assert events == ["commit", "close", "commit", "close", "invoke", "first-chunk"] mock_handle.assert_called_once_with( diff --git a/api/tests/unit_tests/core/app/apps/common/test_workflow_response_converter_truncation.py b/api/tests/unit_tests/core/app/apps/common/test_workflow_response_converter_truncation.py index b3c0eb74faa..ecf58070655 100644 --- a/api/tests/unit_tests/core/app/apps/common/test_workflow_response_converter_truncation.py +++ b/api/tests/unit_tests/core/app/apps/common/test_workflow_response_converter_truncation.py @@ -49,10 +49,11 @@ class TestWorkflowResponseConverter: """Create a WorkflowResponseConverter for testing.""" mock_entity = self.create_mock_generate_entity() - mock_user = Mock(spec=Account) + mock_user = Account( + name="Test User", + email="test@example.com", + ) mock_user.id = "test-user-id" - mock_user.name = "Test User" - mock_user.email = "test@example.com" system_variables = build_system_variables(workflow_id="wf-id", workflow_execution_id="initial-run-id") return WorkflowResponseConverter( diff --git a/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py b/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py index 6de22abbe1b..f7297de29a8 100644 --- a/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py @@ -1,14 +1,18 @@ -from contextlib import contextmanager from types import SimpleNamespace -from unittest.mock import ANY, MagicMock, patch +from unittest.mock import ANY, MagicMock import pytest from pytest_mock import MockerFixture +from sqlalchemy.orm import Session import core.app.apps.completion.app_runner as module from core.app.apps.completion.app_runner import CompletionAppRunner from core.moderation.base import ModerationError from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent +from models.model import App, AppMode, IconType + +APP_ID = "00000000-0000-0000-0000-000000000001" +TENANT_ID = "00000000-0000-0000-0000-000000000002" @pytest.fixture @@ -18,8 +22,8 @@ def runner(): def _build_app_config(dataset=None, external_tools=None, additional_features=None): app_config = MagicMock() - app_config.app_id = "app1" - app_config.tenant_id = "tenant" + app_config.app_id = APP_ID + app_config.tenant_id = TENANT_ID app_config.prompt_template = MagicMock() app_config.dataset = dataset app_config.external_data_variables = external_tools or [] @@ -48,27 +52,33 @@ def _build_generate_entity(app_config, file_upload_config=None): ) -@contextmanager -def patched_create_session(*, return_value=None): - session = MagicMock() - session.scalar.return_value = return_value - session_context = MagicMock() - session_context.__enter__.return_value = session - with patch.object(module, "create_session", return_value=session_context): - yield session +def _persist_app(session: Session) -> App: + app = App( + id=APP_ID, + tenant_id=TENANT_ID, + name="Completion app", + mode=AppMode.COMPLETION, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#ffffff", + enable_site=False, + enable_api=False, + ) + session.add(app) + session.commit() + return app class TestCompletionAppRunner: - def test_run_app_not_found(self, runner, mocker: MockerFixture): + def test_run_app_not_found(self, runner, mocker: MockerFixture, sqlite_session: Session): app_config = _build_app_config() app_generate_entity = _build_generate_entity(app_config) - with patched_create_session(return_value=None): - with pytest.raises(ValueError): - runner.run(app_generate_entity, MagicMock(), MagicMock(), MagicMock()) + with pytest.raises(ValueError): + runner.run(app_generate_entity, MagicMock(), MagicMock(), sqlite_session) - def test_run_moderation_error_outputs_direct(self, runner, mocker: MockerFixture): - app_record = MagicMock(id="app1", tenant_id="tenant") + def test_run_moderation_error_outputs_direct(self, runner, mocker: MockerFixture, sqlite_session: Session): + _persist_app(sqlite_session) app_config = _build_app_config() app_generate_entity = _build_generate_entity(app_config) @@ -78,14 +88,13 @@ class TestCompletionAppRunner: runner.direct_output = MagicMock() runner._handle_invoke_result = MagicMock() - with patched_create_session(return_value=app_record): - runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock()) + runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), sqlite_session) runner.direct_output.assert_called_once() runner._handle_invoke_result.assert_not_called() - def test_run_hosting_moderation_stops(self, runner, mocker: MockerFixture): - app_record = MagicMock(id="app1", tenant_id="tenant") + def test_run_hosting_moderation_stops(self, runner, mocker: MockerFixture, sqlite_session: Session): + _persist_app(sqlite_session) app_config = _build_app_config() app_generate_entity = _build_generate_entity(app_config) @@ -95,13 +104,12 @@ class TestCompletionAppRunner: runner.check_hosting_moderation = MagicMock(return_value=True) runner._handle_invoke_result = MagicMock() - with patched_create_session(return_value=app_record): - runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock()) + runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), sqlite_session) runner._handle_invoke_result.assert_not_called() - def test_run_dataset_and_external_tools_flow(self, runner, mocker: MockerFixture): - app_record = MagicMock(id="app1", tenant_id="tenant") + def test_run_dataset_and_external_tools_flow(self, runner, mocker: MockerFixture, sqlite_session: Session): + _persist_app(sqlite_session) retrieve_config = MagicMock(query_variable="qvar") dataset_config = MagicMock(dataset_ids=["ds"], retrieve_config=retrieve_config) @@ -132,23 +140,35 @@ class TestCompletionAppRunner: model_instance.invoke_llm.return_value = "invoke_result" mocker.patch.object(module, "ModelInstance", return_value=model_instance) - with patched_create_session(return_value=app_record): - runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg", tenant_id="tenant"), MagicMock()) + runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg", tenant_id=TENANT_ID), sqlite_session) dataset_retrieval.retrieve.assert_called_once() assert dataset_retrieval.retrieve.call_args.kwargs["query"] == "query_from_input" runner._handle_invoke_result.assert_called_once() - def test_run_closes_explicit_session_before_stream_consumption(self, runner, mocker: MockerFixture): - app_record = MagicMock(id="app1", tenant_id="tenant") + def test_run_closes_explicit_session_before_stream_consumption( + self, runner, mocker: MockerFixture, sqlite_session: Session + ): + _persist_app(sqlite_session) app_config = _build_app_config() app_generate_entity = _build_generate_entity(app_config) queue_manager = MagicMock() events = [] - session = MagicMock() - session.commit.side_effect = lambda: events.append("commit") - session.close.side_effect = lambda: events.append("close") + session = sqlite_session + original_commit = session.commit + original_close = session.close + + def commit_session() -> None: + events.append("commit") + original_commit() + + def close_session() -> None: + events.append("close") + original_close() + + mocker.patch.object(session, "commit", side_effect=commit_session) + mocker.patch.object(session, "close", side_effect=close_session) runner.organize_prompt_messages = MagicMock(return_value=([], None)) runner.moderation_for_inputs = MagicMock(return_value=(None, app_generate_entity.inputs, "query")) runner.check_hosting_moderation = MagicMock(return_value=False) @@ -168,8 +188,7 @@ class TestCompletionAppRunner: model_instance.invoke_llm.side_effect = invoke_llm mocker.patch.object(module, "ModelInstance", return_value=model_instance) - with patched_create_session(return_value=app_record): - runner.run(app_generate_entity, queue_manager, MagicMock(id="msg"), session) + runner.run(app_generate_entity, queue_manager, MagicMock(id="msg"), session) assert events == ["commit", "close", "invoke", "first-chunk"] runner._handle_invoke_result.assert_called_once_with( @@ -178,11 +197,11 @@ class TestCompletionAppRunner: stream=True, message_id="msg", user_id="user", - tenant_id="tenant", + tenant_id=TENANT_ID, ) - def test_run_uses_low_image_detail_default(self, runner, mocker: MockerFixture): - app_record = MagicMock(id="app1", tenant_id="tenant") + def test_run_uses_low_image_detail_default(self, runner, mocker: MockerFixture, sqlite_session: Session): + _persist_app(sqlite_session) app_config = _build_app_config() app_generate_entity = _build_generate_entity(app_config, file_upload_config=None) @@ -191,8 +210,7 @@ class TestCompletionAppRunner: runner.moderation_for_inputs = MagicMock(return_value=(None, app_generate_entity.inputs, "query")) runner.check_hosting_moderation = MagicMock(return_value=True) - with patched_create_session(return_value=app_record): - runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), MagicMock()) + runner.run(app_generate_entity, MagicMock(), MagicMock(id="msg"), sqlite_session) assert ( runner.organize_prompt_messages.call_args.kwargs["image_detail_config"] diff --git a/api/tests/unit_tests/core/app/apps/completion/test_completion_completion_app_generator.py b/api/tests/unit_tests/core/app/apps/completion/test_completion_completion_app_generator.py index 22a4030183c..4bca34ede39 100644 --- a/api/tests/unit_tests/core/app/apps/completion/test_completion_completion_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/completion/test_completion_completion_app_generator.py @@ -1,20 +1,25 @@ import contextlib +from decimal import Decimal from types import SimpleNamespace -from unittest.mock import MagicMock, call +from unittest.mock import MagicMock import pytest from pydantic import ValidationError from pytest_mock import MockerFixture +from sqlalchemy.orm import Session import core.app.apps.completion.app_generator as module from core.app.apps.completion.app_generator import CompletionAppGenerator from core.app.apps.exc import GenerateTaskStoppedError from core.app.entities.app_invoke_entities import InvokeFrom -from graphon.file import FILE_MODEL_IDENTITY from graphon.model_runtime.errors.invoke import InvokeAuthorizationError +from models.enums import ConversationFromSource +from models.model import AppMode, AppModelConfig, Conversation, Message from services.errors.app import MoreLikeThisDisabledError from services.errors.message import MessageNotExistsError +TABLES = (AppModelConfig, Conversation, Message) + @pytest.fixture def generator(mocker: MockerFixture): @@ -53,11 +58,71 @@ def _build_app_model_config(): return config +def _persist_message( + session: Session, + *, + message_id: str = "msg", + app_id: str = "app1", + app_model_config: AppModelConfig | None = None, +) -> Message: + conversation = Conversation( + app_id=app_id, + app_model_config_id=app_model_config.id if app_model_config else None, + model_provider=None, + override_model_configs=None, + model_id=None, + mode=AppMode.COMPLETION, + name="completion conversation", + summary=None, + inputs={}, + introduction="", + system_instruction="", + invoke_from=InvokeFrom.WEB_APP, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id=None, + from_account_id=None, + read_at=None, + read_account_id=None, + ) + session.add(conversation) + session.flush() + + message = Message( + id=message_id, + app_id=app_id, + model_provider=None, + model_id=None, + override_model_configs=None, + conversation_id=conversation.id, + inputs={"a": 1}, + query="q", + message={}, + message_unit_price=Decimal(0), + answer="", + answer_unit_price=Decimal(0), + parent_message_id=None, + total_price=None, + currency="USD", + error=None, + message_metadata=None, + invoke_from=InvokeFrom.WEB_APP, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id=None, + from_account_id=None, + workflow_run_id=None, + app_mode=AppMode.COMPLETION, + ) + session.add(message) + session.flush() + return message + + +@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True) class TestCompletionAppGenerator: - def test_generate_invalid_query_type(self, generator): + def test_generate_invalid_query_type(self, generator, sqlite_session: Session): with pytest.raises(ValueError): generator.generate( - session=MagicMock(), + session=sqlite_session, app_model=_build_app_model(), user=_build_user(), args={"query": 123, "inputs": {}, "files": []}, @@ -65,10 +130,10 @@ class TestCompletionAppGenerator: streaming=True, ) - def test_generate_override_not_debugger(self, generator): + def test_generate_override_not_debugger(self, generator, sqlite_session: Session): with pytest.raises(ValueError): generator.generate( - session=MagicMock(), + session=sqlite_session, app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {}, "files": [], "model_config": {}}, @@ -76,7 +141,7 @@ class TestCompletionAppGenerator: streaming=False, ) - def test_generate_success_no_file_config(self, generator, mocker: MockerFixture): + def test_generate_success_no_file_config(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model_config = _build_app_model_config() mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config) annotation_reply = {"enabled": False} @@ -105,9 +170,8 @@ class TestCompletionAppGenerator: mocker.patch.object(generator, "_handle_response", return_value="response") mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") - session = MagicMock() result = generator.generate( - session=session, + session=sqlite_session, app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {"a": 1}, "files": [], "trace_session_id": "session-1"}, @@ -118,11 +182,11 @@ class TestCompletionAppGenerator: assert result == "converted" assert generator.generate_entity.call_args.kwargs["extras"]["trace_session_id"] == "session-1" module.file_factory.build_from_mappings.assert_not_called() - load_annotation_reply_config.assert_called_once_with(session, "app1") + load_annotation_reply_config.assert_called_once_with(sqlite_session, "app1") app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) assert get_app_config.call_args.kwargs["annotation_reply"] is annotation_reply - def test_generate_success_with_files(self, generator, mocker: MockerFixture): + def test_generate_success_with_files(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model_config = _build_app_model_config() mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config) @@ -144,7 +208,7 @@ class TestCompletionAppGenerator: mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") result = generator.generate( - session=MagicMock(), + session=sqlite_session, app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {"a": 1}, "files": [{"id": "f"}]}, @@ -155,7 +219,7 @@ class TestCompletionAppGenerator: assert result == "converted" module.file_factory.build_from_mappings.assert_called_once() - def test_generate_override_model_config_debugger(self, generator, mocker: MockerFixture): + def test_generate_override_model_config_debugger(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model_config = _build_app_model_config() mocker.patch.object(generator, "_get_app_model_config", return_value=app_model_config) @@ -180,7 +244,7 @@ class TestCompletionAppGenerator: mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") generator.generate( - session=MagicMock(), + session=sqlite_session, app_model=_build_app_model(), user=_build_user(), args={"query": "q", "inputs": {}, "files": [], "model_config": override_config}, @@ -190,118 +254,90 @@ class TestCompletionAppGenerator: assert get_app_config.call_args.kwargs["override_config_dict"] == override_config - def test_generate_more_like_this_message_not_found(self, generator, mocker: MockerFixture): - session = mocker.MagicMock() - session.scalar.return_value = None + def test_generate_more_like_this_message_not_found(self, generator, sqlite_session: Session): + _persist_message(sqlite_session, app_id="other-app") with pytest.raises(MessageNotExistsError): generator.generate_more_like_this( - session=session, + session=sqlite_session, app_model=_build_app_model(), message_id="msg", user=_build_user(), invoke_from=InvokeFrom.WEB_APP, ) - def test_generate_more_like_this_disabled(self, generator, mocker: MockerFixture): + def test_generate_more_like_this_disabled(self, generator, sqlite_session: Session): app_model = _build_app_model() - current_config = MagicMock(more_like_this=False, more_like_this_dict={"enabled": False}) - - message = MagicMock() - session = mocker.MagicMock() - session.scalar.return_value = message - session.get.return_value = current_config + current_config = AppModelConfig(app_id=app_model.id, more_like_this='{"enabled": false}') + sqlite_session.add(current_config) + sqlite_session.flush() + app_model.app_model_config_id = current_config.id + _persist_message(sqlite_session) with pytest.raises(MoreLikeThisDisabledError): generator.generate_more_like_this( - session=session, + session=sqlite_session, app_model=app_model, message_id="msg", user=_build_user(), invoke_from=InvokeFrom.WEB_APP, ) - def test_generate_more_like_this_app_model_config_missing(self, generator, mocker: MockerFixture): + def test_generate_more_like_this_app_model_config_missing(self, generator, sqlite_session: Session): app_model = _build_app_model() app_model.app_model_config_id = None - message = MagicMock() - session = mocker.MagicMock() - session.scalar.return_value = message + _persist_message(sqlite_session) with pytest.raises(MoreLikeThisDisabledError): generator.generate_more_like_this( - session=session, + session=sqlite_session, app_model=app_model, message_id="msg", user=_build_user(), invoke_from=InvokeFrom.WEB_APP, ) - def test_generate_more_like_this_message_config_none(self, generator, mocker: MockerFixture): + def test_generate_more_like_this_message_config_none(self, generator, sqlite_session: Session): app_model = _build_app_model() - current_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True}) - - message = MagicMock(conversation_id="conv-1") - conversation = MagicMock(app_model_config_id=None) - session = mocker.MagicMock() - session.scalar.return_value = message - session.get.side_effect = [current_config, conversation] + current_config = AppModelConfig(app_id=app_model.id, more_like_this='{"enabled": true}') + sqlite_session.add(current_config) + sqlite_session.flush() + app_model.app_model_config_id = current_config.id + _persist_message(sqlite_session) with pytest.raises(ValueError): generator.generate_more_like_this( - session=session, + session=sqlite_session, app_model=app_model, message_id="msg", user=_build_user(), invoke_from=InvokeFrom.WEB_APP, ) - def test_generate_more_like_this_success(self, generator, mocker: MockerFixture): + def test_generate_more_like_this_success(self, generator, mocker: MockerFixture, sqlite_session: Session): app_model = _build_app_model() - current_config = MagicMock(more_like_this=True, more_like_this_dict={"enabled": True}) - - message = module.Message(id="msg", app_id="app1", conversation_id="conv-1", query="q") - message.inputs = {"attachment": {"dify_model_identity": FILE_MODEL_IDENTITY}} - message_files = [{"id": "f"}] - message_files_with_session = mocker.patch.object( - module.Message, - "message_files_with_session", - return_value=message_files, + app_model_config = AppModelConfig(app_id=app_model.id, more_like_this='{"enabled": true}') + sqlite_session.add(app_model_config) + sqlite_session.flush() + app_model.app_model_config_id = app_model_config.id + _persist_message(sqlite_session, app_model_config=app_model_config) + to_dict = mocker.patch.object( + AppModelConfig, + "to_dict", + return_value={ + "model": {"completion_params": {"temperature": 0.1}}, + "file_upload": {"enabled": True}, + }, ) - - app_model_config = MagicMock(app_id="app1") - app_model_config.to_dict.return_value = { - "model": {"completion_params": {"temperature": 0.1}}, - "file_upload": {"enabled": True}, - } annotation_reply = {"enabled": False} load_annotation_reply_config = mocker.patch.object( module, "load_annotation_reply_config", return_value=annotation_reply, ) - conversation = MagicMock(app_model_config_id="cfg-message") - - session = mocker.MagicMock() - session.scalar.side_effect = [message, "tenant"] - session.get.side_effect = [current_config, conversation, app_model_config] - - global_session = MagicMock() - global_session.scalar.side_effect = AssertionError("global session must not be used") - global_session.scalars.side_effect = AssertionError("global session must not be used") - mocker.patch.object(module.db, "session", global_session) - - def restore_input_file(*, file_mapping, tenant_resolver): - assert file_mapping["dify_model_identity"] == FILE_MODEL_IDENTITY - assert tenant_resolver() == "tenant" - return "input-file" - - mocker.patch("models.model.build_file_from_input_mapping", side_effect=restore_input_file) - - file_extra_config = MagicMock() - mocker.patch.object(module.FileUploadConfigManager, "convert", return_value=file_extra_config) - build_from_mappings = mocker.patch.object(module.file_factory, "build_from_mappings", return_value=["file1"]) + mocker.patch.object(module.FileUploadConfigManager, "convert", return_value=None) + build_from_mappings = mocker.patch.object(module.file_factory, "build_from_mappings") app_config = MagicMock(variables=["v"], to_dict=MagicMock(return_value={})) get_app_config = mocker.patch.object( @@ -320,7 +356,7 @@ class TestCompletionAppGenerator: mocker.patch.object(module.CompletionAppGenerateResponseConverter, "convert", return_value="converted") result = generator.generate_more_like_this( - session=session, + session=sqlite_session, app_model=app_model, message_id="msg", user=_build_user(), @@ -329,25 +365,12 @@ class TestCompletionAppGenerator: ) assert result == "converted" - assert session.get.call_args_list == [ - call(module.AppModelConfig, "cfg-current"), - call(module.Conversation, "conv-1"), - call(module.AppModelConfig, "cfg-message"), - ] - load_annotation_reply_config.assert_called_once_with(session, "app1") - app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) - assert session.scalar.call_count == 2 - message_files_with_session.assert_called_once_with(session=session) - build_from_mappings.assert_called_once_with( - mappings=message_files, - tenant_id="tenant", - config=file_extra_config, - access_controller=generator._file_access_controller, - ) - assert global_session.mock_calls == [] - assert generator.generate_entity.call_args.kwargs["inputs"] == {"attachment": "input-file"} + load_annotation_reply_config.assert_called_once_with(sqlite_session, app_model.id) + to_dict.assert_called_once_with(annotation_reply=annotation_reply) + assert generator.generate_entity.call_args.kwargs["inputs"] == {"a": 1} override_dict = get_app_config.call_args.kwargs["override_config_dict"] assert override_dict["model"]["completion_params"]["temperature"] == 0.9 + build_from_mappings.assert_not_called() @pytest.mark.parametrize( ("error", "should_publish"), @@ -365,13 +388,19 @@ class TestCompletionAppGenerator: (RuntimeError("boom"), True), ], ) - def test_generate_worker_error_handling(self, generator, mocker: MockerFixture, error, should_publish): + def test_generate_worker_error_handling( + self, + generator, + mocker: MockerFixture, + sqlite_session: Session, + error, + should_publish, + ): flask_app = MagicMock() flask_app.app_context.return_value = contextlib.nullcontext() - session = mocker.MagicMock() session_context = mocker.MagicMock() - session_context.__enter__.return_value = session + session_context.__enter__.return_value = sqlite_session create_session = mocker.patch.object(module.session_factory, "create_session") create_session.return_value = session_context mocker.patch.object(module.db, "session") diff --git a/api/tests/unit_tests/core/app/apps/pipeline/test_pipeline_generator.py b/api/tests/unit_tests/core/app/apps/pipeline/test_pipeline_generator.py index f2b8179160b..f1edcd1a1b1 100644 --- a/api/tests/unit_tests/core/app/apps/pipeline/test_pipeline_generator.py +++ b/api/tests/unit_tests/core/app/apps/pipeline/test_pipeline_generator.py @@ -4,12 +4,24 @@ from unittest.mock import MagicMock, PropertyMock import pytest from pytest_mock import MockerFixture +from sqlalchemy import select +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker import core.app.apps.pipeline.pipeline_generator as module from core.app.apps.exc import GenerateTaskStoppedError from core.app.entities.app_invoke_entities import InvokeFrom from core.datasource.entities.datasource_entities import DatasourceProviderType -from models.enums import DataSourceType +from models.dataset import Document, DocumentPipelineExecutionLog +from models.enums import DataSourceType, EndUserType +from models.model import EndUser +from models.workflow import Workflow, WorkflowType + +TENANT_ID = "00000000-0000-0000-0000-000000000001" +PIPELINE_ID = "00000000-0000-0000-0000-000000000002" +DATASET_ID = "00000000-0000-0000-0000-000000000003" +WORKFLOW_ID = "00000000-0000-0000-0000-000000000004" +USER_ID = "00000000-0000-0000-0000-000000000005" class FakeRagPipelineGenerateEntity(SimpleNamespace): @@ -24,9 +36,10 @@ class FakeRagPipelineGenerateEntity(SimpleNamespace): @pytest.fixture -def generator(mocker: MockerFixture): +def generator(mocker: MockerFixture, sqlite_engine: Engine): gen = module.PipelineGenerator() + _patch_sqlite_engine(mocker, sqlite_engine) mocker.patch.object(module, "RagPipelineGenerateEntity", FakeRagPipelineGenerateEntity) mocker.patch.object(module, "RagPipelineInvokeEntity", side_effect=lambda **kwargs: kwargs) mocker.patch.object(module.contexts, "plugin_tool_providers", SimpleNamespace(set=MagicMock())) @@ -37,27 +50,27 @@ def generator(mocker: MockerFixture): def _build_pipeline_dataset(): return SimpleNamespace( - id="ds", + id=DATASET_ID, name="dataset", description="desc", - chunk_structure="chunk", + chunk_structure="text_model", built_in_field_enabled=True, - tenant_id="tenant", + tenant_id=TENANT_ID, ) def _build_pipeline(): - pipeline = MagicMock(tenant_id="tenant", id="pipe") + pipeline = MagicMock(tenant_id=TENANT_ID, id=PIPELINE_ID) pipeline.retrieve_dataset.return_value = _build_pipeline_dataset() return pipeline def _build_workflow(): - return MagicMock(id="wf", graph_dict={"nodes": [], "edges": []}, tenant_id="tenant") + return MagicMock(id=WORKFLOW_ID, graph_dict={"nodes": [], "edges": []}, tenant_id=TENANT_ID) def _build_user(): - return MagicMock(id="user", name="User", session_id="session") + return SimpleNamespace(id=USER_ID, name="User", session_id="session") def _build_args(): @@ -69,39 +82,51 @@ def _build_args(): } -def _patch_session(mocker, session): - mocker.patch.object(module, "Session", return_value=session) - mocker.patch.object(type(module.db), "engine", new_callable=PropertyMock, return_value=MagicMock()) +def _patch_sqlite_engine(mocker: MockerFixture, sqlite_engine: Engine) -> None: + mocker.patch.object(type(module.db), "engine", new_callable=PropertyMock, return_value=sqlite_engine) + + +def _patch_db_session(mocker: MockerFixture, sqlite_session_factory: sessionmaker[Session]) -> scoped_session[Session]: + session_proxy = scoped_session(sqlite_session_factory) + mocker.patch.object(module.db, "session", session_proxy) + return session_proxy + + +def _persist_worker_records(session: Session) -> None: + workflow = Workflow( + id=WORKFLOW_ID, + tenant_id=TENANT_ID, + app_id=PIPELINE_ID, + type=WorkflowType.RAG_PIPELINE, + version=Workflow.VERSION_DRAFT, + graph="{}", + _features="{}", + created_by=USER_ID, + ) + end_user = EndUser( + id=USER_ID, + tenant_id=TENANT_ID, + app_id=PIPELINE_ID, + type=EndUserType.BROWSER, + session_id="session", + name="User", + is_anonymous=True, + ) + session.add_all([workflow, end_user]) + session.commit() def _dummy_preserve(*args, **kwargs): return contextlib.nullcontext() -class DummySession: - def __init__(self): - self.scalar = MagicMock() - self.add = MagicMock() - self.flush = MagicMock() - self.commit = MagicMock() - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb): - return False - - -def test_generate_dataset_missing(generator, mocker: MockerFixture): +def test_generate_dataset_missing(generator, sqlite_session: Session): pipeline = _build_pipeline() pipeline.retrieve_dataset.return_value = None - session = DummySession() - _patch_session(mocker, session) - with pytest.raises(ValueError): generator.generate( - session=session, + session=sqlite_session, pipeline=pipeline, workflow=_build_workflow(), user=_build_user(), @@ -111,13 +136,10 @@ def test_generate_dataset_missing(generator, mocker: MockerFixture): ) -def test_generate_debugger_calls_generate(generator, mocker: MockerFixture): +def test_generate_debugger_calls_generate(generator, mocker: MockerFixture, sqlite_session: Session): pipeline = _build_pipeline() workflow = _build_workflow() - session = DummySession() - _patch_session(mocker, session) - mocker.patch.object( generator, "_format_datasource_info_list", @@ -144,7 +166,7 @@ def test_generate_debugger_calls_generate(generator, mocker: MockerFixture): mocker.patch.object(generator, "_generate", return_value={"result": "ok"}) result = generator.generate( - session=session, + session=sqlite_session, pipeline=pipeline, workflow=workflow, user=_build_user(), @@ -156,13 +178,12 @@ def test_generate_debugger_calls_generate(generator, mocker: MockerFixture): assert result == {"result": "ok"} -def test_generate_published_pipeline_creates_documents_and_delay(generator, mocker: MockerFixture): +def test_generate_published_pipeline_creates_documents_and_delay( + generator, mocker: MockerFixture, sqlite_session: Session +): pipeline = _build_pipeline() workflow = _build_workflow() - session = DummySession() - _patch_session(mocker, session) - datasource_info_list = [{"name": "file1"}, {"name": "file2"}] mocker.patch.object( @@ -179,33 +200,9 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock mocker.patch("services.dataset_service.DocumentService.get_documents_position", return_value=1) features = SimpleNamespace() - mocker.patch("services.feature_service.FeatureService.get_features", return_value=features) + get_features = mocker.patch("services.feature_service.FeatureService.get_features", return_value=features) check_limits = mocker.patch("services.dataset_service.DocumentService.check_document_creation_limits") - document1 = SimpleNamespace( - id="doc1", - position=1, - data_source_type=DatasourceProviderType.LOCAL_FILE, - data_source_info="{}", - name="file1", - indexing_status="", - error=None, - enabled=True, - ) - document2 = SimpleNamespace( - id="doc2", - position=2, - data_source_type=DatasourceProviderType.LOCAL_FILE, - data_source_info="{}", - name="file2", - indexing_status="", - error=None, - enabled=True, - ) - mocker.patch.object(generator, "_build_document", side_effect=[document1, document2]) - - mocker.patch.object(module, "DocumentPipelineExecutionLog", return_value=MagicMock()) - mocker.patch.object( module.DifyCoreRepositoryFactory, "create_workflow_execution_repository", @@ -221,7 +218,7 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock mocker.patch.object(module, "RagPipelineTaskProxy", return_value=task_proxy) result = generator.generate( - session=session, + session=sqlite_session, pipeline=pipeline, workflow=workflow, user=_build_user(), @@ -233,18 +230,24 @@ def test_generate_published_pipeline_creates_documents_and_delay(generator, mock assert result["batch"] assert len(result["documents"]) == 2 check_limits.assert_called_once_with(len(datasource_info_list), features) - session.flush.assert_called_once_with() - session.commit.assert_called_once_with() + persisted_documents = sqlite_session.scalars( + select(Document).where(Document.dataset_id == DATASET_ID).order_by(Document.position) + ).all() + assert [document.name for document in persisted_documents] == ["file1", "file2"] + persisted_logs = sqlite_session.scalars( + select(DocumentPipelineExecutionLog).where(DocumentPipelineExecutionLog.pipeline_id == PIPELINE_ID) + ).all() + assert {log.document_id for log in persisted_logs} == {document.id for document in persisted_documents} task_proxy.delay.assert_called_once() + get_features.assert_called_once_with(TENANT_ID) -def test_generate_published_pipeline_rejects_when_document_creation_limits_exceeded(generator, mocker: MockerFixture): +def test_generate_published_pipeline_rejects_when_document_creation_limits_exceeded( + generator, mocker: MockerFixture, sqlite_session: Session +): pipeline = _build_pipeline() workflow = _build_workflow() - session = DummySession() - _patch_session(mocker, session) - datasource_info_list = [{"name": "file1"}, {"name": "file2"}] mocker.patch.object( generator, @@ -266,7 +269,7 @@ def test_generate_published_pipeline_rejects_when_document_creation_limits_excee with pytest.raises(ValueError, match="document limit exceeded"): generator.generate( - session=session, + session=sqlite_session, pipeline=pipeline, workflow=workflow, user=_build_user(), @@ -276,16 +279,13 @@ def test_generate_published_pipeline_rejects_when_document_creation_limits_excee ) check_limits.assert_called_once_with(len(datasource_info_list), features) - session.add.assert_not_called() + assert sqlite_session.scalars(select(Document)).all() == [] -def test_generate_is_retry_calls_generate(generator, mocker: MockerFixture): +def test_generate_is_retry_calls_generate(generator, mocker: MockerFixture, sqlite_session: Session): pipeline = _build_pipeline() workflow = _build_workflow() - session = DummySession() - _patch_session(mocker, session) - mocker.patch.object( generator, "_format_datasource_info_list", @@ -309,39 +309,48 @@ def test_generate_is_retry_calls_generate(generator, mocker: MockerFixture): return_value=MagicMock(), ) - mocker.patch.object(generator, "_generate", return_value={"result": "ok"}) + generate = mocker.patch.object(generator, "_generate", return_value={"result": "ok"}) + + args = _build_args() + args["original_document_id"] = "document-1" result = generator.generate( - session=session, + session=sqlite_session, pipeline=pipeline, workflow=workflow, user=_build_user(), - args=_build_args(), + args=args, invoke_from=InvokeFrom.PUBLISHED_PIPELINE, streaming=True, is_retry=True, ) assert result == {"result": "ok"} + application_generate_entity = generate.call_args.kwargs["application_generate_entity"] + assert application_generate_entity.document_id == "document-1" + assert application_generate_entity.original_document_id is None -def test_generate_worker_handles_errors(generator, mocker: MockerFixture): +def test_generate_worker_handles_errors( + generator, + mocker: MockerFixture, + sqlite_session: Session, + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], +): flask_app = MagicMock() flask_app.app_context.return_value = contextlib.nullcontext() mocker.patch.object(module, "preserve_flask_contexts", _dummy_preserve) - mocker.patch.object(module.db, "session", MagicMock(close=MagicMock())) - mocker.patch.object(type(module.db), "engine", new_callable=PropertyMock, return_value=MagicMock()) + _persist_worker_records(sqlite_session) + _patch_sqlite_engine(mocker, sqlite_engine) + _patch_db_session(mocker, sqlite_session_factory) application_generate_entity = FakeRagPipelineGenerateEntity( - app_config=SimpleNamespace(tenant_id="tenant", app_id="pipe", workflow_id="wf"), + app_config=SimpleNamespace(tenant_id=TENANT_ID, app_id=PIPELINE_ID, workflow_id=WORKFLOW_ID), invoke_from=InvokeFrom.WEB_APP, - user_id="user", + user_id=USER_ID, ) - session = DummySession() - session.scalar.side_effect = [MagicMock(), MagicMock(session_id="session")] - _patch_session(mocker, session) - runner_instance = MagicMock() runner_instance.run.side_effect = ValueError("bad") mocker.patch.object(module, "PipelineRunner", return_value=runner_instance) @@ -360,23 +369,26 @@ def test_generate_worker_handles_errors(generator, mocker: MockerFixture): queue_manager.publish_error.assert_called_once() -def test_generate_worker_sets_system_user_id_for_external_call(generator, mocker: MockerFixture): +def test_generate_worker_sets_system_user_id_for_external_call( + generator, + mocker: MockerFixture, + sqlite_session: Session, + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], +): flask_app = MagicMock() flask_app.app_context.return_value = contextlib.nullcontext() mocker.patch.object(module, "preserve_flask_contexts", _dummy_preserve) - mocker.patch.object(module.db, "session", MagicMock(close=MagicMock())) - mocker.patch.object(type(module.db), "engine", new_callable=PropertyMock, return_value=MagicMock()) + _persist_worker_records(sqlite_session) + _patch_sqlite_engine(mocker, sqlite_engine) + _patch_db_session(mocker, sqlite_session_factory) application_generate_entity = FakeRagPipelineGenerateEntity( - app_config=SimpleNamespace(tenant_id="tenant", app_id="pipe", workflow_id="wf"), + app_config=SimpleNamespace(tenant_id=TENANT_ID, app_id=PIPELINE_ID, workflow_id=WORKFLOW_ID), invoke_from=InvokeFrom.WEB_APP, - user_id="user", + user_id=USER_ID, ) - session = DummySession() - session.scalar.side_effect = [MagicMock(), MagicMock(session_id="session")] - _patch_session(mocker, session) - runner_instance = MagicMock() mocker.patch.object(module, "PipelineRunner", return_value=runner_instance) @@ -393,13 +405,11 @@ def test_generate_worker_sets_system_user_id_for_external_call(generator, mocker assert module.PipelineRunner.call_args.kwargs["system_user_id"] == "session" -def test_generate_raises_when_workflow_not_found(generator, mocker: MockerFixture): +def test_generate_raises_when_workflow_not_found(generator, mocker: MockerFixture, sqlite_session: Session): flask_app = MagicMock() mocker.patch.object(module, "preserve_flask_contexts", _dummy_preserve) - session = MagicMock() - session.get.return_value = None - mocker.patch.object(module.db, "session", session) + session = sqlite_session with pytest.raises(ValueError): generator._generate( @@ -422,19 +432,29 @@ def test_generate_raises_when_workflow_not_found(generator, mocker: MockerFixtur ) -def test_generate_success_returns_converted(generator, mocker: MockerFixture): +def test_generate_success_returns_converted(generator, mocker: MockerFixture, sqlite_session: Session): flask_app = MagicMock() mocker.patch.object(module, "preserve_flask_contexts", _dummy_preserve) - workflow = MagicMock(id="wf", tenant_id="tenant", app_id="pipe", graph_dict={}) - session = MagicMock() - session.get.return_value = workflow - mocker.patch.object(module.db, "session", session) + workflow = Workflow( + id="00000000-0000-0000-0000-000000000001", + tenant_id="00000000-0000-0000-0000-000000000002", + app_id="00000000-0000-0000-0000-000000000003", + type=WorkflowType.RAG_PIPELINE, + version=Workflow.VERSION_DRAFT, + graph="{}", + _features="{}", + created_by="00000000-0000-0000-0000-000000000004", + ) + sqlite_session.add(workflow) + sqlite_session.commit() + session = sqlite_session queue_manager = MagicMock() mocker.patch.object(module, "PipelineQueueManager", return_value=queue_manager) worker_thread = MagicMock() + worker_thread.is_alive.return_value = False mocker.patch.object(module.threading, "Thread", return_value=worker_thread) mocker.patch.object(generator, "_get_draft_var_saver_factory", return_value=MagicMock()) @@ -446,7 +466,7 @@ def test_generate_success_returns_converted(generator, mocker: MockerFixture): flask_app=flask_app, context=contextlib.nullcontext(), pipeline=_build_pipeline(), - workflow_id="wf", + workflow_id=workflow.id, user=_build_user(), application_generate_entity=FakeRagPipelineGenerateEntity( task_id="t", @@ -461,12 +481,13 @@ def test_generate_success_returns_converted(generator, mocker: MockerFixture): ) assert result == "converted" + worker_thread.join.assert_called_once_with(timeout=300) -def test_single_iteration_generate_validates_inputs(generator, mocker: MockerFixture): +def test_single_iteration_generate_validates_inputs(generator, sqlite_session: Session): with pytest.raises(ValueError): generator.single_iteration_generate( - _build_pipeline(), _build_workflow(), "", _build_user(), {}, session=DummySession() + _build_pipeline(), _build_workflow(), "", _build_user(), {}, session=sqlite_session ) with pytest.raises(ValueError): @@ -476,17 +497,14 @@ def test_single_iteration_generate_validates_inputs(generator, mocker: MockerFix "node", _build_user(), {"inputs": None}, - session=DummySession(), + session=sqlite_session, ) -def test_single_iteration_generate_dataset_required(generator, mocker: MockerFixture): +def test_single_iteration_generate_dataset_required(generator, sqlite_session: Session): pipeline = _build_pipeline() pipeline.retrieve_dataset.return_value = None - session = DummySession() - _patch_session(mocker, session) - with pytest.raises(ValueError): generator.single_iteration_generate( pipeline, @@ -494,16 +512,18 @@ def test_single_iteration_generate_dataset_required(generator, mocker: MockerFix "node", _build_user(), {"inputs": {"a": 1}}, - session=session, + session=sqlite_session, ) -def test_single_iteration_generate_success(generator, mocker: MockerFixture): +def test_single_iteration_generate_success( + generator, + mocker: MockerFixture, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +): pipeline = _build_pipeline() - session = DummySession() - _patch_session(mocker, session) - mocker.patch.object( module.PipelineConfigManager, "get_pipeline_config", @@ -519,7 +539,7 @@ def test_single_iteration_generate_success(generator, mocker: MockerFixture): "create_workflow_node_execution_repository", return_value=MagicMock(), ) - mocker.patch.object(module.db, "session", MagicMock(return_value=MagicMock())) + _patch_db_session(mocker, sqlite_session_factory) mocker.patch.object(module, "WorkflowDraftVariableService", return_value=MagicMock()) mocker.patch.object(module, "DraftVarLoader", return_value=MagicMock()) @@ -533,18 +553,20 @@ def test_single_iteration_generate_success(generator, mocker: MockerFixture): _build_user(), {"inputs": {"a": 1}}, streaming=False, - session=session, + session=sqlite_session, ) assert result == {"ok": True} -def test_single_loop_generate_success(generator, mocker: MockerFixture): +def test_single_loop_generate_success( + generator, + mocker: MockerFixture, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +): pipeline = _build_pipeline() - session = DummySession() - _patch_session(mocker, session) - mocker.patch.object( module.PipelineConfigManager, "get_pipeline_config", @@ -560,7 +582,7 @@ def test_single_loop_generate_success(generator, mocker: MockerFixture): "create_workflow_node_execution_repository", return_value=MagicMock(), ) - mocker.patch.object(module.db, "session", MagicMock(return_value=MagicMock())) + _patch_db_session(mocker, sqlite_session_factory) mocker.patch.object(module, "WorkflowDraftVariableService", return_value=MagicMock()) mocker.patch.object(module, "DraftVarLoader", return_value=MagicMock()) @@ -574,7 +596,7 @@ def test_single_loop_generate_success(generator, mocker: MockerFixture): _build_user(), {"inputs": {"a": 1}}, streaming=False, - session=session, + session=sqlite_session, ) assert result == {"ok": True} diff --git a/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py b/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py index d9e66cd26a6..c1e16084ffb 100644 --- a/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/test_advanced_chat_app_generator.py @@ -4,16 +4,17 @@ from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker from core.app.app_config.entities import AppAdditionalFeatures, WorkflowUIBasedAppConfig -from core.app.apps import message_based_app_generator from core.app.apps.advanced_chat.app_generator import AdvancedChatAppGenerator from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom from core.app.task_pipeline import message_cycle_manager from core.app.task_pipeline.message_cycle_manager import MessageCycleManager from core.ops.ops_trace_manager import TraceQueueManager from models.enums import ConversationFromSource -from models.model import AppMode, Conversation, Message +from models.model import AppMode, Conversation from services.errors.conversation import ConversationNotExistsError @@ -46,25 +47,7 @@ def _make_generate_entity(app_config: WorkflowUIBasedAppConfig) -> AdvancedChatA ) -@pytest.fixture(autouse=True) -def mock_db_session(monkeypatch: pytest.MonkeyPatch): - session = MagicMock() - - def refresh_side_effect(obj): - if isinstance(obj, Conversation) and obj.id is None: - obj.id = "generated-conversation-id" - if isinstance(obj, Message) and obj.id is None: - obj.id = "generated-message-id" - - session.refresh.side_effect = refresh_side_effect - session.add.return_value = None - session.commit.return_value = None - - monkeypatch.setattr(message_based_app_generator, "db", SimpleNamespace(session=session)) - return session - - -def test_init_generate_records_sets_conversation_metadata(mock_db_session): +def test_init_generate_records_sets_conversation_metadata(sqlite_session: Session): app_config = _make_app_config() entity = _make_generate_entity(app_config) @@ -73,15 +56,15 @@ def test_init_generate_records_sets_conversation_metadata(mock_db_session): conversation, _ = generator._init_generate_records( entity, conversation=None, - session=mock_db_session, + session=sqlite_session, ) - assert entity.conversation_id == "generated-conversation-id" - assert conversation.id == "generated-conversation-id" + assert entity.conversation_id == conversation.id + assert conversation.id is not None assert entity.is_new_conversation is True -def test_init_generate_records_marks_existing_conversation(mock_db_session): +def test_init_generate_records_marks_existing_conversation(sqlite_session: Session): app_config = _make_app_config() entity = _make_generate_entity(app_config) @@ -104,13 +87,15 @@ def test_init_generate_records_marks_existing_conversation(mock_db_session): from_account_id=None, ) existing_conversation.id = "existing-conversation-id" + sqlite_session.add(existing_conversation) + sqlite_session.flush() generator = AdvancedChatAppGenerator() conversation, _ = generator._init_generate_records( entity, conversation=existing_conversation, - session=mock_db_session, + session=sqlite_session, ) assert entity.conversation_id == "existing-conversation-id" @@ -118,7 +103,12 @@ def test_init_generate_records_marks_existing_conversation(mock_db_session): assert entity.is_new_conversation is False -def test_generate_falls_back_to_new_conversation_when_conversation_missing(monkeypatch: pytest.MonkeyPatch): +def test_generate_falls_back_to_new_conversation_when_conversation_missing( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +): app_config = _make_app_config() workflow = SimpleNamespace( features_dict={}, @@ -144,9 +134,10 @@ def test_generate_falls_back_to_new_conversation_when_conversation_missing(monke "core.app.apps.advanced_chat.app_generator.AdvancedChatAppConfigManager.get_app_config", lambda **_kwargs: app_config, ) + db_session = scoped_session(sqlite_session_factory) monkeypatch.setattr( "core.app.apps.advanced_chat.app_generator.db", - SimpleNamespace(engine=object(), session=lambda: MagicMock()), + SimpleNamespace(engine=sqlite_engine, session=db_session), ) trace_manager = object.__new__(TraceQueueManager) monkeypatch.setattr( @@ -163,7 +154,7 @@ def test_generate_falls_back_to_new_conversation_when_conversation_missing(monke ) captured: dict[str, object] = {} - session = MagicMock() + session = sqlite_session def fake_generate(self, **kwargs): captured.update(kwargs) @@ -188,6 +179,7 @@ def test_generate_falls_back_to_new_conversation_when_conversation_missing(monke application_generate_entity = captured["application_generate_entity"] assert isinstance(application_generate_entity, AdvancedChatAppGenerateEntity) assert application_generate_entity.conversation_id is None + db_session.remove() def test_message_cycle_manager_uses_new_conversation_flag(monkeypatch: pytest.MonkeyPatch): diff --git a/api/tests/unit_tests/core/app/apps/test_base_app_generator.py b/api/tests/unit_tests/core/app/apps/test_base_app_generator.py index 8e7468bb0b8..04724d701ef 100644 --- a/api/tests/unit_tests/core/app/apps/test_base_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/test_base_app_generator.py @@ -1,7 +1,40 @@ +import logging +from types import SimpleNamespace +from unittest.mock import Mock + import pytest +from sqlalchemy import inspect from core.app.apps.base_app_generator import BaseAppGenerator +from graphon.enums import BuiltinNodeTypes from graphon.variables.input_entities import VariableEntity, VariableEntityType +from models import Workflow, WorkflowRun + + +def test_restore_workflow_run_graph(): + workflow = Workflow(graph='{"nodes": [{"id": "edited"}]}') + session = SimpleNamespace(get=Mock(return_value=SimpleNamespace(graph='{"nodes": [{"id": "paused"}]}'))) + + BaseAppGenerator._restore_workflow_run_graph(session=session, workflow=workflow, workflow_run_id="run-id") + + session.get.assert_called_once_with(WorkflowRun, "run-id") + assert workflow.graph == '{"nodes": [{"id": "paused"}]}' + assert not inspect(workflow).attrs.graph.history.has_changes() + + +@pytest.mark.parametrize( + ("workflow_run_id", "workflow_run"), + [(None, None), ("run-id", None), ("run-id", SimpleNamespace(graph=None))], +) +def test_restore_workflow_run_graph_requires_persisted_snapshot(workflow_run_id, workflow_run): + session = SimpleNamespace(get=Mock(return_value=workflow_run)) + + with pytest.raises(ValueError): + BaseAppGenerator._restore_workflow_run_graph( + session=session, + workflow=Workflow(graph="{}"), + workflow_run_id=workflow_run_id, + ) def test_validate_inputs_with_zero(): @@ -369,6 +402,58 @@ def test_validate_inputs_optional_file_with_empty_string_ignores_default(): class TestBaseAppGeneratorExtras: + def test_wrap_stream_joins_worker_after_stream_exhaustion(self): + base_app_generator = BaseAppGenerator() + worker_thread = Mock() + worker_thread.is_alive.return_value = False + + def response_stream(): + yield {"event": "workflow_finished"} + + managed_stream = base_app_generator._wrap_stream_with_worker_thread_join( + response_stream(), + worker_thread, + ) + + assert next(managed_stream) == {"event": "workflow_finished"} + worker_thread.join.assert_not_called() + + with pytest.raises(StopIteration): + next(managed_stream) + + worker_thread.join.assert_called_once_with(timeout=300) + + def test_wrap_stream_joins_worker_when_stream_closes(self): + base_app_generator = BaseAppGenerator() + worker_thread = Mock() + worker_thread.is_alive.return_value = False + + def response_stream(): + yield {"event": "workflow_started"} + yield {"event": "workflow_finished"} + + managed_stream = base_app_generator._wrap_stream_with_worker_thread_join( + response_stream(), + worker_thread, + ) + + assert next(managed_stream) == {"event": "workflow_started"} + managed_stream.close() + + worker_thread.join.assert_called_once_with(timeout=300) + + def test_join_worker_thread_warns_when_thread_remains_alive(self, caplog: pytest.LogCaptureFixture): + worker_thread = Mock() + worker_thread.name = "leaked-app-worker" + worker_thread.is_alive.return_value = True + + with caplog.at_level(logging.WARNING, logger="core.app.apps.base_app_generator"): + BaseAppGenerator._join_worker_thread(worker_thread) + + worker_thread.join.assert_called_once_with(timeout=300) + assert "Possible app worker thread leak" in caplog.text + assert "leaked-app-worker" in caplog.text + def test_prepare_user_inputs_converts_files_and_lists(self, monkeypatch: pytest.MonkeyPatch): base_app_generator = BaseAppGenerator() @@ -477,7 +562,6 @@ class TestBaseAppGeneratorExtras: def test_get_draft_var_saver_factory_debugger(self): from core.app.entities.app_invoke_entities import InvokeFrom - from graphon.enums import BuiltinNodeTypes from models import Account base_app_generator = BaseAppGenerator() diff --git a/api/tests/unit_tests/core/app/apps/test_message_based_app_generator.py b/api/tests/unit_tests/core/app/apps/test_message_based_app_generator.py index 368be54dbf0..61f7127cca5 100644 --- a/api/tests/unit_tests/core/app/apps/test_message_based_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/test_message_based_app_generator.py @@ -1,9 +1,9 @@ from __future__ import annotations from types import SimpleNamespace -from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from core.app.app_config.entities import ( AppAdditionalFeatures, @@ -12,11 +12,10 @@ from core.app.app_config.entities import ( ModelConfigEntity, PromptTemplateEntity, ) -from core.app.apps import message_based_app_generator from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.message_based_app_generator import MessageBasedAppGenerator from core.app.entities.app_invoke_entities import ChatAppGenerateEntity, InvokeFrom -from models.model import AppMode, Conversation, Message +from models.model import AppMode from services.errors.app_model_config import AppModelConfigBrokenError @@ -84,25 +83,7 @@ def _make_chat_generate_entity(app_config: EasyUIBasedAppConfig) -> ChatAppGener ) -@pytest.fixture(autouse=True) -def mock_db_session(monkeypatch: pytest.MonkeyPatch): - session = MagicMock() - - def refresh_side_effect(obj): - if isinstance(obj, Conversation) and obj.id is None: - obj.id = "generated-conversation-id" - if isinstance(obj, Message) and obj.id is None: - obj.id = "generated-message-id" - - session.refresh.side_effect = refresh_side_effect - session.add.return_value = None - session.commit.return_value = None - - monkeypatch.setattr(message_based_app_generator, "db", SimpleNamespace(session=session)) - return session - - -def test_init_generate_records_skips_conversation_fields_for_non_conversation_entity(mock_db_session): +def test_init_generate_records_skips_conversation_fields_for_non_conversation_entity(sqlite_session: Session): app_config = _make_app_config(AppMode.COMPLETION) entity = DummyCompletionGenerateEntity(app_config=app_config) @@ -111,16 +92,16 @@ def test_init_generate_records_skips_conversation_fields_for_non_conversation_en conversation, message = generator._init_generate_records( entity, conversation=None, - session=mock_db_session, + session=sqlite_session, ) - assert conversation.id == "generated-conversation-id" - assert message.id == "generated-message-id" + assert conversation.id is not None + assert message.id is not None assert hasattr(entity, "conversation_id") is False assert hasattr(entity, "is_new_conversation") is False -def test_init_generate_records_sets_conversation_fields_for_chat_entity(mock_db_session): +def test_init_generate_records_sets_conversation_fields_for_chat_entity(sqlite_session: Session): app_config = _make_app_config(AppMode.CHAT) entity = _make_chat_generate_entity(app_config) @@ -129,12 +110,12 @@ def test_init_generate_records_sets_conversation_fields_for_chat_entity(mock_db_ conversation, _ = generator._init_generate_records( entity, conversation=None, - session=mock_db_session, + session=sqlite_session, ) - assert entity.conversation_id == "generated-conversation-id" + assert entity.conversation_id == conversation.id assert entity.is_new_conversation is True - assert conversation.id == "generated-conversation-id" + assert conversation.id is not None class TestMessageBasedAppGeneratorExtras: @@ -163,17 +144,15 @@ class TestMessageBasedAppGeneratorExtras: stream=False, ) - def test_get_app_model_config_requires_valid_config(self): + def test_get_app_model_config_requires_valid_config(self, sqlite_session: Session): generator = MessageBasedAppGenerator() app_model = SimpleNamespace(id="app", app_model_config_id=None, app_model_config=None) - session = MagicMock() + session = sqlite_session with pytest.raises(AppModelConfigBrokenError): generator._get_app_model_config(app_model, conversation=None, session=session) conversation = SimpleNamespace(app_model_config_id="missing-id") - session.scalar.return_value = None - with pytest.raises(AppModelConfigBrokenError): generator._get_app_model_config( app_model=SimpleNamespace(id="app"), diff --git a/api/tests/unit_tests/core/app/apps/test_message_based_app_queue_manager.py b/api/tests/unit_tests/core/app/apps/test_message_based_app_queue_manager.py index 847ad0ce9bc..0bc4753752d 100644 --- a/api/tests/unit_tests/core/app/apps/test_message_based_app_queue_manager.py +++ b/api/tests/unit_tests/core/app/apps/test_message_based_app_queue_manager.py @@ -6,7 +6,12 @@ from core.app.apps.base_app_queue_manager import PublishFrom from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager from core.app.entities.app_invoke_entities import InvokeFrom -from core.app.entities.queue_entities import QueueErrorEvent, QueueMessageEndEvent, QueueStopEvent +from core.app.entities.queue_entities import ( + QueueErrorEvent, + QueueMessageEndEvent, + QueueStopEvent, + QueueWorkflowPausedEvent, +) class TestMessageBasedAppQueueManager: @@ -63,3 +68,21 @@ class TestMessageBasedAppQueueManager: manager._publish(QueueMessageEndEvent(), PublishFrom.TASK_PIPELINE) assert manager._q.qsize() == 1 + + def test_publish_pause_event_stops_listener_without_aborting_execution(self): + with patch("core.app.apps.base_app_queue_manager.redis_client") as mock_redis: + mock_redis.setex.return_value = True + manager = MessageBasedAppQueueManager( + task_id="t1", + user_id="u1", + invoke_from=InvokeFrom.DEBUGGER, + conversation_id="c1", + app_mode="advanced-chat", + message_id="m1", + ) + manager.stop_listen = Mock() + manager._is_stopped = Mock(return_value=False) + + manager._publish(QueueWorkflowPausedEvent(), PublishFrom.APPLICATION_MANAGER) + + manager.stop_listen.assert_called_once_with(execution_terminal=True) diff --git a/api/tests/unit_tests/core/app/apps/test_workflow_app_generator.py b/api/tests/unit_tests/core/app/apps/test_workflow_app_generator.py index 8f5cb2b8115..43dcab8c241 100644 --- a/api/tests/unit_tests/core/app/apps/test_workflow_app_generator.py +++ b/api/tests/unit_tests/core/app/apps/test_workflow_app_generator.py @@ -1,10 +1,133 @@ +"""SQLite-backed tests for workflow app generation and worker reload behavior.""" + import contextlib +import json from types import SimpleNamespace from unittest.mock import MagicMock -from pytest_mock import MockerFixture +import pytest +from flask import Flask +from sqlalchemy.orm import Session, sessionmaker +import core.app.apps.workflow.app_generator as app_generator_module +from core.app.app_config.entities import WorkflowUIBasedAppConfig +from core.app.apps.draft_variable_saver import DraftVariableSaverFactory from core.app.apps.workflow.app_generator import SKIP_PREPARE_USER_INPUTS_KEY, WorkflowAppGenerator +from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity +from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig, PauseStatePersistenceLayer +from core.ops.ops_trace_manager import TraceQueueManager +from core.repositories import SQLAlchemyWorkflowExecutionRepository, SQLAlchemyWorkflowNodeExecutionRepository +from graphon.enums import WorkflowExecutionStatus +from graphon.runtime import GraphRuntimeState, VariablePool +from models.enums import CreatorUserRole, EndUserType, WorkflowRunTriggeredFrom +from models.model import App, AppMode, EndUser +from models.snippet import CustomizedSnippet +from models.workflow import Workflow, WorkflowKind, WorkflowNodeExecutionTriggeredFrom, WorkflowRun, WorkflowType + + +def _workflow( + *, + workflow_id: str = "workflow", + app_id: str = "app", + tenant_id: str = "tenant", + kind: WorkflowKind = WorkflowKind.STANDARD, +) -> Workflow: + return Workflow( + id=workflow_id, + tenant_id=tenant_id, + app_id=app_id, + type=WorkflowType.WORKFLOW, + kind=kind, + version="1", + graph=json.dumps({"nodes": [], "edges": []}), + features="{}", + created_by="creator", + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + + +def _persist_generator_rows(sqlite_session: Session) -> tuple[App, Workflow, EndUser]: + workflow = _workflow() + app = App( + id="app", + tenant_id="tenant", + name="Workflow app", + description="", + mode=AppMode.WORKFLOW, + workflow_id=workflow.id, + enable_site=True, + enable_api=True, + max_active_requests=0, + ) + end_user = EndUser( + id="user", + tenant_id="tenant", + app_id=app.id, + type=EndUserType.SERVICE_API, + name="End user", + session_id="end-user-session", + ) + sqlite_session.add_all([workflow, app, end_user]) + sqlite_session.commit() + return app, workflow, end_user + + +def _app_config(app: App, workflow: Workflow) -> WorkflowUIBasedAppConfig: + return WorkflowUIBasedAppConfig( + tenant_id=app.tenant_id, + app_id=app.id, + app_mode=app.mode, + workflow_id=workflow.id, + ) + + +def _generate_entity( + app: App, + workflow: Workflow, + end_user: EndUser, + *, + invoke_from: InvokeFrom = InvokeFrom.SERVICE_API, + stream: bool = True, +) -> WorkflowAppGenerateEntity: + return WorkflowAppGenerateEntity( + task_id="task", + app_config=_app_config(app, workflow), + inputs={}, + files=[], + user_id=end_user.id, + stream=stream, + invoke_from=invoke_from, + trace_manager=MagicMock(spec=TraceQueueManager), + workflow_execution_id="run", + ) + + +def _runtime_state() -> GraphRuntimeState: + return GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0) + + +def _repositories( + sqlite_session: Session, app: App, end_user: EndUser +) -> tuple[SQLAlchemyWorkflowExecutionRepository, SQLAlchemyWorkflowNodeExecutionRepository]: + session_factory = sessionmaker(bind=sqlite_session.get_bind(), expire_on_commit=False) + return ( + SQLAlchemyWorkflowExecutionRepository( + session_factory=session_factory, + tenant_id=app.tenant_id, + user=end_user, + app_id=app.id, + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, + ), + SQLAlchemyWorkflowNodeExecutionRepository( + session_factory=session_factory, + tenant_id=app.tenant_id, + user=end_user, + app_id=app.id, + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + ), + ) def test_should_prepare_user_inputs_defaults_to_true(): @@ -25,47 +148,67 @@ def test_should_prepare_user_inputs_keeps_validation_when_flag_false(): assert WorkflowAppGenerator()._should_prepare_user_inputs(args) -def test_ensure_snippet_start_node_in_worker_returns_standard_workflow_without_lookup(): - session = MagicMock() - workflow = SimpleNamespace(kind_or_standard="standard") +def test_ensure_snippet_start_node_in_worker_returns_standard_workflow_without_lookup( + sqlite_session: Session, +) -> None: + workflow = _workflow() + sqlite_session.add(workflow) + sqlite_session.commit() - result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=session, workflow=workflow) + result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=sqlite_session, workflow=workflow) assert result is workflow - session.scalar.assert_not_called() -def test_ensure_snippet_start_node_in_worker_returns_snippet_workflow_when_snippet_missing(): - session = MagicMock() - session.scalar.return_value = None - workflow = SimpleNamespace(kind_or_standard="snippet", app_id="snippet-1", tenant_id="tenant-1") +def test_ensure_snippet_start_node_in_worker_returns_snippet_workflow_when_snippet_missing( + sqlite_session: Session, +) -> None: + workflow = _workflow(app_id="snippet-1", tenant_id="tenant-1", kind=WorkflowKind.SNIPPET) + other_tenant_snippet = CustomizedSnippet( + id="snippet-1", + tenant_id="tenant-2", + name="Other tenant snippet", + description="", + type="node", + ) + sqlite_session.add_all([workflow, other_tenant_snippet]) + sqlite_session.commit() - result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=session, workflow=workflow) + result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=sqlite_session, workflow=workflow) assert result is workflow - session.scalar.assert_called_once() -def test_ensure_snippet_start_node_in_worker_applies_snippet_start_injection(mocker: MockerFixture): - session = MagicMock() - snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") - session.scalar.return_value = snippet - workflow = SimpleNamespace(kind_or_standard="snippet", app_id="snippet-1", tenant_id="tenant-1") - updated_workflow = SimpleNamespace(name="updated-workflow") - ensure_start_node = mocker.patch( +def test_ensure_snippet_start_node_in_worker_applies_snippet_start_injection( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + workflow = _workflow(app_id="snippet-1", tenant_id="tenant-1", kind=WorkflowKind.SNIPPET) + snippet = CustomizedSnippet( + id="snippet-1", + tenant_id="tenant-1", + name="Matching snippet", + description="", + type="node", + ) + sqlite_session.add_all([workflow, snippet]) + sqlite_session.commit() + ensure_start_node = MagicMock(return_value=workflow) + monkeypatch.setattr( "services.snippet_generate_service.SnippetGenerateService.ensure_start_node_for_worker", - return_value=updated_workflow, + ensure_start_node, ) - result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=session, workflow=workflow) + result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=sqlite_session, workflow=workflow) - assert result is updated_workflow - session.scalar.assert_called_once() + assert result is workflow ensure_start_node.assert_called_once_with(workflow, snippet) -def test_generate_includes_parent_trace_context_in_extras(monkeypatch): +def test_generate_includes_parent_trace_context_in_extras( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: generator = WorkflowAppGenerator() + app, workflow, end_user = _persist_generator_rows(sqlite_session) monkeypatch.setattr( "core.app.apps.workflow.app_generator.WorkflowAppGenerator._bind_file_access_scope", @@ -73,46 +216,52 @@ def test_generate_includes_parent_trace_context_in_extras(monkeypatch): ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.WorkflowAppConfigManager.get_app_config", - lambda *args, **kwargs: SimpleNamespace( - app_id="app-1", tenant_id="tenant-1", workflow_id="workflow-1", variables=[] - ), + lambda *args, **kwargs: _app_config(app, workflow), ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.file_factory.build_from_mappings", lambda *args, **kwargs: [] ) - monkeypatch.setattr("core.app.apps.workflow.app_generator.TraceQueueManager", MagicMock()) - workflow_execution_factory = MagicMock(return_value=MagicMock()) - workflow_node_execution_factory = MagicMock(return_value=MagicMock()) + monkeypatch.setattr( + "core.app.apps.workflow.app_generator.TraceQueueManager", + MagicMock(return_value=MagicMock(spec=TraceQueueManager)), + ) + repository_tenant_ids: dict[str, str] = {} + workflow_execution_factory = app_generator_module.DifyCoreRepositoryFactory.create_workflow_execution_repository + workflow_node_execution_factory = ( + app_generator_module.DifyCoreRepositoryFactory.create_workflow_node_execution_repository + ) + + def create_workflow_execution_repository(**kwargs): + repository_tenant_ids["workflow"] = kwargs["tenant_id"] + return workflow_execution_factory(**kwargs) + + def create_workflow_node_execution_repository(**kwargs): + repository_tenant_ids["node"] = kwargs["tenant_id"] + return workflow_node_execution_factory(**kwargs) + monkeypatch.setattr( "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_execution_repository", - workflow_execution_factory, + create_workflow_execution_repository, ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", - workflow_node_execution_factory, + create_workflow_node_execution_repository, ) - monkeypatch.setattr("core.app.apps.workflow.app_generator.db", SimpleNamespace(engine=MagicMock())) + monkeypatch.setattr("core.app.apps.workflow.app_generator.db", SimpleNamespace(engine=sqlite_session.get_bind())) monkeypatch.setattr(generator, "_prepare_user_inputs", lambda *, user_inputs, **kwargs: user_inputs) captured = {} - def fake_workflow_app_generate_entity(**kwargs): - captured["workflow_app_generate_entity_kwargs"] = kwargs - return SimpleNamespace(**kwargs) - def fake_generate(**kwargs): - captured["application_generate_entity"] = kwargs["application_generate_entity"] + captured.update(kwargs) return {"data": {}} - monkeypatch.setattr( - "core.app.apps.workflow.app_generator.WorkflowAppGenerateEntity", fake_workflow_app_generate_entity - ) monkeypatch.setattr(generator, "_generate", fake_generate) result = generator.generate( - app_model=SimpleNamespace(tenant_id="tenant-1", id="app-1"), - workflow=SimpleNamespace(features_dict={}), - user=SimpleNamespace(id="user-1", session_id="session-1"), + app_model=app, + workflow=workflow, + user=end_user, args={ "inputs": {"query": "hello"}, "files": [], @@ -123,42 +272,56 @@ def test_generate_includes_parent_trace_context_in_extras(monkeypatch): }, "trace_session_id": "session-1", }, - invoke_from="service-api", + invoke_from=InvokeFrom.SERVICE_API, streaming=False, call_depth=0, ) assert result == {"data": {}} - extras = captured["workflow_app_generate_entity_kwargs"]["extras"] + application_generate_entity = captured["application_generate_entity"] + assert isinstance(application_generate_entity, WorkflowAppGenerateEntity) + extras = application_generate_entity.extras assert extras["external_trace_id"] == "trace-1" assert extras["parent_trace_context"].model_dump() == { "parent_workflow_run_id": "outer-workflow-run-1", "parent_node_execution_id": "outer-node-execution-1", } assert extras["trace_session_id"] == "session-1" - assert workflow_execution_factory.call_args.kwargs["tenant_id"] == "tenant-1" - assert workflow_node_execution_factory.call_args.kwargs["tenant_id"] == "tenant-1" + assert isinstance(captured["workflow_execution_repository"], SQLAlchemyWorkflowExecutionRepository) + assert isinstance(captured["workflow_node_execution_repository"], SQLAlchemyWorkflowNodeExecutionRepository) + assert repository_tenant_ids == {"workflow": app.tenant_id, "node": app.tenant_id} -def test_resume_delegates_to_generate(mocker: MockerFixture): +def test_resume_delegates_to_generate(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: generator = WorkflowAppGenerator() - mock_generate = mocker.patch.object(generator, "_generate", return_value="ok") + app, workflow, end_user = _persist_generator_rows(sqlite_session) + mock_generate = MagicMock(return_value="ok") + monkeypatch.setattr(generator, "_generate", mock_generate) - application_generate_entity = SimpleNamespace(stream=False, invoke_from="debugger", trace_manager=MagicMock()) - runtime_state = MagicMock(name="runtime-state") - pause_config = MagicMock(name="pause-config") + application_generate_entity = _generate_entity( + app, + workflow, + end_user, + invoke_from=InvokeFrom.DEBUGGER, + stream=False, + ) + runtime_state = _runtime_state() + workflow_execution_repository, workflow_node_execution_repository = _repositories(sqlite_session, app, end_user) + pause_config = PauseStateLayerConfig( + session_factory=sessionmaker(bind=sqlite_session.get_bind(), expire_on_commit=False), + state_owner_user_id="owner", + ) result = generator.resume( - app_model=MagicMock(), - workflow=MagicMock(), - user=MagicMock(), + app_model=app, + workflow=workflow, + user=end_user, application_generate_entity=application_generate_entity, graph_runtime_state=runtime_state, - workflow_execution_repository=MagicMock(), - workflow_node_execution_repository=MagicMock(), + workflow_execution_repository=workflow_execution_repository, + workflow_node_execution_repository=workflow_node_execution_repository, graph_engine_layers=("layer",), pause_state_config=pause_config, - variable_loader=MagicMock(), ) assert result == "ok" @@ -167,39 +330,45 @@ def test_resume_delegates_to_generate(mocker: MockerFixture): assert kwargs["graph_runtime_state"] is runtime_state assert kwargs["pause_state_config"] is pause_config assert kwargs["streaming"] is False - assert kwargs["invoke_from"] == "debugger" + assert kwargs["invoke_from"] == InvokeFrom.DEBUGGER -def test_generate_appends_pause_layer_and_forwards_state(mocker: MockerFixture): +def test_generate_appends_pause_layer_and_forwards_state( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: generator = WorkflowAppGenerator() + app, workflow, end_user = _persist_generator_rows(sqlite_session) - mock_queue_manager = MagicMock() - mocker.patch("core.app.apps.workflow.app_generator.WorkflowAppQueueManager", return_value=mock_queue_manager) - - fake_current_app = MagicMock() - fake_current_app._get_current_object.return_value = MagicMock() - mocker.patch("core.app.apps.workflow.app_generator.current_app", fake_current_app) - - mocker.patch( - "core.app.apps.workflow.app_generator.WorkflowAppGenerateResponseConverter.convert", - return_value="converted", + queue_manager = MagicMock() + monkeypatch.setattr( + "core.app.apps.workflow.app_generator.WorkflowAppQueueManager", + MagicMock(return_value=queue_manager), ) - mocker.patch.object(WorkflowAppGenerator, "_handle_response", return_value="response") - draft_saver_factory = mocker.patch.object( + + monkeypatch.setattr( + "core.app.apps.workflow.app_generator.WorkflowAppGenerateResponseConverter.convert", + MagicMock(return_value="converted"), + ) + monkeypatch.setattr(WorkflowAppGenerator, "_handle_response", MagicMock(return_value="response")) + + draft_factory_tenant_ids: list[str] = [] + get_draft_var_saver_factory = WorkflowAppGenerator._get_draft_var_saver_factory + + def get_recording_draft_var_saver_factory( + invoke_from: InvokeFrom, account: EndUser, *, tenant_id: str + ) -> DraftVariableSaverFactory: + draft_factory_tenant_ids.append(tenant_id) + return get_draft_var_saver_factory(invoke_from, account, tenant_id=tenant_id) + + monkeypatch.setattr( WorkflowAppGenerator, "_get_draft_var_saver_factory", - return_value=MagicMock(), + staticmethod(get_recording_draft_var_saver_factory), ) - pause_layer = MagicMock(name="pause-layer") - mocker.patch( - "core.app.apps.workflow.app_generator.PauseStatePersistenceLayer", - return_value=pause_layer, - ) - - dummy_session = MagicMock() - dummy_session.close = MagicMock() - mocker.patch("core.app.apps.workflow.app_generator.db.session", dummy_session) + engine = sqlite_session.get_bind() + scoped_session = Session(engine, expire_on_commit=False) + monkeypatch.setattr(app_generator_module.db, "session", scoped_session) worker_kwargs: dict[str, object] = {} @@ -211,80 +380,104 @@ def test_generate_appends_pause_layer_and_forwards_state(mocker: MockerFixture): def start(self): return None - mocker.patch("core.app.apps.workflow.app_generator.threading.Thread", DummyThread) + def join(self, timeout): + worker_kwargs["joined"] = True + worker_kwargs["join_timeout"] = timeout - app_model = SimpleNamespace(mode="workflow", tenant_id="tenant") - app_config = SimpleNamespace(app_id="app", tenant_id="tenant", workflow_id="wf") - application_generate_entity = SimpleNamespace( - task_id="task", - user_id="user", - invoke_from="service-api", - app_config=app_config, - files=[], - stream=True, - workflow_execution_id="run", - ) + def is_alive(self): + return False - graph_runtime_state = MagicMock() + monkeypatch.setattr("core.app.apps.workflow.app_generator.threading.Thread", DummyThread) - result = generator._generate( - app_model=app_model, - workflow=MagicMock(), - user=MagicMock(), - application_generate_entity=application_generate_entity, - invoke_from="service-api", - workflow_execution_repository=MagicMock(), - workflow_node_execution_repository=MagicMock(), - streaming=True, - graph_engine_layers=("base-layer",), - graph_runtime_state=graph_runtime_state, - pause_state_config=SimpleNamespace(session_factory=MagicMock(), state_owner_user_id="owner"), - ) + application_generate_entity = _generate_entity(app, workflow, end_user) + graph_runtime_state = _runtime_state() + workflow_execution_repository, workflow_node_execution_repository = _repositories(sqlite_session, app, end_user) + + flask_app = Flask(__name__) + with flask_app.app_context(): + result = generator._generate( + app_model=app, + workflow=workflow, + user=end_user, + application_generate_entity=application_generate_entity, + invoke_from=InvokeFrom.SERVICE_API, + workflow_execution_repository=workflow_execution_repository, + workflow_node_execution_repository=workflow_node_execution_repository, + streaming=True, + graph_engine_layers=("base-layer",), + graph_runtime_state=graph_runtime_state, + pause_state_config=PauseStateLayerConfig( + session_factory=sessionmaker(bind=engine, expire_on_commit=False), + state_owner_user_id="owner", + ), + ) assert result == "converted" - assert worker_kwargs["kwargs"]["graph_engine_layers"] == ("base-layer", pause_layer) + graph_engine_layers = worker_kwargs["kwargs"]["graph_engine_layers"] + assert graph_engine_layers[0] == "base-layer" + assert isinstance(graph_engine_layers[1], PauseStatePersistenceLayer) assert worker_kwargs["kwargs"]["graph_runtime_state"] is graph_runtime_state - assert draft_saver_factory.call_args.kwargs["tenant_id"] == app_model.tenant_id + assert worker_kwargs["joined"] is True + assert worker_kwargs["join_timeout"] == 300 + assert draft_factory_tenant_ids == [app.tenant_id] -def test_resume_path_runs_worker_with_runtime_state(mocker: MockerFixture): +def test_resume_path_runs_worker_with_runtime_state(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: generator = WorkflowAppGenerator() - runtime_state = MagicMock(name="runtime-state") - - pause_layer = MagicMock(name="pause-layer") - mocker.patch("core.app.apps.workflow.app_generator.PauseStatePersistenceLayer", return_value=pause_layer) + app, workflow, end_user = _persist_generator_rows(sqlite_session) + workflow_run = WorkflowRun( + id="run", + tenant_id=workflow.tenant_id, + app_id=app.id, + workflow_id=workflow.id, + type=workflow.type, + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, + version=workflow.version, + graph=workflow.graph, + inputs="{}", + status=WorkflowExecutionStatus.RUNNING, + created_by_role=CreatorUserRole.END_USER, + created_by=end_user.id, + ) + sqlite_session.add(workflow_run) + sqlite_session.commit() + runtime_state = _runtime_state() queue_manager = MagicMock() - mocker.patch("core.app.apps.workflow.app_generator.WorkflowAppQueueManager", return_value=queue_manager) + monkeypatch.setattr( + "core.app.apps.workflow.app_generator.WorkflowAppQueueManager", + MagicMock(return_value=queue_manager), + ) - mocker.patch.object(generator, "_handle_response", return_value="raw-response") - mocker.patch( + monkeypatch.setattr(generator, "_handle_response", MagicMock(return_value="raw-response")) + monkeypatch.setattr( "core.app.apps.workflow.app_generator.WorkflowAppGenerateResponseConverter.convert", - side_effect=lambda response, invoke_from: response, + MagicMock(side_effect=lambda response, invoke_from: response), ) - fake_db = SimpleNamespace(session=MagicMock(), engine=MagicMock()) - mocker.patch("core.app.apps.workflow.app_generator.db", fake_db) - - workflow = SimpleNamespace( - id="workflow", tenant_id="tenant", app_id="app", graph_dict={}, type="workflow", version="1" + engine = sqlite_session.get_bind() + monkeypatch.setattr(app_generator_module.db, "session", Session(engine, expire_on_commit=False)) + monkeypatch.setattr( + app_generator_module.session_factory, + "create_session", + lambda: Session(engine, expire_on_commit=False), ) - end_user = SimpleNamespace(session_id="end-user-session") - app_record = SimpleNamespace(id="app") - - session = MagicMock() - session.__enter__.return_value = session - session.__exit__.return_value = False - session.scalar.side_effect = [workflow, end_user, app_record] - mocker.patch("core.app.apps.workflow.app_generator.session_factory", return_value=session) runner_instance = MagicMock() def runner_ctor(**kwargs): assert kwargs["graph_runtime_state"] is runtime_state + assert kwargs["workflow"].id == workflow.id + assert kwargs["system_user_id"] == end_user.session_id + assert kwargs["queue_manager"] is queue_manager return runner_instance - mocker.patch("core.app.apps.workflow.app_generator.WorkflowAppRunner", side_effect=runner_ctor) + monkeypatch.setattr( + "core.app.apps.workflow.app_generator.WorkflowAppRunner", + MagicMock(side_effect=runner_ctor), + ) + + worker_lifecycle: dict[str, object] = {} class ImmediateThread: def __init__(self, target, kwargs): @@ -293,43 +486,35 @@ def test_resume_path_runs_worker_with_runtime_state(mocker: MockerFixture): def start(self): return None - mocker.patch("core.app.apps.workflow.app_generator.threading.Thread", ImmediateThread) + def join(self, timeout): + worker_lifecycle["joined"] = True + worker_lifecycle["join_timeout"] = timeout - mocker.patch( - "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_execution_repository", - return_value=MagicMock(), - ) - mocker.patch( - "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", - return_value=MagicMock(), + def is_alive(self): + return False + + monkeypatch.setattr("core.app.apps.workflow.app_generator.threading.Thread", ImmediateThread) + + pause_config = PauseStateLayerConfig( + session_factory=sessionmaker(bind=engine, expire_on_commit=False), + state_owner_user_id="owner", ) - pause_config = SimpleNamespace(session_factory=MagicMock(), state_owner_user_id="owner") - - app_model = SimpleNamespace(mode="workflow", tenant_id="tenant") - app_config = SimpleNamespace(app_id="app", tenant_id="tenant", workflow_id="workflow") - application_generate_entity = SimpleNamespace( - task_id="task", - user_id="user", - invoke_from="service-api", - app_config=app_config, - files=[], - stream=True, - workflow_execution_id="run", - trace_manager=MagicMock(), - ) + application_generate_entity = _generate_entity(app, workflow, end_user) + workflow_execution_repository, workflow_node_execution_repository = _repositories(sqlite_session, app, end_user) result = generator.resume( - app_model=app_model, + app_model=app, workflow=workflow, - user=MagicMock(), + user=end_user, application_generate_entity=application_generate_entity, graph_runtime_state=runtime_state, - workflow_execution_repository=MagicMock(), - workflow_node_execution_repository=MagicMock(), + workflow_execution_repository=workflow_execution_repository, + workflow_node_execution_repository=workflow_node_execution_repository, pause_state_config=pause_config, ) assert result == "raw-response" + assert worker_lifecycle["joined"] is True + assert worker_lifecycle["join_timeout"] == 300 runner_instance.run.assert_called_once() - queue_manager.graph_runtime_state = runtime_state diff --git a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py index fd643893f69..460d2943624 100644 --- a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py +++ b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_core.py @@ -23,6 +23,7 @@ from core.app.entities.queue_entities import ( QueueWorkflowStartedEvent, QueueWorkflowSucceededEvent, ) +from core.workflow.nodes.agent.events import NodeRunAgentLogEvent from core.workflow.nodes.human_input.pause_reason import HumanInputRequired from core.workflow.system_variables import default_system_variables from graphon.entities.pause_reason import HitlRequired @@ -32,7 +33,6 @@ from graphon.graph_events import ( GraphRunPausedEvent, GraphRunStartedEvent, GraphRunSucceededEvent, - NodeRunAgentLogEvent, NodeRunExceptionEvent, NodeRunFailedEvent, NodeRunHumanInputFormFilledEvent, @@ -334,7 +334,6 @@ class TestWorkflowBasedAppRunner: variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables()), start_at=0.0, ) - graph_runtime_state.register_paused_node("node-1") workflow_entry = SimpleNamespace(graph_engine=SimpleNamespace(graph_runtime_state=graph_runtime_state)) emails: list[dict] = [] diff --git a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_notifications.py b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_notifications.py index 778c6482635..f7f9bdb4ba2 100644 --- a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_notifications.py +++ b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_notifications.py @@ -20,9 +20,6 @@ class _DummyQueueManager: class _DummyRuntimeState: variable_pool = object() - def get_paused_nodes(self): - return ["node-1"] - class _DummyGraphEngine: def __init__(self): diff --git a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_single_node.py b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_single_node.py index f02994fd61c..bdfedd82fe4 100644 --- a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_single_node.py +++ b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_single_node.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json from typing import Any from unittest.mock import MagicMock, patch @@ -12,7 +13,7 @@ from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerat from core.workflow.system_variables import default_system_variables from graphon.entities.graph_config import NodeConfigDictAdapter from graphon.runtime import GraphRuntimeState, VariablePool -from models.workflow import Workflow +from models.workflow import Workflow, WorkflowKind def _make_graph_state(): @@ -55,13 +56,14 @@ def test_run_uses_single_node_execution_branch( app_generate_entity.single_iteration_run = single_iteration_run app_generate_entity.single_loop_run = single_loop_run - workflow = MagicMock(spec=Workflow) - workflow.tenant_id = "tenant" - workflow.app_id = "app" - workflow.id = "workflow" - workflow.type = "workflow" - workflow.version = "v1" - workflow.graph_dict = {"nodes": [], "edges": []} + workflow = Workflow( + tenant_id="tenant", + app_id="app", + id="workflow", + type="workflow", + version="v1", + graph=json.dumps({"nodes": [], "edges": []}), + ) workflow.environment_variables = [] runner = WorkflowAppRunner( @@ -119,24 +121,28 @@ def test_single_node_run_validates_target_node_config(monkeypatch: pytest.Monkey app_id="app", ) - workflow = MagicMock(spec=Workflow) - workflow.id = "workflow" - workflow.tenant_id = "tenant" - workflow.graph_dict = { - "nodes": [ + workflow = Workflow( + id="workflow", + tenant_id="tenant", + graph=json.dumps( { - "id": "loop-node", - "data": { - "type": "loop", - "title": "Loop", - "loop_count": 1, - "break_conditions": [], - "logical_operator": "and", - }, + "nodes": [ + { + "id": "loop-node", + "data": { + "type": "loop", + "title": "Loop", + "loop_count": 1, + "start_node_id": "loop-start", + "break_conditions": [], + "logical_operator": "and", + }, + } + ], + "edges": [], } - ], - "edges": [], - } + ), + ) _, _, graph_runtime_state = _make_graph_state() seen_configs: list[object] = [] @@ -187,15 +193,16 @@ def test_run_adds_inputs_with_snippet_compatible_start_aliases() -> None: app_generate_entity.single_iteration_run = None app_generate_entity.single_loop_run = None - workflow = MagicMock(spec=Workflow) - workflow.tenant_id = "tenant" - workflow.app_id = "app" - workflow.id = "workflow" - workflow.type = "workflow" - workflow.version = "v1" - workflow.graph_dict = {"nodes": [], "edges": []} + workflow = Workflow( + tenant_id="tenant", + app_id="app", + id="workflow", + type="workflow", + version="v1", + graph=json.dumps({"nodes": [], "edges": []}), + kind=WorkflowKind.SNIPPET, + ) workflow.environment_variables = [] - workflow.kind_or_standard = "snippet" runner = WorkflowAppRunner( application_generate_entity=app_generate_entity, diff --git a/api/tests/unit_tests/core/app/apps/test_workflow_pause_events.py b/api/tests/unit_tests/core/app/apps/test_workflow_pause_events.py index c0052bc5bab..e57b56623c3 100644 --- a/api/tests/unit_tests/core/app/apps/test_workflow_pause_events.py +++ b/api/tests/unit_tests/core/app/apps/test_workflow_pause_events.py @@ -3,6 +3,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from core.app.apps.common import workflow_response_converter from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter @@ -24,27 +25,7 @@ from graphon.entities.pause_reason import HitlRequired from graphon.graph_events import GraphRunPausedEvent from graphon.runtime import GraphRuntimeState, VariablePool from models.account import Account -from models.human_input import RecipientType - - -class _FakeSession: - """Stub session: `execute` feeds the form-expiration query, `scalars` the recipients.""" - - def __init__(self, *, execute_rows=(), scalars_rows=()): - self._execute_rows = execute_rows - self._scalars_rows = scalars_rows - - def execute(self, _stmt): - return list(self._execute_rows) - - def scalars(self, _stmt): - return list(self._scalars_rows) - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb): - return False +from models.human_input import HumanInputForm, HumanInputFormRecipient, RecipientType class _RecordingWorkflowAppRunner(WorkflowAppRunner): @@ -59,8 +40,45 @@ class _RecordingWorkflowAppRunner(WorkflowAppRunner): class _FakeRuntimeState: variable_pool = object() - def get_paused_nodes(self): - return ["node-pause-1"] + +@pytest.fixture +def sqlite_pause_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Bind pause-response queries to the shared SQLite session's database.""" + monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=sqlite_session.get_bind())) + return sqlite_session + + +def _persist_human_input_form( + session: Session, + *, + recipients: list[tuple[RecipientType, str]] | None = None, +) -> datetime: + expiration_time = datetime(2024, 1, 1, tzinfo=UTC) + form = HumanInputForm( + id="form-1", + tenant_id="tenant-id", + app_id="app-id", + workflow_run_id="run-id", + node_id="node-id", + form_definition='{"display_in_ui": true}', + rendered_content="Rendered", + expiration_time=expiration_time, + ) + recipient_models = [ + HumanInputFormRecipient( + id=f"recipient-{index}", + form_id=form.id, + delivery_id=f"delivery-{index}", + recipient_type=recipient_type, + recipient_payload="{}", + access_token=access_token, + ) + for index, (recipient_type, access_token) in enumerate(recipients or ()) + ] + session.add(form) + session.add_all(recipient_models) + session.commit() + return expiration_time def _build_runner(): @@ -119,6 +137,7 @@ def test_graph_run_paused_event_emits_queue_pause_event(monkeypatch: pytest.Monk "core.app.apps.workflow_app_runner.enrich_graph_pause_reasons", lambda **_: [enriched_reason], ) + monkeypatch.setattr("core.app.apps.workflow_app_runner.dispatch_human_input_email_task", MagicMock()) runner._handle_event(workflow_entry, event) @@ -127,7 +146,7 @@ def test_graph_run_paused_event_emits_queue_pause_event(monkeypatch: pytest.Monk assert isinstance(queue_event, QueueWorkflowPausedEvent) assert queue_event.reasons == [enriched_reason] assert queue_event.outputs == {"foo": "bar"} - assert queue_event.paused_nodes == ["node-pause-1"] + assert queue_event.paused_nodes == ["node-human"] def _build_converter(*, invoke_from: InvokeFrom = InvokeFrom.SERVICE_API): @@ -143,10 +162,8 @@ def _build_converter(*, invoke_from: InvokeFrom = InvokeFrom.SERVICE_API): workflow_id="workflow-id", workflow_execution_id="run-id", ) - user = MagicMock(spec=Account) + user = Account(name="Tester", email="tester@example.com") user.id = "account-id" - user.name = "Tester" - user.email = "tester@example.com" return WorkflowResponseConverter( application_generate_entity=application_generate_entity, user=user, @@ -154,7 +171,12 @@ def _build_converter(*, invoke_from: InvokeFrom = InvokeFrom.SERVICE_API): ) -def test_queue_workflow_paused_event_to_stream_responses(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize( + "sqlite_session", + [(HumanInputForm, HumanInputFormRecipient)], + indirect=True, +) +def test_queue_workflow_paused_event_to_stream_responses(sqlite_pause_session: Session): converter = _build_converter() converter.workflow_start_to_stream_response( task_id="task", @@ -163,18 +185,14 @@ def test_queue_workflow_paused_event_to_stream_responses(monkeypatch: pytest.Mon reason=WorkflowStartReason.INITIAL, ) - expiration_time = datetime(2024, 1, 1, tzinfo=UTC) - session = _FakeSession( - execute_rows=[("form-1", expiration_time, '{"display_in_ui": true}')], - scalars_rows=[ - SimpleNamespace(form_id="form-1", recipient_type=RecipientType.CONSOLE, access_token="console-token"), - SimpleNamespace(form_id="form-1", recipient_type=RecipientType.BACKSTAGE, access_token="backstage-token"), + expiration_time = _persist_human_input_form( + sqlite_pause_session, + recipients=[ + (RecipientType.CONSOLE, "console-token"), + (RecipientType.BACKSTAGE, "backstage-token"), ], ) - monkeypatch.setattr(workflow_response_converter, "Session", lambda **_: session) - monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=object())) - reason = HumanInputRequired( form_id="form-1", form_content="Rendered", @@ -216,8 +234,11 @@ def test_queue_workflow_paused_event_to_stream_responses(monkeypatch: pytest.Mon assert hi_resp.data.expiration_time == int(expiration_time.timestamp()) -def _build_paused_human_input_response(monkeypatch, recipients): - """Drive the live OPENAPI pause path with the given recipients via a fake session.""" +def _build_paused_human_input_response( + session: Session, + recipients: list[tuple[RecipientType, str]], +): + """Drive the live OPENAPI pause path with persisted forms and recipients.""" converter = _build_converter(invoke_from=InvokeFrom.OPENAPI) converter.workflow_start_to_stream_response( task_id="task", @@ -226,14 +247,7 @@ def _build_paused_human_input_response(monkeypatch, recipients): reason=WorkflowStartReason.INITIAL, ) - expiration_time = datetime(2024, 1, 1, tzinfo=UTC) - session = _FakeSession( - execute_rows=[("form-1", expiration_time, '{"display_in_ui": true}')], - scalars_rows=list(recipients), - ) - - monkeypatch.setattr(workflow_response_converter, "Session", lambda **_: session) - monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=object())) + _persist_human_input_form(session, recipients=recipients) reason = HumanInputRequired( form_id="form-1", @@ -259,12 +273,17 @@ def _build_paused_human_input_response(monkeypatch, recipients): return responses -def test_openapi_pause_without_web_app_recipient_emits_approval_channels(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize( + "sqlite_session", + [(HumanInputForm, HumanInputFormRecipient)], + indirect=True, +) +def test_openapi_pause_without_web_app_recipient_emits_approval_channels(sqlite_pause_session: Session): responses = _build_paused_human_input_response( - monkeypatch, + sqlite_pause_session, recipients=[ - SimpleNamespace(form_id="form-1", recipient_type=RecipientType.EMAIL_MEMBER, access_token="email-token"), - SimpleNamespace(form_id="form-1", recipient_type=RecipientType.BACKSTAGE, access_token="backstage-token"), + (RecipientType.EMAIL_MEMBER, "email-token"), + (RecipientType.BACKSTAGE, "backstage-token"), ], ) @@ -276,16 +295,17 @@ def test_openapi_pause_without_web_app_recipient_emits_approval_channels(monkeyp assert pause_resp.data.reasons[0]["approval_channels"] == ["console", "email"] -def test_openapi_pause_with_web_app_recipient_sets_token_and_channels(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize( + "sqlite_session", + [(HumanInputForm, HumanInputFormRecipient)], + indirect=True, +) +def test_openapi_pause_with_web_app_recipient_sets_token_and_channels(sqlite_pause_session: Session): responses = _build_paused_human_input_response( - monkeypatch, + sqlite_pause_session, recipients=[ - SimpleNamespace( - form_id="form-1", - recipient_type=RecipientType.STANDALONE_WEB_APP, - access_token="web-app-token", - ), - SimpleNamespace(form_id="form-1", recipient_type=RecipientType.BACKSTAGE, access_token="backstage-token"), + (RecipientType.STANDALONE_WEB_APP, "web-app-token"), + (RecipientType.BACKSTAGE, "backstage-token"), ], ) @@ -297,7 +317,12 @@ def test_openapi_pause_with_web_app_recipient_sets_token_and_channels(monkeypatc assert pause_resp.data.reasons[0]["approval_channels"] == ["console"] -def test_queue_workflow_paused_event_resolves_variable_select_options(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize( + "sqlite_session", + [(HumanInputForm, HumanInputFormRecipient)], + indirect=True, +) +def test_queue_workflow_paused_event_resolves_variable_select_options(sqlite_pause_session: Session): converter = _build_converter() converter.workflow_start_to_stream_response( task_id="task", @@ -306,11 +331,7 @@ def test_queue_workflow_paused_event_resolves_variable_select_options(monkeypatc reason=WorkflowStartReason.INITIAL, ) - expiration_time = datetime(2024, 1, 1, tzinfo=UTC) - session = _FakeSession(execute_rows=[("form-1", expiration_time, '{"display_in_ui": true}')]) - - monkeypatch.setattr(workflow_response_converter, "Session", lambda **_: session) - monkeypatch.setattr(workflow_response_converter, "db", SimpleNamespace(engine=object())) + _persist_human_input_form(sqlite_pause_session) reason = HumanInputRequired( form_id="form-1", diff --git a/api/tests/unit_tests/core/app/apps/workflow/test_active_workflow_tasks.py b/api/tests/unit_tests/core/app/apps/workflow/test_active_workflow_tasks.py index c50b16533ff..769dd0c092d 100644 --- a/api/tests/unit_tests/core/app/apps/workflow/test_active_workflow_tasks.py +++ b/api/tests/unit_tests/core/app/apps/workflow/test_active_workflow_tasks.py @@ -1,5 +1,9 @@ +import threading +from collections.abc import Generator + import pytest +from core.app.apps.base_app_generator import BaseAppGenerator from core.app.apps.workflow.active_workflow_tasks import ( active_workflow_task, get_active_workflow_task_count, @@ -28,3 +32,51 @@ def test_active_workflow_task_rejects_duplicate_task_id() -> None: with pytest.raises(ValueError, match="already active"): with active_workflow_task("task-a"): pass + + +def test_managed_stream_waits_for_active_worker_cleanup() -> None: + worker_started = threading.Event() + release_worker = threading.Event() + stream_exhausted = threading.Event() + consumer_finished = threading.Event() + consumer_errors: list[BaseException] = [] + + def run_worker() -> None: + with active_workflow_task("task-a"): + worker_started.set() + release_worker.wait() + + def response_stream() -> Generator[dict[str, str], None, None]: + yield {"event": "workflow_finished"} + stream_exhausted.set() + + worker_thread = threading.Thread(target=run_worker) + worker_thread.start() + assert worker_started.wait(timeout=2) + + managed_stream = BaseAppGenerator._wrap_stream_with_worker_thread_join(response_stream(), worker_thread) + assert next(managed_stream) == {"event": "workflow_finished"} + + def finish_stream() -> None: + try: + list(managed_stream) + except BaseException as exc: + consumer_errors.append(exc) + finally: + consumer_finished.set() + + consumer_thread = threading.Thread(target=finish_stream) + consumer_thread.start() + try: + assert stream_exhausted.wait(timeout=2) + assert not consumer_finished.is_set() + assert get_active_workflow_task_count() == 1 + finally: + release_worker.set() + consumer_thread.join(timeout=2) + worker_thread.join(timeout=2) + + assert not consumer_thread.is_alive() + assert not worker_thread.is_alive() + assert consumer_errors == [] + assert get_active_workflow_task_count() == 0 diff --git a/api/tests/unit_tests/core/app/apps/workflow/test_app_generator_extra.py b/api/tests/unit_tests/core/app/apps/workflow/test_app_generator_extra.py index 3509c349aef..3a219be9b14 100644 --- a/api/tests/unit_tests/core/app/apps/workflow/test_app_generator_extra.py +++ b/api/tests/unit_tests/core/app/apps/workflow/test_app_generator_extra.py @@ -1,50 +1,242 @@ from __future__ import annotations import contextlib +import json +from collections.abc import Iterator from types import SimpleNamespace from unittest.mock import Mock import pytest +from sqlalchemy import inspect +from sqlalchemy.orm import Session, scoped_session, sessionmaker from core.app.app_config.entities import AppAdditionalFeatures, WorkflowUIBasedAppConfig from core.app.apps.exc import GenerateTaskStoppedError +from core.app.apps.workflow import app_generator as app_generator_module from core.app.apps.workflow.app_generator import SKIP_PREPARE_USER_INPUTS_KEY, WorkflowAppGenerator from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity from core.ops.ops_trace_manager import TraceQueueManager -from models.model import AppMode +from models.enums import EndUserType +from models.model import App, AppMode, EndUser +from models.snippet import CustomizedSnippet +from models.workflow import Workflow, WorkflowKind, WorkflowType + +TENANT_ID = "00000000-0000-0000-0000-000000000001" +OTHER_TENANT_ID = "00000000-0000-0000-0000-000000000002" +APP_ID = "00000000-0000-0000-0000-000000000003" +WORKFLOW_ID = "00000000-0000-0000-0000-000000000004" +END_USER_ID = "00000000-0000-0000-0000-000000000005" +CREATOR_ID = "00000000-0000-0000-0000-000000000006" + + +@pytest.fixture +def sqlite_generator_scoped_session( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> Iterator[scoped_session[Session]]: + """Adapt the shared SQLite session factory to Flask-SQLAlchemy's scoped session API.""" + engine = sqlite_session.get_bind() + request_sessions = scoped_session(sqlite_session_factory) + monkeypatch.setattr( + app_generator_module, + "db", + SimpleNamespace(engine=engine, session=request_sessions), + ) + try: + yield request_sessions + finally: + request_sessions.remove() + + +def _persist_app(session: Session) -> App: + app = App( + id=APP_ID, + tenant_id=TENANT_ID, + name="Workflow app", + description="", + mode=AppMode.WORKFLOW, + icon_type=None, + icon="", + icon_background=None, + app_model_config_id=None, + workflow_id=WORKFLOW_ID, + enable_site=False, + enable_api=True, + max_active_requests=None, + created_by=CREATOR_ID, + ) + session.add(app) + session.commit() + return app + + +def _persist_workflow( + session: Session, + *, + workflow_id: str = WORKFLOW_ID, + app_id: str = APP_ID, + tenant_id: str = TENANT_ID, + kind: WorkflowKind = WorkflowKind.STANDARD, +) -> Workflow: + workflow = Workflow.new( + tenant_id=tenant_id, + app_id=app_id, + type=WorkflowType.WORKFLOW.value, + version="1", + graph=json.dumps({"nodes": [], "edges": []}), + features="{}", + created_by=CREATOR_ID, + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + kind=kind.value, + ) + workflow.id = workflow_id + session.add(workflow) + session.commit() + return workflow + + +def _persist_end_user(session: Session) -> EndUser: + end_user = EndUser( + id=END_USER_ID, + tenant_id=TENANT_ID, + app_id=APP_ID, + type=EndUserType.BROWSER, + name="End user", + session_id="session-id", + ) + session.add(end_user) + session.commit() + return end_user + + +def _persist_snippet( + session: Session, + *, + snippet_id: str, + tenant_id: str = TENANT_ID, +) -> CustomizedSnippet: + snippet = CustomizedSnippet( + id=snippet_id, + tenant_id=tenant_id, + name="Snippet", + description=None, + type="node", + ) + session.add(snippet) + session.commit() + return snippet class TestWorkflowAppGeneratorValidation: - def test_ensure_snippet_start_node_returns_original_for_non_snippet_workflow(self): + @pytest.mark.usefixtures("sqlite_generator_scoped_session") + def test_generate_stream_joins_worker_after_response_exhaustion(self, monkeypatch: pytest.MonkeyPatch): + generator = WorkflowAppGenerator() + worker_thread = Mock() + worker_thread.is_alive.return_value = False + app_config = WorkflowUIBasedAppConfig( + tenant_id="tenant", + app_id="app", + app_mode=AppMode.WORKFLOW, + additional_features=AppAdditionalFeatures(), + variables=[], + workflow_id="workflow-id", + ) + application_generate_entity = WorkflowAppGenerateEntity.model_construct( + task_id="task", + app_config=app_config, + inputs={}, + files=[], + user_id="user", + stream=True, + invoke_from=InvokeFrom.WEB_APP, + extras={}, + ) + + def response_stream(): + yield {"event": "workflow_finished"} + + monkeypatch.setattr(generator, "_bind_file_access_scope", lambda **kwargs: contextlib.nullcontext()) + monkeypatch.setattr( + "core.app.apps.workflow.app_generator.WorkflowAppQueueManager", + lambda **kwargs: SimpleNamespace(**kwargs), + ) + monkeypatch.setattr( + "core.app.apps.workflow.app_generator.current_app", + SimpleNamespace(_get_current_object=lambda: SimpleNamespace(name="flask")), + ) + monkeypatch.setattr("core.app.apps.workflow.app_generator.contextvars.copy_context", lambda: "ctx") + monkeypatch.setattr("core.app.apps.workflow.app_generator.threading.Thread", lambda **kwargs: worker_thread) + monkeypatch.setattr(generator, "_get_draft_var_saver_factory", lambda *args, **kwargs: "draft-factory") + monkeypatch.setattr(generator, "_handle_response", lambda **kwargs: response_stream()) + monkeypatch.setattr( + "core.app.apps.workflow.app_generator.WorkflowAppGenerateResponseConverter.convert", + lambda response, invoke_from: response, + ) + + managed_stream = generator._generate( + app_model=SimpleNamespace(mode=AppMode.WORKFLOW, tenant_id="tenant"), + workflow=SimpleNamespace(id="workflow-id"), + user=SimpleNamespace(id="user"), + application_generate_entity=application_generate_entity, + invoke_from=InvokeFrom.WEB_APP, + workflow_execution_repository=SimpleNamespace(), + workflow_node_execution_repository=SimpleNamespace(), + streaming=True, + ) + + worker_thread.start.assert_called_once_with() + worker_thread.join.assert_not_called() + assert list(managed_stream) == [{"event": "workflow_finished"}] + worker_thread.join.assert_called_once_with(timeout=300) + + def test_ensure_snippet_start_node_returns_original_for_non_snippet_workflow( + self, + unbound_session: Session, + ): workflow = SimpleNamespace(kind_or_standard="workflow") - session = SimpleNamespace(scalar=Mock()) - result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=session, workflow=workflow) + result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker( + session=unbound_session, + workflow=workflow, + ) assert result is workflow - session.scalar.assert_not_called() - def test_ensure_snippet_start_node_returns_original_when_snippet_missing(self): - workflow = SimpleNamespace(kind_or_standard="snippet", app_id="snippet-1", tenant_id="tenant-1") - session = SimpleNamespace(scalar=Mock(return_value=None)) + def test_ensure_snippet_start_node_returns_original_when_snippet_is_from_another_tenant( + self, + sqlite_session: Session, + ): + workflow = _persist_workflow(sqlite_session, kind=WorkflowKind.SNIPPET) + _persist_snippet(sqlite_session, snippet_id=APP_ID, tenant_id=OTHER_TENANT_ID) - result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=session, workflow=workflow) + result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker( + session=sqlite_session, + workflow=workflow, + ) assert result is workflow - session.scalar.assert_called_once() - def test_ensure_snippet_start_node_delegates_when_snippet_exists(self, monkeypatch: pytest.MonkeyPatch): - workflow = SimpleNamespace(kind_or_standard="snippet", app_id="snippet-1", tenant_id="tenant-1") - snippet = SimpleNamespace(id="snippet-1") + def test_ensure_snippet_start_node_delegates_when_snippet_exists( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ): + workflow = _persist_workflow(sqlite_session, kind=WorkflowKind.SNIPPET) + snippet = _persist_snippet(sqlite_session, snippet_id=APP_ID) injected_workflow = SimpleNamespace(id="workflow-injected") - session = SimpleNamespace(scalar=Mock(return_value=snippet)) ensure_start_node = Mock(return_value=injected_workflow) monkeypatch.setattr( "services.snippet_generate_service.SnippetGenerateService.ensure_start_node_for_worker", ensure_start_node, ) - result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(session=session, workflow=workflow) + result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker( + session=sqlite_session, + workflow=workflow, + ) assert result is injected_workflow ensure_start_node.assert_called_once_with(workflow, snippet) @@ -55,7 +247,7 @@ class TestWorkflowAppGeneratorValidation: assert generator._should_prepare_user_inputs({}) is True assert generator._should_prepare_user_inputs({SKIP_PREPARE_USER_INPUTS_KEY: True}) is False - def test_single_iteration_generate_validates_args(self): + def test_single_iteration_generate_validates_args(self, sqlite_session: Session): generator = WorkflowAppGenerator() with pytest.raises(ValueError, match="node_id is required"): @@ -66,7 +258,7 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(), args={"inputs": {}}, streaming=False, - session=Mock(), + session=sqlite_session, ) with pytest.raises(ValueError, match="inputs is required"): @@ -77,10 +269,10 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(), args={}, streaming=False, - session=Mock(), + session=sqlite_session, ) - def test_single_loop_generate_validates_args(self): + def test_single_loop_generate_validates_args(self, sqlite_session: Session): generator = WorkflowAppGenerator() with pytest.raises(ValueError, match="node_id is required"): @@ -91,20 +283,30 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(), args=SimpleNamespace(inputs={}), streaming=False, - session=Mock(), + session=sqlite_session, ) - def test_single_iteration_generate_includes_trace_session_id_in_extras(self, monkeypatch: pytest.MonkeyPatch): + @pytest.mark.usefixtures("sqlite_generator_scoped_session") + def test_single_iteration_generate_includes_trace_session_id_in_extras( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ): generator = WorkflowAppGenerator() + app = _persist_app(sqlite_session) + workflow = _persist_workflow(sqlite_session) + user = _persist_end_user(sqlite_session) app_config = WorkflowUIBasedAppConfig( - tenant_id="tenant", - app_id="app", + tenant_id=TENANT_ID, + app_id=APP_ID, app_mode=AppMode.WORKFLOW, additional_features=AppAdditionalFeatures(), variables=[], - workflow_id="workflow-id", + workflow_id=WORKFLOW_ID, ) captured: dict[str, object] = {} + repository_session_makers: list[sessionmaker[Session]] = [] + draft_sessions: list[Session] = [] monkeypatch.setattr( "core.app.apps.workflow.app_generator.WorkflowAppConfigManager.get_app_config", @@ -112,47 +314,59 @@ class TestWorkflowAppGeneratorValidation: ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_execution_repository", - lambda **kwargs: SimpleNamespace(), + lambda **kwargs: repository_session_makers.append(kwargs["session_factory"]) or SimpleNamespace(), ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", - lambda **kwargs: SimpleNamespace(), + lambda **kwargs: repository_session_makers.append(kwargs["session_factory"]) or SimpleNamespace(), ) monkeypatch.setattr("core.app.apps.workflow.app_generator.DraftVarLoader", lambda **kwargs: SimpleNamespace()) - monkeypatch.setattr("core.app.apps.workflow.app_generator.sessionmaker", lambda **kwargs: SimpleNamespace()) - monkeypatch.setattr( - "core.app.apps.workflow.app_generator.db", - SimpleNamespace(engine=object(), session=lambda: SimpleNamespace()), - ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.WorkflowDraftVariableService", - lambda session: SimpleNamespace(prefill_conversation_variable_default_values=lambda *args, **kwargs: None), + lambda session: ( + draft_sessions.append(session) + or SimpleNamespace(prefill_conversation_variable_default_values=lambda *args, **kwargs: None) + ), ) monkeypatch.setattr(generator, "_generate", lambda **kwargs: captured.update(kwargs) or {"ok": True}) generator.single_iteration_generate( - app_model=SimpleNamespace(id="app", tenant_id="tenant"), - workflow=SimpleNamespace(id="workflow-id"), + app_model=app, + workflow=workflow, node_id="node-1", - user=SimpleNamespace(id="user-id"), + user=user, args={"inputs": {"foo": "bar"}, "trace_session_id": "session-1"}, streaming=False, - session=Mock(), + session=sqlite_session, ) assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1" + assert len(repository_session_makers) == 2 + assert all(factory.kw["bind"] is sqlite_session.get_bind() for factory in repository_session_makers) + assert len(draft_sessions) == 1 + assert draft_sessions[0] is sqlite_session - def test_single_loop_generate_includes_trace_session_id_in_extras(self, monkeypatch: pytest.MonkeyPatch): + @pytest.mark.usefixtures("sqlite_generator_scoped_session") + def test_single_loop_generate_includes_trace_session_id_in_extras( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ): generator = WorkflowAppGenerator() + app = _persist_app(sqlite_session) + workflow = _persist_workflow(sqlite_session) + user = _persist_end_user(sqlite_session) app_config = WorkflowUIBasedAppConfig( - tenant_id="tenant", - app_id="app", + tenant_id=TENANT_ID, + app_id=APP_ID, app_mode=AppMode.WORKFLOW, additional_features=AppAdditionalFeatures(), variables=[], - workflow_id="workflow-id", + workflow_id=WORKFLOW_ID, ) captured: dict[str, object] = {} + repository_session_makers: list[sessionmaker[Session]] = [] + draft_sessions: list[Session] = [] monkeypatch.setattr( "core.app.apps.workflow.app_generator.WorkflowAppConfigManager.get_app_config", @@ -160,35 +374,37 @@ class TestWorkflowAppGeneratorValidation: ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_execution_repository", - lambda **kwargs: SimpleNamespace(), + lambda **kwargs: repository_session_makers.append(kwargs["session_factory"]) or SimpleNamespace(), ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", - lambda **kwargs: SimpleNamespace(), + lambda **kwargs: repository_session_makers.append(kwargs["session_factory"]) or SimpleNamespace(), ) monkeypatch.setattr("core.app.apps.workflow.app_generator.DraftVarLoader", lambda **kwargs: SimpleNamespace()) - monkeypatch.setattr("core.app.apps.workflow.app_generator.sessionmaker", lambda **kwargs: SimpleNamespace()) - monkeypatch.setattr( - "core.app.apps.workflow.app_generator.db", - SimpleNamespace(engine=object(), session=lambda: SimpleNamespace()), - ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.WorkflowDraftVariableService", - lambda session: SimpleNamespace(prefill_conversation_variable_default_values=lambda *args, **kwargs: None), + lambda session: ( + draft_sessions.append(session) + or SimpleNamespace(prefill_conversation_variable_default_values=lambda *args, **kwargs: None) + ), ) monkeypatch.setattr(generator, "_generate", lambda **kwargs: captured.update(kwargs) or {"ok": True}) generator.single_loop_generate( - app_model=SimpleNamespace(id="app", tenant_id="tenant"), - workflow=SimpleNamespace(id="workflow-id"), + app_model=app, + workflow=workflow, node_id="node-2", - user=SimpleNamespace(id="user-id"), + user=user, args=SimpleNamespace(inputs={"foo": "bar"}, trace_session_id="session-1"), streaming=False, - session=Mock(), + session=sqlite_session, ) assert captured["application_generate_entity"].extras["trace_session_id"] == "session-1" + assert len(repository_session_makers) == 2 + assert all(factory.kw["bind"] is sqlite_session.get_bind() for factory in repository_session_makers) + assert len(draft_sessions) == 1 + assert draft_sessions[0] is sqlite_session with pytest.raises(ValueError, match="inputs is required"): generator.single_loop_generate( @@ -198,7 +414,7 @@ class TestWorkflowAppGeneratorValidation: user=SimpleNamespace(), args=SimpleNamespace(inputs=None), streaming=False, - session=Mock(), + session=sqlite_session, ) @@ -252,17 +468,26 @@ class TestWorkflowAppGeneratorHandleResponse: class TestWorkflowAppGeneratorGenerate: - def test_generate_skips_prepare_inputs_when_flag_set(self, monkeypatch: pytest.MonkeyPatch): + @pytest.mark.usefixtures("sqlite_generator_scoped_session") + def test_generate_skips_prepare_inputs_when_flag_set( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ): generator = WorkflowAppGenerator() + app = _persist_app(sqlite_session) + workflow = _persist_workflow(sqlite_session) + user = _persist_end_user(sqlite_session) app_config = WorkflowUIBasedAppConfig( - tenant_id="tenant", - app_id="app", + tenant_id=TENANT_ID, + app_id=APP_ID, app_mode=AppMode.WORKFLOW, additional_features=AppAdditionalFeatures(), variables=[], - workflow_id="workflow-id", + workflow_id=WORKFLOW_ID, ) + repository_session_makers: list[sessionmaker[Session]] = [] monkeypatch.setattr( "core.app.apps.workflow.app_generator.WorkflowAppConfigManager.get_app_config", @@ -291,19 +516,11 @@ class TestWorkflowAppGeneratorGenerate: ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_execution_repository", - lambda **kwargs: SimpleNamespace(), + lambda **kwargs: repository_session_makers.append(kwargs["session_factory"]) or SimpleNamespace(), ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", - lambda **kwargs: SimpleNamespace(), - ) - monkeypatch.setattr( - "core.app.apps.workflow.app_generator.db", - SimpleNamespace(engine=object(), session=SimpleNamespace(close=lambda: None)), - ) - monkeypatch.setattr( - "core.app.apps.workflow.app_generator.sessionmaker", - lambda **kwargs: SimpleNamespace(), + lambda **kwargs: repository_session_makers.append(kwargs["session_factory"]) or SimpleNamespace(), ) prepare_inputs = pytest.fail @@ -312,9 +529,9 @@ class TestWorkflowAppGeneratorGenerate: monkeypatch.setattr(generator, "_generate", lambda **kwargs: {"ok": True}) result = generator.generate( - app_model=SimpleNamespace(id="app", tenant_id="tenant"), - workflow=SimpleNamespace(features_dict={}), - user=SimpleNamespace(id="user", session_id="session"), + app_model=app, + workflow=workflow, + user=user, args={"inputs": {}, SKIP_PREPARE_USER_INPUTS_KEY: True}, invoke_from=InvokeFrom.WEB_APP, streaming=False, @@ -322,6 +539,8 @@ class TestWorkflowAppGeneratorGenerate: ) assert result == {"ok": True} + assert len(repository_session_makers) == 2 + assert all(factory.kw["bind"] is sqlite_session.get_bind() for factory in repository_session_makers) class TestWorkflowAppGeneratorResume: @@ -436,26 +655,16 @@ class TestWorkflowAppGeneratorResume: class TestWorkflowAppGeneratorWorker: - def test_generate_worker_uses_end_user_session_for_external_invocation(self, monkeypatch: pytest.MonkeyPatch): + def test_generate_worker_uses_end_user_session_for_external_invocation( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + ): generator = WorkflowAppGenerator() - - workflow = SimpleNamespace( - id="workflow-id", - tenant_id="tenant", - app_id="app", - graph_dict={}, - type="workflow", - version="1", - ) - end_user = SimpleNamespace(id="end-user-id", session_id="session-id") - session = SimpleNamespace(scalar=Mock(side_effect=[workflow, end_user])) - - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False + _persist_app(sqlite_session) + _persist_workflow(sqlite_session) + _persist_end_user(sqlite_session) runner_kwargs = {} @@ -472,28 +681,26 @@ class TestWorkflowAppGeneratorWorker: ) monkeypatch.setattr( "core.app.apps.workflow.app_generator.session_factory.create_session", - lambda: _SessionContext(), - ) - monkeypatch.setattr( - "core.app.apps.workflow.app_generator.WorkflowAppGenerator._ensure_snippet_start_node_in_worker", - lambda self, *, session, workflow: workflow, + sqlite_session_factory, ) monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppRunner", _Runner) + restore_workflow_run_graph = Mock() + monkeypatch.setattr(generator, "_restore_workflow_run_graph", restore_workflow_run_graph) app_config = WorkflowUIBasedAppConfig( - tenant_id="tenant", - app_id="app", + tenant_id=TENANT_ID, + app_id=APP_ID, app_mode=AppMode.WORKFLOW, additional_features=AppAdditionalFeatures(), variables=[], - workflow_id="workflow-id", + workflow_id=WORKFLOW_ID, ) application_generate_entity = WorkflowAppGenerateEntity.model_construct( task_id="task", app_config=app_config, inputs={}, files=[], - user_id="end-user-id", + user_id=END_USER_ID, stream=False, invoke_from=InvokeFrom.WEB_APP, extras={}, @@ -510,6 +717,13 @@ class TestWorkflowAppGeneratorWorker: variable_loader=SimpleNamespace(), workflow_execution_repository=SimpleNamespace(), workflow_node_execution_repository=SimpleNamespace(), + graph_runtime_state=SimpleNamespace(), ) assert runner_kwargs["system_user_id"] == "session-id" + restore_workflow_run_graph.assert_called_once() + restore_kwargs = restore_workflow_run_graph.call_args.kwargs + assert isinstance(restore_kwargs["session"], Session) + assert restore_kwargs["workflow"] is runner_kwargs["workflow"] + assert restore_kwargs["workflow_run_id"] == "run-id" + assert inspect(runner_kwargs["workflow"]).detached is True diff --git a/api/tests/unit_tests/core/app/apps/workflow/test_app_queue_manager.py b/api/tests/unit_tests/core/app/apps/workflow/test_app_queue_manager.py index 5e9a69738c4..c33860b3c14 100644 --- a/api/tests/unit_tests/core/app/apps/workflow/test_app_queue_manager.py +++ b/api/tests/unit_tests/core/app/apps/workflow/test_app_queue_manager.py @@ -1,11 +1,16 @@ from __future__ import annotations -from unittest.mock import patch +from unittest.mock import Mock, patch from core.app.apps.base_app_queue_manager import PublishFrom from core.app.apps.workflow.app_queue_manager import WorkflowAppQueueManager from core.app.entities.app_invoke_entities import InvokeFrom -from core.app.entities.queue_entities import QueueMessageEndEvent, QueuePingEvent, QueueStopEvent +from core.app.entities.queue_entities import ( + QueueMessageEndEvent, + QueuePingEvent, + QueueStopEvent, + QueueWorkflowPausedEvent, +) class TestWorkflowAppQueueManager: @@ -36,6 +41,19 @@ class TestWorkflowAppQueueManager: manager._publish(QueuePingEvent(), PublishFrom.TASK_PIPELINE) + def test_publish_pause_event_stops_listener_without_aborting_execution(self): + manager = WorkflowAppQueueManager( + task_id="task", + user_id="user", + invoke_from=InvokeFrom.DEBUGGER, + app_mode="workflow", + ) + manager.stop_listen = Mock() + + manager._publish(QueueWorkflowPausedEvent(), PublishFrom.APPLICATION_MANAGER) + + manager.stop_listen.assert_called_once_with(execution_terminal=True) + def test_listener_close_aborts_unfinished_execution(self): with ( patch("core.app.apps.base_app_queue_manager.redis_client") as redis_client, @@ -99,3 +117,23 @@ class TestWorkflowAppQueueManager: _ = list(manager.listen()) graph_engine_manager.return_value.send_stop_command.assert_not_called() + + def test_workflow_pause_does_not_abort_execution(self): + with ( + patch("core.app.apps.base_app_queue_manager.redis_client") as redis_client, + patch("core.app.apps.base_app_queue_manager.GraphEngineManager") as graph_engine_manager, + ): + redis_client.get.return_value = None + manager = WorkflowAppQueueManager( + task_id="task", + user_id="user", + invoke_from=InvokeFrom.DEBUGGER, + app_mode="workflow", + ) + manager.publish(QueueWorkflowPausedEvent(), PublishFrom.APPLICATION_MANAGER) + listener = manager.listen() + + assert isinstance(next(listener).event, QueueWorkflowPausedEvent) + listener.close() + + graph_engine_manager.return_value.send_stop_command.assert_not_called() diff --git a/api/tests/unit_tests/core/app/apps/workflow/test_generate_task_pipeline_core.py b/api/tests/unit_tests/core/app/apps/workflow/test_generate_task_pipeline_core.py index 04fe7a2ebed..884b7a7a8aa 100644 --- a/api/tests/unit_tests/core/app/apps/workflow/test_generate_task_pipeline_core.py +++ b/api/tests/unit_tests/core/app/apps/workflow/test_generate_task_pipeline_core.py @@ -1,11 +1,11 @@ from __future__ import annotations import logging -from contextlib import contextmanager from types import SimpleNamespace -from unittest.mock import MagicMock import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session from core.app.app_config.entities import AppAdditionalFeatures, WorkflowUIBasedAppConfig from core.app.apps.workflow.generate_task_pipeline import WorkflowAppGenerateTaskPipeline @@ -50,10 +50,12 @@ from core.app.entities.task_entities import ( from core.base.tts.app_generator_tts_publisher import AudioTrunk from core.workflow.system_variables import build_system_variables, system_variables_to_mapping from graphon.enums import BuiltinNodeTypes, WorkflowExecutionStatus +from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.runtime import GraphRuntimeState, VariablePool from libs.datetime_utils import naive_utc_now from models.enums import CreatorUserRole from models.model import AppMode, EndUser +from models.workflow import WorkflowAppLog from tests.workflow_test_utils import build_test_variable_pool @@ -102,7 +104,7 @@ class TestWorkflowGenerateTaskPipeline: variables=build_system_variables(workflow_execution_id="run-id"), ), start_at=0.0, - total_tokens=5, + llm_usage=LLMUsage.empty_usage().model_copy(update={"total_tokens": 5}), node_run_steps=2, ) @@ -193,7 +195,7 @@ class TestWorkflowGenerateTaskPipeline: assert isinstance(responses[0], ValueError) - def test_handle_workflow_started_event_sets_run_id(self, monkeypatch: pytest.MonkeyPatch): + def test_handle_workflow_started_event_sets_run_id(self, monkeypatch: pytest.MonkeyPatch, sqlite_engine): pipeline = _make_pipeline() pipeline._graph_runtime_state = GraphRuntimeState( variable_pool=build_test_variable_pool(variables=build_system_variables(workflow_execution_id="run-id")), @@ -201,11 +203,10 @@ class TestWorkflowGenerateTaskPipeline: ) pipeline._workflow_response_converter.workflow_start_to_stream_response = lambda **kwargs: "started" - @contextmanager - def _fake_session(): - yield SimpleNamespace() - - monkeypatch.setattr(pipeline, "_database_session", _fake_session) + monkeypatch.setattr( + "core.app.apps.workflow.generate_task_pipeline.db", + SimpleNamespace(engine=sqlite_engine), + ) monkeypatch.setattr(pipeline, "_save_workflow_app_log", lambda **kwargs: None) responses = list(pipeline._handle_workflow_started_event(QueueWorkflowStartedEvent())) @@ -339,19 +340,18 @@ class TestWorkflowGenerateTaskPipeline: assert responses == ["finish"] - def test_save_workflow_app_log_created_from(self): + @pytest.mark.parametrize("sqlite_session", [(WorkflowAppLog,)], indirect=True) + def test_save_workflow_app_log_created_from(self, sqlite_session: Session): pipeline = _make_pipeline() pipeline._application_generate_entity.invoke_from = InvokeFrom.SERVICE_API pipeline._user_id = "user" - added: list[object] = [] + pipeline._save_workflow_app_log(session=sqlite_session, workflow_run_id="run-id") + sqlite_session.flush() - class _Session: - def add(self, item): - added.append(item) - - pipeline._save_workflow_app_log(session=_Session(), workflow_run_id="run-id") - - assert added + saved_log = sqlite_session.scalar(select(WorkflowAppLog)) + assert saved_log is not None + assert saved_log.workflow_run_id == "run-id" + assert saved_log.created_from == "service-api" def test_iteration_loop_and_human_input_handlers(self): pipeline = _make_pipeline() @@ -674,35 +674,29 @@ class TestWorkflowGenerateTaskPipeline: assert "Fails to get audio trunk, task_id: task" in caplog.messages assert any(isinstance(item, MessageAudioEndStreamResponse) for item in responses) - def test_database_session_rolls_back_on_error(self, monkeypatch: pytest.MonkeyPatch): + @pytest.mark.parametrize("sqlite_session", [(WorkflowAppLog,)], indirect=True) + def test_database_session_rolls_back_on_error( + self, monkeypatch: pytest.MonkeyPatch, sqlite_engine, sqlite_session: Session + ): pipeline = _make_pipeline() - calls = {"enter": 0, "exit_exc": None} + pipeline._application_generate_entity.invoke_from = InvokeFrom.SERVICE_API + pipeline._user_id = "user" + monkeypatch.setattr( + "core.app.apps.workflow.generate_task_pipeline.db", + SimpleNamespace(engine=sqlite_engine), + ) - class _BeginContext: - def __enter__(self): - calls["enter"] += 1 - return MagicMock() - - def __exit__(self, exc_type, exc, tb): - calls["exit_exc"] = exc_type - return False - - class _Sessionmaker: - def __init__(self, *args, **kwargs): - pass - - def begin(self): - return _BeginContext() - - monkeypatch.setattr("core.app.apps.workflow.generate_task_pipeline.sessionmaker", _Sessionmaker) - monkeypatch.setattr("core.app.apps.workflow.generate_task_pipeline.db", SimpleNamespace(engine=object())) - - with pytest.raises(RuntimeError, match="db error"): - with pipeline._database_session(): + def persist_then_fail() -> None: + with pipeline._database_session() as session: + pipeline._save_workflow_app_log(session=session, workflow_run_id="run-id") + session.flush() raise RuntimeError("db error") - assert calls["enter"] == 1 - assert calls["exit_exc"] is RuntimeError + with pytest.raises(RuntimeError, match="db error"): + persist_then_fail() + + sqlite_session.expire_all() + assert sqlite_session.scalar(select(WorkflowAppLog)) is None def test_node_retry_and_started_handlers_cover_none_and_value(self): pipeline = _make_pipeline() @@ -862,31 +856,30 @@ class TestWorkflowGenerateTaskPipeline: pipeline._handle_workflow_failed_and_stop_events = lambda event, **kwargs: iter(["stopped"]) assert list(pipeline._process_stream_response()) == ["stopped"] - def test_save_workflow_app_log_covers_invoke_from_variants(self): + @pytest.mark.parametrize("sqlite_session", [(WorkflowAppLog,)], indirect=True) + def test_save_workflow_app_log_covers_invoke_from_variants(self, sqlite_session: Session): pipeline = _make_pipeline() pipeline._user_id = "user-id" - added: list[object] = [] - - class _Session: - def add(self, item): - added.append(item) pipeline._application_generate_entity.invoke_from = InvokeFrom.EXPLORE - pipeline._save_workflow_app_log(session=_Session(), workflow_run_id="run-id") - assert added[-1].created_from == "installed-app" + pipeline._save_workflow_app_log(session=sqlite_session, workflow_run_id="run-id") pipeline._application_generate_entity.invoke_from = InvokeFrom.WEB_APP - pipeline._save_workflow_app_log(session=_Session(), workflow_run_id="run-id") - assert added[-1].created_from == "web-app" + pipeline._save_workflow_app_log(session=sqlite_session, workflow_run_id="run-id-2") + sqlite_session.flush() + saved_logs = sqlite_session.scalars(select(WorkflowAppLog).order_by(WorkflowAppLog.workflow_run_id)).all() + assert [log.created_from for log in saved_logs] == ["installed-app", "web-app"] - count_before = len(added) + count_before = len(saved_logs) pipeline._application_generate_entity.invoke_from = InvokeFrom.DEBUGGER - pipeline._save_workflow_app_log(session=_Session(), workflow_run_id="run-id") - assert len(added) == count_before + pipeline._save_workflow_app_log(session=sqlite_session, workflow_run_id="run-id-3") + sqlite_session.flush() + assert len(sqlite_session.scalars(select(WorkflowAppLog)).all()) == count_before pipeline._application_generate_entity.invoke_from = InvokeFrom.WEB_APP - pipeline._save_workflow_app_log(session=_Session(), workflow_run_id=None) - assert len(added) == count_before + pipeline._save_workflow_app_log(session=sqlite_session, workflow_run_id=None) + sqlite_session.flush() + assert len(sqlite_session.scalars(select(WorkflowAppLog)).all()) == count_before def test_save_output_for_event_writes_draft_variables(self): pipeline = _make_pipeline() diff --git a/api/tests/unit_tests/core/app/layers/test_trigger_post_layer.py b/api/tests/unit_tests/core/app/layers/test_trigger_post_layer.py index ccdb658b491..150f05081ea 100644 --- a/api/tests/unit_tests/core/app/layers/test_trigger_post_layer.py +++ b/api/tests/unit_tests/core/app/layers/test_trigger_post_layer.py @@ -1,9 +1,13 @@ import logging +from collections.abc import Iterator +from dataclasses import dataclass from datetime import UTC, datetime, timedelta from types import SimpleNamespace from unittest.mock import Mock, patch import pytest +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session, sessionmaker from core.app.layers.trigger_post_layer import TriggerPostLayer from core.workflow.system_variables import build_system_variables @@ -13,19 +17,63 @@ from graphon.graph_events import ( GraphRunSucceededEvent, ) from graphon.runtime import VariablePool -from models.enums import WorkflowTriggerStatus +from models.enums import AppTriggerType, CreatorUserRole, WorkflowTriggerStatus +from models.trigger import WorkflowTriggerLog + + +@dataclass(frozen=True) +class TriggerDatabase: + session: Session + statements: list[str] + + +@pytest.fixture(autouse=True) +def trigger_database(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> Iterator[TriggerDatabase]: + """Create the trigger-log table and bind layer-owned sessions to SQLite.""" + WorkflowTriggerLog.metadata.create_all(sqlite_engine, tables=[WorkflowTriggerLog.__table__]) + sqlite_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr("core.db.session_factory._session_maker", sqlite_session_maker) + statements: list[str] = [] + + def record_statement(_connection, _cursor, statement, _parameters, _context, _executemany) -> None: + statements.append(statement) + + event.listen(sqlite_engine, "before_cursor_execute", record_statement) + with sqlite_session_maker() as session: + try: + yield TriggerDatabase(session=session, statements=statements) + finally: + event.remove(sqlite_engine, "before_cursor_execute", record_statement) + + +def _persist_trigger_log(database: TriggerDatabase, *, trigger_log_id: str = "log-1") -> WorkflowTriggerLog: + trigger_log = WorkflowTriggerLog( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_run_id=None, + root_node_id=None, + trigger_metadata="{}", + trigger_type=AppTriggerType.TRIGGER_WEBHOOK, + trigger_data="{}", + inputs="{}", + outputs=None, + status=WorkflowTriggerStatus.RUNNING, + error=None, + queue_name="workflow", + celery_task_id=None, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + ) + trigger_log.id = trigger_log_id + database.session.add(trigger_log) + database.session.commit() + return trigger_log class TestTriggerPostLayer: - def test_on_event_updates_trigger_log(self): - trigger_log = SimpleNamespace( - status=None, - workflow_run_id=None, - outputs=None, - elapsed_time=None, - total_tokens=None, - finished_at=None, - ) + def test_on_event_updates_trigger_log(self, trigger_database: TriggerDatabase): + trigger_log = _persist_trigger_log(trigger_database) runtime_state = SimpleNamespace( outputs={"answer": "ok"}, variable_pool=VariablePool.from_bootstrap( @@ -35,19 +83,10 @@ class TestTriggerPostLayer: ) with ( - patch("core.app.layers.trigger_post_layer.session_factory") as mock_session_factory, - patch("core.app.layers.trigger_post_layer.SQLAlchemyWorkflowTriggerLogRepository") as mock_repo_cls, patch("core.app.layers.trigger_post_layer.datetime") as mock_datetime, ): mock_datetime.now.return_value = datetime(2026, 2, 20, tzinfo=UTC) - session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = session - - repo = Mock() - repo.get_by_id.return_value = trigger_log - mock_repo_cls.return_value = repo - layer = TriggerPostLayer( cfs_plan_scheduler_entity=Mock(), start_time=datetime(2026, 2, 20, tzinfo=UTC) - timedelta(seconds=10), @@ -57,25 +96,18 @@ class TestTriggerPostLayer: layer.on_event(GraphRunSucceededEvent()) - assert trigger_log.status == WorkflowTriggerStatus.SUCCEEDED - assert trigger_log.workflow_run_id == "run-1" - assert trigger_log.outputs is not None - assert trigger_log.elapsed_time is not None - assert trigger_log.total_tokens == 12 - assert trigger_log.finished_at is not None - repo.update.assert_called_once_with(trigger_log) - session.commit.assert_called_once() + trigger_database.session.expire_all() + persisted_log = trigger_database.session.get(WorkflowTriggerLog, trigger_log.id) + assert persisted_log is not None + assert persisted_log.status == WorkflowTriggerStatus.SUCCEEDED + assert persisted_log.workflow_run_id == "run-1" + assert persisted_log.outputs == '{"answer":"ok"}' + assert persisted_log.elapsed_time == 10 + assert persisted_log.total_tokens == 12 + assert persisted_log.finished_at is not None - def test_on_event_updates_trigger_log_for_aborted_event(self): - trigger_log = SimpleNamespace( - status=None, - workflow_run_id=None, - outputs=None, - error=None, - elapsed_time=None, - total_tokens=None, - finished_at=None, - ) + def test_on_event_updates_trigger_log_for_aborted_event(self, trigger_database: TriggerDatabase): + trigger_log = _persist_trigger_log(trigger_database) runtime_state = SimpleNamespace( outputs={"partial": "ok"}, variable_pool=VariablePool.from_bootstrap( @@ -85,19 +117,10 @@ class TestTriggerPostLayer: ) with ( - patch("core.app.layers.trigger_post_layer.session_factory") as mock_session_factory, - patch("core.app.layers.trigger_post_layer.SQLAlchemyWorkflowTriggerLogRepository") as mock_repo_cls, patch("core.app.layers.trigger_post_layer.datetime") as mock_datetime, ): mock_datetime.now.return_value = datetime(2026, 2, 20, tzinfo=UTC) - session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = session - - repo = Mock() - repo.get_by_id.return_value = trigger_log - mock_repo_cls.return_value = repo - layer = TriggerPostLayer( cfs_plan_scheduler_entity=Mock(), start_time=datetime(2026, 2, 20, tzinfo=UTC) - timedelta(seconds=10), @@ -107,17 +130,22 @@ class TestTriggerPostLayer: layer.on_event(GraphRunAbortedEvent(reason="timeout")) - assert trigger_log.status == WorkflowTriggerStatus.FAILED - assert trigger_log.workflow_run_id == "run-1" - assert trigger_log.outputs is not None - assert trigger_log.error == "timeout" - assert trigger_log.elapsed_time is not None - assert trigger_log.total_tokens == 7 - assert trigger_log.finished_at is not None - repo.update.assert_called_once_with(trigger_log) - session.commit.assert_called_once() + trigger_database.session.expire_all() + persisted_log = trigger_database.session.get(WorkflowTriggerLog, trigger_log.id) + assert persisted_log is not None + assert persisted_log.status == WorkflowTriggerStatus.FAILED + assert persisted_log.workflow_run_id == "run-1" + assert persisted_log.outputs == '{"partial":"ok"}' + assert persisted_log.error == "timeout" + assert persisted_log.elapsed_time == 10 + assert persisted_log.total_tokens == 7 + assert persisted_log.finished_at is not None - def test_on_event_handles_missing_trigger_log(self, caplog: pytest.LogCaptureFixture): + def test_on_event_handles_missing_trigger_log( + self, + caplog: pytest.LogCaptureFixture, + trigger_database: TriggerDatabase, + ): runtime_state = SimpleNamespace( outputs={}, variable_pool=VariablePool.from_bootstrap( @@ -126,31 +154,20 @@ class TestTriggerPostLayer: total_tokens=0, ) - with ( - patch("core.app.layers.trigger_post_layer.session_factory") as mock_session_factory, - patch("core.app.layers.trigger_post_layer.SQLAlchemyWorkflowTriggerLogRepository") as mock_repo_cls, - ): - session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = session + layer = TriggerPostLayer( + cfs_plan_scheduler_entity=Mock(), + start_time=datetime(2026, 2, 20, tzinfo=UTC), + trigger_log_id="missing", + ) + layer.initialize(runtime_state, Mock()) - repo = Mock() - repo.get_by_id.return_value = None - mock_repo_cls.return_value = repo - - layer = TriggerPostLayer( - cfs_plan_scheduler_entity=Mock(), - start_time=datetime(2026, 2, 20, tzinfo=UTC), - trigger_log_id="missing", - ) - layer.initialize(runtime_state, Mock()) - - with caplog.at_level(logging.ERROR, logger="core.app.layers.trigger_post_layer"): - layer.on_event(GraphRunFailedEvent(error="boom")) + with caplog.at_level(logging.ERROR, logger="core.app.layers.trigger_post_layer"): + layer.on_event(GraphRunFailedEvent(error="boom")) assert any(record.levelno == logging.ERROR for record in caplog.records) - session.commit.assert_not_called() + assert trigger_database.session.get(WorkflowTriggerLog, "missing") is None - def test_on_event_ignores_non_status_events(self): + def test_on_event_ignores_non_status_events(self, trigger_database: TriggerDatabase): runtime_state = SimpleNamespace( outputs={}, variable_pool=VariablePool.from_bootstrap( @@ -159,14 +176,14 @@ class TestTriggerPostLayer: total_tokens=0, ) - with patch("core.app.layers.trigger_post_layer.session_factory") as mock_session_factory: - layer = TriggerPostLayer( - cfs_plan_scheduler_entity=Mock(), - start_time=datetime(2026, 2, 20, tzinfo=UTC), - trigger_log_id="log-1", - ) - layer.initialize(runtime_state, Mock()) + layer = TriggerPostLayer( + cfs_plan_scheduler_entity=Mock(), + start_time=datetime(2026, 2, 20, tzinfo=UTC), + trigger_log_id="log-1", + ) + layer.initialize(runtime_state, Mock()) - layer.on_event(Mock()) + trigger_database.statements.clear() + layer.on_event(Mock()) - mock_session_factory.create_session.assert_not_called() + assert trigger_database.statements == [] diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_based_generate_task_pipeline.py b/api/tests/unit_tests/core/app/task_pipeline/test_based_generate_task_pipeline.py index 527d2c76cfa..18c2fdd97a2 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_based_generate_task_pipeline.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_based_generate_task_pipeline.py @@ -3,6 +3,7 @@ from types import SimpleNamespace from unittest.mock import Mock import pytest +from dify_agent.protocol import RunFailureType from sqlalchemy.orm import Session from clients.agent_backend.errors import AgentBackendRunFailedError @@ -151,6 +152,22 @@ class TestBasedGenerateTaskPipeline: "message": "Knowledge retrieval failed (agent_run_id=run-1)", } + def test_stream_converter_maps_agent_run_limit_error(self): + data = AppGenerateResponseConverter._error_to_stream_response( + AgentBackendRunFailedError( + "run-1", + {}, + message="run limit reached", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + ) + ) + + assert data == { + "code": "agent_run_limit_exceeded", + "status": 400, + "message": "run limit reached (agent_run_id=run-1)", + } + def test_handle_output_moderation_when_flagged(self, pipeline): handler = Mock() handler.moderation_completion.return_value = ("filtered", True) diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py index 53a70553e8d..7d0c61fe89a 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py @@ -1,7 +1,10 @@ +from collections.abc import Iterator from types import SimpleNamespace from unittest.mock import Mock, patch import pytest +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session from core.app.apps.base_app_queue_manager import AppQueueManager from core.app.entities.app_invoke_entities import ChatAppGenerateEntity @@ -31,15 +34,17 @@ from graphon.model_runtime.entities.message_entities import TextPromptMessageCon from models.model import AppMode -def _patch_stream_session(): - session = Mock() - session_cm = Mock() - session_cm.__enter__ = Mock(return_value=session) - session_cm.__exit__ = Mock(return_value=False) - return session, patch( - "core.app.task_pipeline.easy_ui_based_generate_task_pipeline.session_factory.create_session", - return_value=session_cm, - ) +@pytest.fixture +def committed_sessions(sqlite_engine: Engine) -> Iterator[list[Session]]: + sessions: list[Session] = [] + + def record_commit(session: Session) -> None: + if session.get_bind() is sqlite_engine: + sessions.append(session) + + event.listen(Session, "after_commit", record_commit) + yield sessions + event.remove(Session, "after_commit", record_commit) class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: @@ -245,7 +250,7 @@ class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: answer="Agent response", message_id="test-message-id" ) - def test_message_end_event(self, pipeline, mock_message_cycle_manager, mock_task_state): + def test_message_end_event(self, pipeline, mock_message_cycle_manager, mock_task_state, committed_sessions): """Test handling of message end events.""" # Setup llm_result = Mock(spec=RuntimeLLMResult) @@ -262,19 +267,17 @@ class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: pipeline._save_message = Mock() pipeline._message_end_to_stream_response = Mock(return_value=Mock(spec=MessageEndStreamResponse)) - session, patch_session = _patch_stream_session() - with patch_session: - # Execute - responses = list(pipeline._process_stream_response(publisher=None, trace_manager=None)) + responses = list(pipeline._process_stream_response(publisher=None, trace_manager=None)) # Assert assert len(responses) == 1 assert mock_task_state.llm_result == llm_result - pipeline._save_message.assert_called_once_with(session=session, trace_manager=None) - session.commit.assert_called_once() + session = pipeline._save_message.call_args.kwargs["session"] + assert isinstance(session, Session) + assert committed_sessions == [session] pipeline._message_end_to_stream_response.assert_called_once() - def test_error_event(self, pipeline): + def test_error_event(self, pipeline, committed_sessions): """Test handling of error events.""" # Setup error_event = Mock(spec=QueueErrorEvent) @@ -287,15 +290,14 @@ class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: pipeline.handle_error = Mock(return_value=Exception("Test error")) pipeline.error_to_stream_response = Mock(return_value=Mock(spec=ErrorStreamResponse)) - session, patch_session = _patch_stream_session() - with patch_session: - # Execute - responses = list(pipeline._process_stream_response(publisher=None, trace_manager=None)) + responses = list(pipeline._process_stream_response(publisher=None, trace_manager=None)) # Assert assert len(responses) == 1 + session = pipeline.handle_error.call_args.kwargs["session"] + assert isinstance(session, Session) + assert committed_sessions == [session] pipeline.handle_error.assert_called_once_with(event=error_event, session=session, message_id="test-message-id") - session.commit.assert_called_once() pipeline.error_to_stream_response.assert_called_once() def test_ping_event(self, pipeline): @@ -358,7 +360,7 @@ class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: publisher.publish.assert_any_call(mock_queue_message) publisher.publish.assert_any_call(None) - def test_trace_manager_passed_to_save_message(self, pipeline): + def test_trace_manager_passed_to_save_message(self, pipeline, committed_sessions): """Test that trace manager is passed to _save_message.""" # Setup trace_manager = Mock(spec=TraceQueueManager) @@ -373,14 +375,13 @@ class TestEasyUIBasedGenerateTaskPipelineProcessStreamResponse: pipeline._save_message = Mock() pipeline._message_end_to_stream_response = Mock(return_value=Mock(spec=MessageEndStreamResponse)) - session, patch_session = _patch_stream_session() - with patch_session: - # Execute - list(pipeline._process_stream_response(publisher=None, trace_manager=trace_manager)) + list(pipeline._process_stream_response(publisher=None, trace_manager=trace_manager)) # Assert + session = pipeline._save_message.call_args.kwargs["session"] + assert isinstance(session, Session) + assert committed_sessions == [session] pipeline._save_message.assert_called_once_with(session=session, trace_manager=trace_manager) - session.commit.assert_called_once() def test_multiple_events_sequence(self, pipeline, mock_message_cycle_manager, mock_task_state): """Test handling multiple events in sequence.""" diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py index 49ecc9358c8..69b60b14ab3 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py @@ -7,6 +7,7 @@ from typing import cast from unittest.mock import Mock import pytest +from sqlalchemy.orm import Session from core.app.app_config.entities import ( AppAdditionalFeatures, @@ -56,7 +57,7 @@ from extensions.storage.storage_type import StorageType from graphon.file import FileTransferMethod, FileType from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, TextPromptMessageContent -from models.enums import CreatorUserRole +from models.enums import ConversationFromSource, CreatorUserRole from models.model import AppMode, Conversation, Message, MessageAgentThought, MessageFile, UploadFile @@ -136,14 +137,52 @@ def _unknown_queue_message() -> MessageQueueMessage: def _make_conversation(app_mode: AppMode) -> Conversation: - conversation = Conversation() + conversation = Conversation( + app_id="app", + app_model_config_id=None, + model_provider=None, + override_model_configs=None, + model_id=None, + mode=app_mode, + name="conversation", + inputs={}, + introduction="", + system_instruction="", + system_instruction_tokens=0, + status="normal", + invoke_from=InvokeFrom.WEB_APP, + from_source=ConversationFromSource.API, + from_end_user_id="user", + from_account_id=None, + ) conversation.id = "conv" conversation.mode = app_mode return conversation def _make_message() -> Message: - message = Message() + message = Message( + app_id="app", + conversation_id="conv", + inputs={}, + query="query", + message="", + message_tokens=0, + message_unit_price=0, + message_price_unit=0, + answer="", + answer_tokens=0, + answer_unit_price=0, + answer_price_unit=0, + provider_response_latency=0, + total_price=0, + currency="USD", + invoke_from=InvokeFrom.WEB_APP, + from_source=ConversationFromSource.API, + from_end_user_id="user", + from_account_id=None, + app_mode=AppMode.CHAT, + ) message.id = "msg" message.created_at = datetime.now(UTC) return message @@ -1173,7 +1212,9 @@ class TestEasyUiBasedGenerateTaskPipeline: assert list(pipeline._process_stream_response(publisher=None)) == [] - def test_save_message_persists_fields_and_emits_trace(self, monkeypatch: pytest.MonkeyPatch): + def test_save_message_persists_fields_and_emits_trace( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + ): conversation = _make_conversation(AppMode.CHAT) message = _make_message() application_generate_entity = _make_entity(ChatAppGenerateEntity, AppMode.CHAT) @@ -1195,8 +1236,9 @@ class TestEasyUiBasedGenerateTaskPipeline: message_obj = _make_message() conversation_obj = _make_conversation(AppMode.CHAT) - session = Mock() - session.scalar.side_effect = [message_obj, conversation_obj] + session = sqlite_session + session.add_all([conversation_obj, message_obj]) + session.flush() trace_manager_double = _TraceManagerDouble() trace_manager = cast(TraceQueueManager, trace_manager_double) sent_payloads: list[tuple[tuple[object, ...], dict[str, object]]] = [] @@ -1234,7 +1276,7 @@ class TestEasyUiBasedGenerateTaskPipeline: assert trace_task.kwargs["trace_session_id"] == "session-1" assert len(sent_payloads) == 1 - def test_save_message_raises_when_message_not_found(self): + def test_save_message_raises_when_message_not_found(self, sqlite_session: Session): conversation = _make_conversation(AppMode.CHAT) message = _make_message() pipeline = EasyUIBasedGenerateTaskPipeline( @@ -1244,13 +1286,12 @@ class TestEasyUiBasedGenerateTaskPipeline: message=message, stream=False, ) - session = Mock() - session.scalar.return_value = None + session = sqlite_session with pytest.raises(ValueError, match="message msg not found"): pipeline._save_message(session=session) - def test_save_message_raises_when_conversation_not_found(self): + def test_save_message_raises_when_conversation_not_found(self, sqlite_session: Session): conversation = _make_conversation(AppMode.CHAT) message = _make_message() pipeline = EasyUIBasedGenerateTaskPipeline( @@ -1260,8 +1301,9 @@ class TestEasyUiBasedGenerateTaskPipeline: message=message, stream=False, ) - session = Mock() - session.scalar.side_effect = [_make_message(), None] + session = sqlite_session + session.add(_make_message()) + session.flush() with pytest.raises(ValueError, match="Conversation conv not found"): pipeline._save_message(session=session) diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_message_end_files.py b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_message_end_files.py index c57e84b6f2b..c353fca675f 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_message_end_files.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_message_end_files.py @@ -169,7 +169,7 @@ class TestMessageEndStreamResponseFiles: assert file_dict["url"].startswith("https://example.com/signed-url") assert file_dict["upload_file_id"] == message_file_local.upload_file_id assert file_dict["remote_url"] == "" - get_signed_url.assert_called_once_with(upload_file_id=str(upload_file.id)) + get_signed_url.assert_called_once_with(upload_file_id=upload_file.id) def test_message_end_with_remote_url( self, sqlite_session: Session, mock_pipeline: Mock, message_file_remote: MessageFile diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_message_cycle_manager_optimization.py b/api/tests/unit_tests/core/app/task_pipeline/test_message_cycle_manager_optimization.py index c0cc2d75386..e9fca80d61f 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_message_cycle_manager_optimization.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_message_cycle_manager_optimization.py @@ -1,24 +1,95 @@ """Unit tests for the message cycle manager optimization.""" import logging +from collections.abc import Iterator +from dataclasses import dataclass from types import SimpleNamespace from unittest.mock import Mock, patch import pytest from flask import Flask, current_app +from sqlalchemy import Engine, event, select +from sqlalchemy.orm import Session, sessionmaker from core.app.entities.queue_entities import QueueAnnotationReplyEvent, QueueRetrieverResourcesEvent from core.app.entities.task_entities import MessageStreamResponse, StreamEvent, TaskStateMetadata +from core.app.task_pipeline import message_cycle_manager as message_cycle_manager_module from core.app.task_pipeline.message_cycle_manager import MessageCycleManager from core.rag.entities import RetrievalSourceMetadata -from models.model import App, AppMode +from graphon.file import FileTransferMethod, FileType +from models import model as model_module +from models.base import TypeBase +from models.enums import ConversationFromSource, CreatorUserRole, MessageFileBelongsTo +from models.model import App, AppMode, Conversation, MessageFile -def _patch_create_session(mock_session): - session_cm = Mock() - session_cm.__enter__ = Mock(return_value=mock_session) - session_cm.__exit__ = Mock(return_value=False) - return patch("core.app.task_pipeline.message_cycle_manager.session_factory.create_session", return_value=session_cm) +@dataclass(frozen=True) +class _SQLiteDb: + engine: Engine + session: Session + + +@pytest.fixture +def cycle_db(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]: + """Bind request-owned and cycle-manager-owned sessions to isolated SQLite.""" + TypeBase.metadata.create_all( + sqlite_engine, + tables=[App.__table__, Conversation.__table__, MessageFile.__table__], + ) + owned_session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + with owned_session_factory() as request_session: + sqlite_db = _SQLiteDb(engine=sqlite_engine, session=request_session) + monkeypatch.setattr(message_cycle_manager_module, "db", sqlite_db) + monkeypatch.setattr(model_module, "db", sqlite_db) + monkeypatch.setattr(message_cycle_manager_module.session_factory, "create_session", owned_session_factory) + yield request_session + + +def _app(*, app_id: str = "app-id", tenant_id: str = "tenant-1") -> App: + return App( + id=app_id, + tenant_id=tenant_id, + name="Test App", + description="", + mode=AppMode.CHAT, + enable_site=True, + enable_api=True, + max_active_requests=0, + ) + + +def _conversation(*, conversation_id: str = "conv-1", app_id: str = "app-id") -> Conversation: + conversation = Conversation( + app_id=app_id, + mode=AppMode.CHAT, + name="", + status="normal", + from_source=ConversationFromSource.API, + inputs={}, + ) + conversation.id = conversation_id + return conversation + + +def _message_file( + *, + file_id: str = "file-1", + message_id: str = "test-message-id", + belongs_to: MessageFileBelongsTo | None = MessageFileBelongsTo.ASSISTANT, + url: str | None = "http://example.com/image.png", + file_type: FileType = FileType.IMAGE, +) -> MessageFile: + message_file = MessageFile( + message_id=message_id, + type=file_type, + transfer_method=FileTransferMethod.TOOL_FILE, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-id", + belongs_to=belongs_to, + url=url, + ) + message_file.id = file_id + return message_file class TestMessageCycleManagerOptimization: @@ -37,30 +108,22 @@ class TestMessageCycleManagerOptimization: task_state = Mock() return MessageCycleManager(application_generate_entity=mock_application_generate_entity, task_state=task_state) - def test_get_message_event_type_with_assistant_file(self, message_cycle_manager): + def test_get_message_event_type_with_assistant_file(self, message_cycle_manager, cycle_db: Session): """Test get_message_event_type returns MESSAGE_FILE when message has assistant-generated files. This ensures that AI-generated images (belongs_to='assistant') trigger the MESSAGE_FILE event, allowing the frontend to properly display generated image files with url field. """ - with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory: - # Setup mock session and message file - mock_session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session + cycle_db.add(_message_file()) + cycle_db.commit() - mock_message_file = Mock() - mock_message_file.belongs_to = "assistant" - mock_session.scalar.return_value = mock_message_file + with current_app.app_context(): + result = message_cycle_manager.get_message_event_type("test-message-id") - # Execute - with current_app.app_context(): - result = message_cycle_manager.get_message_event_type("test-message-id") + assert result == StreamEvent.MESSAGE_FILE + assert "test-message-id" in message_cycle_manager._message_has_file - # Assert - assert result == StreamEvent.MESSAGE_FILE - mock_session.scalar.assert_called_once() - - def test_get_message_event_type_with_user_file(self, message_cycle_manager): + def test_get_message_event_type_with_user_file(self, message_cycle_manager, cycle_db: Session): """Test get_message_event_type returns MESSAGE when message only has user-uploaded files. This is a regression test for the issue where user-uploaded images (belongs_to='user') @@ -68,90 +131,81 @@ class TestMessageCycleManagerOptimization: resulting in broken images in the chat UI. The query filters for belongs_to='assistant', so when only user files exist, the database query returns None, resulting in MESSAGE event type. """ - with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory: - # Setup mock session and message file - mock_session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session + cycle_db.add(_message_file(belongs_to=MessageFileBelongsTo.USER)) + cycle_db.commit() - # When querying for assistant files with only user files present, return None - # (simulates database query with belongs_to='assistant' filter returning no results) - mock_session.scalar.return_value = None + with current_app.app_context(): + result = message_cycle_manager.get_message_event_type("test-message-id") - # Execute - with current_app.app_context(): - result = message_cycle_manager.get_message_event_type("test-message-id") + assert result == StreamEvent.MESSAGE + assert "test-message-id" not in message_cycle_manager._message_has_file - # Assert - assert result == StreamEvent.MESSAGE - mock_session.scalar.assert_called_once() - - def test_get_message_event_type_without_message_file(self, message_cycle_manager): + def test_get_message_event_type_without_message_file(self, message_cycle_manager, cycle_db: Session): """Test get_message_event_type returns MESSAGE when message has no files.""" - with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory: - # Setup mock session and no message file - mock_session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session - # Current implementation uses session.scalar(select(...)) - mock_session.scalar.return_value = None + assert list(cycle_db.scalars(select(MessageFile)).all()) == [] - # Execute - with current_app.app_context(): - result = message_cycle_manager.get_message_event_type("test-message-id") + with current_app.app_context(): + result = message_cycle_manager.get_message_event_type("test-message-id") - # Assert - assert result == StreamEvent.MESSAGE - mock_session.scalar.assert_called_once() + assert result == StreamEvent.MESSAGE - def test_get_message_event_type_uses_cache_without_query(self, message_cycle_manager): + def test_get_message_event_type_uses_cache_without_query( + self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine + ): """Return MESSAGE_FILE directly from in-memory cache without opening a DB session.""" message_cycle_manager._message_has_file.add("cached-message") + statements: list[str] = [] - with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory: + def record_statement(_conn, _cursor, statement, _parameters, _context, _executemany) -> None: + statements.append(statement) + + event.listen(sqlite_engine, "before_cursor_execute", record_statement) + try: result = message_cycle_manager.get_message_event_type("cached-message") + finally: + event.remove(sqlite_engine, "before_cursor_execute", record_statement) assert result == StreamEvent.MESSAGE_FILE - mock_session_factory.create_session.assert_not_called() + assert statements == [] - def test_message_to_stream_response_with_precomputed_event_type(self, message_cycle_manager): + def test_message_to_stream_response_with_precomputed_event_type(self, message_cycle_manager, cycle_db: Session): """MessageCycleManager.message_to_stream_response expects a valid event_type; callers should precompute it.""" - with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory: - # Setup mock session and message file - mock_session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session + cycle_db.add(_message_file()) + cycle_db.commit() - mock_message_file = Mock() - mock_message_file.belongs_to = "assistant" - mock_session.scalar.return_value = mock_message_file + with current_app.app_context(): + event_type = message_cycle_manager.get_message_event_type("test-message-id") + result = message_cycle_manager.message_to_stream_response( + answer="Hello world", message_id="test-message-id", event_type=event_type + ) - # Execute: compute event type once, then pass to message_to_stream_response - with current_app.app_context(): - event_type = message_cycle_manager.get_message_event_type("test-message-id") - result = message_cycle_manager.message_to_stream_response( - answer="Hello world", message_id="test-message-id", event_type=event_type - ) + assert isinstance(result, MessageStreamResponse) + assert result.answer == "Hello world" + assert result.id == "test-message-id" + assert result.event == StreamEvent.MESSAGE_FILE - # Assert - assert isinstance(result, MessageStreamResponse) - assert result.answer == "Hello world" - assert result.id == "test-message-id" - assert result.event == StreamEvent.MESSAGE_FILE - mock_session.scalar.assert_called_once() - - def test_message_to_stream_response_with_event_type_skips_query(self, message_cycle_manager): + def test_message_to_stream_response_with_event_type_skips_query( + self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine + ): """Test that message_to_stream_response skips database query when event_type is provided.""" - with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory: - # Execute with event_type provided + statements: list[str] = [] + + def record_statement(_conn, _cursor, statement, _parameters, _context, _executemany) -> None: + statements.append(statement) + + event.listen(sqlite_engine, "before_cursor_execute", record_statement) + try: result = message_cycle_manager.message_to_stream_response( answer="Hello world", message_id="test-message-id", event_type=StreamEvent.MESSAGE ) + finally: + event.remove(sqlite_engine, "before_cursor_execute", record_statement) - # Assert - assert isinstance(result, MessageStreamResponse) - assert result.answer == "Hello world" - assert result.id == "test-message-id" - assert result.event == StreamEvent.MESSAGE - # Should not open a session when event_type is provided - mock_session_factory.create_session.assert_not_called() + assert isinstance(result, MessageStreamResponse) + assert result.answer == "Hello world" + assert result.id == "test-message-id" + assert result.event == StreamEvent.MESSAGE + assert statements == [] def test_message_to_stream_response_with_from_variable_selector(self, message_cycle_manager): """Test message_to_stream_response with from_variable_selector parameter.""" @@ -168,40 +222,32 @@ class TestMessageCycleManagerOptimization: assert result.from_variable_selector == ["var1", "var2"] assert result.event == StreamEvent.MESSAGE - def test_optimization_usage_example(self, message_cycle_manager): + def test_optimization_usage_example(self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine): """Test the optimization pattern that should be used by callers.""" - # Step 1: Get event type once (this queries database) - with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory: - mock_session = Mock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session - # Current implementation uses session.scalar(select(...)) - mock_session.scalar.return_value = None # No files + statements: list[str] = [] + + def record_statement(_conn, _cursor, statement, _parameters, _context, _executemany) -> None: + statements.append(statement) + + event.listen(sqlite_engine, "before_cursor_execute", record_statement) + try: with current_app.app_context(): event_type = message_cycle_manager.get_message_event_type("test-message-id") - - # Should open session once - mock_session_factory.create_session.assert_called_once() - assert event_type == StreamEvent.MESSAGE - - # Step 2: Use event_type for multiple calls (no additional queries) - with patch("core.app.task_pipeline.message_cycle_manager.session_factory") as mock_session_factory: - mock_session_factory.create_session.return_value.__enter__.return_value = Mock() - chunk1_response = message_cycle_manager.message_to_stream_response( answer="Chunk 1", message_id="test-message-id", event_type=event_type ) - chunk2_response = message_cycle_manager.message_to_stream_response( answer="Chunk 2", message_id="test-message-id", event_type=event_type ) + finally: + event.remove(sqlite_engine, "before_cursor_execute", record_statement) - # Should not open session again when event_type provided - mock_session_factory.create_session.assert_not_called() - - assert chunk1_response.event == StreamEvent.MESSAGE - assert chunk2_response.event == StreamEvent.MESSAGE - assert chunk1_response.answer == "Chunk 1" - assert chunk2_response.answer == "Chunk 2" + assert event_type == StreamEvent.MESSAGE + assert len([statement for statement in statements if statement.lstrip().upper().startswith("SELECT")]) == 1 + assert chunk1_response.event == StreamEvent.MESSAGE + assert chunk2_response.event == StreamEvent.MESSAGE + assert chunk1_response.answer == "Chunk 1" + assert chunk2_response.answer == "Chunk 2" def test_generate_conversation_name_returns_none_for_completion(self, message_cycle_manager): """Return None when completion entities are used for conversation naming. @@ -269,51 +315,38 @@ class TestMessageCycleManagerOptimization: assert message_cycle_manager._application_generate_entity.is_new_conversation is False mock_timer.assert_not_called() - def test_generate_conversation_name_worker_returns_when_conversation_missing(self, message_cycle_manager): + def test_generate_conversation_name_worker_returns_when_conversation_missing( + self, message_cycle_manager, cycle_db: Session + ): """Return early when the conversation cannot be found.""" flask_app = Flask(__name__) - db_session = Mock() - db_session.scalar.return_value = None + assert list(cycle_db.scalars(select(Conversation)).all()) == [] - with _patch_create_session(db_session): - message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-missing", "hello") + message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-missing", "hello") - db_session.commit.assert_not_called() + assert list(cycle_db.scalars(select(Conversation)).all()) == [] - def test_generate_conversation_name_worker_returns_when_app_missing(self, message_cycle_manager): + def test_generate_conversation_name_worker_returns_when_app_missing(self, message_cycle_manager, cycle_db: Session): """Return early when non-completion conversation has no app relation.""" flask_app = Flask(__name__) - conversation = SimpleNamespace(mode=AppMode.CHAT, app=None, app_id="app-id") - db_session = Mock() - db_session.scalar.return_value = conversation - db_session.get.return_value = None + conversation = _conversation() + cycle_db.add(conversation) + cycle_db.commit() - with _patch_create_session(db_session): - message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello") + message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello") - db_session.commit.assert_not_called() + assert cycle_db.get(Conversation, "conv-1").name == "" + assert cycle_db.get(App, "app-id") is None - def test_generate_conversation_name_worker_uses_cached_name(self, message_cycle_manager): + def test_generate_conversation_name_worker_uses_cached_name( + self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine + ): """Use cached conversation name when present and avoid LLM call.""" flask_app = Flask(__name__) - - class ConversationWithPoisonedApp: - mode = AppMode.CHAT - app_id = "app-id" - name = "" - - @property - def app(self): - raise AssertionError("conversation.app must not open an implicit session") - - conversation = ConversationWithPoisonedApp() - app_model = SimpleNamespace(tenant_id="tenant-1") - db_session = Mock() - db_session.scalar.return_value = conversation - db_session.get.return_value = app_model + cycle_db.add_all([_app(), _conversation()]) + cycle_db.commit() with ( - _patch_create_session(db_session) as create_session, patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis, patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator, ): @@ -321,27 +354,23 @@ class TestMessageCycleManagerOptimization: message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello") + assert cycle_db.in_transaction() is False + with Session(sqlite_engine) as verification_session: + conversation = verification_session.get(Conversation, "conv-1") + assert conversation is not None assert conversation.name == "cached-title" - create_session.assert_called_once_with() - db_session.get.assert_called_once_with(App, "app-id") - db_session.commit.assert_called_once() mock_llm_generator.generate_conversation_name.assert_not_called() mock_redis.setex.assert_not_called() - def test_generate_conversation_name_worker_generates_and_caches_name(self, message_cycle_manager): + def test_generate_conversation_name_worker_generates_and_caches_name( + self, message_cycle_manager, cycle_db: Session, sqlite_engine: Engine + ): """Generate conversation name and write it to redis cache on cache miss.""" flask_app = Flask(__name__) - conversation = SimpleNamespace( - mode=AppMode.CHAT, - app=SimpleNamespace(tenant_id="tenant-1"), - app_id="app-id", - name="", - ) - db_session = Mock() - db_session.scalar.return_value = conversation + cycle_db.add_all([_app(), _conversation()]) + cycle_db.commit() with ( - _patch_create_session(db_session), patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis, patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator, ): @@ -350,27 +379,27 @@ class TestMessageCycleManagerOptimization: message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", "hello") + assert cycle_db.in_transaction() is False + with Session(sqlite_engine) as verification_session: + conversation = verification_session.get(Conversation, "conv-1") + assert conversation is not None assert conversation.name == "generated-title" - db_session.commit.assert_called_once() mock_redis.setex.assert_called_once() def test_generate_conversation_name_worker_falls_back_when_generation_fails( - self, message_cycle_manager, caplog: pytest.LogCaptureFixture + self, + message_cycle_manager, + cycle_db: Session, + sqlite_engine: Engine, + caplog: pytest.LogCaptureFixture, ): """Fallback to truncated query when LLM generation fails.""" flask_app = Flask(__name__) - conversation = SimpleNamespace( - mode=AppMode.CHAT, - app=SimpleNamespace(tenant_id="tenant-1"), - app_id="app-id", - name="", - ) - db_session = Mock() - db_session.scalar.return_value = conversation + cycle_db.add_all([_app(), _conversation()]) + cycle_db.commit() long_query = "q" * 60 with ( - _patch_create_session(db_session), patch("core.app.task_pipeline.message_cycle_manager.redis_client") as mock_redis, patch("core.app.task_pipeline.message_cycle_manager.LLMGenerator") as mock_llm_generator, patch("core.app.task_pipeline.message_cycle_manager.dify_config") as mock_dify_config, @@ -382,11 +411,14 @@ class TestMessageCycleManagerOptimization: with caplog.at_level(logging.ERROR, logger="core.app.task_pipeline.message_cycle_manager"): message_cycle_manager._generate_conversation_name_worker(flask_app, "conv-1", long_query) + assert cycle_db.in_transaction() is False + with Session(sqlite_engine) as verification_session: + conversation = verification_session.get(Conversation, "conv-1") + assert conversation is not None assert conversation.name == (long_query[:47] + "...") - db_session.commit.assert_called_once() assert any(record.levelno == logging.ERROR for record in caplog.records) - def test_handle_annotation_reply_sets_metadata(self, message_cycle_manager): + def test_handle_annotation_reply_sets_metadata(self, message_cycle_manager, unbound_session: Session): """Populate task metadata from annotation reply events. Args: message_cycle_manager with TaskStateMetadata and a mocked AppAnnotationService. @@ -399,7 +431,7 @@ class TestMessageCycleManagerOptimization: id="ann-1", account_id="acct-1", ) - session = Mock() + session = unbound_session with ( patch("core.app.task_pipeline.message_cycle_manager.AppAnnotationService") as mock_service, @@ -454,33 +486,25 @@ class TestMessageCycleManagerOptimization: assert message_cycle_manager._task_state.metadata.retriever_resources[0].position == 1 assert message_cycle_manager._task_state.metadata.retriever_resources[1].position == 2 - def test_message_file_to_stream_response_builds_signed_url(self, message_cycle_manager): + def test_message_file_to_stream_response_builds_signed_url(self, message_cycle_manager, cycle_db: Session): """Build a stream response with a signed tool file URL. - Args: message_cycle_manager with mocked Session/db and sign_tool_file. + Args: message_cycle_manager with a persisted MessageFile and mocked sign_tool_file. Returns: MessageStreamResponse with signed url and belongs_to normalized to user. Side effects: Calls sign_tool_file for tool file ids. """ message_cycle_manager._application_generate_entity.task_id = "task-1" - - message_file = SimpleNamespace( - id="file-1", - type="image", - belongs_to=None, - url="tool://file.verylongextension", - message_id="msg-1", + cycle_db.add( + _message_file( + file_id="file-1", + message_id="msg-1", + belongs_to=None, + url="tool://file.verylongextension", + ) ) + cycle_db.commit() - session = Mock() - session.scalar.return_value = message_file - - with ( - patch("core.app.task_pipeline.message_cycle_manager.Session") as mock_session_cls, - patch("core.app.task_pipeline.message_cycle_manager.sign_tool_file") as mock_sign, - patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db, - ): - mock_db.engine = Mock() - mock_session_cls.return_value.__enter__.return_value = session + with patch("core.app.task_pipeline.message_cycle_manager.sign_tool_file") as mock_sign: mock_sign.return_value = "signed-url" response = message_cycle_manager.message_file_to_stream_response(SimpleNamespace(message_file_id="file-1")) @@ -514,56 +538,42 @@ class TestMessageCycleManagerOptimization: assert len(message_cycle_manager._task_state.metadata.retriever_resources) == 1 assert message_cycle_manager._task_state.metadata.retriever_resources[0].position == 1 - def test_message_file_to_stream_response_uses_http_url_directly(self, message_cycle_manager): + def test_message_file_to_stream_response_uses_http_url_directly(self, message_cycle_manager, cycle_db: Session): """Use original URL when message file URL is already HTTP.""" message_cycle_manager._application_generate_entity.task_id = "task-http" - message_file = SimpleNamespace( - id="file-http", - type="image", - belongs_to="assistant", - url="http://example.com/pic.png", - message_id="msg-http", - ) - - session = Mock() - session.scalar.return_value = message_file - - with ( - patch("core.app.task_pipeline.message_cycle_manager.Session") as mock_session_cls, - patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db, - ): - mock_db.engine = Mock() - mock_session_cls.return_value.__enter__.return_value = session - - response = message_cycle_manager.message_file_to_stream_response( - SimpleNamespace(message_file_id="file-http") + cycle_db.add( + _message_file( + file_id="file-http", + message_id="msg-http", + belongs_to=MessageFileBelongsTo.ASSISTANT, + url="http://example.com/pic.png", ) + ) + cycle_db.commit() + + response = message_cycle_manager.message_file_to_stream_response(SimpleNamespace(message_file_id="file-http")) assert response is not None assert response.url == "http://example.com/pic.png" assert "msg-http" in message_cycle_manager._message_has_file - def test_message_file_to_stream_response_defaults_extension_to_bin_without_dot(self, message_cycle_manager): + def test_message_file_to_stream_response_defaults_extension_to_bin_without_dot( + self, message_cycle_manager, cycle_db: Session + ): """Default tool file extension to .bin when URL has no extension part.""" message_cycle_manager._application_generate_entity.task_id = "task-bin" - message_file = SimpleNamespace( - id="file-bin", - type="file", - belongs_to="assistant", - url="tool-file-id", - message_id="msg-bin", + cycle_db.add( + _message_file( + file_id="file-bin", + message_id="msg-bin", + belongs_to=MessageFileBelongsTo.ASSISTANT, + url="tool-file-id", + file_type=FileType.CUSTOM, + ) ) + cycle_db.commit() - session = Mock() - session.scalar.return_value = message_file - - with ( - patch("core.app.task_pipeline.message_cycle_manager.Session") as mock_session_cls, - patch("core.app.task_pipeline.message_cycle_manager.sign_tool_file") as mock_sign, - patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db, - ): - mock_db.engine = Mock() - mock_session_cls.return_value.__enter__.return_value = session + with patch("core.app.task_pipeline.message_cycle_manager.sign_tool_file") as mock_sign: mock_sign.return_value = "signed-bin-url" response = message_cycle_manager.message_file_to_stream_response( @@ -574,19 +584,13 @@ class TestMessageCycleManagerOptimization: assert response.url == "signed-bin-url" mock_sign.assert_called_once_with(tool_file_id="tool-file-id", extension=".bin") - def test_message_file_to_stream_response_returns_none_when_file_missing(self, message_cycle_manager): + def test_message_file_to_stream_response_returns_none_when_file_missing( + self, message_cycle_manager, cycle_db: Session + ): """Return None when message file lookup does not find a record.""" - session = Mock() - session.scalar.return_value = None + assert list(cycle_db.scalars(select(MessageFile)).all()) == [] - with ( - patch("core.app.task_pipeline.message_cycle_manager.Session") as mock_session_cls, - patch("core.app.task_pipeline.message_cycle_manager.db") as mock_db, - ): - mock_db.engine = Mock() - mock_session_cls.return_value.__enter__.return_value = session - - response = message_cycle_manager.message_file_to_stream_response(SimpleNamespace(message_file_id="missing")) + response = message_cycle_manager.message_file_to_stream_response(SimpleNamespace(message_file_id="missing")) assert response is None diff --git a/api/tests/unit_tests/core/app/test_llm_quota.py b/api/tests/unit_tests/core/app/test_llm_quota.py index ba691bb298a..e7297639610 100644 --- a/api/tests/unit_tests/core/app/test_llm_quota.py +++ b/api/tests/unit_tests/core/app/test_llm_quota.py @@ -10,10 +10,12 @@ from sqlalchemy.orm import sessionmaker from configs import dify_config from core.app.llm.quota import ( + LLMQuotaReservationState, deduct_llm_quota, deduct_llm_quota_for_model, ensure_llm_quota_available, ensure_llm_quota_available_for_model, + reserve_llm_quota_for_model, ) from core.entities.model_entities import ModelStatus from core.entities.provider_entities import ProviderQuotaType, QuotaUnit @@ -100,6 +102,111 @@ def test_ensure_llm_quota_available_for_model_ignores_custom_provider_configurat provider_configuration.get_provider_model.assert_not_called() +def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> None: + credit_reservation = MagicMock() + provider_configuration = SimpleNamespace( + using_provider_type=ProviderType.SYSTEM, + get_provider_model=MagicMock(return_value=SimpleNamespace(status=ModelStatus.ACTIVE)), + system_configuration=SimpleNamespace( + current_quota_type=ProviderQuotaType.TRIAL, + quota_configurations=[ + SimpleNamespace( + quota_type=ProviderQuotaType.TRIAL, + quota_unit=QuotaUnit.CREDITS, + quota_limit=100, + ) + ], + ), + ) + provider_manager = MagicMock() + provider_manager.get_configurations.return_value.get.return_value = provider_configuration + + with ( + patch("core.app.llm.quota.create_plugin_provider_manager", return_value=provider_manager), + patch.object(type(dify_config), "get_model_credits", return_value=9), + patch("core.app.llm.quota.CreditPoolService.reserve_credits", return_value=credit_reservation) as reserve, + ): + reservation = reserve_llm_quota_for_model( + tenant_id="tenant-id", + provider="openai", + model="gpt-4o", + ) + reservation.commit(LLMUsage.empty_usage()) + reservation.release() + + assert reservation.state == LLMQuotaReservationState.COMMITTED + assert reservation.commit_before_delivery is True + reserve.assert_called_once_with( + tenant_id="tenant-id", + credits_required=9, + pool_type="trial", + request_id=ANY, + session_factory=ANY, + meta={"source": "llm.invoke", "provider": "openai", "model": "gpt-4o"}, + ) + credit_reservation.commit.assert_called_once_with() + credit_reservation.release.assert_not_called() + + +def test_reserve_llm_quota_requires_accurate_usage_for_free_tokens() -> None: + provider_configuration = SimpleNamespace( + using_provider_type=ProviderType.SYSTEM, + get_provider_model=MagicMock(return_value=SimpleNamespace(status=ModelStatus.ACTIVE)), + system_configuration=SimpleNamespace( + current_quota_type=ProviderQuotaType.FREE, + quota_configurations=[ + SimpleNamespace( + quota_type=ProviderQuotaType.FREE, + quota_unit=QuotaUnit.TOKENS, + quota_limit=100, + ) + ], + ), + ) + provider_manager = MagicMock() + provider_manager.get_configurations.return_value.get.return_value = provider_configuration + + with patch("core.app.llm.quota.create_plugin_provider_manager", return_value=provider_manager): + reservation = reserve_llm_quota_for_model( + tenant_id="tenant-id", + provider="openai", + model="gpt-4o", + ) + + assert reservation.commit_before_delivery is False + with pytest.raises(ValueError, match="Accurate terminal usage"): + reservation.commit() + + +def test_reserve_llm_quota_rejects_token_based_credit_pool() -> None: + provider_configuration = SimpleNamespace( + using_provider_type=ProviderType.SYSTEM, + get_provider_model=MagicMock(return_value=SimpleNamespace(status=ModelStatus.ACTIVE)), + system_configuration=SimpleNamespace( + current_quota_type=ProviderQuotaType.TRIAL, + quota_configurations=[ + SimpleNamespace( + quota_type=ProviderQuotaType.TRIAL, + quota_unit=QuotaUnit.TOKENS, + quota_limit=100, + ) + ], + ), + ) + provider_manager = MagicMock() + provider_manager.get_configurations.return_value.get.return_value = provider_configuration + + with ( + patch("core.app.llm.quota.create_plugin_provider_manager", return_value=provider_manager), + pytest.raises(ValueError, match="do not support pre-invocation reservation"), + ): + reserve_llm_quota_for_model( + tenant_id="tenant-id", + provider="openai", + model="gpt-4o", + ) + + def test_deduct_llm_quota_for_model_uses_identity_based_trial_billing() -> None: usage = LLMUsage.empty_usage() usage.total_tokens = 42 diff --git a/api/tests/unit_tests/core/app/workflow/layers/test_persistence.py b/api/tests/unit_tests/core/app/workflow/layers/test_persistence.py index 5c50cb78dae..b7a8754d92c 100644 --- a/api/tests/unit_tests/core/app/workflow/layers/test_persistence.py +++ b/api/tests/unit_tests/core/app/workflow/layers/test_persistence.py @@ -34,6 +34,7 @@ def test_update_node_execution_prefers_event_finished_at(monkeypatch: pytest.Mon node_execution = Mock() node_execution.id = "node-exec-1" node_execution.created_at = datetime(2024, 1, 1, 0, 0, 0, tzinfo=UTC).replace(tzinfo=None) + node_execution.process_data = None node_execution.update_from_mapping = Mock() layer._node_snapshots[node_execution.id] = _NodeRuntimeSnapshot( @@ -66,6 +67,7 @@ def test_update_node_execution_projects_start_outputs() -> None: node_execution.id = "node-exec-2" node_execution.node_type = BuiltinNodeTypes.START node_execution.created_at = datetime(2024, 1, 1, 0, 0, 0, tzinfo=UTC).replace(tzinfo=None) + node_execution.process_data = None node_execution.update_from_mapping = Mock() layer._node_snapshots[node_execution.id] = _NodeRuntimeSnapshot( diff --git a/api/tests/unit_tests/core/app/workflow/test_file_runtime.py b/api/tests/unit_tests/core/app/workflow/test_file_runtime.py index 0025c21f437..02b4ad362ec 100644 --- a/api/tests/unit_tests/core/app/workflow/test_file_runtime.py +++ b/api/tests/unit_tests/core/app/workflow/test_file_runtime.py @@ -3,19 +3,85 @@ from __future__ import annotations import base64 import hashlib import hmac +from collections.abc import Iterator +from datetime import UTC, datetime from types import SimpleNamespace from unittest.mock import MagicMock, patch from urllib.parse import parse_qs, urlparse import pytest +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom from core.app.file_access import DatabaseFileAccessController, FileAccessScope from core.app.workflow import file_runtime from core.app.workflow.file_runtime import DifyWorkflowFileRuntime, bind_dify_workflow_file_runtime from core.workflow.file_reference import build_file_reference +from extensions.storage.storage_type import StorageType from graphon.file import File, FileTransferMethod, FileType from models import ToolFile, UploadFile +from models.base import TypeBase +from models.enums import CreatorUserRole + + +@pytest.fixture +def file_session(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]: + """Bind runtime-owned sessions to SQLite with only the two file tables present.""" + tables = [TypeBase.metadata.tables[model.__tablename__] for model in (UploadFile, ToolFile)] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr(file_runtime.session_factory, "create_session", session_maker) + with session_maker() as session: + yield session + + +def _persist_upload_file( + session: Session, + *, + file_id: str = "upload-file-id", + key: str = "canonical-storage-key", + tenant_id: str = "tenant-id", + created_by: str = "end-user-id", +) -> UploadFile: + upload_file = UploadFile( + tenant_id=tenant_id, + storage_type=StorageType.LOCAL, + key=key, + name="diagram.png", + size=128, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.END_USER, + created_by=created_by, + created_at=datetime(2024, 1, 1, tzinfo=UTC), + used=False, + ) + upload_file.id = file_id + session.add(upload_file) + session.commit() + return upload_file + + +def _persist_tool_file( + session: Session, + *, + file_id: str = "tool-file-id", + key: str = "tool-storage-key", +) -> ToolFile: + tool_file = ToolFile( + user_id="end-user-id", + tenant_id="tenant-id", + conversation_id=None, + file_key=key, + mimetype="image/png", + name="diagram.png", + size=128, + ) + tool_file.id = file_id + session.add(tool_file) + session.commit() + return tool_file def _build_file( @@ -75,8 +141,9 @@ def test_resolve_file_url_requires_extension_for_tool_files() -> None: def test_resolve_file_url_uses_tool_signatures_for_tool_and_datasource_files( monkeypatch: pytest.MonkeyPatch, ) -> None: - sign_tool_file = MagicMock(return_value="https://signed.example.com/file") - monkeypatch.setattr(file_runtime, "sign_tool_file", sign_tool_file) + sign_tool_file_uri = MagicMock(return_value="/files/signed") + monkeypatch.setattr(file_runtime, "sign_tool_file_uri", sign_tool_file_uri) + monkeypatch.setattr(file_runtime.dify_config, "FILES_URL", "https://files.example.com") runtime = _build_runtime() tool_file = _build_file( @@ -90,9 +157,35 @@ def test_resolve_file_url_uses_tool_signatures_for_tool_and_datasource_files( extension=".png", ) - assert runtime.resolve_file_url(file=tool_file) == "https://signed.example.com/file" - assert runtime.resolve_file_url(file=datasource_file) == "https://signed.example.com/file" - assert sign_tool_file.call_count == 2 + assert runtime.resolve_file_url(file=tool_file) == "https://files.example.com/files/signed" + assert runtime.resolve_file_url(file=datasource_file) == "https://files.example.com/files/signed" + assert sign_tool_file_uri.call_count == 2 + + +def test_resolve_file_uri_keeps_dify_owned_file_origin_free(monkeypatch: pytest.MonkeyPatch) -> None: + sign_tool_file_uri = MagicMock(return_value="/files/tools/tool-file-id.png?sign=1") + monkeypatch.setattr(file_runtime, "sign_tool_file_uri", sign_tool_file_uri) + runtime = _build_runtime() + file = _build_file( + transfer_method=FileTransferMethod.TOOL_FILE, + reference=build_file_reference(record_id="tool-file-id"), + extension=".png", + ) + + assert runtime.resolve_file_uri(file=file) == "/files/tools/tool-file-id.png?sign=1" + + +def test_resolve_file_url_returns_relative_uri_when_files_url_is_empty(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(file_runtime, "sign_tool_file_uri", lambda **_: "/files/tools/tool-file-id.png?sign=1") + monkeypatch.setattr(file_runtime.dify_config, "FILES_URL", "") + runtime = _build_runtime() + file = _build_file( + transfer_method=FileTransferMethod.TOOL_FILE, + reference=build_file_reference(record_id="tool-file-id"), + extension=".png", + ) + + assert runtime.resolve_file_url(file=file, for_external=True) == "/files/tools/tool-file-id.png?sign=1" def test_resolve_upload_file_url_signs_internal_urls_and_supports_attachments( @@ -164,56 +257,37 @@ def test_verify_preview_signature_validates_signature_and_expiration(monkeypatch ) -def test_load_file_bytes_returns_bytes_and_rejects_non_bytes(monkeypatch: pytest.MonkeyPatch) -> None: +def test_load_file_bytes_returns_bytes_and_rejects_non_bytes( + monkeypatch: pytest.MonkeyPatch, file_session: Session +) -> None: runtime = _build_runtime() file = _build_file( transfer_method=FileTransferMethod.LOCAL_FILE, reference=build_file_reference(record_id="upload-file-id"), ) - session = MagicMock() - session.get.return_value = SimpleNamespace(key="canonical-storage-key") - - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - - monkeypatch.setattr(file_runtime.session_factory, "create_session", lambda: _SessionContext()) + _persist_upload_file(file_session) monkeypatch.setattr(file_runtime.storage, "load", lambda *args, **kwargs: b"image-bytes") assert runtime.load_file_bytes(file=file) == b"image-bytes" - session.get.assert_called_with(UploadFile, "upload-file-id") monkeypatch.setattr(file_runtime.storage, "load", lambda *args, **kwargs: "not-bytes") with pytest.raises(ValueError, match="is not a bytes object"): runtime.load_file_bytes(file=file) -def test_resolve_storage_key_ignores_encoded_reference_when_unscoped(monkeypatch: pytest.MonkeyPatch) -> None: +def test_resolve_storage_key_ignores_encoded_reference_when_unscoped(file_session: Session) -> None: runtime = _build_runtime() file = _build_file( transfer_method=FileTransferMethod.LOCAL_FILE, reference=build_file_reference(record_id="upload-file-id", storage_key="tampered-storage-key"), ) - session = MagicMock() - session.get.return_value = SimpleNamespace(key="canonical-storage-key") - - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - - monkeypatch.setattr(file_runtime.session_factory, "create_session", lambda: _SessionContext()) + _persist_upload_file(file_session) assert runtime._resolve_storage_key(file=file) == "canonical-storage-key" - session.get.assert_called_once_with(UploadFile, "upload-file-id") -def test_resolve_storage_key_uses_canonical_record_when_scope_is_bound(monkeypatch: pytest.MonkeyPatch) -> None: +def test_resolve_storage_key_uses_canonical_record_when_scope_is_bound(file_session: Session) -> None: + upload_file = _persist_upload_file(file_session) controller = MagicMock() controller.current_scope.return_value = FileAccessScope( tenant_id="tenant-id", @@ -221,28 +295,19 @@ def test_resolve_storage_key_uses_canonical_record_when_scope_is_bound(monkeypat user_from=UserFrom.END_USER, invoke_from=InvokeFrom.WEB_APP, ) - controller.get_upload_file.return_value = SimpleNamespace(key="canonical-storage-key") + controller.get_upload_file.return_value = upload_file runtime = DifyWorkflowFileRuntime(file_access_controller=controller) file = _build_file( transfer_method=FileTransferMethod.LOCAL_FILE, reference=build_file_reference(record_id="upload-file-id", storage_key="tampered-storage-key"), ) - session = MagicMock() - - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - - monkeypatch.setattr(file_runtime.session_factory, "create_session", lambda: _SessionContext()) - assert runtime._resolve_storage_key(file=file) == "canonical-storage-key" - controller.get_upload_file.assert_called_once_with(session=session, file_id="upload-file-id") + controller.get_upload_file.assert_called_once() + assert isinstance(controller.get_upload_file.call_args.kwargs["session"], Session) + assert controller.get_upload_file.call_args.kwargs["file_id"] == "upload-file-id" -def test_resolve_upload_file_url_rejects_unauthorized_scoped_access(monkeypatch: pytest.MonkeyPatch) -> None: +def test_resolve_upload_file_url_rejects_unauthorized_scoped_access(file_session: Session) -> None: controller = MagicMock() controller.current_scope.return_value = FileAccessScope( tenant_id="tenant-id", @@ -252,17 +317,6 @@ def test_resolve_upload_file_url_rejects_unauthorized_scoped_access(monkeypatch: ) controller.get_upload_file.return_value = None runtime = DifyWorkflowFileRuntime(file_access_controller=controller) - session = MagicMock() - - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - - monkeypatch.setattr(file_runtime.session_factory, "create_session", lambda: _SessionContext()) - with pytest.raises(ValueError, match="Upload file upload-file-id not found"): runtime.resolve_upload_file_url(upload_file_id="upload-file-id") @@ -276,7 +330,7 @@ def test_resolve_upload_file_url_rejects_unauthorized_scoped_access(monkeypatch: ], ) def test_resolve_storage_key_loads_database_records( - monkeypatch: pytest.MonkeyPatch, + file_session: Session, transfer_method: FileTransferMethod, record_id: str, expected_storage_key: str, @@ -287,25 +341,10 @@ def test_resolve_storage_key_loads_database_records( reference=build_file_reference(record_id=record_id), extension=".png", ) - session = MagicMock() - - def get(model_class, value): - if transfer_method in {FileTransferMethod.LOCAL_FILE, FileTransferMethod.DATASOURCE_FILE}: - assert model_class is UploadFile - return SimpleNamespace(key="upload-storage-key") - assert model_class is ToolFile - return SimpleNamespace(file_key="tool-storage-key") - - session.get.side_effect = get - - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - - monkeypatch.setattr(file_runtime.session_factory, "create_session", lambda: _SessionContext()) + if transfer_method in {FileTransferMethod.LOCAL_FILE, FileTransferMethod.DATASOURCE_FILE}: + _persist_upload_file(file_session, key="upload-storage-key") + else: + _persist_tool_file(file_session) assert runtime._resolve_storage_key(file=file) == expected_storage_key @@ -318,7 +357,7 @@ def test_resolve_storage_key_loads_database_records( ], ) def test_resolve_storage_key_raises_when_records_are_missing( - monkeypatch: pytest.MonkeyPatch, + file_session: Session, transfer_method: FileTransferMethod, expected_message: str, ) -> None: @@ -329,18 +368,6 @@ def test_resolve_storage_key_raises_when_records_are_missing( reference=build_file_reference(record_id=record_id), extension=".png", ) - session = MagicMock() - session.get.return_value = None - - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return False - - monkeypatch.setattr(file_runtime.session_factory, "create_session", lambda: _SessionContext()) - with pytest.raises(ValueError, match=expected_message): runtime._resolve_storage_key(file=file) diff --git a/api/tests/unit_tests/core/app/workflow/test_persistence_layer.py b/api/tests/unit_tests/core/app/workflow/test_persistence_layer.py index 31b9bb28940..bc67bfc12d5 100644 --- a/api/tests/unit_tests/core/app/workflow/test_persistence_layer.py +++ b/api/tests/unit_tests/core/app/workflow/test_persistence_layer.py @@ -9,7 +9,7 @@ from core.app.entities.app_invoke_entities import WorkflowAppGenerateEntity from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer from core.ops.ops_trace_manager import TraceTask, TraceTaskName from core.workflow.system_variables import SystemVariableKey, build_system_variables -from graphon.entities import WorkflowNodeExecution +from graphon.entities import WorkflowNodeExecution, WorkflowStartReason from graphon.entities.pause_reason import SchedulingPause from graphon.enums import ( BuiltinNodeTypes, @@ -32,6 +32,7 @@ from graphon.graph_events import ( NodeRunStartedEvent, NodeRunSucceededEvent, ) +from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.node_events import NodeRunResult from graphon.runtime import GraphRuntimeState, ReadOnlyGraphRuntimeStateWrapper, VariablePool @@ -39,14 +40,22 @@ from graphon.runtime import GraphRuntimeState, ReadOnlyGraphRuntimeStateWrapper, class _RepoRecorder: def __init__(self) -> None: self.saved: list[object] = [] + self.synchronously_saved: list[object] = [] self.saved_exec_data: list[object] = [] + self.loaded: list[object] = [] def save(self, entity): self.saved.append(entity) + def save_synchronously(self, entity): + self.synchronously_saved.append(entity) + def save_execution_data(self, entity): self.saved_exec_data.append(entity) + def get_by_workflow_execution(self, _workflow_execution_id): + return self.loaded + def _naive_utc_now() -> datetime: return datetime.now(UTC).replace(tzinfo=None) @@ -165,12 +174,45 @@ class TestWorkflowPersistenceLayer: assert exec_repo.saved + def test_resumption_restores_container_execution_before_terminal_event(self): + layer, _, node_repo, _ = _make_layer() + started_at = _naive_utc_now() + execution = WorkflowNodeExecution( + id="loop-exec", + workflow_id="workflow-id", + workflow_execution_id="run-id", + index=4, + node_id="loop", + node_type=BuiltinNodeTypes.LOOP, + title="Loop", + status=WorkflowNodeExecutionStatus.RUNNING, + created_at=started_at, + ) + node_repo.loaded = [execution] + + layer.on_event(GraphRunStartedEvent(reason=WorkflowStartReason.RESUMPTION)) + layer.on_event( + NodeRunSucceededEvent( + id=execution.id, + node_id=execution.node_id, + node_type=execution.node_type, + start_at=started_at, + node_run_result=NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED), + ) + ) + + assert execution.status == WorkflowNodeExecutionStatus.SUCCEEDED + assert layer._next_node_sequence() == 5 + def test_handle_graph_run_succeeded_updates_execution(self): layer, exec_repo, _, runtime_state = _make_layer() layer._handle_graph_run_started() - runtime_state.total_tokens = 3 - runtime_state.node_run_steps = 2 - runtime_state.outputs = {"out": "v"} + usage = LLMUsage.empty_usage() + usage.total_tokens = 3 + runtime_state.add_llm_usage(usage) + for _ in range(2): + runtime_state.increment_node_run_steps() + runtime_state.set_output("out", "v") layer._handle_graph_run_succeeded(GraphRunSucceededEvent(outputs={"ok": True})) @@ -182,8 +224,11 @@ class TestWorkflowPersistenceLayer: def test_handle_graph_run_partial_succeeded_updates_execution(self): layer, exec_repo, _, runtime_state = _make_layer() layer._handle_graph_run_started() - runtime_state.total_tokens = 5 - runtime_state.node_run_steps = 4 + usage = LLMUsage.empty_usage() + usage.total_tokens = 5 + runtime_state.add_llm_usage(usage) + for _ in range(4): + runtime_state.increment_node_run_steps() runtime_state._graph_execution = SimpleNamespace(exceptions_count=2) layer._handle_graph_run_partial_succeeded( @@ -289,8 +334,11 @@ class TestWorkflowPersistenceLayer: def test_handle_graph_run_paused_updates_outputs(self): layer, exec_repo, _, runtime_state = _make_layer() layer._handle_graph_run_started() - runtime_state.total_tokens = 7 - runtime_state.node_run_steps = 5 + usage = LLMUsage.empty_usage() + usage.total_tokens = 7 + runtime_state.add_llm_usage(usage) + for _ in range(5): + runtime_state.increment_node_run_steps() layer._handle_graph_run_paused(GraphRunPausedEvent(outputs={"pause": True})) @@ -331,6 +379,24 @@ class TestWorkflowPersistenceLayer: layer._handle_node_retry(retry_event) assert node_repo.saved_exec_data + def test_agent_v2_caller_row_is_saved_synchronously_before_node_run(self): + layer, _, node_repo, _ = _make_layer() + layer._handle_graph_run_started() + + layer._handle_node_started( + NodeRunStartedEvent( + id="agent-exec", + node_id="agent-node", + node_type=BuiltinNodeTypes.AGENT, + node_version="2", + node_title="Agent", + start_at=_naive_utc_now(), + ) + ) + + assert [execution.id for execution in node_repo.synchronously_saved] == ["agent-exec"] + assert node_repo.saved == [] + def test_retry_history_is_preserved_after_node_succeeds(self): layer, _, node_repo, _ = _make_layer() layer._handle_graph_run_started() @@ -501,7 +567,12 @@ class TestWorkflowPersistenceLayer: domain_execution = layer._node_execution_cache["exec"] domain_execution.inputs = {"old": True} - result = NodeRunResult(inputs={"new": True}, outputs={"out": 1}, process_data={"p": 1}, metadata={}) + result = NodeRunResult( + inputs={"new": True}, + outputs={"out": 1}, + process_data={"p": 1, "workflow_agent_binding_id": "workflow-binding-1"}, + metadata={}, + ) pause_event = NodeRunPauseRequestedEvent( id="exec", node_id="node", @@ -513,6 +584,38 @@ class TestWorkflowPersistenceLayer: assert domain_execution.status == WorkflowNodeExecutionStatus.PAUSED assert domain_execution.inputs == {"old": True} + assert domain_execution.process_data == {"workflow_agent_binding_id": "workflow-binding-1"} + + def test_handle_node_retry_preserves_workflow_agent_binding_identity(self): + layer, _, _, _ = _make_layer() + layer._handle_graph_run_started() + started_at = _naive_utc_now() + layer._handle_node_started( + NodeRunStartedEvent( + id="exec", + node_id="node", + node_type=BuiltinNodeTypes.AGENT, + node_title="Agent", + start_at=started_at, + ) + ) + + layer._handle_node_retry( + NodeRunRetryEvent( + id="exec", + node_id="node", + node_type=BuiltinNodeTypes.AGENT, + node_title="Agent", + start_at=started_at, + error="retry", + retry_index=1, + node_run_result=NodeRunResult( + process_data={"workflow_agent_binding_id": "workflow-binding-1"}, + ), + ) + ) + + assert layer._node_execution_cache["exec"].process_data["workflow_agent_binding_id"] == "workflow-binding-1" def test_get_node_execution_raises_for_missing(self): layer, _, _, _ = _make_layer() diff --git a/api/tests/unit_tests/core/callback_handler/test_index_tool_callback_handler.py b/api/tests/unit_tests/core/callback_handler/test_index_tool_callback_handler.py index 62c4ae9d411..2808da27236 100644 --- a/api/tests/unit_tests/core/callback_handler/test_index_tool_callback_handler.py +++ b/api/tests/unit_tests/core/callback_handler/test_index_tool_callback_handler.py @@ -1,10 +1,38 @@ +from collections.abc import Iterator +from dataclasses import dataclass +from uuid import uuid4 + import pytest from pytest_mock import MockerFixture +from sqlalchemy import select +from sqlalchemy.engine import Engine +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session, SessionTransaction +import core.callback_handler.index_tool_callback_handler as callback_module from core.app.entities.app_invoke_entities import InvokeFrom -from core.callback_handler.index_tool_callback_handler import ( - DatasetIndexToolCallbackHandler, -) +from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler +from core.rag.index_processor.constant.index_type import IndexStructureType +from core.rag.models.document import Document +from models.dataset import ChildChunk, DatasetQuery, DocumentSegment +from models.dataset import Document as DatasetDocument +from models.enums import CreatorUserRole, DatasetQuerySource, DataSourceType, DocumentCreatedFrom + + +class _DatabaseBinding: + engine: Engine + + def __init__(self, engine: Engine) -> None: + self.engine = engine + + +@dataclass(frozen=True) +class _CallerSessionBoundary: + """Caller session state that callback-owned transactions must not disturb.""" + + session: Session + transaction: SessionTransaction + pending_query: DatasetQuery @pytest.fixture @@ -13,16 +41,75 @@ def mock_queue_manager(mocker: MockerFixture): @pytest.fixture -def handler(mock_queue_manager, mocker: MockerFixture): - mocker.patch( - "core.callback_handler.index_tool_callback_handler.db", - ) +def handler(mock_queue_manager, sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(callback_module, "db", _DatabaseBinding(sqlite_engine)) return DatasetIndexToolCallbackHandler( queue_manager=mock_queue_manager, - app_id="app-1", - message_id="msg-1", - user_id="user-1", - invoke_from=mocker.Mock(), + app_id=str(uuid4()), + message_id=str(uuid4()), + user_id=str(uuid4()), + invoke_from=InvokeFrom.DEBUGGER, + ) + + +@pytest.fixture +def caller_session_boundary(sqlite_engine: Engine) -> Iterator[_CallerSessionBoundary]: + """Keep an unflushed caller transaction open to prove callback isolation.""" + + with Session(sqlite_engine, expire_on_commit=False) as session: + pending_query = DatasetQuery( + dataset_id=str(uuid4()), + content="caller-owned pending query", + source=DatasetQuerySource.APP, + source_app_id=str(uuid4()), + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + ) + session.add(pending_query) + transaction = session.get_transaction() + assert transaction is not None + + yield _CallerSessionBoundary( + session=session, + transaction=transaction, + pending_query=pending_query, + ) + + assert session.get_transaction() is transaction + assert list(session.new) == [pending_query] + assert not session.dirty + assert not session.deleted + + with Session(sqlite_engine) as observer_session: + assert observer_session.get(DatasetQuery, pending_query.id) is None + + +def _dataset_document(*, doc_form: IndexStructureType) -> DatasetDocument: + return DatasetDocument( + tenant_id=str(uuid4()), + dataset_id=str(uuid4()), + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch", + name="Document", + created_from=DocumentCreatedFrom.API, + created_by=str(uuid4()), + doc_form=doc_form, + ) + + +def _segment(document: DatasetDocument, *, index_node_id: str) -> DocumentSegment: + return DocumentSegment( + tenant_id=document.tenant_id, + dataset_id=document.dataset_id, + document_id=document.id, + position=1, + content="content", + word_count=1, + tokens=1, + created_by=document.created_by, + index_node_id=index_node_id, + hit_count=0, ) @@ -35,53 +122,40 @@ class TestOnQuery: (InvokeFrom.WEB_APP, "end_user"), ], ) - def test_on_query_success_roles(self, mocker: MockerFixture, mock_queue_manager, invoke_from, expected_role): - # Arrange — the caller passes a session, but our fix uses an independent one - caller_session = mocker.Mock() - - independent_session = mocker.MagicMock() - mock_session_factory = mocker.MagicMock() - mock_session_factory.begin.return_value.__enter__ = mocker.MagicMock(return_value=independent_session) - mock_session_factory.begin.return_value.__exit__ = mocker.MagicMock(return_value=False) - mocker.patch( - "core.callback_handler.index_tool_callback_handler.sessionmaker", - return_value=mock_session_factory, - ) - mocker.patch("core.callback_handler.index_tool_callback_handler.db") - + def test_on_query_success_roles( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + caller_session_boundary: _CallerSessionBoundary, + mock_queue_manager, + invoke_from: InvokeFrom, + expected_role: str, + ) -> None: + monkeypatch.setattr(callback_module, "db", _DatabaseBinding(sqlite_engine)) handler = DatasetIndexToolCallbackHandler( queue_manager=mock_queue_manager, - app_id="app-1", - message_id="msg-1", - user_id="user-1", - invoke_from=mocker.Mock(), + app_id=str(uuid4()), + message_id=str(uuid4()), + user_id=str(uuid4()), + invoke_from=invoke_from, ) - handler._invoke_from = invoke_from + handler.on_query("test query", str(uuid4()), caller_session_boundary.session) - # Act — pass caller_session as required by signature - handler.on_query("test query", "dataset-1", caller_session) - - # Assert — independent session used, not the caller's session - independent_session.add.assert_called_once() - dataset_query = independent_session.add.call_args.args[0] + dataset_query = sqlite_session.scalar(select(DatasetQuery)) + assert dataset_query is not None assert dataset_query.created_by_role == expected_role - caller_session.add.assert_not_called() - caller_session.commit.assert_not_called() - - def test_on_query_none_values(self, mocker: MockerFixture, mock_queue_manager): - caller_session = mocker.Mock() - - independent_session = mocker.MagicMock() - mock_session_factory = mocker.MagicMock() - mock_session_factory.begin.return_value.__enter__ = mocker.MagicMock(return_value=independent_session) - mock_session_factory.begin.return_value.__exit__ = mocker.MagicMock(return_value=False) - mocker.patch( - "core.callback_handler.index_tool_callback_handler.sessionmaker", - return_value=mock_session_factory, - ) - mocker.patch("core.callback_handler.index_tool_callback_handler.db") + def test_on_query_none_values_roll_back_independent_transaction( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + caller_session_boundary: _CallerSessionBoundary, + mock_queue_manager, + ) -> None: + monkeypatch.setattr(callback_module, "db", _DatabaseBinding(sqlite_engine)) handler = DatasetIndexToolCallbackHandler( queue_manager=mock_queue_manager, app_id=None, @@ -90,136 +164,109 @@ class TestOnQuery: invoke_from=None, ) - handler.on_query(None, None, caller_session) + with pytest.raises(IntegrityError): + handler.on_query(None, None, caller_session_boundary.session) # type: ignore[arg-type] - independent_session.add.assert_called_once() - caller_session.add.assert_not_called() + assert sqlite_session.scalar(select(DatasetQuery)) is None class TestOnToolEnd: - def test_on_tool_end_no_metadata(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): - caller_session = mocker.Mock() + def test_on_tool_end_no_metadata( + self, + handler: DatasetIndexToolCallbackHandler, + caller_session_boundary: _CallerSessionBoundary, + ) -> None: + document = Document.model_construct(page_content="content", metadata=None, provider="dify") - independent_session = mocker.MagicMock() - mocker.patch( - "core.callback_handler.index_tool_callback_handler.Session", - return_value=independent_session, - ) - independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) - independent_session.__exit__ = mocker.MagicMock(return_value=False) - - document = mocker.Mock() - document.metadata = None - - handler.on_tool_end([document], caller_session) - - independent_session.commit.assert_called_once() - independent_session.execute.assert_not_called() - caller_session.commit.assert_not_called() + handler.on_tool_end([document], caller_session_boundary.session) def test_on_tool_end_dataset_document_not_found( - self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture - ): - caller_session = mocker.Mock() - - independent_session = mocker.MagicMock() - mocker.patch( - "core.callback_handler.index_tool_callback_handler.Session", - return_value=independent_session, + self, + handler: DatasetIndexToolCallbackHandler, + sqlite_session: Session, + caller_session_boundary: _CallerSessionBoundary, + ) -> None: + document = Document( + page_content="content", + metadata={"document_id": str(uuid4()), "doc_id": "node-1"}, ) - independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) - independent_session.__exit__ = mocker.MagicMock(return_value=False) - independent_session.scalar.return_value = None - document = mocker.Mock() - document.metadata = {"document_id": "doc-1", "doc_id": "node-1"} + handler.on_tool_end([document], caller_session_boundary.session) - handler.on_tool_end([document], caller_session) - - independent_session.scalar.assert_called_once() - caller_session.scalar.assert_not_called() + assert sqlite_session.scalar(select(DatasetDocument)) is None def test_on_tool_end_parent_child_index_with_child( - self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture - ): - caller_session = mocker.Mock() - - independent_session = mocker.MagicMock() - mocker.patch( - "core.callback_handler.index_tool_callback_handler.Session", - return_value=independent_session, + self, + handler: DatasetIndexToolCallbackHandler, + sqlite_session: Session, + caller_session_boundary: _CallerSessionBoundary, + ) -> None: + dataset_document = _dataset_document(doc_form=IndexStructureType.PARENT_CHILD_INDEX) + sqlite_session.add(dataset_document) + sqlite_session.flush() + segment = _segment(dataset_document, index_node_id="parent-node") + sqlite_session.add(segment) + sqlite_session.flush() + child = ChildChunk( + tenant_id=dataset_document.tenant_id, + dataset_id=dataset_document.dataset_id, + document_id=dataset_document.id, + segment_id=segment.id, + position=1, + content="child", + word_count=1, + created_by=dataset_document.created_by, + index_node_id="child-node", ) - independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) - independent_session.__exit__ = mocker.MagicMock(return_value=False) - - mock_dataset_doc = mocker.Mock() - from core.callback_handler.index_tool_callback_handler import IndexStructureType - - mock_dataset_doc.doc_form = IndexStructureType.PARENT_CHILD_INDEX - mock_dataset_doc.dataset_id = "dataset-1" - mock_dataset_doc.id = "doc-1" - - mock_child_chunk = mocker.Mock() - mock_child_chunk.segment_id = "segment-1" - - independent_session.scalar.side_effect = [mock_dataset_doc, mock_child_chunk] - - document = mocker.Mock() - document.metadata = {"document_id": "doc-1", "doc_id": "node-1"} - - handler.on_tool_end([document], caller_session) - - independent_session.execute.assert_called_once() - independent_session.commit.assert_called_once() - caller_session.execute.assert_not_called() - - def test_on_tool_end_non_parent_child_index(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): - caller_session = mocker.Mock() - - independent_session = mocker.MagicMock() - mocker.patch( - "core.callback_handler.index_tool_callback_handler.Session", - return_value=independent_session, + sqlite_session.add(child) + sqlite_session.commit() + document = Document( + page_content="content", + metadata={"document_id": dataset_document.id, "doc_id": child.index_node_id}, ) - independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) - independent_session.__exit__ = mocker.MagicMock(return_value=False) - mock_dataset_doc = mocker.Mock() - mock_dataset_doc.doc_form = "OTHER" + handler.on_tool_end([document], caller_session_boundary.session) - independent_session.scalar.return_value = mock_dataset_doc + sqlite_session.expire_all() + assert sqlite_session.get(DocumentSegment, segment.id).hit_count == 1 # type: ignore[union-attr] - document = mocker.Mock() - document.metadata = { - "document_id": "doc-1", - "doc_id": "node-1", - "dataset_id": "dataset-1", - } - - handler.on_tool_end([document], caller_session) - - independent_session.execute.assert_called_once() - independent_session.commit.assert_called_once() - caller_session.execute.assert_not_called() - - def test_on_tool_end_empty_documents(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): - caller_session = mocker.Mock() - - independent_session = mocker.MagicMock() - mocker.patch( - "core.callback_handler.index_tool_callback_handler.Session", - return_value=independent_session, + def test_on_tool_end_non_parent_child_index( + self, + handler: DatasetIndexToolCallbackHandler, + sqlite_session: Session, + caller_session_boundary: _CallerSessionBoundary, + ) -> None: + dataset_document = _dataset_document(doc_form=IndexStructureType.PARAGRAPH_INDEX) + sqlite_session.add(dataset_document) + sqlite_session.flush() + segment = _segment(dataset_document, index_node_id="node-1") + sqlite_session.add(segment) + sqlite_session.commit() + document = Document( + page_content="content", + metadata={ + "document_id": dataset_document.id, + "doc_id": segment.index_node_id, + "dataset_id": dataset_document.dataset_id, + }, ) - independent_session.__enter__ = mocker.MagicMock(return_value=independent_session) - independent_session.__exit__ = mocker.MagicMock(return_value=False) - handler.on_tool_end([], caller_session) + handler.on_tool_end([document], caller_session_boundary.session) + + sqlite_session.expire_all() + assert sqlite_session.get(DocumentSegment, segment.id).hit_count == 1 # type: ignore[union-attr] + + def test_on_tool_end_empty_documents( + self, + handler: DatasetIndexToolCallbackHandler, + caller_session_boundary: _CallerSessionBoundary, + ) -> None: + handler.on_tool_end([], caller_session_boundary.session) class TestReturnRetrieverResourceInfo: def test_publish_called(self, handler: DatasetIndexToolCallbackHandler, mock_queue_manager, mocker: MockerFixture): mock_event = mocker.patch("core.callback_handler.index_tool_callback_handler.QueueRetrieverResourcesEvent") - resources = [mocker.Mock()] handler.return_retriever_resource_info(resources) diff --git a/api/tests/unit_tests/core/datasource/test_datasource_file_manager.py b/api/tests/unit_tests/core/datasource/test_datasource_file_manager.py index 20730e2e8a9..3bbc8db1b66 100644 --- a/api/tests/unit_tests/core/datasource/test_datasource_file_manager.py +++ b/api/tests/unit_tests/core/datasource/test_datasource_file_manager.py @@ -1,13 +1,62 @@ +from datetime import UTC, datetime from unittest.mock import MagicMock, patch import httpx import pytest +from sqlalchemy.orm import Session from core.datasource.datasource_file_manager import DatasourceFileManager +from extensions.storage.storage_type import StorageType +from models.enums import CreatorUserRole from models.model import MessageFile, UploadFile from models.tools import ToolFile +def _upload_file(id: str, *, key: str, mime_type: str) -> UploadFile: + upload_file = UploadFile( + tenant_id="tenant-1", + storage_type=StorageType.LOCAL, + key=key, + name="file.png", + size=4, + extension=".png", + mime_type=mime_type, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + created_at=datetime.now(UTC).replace(tzinfo=None), + used=False, + ) + upload_file.id = id + return upload_file + + +def _tool_file(id: str, *, key: str = "tool_key", mimetype: str = "image/png") -> ToolFile: + tool_file = ToolFile( + tenant_id="tenant-1", + user_id="user-1", + conversation_id=None, + file_key=key, + mimetype=mimetype, + name="tool.png", + size=4, + ) + tool_file.id = id + return tool_file + + +def _message_file(id: str, *, url: str | None) -> MessageFile: + message_file = MessageFile( + message_id="message-1", + type="image", + transfer_method="remote_url", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + url=url, + ) + message_file.id = id + return message_file + + class TestDatasourceFileManager: @patch("core.datasource.datasource_file_manager.time.time") @patch("core.datasource.datasource_file_manager.os.urandom") @@ -31,11 +80,10 @@ class TestDatasourceFileManager: assert f"nonce={mock_urandom.return_value.hex()}" in signed_url assert "sign=" in signed_url - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") @patch("core.datasource.datasource_file_manager.dify_config") - def test_create_file_by_raw(self, mock_config, mock_uuid, mock_storage, mock_db): + def test_create_file_by_raw(self, mock_config, mock_uuid, mock_storage, sqlite_session: Session): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_config.STORAGE_TYPE = "local" @@ -63,14 +111,14 @@ class TestDatasourceFileManager: assert upload_file.key == f"datasources/{tenant_id}/unique_hex.png" mock_storage.save.assert_called_once_with(upload_file.key, file_binary) - mock_db.session.add.assert_called_once() - mock_db.session.commit.assert_called_once() + persisted_file = sqlite_session.get(UploadFile, upload_file.id) + assert persisted_file is not None + assert persisted_file.key == upload_file.key - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") @patch("core.datasource.datasource_file_manager.dify_config") - def test_create_file_by_raw_filename_no_extension(self, mock_config, mock_uuid, mock_storage, mock_db): + def test_create_file_by_raw_filename_no_extension(self, mock_config, mock_uuid, mock_storage): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_config.STORAGE_TYPE = "local" @@ -93,12 +141,11 @@ class TestDatasourceFileManager: # Verify assert upload_file.name == "test.png" # Should append extension - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") @patch("core.datasource.datasource_file_manager.dify_config") @patch("core.datasource.datasource_file_manager.guess_extension") - def test_create_file_by_raw_unknown_extension(self, mock_guess_ext, mock_config, mock_uuid, mock_storage, mock_db): + def test_create_file_by_raw_unknown_extension(self, mock_guess_ext, mock_config, mock_uuid, mock_storage): # Setup mock_guess_ext.return_value = None # Cannot guess mock_uuid.return_value = MagicMock(hex="unique_hex") @@ -117,11 +164,10 @@ class TestDatasourceFileManager: assert upload_file.extension == ".bin" assert upload_file.name == "unique_hex.bin" - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") @patch("core.datasource.datasource_file_manager.dify_config") - def test_create_file_by_raw_no_filename(self, mock_config, mock_uuid, mock_storage, mock_db): + def test_create_file_by_raw_no_filename(self, mock_config, mock_uuid, mock_storage): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_config.STORAGE_TYPE = "local" @@ -140,10 +186,9 @@ class TestDatasourceFileManager: assert upload_file.extension == ".pdf" @patch("core.datasource.datasource_file_manager.remote_fetcher") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") - def test_create_file_by_url_mimetype_from_guess(self, mock_uuid, mock_storage, mock_db, mock_ssrf): + def test_create_file_by_url_mimetype_from_guess(self, mock_uuid, mock_storage, mock_ssrf): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_response = MagicMock() @@ -153,17 +198,18 @@ class TestDatasourceFileManager: # Execute tool_file = DatasourceFileManager.create_file_by_url( - user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/photo.png" + user_id="user_123", + tenant_id="tenant_456", + file_url="https://example.com/photo.png", ) # Verify assert tool_file.mimetype == "image/png" # Guessed from .png in URL @patch("core.datasource.datasource_file_manager.remote_fetcher") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") - def test_create_file_by_url_mimetype_default(self, mock_uuid, mock_storage, mock_db, mock_ssrf): + def test_create_file_by_url_mimetype_default(self, mock_uuid, mock_storage, mock_ssrf): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_response = MagicMock() @@ -182,10 +228,9 @@ class TestDatasourceFileManager: assert tool_file.mimetype == "application/octet-stream" @patch("core.datasource.datasource_file_manager.remote_fetcher") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") @patch("core.datasource.datasource_file_manager.uuid4") - def test_create_file_by_url_success(self, mock_uuid, mock_storage, mock_db, mock_ssrf): + def test_create_file_by_url_success(self, mock_uuid, mock_storage, mock_ssrf): # Setup mock_uuid.return_value = MagicMock(hex="unique_hex") mock_response = MagicMock() @@ -195,7 +240,9 @@ class TestDatasourceFileManager: # Execute tool_file = DatasourceFileManager.create_file_by_url( - user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/photo.jpg" + user_id="user_123", + tenant_id="tenant_456", + file_url="https://example.com/photo.jpg", ) # Verify @@ -215,106 +262,59 @@ class TestDatasourceFileManager: user_id="user_123", tenant_id="tenant_456", file_url="https://example.com/large.file" ) - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_binary(self, mock_storage, mock_db): - # Setup - mock_upload_file = MagicMock(spec=UploadFile) - mock_upload_file.key = "some_key" - mock_upload_file.mime_type = "image/png" - - mock_db.session.get.return_value = mock_upload_file + def test_get_file_binary(self, mock_storage, sqlite_session: Session): + sqlite_session.add(_upload_file("file_id", key="some_key", mime_type="image/png")) + sqlite_session.commit() mock_storage.load_once.return_value = b"file content" - # Execute result = DatasourceFileManager.get_file_binary("file_id") # Verify assert result == (b"file content", "image/png") - # Case: Not found - mock_db.session.get.return_value = None assert DatasourceFileManager.get_file_binary("unknown") is None - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_binary_by_message_file_id(self, mock_storage, mock_db): - # Setup - mock_message_file = MagicMock(spec=MessageFile) - mock_message_file.url = "http://localhost/files/tools/tool_id.png" - - mock_tool_file = MagicMock(spec=ToolFile) - mock_tool_file.file_key = "tool_key" - mock_tool_file.mimetype = "image/png" - - def mock_get(model, id): - if model == MessageFile: - return mock_message_file - elif model == ToolFile: - return mock_tool_file - return None - - mock_db.session.get.side_effect = mock_get + def test_get_file_binary_by_message_file_id(self, mock_storage, sqlite_session: Session): + sqlite_session.add_all( + [ + _message_file("msg_file_id", url="http://localhost/files/tools/tool_id.png"), + _tool_file("tool_id"), + ] + ) + sqlite_session.commit() mock_storage.load_once.return_value = b"tool content" - # Execute result = DatasourceFileManager.get_file_binary_by_message_file_id("msg_file_id") # Verify assert result == (b"tool content", "image/png") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_binary_by_message_file_id_with_extension(self, mock_storage, mock_db): - # Test that it correctly parses tool_id even with extension in URL - mock_message_file = MagicMock(spec=MessageFile) - mock_message_file.url = "http://localhost/files/tools/abcdef.png" - - mock_tool_file = MagicMock(spec=ToolFile) - mock_tool_file.id = "abcdef" - mock_tool_file.file_key = "tk" - mock_tool_file.mimetype = "image/png" - - def mock_get(model, id): - if model == MessageFile: - return mock_message_file - return mock_tool_file - - mock_db.session.get.side_effect = mock_get + def test_get_file_binary_by_message_file_id_with_extension(self, mock_storage, sqlite_session: Session): + sqlite_session.add_all( + [_message_file("m", url="http://localhost/files/tools/abcdef.png"), _tool_file("abcdef", key="tk")] + ) + sqlite_session.commit() mock_storage.load_once.return_value = b"bits" result = DatasourceFileManager.get_file_binary_by_message_file_id("m") assert result == (b"bits", "image/png") - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_binary_by_message_file_id_failures(self, mock_storage, mock_db): - # Case 1: Message file not found - mock_db.session.get.return_value = None + def test_get_file_binary_by_message_file_id_failures(self, mock_storage, sqlite_session: Session): assert DatasourceFileManager.get_file_binary_by_message_file_id("none") is None - # Case 2: Message file found but tool file not found - mock_message_file = MagicMock(spec=MessageFile) - mock_message_file.url = None - - def mock_get_v2(model, id): - if model == MessageFile: - return mock_message_file - return None - - mock_db.session.get.side_effect = mock_get_v2 + sqlite_session.add(_message_file("msg_id", url=None)) + sqlite_session.commit() assert DatasourceFileManager.get_file_binary_by_message_file_id("msg_id") is None - @patch("core.datasource.datasource_file_manager.db") @patch("core.datasource.datasource_file_manager.storage") - def test_get_file_generator_by_upload_file_id(self, mock_storage, mock_db): - # Setup - mock_upload_file = MagicMock(spec=UploadFile) - mock_upload_file.key = "upload_key" - mock_upload_file.mime_type = "text/plain" - - mock_db.session.get.return_value = mock_upload_file + def test_get_file_generator_by_upload_file_id(self, mock_storage, sqlite_session: Session): + sqlite_session.add(_upload_file("upload_id", key="upload_key", mime_type="text/plain")) + sqlite_session.commit() mock_storage.load_stream.return_value = iter([b"chunk1", b"chunk2"]) @@ -325,8 +325,6 @@ class TestDatasourceFileManager: assert mimetype == "text/plain" assert list(stream) == [b"chunk1", b"chunk2"] - # Case: Not found - mock_db.session.get.return_value = None stream, mimetype = DatasourceFileManager.get_file_generator_by_upload_file_id("none") assert stream is None assert mimetype is None diff --git a/api/tests/unit_tests/core/datasource/test_datasource_manager.py b/api/tests/unit_tests/core/datasource/test_datasource_manager.py index baf51489dfb..5308a4c1922 100644 --- a/api/tests/unit_tests/core/datasource/test_datasource_manager.py +++ b/api/tests/unit_tests/core/datasource/test_datasource_manager.py @@ -1,10 +1,13 @@ import types -from collections.abc import Generator +from collections.abc import Generator, Iterator import pytest from pytest_mock import MockerFixture +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, sessionmaker from contexts.wrapper import RecyclableContextVar +from core.datasource import datasource_manager as datasource_manager_module from core.datasource.datasource_manager import DatasourceManager from core.datasource.entities.datasource_entities import DatasourceMessage, DatasourceProviderType from core.datasource.errors import DatasourceProviderNotFoundError @@ -12,6 +15,34 @@ from core.workflow.file_reference import parse_file_reference from graphon.enums import WorkflowNodeExecutionStatus from graphon.file import File, FileTransferMethod, FileType from graphon.node_events import StreamChunkEvent, StreamCompletedEvent +from models.base import TypeBase +from models.tools import ToolFile + + +@pytest.fixture +def tool_file_session(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]: + """Bind datasource-owned lookups to a SQLite ToolFile table.""" + TypeBase.metadata.create_all(sqlite_engine, tables=[TypeBase.metadata.tables[ToolFile.__tablename__]]) + session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr(datasource_manager_module.session_factory, "create_session", session_maker) + with session_maker() as session: + yield session + + +def _persist_tool_file(session: Session, *, file_id: str, tenant_id: str) -> ToolFile: + tool_file = ToolFile( + user_id="user-1", + tenant_id=tenant_id, + conversation_id=None, + file_key="files/image.png", + mimetype="image/png", + name="image.png", + size=10, + ) + tool_file.id = file_id + session.add(tool_file) + session.commit() + return tool_file def _gen_messages_text_only(text: str) -> Generator[DatasourceMessage, None, None]: @@ -373,7 +404,8 @@ def test_stream_node_events_emits_events_online_document(mocker: MockerFixture): assert events[-1].node_run_result.status == WorkflowNodeExecutionStatus.SUCCEEDED -def test_stream_node_events_builds_file_and_variables_from_messages(mocker: MockerFixture): +def test_stream_node_events_builds_file_and_variables_from_messages(mocker: MockerFixture, tool_file_session: Session): + _persist_tool_file(tool_file_session, file_id="tool_file_1", tenant_id="t1") mocker.patch.object(DatasourceManager, "stream_online_results", return_value=_gen_messages_text_only("ignored")) def _transformed(**_kwargs): @@ -418,19 +450,6 @@ def test_stream_node_events_builds_file_and_variables_from_messages(mocker: Mock side_effect=_transformed, ) - fake_tool_file = types.SimpleNamespace(mimetype="image/png") - - class _Session: - def __enter__(self): - return self - - def __exit__(self, *exc): - return False - - def scalar(self, _stmt): - return fake_tool_file - - mocker.patch("core.datasource.datasource_manager.session_factory.create_session", return_value=_Session()) mocker.patch("core.datasource.datasource_manager.get_file_type_by_mime_type", return_value=FileType.IMAGE) built = File( file_type=FileType.IMAGE, @@ -481,7 +500,8 @@ def test_stream_node_events_builds_file_and_variables_from_messages(mocker: Mock assert events[-1].node_run_result.outputs["x"] == 1 -def test_stream_node_events_raises_when_toolfile_missing(mocker: MockerFixture): +def test_stream_node_events_raises_when_toolfile_missing(mocker: MockerFixture, tool_file_session: Session): + _persist_tool_file(tool_file_session, file_id="missing", tenant_id="other-tenant") mocker.patch.object(DatasourceManager, "stream_online_results", return_value=_gen_messages_text_only("ignored")) def _transformed(**_kwargs): @@ -496,18 +516,6 @@ def test_stream_node_events_raises_when_toolfile_missing(mocker: MockerFixture): side_effect=_transformed, ) - class _Session: - def __enter__(self): - return self - - def __exit__(self, *exc): - return False - - def scalar(self, _stmt): - return None - - mocker.patch("core.datasource.datasource_manager.session_factory.create_session", return_value=_Session()) - with pytest.raises(ValueError, match="ToolFile not found for file_id=missing, tenant_id=t1"): list( DatasourceManager.stream_node_events( diff --git a/api/tests/unit_tests/core/datasource/test_notion_provider.py b/api/tests/unit_tests/core/datasource/test_notion_provider.py index ecbd9691e98..d68187cfb34 100644 --- a/api/tests/unit_tests/core/datasource/test_notion_provider.py +++ b/api/tests/unit_tests/core/datasource/test_notion_provider.py @@ -14,18 +14,64 @@ Tests follow the Arrange-Act-Assert pattern for clarity. """ import json +from collections.abc import Iterator +from dataclasses import dataclass from typing import Any from unittest.mock import Mock, patch +from uuid import uuid4 import httpx import pytest +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from core.datasource.entities.datasource_entities import DatasourceProviderType from core.datasource.online_document.online_document_provider import ( OnlineDocumentDatasourcePluginProviderController, ) +from core.rag.extractor import notion_extractor as notion_extractor_module from core.rag.extractor.notion_extractor import NotionExtractor from core.rag.models.document import Document +from models.base import TypeBase +from models.dataset import Document as DocumentModel +from models.enums import DataSourceType, DocumentCreatedFrom + + +@dataclass(frozen=True) +class _Database: + """Expose the real SQLite session used by the extractor update.""" + + session: Session + + +@pytest.fixture +def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[_Database]: + """Bind a real session for Notion document metadata persistence.""" + + TypeBase.metadata.create_all(sqlite_engine, tables=[DocumentModel.__table__]) + with Session(sqlite_engine, expire_on_commit=False) as session: + database = _Database(session) + monkeypatch.setattr(notion_extractor_module, "db", database) + yield database + + +@pytest.fixture +def persisted_document(database: _Database) -> DocumentModel: + document = DocumentModel( + id=str(uuid4()), + tenant_id=str(uuid4()), + dataset_id=str(uuid4()), + position=1, + data_source_type=DataSourceType.NOTION_IMPORT, + data_source_info=json.dumps({"last_edited_time": "2024-01-01T00:00:00.000Z"}), + batch="batch", + name="Notion page", + created_from=DocumentCreatedFrom.WEB, + created_by=str(uuid4()), + ) + database.session.add(document) + database.session.commit() + return document class TestNotionExtractorAuthentication: @@ -763,9 +809,14 @@ class TestNotionExtractorLastEditedTime: call_args = mock_request.call_args assert "databases/database-789" in call_args[0][1] - @patch("core.rag.extractor.notion_extractor.db") @patch("httpx.request") - def test_update_last_edited_time(self, mock_request, mock_db, extractor_page, mock_document_model): + def test_update_last_edited_time( + self, + mock_request: Mock, + extractor_page: NotionExtractor, + database: _Database, + persisted_document: DocumentModel, + ): """Test updating document model with last edited time.""" # Arrange mock_response = Mock() @@ -777,11 +828,11 @@ class TestNotionExtractorLastEditedTime: mock_request.return_value = mock_response # Act - extractor_page.update_last_edited_time(mock_document_model) + extractor_page.update_last_edited_time(persisted_document) # Assert - assert mock_document_model.data_source_info_dict["last_edited_time"] == "2024-11-27T18:00:00.000Z" - mock_db.session.commit.assert_called_once() + database.session.expire(persisted_document) + assert persisted_document.data_source_info_dict["last_edited_time"] == "2024-11-27T18:00:00.000Z" def test_update_last_edited_time_no_document(self, extractor_page): """Test update_last_edited_time with None document model.""" @@ -807,9 +858,10 @@ class TestNotionExtractorIntegration: mock_doc.data_source_info_dict = {"last_edited_time": "2024-01-01T00:00:00.000Z"} return mock_doc - @patch("core.rag.extractor.notion_extractor.db") @patch("httpx.request") - def test_extract_page_complete_workflow(self, mock_request, mock_db, mock_document_model): + def test_extract_page_complete_workflow( + self, mock_request: Mock, database: _Database, persisted_document: DocumentModel + ): """Test complete page extraction workflow.""" # Arrange extractor = NotionExtractor( @@ -818,7 +870,7 @@ class TestNotionExtractorIntegration: notion_page_type="page", tenant_id="tenant-789", notion_access_token="test-token", - document_model=mock_document_model, + document_model=persisted_document, ) # Mock last edited time request @@ -869,11 +921,18 @@ class TestNotionExtractorIntegration: assert isinstance(documents[0], Document) assert "# Test Page" in documents[0].page_content assert "Test content" in documents[0].page_content + database.session.expire(persisted_document) + assert persisted_document.data_source_info_dict["last_edited_time"] == "2024-11-27T20:00:00.000Z" - @patch("core.rag.extractor.notion_extractor.db") @patch("httpx.post") @patch("httpx.request") - def test_extract_database_complete_workflow(self, mock_request, mock_post, mock_db, mock_document_model): + def test_extract_database_complete_workflow( + self, + mock_request: Mock, + mock_post: Mock, + database: _Database, + persisted_document: DocumentModel, + ): """Test complete database extraction workflow.""" # Arrange extractor = NotionExtractor( @@ -882,7 +941,7 @@ class TestNotionExtractorIntegration: notion_page_type="database", tenant_id="tenant-789", notion_access_token="test-token", - document_model=mock_document_model, + document_model=persisted_document, ) # Mock last edited time request @@ -921,6 +980,8 @@ class TestNotionExtractorIntegration: assert isinstance(documents[0], Document) assert "Name:Item 1" in documents[0].page_content assert "Status:Active" in documents[0].page_content + database.session.expire(persisted_document) + assert persisted_document.data_source_info_dict["last_edited_time"] == "2024-11-27T20:00:00.000Z" def test_extract_invalid_page_type(self): """Test extract with invalid page type.""" diff --git a/api/tests/unit_tests/core/datasource/utils/test_message_transformer.py b/api/tests/unit_tests/core/datasource/utils/test_message_transformer.py index 0fca43cd0b1..30ecb5b4003 100644 --- a/api/tests/unit_tests/core/datasource/utils/test_message_transformer.py +++ b/api/tests/unit_tests/core/datasource/utils/test_message_transformer.py @@ -40,9 +40,14 @@ class TestDatasourceFileMessageTransformer: def test_transform_image_message_success(self, mock_guess_ext, mock_tool_file_manager_cls): # Setup mock_manager = mock_tool_file_manager_cls.return_value - mock_tool_file = MagicMock(spec=ToolFile) + mock_tool_file = ToolFile( + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + file_key="test-key", + mimetype="image/png", + ) mock_tool_file.id = "file_id_123" - mock_tool_file.mimetype = "image/png" mock_manager.create_file_by_url.return_value = mock_tool_file mock_guess_ext.return_value = ".png" @@ -101,9 +106,14 @@ class TestDatasourceFileMessageTransformer: def test_transform_blob_message_image(self, mock_guess_ext, mock_tool_file_manager_cls): # Setup mock_manager = mock_tool_file_manager_cls.return_value - mock_tool_file = MagicMock(spec=ToolFile) + mock_tool_file = ToolFile( + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + file_key="test-key", + mimetype="image/jpeg", + ) mock_tool_file.id = "blob_id_456" - mock_tool_file.mimetype = "image/jpeg" mock_manager.create_file_by_raw.return_value = mock_tool_file mock_guess_ext.return_value = ".jpg" @@ -137,9 +147,14 @@ class TestDatasourceFileMessageTransformer: ): # Setup mock_manager = mock_tool_file_manager_cls.return_value - mock_tool_file = MagicMock(spec=ToolFile) + mock_tool_file = ToolFile( + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + file_key="test-key", + mimetype="application/pdf", + ) mock_tool_file.id = "blob_id_789" - mock_tool_file.mimetype = "application/pdf" mock_manager.create_file_by_raw.return_value = mock_tool_file mock_guess_type.return_value = ("application/pdf", None) mock_guess_ext.return_value = ".pdf" @@ -291,7 +306,13 @@ class TestDatasourceFileMessageTransformer: # This tests line 70 where filename might be None with patch("core.datasource.utils.message_transformer.ToolFileManager") as mock_tool_file_manager_cls: mock_manager = mock_tool_file_manager_cls.return_value - mock_tool_file = MagicMock(spec=ToolFile) + mock_tool_file = ToolFile( + user_id="user-id", + tenant_id="tenant-id", + conversation_id="conversation-id", + file_key="test-key", + mimetype="text/plain", + ) mock_tool_file.id = "blob_id_no_name" mock_tool_file.mimetype = "application/octet-stream" mock_manager.create_file_by_raw.return_value = mock_tool_file diff --git a/api/tests/unit_tests/core/entities/test_entities_parameter_entities.py b/api/tests/unit_tests/core/entities/test_entities_parameter_entities.py index 20b7bf2a9fd..0ba8805e4c8 100644 --- a/api/tests/unit_tests/core/entities/test_entities_parameter_entities.py +++ b/api/tests/unit_tests/core/entities/test_entities_parameter_entities.py @@ -1,9 +1,12 @@ +import pytest + from core.entities.parameter_entities import ( AppSelectorScope, CommonParameterType, ModelSelectorScope, ToolSelectorScope, ) +from core.plugin.entities.parameters import PluginParameterType, cast_parameter_value def test_common_parameter_type_values_are_stable() -> None: @@ -13,6 +16,42 @@ def test_common_parameter_type_values_are_stable() -> None: assert CommonParameterType.DYNAMIC_SELECT.value == "dynamic-select" assert CommonParameterType.ARRAY.value == "array" assert CommonParameterType.OBJECT.value == "object" + assert CommonParameterType.DATE.value == "date" + assert CommonParameterType.DATE_RANGE.value == "date-range" + with pytest.raises(ValueError): + PluginParameterType("date-picker") + + +def test_cast_date_accepts_only_canonical_calendar_dates() -> None: + assert cast_parameter_value(PluginParameterType.DATE, None) == "" + assert cast_parameter_value(PluginParameterType.DATE, "") == "" + assert cast_parameter_value(PluginParameterType.DATE, "2024-02-29") == "2024-02-29" + + for value in ("a", "2024-02-30", "20240229", 20240229): + with pytest.raises(ValueError): + cast_parameter_value(PluginParameterType.DATE, value) + + +def test_cast_date_range_validates_optional_range() -> None: + assert cast_parameter_value(PluginParameterType.DATE_RANGE, "") == {} + assert cast_parameter_value(PluginParameterType.DATE_RANGE, {}) == {} + assert cast_parameter_value(PluginParameterType.DATE_RANGE, "2024-01-01") == {"start": "2024-01-01"} + assert cast_parameter_value(PluginParameterType.DATE_RANGE, {"end": "2024-01-02"}) == {"end": "2024-01-02"} + assert cast_parameter_value( + PluginParameterType.DATE_RANGE, + '{"start":"2024-01-01","end":"2024-01-02"}', + ) == {"start": "2024-01-01", "end": "2024-01-02"} + + for value in ( + "a", + {"start": "a"}, + {"start": "2024-02-30"}, + {"start": 20240101}, + {"start": "2024-01-02", "end": "2024-01-01"}, + '{"start":"2024-01-02","end":"2024-01-01"}', + ): + with pytest.raises(ValueError): + cast_parameter_value(PluginParameterType.DATE_RANGE, value) def test_selector_scope_values_are_stable() -> None: diff --git a/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py b/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py index b5a48918079..b5489c88a03 100644 --- a/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py +++ b/api/tests/unit_tests/core/entities/test_entities_provider_configuration.py @@ -1,12 +1,16 @@ from __future__ import annotations +import json import logging +from collections.abc import Iterator from contextlib import contextmanager from types import SimpleNamespace from typing import Any -from unittest.mock import Mock, patch +from unittest.mock import Mock, PropertyMock, patch import pytest +from sqlalchemy import Engine, event, select +from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from core.entities.model_entities import ModelStatus @@ -26,6 +30,7 @@ from core.entities.provider_entities import ( SystemConfigurationStatus, ) from core.helper.model_provider_cache import ProviderCredentialsCacheType +from extensions.ext_database import db from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelType from graphon.model_runtime.entities.provider_entities import ( @@ -38,7 +43,16 @@ from graphon.model_runtime.entities.provider_entities import ( ProviderEntity, ) from models.enums import CredentialSourceType -from models.provider import ProviderType +from models.provider import ( + LoadBalancingModelConfig, + Provider, + ProviderCredential, + ProviderModel, + ProviderModelCredential, + ProviderModelSetting, + ProviderType, + TenantPreferredModelProvider, +) from models.provider_ids import ModelProviderID _UNSET = object() @@ -1156,6 +1170,41 @@ def test_get_custom_model_record_supports_plugin_id_alias() -> None: assert result is custom_model_record +def test_model_type_db_values_includes_pre_1_15_aliases() -> None: + from core.entities.provider_configuration import _model_type_db_values + + assert _model_type_db_values(ModelType.LLM) == ("llm", "text-generation") + assert _model_type_db_values(ModelType.TEXT_EMBEDDING) == ("text-embedding", "embeddings") + assert _model_type_db_values(ModelType.RERANK) == ("rerank", "reranking") + assert _model_type_db_values(ModelType.TTS) == ("tts",) + + +def test_get_custom_model_record_uses_legacy_aware_model_type_filter( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Regression for #39559: lookups must include pre-1.15 model_type aliases.""" + import core.entities.provider_configuration as provider_configuration_module + + captured: dict[str, tuple[str, ...]] = {} + original = provider_configuration_module._model_type_db_values + + def _capture(model_type: ModelType) -> tuple[str, ...]: + values = original(model_type) + captured["values"] = values + return values + + monkeypatch.setattr(provider_configuration_module, "_model_type_db_values", _capture) + + configuration = _build_provider_configuration(provider_name="langgenius/ollama/ollama") + session = Mock() + session.execute.return_value.scalar_one_or_none.return_value = SimpleNamespace(id="legacy-model") + + result = configuration._get_custom_model_record(ModelType.LLM, "llama3", session) + + assert result.id == "legacy-model" + assert captured["values"] == ("llm", "text-generation") + + def test_get_specific_custom_model_credential_success_and_not_found() -> None: configuration = _build_provider_configuration() configuration.provider.model_credential_schema = _build_secret_model_schema() @@ -2123,3 +2172,768 @@ def test_get_custom_provider_models_skips_custom_models_on_schema_error_or_none( assert "get custom model schema failed, boom" in caplog.messages assert any(model.model == "ok-custom" for model in models) assert all(model.model != "none-custom" for model in models) + + +@pytest.fixture +def sqlite_provider_session( + sqlite_session: Session, + sqlite_engine: Engine, +) -> Iterator[Session]: + """Bind provider-owned sessions to the same isolated SQLite database as the test.""" + with patch.object(type(db), "engine", new_callable=PropertyMock, return_value=sqlite_engine): + yield sqlite_session + + +def _provider_credential( + session: Session, + *, + name: str = "API KEY 1", + tenant_id: str = "tenant-1", + provider_name: str = "openai", + encrypted_config: str = "{}", +) -> ProviderCredential: + record = ProviderCredential( + tenant_id=tenant_id, + provider_name=provider_name, + credential_name=name, + encrypted_config=encrypted_config, + ) + session.add(record) + session.commit() + return record + + +def _provider_record( + session: Session, + *, + credential_id: str | None = None, + tenant_id: str = "tenant-1", + provider_name: str = "openai", +) -> Provider: + record = Provider( + tenant_id=tenant_id, + provider_name=provider_name, + provider_type=ProviderType.CUSTOM, + credential_id=credential_id, + is_valid=True, + ) + session.add(record) + session.commit() + return record + + +def _model_credential( + session: Session, + *, + name: str = "API KEY 1", + tenant_id: str = "tenant-1", + provider_name: str = "openai", + model: str = "gpt-4o", + encrypted_config: str = "{}", +) -> ProviderModelCredential: + record = ProviderModelCredential( + tenant_id=tenant_id, + provider_name=provider_name, + model_name=model, + model_type=ModelType.LLM, + credential_name=name, + encrypted_config=encrypted_config, + ) + session.add(record) + session.commit() + return record + + +def _provider_model_record( + session: Session, + *, + credential_id: str | None = None, + tenant_id: str = "tenant-1", + provider_name: str = "openai", + model: str = "gpt-4o", +) -> ProviderModel: + record = ProviderModel( + tenant_id=tenant_id, + provider_name=provider_name, + model_name=model, + model_type=ModelType.LLM, + credential_id=credential_id, + is_valid=True, + ) + session.add(record) + session.commit() + return record + + +def _load_balancing_config( + session: Session, + *, + credential_id: str, + source: CredentialSourceType, + name: str = "Old", +) -> LoadBalancingModelConfig: + record = LoadBalancingModelConfig( + tenant_id="tenant-1", + provider_name="openai", + model_name="gpt-4o", + model_type=ModelType.LLM, + name=name, + encrypted_config="{}", + credential_id=credential_id, + credential_source_type=source, + ) + session.add(record) + session.commit() + return record + + +@contextmanager +def _raise_on_sql(engine: Engine, table_name: str, operation: str) -> Iterator[None]: + """Fail one table operation while production still owns a real transaction.""" + + def fail_target(_conn, _cursor, statement, _parameters, _context, _executemany): + if statement.lstrip().upper().startswith(operation) and table_name in statement: + raise RuntimeError(f"forced {operation} failure for {table_name}") + + event.listen(engine, "before_cursor_execute", fail_target) + try: + yield + finally: + event.remove(engine, "before_cursor_execute", fail_target) + + +@contextmanager +def _mock_cache_boundaries() -> Iterator[tuple[Mock, Mock]]: + with ( + patch("core.entities.provider_configuration.ProviderCredentialsCache") as credentials_cache, + patch.object(ProviderConfiguration, "_invalidate_provider_configuration_cache") as configuration_cache, + ): + yield credentials_cache, configuration_cache + + +def test_generate_credential_names_from_real_rows_and_tenant_isolation( + sqlite_provider_session: Session, +) -> None: + configuration = _build_provider_configuration() + _provider_credential(sqlite_provider_session, name="API KEY 9") + _provider_credential(sqlite_provider_session, name="legacy") + _provider_credential(sqlite_provider_session, name="API KEY 50", tenant_id="other-tenant") + _model_credential(sqlite_provider_session, name="API KEY 4") + assert configuration._generate_provider_credential_name(sqlite_provider_session) == "API KEY 10" + assert ( + configuration._generate_custom_model_credential_name("gpt-4o", ModelType.LLM, sqlite_provider_session) + == "API KEY 5" + ) + + +def test_validate_provider_credentials_reuses_hidden_secret(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + configuration.provider.provider_credential_schema = _build_secret_provider_schema() + credential = _provider_credential(sqlite_provider_session, encrypted_config='{"openai_api_key":"enc-old"}') + factory = Mock() + factory.provider_credentials_validate.return_value = {"openai_api_key": "raw"} + with ( + patch( + "core.entities.provider_configuration.create_plugin_model_assembly", + return_value=SimpleNamespace(model_runtime=Mock(), model_provider_factory=factory), + ), + patch("core.entities.provider_configuration.encrypter.decrypt_token", return_value="raw"), + patch("core.entities.provider_configuration.encrypter.encrypt_token", return_value="enc-new"), + ): + result = configuration.validate_provider_credentials( + {"openai_api_key": HIDDEN_VALUE}, credential_id=credential.id + ) + assert result == {"openai_api_key": "enc-new"} + + +def test_preferred_provider_state_updates_and_is_tenant_scoped(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + configuration.preferred_provider_type = ProviderType.CUSTOM + other = TenantPreferredModelProvider( + tenant_id="other-tenant", provider_name="openai", preferred_provider_type=ProviderType.CUSTOM + ) + current = TenantPreferredModelProvider( + tenant_id="tenant-1", provider_name="openai", preferred_provider_type=ProviderType.CUSTOM + ) + sqlite_provider_session.add_all([other, current]) + sqlite_provider_session.commit() + assert configuration.switch_preferred_provider_type(ProviderType.SYSTEM, session=sqlite_provider_session) + sqlite_provider_session.refresh(current) + sqlite_provider_session.refresh(other) + assert current.preferred_provider_type == ProviderType.SYSTEM + assert other.preferred_provider_type == ProviderType.CUSTOM + + +def test_provider_record_duplicate_and_setting_helpers_use_real_session(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + provider = _provider_record(sqlite_provider_session) + _provider_record(sqlite_provider_session, tenant_id="other-tenant") + credential = _provider_credential(sqlite_provider_session, name="Main") + _provider_credential(sqlite_provider_session, name="Main", tenant_id="other-tenant") + setting = ProviderModelSetting( + tenant_id="tenant-1", + provider_name="openai", + model_name="gpt-4o", + model_type=ModelType.LLM, + ) + sqlite_provider_session.add(setting) + sqlite_provider_session.commit() + assert configuration._get_provider_record(sqlite_provider_session).id == provider.id + assert configuration._check_provider_credential_name_exists("Main", sqlite_provider_session) + assert not configuration._check_provider_credential_name_exists( + "Main", sqlite_provider_session, exclude_id=credential.id + ) + assert configuration._get_provider_model_setting(ModelType.LLM, "gpt-4o", sqlite_provider_session).id == setting.id + + +def test_create_provider_credential_persists_provider_and_rejects_duplicate( + sqlite_provider_session: Session, +) -> None: + configuration = _build_provider_configuration() + with ( + patch.object(ProviderConfiguration, "validate_provider_credentials", return_value={"api_key": "enc"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + ): + configuration.create_provider_credential({"api_key": "raw"}, "Main") + credential = sqlite_provider_session.scalar( + select(ProviderCredential).where(ProviderCredential.credential_name == "Main") + ) + provider = sqlite_provider_session.scalar(select(Provider).where(Provider.tenant_id == "tenant-1")) + assert credential is not None + assert provider is not None + assert provider.credential_id == credential.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=provider.id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_called_once_with( + preferred_model_providers=True, + provider_credentials=True, + ) + with pytest.raises(ValueError, match="already exists"): + configuration.create_provider_credential({"api_key": "raw"}, "Main") + + +def test_update_provider_credential_propagates_to_load_balancing(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + credential = _provider_credential(sqlite_provider_session, name="Old") + provider = _provider_record(sqlite_provider_session, credential_id=credential.id) + lb_config = _load_balancing_config( + sqlite_provider_session, credential_id=credential.id, source=CredentialSourceType.PROVIDER + ) + with ( + patch.object(ProviderConfiguration, "validate_provider_credentials", return_value={"api_key": "enc-new"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + ): + configuration.update_provider_credential({"api_key": "raw"}, credential.id, "New") + sqlite_provider_session.expire_all() + persisted_credential = sqlite_provider_session.get(ProviderCredential, credential.id) + persisted_lb = sqlite_provider_session.get(LoadBalancingModelConfig, lb_config.id) + assert persisted_credential is not None + assert persisted_credential.credential_name == "New" + assert persisted_lb is not None + assert persisted_lb.name == "New" + assert json.loads(persisted_lb.encrypted_config) == {"api_key": "enc-new"} + assert {cache_call.kwargs["identity_id"] for cache_call in credentials_cache.call_args_list} == { + provider.id, + lb_config.id, + } + assert {cache_call.kwargs["cache_type"] for cache_call in credentials_cache.call_args_list} == { + ProviderCredentialsCacheType.PROVIDER, + ProviderCredentialsCacheType.LOAD_BALANCING_MODEL, + } + assert credentials_cache.return_value.delete.call_count == 2 + configuration_cache.assert_called_once_with( + provider_credentials=True, + provider_load_balancing_configs=True, + ) + + +def test_switch_active_provider_credential_updates_persisted_state_and_cache( + sqlite_provider_session: Session, +) -> None: + configuration = _build_provider_configuration() + configuration.preferred_provider_type = ProviderType.CUSTOM + first = _provider_credential(sqlite_provider_session, name="First") + second = _provider_credential(sqlite_provider_session, name="Second") + provider = _provider_record(sqlite_provider_session, credential_id=first.id) + provider_id = provider.id + + with _mock_cache_boundaries() as (credentials_cache, configuration_cache): + configuration.switch_active_provider_credential(second.id) + + sqlite_provider_session.expire_all() + persisted_provider = sqlite_provider_session.get(Provider, provider_id) + assert persisted_provider is not None + assert persisted_provider.credential_id == second.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=provider_id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_not_called() + + +def test_deleting_active_provider_credential_switches_preference_to_system( + sqlite_provider_session: Session, +) -> None: + configuration = _build_provider_configuration() + configuration.preferred_provider_type = ProviderType.CUSTOM + first = _provider_credential(sqlite_provider_session, name="First") + active = _provider_credential(sqlite_provider_session, name="Active") + provider = _provider_record(sqlite_provider_session, credential_id=active.id) + preferred_provider = TenantPreferredModelProvider( + tenant_id="tenant-1", + provider_name="openai", + preferred_provider_type=ProviderType.CUSTOM, + ) + sqlite_provider_session.add(preferred_provider) + sqlite_provider_session.commit() + first_id = first.id + active_id = active.id + provider_id = provider.id + preferred_provider_id = preferred_provider.id + + with _mock_cache_boundaries() as (credentials_cache, configuration_cache): + configuration.delete_provider_credential(active_id) + + sqlite_provider_session.expire_all() + assert sqlite_provider_session.get(ProviderCredential, active_id) is None + assert sqlite_provider_session.get(ProviderCredential, first_id) is not None + persisted_provider = sqlite_provider_session.get(Provider, provider_id) + persisted_preference = sqlite_provider_session.get(TenantPreferredModelProvider, preferred_provider_id) + assert persisted_provider is not None + assert persisted_provider.credential_id is None + assert persisted_preference is not None + assert persisted_preference.preferred_provider_type == ProviderType.SYSTEM + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=provider_id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_called_once_with( + preferred_model_providers=True, + provider_credentials=True, + provider_load_balancing_configs=False, + ) + + +def test_specific_provider_credential_decrypts_and_obfuscates(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + configuration.provider.provider_credential_schema = _build_secret_provider_schema() + credential = _provider_credential(sqlite_provider_session, encrypted_config='{"openai_api_key":"enc"}') + with ( + patch("core.entities.provider_configuration.encrypter.decrypt_token", return_value="raw"), + patch("core.entities.provider_configuration.encrypter.obfuscated_token", return_value="masked"), + ): + result = configuration._get_specific_provider_credential(credential.id) + assert result == {"openai_api_key": "masked"} + with pytest.raises(ValueError, match="not found"): + configuration._get_specific_provider_credential("missing") + + +def test_validate_custom_model_credentials_reuses_hidden_secret(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + configuration.provider.model_credential_schema = _build_secret_model_schema() + credential = _model_credential(sqlite_provider_session, encrypted_config='{"openai_api_key":"enc-old"}') + factory = Mock() + factory.model_credentials_validate.return_value = {"openai_api_key": "raw"} + with ( + patch( + "core.entities.provider_configuration.create_plugin_model_assembly", + return_value=SimpleNamespace(model_runtime=Mock(), model_provider_factory=factory), + ), + patch("core.entities.provider_configuration.encrypter.decrypt_token", return_value="raw"), + patch("core.entities.provider_configuration.encrypter.encrypt_token", return_value="enc-new"), + ): + result = configuration.validate_custom_model_credentials( + ModelType.LLM, + "gpt-4o", + {"openai_api_key": HIDDEN_VALUE}, + credential_id=credential.id, + ) + assert result == {"openai_api_key": "enc-new"} + + +def test_specific_custom_model_credential_preserves_secret_when_decryption_fails( + sqlite_provider_session: Session, + caplog: pytest.LogCaptureFixture, +) -> None: + configuration = _build_provider_configuration() + configuration.provider.model_credential_schema = _build_secret_model_schema() + credential = _model_credential( + sqlite_provider_session, + name="Main", + encrypted_config='{"openai_api_key":"enc-secret"}', + ) + + with ( + caplog.at_level(logging.ERROR, logger="core.entities.provider_configuration"), + patch("core.entities.provider_configuration.encrypter.decrypt_token", side_effect=RuntimeError("boom")), + patch.object( + ProviderConfiguration, + "obfuscated_credentials", + side_effect=lambda credentials, credential_form_schemas: credentials, + ), + ): + result = configuration._get_specific_custom_model_credential(ModelType.LLM, "gpt-4o", credential.id) + + assert result == { + "current_credential_id": credential.id, + "current_credential_name": "Main", + "credentials": {"openai_api_key": "enc-secret"}, + } + assert caplog.messages.count("Failed to decrypt model credential secret variable openai_api_key") == 1 + + +def test_create_update_and_delete_custom_model_credential(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + with ( + patch.object(ProviderConfiguration, "validate_custom_model_credentials", return_value={"api_key": "enc"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + ): + configuration.create_custom_model_credential(ModelType.LLM, "gpt-4o", {"api_key": "raw"}, "Main") + credential = sqlite_provider_session.scalar(select(ProviderModelCredential)) + model = sqlite_provider_session.scalar(select(ProviderModel)) + assert credential is not None + assert model is not None + assert model.credential_id == credential.id + credential_id = credential.id + model_id = model.id + lb_config = _load_balancing_config( + sqlite_provider_session, + credential_id=credential_id, + source=CredentialSourceType.CUSTOM_MODEL, + ) + lb_config_id = lb_config.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=model_id, + cache_type=ProviderCredentialsCacheType.MODEL, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_called_once_with( + provider_models=True, + provider_model_credentials=True, + ) + + with ( + patch.object(ProviderConfiguration, "validate_custom_model_credentials", return_value={"api_key": "enc-2"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + ): + configuration.update_custom_model_credential( + ModelType.LLM, "gpt-4o", {"api_key": "raw"}, "Renamed", credential_id + ) + sqlite_provider_session.expire_all() + persisted_credential = sqlite_provider_session.get(ProviderModelCredential, credential_id) + persisted_lb = sqlite_provider_session.get(LoadBalancingModelConfig, lb_config_id) + assert persisted_credential is not None + assert persisted_credential.credential_name == "Renamed" + assert json.loads(persisted_credential.encrypted_config) == {"api_key": "enc-2"} + assert persisted_lb is not None + assert persisted_lb.name == "Renamed" + assert json.loads(persisted_lb.encrypted_config) == {"api_key": "enc-2"} + assert {cache_call.kwargs["identity_id"] for cache_call in credentials_cache.call_args_list} == { + model_id, + lb_config_id, + } + assert {cache_call.kwargs["cache_type"] for cache_call in credentials_cache.call_args_list} == { + ProviderCredentialsCacheType.MODEL, + ProviderCredentialsCacheType.LOAD_BALANCING_MODEL, + } + assert credentials_cache.return_value.delete.call_count == 2 + configuration_cache.assert_called_once_with( + provider_models=True, + provider_model_credentials=True, + provider_load_balancing_configs=True, + ) + + with _mock_cache_boundaries() as (credentials_cache, configuration_cache): + configuration.delete_custom_model_credential(ModelType.LLM, "gpt-4o", credential_id) + sqlite_provider_session.expire_all() + assert sqlite_provider_session.get(ProviderModelCredential, credential_id) is None + assert sqlite_provider_session.get(ProviderModel, model_id) is None + assert sqlite_provider_session.get(LoadBalancingModelConfig, lb_config_id) is None + assert {cache_call.kwargs["identity_id"] for cache_call in credentials_cache.call_args_list} == { + model_id, + lb_config_id, + } + assert {cache_call.kwargs["cache_type"] for cache_call in credentials_cache.call_args_list} == { + ProviderCredentialsCacheType.MODEL, + ProviderCredentialsCacheType.LOAD_BALANCING_MODEL, + } + assert credentials_cache.return_value.delete.call_count == 2 + configuration_cache.assert_called_once_with( + provider_models=True, + provider_model_credentials=True, + provider_load_balancing_configs=True, + ) + + +def test_add_and_switch_custom_model_credential(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + first = _model_credential(sqlite_provider_session, name="First") + second = _model_credential(sqlite_provider_session, name="Second") + with _mock_cache_boundaries() as (credentials_cache, configuration_cache): + configuration.add_model_credential_to_model(ModelType.LLM, "gpt-4o", first.id) + configuration.switch_custom_model_credential(ModelType.LLM, "gpt-4o", second.id) + model = sqlite_provider_session.scalar(select(ProviderModel)) + assert model is not None + assert model.credential_id == second.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=model.id, + cache_type=ProviderCredentialsCacheType.MODEL, + ) + credentials_cache.return_value.delete.assert_called_once_with() + assert configuration_cache.call_count == 2 + for cache_call in configuration_cache.call_args_list: + assert cache_call.kwargs == {"provider_models": True} + with pytest.raises(ValueError, match="Can't add same credential"): + configuration.add_model_credential_to_model(ModelType.LLM, "gpt-4o", second.id) + + +def test_model_settings_and_load_balancing_persist(sqlite_provider_session: Session) -> None: + configuration = _build_provider_configuration() + with patch.object(configuration, "_invalidate_provider_configuration_cache") as configuration_cache: + configuration.disable_model(ModelType.LLM, "gpt-4o") + configuration_cache.assert_called_once_with(provider_model_settings=True) + persisted_setting = sqlite_provider_session.scalar(select(ProviderModelSetting)) + assert persisted_setting is not None + assert persisted_setting.enabled is False + with patch.object(configuration, "_invalidate_provider_configuration_cache") as configuration_cache: + configuration.enable_model(ModelType.LLM, "gpt-4o") + configuration_cache.assert_called_once_with(provider_model_settings=True) + sqlite_provider_session.expire_all() + refreshed_setting = sqlite_provider_session.get(ProviderModelSetting, persisted_setting.id) + assert refreshed_setting is not None + assert refreshed_setting.enabled is True + + first = _provider_credential(sqlite_provider_session, name="First") + second = _provider_credential(sqlite_provider_session, name="Second") + _load_balancing_config(sqlite_provider_session, credential_id=first.id, source=CredentialSourceType.PROVIDER) + _load_balancing_config(sqlite_provider_session, credential_id=second.id, source=CredentialSourceType.PROVIDER) + with patch.object(configuration, "_invalidate_provider_configuration_cache") as configuration_cache: + configuration.enable_model_load_balancing(ModelType.LLM, "gpt-4o") + configuration_cache.assert_called_once_with(provider_model_settings=True) + sqlite_provider_session.expire_all() + refreshed_setting = sqlite_provider_session.get(ProviderModelSetting, persisted_setting.id) + assert refreshed_setting is not None + assert refreshed_setting.load_balancing_enabled is True + with patch.object(configuration, "_invalidate_provider_configuration_cache") as configuration_cache: + configuration.disable_model_load_balancing(ModelType.LLM, "gpt-4o") + configuration_cache.assert_called_once_with(provider_model_settings=True) + sqlite_provider_session.expire_all() + refreshed_setting = sqlite_provider_session.get(ProviderModelSetting, persisted_setting.id) + assert refreshed_setting is not None + assert refreshed_setting.load_balancing_enabled is False + + +def test_provider_create_rolls_back_on_insert_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + with ( + patch.object(ProviderConfiguration, "validate_provider_credentials", return_value={"api_key": "enc"}), + _mock_cache_boundaries(), + _raise_on_sql(sqlite_engine, "provider_credentials", "INSERT"), + pytest.raises(RuntimeError, match="forced INSERT"), + ): + configuration.create_provider_credential({"api_key": "raw"}, "Main") + assert sqlite_provider_session.scalar(select(ProviderCredential)) is None + assert sqlite_provider_session.scalar(select(Provider)) is None + + +def test_custom_model_create_rolls_back_on_insert_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + with ( + patch.object(ProviderConfiguration, "validate_custom_model_credentials", return_value={"api_key": "enc"}), + _mock_cache_boundaries(), + _raise_on_sql(sqlite_engine, "provider_model_credentials", "INSERT"), + pytest.raises(RuntimeError, match="forced INSERT"), + ): + configuration.create_custom_model_credential(ModelType.LLM, "gpt-4o", {"api_key": "raw"}, "Main") + assert sqlite_provider_session.scalar(select(ProviderModelCredential)) is None + assert sqlite_provider_session.scalar(select(ProviderModel)) is None + + +def test_provider_update_rolls_back_on_update_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + credential = _provider_credential( + sqlite_provider_session, + name="Old", + encrypted_config='{"api_key":"enc-old"}', + ) + _provider_record(sqlite_provider_session, credential_id=credential.id) + + with ( + patch.object(ProviderConfiguration, "validate_provider_credentials", return_value={"api_key": "enc-new"}), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_credentials", "UPDATE"), + pytest.raises(RuntimeError, match="forced UPDATE"), + ): + configuration.update_provider_credential({"api_key": "raw"}, credential.id, "New") + + sqlite_provider_session.expire_all() + persisted = sqlite_provider_session.get(ProviderCredential, credential.id) + assert persisted is not None + assert persisted.credential_name == "Old" + assert json.loads(persisted.encrypted_config) == {"api_key": "enc-old"} + credentials_cache.assert_not_called() + configuration_cache.assert_not_called() + + +def test_provider_delete_rolls_back_on_delete_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + credential = _provider_credential(sqlite_provider_session) + provider = _provider_record(sqlite_provider_session, credential_id=credential.id) + + with ( + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_credentials", "DELETE"), + pytest.raises(RuntimeError, match="forced DELETE"), + ): + configuration.delete_provider_credential(credential.id) + + sqlite_provider_session.expire_all() + assert sqlite_provider_session.get(ProviderCredential, credential.id) is not None + persisted_provider = sqlite_provider_session.get(Provider, provider.id) + assert persisted_provider is not None + assert persisted_provider.credential_id == credential.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=provider.id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_not_called() + + +def test_provider_switch_rolls_back_on_update_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + first = _provider_credential(sqlite_provider_session, name="First") + second = _provider_credential(sqlite_provider_session, name="Second") + provider = _provider_record(sqlite_provider_session, credential_id=first.id) + + with ( + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "providers", "UPDATE"), + pytest.raises(RuntimeError, match="forced UPDATE"), + ): + configuration.switch_active_provider_credential(second.id) + + sqlite_provider_session.expire_all() + persisted_provider = sqlite_provider_session.get(Provider, provider.id) + assert persisted_provider is not None + assert persisted_provider.credential_id == first.id + credentials_cache.assert_not_called() + configuration_cache.assert_not_called() + + +def test_custom_model_update_rolls_back_on_update_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + credential = _model_credential( + sqlite_provider_session, + name="Old", + encrypted_config='{"api_key":"enc-old"}', + ) + _provider_model_record(sqlite_provider_session, credential_id=credential.id) + + with ( + patch.object( + ProviderConfiguration, + "validate_custom_model_credentials", + return_value={"api_key": "enc-new"}, + ), + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_model_credentials", "UPDATE"), + pytest.raises(RuntimeError, match="forced UPDATE"), + ): + configuration.update_custom_model_credential( + ModelType.LLM, + "gpt-4o", + {"api_key": "raw"}, + "New", + credential.id, + ) + + sqlite_provider_session.expire_all() + persisted = sqlite_provider_session.get(ProviderModelCredential, credential.id) + assert persisted is not None + assert persisted.credential_name == "Old" + assert json.loads(persisted.encrypted_config) == {"api_key": "enc-old"} + credentials_cache.assert_not_called() + configuration_cache.assert_not_called() + + +def test_custom_model_delete_rolls_back_on_delete_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + credential = _model_credential(sqlite_provider_session) + model = _provider_model_record(sqlite_provider_session, credential_id=credential.id) + + with ( + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_model_credentials", "DELETE"), + pytest.raises(RuntimeError, match="forced DELETE"), + ): + configuration.delete_custom_model_credential(ModelType.LLM, "gpt-4o", credential.id) + + sqlite_provider_session.expire_all() + assert sqlite_provider_session.get(ProviderModelCredential, credential.id) is not None + persisted_model = sqlite_provider_session.get(ProviderModel, model.id) + assert persisted_model is not None + assert persisted_model.credential_id == credential.id + credentials_cache.assert_called_once_with( + tenant_id="tenant-1", + identity_id=model.id, + cache_type=ProviderCredentialsCacheType.MODEL, + ) + credentials_cache.return_value.delete.assert_called_once_with() + configuration_cache.assert_not_called() + + +def test_custom_model_switch_rolls_back_on_update_failure( + sqlite_provider_session: Session, + sqlite_engine: Engine, +) -> None: + configuration = _build_provider_configuration() + first = _model_credential(sqlite_provider_session, name="First") + second = _model_credential(sqlite_provider_session, name="Second") + model = _provider_model_record(sqlite_provider_session, credential_id=first.id) + + with ( + _mock_cache_boundaries() as (credentials_cache, configuration_cache), + _raise_on_sql(sqlite_engine, "provider_models", "UPDATE"), + pytest.raises(RuntimeError, match="forced UPDATE"), + ): + configuration.switch_custom_model_credential(ModelType.LLM, "gpt-4o", second.id) + + sqlite_provider_session.expire_all() + persisted_model = sqlite_provider_session.get(ProviderModel, model.id) + assert persisted_model is not None + assert persisted_model.credential_id == first.id + credentials_cache.assert_not_called() + configuration_cache.assert_not_called() diff --git a/api/tests/unit_tests/core/helper/code_executor/jinja2/test_jinja2_sandbox.py b/api/tests/unit_tests/core/helper/code_executor/jinja2/test_jinja2_sandbox.py new file mode 100644 index 00000000000..7ff7aa166c7 --- /dev/null +++ b/api/tests/unit_tests/core/helper/code_executor/jinja2/test_jinja2_sandbox.py @@ -0,0 +1,67 @@ +"""Verify that Jinja2 transform_caller scripts block unsafe template attacks via SandboxedEnvironment.""" + +import io +import sys + +import pytest + +from core.helper.code_executor.jinja2.jinja2_transformer import Jinja2TemplateTransformer + +MALICIOUS_TEMPLATES = [ + pytest.param( + "{{ lipsum.__globals__.__builtins__.__import__('os').popen('id').read() }}", + id="lipsum_globals_builtins", + ), + pytest.param( + "{{ ''.__class__.__mro__[1].__subclasses__() }}", + id="string_class_mro", + ), + pytest.param( + "{{ cycler.__init__.__globals__.os.popen('whoami').read() }}", + id="cycler_init_globals", + ), + pytest.param( + "{{ namespace.__init__.__globals__['__builtins__']['__import__']('os').system('id') }}", + id="namespace_init_globals", + ), +] + + +def _exec_scripts(runner: str, preload: str) -> str: + """Execute preload then runner in a shared namespace, return captured stdout.""" + ns: dict = {} + exec(compile(preload, "", "exec"), ns) # noqa: S102 + captured = io.StringIO() + old_stdout = sys.stdout + sys.stdout = captured + try: + exec(compile(runner, "", "exec"), ns) # noqa: S102 + finally: + sys.stdout = old_stdout + return captured.getvalue() + + +class TestJinja2TransformCallerSandbox: + """Test transform_caller output (runner + preload) blocks attacks and allows safe templates.""" + + @pytest.mark.parametrize("malicious_template", MALICIOUS_TEMPLATES) + def test_blocks_unsafe_template(self, malicious_template: str) -> None: + runner, preload = Jinja2TemplateTransformer.transform_caller(malicious_template, {}) + ns: dict = {} + exec(compile(preload, "", "exec"), ns) # noqa: S102 + with pytest.raises(Exception) as exc_info: + exec(compile(runner, "", "exec"), ns) # noqa: S102 + assert "unsafe" in str(exc_info.value).lower() or "security" in str(exc_info.value).lower() + + def test_renders_safe_template(self) -> None: + runner, preload = Jinja2TemplateTransformer.transform_caller( + "Hello {{ name }}, you are {{ age }} years old!", + {"name": "Alice", "age": 30}, + ) + output = _exec_scripts(runner, preload) + assert "Hello Alice, you are 30 years old!" in output + + def test_scripts_use_sandboxed_environment(self) -> None: + runner, preload = Jinja2TemplateTransformer.transform_caller("{{ x }}", {"x": 1}) + assert "SandboxedEnvironment" in runner + assert "SandboxedEnvironment" in preload diff --git a/api/tests/unit_tests/core/helper/test_ssrf_proxy.py b/api/tests/unit_tests/core/helper/test_ssrf_proxy.py index f065728b227..5b51a99aa58 100644 --- a/api/tests/unit_tests/core/helper/test_ssrf_proxy.py +++ b/api/tests/unit_tests/core/helper/test_ssrf_proxy.py @@ -19,6 +19,7 @@ from core.helper.ssrf_proxy import ( max_retries_exceeded_error, request_error, ) +from core.tools.errors import ToolSSRFError @patch("core.helper.ssrf_proxy._get_ssrf_client", autospec=True) @@ -360,3 +361,96 @@ def test_graphon_ssrf_proxy_wraps_module_requests(method_name: str) -> None: assert wrapped.status_code == 200 assert wrapped.url == "https://example.com/resource" assert wrapped.content == b"ok" + + +# --------------------------------------------------------------------------- +# Squid-blocked 403 regression tests (issue #38443) +# --------------------------------------------------------------------------- +# When the SSRF proxy denies a request to a private/internal network address, +# Squid returns 401/403 with itself identified in the Server or Via header. +# The Python client must raise ToolSSRFError with a message that tells the +# user exactly which env var to set, instead of just "blocked by SSRF +# protection" (the pre-#38443 message gave no actionable guidance). + + +def _build_squid_blocked_response(status_code: int = 403) -> MagicMock: + """Construct a mock httpx.Response that looks like Squid's ACL deny.""" + response = MagicMock() + response.status_code = status_code + response.headers = {"server": "squid/4.10", "via": "1.1 squid (squid/4.10)"} + return response + + +@patch("core.helper.ssrf_proxy._get_ssrf_client", autospec=True) +def test_squid_block_raises_actionable_tool_ssrf_error(mock_get_client) -> None: + """A 403 from Squid must raise ToolSSRFError whose message tells the user + exactly which env var to set. Pre-#38443 the message had no remediation + hint, so users hit dead ends when their internal API was blocked.""" + mock_client = MagicMock() + mock_client.send.return_value = _build_squid_blocked_response(status_code=403) + mock_get_client.return_value = mock_client + + with pytest.raises(ToolSSRFError) as exc_info: + make_request("GET", "http://172.21.0.5/api/health") + + msg = str(exc_info.value) + assert "172.21.0.5" in msg, f"URL should appear in the error, got: {msg!r}" + assert "SSRF_PROXY_ALLOW_PRIVATE_IPS" in msg, "Error must tell the user which env var to set; got: " + msg + # The remediation hint must include a concrete example, otherwise users + # still have to grep the squid config to figure out the syntax. + assert "172.21.0.0/16" in msg + # And it must point to the issue so maintainers can find context. + assert "issues/38443" in msg + + +@patch("core.helper.ssrf_proxy._get_ssrf_client", autospec=True) +def test_squid_401_via_header_also_triggers_actionable_error(mock_get_client) -> None: + """Squid can return 401 with only the Via header set (no Server header) + on some configurations. The detection must work for both.""" + mock_client = MagicMock() + response = MagicMock() + response.status_code = 401 + # Server header absent — only Via identifies Squid. + response.headers = {"server": "", "via": "1.1 squid (squid/4.10)"} + mock_client.send.return_value = response + mock_get_client.return_value = mock_client + + with pytest.raises(ToolSSRFError) as exc_info: + make_request("GET", "http://10.0.0.1/internal") + + assert "SSRF_PROXY_ALLOW_PRIVATE_IPS" in str(exc_info.value) + + +@patch("core.helper.ssrf_proxy._get_ssrf_client", autospec=True) +def test_non_squid_403_is_not_treated_as_ssrf_block(mock_get_client) -> None: + """A 403 from the *target server* (not Squid) must NOT be re-raised as + a ToolSSRFError — that would mislead the user into editing SSRF config + when the real problem is application-level authorization on the target. + Pre-#38443 we didn't have this guard at all; the new wording only changes + the Squid path, so verify we don't accidentally widen it.""" + mock_client = MagicMock() + response = MagicMock() + response.status_code = 403 + response.headers = {"server": "nginx/1.21", "via": "1.1 varnish"} + mock_client.send.return_value = response + mock_get_client.return_value = mock_client + + # Should return the response, not raise. + returned = make_request("GET", "http://public.example.com/admin") + assert returned.status_code == 403 + + +@patch("core.helper.ssrf_proxy._get_ssrf_client", autospec=True) +def test_squid_block_with_internal_10_x_url_mentions_allowlist(mock_get_client) -> None: + """Real-world repro from #38443: 10.x.x.x internal API blocked. The error + message must still point at SSRF_PROXY_ALLOW_PRIVATE_IPS, not just say + "private address" without telling the user what to do.""" + mock_client = MagicMock() + mock_client.send.return_value = _build_squid_blocked_response(status_code=403) + mock_get_client.return_value = mock_client + + with pytest.raises(ToolSSRFError) as exc_info: + make_request("POST", "http://10.0.0.42/v1/chat/completions") + + assert "10.0.0.42" in str(exc_info.value) + assert "SSRF_PROXY_ALLOW_PRIVATE_IPS" in str(exc_info.value) diff --git a/api/tests/unit_tests/core/llm_generator/test_llm_generator.py b/api/tests/unit_tests/core/llm_generator/test_llm_generator.py index c37becd5f06..531e07ef284 100644 --- a/api/tests/unit_tests/core/llm_generator/test_llm_generator.py +++ b/api/tests/unit_tests/core/llm_generator/test_llm_generator.py @@ -1,14 +1,170 @@ +"""Tests for LLM generation and database-backed instruction modification.""" + import json -from unittest.mock import MagicMock, patch +from collections.abc import Iterator +from datetime import datetime +from decimal import Decimal +from unittest.mock import MagicMock, Mock, patch +from uuid import uuid4 import pytest +from sqlalchemy.orm import Session, scoped_session, sessionmaker from core.app.app_config.entities import ModelConfig +from core.llm_generator import llm_generator as llm_generator_module from core.llm_generator.entities import RuleCodeGeneratePayload, RuleGeneratePayload, RuleStructuredOutputPayload from core.llm_generator.llm_generator import LLMGenerator -from graphon.model_runtime.entities.llm_entities import LLMMode, LLMResult +from core.model_manager import ModelInstance, ModelManager +from graphon.enums import WorkflowNodeExecutionStatus +from graphon.model_runtime.entities.llm_entities import LLMMode, LLMResult, LLMUsage +from graphon.model_runtime.entities.message_entities import AssistantPromptMessage from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.errors.invoke import InvokeAuthorizationError, InvokeError +from models.enums import ConversationFromSource, CreatorUserRole +from models.model import App, AppMode, Message +from models.workflow import ( + Workflow, + WorkflowNodeExecutionModel, + WorkflowNodeExecutionTriggeredFrom, +) +from services.workflow_service import WorkflowService + + +@pytest.fixture +def database(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]: + """Bind the shared SQLite session to the generator's scoped-session interface.""" + + registry = scoped_session(lambda: sqlite_session) + monkeypatch.setattr(llm_generator_module.db, "session", registry) + try: + yield sqlite_session + finally: + registry.remove() + + +@pytest.fixture +def recording_model_instance(monkeypatch: pytest.MonkeyPatch) -> Mock: + model_instance = Mock(spec=ModelInstance) + model_instance.invoke_llm.return_value = _llm_result('{"modified": "workflow"}') + model_manager = Mock(spec=ModelManager) + model_manager.get_model_instance.return_value = model_instance + monkeypatch.setattr( + llm_generator_module.ModelManager, + "for_tenant", + Mock(return_value=model_manager), + ) + return model_instance + + +def _llm_result(content: str) -> LLMResult: + return LLMResult( + model="test-model", + message=AssistantPromptMessage(content=content), + usage=LLMUsage.empty_usage(), + ) + + +def _persist_app(database: Session, *, tenant_id: str | None = None) -> App: + app = App( + id=str(uuid4()), + tenant_id=tenant_id or str(uuid4()), + name="Generator app", + description="", + mode=AppMode.WORKFLOW, + icon_type=None, + icon="", + icon_background=None, + enable_site=True, + enable_api=True, + ) + database.add(app) + database.commit() + return app + + +def _persist_message( + database: Session, + app: App, + *, + query: str = "q", + answer: str = "a", + created_at: datetime | None = None, +) -> Message: + message = Message( + id=str(uuid4()), + app_id=app.id, + conversation_id=str(uuid4()), + _inputs={}, + query=query, + message={}, + message_unit_price=Decimal(0), + answer=answer, + answer_unit_price=Decimal(0), + currency="USD", + from_source=ConversationFromSource.API, + error="e", + created_at=created_at, + ) + database.add(message) + database.commit() + return message + + +def _persist_workflow(database: Session, app: App, *, node_type: str | None) -> Workflow: + nodes = [] if node_type is None else [{"id": "node", "data": {"type": node_type}}] + workflow = Workflow.new( + tenant_id=app.tenant_id, + app_id=app.id, + type="workflow", + version=Workflow.VERSION_DRAFT, + graph=json.dumps({"graph": {"nodes": nodes}}), + features="{}", + created_by=str(uuid4()), + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + database.add(workflow) + database.commit() + return workflow + + +def _persist_node_execution( + database: Session, + app: App, + workflow: Workflow, + *, + agent_log: list[dict[str, object]], + inputs: dict[str, object] | None = None, +) -> WorkflowNodeExecutionModel: + execution = WorkflowNodeExecutionModel( + id=str(uuid4()), + tenant_id=app.tenant_id, + app_id=app.id, + workflow_id=workflow.id, + triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP, + workflow_run_id=None, + index=1, + predecessor_node_id=None, + node_execution_id=str(uuid4()), + node_id="node", + node_type="llm", + title="LLM", + inputs=json.dumps(inputs or {}), + process_data=None, + outputs=None, + status=WorkflowNodeExecutionStatus.SUCCEEDED, + error="", + elapsed_time=0.1, + execution_metadata=json.dumps({"agent_log": agent_log}), + created_at=datetime(2026, 1, 1), + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + finished_at=datetime(2026, 1, 1), + ) + database.add(execution) + database.commit() + return execution class TestLLMGenerator: @@ -409,298 +565,296 @@ class TestLLMGenerator: result = LLMGenerator.generate_structured_output("tenant_id", payload) assert "An unexpected error occurred" in result["error"] - def test_instruction_modify_legacy_no_last_run(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session.scalar") as mock_scalar: - mock_scalar.return_value = None + def test_instruction_modify_legacy_without_last_run_uses_real_empty_query( + self, + database: Session, + recording_model_instance: Mock, + model_config_entity: ModelConfig, + ): + app = _persist_app(database) + recording_model_instance.invoke_llm.return_value = _llm_result('{"modified": "prompt"}') - # Mock __instruction_modify_common call via invoke_llm - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = '{"modified": "prompt"}' - mock_model_instance.invoke_llm.return_value = mock_response + result = LLMGenerator.instruction_modify_legacy( + app.tenant_id, + app.id, + "current_val", + "Test {{#last_run#}} and {{#current#}} and {{#error_message#}}", + model_config_entity, + "ideal", + ) - result = LLMGenerator.instruction_modify_legacy( - "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" - ) - assert result == {"modified": "prompt"} - stmt = mock_scalar.call_args.args[0] - compiled = stmt.compile() - statement = str(compiled) - assert "messages.app_id" in statement - assert "apps.tenant_id" in statement - assert "flow_id" in compiled.params.values() - assert "tenant_id" in compiled.params.values() + assert result == {"modified": "prompt"} + user_payload = json.loads(recording_model_instance.invoke_llm.call_args.kwargs["prompt_messages"][1].content) + assert "null" in user_payload["instruction"] + assert "current_val" in user_payload["instruction"] - def test_instruction_modify_legacy_with_last_run(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session.scalar") as mock_scalar: - last_run = MagicMock() - last_run.query = "q" - last_run.answer = "a" - last_run.error = "e" - mock_scalar.return_value = last_run + def test_instruction_modify_legacy_reads_latest_tenant_scoped_message( + self, + database: Session, + recording_model_instance: Mock, + model_config_entity: ModelConfig, + ): + app = _persist_app(database) + _persist_message( + database, + app, + query="older question", + answer="older answer", + created_at=datetime(2026, 1, 1), + ) + _persist_message( + database, + app, + query="latest question", + answer="latest answer", + created_at=datetime(2026, 1, 2), + ) + other_app = _persist_app(database) + _persist_message( + database, + other_app, + query="other tenant question", + created_at=datetime(2026, 1, 3), + ) + recording_model_instance.invoke_llm.return_value = _llm_result('{"modified": "prompt"}') - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = '{"modified": "prompt"}' - mock_model_instance.invoke_llm.return_value = mock_response + result = LLMGenerator.instruction_modify_legacy( + app.tenant_id, app.id, "current", "instruction", model_config_entity, "ideal" + ) - result = LLMGenerator.instruction_modify_legacy( - "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" - ) - assert result == {"modified": "prompt"} - stmt = mock_scalar.call_args.args[0] - compiled = stmt.compile() - statement = str(compiled) - assert "messages.app_id" in statement - assert "apps.tenant_id" in statement - assert "flow_id" in compiled.params.values() - assert "tenant_id" in compiled.params.values() + assert result == {"modified": "prompt"} + user_payload = json.loads(recording_model_instance.invoke_llm.call_args.kwargs["prompt_messages"][1].content) + assert user_payload["last_run"]["query"] == "latest question" + assert user_payload["last_run"]["answer"] == "latest answer" + assert "older question" not in json.dumps(user_payload) + assert "other tenant question" not in json.dumps(user_payload) - def test_instruction_modify_workflow_app_not_found(self): - with patch("extensions.ext_database.db.session") as mock_session: - mock_session.return_value.scalar.return_value = None - with pytest.raises(ValueError, match="App not found."): - LLMGenerator.instruction_modify_workflow("t", "f", "n", "c", "i", MagicMock(), "o", MagicMock()) - stmt = mock_session.return_value.scalar.call_args.args[0] - compiled = stmt.compile() - statement = str(compiled) - assert "apps.id" in statement - assert "apps.tenant_id" in statement - assert "f" in compiled.params.values() - assert "t" in compiled.params.values() + def test_instruction_modify_workflow_app_not_found( + self, + database: Session, + sqlite_session_factory: sessionmaker[Session], + model_config_entity: ModelConfig, + ): + workflow_service = WorkflowService(sqlite_session_factory) - def test_instruction_modify_workflow_no_workflow(self): - with patch("extensions.ext_database.db.session") as mock_session: - mock_session.return_value.scalar.return_value = MagicMock() - workflow_service = MagicMock() - workflow_service.get_draft_workflow.return_value = None - with pytest.raises(ValueError, match="Workflow not found for the given app model."): - LLMGenerator.instruction_modify_workflow("t", "f", "n", "c", "i", MagicMock(), "o", workflow_service) - - def test_instruction_modify_workflow_success(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session") as mock_session: - mock_session.return_value.scalar.return_value = MagicMock() - workflow = MagicMock() - workflow.graph_dict = {"graph": {"nodes": [{"id": "node_id", "data": {"type": "llm"}}]}} - - workflow_service = MagicMock() - workflow_service.get_draft_workflow.return_value = workflow - - last_run = MagicMock() - last_run.node_type = "llm" - last_run.status = "s" - last_run.error = "e" - # Return regular values, not Mocks - last_run.execution_metadata_dict = {"agent_log": [{"status": "s", "error": "e", "data": {}}]} - last_run.load_full_inputs.return_value = {"in": "val"} - - workflow_service.get_node_last_run.return_value = last_run - - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = '{"modified": "workflow"}' - mock_model_instance.invoke_llm.return_value = mock_response - - result = LLMGenerator.instruction_modify_workflow( - "tenant_id", - "flow_id", - "node_id", + with pytest.raises(ValueError, match="App not found"): + LLMGenerator.instruction_modify_workflow( + str(uuid4()), + str(uuid4()), + "node", "current", "instruction", model_config_entity, "ideal", workflow_service, ) - assert result == {"modified": "workflow"} - def test_instruction_modify_workflow_no_last_run_fallback(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session") as mock_session: - mock_session.return_value.scalar.return_value = MagicMock() - workflow = MagicMock() - workflow.graph_dict = {"graph": {"nodes": [{"id": "node_id", "data": {"type": "code"}}]}} + def test_instruction_modify_workflow_rejects_app_from_another_tenant( + self, + database: Session, + sqlite_session_factory: sessionmaker[Session], + model_config_entity: ModelConfig, + ): + app = _persist_app(database) + workflow_service = WorkflowService(sqlite_session_factory) - workflow_service = MagicMock() - workflow_service.get_draft_workflow.return_value = workflow - workflow_service.get_node_last_run.return_value = None - - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = '{"modified": "fallback"}' - mock_model_instance.invoke_llm.return_value = mock_response - - result = LLMGenerator.instruction_modify_workflow( - "tenant_id", - "flow_id", - "node_id", + with pytest.raises(ValueError, match="App not found"): + LLMGenerator.instruction_modify_workflow( + str(uuid4()), + app.id, + "node", "current", "instruction", model_config_entity, "ideal", workflow_service, ) - assert result == {"modified": "fallback"} - def test_instruction_modify_workflow_node_type_fallback(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session") as mock_session: - mock_session.return_value.scalar.return_value = MagicMock() - workflow = MagicMock() - # Cause exception in node_type logic - workflow.graph_dict = {"graph": {"nodes": []}} + def test_instruction_modify_workflow_requires_draft_workflow( + self, + database: Session, + sqlite_session_factory: sessionmaker[Session], + model_config_entity: ModelConfig, + ): + app = _persist_app(database) + workflow_service = WorkflowService(sqlite_session_factory) - workflow_service = MagicMock() - workflow_service.get_draft_workflow.return_value = workflow - workflow_service.get_node_last_run.return_value = None - - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = '{"modified": "fallback"}' - mock_model_instance.invoke_llm.return_value = mock_response - - result = LLMGenerator.instruction_modify_workflow( - "tenant_id", - "flow_id", - "node_id", + with pytest.raises(ValueError, match="Workflow not found"): + LLMGenerator.instruction_modify_workflow( + app.tenant_id, + app.id, + "node", "current", "instruction", model_config_entity, "ideal", workflow_service, ) - assert result == {"modified": "fallback"} - def test_instruction_modify_workflow_empty_agent_log(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session") as mock_session: - mock_session.return_value.scalar.return_value = MagicMock() - workflow = MagicMock() - workflow.graph_dict = {"graph": {"nodes": [{"id": "node_id", "data": {"type": "llm"}}]}} + def test_instruction_modify_workflow_uses_last_run( + self, + database: Session, + sqlite_session_factory: sessionmaker[Session], + recording_model_instance: Mock, + model_config_entity: ModelConfig, + ): + app = _persist_app(database) + workflow = _persist_workflow(database, app, node_type="llm") + _persist_node_execution( + database, + app, + workflow, + inputs={"input": "value"}, + agent_log=[{"status": "s", "error": "", "data": {"step": 1}}], + ) + workflow_service = WorkflowService(sqlite_session_factory) - workflow_service = MagicMock() - workflow_service.get_draft_workflow.return_value = workflow + result = LLMGenerator.instruction_modify_workflow( + app.tenant_id, + app.id, + "node", + "current", + "instruction", + model_config_entity, + "ideal", + workflow_service, + ) - last_run = MagicMock() - last_run.node_type = "llm" - last_run.status = "s" - last_run.error = "e" - # Return regular empty list, not a Mock - last_run.execution_metadata_dict = {"agent_log": []} - last_run.load_full_inputs.return_value = {} + assert result == {"modified": "workflow"} + user_payload = json.loads(recording_model_instance.invoke_llm.call_args.kwargs["prompt_messages"][1].content) + assert user_payload["last_run"]["inputs"] == {"input": "value"} + assert user_payload["last_run"]["agent_log"][0]["data"] == {"step": 1} - workflow_service.get_node_last_run.return_value = last_run + def test_instruction_modify_workflow_accepts_empty_agent_log( + self, + database: Session, + sqlite_session_factory: sessionmaker[Session], + recording_model_instance: Mock, + model_config_entity: ModelConfig, + ): + app = _persist_app(database) + workflow = _persist_workflow(database, app, node_type="llm") + _persist_node_execution(database, app, workflow, agent_log=[]) + workflow_service = WorkflowService(sqlite_session_factory) - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = '{"modified": "workflow"}' - mock_model_instance.invoke_llm.return_value = mock_response + result = LLMGenerator.instruction_modify_workflow( + app.tenant_id, + app.id, + "node", + "current", + "instruction", + model_config_entity, + "ideal", + workflow_service, + ) - result = LLMGenerator.instruction_modify_workflow( - "tenant_id", - "flow_id", - "node_id", - "current", - "instruction", - model_config_entity, - "ideal", - workflow_service, - ) - assert result == {"modified": "workflow"} + assert result == {"modified": "workflow"} + user_payload = json.loads(recording_model_instance.invoke_llm.call_args.kwargs["prompt_messages"][1].content) + assert user_payload["last_run"]["agent_log"] == [] - def test_instruction_modify_common_placeholders(self, mock_model_instance, model_config_entity): - # Testing placeholders replacement via instruction_modify_legacy for convenience - with patch("extensions.ext_database.db.session.scalar") as mock_scalar: - mock_scalar.return_value = None + @pytest.mark.parametrize( + "node_type", + [ + "code", + None, + ], + ) + def test_instruction_modify_workflow_falls_back_without_last_run( + self, + database: Session, + sqlite_session_factory: sessionmaker[Session], + recording_model_instance: Mock, + model_config_entity: ModelConfig, + node_type: str | None, + ): + app = _persist_app(database) + _persist_workflow(database, app, node_type=node_type) + workflow_service = WorkflowService(sqlite_session_factory) + recording_model_instance.invoke_llm.return_value = _llm_result('{"modified": "fallback"}') - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = '{"ok": true}' - mock_model_instance.invoke_llm.return_value = mock_response + result = LLMGenerator.instruction_modify_workflow( + app.tenant_id, + app.id, + "node", + "current", + "instruction", + model_config_entity, + "ideal", + workflow_service, + ) - instruction = "Test {{#last_run#}} and {{#current#}} and {{#error_message#}}" - LLMGenerator.instruction_modify_legacy( - "tenant_id", "flow_id", "current_val", instruction, model_config_entity, "ideal" - ) + assert result == {"modified": "fallback"} - # Verify the call to invoke_llm contains replaced instruction - args, kwargs = mock_model_instance.invoke_llm.call_args - prompt_messages = kwargs["prompt_messages"] - user_msg = prompt_messages[1].content - user_msg_dict = json.loads(user_msg) - assert "null" in user_msg_dict["instruction"] # because last_run is None and current is current_val etc. - assert "current_val" in user_msg_dict["instruction"] + def test_instruction_modify_workflow_falls_back_for_unknown_node_type( + self, + database: Session, + sqlite_session_factory: sessionmaker[Session], + recording_model_instance: Mock, + model_config_entity: ModelConfig, + ): + app = _persist_app(database) + _persist_workflow(database, app, node_type="unknown") + workflow_service = WorkflowService(sqlite_session_factory) - def test_instruction_modify_common_no_braces(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session.scalar") as mock_scalar: - mock_scalar.return_value = None - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = "No braces here" - mock_model_instance.invoke_llm.return_value = mock_response - result = LLMGenerator.instruction_modify_legacy( - "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" - ) - assert "An unexpected error occurred" in result["error"] - assert "Could not find a valid JSON object" in result["error"] + result = LLMGenerator.instruction_modify_workflow( + app.tenant_id, + app.id, + "node", + "current", + "instruction", + model_config_entity, + "ideal", + workflow_service, + ) - def test_instruction_modify_common_not_dict(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session.scalar") as mock_scalar: - mock_scalar.return_value = None - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = "[1, 2, 3]" - mock_model_instance.invoke_llm.return_value = mock_response - result = LLMGenerator.instruction_modify_legacy( - "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" - ) - # The exception message is "Expected a JSON object, but got list" - assert "An unexpected error occurred" in result["error"] + assert result == {"modified": "workflow"} + system_prompt = recording_model_instance.invoke_llm.call_args.kwargs["prompt_messages"][0].content + assert system_prompt == llm_generator_module.LLM_MODIFY_PROMPT_SYSTEM - def test_instruction_modify_common_other_node_type(self, mock_model_instance, model_config_entity): - with patch("core.llm_generator.llm_generator.ModelManager.for_tenant") as mock_manager: - instance = MagicMock() - mock_manager.return_value.get_model_instance.return_value = instance - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = '{"ok": true}' - instance.invoke_llm.return_value = mock_response + @pytest.mark.parametrize( + ("raw_output", "error_fragment"), + [ + ("No braces here", "Could not find a valid JSON object"), + ("[1, 2, 3]", "Could not find a valid JSON object"), + ], + ) + def test_instruction_modify_rejects_invalid_model_output( + self, + database: Session, + mock_model_instance: MagicMock, + model_config_entity: ModelConfig, + raw_output: str, + error_fragment: str, + ): + app = _persist_app(database) + response = MagicMock() + response.message.get_text_content.return_value = raw_output + mock_model_instance.invoke_llm.return_value = response - with patch("extensions.ext_database.db.session") as mock_session: - mock_session.return_value.scalar.return_value = MagicMock() - workflow = MagicMock() - workflow.graph_dict = {"graph": {"nodes": [{"id": "node_id", "data": {"type": "other"}}]}} + result = LLMGenerator.instruction_modify_legacy( + app.tenant_id, app.id, "current", "instruction", model_config_entity, "ideal" + ) - workflow_service = MagicMock() - workflow_service.get_draft_workflow.return_value = workflow - workflow_service.get_node_last_run.return_value = None + assert "An unexpected error occurred" in result["error"] + assert error_fragment in result["error"] - LLMGenerator.instruction_modify_workflow( - "tenant_id", - "flow_id", - "node_id", - "current", - "instruction", - model_config_entity, - "ideal", - workflow_service, - ) + @pytest.mark.parametrize( + ("model_error", "error_fragment"), + [(InvokeError("invoke failed"), "Failed to generate code"), (RuntimeError("boom"), "unexpected error")], + ) + def test_instruction_modify_handles_model_errors( + self, + database: Session, + mock_model_instance: MagicMock, + model_config_entity: ModelConfig, + model_error: Exception, + error_fragment: str, + ): + app = _persist_app(database) + mock_model_instance.invoke_llm.side_effect = model_error - def test_instruction_modify_common_invoke_error(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session.scalar") as mock_scalar: - mock_scalar.return_value = None - mock_model_instance.invoke_llm.side_effect = InvokeError("Invoke Failed") + result = LLMGenerator.instruction_modify_legacy( + app.tenant_id, app.id, "current", "instruction", model_config_entity, "ideal" + ) - result = LLMGenerator.instruction_modify_legacy( - "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" - ) - assert "Failed to generate code" in result["error"] - - def test_instruction_modify_common_exception(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session.scalar") as mock_scalar: - mock_scalar.return_value = None - mock_model_instance.invoke_llm.side_effect = Exception("Random error") - - result = LLMGenerator.instruction_modify_legacy( - "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" - ) - assert "An unexpected error occurred" in result["error"] - - def test_instruction_modify_common_json_error(self, mock_model_instance, model_config_entity): - with patch("extensions.ext_database.db.session.scalar") as mock_scalar: - mock_scalar.return_value = None - - mock_response = MagicMock() - mock_response.message.get_text_content.return_value = "No JSON here" - mock_model_instance.invoke_llm.return_value = mock_response - - result = LLMGenerator.instruction_modify_legacy( - "tenant_id", "flow_id", "current", "instruction", model_config_entity, "ideal" - ) - assert "An unexpected error occurred" in result["error"] + assert error_fragment.lower() in result["error"].lower() diff --git a/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py b/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py index d193edc9fd5..fc890f778c9 100644 --- a/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py +++ b/api/tests/unit_tests/core/llm_generator/test_llm_generator_missing.py @@ -1,7 +1,57 @@ -import sys +from datetime import datetime, timedelta from unittest.mock import MagicMock, patch +import pytest +from sqlalchemy import event +from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy.orm import Session, sessionmaker + +import core.llm_generator.llm_generator as generator_module from core.llm_generator.llm_generator import LLMGenerator, _parse_string_list +from core.model_manager import ModelInstance, ModelManager +from core.workflow.generator import tool_catalogue as tool_catalogue_module +from core.workflow.generator.tool_catalogue import ToolCatalogueEntry +from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage +from graphon.model_runtime.entities.message_entities import AssistantPromptMessage +from models.dataset import Dataset +from services.workflow_service import WorkflowService + + +@pytest.fixture +def dataset_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Bind the real SQLite session to the production database extension.""" + + monkeypatch.setattr(generator_module.db, "session", sqlite_session) + return sqlite_session + + +def _llm_result(content: str) -> LLMResult: + """Build a real non-streaming LLM response around deterministic test content.""" + + return LLMResult( + model="test-model", + message=AssistantPromptMessage(content=content), + usage=LLMUsage.empty_usage(), + ) + + +def _model_manager() -> tuple[MagicMock, MagicMock]: + """Build spec-constrained mocks for the model-manager boundary and its default model.""" + + model_manager = MagicMock(spec=ModelManager) + model_instance = MagicMock(spec=ModelInstance) + model_manager.get_default_model_instance.return_value = model_instance + return model_manager, model_instance + + +def _dataset(*, dataset_id: str, tenant_id: str, name: str, created_at: datetime) -> Dataset: + return Dataset( + id=dataset_id, + tenant_id=tenant_id, + name=name, + created_by="account-id", + created_at=created_at, + ) class TestParseStringList: @@ -34,95 +84,115 @@ class TestParseStringList: class TestGenerateWorkflowInstructionSuggestions: @patch("core.llm_generator.llm_generator.ModelManager.for_tenant") def test_no_default_model(self, mock_for_tenant): - mock_for_tenant.return_value.get_default_model_instance.side_effect = Exception("No model") + model_manager, _ = _model_manager() + model_manager.get_default_model_instance.side_effect = RuntimeError("no default model") + mock_for_tenant.return_value = model_manager + assert LLMGenerator.generate_workflow_instruction_suggestions("tenant", mode="workflow") == [] @patch("core.llm_generator.llm_generator.ModelManager.for_tenant") @patch("core.llm_generator.llm_generator.LLMGenerator._build_suggestion_context") def test_llm_success(self, mock_build_context, mock_for_tenant): mock_build_context.return_value = "context" - - mock_model = MagicMock() - mock_model.invoke_llm.return_value = MagicMock() - mock_model.invoke_llm.return_value.message.get_text_content.return_value = '["idea 1", "idea 2"]' - - mock_for_tenant.return_value.get_default_model_instance.return_value = mock_model + model_manager, model_instance = _model_manager() + model_instance.invoke_llm.return_value = _llm_result('["idea 1", "idea 2"]') + mock_for_tenant.return_value = model_manager result = LLMGenerator.generate_workflow_instruction_suggestions("tenant", mode="workflow") assert result == ["idea 1", "idea 2"] + model_instance.invoke_llm.assert_called_once() @patch("core.llm_generator.llm_generator.ModelManager.for_tenant") @patch("core.llm_generator.llm_generator.LLMGenerator._build_suggestion_context") def test_llm_error(self, mock_build_context, mock_for_tenant): mock_build_context.return_value = "context" + model_manager, model_instance = _model_manager() + model_instance.invoke_llm.side_effect = RuntimeError("API error") + mock_for_tenant.return_value = model_manager - mock_model = MagicMock() - mock_model.invoke_llm.side_effect = Exception("API error") - - mock_for_tenant.return_value.get_default_model_instance.return_value = mock_model - - assert LLMGenerator.generate_workflow_instruction_suggestions("tenant", mode="workflow") == [] + result = LLMGenerator.generate_workflow_instruction_suggestions("tenant", mode="workflow") + assert result == [] + model_instance.invoke_llm.assert_called_once() @patch("core.llm_generator.llm_generator.ModelManager.for_tenant") @patch("core.llm_generator.llm_generator.LLMGenerator._build_suggestion_context") def test_llm_bad_output(self, mock_build_context, mock_for_tenant): mock_build_context.return_value = "context" + model_manager, model_instance = _model_manager() + model_instance.invoke_llm.return_value = _llm_result("Not a list") + mock_for_tenant.return_value = model_manager - mock_model = MagicMock() - mock_model.invoke_llm.return_value = MagicMock() - mock_model.invoke_llm.return_value.message.get_text_content.return_value = "Not a list" - - mock_for_tenant.return_value.get_default_model_instance.return_value = mock_model - - assert LLMGenerator.generate_workflow_instruction_suggestions("tenant", mode="workflow") == [] + result = LLMGenerator.generate_workflow_instruction_suggestions("tenant", mode="workflow") + assert result == [] + model_instance.invoke_llm.assert_called_once() +@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) class TestBuildSuggestionContext: - @patch("core.llm_generator.llm_generator.db.session.scalars") - def test_both_success(self, mock_scalars, monkeypatch): - mock_scalars.return_value.all.return_value = ["kb1", "kb2"] + def test_both_success(self, dataset_session: Session, monkeypatch: pytest.MonkeyPatch): + now = datetime.now() + dataset_session.add_all( + ( + _dataset(dataset_id="kb-1", tenant_id="tenant", name="kb1", created_at=now), + _dataset( + dataset_id="kb-2", + tenant_id="tenant", + name="kb2", + created_at=now - timedelta(seconds=1), + ), + _dataset(dataset_id="other-kb", tenant_id="other", name="private", created_at=now), + ) + ) + dataset_session.commit() - # ``_build_suggestion_context`` imports the tool catalogue lazily, so we - # stub the module in ``sys.modules``. Use ``monkeypatch.setitem`` so the - # ORIGINAL module is RESTORED on teardown — a bare ``del`` would evict it - # from sys.modules entirely, after which a sibling test that imported - # ``build_tool_catalogue`` at collection time (e.g. test_tool_catalogue) - # diverges from a freshly re-imported module and its @patch targets stop - # applying, silently breaking it under xdist. - mock_tool_catalogue = MagicMock() - mock_tool_catalogue.build_tool_catalogue.return_value = "catalog" - mock_tool_catalogue.format_tool_catalogue.return_value = "tool1\ntool2" - monkeypatch.setitem(sys.modules, "core.workflow.generator.tool_catalogue", mock_tool_catalogue) + def build_tool_catalogue(_tenant_id: str) -> list[ToolCatalogueEntry]: + return [ + ToolCatalogueEntry( + provider_name="provider", + provider_type="builtin", + plugin_id="", + tool_name="tool1", + tool_label="tool1", + description="First tool", + ), + ToolCatalogueEntry( + provider_name="provider", + provider_type="builtin", + plugin_id="", + tool_name="tool2", + tool_label="tool2", + description="Second tool", + ), + ] + + # Keep the real module and formatter; only isolate provider/plugin discovery. + monkeypatch.setattr(tool_catalogue_module, "build_tool_catalogue", build_tool_catalogue) result = LLMGenerator._build_suggestion_context("tenant") assert "Knowledge bases:\n- kb1\n- kb2" in result - assert "Installed tools:\ntool1\ntool2" in result + assert "Installed tools:\n- provider/tool1 — First tool\n- provider/tool2 — Second tool" in result - @patch("core.llm_generator.llm_generator.db.session.scalars") - def test_both_fail(self, mock_scalars, monkeypatch): - mock_scalars.side_effect = Exception("DB error") + def test_both_fail(self, dataset_session: Session, monkeypatch: pytest.MonkeyPatch): + def fail_query(_orm_execute_state: object) -> None: + raise SQLAlchemyError("DB error") - # See ``test_both_success``: restore the original module via monkeypatch - # rather than ``del``-ing it, so we don't evict it for sibling tests. - mock_tool_catalogue = MagicMock() - mock_tool_catalogue.build_tool_catalogue.side_effect = Exception("Tool error") - monkeypatch.setitem(sys.modules, "core.workflow.generator.tool_catalogue", mock_tool_catalogue) + def fail_tool_catalogue(_tenant_id: str) -> list[ToolCatalogueEntry]: + raise RuntimeError("Tool error") - assert LLMGenerator._build_suggestion_context("tenant") == "" + event.listen(dataset_session, "do_orm_execute", fail_query) + monkeypatch.setattr(tool_catalogue_module, "build_tool_catalogue", fail_tool_catalogue) + + try: + assert LLMGenerator._build_suggestion_context("tenant") == "" + finally: + event.remove(dataset_session, "do_orm_execute", fail_query) class TestWorkflowServiceInterface: - def test_protocol_methods(self): - # Just to cover the 'pass' statements in the Protocol definition + def test_real_workflow_service_exposes_protocol_methods(self): from core.llm_generator.llm_generator import WorkflowServiceInterface - class MockService(WorkflowServiceInterface): - def get_draft_workflow(self, app_model, workflow_id=None, *, session): - return super().get_draft_workflow(app_model, workflow_id, session=session) + service: WorkflowServiceInterface = WorkflowService(sessionmaker()) - def get_node_last_run(self, app_model, workflow, node_id): - return super().get_node_last_run(app_model, workflow, node_id) - - service = MockService() - service.get_draft_workflow(None, session=None) - service.get_node_last_run(None, None, "node") + assert callable(service.get_draft_workflow) + assert callable(service.get_node_last_run) diff --git a/api/tests/unit_tests/core/mcp/server/test_streamable_http.py b/api/tests/unit_tests/core/mcp/server/test_streamable_http.py index cc00e02252e..75b01a1397b 100644 --- a/api/tests/unit_tests/core/mcp/server/test_streamable_http.py +++ b/api/tests/unit_tests/core/mcp/server/test_streamable_http.py @@ -29,15 +29,22 @@ class TestHandleMCPRequest: def setup_method(self): """Setup test fixtures""" - self.app = Mock(spec=App) + self.app = App() self.app.name = "test_app" self.app.mode = AppMode.CHAT - self.mcp_server = Mock(spec=AppMCPServer) + self.mcp_server = AppMCPServer( + tenant_id="tenant-id", + app_id="app-id", + name="Test Server", + description="", + server_code="test-server", + status="active", + parameters="{}", + ) self.mcp_server.description = "Test server" - self.mcp_server.parameters_dict = {} - self.end_user = Mock(spec=EndUser) + self.end_user = EndUser() self.user_input_form = [] # Create mock request @@ -336,8 +343,9 @@ class TestIndividualHandlers: @patch("core.mcp.server.streamable_http.AppGenerateService") def test_handle_call_tool(self, mock_app_generate): """Test call tool handler""" - app = Mock(spec=App) - app.mode = AppMode.CHAT + app = App( + mode=AppMode.CHAT, + ) # Create mock request mock_request = Mock() @@ -347,7 +355,7 @@ class TestIndividualHandlers: mock_request.root = mock_call_request user_input_form: list[VariableEntity] = [] - end_user = Mock(spec=EndUser) + end_user = EndUser() # Mock app generate service response mock_response = {"answer": "test answer"} @@ -365,8 +373,9 @@ class TestIndividualHandlers: @patch("core.mcp.server.streamable_http.AppGenerateService") def test_handle_call_tool_structured_output_modern_client(self, mock_app_generate): """structuredContent is attached alongside TextContent for >= 2025-06-18.""" - app = Mock(spec=App) - app.mode = AppMode.CHAT + app = App( + mode=AppMode.CHAT, + ) mock_request = Mock() mock_call_request = Mock(spec=types.CallToolRequest) @@ -376,7 +385,7 @@ class TestIndividualHandlers: mock_app_generate.generate.return_value = {"answer": "test answer"} - result = handle_call_tool(Mock(), app, mock_request, [], Mock(spec=EndUser), "2025-06-18") + result = handle_call_tool(Mock(), app, mock_request, [], EndUser(), "2025-06-18") assert result.structuredContent == {"answer": "test answer"} assert result.content[0].text == "test answer" @@ -384,8 +393,9 @@ class TestIndividualHandlers: @patch("core.mcp.server.streamable_http.AppGenerateService") def test_handle_call_tool_no_structured_output_legacy_client(self, mock_app_generate): """structuredContent is omitted for 2024-11-05 clients.""" - app = Mock(spec=App) - app.mode = AppMode.CHAT + app = App( + mode=AppMode.CHAT, + ) mock_request = Mock() mock_call_request = Mock(spec=types.CallToolRequest) @@ -395,14 +405,14 @@ class TestIndividualHandlers: mock_app_generate.generate.return_value = {"answer": "test answer"} - result = handle_call_tool(Mock(), app, mock_request, [], Mock(spec=EndUser), "2024-11-05") + result = handle_call_tool(Mock(), app, mock_request, [], EndUser(), "2024-11-05") assert result.structuredContent is None assert result.content[0].text == "test answer" def test_handle_call_tool_no_end_user(self): """Test call tool handler without end user""" - app = Mock(spec=App) + app = App() mock_request = Mock() user_input_form: list[VariableEntity] = [] @@ -460,8 +470,9 @@ class TestUtilityFunctions: def test_prepare_tool_arguments_chat_mode(self): """Test preparing tool arguments for chat mode""" - app = Mock(spec=App) - app.mode = AppMode.CHAT + app = App( + mode=AppMode.CHAT, + ) arguments = {"query": "test question", "name": "John"} @@ -474,8 +485,9 @@ class TestUtilityFunctions: def test_prepare_tool_arguments_workflow_mode(self): """Test preparing tool arguments for workflow mode""" - app = Mock(spec=App) - app.mode = AppMode.WORKFLOW + app = App( + mode=AppMode.WORKFLOW, + ) arguments = {"input_text": "test input"} @@ -486,8 +498,9 @@ class TestUtilityFunctions: def test_prepare_tool_arguments_completion_mode(self): """Test preparing tool arguments for completion mode""" - app = Mock(spec=App) - app.mode = AppMode.COMPLETION + app = App( + mode=AppMode.COMPLETION, + ) arguments = {"name": "John"} @@ -498,8 +511,9 @@ class TestUtilityFunctions: def test_extract_answer_from_mapping_response_chat(self): """Test extracting answer from mapping response for chat mode""" - app = Mock(spec=App) - app.mode = AppMode.CHAT + app = App( + mode=AppMode.CHAT, + ) response = {"answer": "test answer", "other": "data"} @@ -509,8 +523,9 @@ class TestUtilityFunctions: def test_extract_answer_from_mapping_response_workflow(self): """Test extracting answer from mapping response for workflow mode""" - app = Mock(spec=App) - app.mode = AppMode.WORKFLOW + app = App( + mode=AppMode.WORKFLOW, + ) response = {"data": {"outputs": {"result": "test result"}}} @@ -521,7 +536,7 @@ class TestUtilityFunctions: def test_extract_answer_from_streaming_response(self): """Test extracting answer from streaming response""" - app = Mock(spec=App) + app = App() # Mock RateLimitGenerator mock_generator = Mock(spec=RateLimitGenerator) @@ -538,8 +553,9 @@ class TestUtilityFunctions: def test_extract_structured_output_workflow(self): """Workflow mode exposes the raw outputs mapping as structured content.""" - app = Mock(spec=App) - app.mode = AppMode.WORKFLOW + app = App( + mode=AppMode.WORKFLOW, + ) response = {"data": {"outputs": {"result": "test result"}}} @@ -547,58 +563,66 @@ class TestUtilityFunctions: def test_extract_structured_output_chat(self): """Chat mode wraps the answer string under an 'answer' key.""" - app = Mock(spec=App) - app.mode = AppMode.CHAT + app = App( + mode=AppMode.CHAT, + ) assert extract_structured_output(app, {"answer": "hi"}, "hi") == {"answer": "hi"} def test_extract_structured_output_workflow_missing_outputs(self): """Missing or malformed outputs fall back to None.""" - app = Mock(spec=App) - app.mode = AppMode.WORKFLOW + app = App( + mode=AppMode.WORKFLOW, + ) assert extract_structured_output(app, {"data": {}}, "ignored") is None def test_extract_structured_output_workflow_non_mapping_response(self): """A non-mapping workflow response yields no structured output.""" - app = Mock(spec=App) - app.mode = AppMode.WORKFLOW + app = App( + mode=AppMode.WORKFLOW, + ) assert extract_structured_output(app, None, "ignored") is None def test_extract_structured_output_workflow_non_mapping_data(self): """A non-mapping 'data' entry yields no structured output.""" - app = Mock(spec=App) - app.mode = AppMode.WORKFLOW + app = App( + mode=AppMode.WORKFLOW, + ) assert extract_structured_output(app, {"data": "not a mapping"}, "ignored") is None def test_extract_structured_output_workflow_non_mapping_outputs(self): """A non-mapping 'outputs' entry yields no structured output.""" - app = Mock(spec=App) - app.mode = AppMode.WORKFLOW + app = App( + mode=AppMode.WORKFLOW, + ) assert extract_structured_output(app, {"data": {"outputs": ["not", "a", "mapping"]}}, "ignored") is None @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.AGENT_CHAT, AppMode.COMPLETION]) def test_extract_structured_output_other_answer_modes(self, mode): """Every chat-style mode wraps the answer string under an 'answer' key.""" - app = Mock(spec=App) - app.mode = mode + app = App( + mode=mode, + ) assert extract_structured_output(app, {"answer": "hi"}, "hi") == {"answer": "hi"} def test_extract_structured_output_unknown_mode(self): """Modes outside the MCP surface produce no structured output.""" - app = Mock(spec=App) - app.mode = AppMode.CHANNEL + app = App( + mode=AppMode.CHANNEL, + ) assert extract_structured_output(app, {"answer": "hi"}, "hi") is None def test_process_mapping_response_invalid_mode(self): """Test processing mapping response with invalid app mode""" - app = Mock(spec=App) - app.mode = "invalid_mode" + app = App( + mode="invalid_mode", + ) response = {"answer": "test"} diff --git a/api/tests/unit_tests/core/memory/test_token_buffer_memory.py b/api/tests/unit_tests/core/memory/test_token_buffer_memory.py index 007486f3c34..db5498bddb4 100644 --- a/api/tests/unit_tests/core/memory/test_token_buffer_memory.py +++ b/api/tests/unit_tests/core/memory/test_token_buffer_memory.py @@ -1,11 +1,19 @@ -"""Comprehensive unit tests for core/memory/token_buffer_memory.py""" +"""Comprehensive SQLite-backed tests for token-buffer memory.""" +from collections.abc import Iterator +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from decimal import Decimal from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session +from core.memory import token_buffer_memory as memory_module from core.memory.token_buffer_memory import TokenBufferMemory +from graphon.file import FileTransferMethod, FileType from graphon.model_runtime.entities import ( AssistantPromptMessage, ImagePromptMessageContent, @@ -13,13 +21,44 @@ from graphon.model_runtime.entities import ( TextPromptMessageContent, UserPromptMessage, ) -from models.model import AppMode +from models.base import TypeBase +from models.enums import ConversationFromSource, CreatorUserRole, MessageFileBelongsTo +from models.model import AppMode, Message, MessageFile +from models.workflow import Workflow, WorkflowType # --------------------------------------------------------------------------- # Helpers / shared fixtures # --------------------------------------------------------------------------- +@dataclass(frozen=True) +class Database: + """Typed SQLite binding plus executed SQL for query-count assertions.""" + + engine: Engine + session: Session + statements: list[tuple[str, object]] + + +@pytest.fixture +def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Database]: + TypeBase.metadata.create_all( + sqlite_engine, + tables=[Message.__table__, MessageFile.__table__, Workflow.__table__], + ) + statements: list[tuple[str, object]] = [] + + def record_statement(_connection, _cursor, statement, parameters, _context, _executemany) -> None: + statements.append((statement, parameters)) + + event.listen(sqlite_engine, "before_cursor_execute", record_statement) + with Session(sqlite_engine, expire_on_commit=False) as session: + database = Database(engine=sqlite_engine, session=session, statements=statements) + monkeypatch.setattr(memory_module, "db", database) + yield database + event.remove(sqlite_engine, "before_cursor_execute", record_statement) + + def _make_conversation(mode: AppMode = AppMode.CHAT) -> MagicMock: """Return a minimal Conversation mock.""" conv = MagicMock() @@ -47,6 +86,73 @@ def _make_message(answer: str = "hello", answer_tokens: int = 5) -> MagicMock: return msg +def _persist_message( + database: Database, + conversation_id: str, + *, + query: str = "user query", + answer: str = "hello", + answer_tokens: int = 5, + created_at: datetime | None = None, + workflow_run_id: str | None = None, +) -> Message: + message = Message( + id=str(uuid4()), + app_id="app-1", + conversation_id=conversation_id, + _inputs={}, + query=query, + message={}, + message_unit_price=Decimal(0), + answer=answer, + answer_tokens=answer_tokens, + answer_unit_price=Decimal(0), + currency="USD", + from_source=ConversationFromSource.API, + workflow_run_id=workflow_run_id, + created_at=created_at or datetime.now(UTC).replace(tzinfo=None), + ) + database.session.add(message) + database.session.commit() + return message + + +def _persist_message_file( + database: Database, + message: Message, + *, + belongs_to: MessageFileBelongsTo | None, +) -> MessageFile: + message_file = MessageFile( + message_id=message.id, + type=FileType.IMAGE, + transfer_method=FileTransferMethod.REMOTE_URL, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + belongs_to=belongs_to, + url="https://example.com/image.png", + ) + database.session.add(message_file) + database.session.commit() + return message_file + + +def _persist_workflow(database: Database, *, workflow_id: str) -> Workflow: + workflow = Workflow( + id=workflow_id, + tenant_id="tenant-1", + app_id="app-1", + type=WorkflowType.CHAT, + version="1", + graph="{}", + features="{}", + created_by="account-1", + ) + database.session.add(workflow) + database.session.commit() + return workflow + + # =========================================================================== # Tests for __init__ and workflow_run_repo property # =========================================================================== @@ -61,25 +167,25 @@ class TestInit: assert mem.model_instance is mi assert mem._workflow_run_repo is None - def test_workflow_run_repo_is_created_lazily(self): + def test_workflow_run_repo_is_created_lazily(self, database: Database): conv = _make_conversation() mi = _make_model_instance() mem = TokenBufferMemory(conversation=conv, model_instance=mi) mock_repo = MagicMock() - with ( - patch("core.memory.token_buffer_memory.sessionmaker") as mock_sm, - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.DifyAPIRepositoryFactory.create_api_workflow_run_repository", - return_value=mock_repo, - ), - ): - mock_db.engine = MagicMock() + with patch( + "core.memory.token_buffer_memory.DifyAPIRepositoryFactory.create_api_workflow_run_repository", + return_value=mock_repo, + ) as repository_factory: repo = mem.workflow_run_repo assert repo is mock_repo assert mem._workflow_run_repo is mock_repo + session_factory = repository_factory.call_args.args[0] + with session_factory() as session: + assert isinstance(session, Session) + assert session.get_bind() is database.engine + def test_workflow_run_repo_cached_after_first_access(self): conv = _make_conversation() mi = _make_model_instance() @@ -410,7 +516,7 @@ class TestBuildPromptMessageWithFiles: ) @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) - def test_workflow_mode_workflow_not_found_raises(self, mode): + def test_workflow_mode_workflow_not_found_raises(self, mode, database: Database): """Raises ValueError when Workflow lookup returns None.""" conv = _make_conversation(mode) conv.app = MagicMock() @@ -422,22 +528,17 @@ class TestBuildPromptMessageWithFiles: mem._workflow_run_repo = MagicMock() mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - ): - mock_db.session.scalar.return_value = None # workflow not found - - with pytest.raises(ValueError, match="Workflow not found"): - mem._build_prompt_message_with_files( - message_files=[], - text_content="text", - message=_make_message(), - app_record=MagicMock(), - is_user_message=True, - ) + with pytest.raises(ValueError, match="Workflow not found"): + mem._build_prompt_message_with_files( + message_files=[], + text_content="text", + message=_make_message(), + app_record=MagicMock(), + is_user_message=True, + ) @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) - def test_workflow_mode_success_no_files_user(self, mode): + def test_workflow_mode_success_no_files_user(self, mode, database: Database): """Happy path: workflow mode, no message files → plain UserPromptMessage.""" conv = _make_conversation(mode) conv.app = MagicMock() @@ -445,22 +546,16 @@ class TestBuildPromptMessageWithFiles: mock_workflow_run = MagicMock() mock_workflow_run.workflow_id = str(uuid4()) - mock_workflow = MagicMock() - mock_workflow.features_dict = {} + workflow = _persist_workflow(database, workflow_id=mock_workflow_run.workflow_id) mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance()) mem._workflow_run_repo = MagicMock() mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.FileUploadConfigManager.convert", - return_value=None, - ), + with patch( + "core.memory.token_buffer_memory.FileUploadConfigManager.convert", + return_value=None, ): - mock_db.session.scalar.return_value = mock_workflow - result = mem._build_prompt_message_with_files( message_files=[], text_content="wf text", @@ -471,6 +566,7 @@ class TestBuildPromptMessageWithFiles: assert isinstance(result, UserPromptMessage) assert result.content == "wf text" + assert database.session.get(Workflow, workflow.id) is workflow # ------------------------------------------------------------------ # Invalid mode @@ -498,417 +594,140 @@ class TestBuildPromptMessageWithFiles: class TestGetHistoryPromptMessages: - """Tests for get_history_prompt_messages.""" + """Tests for persisted history retrieval, file batching, and pruning.""" def _make_memory(self, mode: AppMode = AppMode.CHAT) -> TokenBufferMemory: conv = _make_conversation(mode) conv.app = MagicMock() return TokenBufferMemory(conversation=conv, model_instance=_make_model_instance()) - def test_returns_empty_when_no_messages(self): + def test_returns_empty_when_no_messages(self, database: Database) -> None: + assert self._make_memory().get_history_prompt_messages() == [] + + def test_skips_newest_message_without_answer(self, database: Database) -> None: mem = self._make_memory() - with patch("core.memory.token_buffer_memory.db") as mock_db: - mock_db.session.scalars.return_value.all.return_value = [] - result = mem.get_history_prompt_messages() - assert result == [] + message = _persist_message(database, mem.conversation.id, answer="", answer_tokens=0) - def test_skips_first_message_without_answer(self): - """The newest message (index 0 after extraction) without answer and tokens==0 is skipped.""" + assert mem.get_history_prompt_messages() == [] + assert database.session.get(Message, message.id) is message + + def test_message_with_answer_returns_user_and_assistant_prompts(self, database: Database) -> None: mem = self._make_memory() + _persist_message(database, mem.conversation.id, query="My query", answer="My answer", answer_tokens=10) - msg_no_answer = _make_message(answer="", answer_tokens=0) - msg_no_answer.parent_message_id = None # ensures extract_thread_messages returns it - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.extract_thread_messages", - return_value=[msg_no_answer], - ), - ): - mock_db.session.scalars.return_value.all.side_effect = [ - [msg_no_answer], # first call: messages query - [], # second call: user files query (never hit, but safe) - ] - result = mem.get_history_prompt_messages() - - assert result == [] - - def test_message_with_answer_not_skipped(self): - """A message with a non-empty answer is NOT popped.""" - mem = self._make_memory() - - msg = _make_message(answer="some answer", answer_tokens=10) - msg.parent_message_id = None - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.extract_thread_messages", - return_value=[msg], - ), - patch( - "core.memory.token_buffer_memory.FileUploadConfigManager.convert", - return_value=None, - ), - ): - # user files query → empty; assistant files query → empty - mock_db.session.scalars.return_value.all.return_value = [] - result = mem.get_history_prompt_messages() - - assert len(result) == 2 # one user + one assistant - - def test_message_limit_default_is_500(self): - """When message_limit is None the stmt is limited to 500.""" - mem = self._make_memory() - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch("core.memory.token_buffer_memory.select") as mock_select, - patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]), - ): - mock_stmt = MagicMock() - mock_select.return_value.where.return_value.order_by.return_value = mock_stmt - mock_stmt.limit.return_value = mock_stmt - mock_db.session.scalars.return_value.all.return_value = [] - - mem.get_history_prompt_messages(message_limit=None) - mock_stmt.limit.assert_called_with(500) - - def test_message_limit_clipped_to_500(self): - """A message_limit > 500 is clamped to 500.""" - mem = self._make_memory() - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch("core.memory.token_buffer_memory.select") as mock_select, - patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]), - ): - mock_stmt = MagicMock() - mock_select.return_value.where.return_value.order_by.return_value = mock_stmt - mock_stmt.limit.return_value = mock_stmt - mock_db.session.scalars.return_value.all.return_value = [] - - mem.get_history_prompt_messages(message_limit=9999) - mock_stmt.limit.assert_called_with(500) - - def test_message_limit_positive_used(self): - """A positive message_limit < 500 is used as-is.""" - mem = self._make_memory() - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch("core.memory.token_buffer_memory.select") as mock_select, - patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]), - ): - mock_stmt = MagicMock() - mock_select.return_value.where.return_value.order_by.return_value = mock_stmt - mock_stmt.limit.return_value = mock_stmt - mock_db.session.scalars.return_value.all.return_value = [] - - mem.get_history_prompt_messages(message_limit=10) - mock_stmt.limit.assert_called_with(10) - - def test_message_limit_zero_uses_default(self): - """message_limit=0 triggers the else branch → default 500.""" - mem = self._make_memory() - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch("core.memory.token_buffer_memory.select") as mock_select, - patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=[]), - ): - mock_stmt = MagicMock() - mock_select.return_value.where.return_value.order_by.return_value = mock_stmt - mock_stmt.limit.return_value = mock_stmt - mock_db.session.scalars.return_value.all.return_value = [] - - mem.get_history_prompt_messages(message_limit=0) - mock_stmt.limit.assert_called_with(500) - - def test_user_files_cause_build_with_files_call(self): - """When user_files is non-empty _build_prompt_message_with_files is invoked.""" - mem = self._make_memory() - msg = _make_message() - msg.parent_message_id = None - - mock_user_file = MagicMock() - mock_user_file.message_id = msg.id # must match so batched grouping keys it to this message - mock_user_prompt = UserPromptMessage(content="from build") - mock_assistant_prompt = AssistantPromptMessage(content="answer") - - call_count = {"n": 0} - - def scalars_side_effect(stmt): - r = MagicMock() - if call_count["n"] == 0: - # messages query - r.all.return_value = [msg] - elif call_count["n"] == 1: - # user files - r.all.return_value = [mock_user_file] - else: - # assistant files - r.all.return_value = [] - call_count["n"] += 1 - return r - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.extract_thread_messages", - return_value=[msg], - ), - patch.object( - mem, - "_build_prompt_message_with_files", - side_effect=[mock_user_prompt, mock_assistant_prompt], - ) as mock_build, - patch( - "core.memory.token_buffer_memory.FileUploadConfigManager.convert", - return_value=None, - ), - ): - mock_db.session.scalars.side_effect = scalars_side_effect - result = mem.get_history_prompt_messages() - - assert mock_build.call_count >= 1 - # First call should be user message - first_call_kwargs = mock_build.call_args_list[0][1] - assert first_call_kwargs["is_user_message"] is True - - def test_assistant_files_cause_build_with_files_call(self): - """When assistant_files is non-empty, build is called with is_user_message=False.""" - mem = self._make_memory() - msg = _make_message() - msg.parent_message_id = None - - mock_assistant_file = MagicMock() - mock_assistant_file.message_id = msg.id # must match so batched grouping keys it to this message - mock_user_prompt = UserPromptMessage(content="query") - mock_assistant_prompt = AssistantPromptMessage(content="built") - - call_count = {"n": 0} - - def scalars_side_effect(stmt): - r = MagicMock() - if call_count["n"] == 0: - r.all.return_value = [msg] - elif call_count["n"] == 1: - r.all.return_value = [] # no user files - else: - r.all.return_value = [mock_assistant_file] - call_count["n"] += 1 - return r - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.extract_thread_messages", - return_value=[msg], - ), - patch.object( - mem, - "_build_prompt_message_with_files", - return_value=mock_assistant_prompt, - ) as mock_build, - ): - mock_db.session.scalars.side_effect = scalars_side_effect - result = mem.get_history_prompt_messages() - - mock_build.assert_called_once() - call_kwargs = mock_build.call_args[1] - assert call_kwargs["is_user_message"] is False - - def test_message_files_loaded_with_constant_query_count(self): - """Regression guard against N+1: message files must be batch-loaded. - - Regardless of the number of messages in the thread, file loading must use a - constant number of queries (1 messages query + 2 batched file queries), - never 2 queries per message. - """ - mem = self._make_memory() - - messages = [_make_message() for _ in range(5)] - for m in messages: - m.parent_message_id = None - - scalars_calls = {"n": 0} - - def scalars_side_effect(stmt): - r = MagicMock() - # First call returns the thread messages; the batched file queries return none. - r.all.return_value = messages if scalars_calls["n"] == 0 else [] - scalars_calls["n"] += 1 - return r - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=messages), - patch("core.memory.token_buffer_memory.FileUploadConfigManager.convert", return_value=None), - ): - mock_db.session.scalars.side_effect = scalars_side_effect - mem.get_history_prompt_messages() - - # 1 (messages) + 2 (batched user/assistant files) = 3, independent of message count. - # Before this fix it would have been 1 + 2 * 5 = 11 (an N+1 pattern). - assert scalars_calls["n"] == 3 - - def test_token_pruning_removes_oldest_messages(self): - """If tokens exceed limit, oldest messages are removed until within limit.""" - conv = _make_conversation() - conv.app = MagicMock() - - # Model returns tokens that decrease only after removing pairs - token_values = [3000, 1500] # first call over limit, second within - mi = MagicMock() - mi.get_llm_num_tokens.side_effect = token_values - - mem = TokenBufferMemory(conversation=conv, model_instance=mi) - - msg = _make_message() - msg.parent_message_id = None - - call_count = {"n": 0} - - def scalars_side_effect(stmt): - r = MagicMock() - if call_count["n"] == 0: - r.all.return_value = [msg] - else: - r.all.return_value = [] - call_count["n"] += 1 - return r - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.extract_thread_messages", - return_value=[msg], - ), - patch( - "core.memory.token_buffer_memory.FileUploadConfigManager.convert", - return_value=None, - ), - ): - mock_db.session.scalars.side_effect = scalars_side_effect - result = mem.get_history_prompt_messages(max_token_limit=2000) - - # After pruning, we should have fewer than the 2 initial messages - assert len(result) <= 1 - - def test_token_pruning_stops_at_single_message(self): - """Pruning stops when only 1 message remains (to prevent empty list).""" - conv = _make_conversation() - conv.app = MagicMock() - - # Always over limit - mi = MagicMock() - mi.get_llm_num_tokens.return_value = 99999 - - mem = TokenBufferMemory(conversation=conv, model_instance=mi) - - msg = _make_message() - msg.parent_message_id = None - - call_count = {"n": 0} - - def scalars_side_effect(stmt): - r = MagicMock() - if call_count["n"] == 0: - r.all.return_value = [msg] - else: - r.all.return_value = [] - call_count["n"] += 1 - return r - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.extract_thread_messages", - return_value=[msg], - ), - patch( - "core.memory.token_buffer_memory.FileUploadConfigManager.convert", - return_value=None, - ), - ): - mock_db.session.scalars.side_effect = scalars_side_effect - result = mem.get_history_prompt_messages(max_token_limit=1) - - # At least 1 message should remain - assert len(result) >= 1 - - def test_no_pruning_when_within_limit(self): - """When tokens ≤ limit, no pruning occurs.""" - mem = self._make_memory() - mem.model_instance.get_llm_num_tokens.return_value = 50 # well under default 2000 - - msg = _make_message() - msg.parent_message_id = None - - call_count = {"n": 0} - - def scalars_side_effect(stmt): - r = MagicMock() - if call_count["n"] == 0: - r.all.return_value = [msg] - else: - r.all.return_value = [] - call_count["n"] += 1 - return r - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.extract_thread_messages", - return_value=[msg], - ), - patch( - "core.memory.token_buffer_memory.FileUploadConfigManager.convert", - return_value=None, - ), - ): - mock_db.session.scalars.side_effect = scalars_side_effect - result = mem.get_history_prompt_messages(max_token_limit=2000) - - assert len(result) == 2 # user + assistant - - def test_plain_user_and_assistant_messages_returned(self): - """Without files, plain UserPromptMessage and AssistantPromptMessage appear.""" - mem = self._make_memory() - - msg = _make_message(answer="My answer") - msg.query = "My query" - msg.parent_message_id = None - - call_count = {"n": 0} - - def scalars_side_effect(stmt): - r = MagicMock() - if call_count["n"] == 0: - r.all.return_value = [msg] - else: - r.all.return_value = [] - call_count["n"] += 1 - return r - - with ( - patch("core.memory.token_buffer_memory.db") as mock_db, - patch( - "core.memory.token_buffer_memory.extract_thread_messages", - return_value=[msg], - ), - patch( - "core.memory.token_buffer_memory.FileUploadConfigManager.convert", - return_value=None, - ), - ): - mock_db.session.scalars.side_effect = scalars_side_effect - result = mem.get_history_prompt_messages() + result = mem.get_history_prompt_messages() assert len(result) == 2 - user_msg, ai_msg = result - assert isinstance(user_msg, UserPromptMessage) - assert user_msg.content == "My query" - assert isinstance(ai_msg, AssistantPromptMessage) - assert ai_msg.content == "My answer" + assert isinstance(result[0], UserPromptMessage) + assert result[0].content == "My query" + assert isinstance(result[1], AssistantPromptMessage) + assert result[1].content == "My answer" + + def test_history_is_conversation_scoped(self, database: Database) -> None: + mem = self._make_memory() + _persist_message(database, mem.conversation.id, answer="visible") + _persist_message(database, "other-conversation", answer="hidden") + + result = mem.get_history_prompt_messages() + + assert [prompt.content for prompt in result] == ["user query", "visible"] + + @pytest.mark.parametrize( + ("message_limit", "expected_limit"), + [(None, 500), (9999, 500), (10, 10), (0, 500)], + ) + def test_message_limit_is_applied_to_executable_query( + self, + database: Database, + message_limit: int | None, + expected_limit: int, + ) -> None: + mem = self._make_memory() + before = len(database.statements) + + mem.get_history_prompt_messages(message_limit=message_limit) + + statements = database.statements[before:] + assert len(statements) == 1 + sql, parameters = statements[0] + assert "LIMIT" in sql + assert expected_limit in parameters + + @pytest.mark.parametrize( + ("belongs_to", "is_user_message"), + [ + (MessageFileBelongsTo.USER, True), + (None, True), + (MessageFileBelongsTo.ASSISTANT, False), + ], + ) + def test_message_files_use_persisted_ownership( + self, + database: Database, + belongs_to: MessageFileBelongsTo | None, + is_user_message: bool, + ) -> None: + mem = self._make_memory() + message = _persist_message(database, mem.conversation.id) + message_file = _persist_message_file(database, message, belongs_to=belongs_to) + built_prompt = ( + UserPromptMessage(content="built user") + if is_user_message + else AssistantPromptMessage(content="built assistant") + ) + + with patch.object(mem, "_build_prompt_message_with_files", return_value=built_prompt) as build_prompt: + result = mem.get_history_prompt_messages() + + build_prompt.assert_called_once() + assert build_prompt.call_args.kwargs["message_files"] == [message_file] + assert build_prompt.call_args.kwargs["is_user_message"] is is_user_message + assert built_prompt in result + + def test_message_files_are_batch_loaded_with_constant_query_count(self, database: Database) -> None: + mem = self._make_memory() + base_time = datetime.now(UTC).replace(tzinfo=None) + messages = [ + _persist_message( + database, + mem.conversation.id, + query=f"query-{index}", + answer=f"answer-{index}", + created_at=base_time + timedelta(seconds=index), + ) + for index in range(5) + ] + before = len(database.statements) + + with patch("core.memory.token_buffer_memory.extract_thread_messages", return_value=messages): + result = mem.get_history_prompt_messages() + + selects = [sql for sql, _ in database.statements[before:] if sql.lstrip().upper().startswith("SELECT")] + assert len(selects) == 3 + assert len(result) == 10 + + @pytest.mark.parametrize( + ("token_values", "max_token_limit", "expected_length"), + [ + ([3000, 1500], 2000, 1), + ([99999, 99999], 1, 1), + ([50], 2000, 2), + ], + ) + def test_token_pruning_uses_persisted_history( + self, + database: Database, + token_values: list[int], + max_token_limit: int, + expected_length: int, + ) -> None: + mem = self._make_memory() + mem.model_instance.get_llm_num_tokens.side_effect = token_values + _persist_message(database, mem.conversation.id) + + result = mem.get_history_prompt_messages(max_token_limit=max_token_limit) + + assert len(result) == expected_length # =========================================================================== diff --git a/api/tests/unit_tests/core/moderation/api/test_api.py b/api/tests/unit_tests/core/moderation/api/test_api.py index 558b20e5f88..99eded2a5c4 100644 --- a/api/tests/unit_tests/core/moderation/api/test_api.py +++ b/api/tests/unit_tests/core/moderation/api/test_api.py @@ -2,10 +2,12 @@ from unittest.mock import MagicMock, patch import pytest from pydantic import ValidationError +from sqlalchemy.orm import Session from core.extension.api_based_extension_requestor import APIBasedExtensionPoint from core.moderation.api.api import ApiModeration, ModerationInputParams, ModerationOutputParams from core.moderation.base import ModerationAction, ModerationInputsResult, ModerationOutputsResult +from extensions.ext_database import db from models.api_based_extension import APIBasedExtension @@ -48,9 +50,11 @@ class TestApiModeration: @patch("core.moderation.api.api.ApiModeration._get_api_based_extension") def test_validate_config_success(self, mock_get_extension, api_config): - mock_get_extension.return_value = MagicMock(spec=APIBasedExtension) + mock_get_extension.return_value = APIBasedExtension( + tenant_id="tenant-id", name="Test Extension", api_endpoint="https://example.com", api_key="test-key" + ) ApiModeration.validate_config("test-tenant-id", api_config) - mock_get_extension.assert_called_once_with("test-tenant-id", "test-extension-id") + mock_get_extension.assert_called_once_with("test-tenant-id", "test-extension-id", db.session) def test_validate_config_missing_extension_id(self): config = { @@ -134,9 +138,12 @@ class TestApiModeration: @patch("core.moderation.api.api.decrypt_token") @patch("core.moderation.api.api.APIBasedExtensionRequestor") def test_get_config_by_requestor_success(self, mock_requestor_cls, mock_decrypt, mock_get_ext, api_moderation): - mock_ext = MagicMock(spec=APIBasedExtension) - mock_ext.api_endpoint = "http://api.test" - mock_ext.api_key = "encrypted-key" + mock_ext = APIBasedExtension( + tenant_id="tenant-id", + name="Test Extension", + api_endpoint="http://api.test", + api_key="encrypted-key", + ) mock_get_ext.return_value = mock_ext mock_decrypt.return_value = "decrypted-key" @@ -149,7 +156,7 @@ class TestApiModeration: result = api_moderation._get_config_by_requestor(APIBasedExtensionPoint.APP_MODERATION_INPUT, params) assert result == {"flagged": True} - mock_get_ext.assert_called_once_with("test-tenant-id", "test-extension-id") + mock_get_ext.assert_called_once_with("test-tenant-id", "test-extension-id", db.session) mock_decrypt.assert_called_once_with("test-tenant-id", "encrypted-key") mock_requestor_cls.assert_called_once_with("http://api.test", "decrypted-key") mock_requestor.request.assert_called_once_with(APIBasedExtensionPoint.APP_MODERATION_INPUT, params) @@ -165,17 +172,25 @@ class TestApiModeration: with pytest.raises(ValueError, match="API-based Extension not found"): api_moderation._get_config_by_requestor(APIBasedExtensionPoint.APP_MODERATION_INPUT, {}) - @patch("core.moderation.api.api.db.session.scalar") - def test_get_api_based_extension(self, mock_scalar): - mock_ext = MagicMock(spec=APIBasedExtension) - mock_scalar.return_value = mock_ext + def test_get_api_based_extension(self, sqlite_session: Session) -> None: + target = APIBasedExtension( + tenant_id="tenant-1", + name="Target extension", + api_endpoint="https://example.com/moderate", + api_key="encrypted-key", + ) + target.id = "ext-1" + other_tenant = APIBasedExtension( + tenant_id="tenant-2", + name="Other extension", + api_endpoint="https://example.com/other", + api_key="other-key", + ) + other_tenant.id = "ext-2" + sqlite_session.add_all((target, other_tenant)) + sqlite_session.commit() - result = ApiModeration._get_api_based_extension("tenant-1", "ext-1") + result = ApiModeration._get_api_based_extension("tenant-1", "ext-1", sqlite_session) - assert result == mock_ext - mock_scalar.assert_called_once() - # Verify the call has the correct filters - args, kwargs = mock_scalar.call_args - stmt = args[0] - # We can't easily inspect the statement without complex sqlalchemy tricks, - # but calling it is usually enough for unit tests if we mock the result. + assert result is target + assert ApiModeration._get_api_based_extension("tenant-1", "ext-2", sqlite_session) is None diff --git a/api/tests/unit_tests/core/ops/test_lookup_helpers.py b/api/tests/unit_tests/core/ops/test_lookup_helpers.py index 86aa68643da..90c2127b6ad 100644 --- a/api/tests/unit_tests/core/ops/test_lookup_helpers.py +++ b/api/tests/unit_tests/core/ops/test_lookup_helpers.py @@ -7,36 +7,212 @@ Covers: - TraceTask._get_user_id_from_metadata """ -from unittest.mock import MagicMock, patch +import uuid +from collections.abc import Iterator +from contextlib import contextmanager +from unittest.mock import PropertyMock, patch import pytest +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session + +from core.tools.entities.tool_entities import ApiProviderSchemaType +from extensions.ext_database import db +from graphon.model_runtime.entities.model_entities import ModelType +from models.account import Tenant +from models.base import TypeBase +from models.model import App, AppMode, IconType +from models.provider import Provider, ProviderCredential, ProviderModel, ProviderModelCredential, ProviderType +from models.tools import ApiToolProvider, BuiltinToolProvider, MCPToolProvider, WorkflowToolProvider # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- -def _make_db_and_session_patches(scalar_side_effect=None, scalar_return_value=None): - """Return (mock_db, cm, session) ready to patch 'core.ops.ops_trace_manager.db' - and 'core.ops.ops_trace_manager.Session'. +@pytest.fixture +def orm_session(sqlite_engine: Engine) -> Iterator[Session]: + models = ( + App, + Tenant, + Provider, + ProviderCredential, + ProviderModel, + ProviderModelCredential, + BuiltinToolProvider, + ApiToolProvider, + WorkflowToolProvider, + MCPToolProvider, + ) + tables = [model.metadata.tables[model.__tablename__] for model in models] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) - Provide either scalar_side_effect (list, for multiple calls) or - scalar_return_value (single value). - """ - mock_db = MagicMock() - mock_db.engine = MagicMock() + with patch.object(type(db), "engine", new_callable=PropertyMock, return_value=sqlite_engine): + with Session(sqlite_engine, expire_on_commit=False) as session: + yield session - session = MagicMock() - if scalar_side_effect is not None: - session.scalar.side_effect = scalar_side_effect + +def _persist_app(session: Session, *, tenant_id: str, name: str = "MyApp") -> App: + app = App( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name=name, + mode=AppMode.WORKFLOW, + icon_type=IconType.EMOJI, + icon="workflow", + icon_background="#FFFFFF", + enable_site=True, + enable_api=False, + ) + session.add(app) + session.commit() + return app + + +def _persist_tenant(session: Session, *, name: str = "MyWorkspace") -> Tenant: + tenant = Tenant(name=name) + session.add(tenant) + session.commit() + return tenant + + +def _persist_tool_provider( + session: Session, provider_type: str +) -> BuiltinToolProvider | ApiToolProvider | WorkflowToolProvider | MCPToolProvider: + tenant_id = str(uuid.uuid4()) + user_id = str(uuid.uuid4()) + if provider_type in {"builtin", "plugin"}: + provider = BuiltinToolProvider( + name="CredentialA", + tenant_id=tenant_id, + user_id=user_id, + provider="test/provider", + ) + elif provider_type == "api": + provider = ApiToolProvider( + name="CredentialA", + icon="icon.svg", + schema="{}", + schema_type_str=ApiProviderSchemaType.OPENAPI, + user_id=user_id, + tenant_id=tenant_id, + description="API provider", + tools_str="[]", + credentials_str="{}", + ) + elif provider_type == "workflow": + provider = WorkflowToolProvider( + name="CredentialA", + label="CredentialA", + icon="icon.svg", + app_id=str(uuid.uuid4()), + version="1", + user_id=user_id, + tenant_id=tenant_id, + description="Workflow provider", + ) + elif provider_type == "mcp": + provider = MCPToolProvider( + name="CredentialA", + server_identifier="credential-a", + server_url="https://example.com/mcp", + server_url_hash="credential-a-hash", + icon="icon.svg", + tenant_id=tenant_id, + user_id=user_id, + ) else: - session.scalar.return_value = scalar_return_value + raise ValueError(f"unsupported provider type: {provider_type}") - cm = MagicMock() - cm.__enter__ = MagicMock(return_value=session) - cm.__exit__ = MagicMock(return_value=False) + session.add(provider) + session.commit() + return provider - return mock_db, cm, session + +def _persist_provider_credential( + session: Session, + *, + tenant_id: str, + credential_name: str = "ProvCredName", +) -> ProviderCredential: + credential = ProviderCredential( + tenant_id=tenant_id, + provider_name="openai", + credential_name=credential_name, + encrypted_config="{}", + ) + session.add(credential) + session.commit() + return credential + + +def _persist_model_credential( + session: Session, + *, + tenant_id: str, + credential_name: str = "ModelCredName", +) -> ProviderModelCredential: + credential = ProviderModelCredential( + tenant_id=tenant_id, + provider_name="openai", + model_name="gpt-4", + model_type=ModelType.LLM, + credential_name=credential_name, + encrypted_config="{}", + ) + session.add(credential) + session.commit() + return credential + + +def _persist_provider( + session: Session, + *, + tenant_id: str, + credential_id: str | None, +) -> Provider: + provider = Provider( + tenant_id=tenant_id, + provider_name="openai", + provider_type=ProviderType.CUSTOM, + credential_id=credential_id, + ) + session.add(provider) + session.commit() + return provider + + +def _persist_provider_model( + session: Session, + *, + tenant_id: str, + credential_id: str | None, +) -> ProviderModel: + model = ProviderModel( + tenant_id=tenant_id, + provider_name="openai", + model_name="gpt-4", + model_type=ModelType.LLM, + credential_id=credential_id, + ) + session.add(model) + session.commit() + return model + + +@contextmanager +def _raise_on_table(engine: Engine, table_name: str) -> Iterator[None]: + """Raise only when SQL targets the named table, leaving other real lookups intact.""" + + def fail_target_query(_conn, _cursor, statement, _parameters, _context, _executemany): + if f"FROM {table_name}" in statement: + raise RuntimeError(f"forced failure for {table_name}") + + event.listen(engine, "before_cursor_execute", fail_target_query) + try: + yield + finally: + event.remove(engine, "before_cursor_execute", fail_target_query) # --------------------------------------------------------------------------- @@ -47,62 +223,42 @@ def _make_db_and_session_patches(scalar_side_effect=None, scalar_return_value=No class TestLookupAppAndWorkspaceNames: """Tests for _lookup_app_and_workspace_names(app_id, tenant_id).""" - def test_both_found(self): + def test_both_found(self, orm_session: Session): """Returns (app_name, workspace_name) when both records exist.""" from core.ops.ops_trace_manager import _lookup_app_and_workspace_names - mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=["MyApp", "MyWorkspace"]) - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - app_name, workspace_name = _lookup_app_and_workspace_names("app-123", "tenant-456") + tenant = _persist_tenant(orm_session) + app = _persist_app(orm_session, tenant_id=tenant.id) + app_name, workspace_name = _lookup_app_and_workspace_names(app.id, tenant.id) assert app_name == "MyApp" assert workspace_name == "MyWorkspace" - def test_app_only_found(self): + def test_app_only_found(self, orm_session: Session): """Returns (app_name, '') when tenant record is absent.""" from core.ops.ops_trace_manager import _lookup_app_and_workspace_names - mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=["MyApp", None]) - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - app_name, workspace_name = _lookup_app_and_workspace_names("app-123", "tenant-456") + app = _persist_app(orm_session, tenant_id=str(uuid.uuid4())) + app_name, workspace_name = _lookup_app_and_workspace_names(app.id, str(uuid.uuid4())) assert app_name == "MyApp" assert workspace_name == "" - def test_tenant_only_found(self): + def test_tenant_only_found(self, orm_session: Session): """Returns ('', workspace_name) when app record is absent.""" from core.ops.ops_trace_manager import _lookup_app_and_workspace_names - mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=[None, "MyWorkspace"]) - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - app_name, workspace_name = _lookup_app_and_workspace_names("app-123", "tenant-456") + tenant = _persist_tenant(orm_session) + app_name, workspace_name = _lookup_app_and_workspace_names(str(uuid.uuid4()), tenant.id) assert app_name == "" assert workspace_name == "MyWorkspace" - def test_neither_found(self): + def test_neither_found(self, orm_session: Session): """Returns ('', '') when both DB lookups return None.""" from core.ops.ops_trace_manager import _lookup_app_and_workspace_names - mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=[None, None]) - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - app_name, workspace_name = _lookup_app_and_workspace_names("app-123", "tenant-456") + app_name, workspace_name = _lookup_app_and_workspace_names(str(uuid.uuid4()), str(uuid.uuid4())) assert app_name == "" assert workspace_name == "" @@ -111,50 +267,30 @@ class TestLookupAppAndWorkspaceNames: """Returns ('', '') immediately when both IDs are None — no DB access.""" from core.ops.ops_trace_manager import _lookup_app_and_workspace_names - mock_db = MagicMock() - mock_session_cls = MagicMock() + app_name, workspace_name = _lookup_app_and_workspace_names(None, None) - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", mock_session_cls), - ): - app_name, workspace_name = _lookup_app_and_workspace_names(None, None) - - mock_session_cls.assert_not_called() assert app_name == "" assert workspace_name == "" - def test_app_id_none_only_queries_tenant(self): + def test_app_id_none_only_queries_tenant(self, orm_session: Session): """When app_id is None, only the tenant query is issued.""" from core.ops.ops_trace_manager import _lookup_app_and_workspace_names - mock_db, cm, session = _make_db_and_session_patches(scalar_return_value="OnlyWorkspace") - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - app_name, workspace_name = _lookup_app_and_workspace_names(None, "tenant-456") + tenant = _persist_tenant(orm_session, name="OnlyWorkspace") + app_name, workspace_name = _lookup_app_and_workspace_names(None, tenant.id) assert app_name == "" assert workspace_name == "OnlyWorkspace" - assert session.scalar.call_count == 1 - def test_tenant_id_none_only_queries_app(self): + def test_tenant_id_none_only_queries_app(self, orm_session: Session): """When tenant_id is None, only the app query is issued.""" from core.ops.ops_trace_manager import _lookup_app_and_workspace_names - mock_db, cm, session = _make_db_and_session_patches(scalar_return_value="OnlyApp") - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - app_name, workspace_name = _lookup_app_and_workspace_names("app-123", None) + app = _persist_app(orm_session, tenant_id=str(uuid.uuid4()), name="OnlyApp") + app_name, workspace_name = _lookup_app_and_workspace_names(app.id, None) assert app_name == "OnlyApp" assert workspace_name == "" - assert session.scalar.call_count == 1 # --------------------------------------------------------------------------- @@ -166,32 +302,20 @@ class TestLookupCredentialName: """Tests for _lookup_credential_name(credential_id, provider_type).""" @pytest.mark.parametrize("provider_type", ["builtin", "plugin", "api", "workflow", "mcp"]) - def test_known_provider_types_return_name(self, provider_type): + def test_known_provider_types_return_name(self, provider_type: str, orm_session: Session): """Each valid provider_type results in a DB query and returns the credential name.""" from core.ops.ops_trace_manager import _lookup_credential_name - mock_db, cm, session = _make_db_and_session_patches(scalar_return_value="CredentialA") - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - result = _lookup_credential_name("cred-123", provider_type) + provider = _persist_tool_provider(orm_session, provider_type) + result = _lookup_credential_name(provider.id, provider_type) assert result == "CredentialA" - session.scalar.assert_called_once() - def test_credential_not_found_returns_empty_string(self): + def test_credential_not_found_returns_empty_string(self, orm_session: Session): """Returns '' when DB yields None for the given credential_id.""" from core.ops.ops_trace_manager import _lookup_credential_name - mock_db, cm, _session = _make_db_and_session_patches(scalar_return_value=None) - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - result = _lookup_credential_name("cred-999", "api") + result = _lookup_credential_name(str(uuid.uuid4()), "api") assert result == "" @@ -199,48 +323,24 @@ class TestLookupCredentialName: """Returns '' immediately for an unrecognised provider_type — no DB access.""" from core.ops.ops_trace_manager import _lookup_credential_name - mock_db = MagicMock() - mock_session_cls = MagicMock() + result = _lookup_credential_name(str(uuid.uuid4()), "unknown_type") - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", mock_session_cls), - ): - result = _lookup_credential_name("cred-123", "unknown_type") - - mock_session_cls.assert_not_called() assert result == "" def test_none_credential_id_returns_empty_string_without_db(self): """Returns '' immediately when credential_id is None — no DB access.""" from core.ops.ops_trace_manager import _lookup_credential_name - mock_db = MagicMock() - mock_session_cls = MagicMock() + result = _lookup_credential_name(None, "api") - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", mock_session_cls), - ): - result = _lookup_credential_name(None, "api") - - mock_session_cls.assert_not_called() assert result == "" def test_none_provider_type_returns_empty_string_without_db(self): """Returns '' immediately when provider_type is None — no DB access.""" from core.ops.ops_trace_manager import _lookup_credential_name - mock_db = MagicMock() - mock_session_cls = MagicMock() + result = _lookup_credential_name(str(uuid.uuid4()), None) - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", mock_session_cls), - ): - result = _lookup_credential_name("cred-123", None) - - mock_session_cls.assert_not_called() assert result == "" def test_builtin_and_plugin_map_to_same_model(self): @@ -281,106 +381,78 @@ class TestLookupCredentialName: class TestLookupLlmCredentialInfo: """Tests for _lookup_llm_credential_info(tenant_id, provider, model, model_type).""" - def _provider_record(self, credential_id: str | None = None) -> MagicMock: - record = MagicMock() - record.credential_id = credential_id - return record - - def _model_record(self, credential_id: str | None = None) -> MagicMock: - record = MagicMock() - record.credential_id = credential_id - return record - - def test_model_level_credential_found(self): + def test_model_level_credential_found(self, orm_session: Session): """Returns model-level credential_id and name when ProviderModel has a credential.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - provider_record = self._provider_record(credential_id=None) - model_record = self._model_record(credential_id="model-cred-id") + tenant_id = str(uuid.uuid4()) + model_credential = _persist_model_credential(orm_session, tenant_id=tenant_id) + _persist_provider(orm_session, tenant_id=tenant_id, credential_id=None) + _persist_provider_model(orm_session, tenant_id=tenant_id, credential_id=model_credential.id) - # scalar calls: (1) Provider, (2) ProviderModel, (3) ProviderModelCredential.credential_name - mock_db, cm, _session = _make_db_and_session_patches( - scalar_side_effect=[provider_record, model_record, "ModelCredName"] + decoy_tenant_id = str(uuid.uuid4()) + decoy_credential = _persist_model_credential( + orm_session, + tenant_id=decoy_tenant_id, + credential_name="WrongTenantCredential", ) + _persist_provider(orm_session, tenant_id=decoy_tenant_id, credential_id=None) + _persist_provider_model(orm_session, tenant_id=decoy_tenant_id, credential_id=decoy_credential.id) - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4") + cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4") - assert cred_id == "model-cred-id" + assert cred_id == model_credential.id assert cred_name == "ModelCredName" - def test_provider_level_fallback_when_no_model_credential(self): + def test_provider_level_fallback_when_no_model_credential(self, orm_session: Session): """Falls back to provider-level credential when ProviderModel has no credential_id.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - provider_record = self._provider_record(credential_id="prov-cred-id") - model_record = self._model_record(credential_id=None) + tenant_id = str(uuid.uuid4()) + provider_credential = _persist_provider_credential(orm_session, tenant_id=tenant_id) + _persist_provider(orm_session, tenant_id=tenant_id, credential_id=provider_credential.id) + _persist_provider_model(orm_session, tenant_id=tenant_id, credential_id=None) - # scalar calls: (1) Provider, (2) ProviderModel (no cred), (3) ProviderCredential.credential_name - mock_db, cm, _session = _make_db_and_session_patches( - scalar_side_effect=[provider_record, model_record, "ProvCredName"] - ) + cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4") - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4") - - assert cred_id == "prov-cred-id" + assert cred_id == provider_credential.id assert cred_name == "ProvCredName" - def test_provider_level_fallback_when_no_model_record(self): + def test_provider_level_fallback_when_no_model_record(self, orm_session: Session): """Falls back to provider-level credential when no ProviderModel row exists.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - provider_record = self._provider_record(credential_id="prov-cred-id") + tenant_id = str(uuid.uuid4()) + provider_credential = _persist_provider_credential(orm_session, tenant_id=tenant_id) + _persist_provider(orm_session, tenant_id=tenant_id, credential_id=provider_credential.id) - # scalar calls: (1) Provider, (2) ProviderModel → None, (3) ProviderCredential.credential_name - mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=[provider_record, None, "ProvCredName"]) + cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4") - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4") - - assert cred_id == "prov-cred-id" + assert cred_id == provider_credential.id assert cred_name == "ProvCredName" - def test_no_model_arg_uses_provider_level_only(self): + def test_no_model_arg_uses_provider_level_only(self, orm_session: Session): """When model is None, skips ProviderModel query and uses provider credential.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - provider_record = self._provider_record(credential_id="prov-cred-id") + tenant_id = str(uuid.uuid4()) + provider_credential = _persist_provider_credential(orm_session, tenant_id=tenant_id) + _persist_provider(orm_session, tenant_id=tenant_id, credential_id=provider_credential.id) - # scalar calls: (1) Provider, (2) ProviderCredential.credential_name — no ProviderModel - mock_db, cm, session = _make_db_and_session_patches(scalar_side_effect=[provider_record, "ProvCredName"]) + cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", None) - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", None) - - assert cred_id == "prov-cred-id" + assert cred_id == provider_credential.id assert cred_name == "ProvCredName" - assert session.scalar.call_count == 2 - def test_provider_not_found_returns_none_and_empty(self): + def test_provider_not_found_returns_none_and_empty(self, orm_session: Session): """Returns (None, '') when Provider record does not exist.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - mock_db, cm, _session = _make_db_and_session_patches(scalar_return_value=None) + other_tenant_id = str(uuid.uuid4()) + _persist_provider(orm_session, tenant_id=other_tenant_id, credential_id=None) + tenant_id = str(uuid.uuid4()) - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4") + cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4") assert cred_id is None assert cred_name == "" @@ -389,16 +461,8 @@ class TestLookupLlmCredentialInfo: """Returns (None, '') immediately when tenant_id is None — no DB access.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - mock_db = MagicMock() - mock_session_cls = MagicMock() + cred_id, cred_name = _lookup_llm_credential_info(None, "openai", "gpt-4") - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", mock_session_cls), - ): - cred_id, cred_name = _lookup_llm_credential_info(None, "openai", "gpt-4") - - mock_session_cls.assert_not_called() assert cred_id is None assert cred_name == "" @@ -406,69 +470,46 @@ class TestLookupLlmCredentialInfo: """Returns (None, '') immediately when provider is None — no DB access.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - mock_db = MagicMock() - mock_session_cls = MagicMock() + cred_id, cred_name = _lookup_llm_credential_info(str(uuid.uuid4()), None, "gpt-4") - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", mock_session_cls), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", None, "gpt-4") - - mock_session_cls.assert_not_called() assert cred_id is None assert cred_name == "" - def test_db_error_on_outer_query_returns_none_and_empty(self): + def test_db_error_on_outer_query_returns_none_and_empty(self, orm_session: Session, sqlite_engine: Engine): """Returns (None, '') and logs a warning when the outer DB query raises.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - mock_db, cm, session = _make_db_and_session_patches() - session.scalar.side_effect = Exception("DB connection failed") - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4") + with _raise_on_table(sqlite_engine, "providers"): + cred_id, cred_name = _lookup_llm_credential_info(str(uuid.uuid4()), "openai", "gpt-4") assert cred_id is None assert cred_name == "" - def test_credential_name_lookup_failure_returns_id_with_empty_name(self): + def test_credential_name_lookup_failure_returns_id_with_empty_name( + self, orm_session: Session, sqlite_engine: Engine + ): """When credential name sub-query fails, returns cred_id but '' for name.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - provider_record = self._provider_record(credential_id="prov-cred-id") + tenant_id = str(uuid.uuid4()) + provider_credential = _persist_provider_credential(orm_session, tenant_id=tenant_id) + _persist_provider(orm_session, tenant_id=tenant_id, credential_id=provider_credential.id) - # Provider found, no model record, then name lookup raises - mock_db, cm, _session = _make_db_and_session_patches( - scalar_side_effect=[provider_record, None, Exception("deleted")] - ) + with _raise_on_table(sqlite_engine, "provider_credentials"): + cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4") - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4") - - assert cred_id == "prov-cred-id" + assert cred_id == provider_credential.id assert cred_name == "" - def test_no_credential_on_provider_or_model_returns_none_id(self): + def test_no_credential_on_provider_or_model_returns_none_id(self, orm_session: Session): """Returns (None, '') when neither provider nor model has a credential_id.""" from core.ops.ops_trace_manager import _lookup_llm_credential_info - provider_record = self._provider_record(credential_id=None) - model_record = self._model_record(credential_id=None) + tenant_id = str(uuid.uuid4()) + _persist_provider(orm_session, tenant_id=tenant_id, credential_id=None) + _persist_provider_model(orm_session, tenant_id=tenant_id, credential_id=None) - mock_db, cm, _session = _make_db_and_session_patches(scalar_side_effect=[provider_record, model_record]) - - with ( - patch("core.ops.ops_trace_manager.db", mock_db), - patch("core.ops.ops_trace_manager.Session", return_value=cm), - ): - cred_id, cred_name = _lookup_llm_credential_info("tenant-1", "openai", "gpt-4") + cred_id, cred_name = _lookup_llm_credential_info(tenant_id, "openai", "gpt-4") assert cred_id is None assert cred_name == "" diff --git a/api/tests/unit_tests/core/ops/test_ops_trace_manager.py b/api/tests/unit_tests/core/ops/test_ops_trace_manager.py index 704f5d362c6..3c3011597f0 100644 --- a/api/tests/unit_tests/core/ops/test_ops_trace_manager.py +++ b/api/tests/unit_tests/core/ops/test_ops_trace_manager.py @@ -1,18 +1,30 @@ -import contextlib +"""SQLite-backed tests for :mod:`core.ops.ops_trace_manager`.""" + +from __future__ import annotations + import json import queue +from collections.abc import Iterator from datetime import datetime, timedelta +from decimal import Decimal from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import PropertyMock, patch +from uuid import UUID import pytest +from flask import Flask +from sqlalchemy import Engine +from sqlalchemy.orm import Session, sessionmaker -from core.ops.ops_trace_manager import ( - OpsTraceManager, - TraceQueueManager, - TraceTask, - TraceTaskName, -) +import core.ops.ops_trace_manager as module +from core.ops.ops_trace_manager import OpsTraceManager, TraceQueueManager, TraceTask, TraceTaskName +from core.rag.models.document import Document as RetrievalDocument +from graphon.enums import WorkflowExecutionStatus +from graphon.file import FileTransferMethod, FileType +from models.enums import ConversationFromSource, CreatorUserRole, MessageStatus, WorkflowRunTriggeredFrom +from models.model import App, AppMode, AppModelConfig, Conversation, Message, MessageFile, TraceAppConfig +from models.workflow import WorkflowAppLog, WorkflowAppLogCreatedFrom, WorkflowRun, WorkflowType +from repositories.sqlalchemy_api_workflow_run_repository import DifyAPISQLAlchemyWorkflowRunRepository class DummyConfig: @@ -24,11 +36,8 @@ class DummyConfig: class DummyTraceInstance: - instances: list["DummyTraceInstance"] = [] - def __init__(self, config): self.config = config - DummyTraceInstance.instances.append(self) def api_check(self): return True @@ -40,14 +49,6 @@ class DummyTraceInstance: return "https://project.fake" -FAKE_PROVIDER_ENTRY = { - "config_class": DummyConfig, - "secret_keys": ["secret_value"], - "other_keys": ["other_value"], - "trace_instance": DummyTraceInstance, -} - - class FakeProviderMap: def __init__(self, data): self._data = data @@ -55,7 +56,15 @@ class FakeProviderMap: def __getitem__(self, key): if key in self._data: return self._data[key] - raise KeyError(f"Unsupported tracing provider: {key}") + raise KeyError(key) + + +PROVIDER_ENTRY = { + "config_class": DummyConfig, + "secret_keys": ["secret_value"], + "other_keys": ["other_value"], + "trace_instance": DummyTraceInstance, +} class DummyTimer: @@ -73,21 +82,172 @@ class DummyTimer: return False -class FakeMessageFile: - def __init__(self): - self.url = "path/to/file" - self.id = "file-id" - self.type = "document" - self.created_by_role = "role" - self.created_by = "user" +class EncryptTokenRecorder: + def __init__(self) -> None: + self.calls: list[tuple[str, str]] = [] + + def __call__(self, tenant_id: str, value: str) -> str: + self.calls.append((tenant_id, value)) + return f"enc-{value}" -def make_message_data(**overrides): +class BatchDecryptTokenRecorder: + def __init__(self) -> None: + self.calls: list[tuple[str, list[str]]] = [] + + def __call__(self, tenant_id: str, values: list[str]) -> list[str]: + self.calls.append((tenant_id, values)) + return [f"dec-{value}" for value in values] + + +class ObfuscatedTokenRecorder: + def __init__(self) -> None: + self.calls: list[str] = [] + + def __call__(self, value: str) -> str: + self.calls.append(value) + return f"ob-{value}" + + +class RecordingStorage: + def __init__(self) -> None: + self.writes: list[tuple[str, bytes]] = [] + + def save(self, path: str, data: bytes) -> None: + self.writes.append((path, data)) + + +class RecordingDispatcher: + def __init__(self) -> None: + self.payloads: list[dict[str, str]] = [] + + def delay(self, payload: dict[str, str]) -> None: + self.payloads.append(payload) + + +@pytest.fixture +def database(sqlite_engine: Engine, sqlite_session: Session) -> Iterator[Session]: + with ( + patch.object(module.db, "session", sqlite_session), + patch.object(type(module.db), "engine", new_callable=PropertyMock, return_value=sqlite_engine), + ): + yield sqlite_session + + +@pytest.fixture +def trace_environment(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + monkeypatch.setattr(module, "provider_config_map", FakeProviderMap({"dummy": PROVIDER_ENTRY})) + OpsTraceManager.ops_trace_instances_cache.clear() + OpsTraceManager.decrypted_configs_cache.clear() + monkeypatch.setattr(module.threading, "Timer", DummyTimer) + monkeypatch.setattr(module, "trace_manager_queue", queue.Queue()) + monkeypatch.setattr(module, "trace_manager_timer", None) + monkeypatch.setattr("core.telemetry.gateway.is_enterprise_telemetry_enabled", lambda: False) + + app = Flask(__name__) + with app.app_context(): + yield + + +@pytest.fixture +def encryption_functions( + monkeypatch: pytest.MonkeyPatch, +) -> tuple[EncryptTokenRecorder, BatchDecryptTokenRecorder, ObfuscatedTokenRecorder]: + encrypt = EncryptTokenRecorder() + decrypt = BatchDecryptTokenRecorder() + obfuscate = ObfuscatedTokenRecorder() + monkeypatch.setattr(module, "encrypt_token", encrypt) + monkeypatch.setattr(module, "batch_decrypt_token", decrypt) + monkeypatch.setattr(module, "obfuscated_token", obfuscate) + return encrypt, decrypt, obfuscate + + +def _app(session: Session, *, app_id: str = "app-id", tracing: str | None = None) -> App: + app = App( + id=app_id, + tenant_id="tenant-1", + name="App", + description="description", + mode=AppMode.CHAT, + icon_type=None, + icon=None, + icon_background=None, + enable_site=True, + enable_api=True, + max_active_requests=None, + tracing=tracing, + ) + session.add(app) + session.commit() + return app + + +def _conversation_message( + session: Session, app: App, *, config: AppModelConfig | None = None +) -> tuple[Conversation, Message]: + conversation = Conversation( + id="conversation-1", + app_id=app.id, + app_model_config_id=config.id if config else None, + model_provider="provider", + override_model_configs=None, + model_id="model", + mode=AppMode.CHAT, + name="Conversation", + summary="", + _inputs={}, + introduction="", + system_instruction="", + invoke_from=None, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id="end-user-1", + from_account_id=None, + read_at=None, + read_account_id=None, + ) + message = Message( + id="message-1", + app_id=app.id, + model_provider="provider", + model_id="model", + override_model_configs=None, + conversation_id=conversation.id, + _inputs={}, + query="query", + message={"text": "hello"}, + message_tokens=5, + message_unit_price=Decimal(0), + message_price_unit=Decimal("0.001"), + answer="world", + answer_tokens=7, + answer_unit_price=Decimal(0), + answer_price_unit=Decimal("0.001"), + parent_message_id=None, + provider_response_latency=1, + total_price=Decimal(0), + currency="USD", + status=MessageStatus.NORMAL, + error=None, + message_metadata=None, + invoke_from=None, + from_source=ConversationFromSource.CONSOLE, + from_end_user_id="end-user-1", + from_account_id=None, + agent_based=False, + workflow_run_id="run-1", + app_mode=AppMode.CHAT, + ) + session.add_all([conversation, message]) + session.commit() + return conversation, message + + +def _message_data(**overrides): created_at = datetime(2025, 2, 20, 12, 0, 0) - base = { - "id": "msg-id", + data = { + "id": "message-1", "app_id": "app-id", - "conversation_id": "conv-id", + "conversation_id": "conversation-1", "created_at": created_at, "updated_at": created_at + timedelta(seconds=3), "message": "hello", @@ -99,487 +259,352 @@ def make_message_data(**overrides): "status": "complete", "model_provider": "provider", "model_id": "model", - "from_end_user_id": "end-user", - "from_account_id": "account", + "from_end_user_id": "end-user-1", + "from_account_id": None, "agent_based": False, - "workflow_run_id": "workflow-run", - "from_source": "source", + "workflow_run_id": "run-1", + "from_source": "console", "message_metadata": json.dumps({"usage": {"time_to_first_token": 1, "time_to_generate": 2}}), "agent_thoughts": [], - "query": "sample-query", - "inputs": "sample-input", + "query": "query", + "inputs": "inputs", } - base.update(overrides) - - class MessageData: - def __init__(self, data): - self.__dict__.update(data) - - def to_dict(self): - return dict(self.__dict__) - - return MessageData(base) + data.update(overrides) + return SimpleNamespace(**data, to_dict=lambda: data) -def make_agent_thought(tool_name, created_at): - return SimpleNamespace( - tools=[tool_name], - created_at=created_at, - tool_meta={ - tool_name: { - "tool_config": {"foo": "bar"}, - "time_cost": 5, - "error": "", - "tool_parameters": {"x": 1}, - } - }, +def test_encrypt_decrypt_obfuscate_and_cache( + trace_environment: None, + encryption_functions: tuple[EncryptTokenRecorder, BatchDecryptTokenRecorder, ObfuscatedTokenRecorder], +) -> None: + encrypted = OpsTraceManager.encrypt_tracing_config( + "tenant-1", "dummy", {"secret_value": "value", "other_value": "info"} + ) + assert encrypted == {"secret_value": "enc-value", "other_value": "info"} + preserved = OpsTraceManager.encrypt_tracing_config( + "tenant-1", "dummy", {"secret_value": "*"}, current_trace_config={"secret_value": "keep"} + ) + assert preserved["secret_value"] == "keep" + first = OpsTraceManager.decrypt_tracing_config("tenant-1", "dummy", encrypted) + second = OpsTraceManager.decrypt_tracing_config("tenant-1", "dummy", encrypted) + assert first == second + assert len(encryption_functions[1].calls) == 1 + obfuscated = OpsTraceManager.obfuscated_decrypt_token("dummy", first) + assert obfuscated["secret_value"] == "ob-dec-enc-value" + assert encryption_functions[2].calls == ["dec-enc-value"] + + +def test_decrypted_config_reads_real_trace_and_app_rows( + trace_environment: None, + encryption_functions, + database: Session, +) -> None: + app = _app(database) + trace = TraceAppConfig( + app_id=app.id, + tracing_provider="dummy", + tracing_config={"secret_value": "encrypted", "other_value": "info"}, + ) + database.add(trace) + database.commit() + result = OpsTraceManager.get_decrypted_tracing_config(app.id, "dummy") + assert result == {"secret_value": "dec-encrypted", "other_value": "info"} + assert OpsTraceManager.get_decrypted_tracing_config(app.id, "missing") is None + + null_config_app = _app(database, app_id="app-null-config") + database.add(TraceAppConfig(app_id=null_config_app.id, tracing_provider="dummy", tracing_config=None)) + database.commit() + with pytest.raises(ValueError, match="Tracing config cannot be None"): + OpsTraceManager.get_decrypted_tracing_config(null_config_app.id, "dummy") + + database.delete(app) + database.commit() + with pytest.raises(ValueError, match="App not found"): + OpsTraceManager.get_decrypted_tracing_config("app-id", "dummy") + + +def test_ops_trace_instance_uses_persisted_enabled_state_and_cache( + trace_environment: None, + encryption_functions, + database: Session, +) -> None: + app = _app(database, tracing=json.dumps({"enabled": False, "tracing_provider": "dummy"})) + assert OpsTraceManager.get_ops_trace_instance(app.id) is None + app.tracing = json.dumps({"enabled": True, "tracing_provider": "dummy"}) + database.add(TraceAppConfig(app_id=app.id, tracing_provider="dummy", tracing_config={"secret_value": "encrypted"})) + database.commit() + instance = OpsTraceManager.get_ops_trace_instance(app.id) + assert isinstance(instance, DummyTraceInstance) + assert OpsTraceManager.get_ops_trace_instance(app.id) is instance + + app.tracing = json.dumps({"enabled": True, "tracing_provider": "missing"}) + database.commit() + assert OpsTraceManager.get_ops_trace_instance(app.id) is None + + assert OpsTraceManager.get_ops_trace_instance(None) is None + assert OpsTraceManager.get_ops_trace_instance("tenant-storage-id") is None + assert OpsTraceManager.get_ops_trace_instance("missing") is None + + +def test_message_config_lookup_uses_real_conversation_and_model_config(database: Session) -> None: + app = _app(database) + config = AppModelConfig(app_id=app.id, model='{"provider":"openai"}') + database.add(config) + database.commit() + conversation, message = _conversation_message(database, app, config=config) + result = OpsTraceManager.get_app_config_through_message_id(message.id) + assert result.id == config.id + + conversation.app_model_config_id = None + conversation.override_model_configs = json.dumps({"provider": "override"}) + database.commit() + override = OpsTraceManager.get_app_config_through_message_id(message.id) + assert json.loads(override) == {"provider": "override"} + + assert OpsTraceManager.get_app_config_through_message_id("missing") is None + + +def test_update_and_get_app_tracing_config_persist_state(trace_environment: None, database: Session) -> None: + app = _app(database) + assert OpsTraceManager.get_app_tracing_config(app.id, database) == { + "enabled": False, + "tracing_provider": None, + } + OpsTraceManager.update_app_tracing_config(app.id, True, "dummy") + database.expire_all() + assert OpsTraceManager.get_app_tracing_config(app.id, database) == { + "enabled": True, + "tracing_provider": "dummy", + } + with pytest.raises(ValueError, match="Invalid tracing provider"): + OpsTraceManager.update_app_tracing_config(app.id, True, "missing") + with pytest.raises(ValueError, match="App not found"): + OpsTraceManager.update_app_tracing_config("missing", False, None) + with pytest.raises(ValueError, match="App not found"): + OpsTraceManager.get_app_tracing_config("missing", database) + + +def test_message_trace_reads_real_conversation_app_and_message_file( + monkeypatch: pytest.MonkeyPatch, + trace_environment: None, + database: Session, +) -> None: + app = _app(database) + _, message = _conversation_message(database, app) + file = MessageFile( + message_id=message.id, + type=FileType.DOCUMENT, + transfer_method=FileTransferMethod.REMOTE_URL, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + url="path/to/file", + ) + database.add(file) + database.commit() + monkeypatch.setattr(module, "get_message_data", lambda _message_id: _message_data()) + result = TraceTask(trace_type=TraceTaskName.MESSAGE_TRACE, message_id=message.id).message_trace(message.id) + assert result.message_id == message.id + assert result.conversation_mode == AppMode.CHAT + assert result.file_list[0].endswith("path/to/file") + assert result.metadata["tenant_id"] == "tenant-1" + + +def test_workflow_log_enriches_moderation_and_suggested_question_traces( + monkeypatch: pytest.MonkeyPatch, + database: Session, +) -> None: + log = WorkflowAppLog( + tenant_id="tenant-1", + app_id="app-id", + workflow_id="workflow-1", + workflow_run_id="run-1", + created_from=WorkflowAppLogCreatedFrom.WEB_APP, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + ) + database.add(log) + database.commit() + monkeypatch.setattr(module, "get_message_data", lambda _message_id: _message_data()) + task = TraceTask(trace_type=TraceTaskName.MODERATION_TRACE, message_id="message-1") + moderation = SimpleNamespace(action="block", preset_response="no", query="q", flagged=True) + result = task.moderation_trace( + "message-1", + {"start": 1, "end": 2}, + moderation_result=moderation, + inputs={"source": "payload"}, + ) + assert result.message_id == log.id + assert result.flagged is True + assert result.inputs == {"source": "payload"} + suggested = task.suggested_question_trace("message-1", {"start": 1, "end": 2}, suggested_question=["q1"]) + assert suggested.message_id == log.id + assert suggested.suggested_question == ["q1"] + + +def test_dataset_retrieval_trace_serializes_documents( + monkeypatch: pytest.MonkeyPatch, + trace_environment: None, + database: Session, +) -> None: + _app(database) + monkeypatch.setattr(module, "get_message_data", lambda _message_id: _message_data()) + document = RetrievalDocument(page_content="value") + + result = TraceTask(trace_type=TraceTaskName.DATASET_RETRIEVAL_TRACE).dataset_retrieval_trace( + "message-1", + {"start": 1, "end": 2}, + documents=[document], ) + assert result.documents == [document.model_dump()] + assert result.documents[0]["page_content"] == "value" + assert result.metadata["tenant_id"] == "tenant-1" -def make_workflow_run(): - return SimpleNamespace( - workflow_id="wf-1", - tenant_id="tenant", - id="run-id", - elapsed_time=10, - status="finished", - inputs_dict={"sys.file": ["f1"], "query": "search"}, - outputs_dict={"out": "value"}, + +def test_workflow_trace_reads_real_workflow_log_from_owned_session( + monkeypatch: pytest.MonkeyPatch, + trace_environment: None, + database: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + app = _app(database) + log = WorkflowAppLog( + tenant_id=app.tenant_id, + app_id=app.id, + workflow_id="workflow-1", + workflow_run_id="run-1", + created_from=WorkflowAppLogCreatedFrom.WEB_APP, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + ) + workflow_run = WorkflowRun( + id="run-1", + tenant_id=app.tenant_id, + app_id=app.id, + workflow_id="workflow-1", + type=WorkflowType.WORKFLOW, + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, version="3", + graph="{}", + inputs=json.dumps({"query": "search"}), + status=WorkflowExecutionStatus.SUCCEEDED, + outputs=json.dumps({"out": "value"}), error=None, + elapsed_time=10, total_tokens=12, - workflow_run_id="run-id", + total_steps=1, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", created_at=datetime(2025, 2, 20, 10, 0, 0), finished_at=datetime(2025, 2, 20, 10, 0, 5), - triggered_from="user", - app_id="app-id", - to_dict=lambda self=None: {"run": "value"}, ) - - -def configure_db_scalar(session, *, message_file=None, workflow_app_log=None): - """Configure session.scalar to return appropriate values for MessageFile/WorkflowAppLog lookups.""" - original_scalar = session.scalar - - def _side_effect(stmt): - stmt_str = str(stmt) - if "message_file" in stmt_str.lower(): - return message_file - if "workflow_app_log" in stmt_str.lower(): - return workflow_app_log - return original_scalar(stmt) - - session.scalar.side_effect = _side_effect - - -class DummySessionContext: - scalar_values = [] - - def __init__(self, engine): - self._values = list(self.scalar_values) - self._index = 0 - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - return False - - def execute(self, *args, **kwargs): - return self - - def scalar(self, *args, **kwargs): - if self._index >= len(self._values): - return None - value = self._values[self._index] - self._index += 1 - return value - - def scalars(self, *args, **kwargs): - return self - - def all(self): - return [] - - -@pytest.fixture(autouse=True) -def patch_provider_map(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr( - "core.ops.ops_trace_manager.provider_config_map", FakeProviderMap({"dummy": FAKE_PROVIDER_ENTRY}) - ) - OpsTraceManager.ops_trace_instances_cache.clear() - OpsTraceManager.decrypted_configs_cache.clear() - - -@pytest.fixture(autouse=True) -def patch_timer_and_current_app(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr("core.ops.ops_trace_manager.threading.Timer", DummyTimer) - monkeypatch.setattr("core.ops.ops_trace_manager.trace_manager_queue", queue.Queue()) - monkeypatch.setattr("core.ops.ops_trace_manager.trace_manager_timer", None) - - class FakeApp: - def app_context(self): - return contextlib.nullcontext() - - fake_current = MagicMock() - fake_current._get_current_object.return_value = FakeApp() - monkeypatch.setattr("core.ops.ops_trace_manager.current_app", fake_current) - - -@pytest.fixture(autouse=True) -def patch_sqlalchemy_session(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr("core.ops.ops_trace_manager.Session", DummySessionContext) - - -@pytest.fixture -def encryption_mocks(monkeypatch: pytest.MonkeyPatch): - encrypt_mock = MagicMock(side_effect=lambda tenant, value: f"enc-{value}") - batch_decrypt_mock = MagicMock(side_effect=lambda tenant, values: [f"dec-{value}" for value in values]) - obfuscate_mock = MagicMock(side_effect=lambda value: f"ob-{value}") - monkeypatch.setattr("core.ops.ops_trace_manager.encrypt_token", encrypt_mock) - monkeypatch.setattr("core.ops.ops_trace_manager.batch_decrypt_token", batch_decrypt_mock) - monkeypatch.setattr("core.ops.ops_trace_manager.obfuscated_token", obfuscate_mock) - return encrypt_mock, batch_decrypt_mock, obfuscate_mock - - -@pytest.fixture -def mock_db(monkeypatch: pytest.MonkeyPatch): - session = MagicMock() - session.scalars.return_value.all.return_value = ["chat"] - db_mock = MagicMock() - db_mock.session = session - db_mock.engine = MagicMock() - monkeypatch.setattr("core.ops.ops_trace_manager.db", db_mock) - return session - - -@pytest.fixture -def workflow_repo_fixture(monkeypatch: pytest.MonkeyPatch): - repo = MagicMock() - repo.get_workflow_run_by_id_without_tenant.return_value = make_workflow_run() + database.add_all([log, workflow_run]) + database.commit() + repo = DifyAPISQLAlchemyWorkflowRunRepository(sqlite_session_factory) monkeypatch.setattr(TraceTask, "_get_workflow_run_repo", classmethod(lambda cls: repo)) - return repo - - -@pytest.fixture -def trace_task_message(monkeypatch: pytest.MonkeyPatch, mock_db): - message_data = make_message_data() - monkeypatch.setattr("core.ops.ops_trace_manager.get_message_data", lambda msg_id: message_data) - configure_db_scalar(mock_db, message_file=FakeMessageFile(), workflow_app_log=SimpleNamespace(id="log-id")) - return message_data - - -def test_encrypt_tracing_config_handles_star_and_encrypt(encryption_mocks): - encrypted = OpsTraceManager.encrypt_tracing_config( - "tenant", - "dummy", - {"secret_value": "value", "other_value": "info"}, - current_trace_config={"secret_value": "keep"}, + monkeypatch.setattr(TraceTask, "_calculate_workflow_token_split", classmethod(lambda cls, *_a, **_k: (5, 7))) + result = TraceTask(trace_type=TraceTaskName.WORKFLOW_TRACE).workflow_trace( + workflow_run_id="run-1", conversation_id=None, user_id="user-1" ) - assert encrypted["secret_value"] == "enc-value" - assert encrypted["other_value"] == "info" + assert result.workflow_run_id == "run-1" + assert result.workflow_id == "workflow-1" + assert result.workflow_app_log_id == log.id + assert result.prompt_tokens == 5 + assert result.completion_tokens == 7 -def test_encrypt_tracing_config_preserves_star(encryption_mocks): - encrypted = OpsTraceManager.encrypt_tracing_config( - "tenant", - "dummy", - {"secret_value": "*", "other_value": "info"}, - current_trace_config={"secret_value": "keep"}, +def test_tool_trace_reads_real_message_file(monkeypatch: pytest.MonkeyPatch, database: Session) -> None: + file = MessageFile( + message_id="message-1", + type=FileType.DOCUMENT, + transfer_method=FileTransferMethod.REMOTE_URL, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + url="tool/file", ) - assert encrypted["secret_value"] == "keep" - - -def test_decrypt_tracing_config_caches(encryption_mocks): - _, decrypt_mock, _ = encryption_mocks - payload = {"secret_value": "enc", "other_value": "info"} - first = OpsTraceManager.decrypt_tracing_config("tenant", "dummy", payload) - second = OpsTraceManager.decrypt_tracing_config("tenant", "dummy", payload) - assert first == second - assert decrypt_mock.call_count == 1 - - -def test_obfuscated_decrypt_token(encryption_mocks): - _, _, obfuscate_mock = encryption_mocks - result = OpsTraceManager.obfuscated_decrypt_token("dummy", {"secret_value": "value", "other_value": "info"}) - assert "secret_value" in result - assert result["secret_value"] == "ob-value" - obfuscate_mock.assert_called_once() - - -def test_get_decrypted_tracing_config_returns_config(encryption_mocks, mock_db): - trace_config_data = SimpleNamespace(tracing_config={"secret_value": "enc", "other_value": "info"}) - app = SimpleNamespace(id="app-id", tenant_id="tenant") - mock_db.scalar.side_effect = [trace_config_data, app] - - decrypted = OpsTraceManager.get_decrypted_tracing_config("app-id", "dummy") - assert decrypted["other_value"] == "info" - - -def test_get_decrypted_tracing_config_missing_trace_config(mock_db): - mock_db.scalar.return_value = None - assert OpsTraceManager.get_decrypted_tracing_config("app-id", "dummy") is None - - -def test_get_decrypted_tracing_config_raises_for_missing_app(mock_db): - trace_config_data = SimpleNamespace(tracing_config={"secret_value": "enc"}) - mock_db.scalar.side_effect = [trace_config_data, None] - with pytest.raises(ValueError, match="App not found"): - OpsTraceManager.get_decrypted_tracing_config("app-id", "dummy") - - -def test_get_decrypted_tracing_config_raises_for_none_config(mock_db): - trace_config_data = SimpleNamespace(tracing_config=None) - mock_db.scalar.side_effect = [trace_config_data, SimpleNamespace(tenant_id="tenant")] - with pytest.raises(ValueError, match="Tracing config cannot be None"): - OpsTraceManager.get_decrypted_tracing_config("app-id", "dummy") - - -def test_get_ops_trace_instance_handles_none_app(mock_db): - mock_db.get.return_value = None - assert OpsTraceManager.get_ops_trace_instance("app-id") is None - - -def test_get_ops_trace_instance_returns_none_when_disabled(mock_db, monkeypatch: pytest.MonkeyPatch): - app = SimpleNamespace(id="app-id", tracing=json.dumps({"enabled": False})) - mock_db.get.return_value = app - assert OpsTraceManager.get_ops_trace_instance("app-id") is None - - -def test_get_ops_trace_instance_invalid_provider(mock_db, monkeypatch: pytest.MonkeyPatch): - app = SimpleNamespace(id="app-id", tracing=json.dumps({"enabled": True, "tracing_provider": "missing"})) - mock_db.get.return_value = app - monkeypatch.setattr("core.ops.ops_trace_manager.provider_config_map", FakeProviderMap({})) - assert OpsTraceManager.get_ops_trace_instance("app-id") is None - - -def test_get_ops_trace_instance_success(monkeypatch: pytest.MonkeyPatch, mock_db): - app = SimpleNamespace(id="app-id", tracing=json.dumps({"enabled": True, "tracing_provider": "dummy"})) - mock_db.get.return_value = app - monkeypatch.setattr( - "core.ops.ops_trace_manager.OpsTraceManager.get_decrypted_tracing_config", - classmethod(lambda cls, aid, provider: {"secret_value": "decrypted", "other_value": "info"}), + database.add(file) + database.commit() + thought = SimpleNamespace( + tools=["tool-a"], + created_at=datetime(2025, 2, 20, 12, 1), + tool_meta={"tool-a": {"tool_config": {}, "time_cost": 5, "error": "", "tool_parameters": {}}}, ) - instance = OpsTraceManager.get_ops_trace_instance("app-id") - assert instance is not None - cached_instance = OpsTraceManager.get_ops_trace_instance("app-id") - assert instance is cached_instance - - -def test_get_app_config_through_message_id_returns_none(mock_db): - mock_db.scalar.return_value = None - assert OpsTraceManager.get_app_config_through_message_id("m") is None - - -def test_get_app_config_through_message_id_prefers_override(mock_db): - message = SimpleNamespace(conversation_id="conv") - conversation = SimpleNamespace(app_model_config_id=None, override_model_configs={"foo": "bar"}) - app_config = SimpleNamespace(id="config-id") - mock_db.scalar.side_effect = [message, conversation] - result = OpsTraceManager.get_app_config_through_message_id("m") - assert result == {"foo": "bar"} - - -def test_get_app_config_through_message_id_app_model_config(mock_db): - message = SimpleNamespace(conversation_id="conv") - conversation = SimpleNamespace(app_model_config_id="cfg", override_model_configs=None) - mock_db.scalar.side_effect = [message, conversation, SimpleNamespace(id="cfg")] - result = OpsTraceManager.get_app_config_through_message_id("m") - assert result.id == "cfg" - - -def test_update_app_tracing_config_invalid_provider(mock_db, monkeypatch: pytest.MonkeyPatch): - mock_db.get.return_value = None - with pytest.raises(ValueError, match="Invalid tracing provider"): - OpsTraceManager.update_app_tracing_config("app", True, "bad") - with pytest.raises(ValueError, match="App not found"): - OpsTraceManager.update_app_tracing_config("app", True, None) - - -def test_update_app_tracing_config_success(mock_db): - app = SimpleNamespace(id="app-id", tracing="{}") - mock_db.get.return_value = app - OpsTraceManager.update_app_tracing_config("app-id", True, "dummy") - assert app.tracing is not None - mock_db.commit.assert_called_once() - - -def test_get_app_tracing_config_errors_when_missing(mock_db): - mock_db.get.return_value = None - with pytest.raises(ValueError, match="App not found"): - OpsTraceManager.get_app_tracing_config("app", mock_db) - - -def test_get_app_tracing_config_returns_defaults(mock_db): - mock_db.get.return_value = SimpleNamespace(tracing=None) - assert OpsTraceManager.get_app_tracing_config("app-id", mock_db) == {"enabled": False, "tracing_provider": None} - - -def test_get_app_tracing_config_returns_payload(mock_db): - payload = {"enabled": True, "tracing_provider": "dummy"} - mock_db.get.return_value = SimpleNamespace(tracing=json.dumps(payload)) - assert OpsTraceManager.get_app_tracing_config("app-id", mock_db) == payload - - -def test_check_and_project_helpers(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr( - "core.ops.ops_trace_manager.provider_config_map", - FakeProviderMap( - { - "dummy": { - "config_class": DummyConfig, - "trace_instance": type( - "Trace", - (), - { - "__init__": lambda self, cfg: None, - "api_check": lambda self: True, - "get_project_key": lambda self: "key", - "get_project_url": lambda self: "url", - }, - ), - "secret_keys": [], - "other_keys": [], - } - } - ), + monkeypatch.setattr(module, "get_message_data", lambda _message_id: _message_data(agent_thoughts=[thought])) + result = TraceTask(trace_type=TraceTaskName.TOOL_TRACE).tool_trace( + "message-1", {"start": 1, "end": 2}, tool_name="tool-a", tool_inputs={}, tool_outputs="result" ) - assert OpsTraceManager.check_trace_config_is_effective({}, "dummy") - assert OpsTraceManager.get_trace_config_project_key({}, "dummy") == "key" - assert OpsTraceManager.get_trace_config_project_url({}, "dummy") == "url" - - -def test_trace_task_conversation_and_extract(monkeypatch: pytest.MonkeyPatch): - task = TraceTask(trace_type=TraceTaskName.CONVERSATION_TRACE, message_id="msg") - assert task.conversation_trace(foo="bar") == {"foo": "bar"} - assert task._extract_streaming_metrics(make_message_data(message_metadata="not json")) == {} - - -def test_trace_task_message_trace(trace_task_message, mock_db): - task = TraceTask(trace_type=TraceTaskName.MESSAGE_TRACE, message_id="msg-id") - result = task.message_trace("msg-id") - assert result.message_id == "msg-id" - - -def test_trace_task_workflow_trace(workflow_repo_fixture, mock_db): - DummySessionContext.scalar_values = ["wf-app-log", "message-ref"] - execution = SimpleNamespace(id_="run-id", total_tokens=0) - task = TraceTask( - trace_type=TraceTaskName.WORKFLOW_TRACE, workflow_execution=execution, conversation_id="conv", user_id="user" - ) - result = task.workflow_trace(workflow_run_id="run-id", conversation_id="conv", user_id="user") - assert result.workflow_run_id == "run-id" - assert result.workflow_id == "wf-1" - - -def test_trace_task_moderation_trace(trace_task_message): - task = TraceTask(trace_type=TraceTaskName.MODERATION_TRACE, message_id="msg-id") - moderation_result = SimpleNamespace(action="block", preset_response="no", query="q", flagged=True) - timer = {"start": 1, "end": 2} - result = task.moderation_trace("msg-id", timer, moderation_result=moderation_result, inputs={"src": "payload"}) - assert result.flagged is True - assert result.message_id == "log-id" - - -def test_trace_task_suggested_question_trace(trace_task_message): - task = TraceTask(trace_type=TraceTaskName.SUGGESTED_QUESTION_TRACE, message_id="msg-id") - timer = {"start": 1, "end": 2} - result = task.suggested_question_trace("msg-id", timer, suggested_question=["q1"]) - assert result.message_id == "log-id" - assert "suggested_question" in result.__dict__ - - -def test_trace_task_dataset_retrieval_trace(trace_task_message): - task = TraceTask(trace_type=TraceTaskName.DATASET_RETRIEVAL_TRACE, message_id="msg-id") - timer = {"start": 1, "end": 2} - mock_doc = SimpleNamespace(model_dump=lambda: {"doc": "value"}) - result = task.dataset_retrieval_trace("msg-id", timer, documents=[mock_doc]) - assert result.documents == [{"doc": "value"}] - - -def test_trace_task_tool_trace(monkeypatch: pytest.MonkeyPatch, mock_db): - custom_message = make_message_data(agent_thoughts=[make_agent_thought("tool-a", datetime(2025, 2, 20, 12, 1, 0))]) - monkeypatch.setattr("core.ops.ops_trace_manager.get_message_data", lambda _: custom_message) - configure_db_scalar(mock_db, message_file=FakeMessageFile()) - task = TraceTask(trace_type=TraceTaskName.TOOL_TRACE, message_id="msg-id") - timer = {"start": 1, "end": 5} - result = task.tool_trace("msg-id", timer, tool_name="tool-a", tool_inputs={"foo": 1}, tool_outputs="result") assert result.tool_name == "tool-a" assert result.time_cost == 5 + assert result.message_file_data.id == file.id -def test_trace_task_generate_name_trace(): - task = TraceTask(trace_type=TraceTaskName.GENERATE_NAME_TRACE, conversation_id="conv-id") - timer = {"start": 1, "end": 2} - assert task.generate_name_trace("conv-id", timer, tenant_id=None) == {} - result = task.generate_name_trace( - "conv-id", timer, tenant_id="tenant", generate_conversation_name="name", inputs="q" +def test_node_execution_trace_resolves_real_message_by_conversation_and_run( + trace_environment: None, database: Session +) -> None: + app = _app(database) + conversation, message = _conversation_message(database, app) + result = TraceTask(trace_type=TraceTaskName.NODE_EXECUTION_TRACE).node_execution_trace( + node_execution_data={ + "tenant_id": app.tenant_id, + "app_id": app.id, + "conversation_id": conversation.id, + "workflow_execution_id": message.workflow_run_id, + "workflow_id": "workflow-1", + "node_execution_id": "node-execution-1", + "node_id": "node-1", + "node_type": "llm", + "title": "Node", + "status": "succeeded", + } ) - assert result.outputs == "name" - assert result.tenant_id == "tenant" + assert result.message_id == message.id + assert result.metadata["conversation_id"] == conversation.id -def test_extract_streaming_metrics_invalid_json(): - task = TraceTask(trace_type=TraceTaskName.MESSAGE_TRACE, message_id="msg-id") - fake_message = make_message_data(message_metadata="invalid") - assert task._extract_streaming_metrics(fake_message) == {} - - -def test_trace_queue_manager_add_and_collect(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr( - "core.ops.ops_trace_manager.OpsTraceManager.get_ops_trace_instance", classmethod(lambda cls, aid: True) +def test_trace_helpers_and_streaming_metrics(trace_environment: None) -> None: + assert OpsTraceManager.check_trace_config_is_effective({}, "dummy") + assert OpsTraceManager.get_trace_config_project_key({}, "dummy") == "fake-key" + assert OpsTraceManager.get_trace_config_project_url({}, "dummy") == "https://project.fake" + task = TraceTask(trace_type=TraceTaskName.MESSAGE_TRACE) + assert task.conversation_trace(foo="bar") == {"foo": "bar"} + assert task._extract_streaming_metrics(_message_data(message_metadata="invalid")) == {} + assert task.generate_name_trace("conversation", {"start": 1, "end": 2}, tenant_id=None) == {} + generated = task.generate_name_trace( + "conversation", + {"start": 1, "end": 2}, + tenant_id="tenant-1", + generate_conversation_name="name", + inputs="query", + ) + assert generated.outputs == "name" + assert generated.tenant_id == "tenant-1" + + +def test_trace_queue_collect_run_and_storage_boundary(monkeypatch: pytest.MonkeyPatch, trace_environment: None) -> None: + monkeypatch.setattr(OpsTraceManager, "get_ops_trace_instance", classmethod(lambda cls, _app_id: True)) + manager = TraceQueueManager(app_id="app-id", user_id="user-1") + task = TraceTask( + trace_type=TraceTaskName.GENERATE_NAME_TRACE, + conversation_id="conversation-1", + timer={"start": 1, "end": 2}, + tenant_id="tenant-1", + generate_conversation_name="name", + inputs="query", ) - manager = TraceQueueManager(app_id="app-id", user_id="user") - task = TraceTask(trace_type=TraceTaskName.CONVERSATION_TRACE) manager.add_trace_task(task) - tasks = manager.collect_tasks() - assert tasks == [task] + assert manager.collect_tasks() == [task] - -def test_trace_queue_manager_run_invokes_send(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr( - "core.ops.ops_trace_manager.OpsTraceManager.get_ops_trace_instance", classmethod(lambda cls, aid: True) - ) - manager = TraceQueueManager(app_id="app-id", user_id="user") - task = TraceTask(trace_type=TraceTaskName.CONVERSATION_TRACE) - called = {} - - def fake_collect(): - return [task] - - def fake_send(tasks): - called["tasks"] = tasks - - monkeypatch.setattr(TraceQueueManager, "collect_tasks", lambda self: fake_collect()) - monkeypatch.setattr(TraceQueueManager, "send_to_celery", lambda self, t: fake_send(t)) + recording_storage = RecordingStorage() + dispatcher = RecordingDispatcher() + monkeypatch.setattr(module.storage, "save", recording_storage.save) + monkeypatch.setattr(module.process_trace_tasks, "delay", dispatcher.delay) + file_id = UUID("00000000-0000-0000-0000-000000000123") + monkeypatch.setattr(module, "uuid4", lambda: file_id) + manager.add_trace_task(task) manager.run() - assert called["tasks"] == [task] - -def test_trace_queue_manager_send_to_celery(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr( - "core.ops.ops_trace_manager.OpsTraceManager.get_ops_trace_instance", classmethod(lambda cls, aid: True) - ) - storage_save = MagicMock() - process_delay = MagicMock() - monkeypatch.setattr("core.ops.ops_trace_manager.storage.save", storage_save) - monkeypatch.setattr("core.ops.ops_trace_manager.process_trace_tasks.delay", process_delay) - monkeypatch.setattr("core.ops.ops_trace_manager.uuid4", MagicMock(return_value=SimpleNamespace(hex="file-123"))) - - manager = TraceQueueManager(app_id="app-id", user_id="user") - - class DummyTraceInfo: - def model_dump(self): - return {"trace": "info"} - - class DummyTask: - def __init__(self): - self.app_id = "app-id" - - def execute(self): - return DummyTraceInfo() - - task = DummyTask() - manager.send_to_celery([task]) - storage_save.assert_called_once() - process_delay.assert_called_once_with({"file_id": "file-123", "app_id": "app-id"}) + assert len(recording_storage.writes) == 1 + path, data = recording_storage.writes[0] + assert path.endswith(f"app-id/{file_id.hex}.json") + assert json.loads(data)["app_id"] == "app-id" + assert dispatcher.payloads == [{"file_id": file_id.hex, "app_id": "app-id"}] diff --git a/api/tests/unit_tests/core/ops/test_utils.py b/api/tests/unit_tests/core/ops/test_utils.py index 8a89422782e..6960a11f415 100644 --- a/api/tests/unit_tests/core/ops/test_utils.py +++ b/api/tests/unit_tests/core/ops/test_utils.py @@ -1,9 +1,11 @@ import re from datetime import datetime -from unittest.mock import MagicMock, patch +from decimal import Decimal import pytest +from sqlalchemy.orm import Session +import core.ops.utils as utils_module from core.ops.utils import ( filter_none_values, generate_dotted_order, @@ -15,6 +17,42 @@ from core.ops.utils import ( validate_url, validate_url_with_path, ) +from models.enums import ConversationFromSource +from models.model import Message + + +class _DatabaseBinding: + """Expose the real SQLite session used by the message lookup helper.""" + + session: Session + + def __init__(self, session: Session) -> None: + self.session = session + + +@pytest.fixture +def message_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Bind the message lookup helper to the shared SQLite test session.""" + + monkeypatch.setattr(utils_module, "db", _DatabaseBinding(sqlite_session)) + return sqlite_session + + +def _message(message_id: str) -> Message: + message = Message( + id=message_id, + app_id="app-id", + conversation_id="conversation-id", + query="question", + message={"role": "user", "content": "question"}, + answer="answer", + message_unit_price=Decimal("0.0001"), + answer_unit_price=Decimal("0.0001"), + currency="USD", + from_source=ConversationFromSource.API, + ) + message._inputs = {} + return message class TestValidateUrl: @@ -220,22 +258,20 @@ class TestFilterNoneValues: assert filter_none_values({}) == {} +@pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) class TestGetMessageData: """Test cases for get_message_data function""" - @patch("core.ops.utils.db") - @patch("core.ops.utils.Message") - @patch("core.ops.utils.select") - def test_get_message_data(self, mock_select, mock_message, mock_db): - mock_scalar = mock_db.session.scalar - mock_msg_instance = MagicMock() - mock_scalar.return_value = mock_msg_instance + def test_get_message_data(self, message_session: Session): + target = _message("message-id") + unrelated = _message("other-message-id") + message_session.add_all((target, unrelated)) + message_session.commit() result = get_message_data("message-id") - assert result == mock_msg_instance - mock_select.assert_called_once() - mock_scalar.assert_called_once() + assert result is target + assert result.id == "message-id" class TestMeasureTime: diff --git a/api/tests/unit_tests/core/plugin/impl/test_model_client.py b/api/tests/unit_tests/core/plugin/impl/test_model_client.py index c707b52ccaf..70c26a6cc0e 100644 --- a/api/tests/unit_tests/core/plugin/impl/test_model_client.py +++ b/api/tests/unit_tests/core/plugin/impl/test_model_client.py @@ -28,6 +28,19 @@ class TestPluginModelClient: ) assert request_mock.call_args.kwargs["params"] == {"page": 1, "page_size": 256} + def test_fetch_model_provider_bindings(self, mocker: MockerFixture): + client = PluginModelClient() + request_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response", return_value=["binding-a"]) + + result = client.fetch_model_provider_bindings("tenant-1") + + assert result == ["binding-a"] + assert request_mock.call_args.args[:2] == ( + "GET", + "plugin/tenant-1/management/models/bindings", + ) + assert "params" not in request_mock.call_args.kwargs + def test_get_model_schema(self, mocker: MockerFixture): client = PluginModelClient() schema = SimpleNamespace(name="schema") diff --git a/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py b/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py index 6fa94fbbca5..3716628927b 100644 --- a/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py +++ b/api/tests/unit_tests/core/plugin/test_backwards_invocation_app.py @@ -1,30 +1,80 @@ import json -from types import SimpleNamespace -from unittest.mock import MagicMock +from typing import cast import pytest from pydantic import BaseModel from pytest_mock import MockerFixture -from sqlalchemy.dialects import postgresql +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session, sessionmaker from core.app.layers.pause_state_persist_layer import PauseStateLayerConfig from core.plugin.backwards_invocation.app import PluginAppBackwardsInvocation from core.plugin.backwards_invocation.base import BaseBackwardsInvocation -from models.model import AppMode +from models import Account, Tenant, TenantAccountJoin +from models.enums import EndUserType +from models.model import App, AppMode, AppModelConfig, EndUser +from models.workflow import Workflow, WorkflowType class _Chunk(BaseModel): value: int -def _build_app_model_config(result: dict | None = None): - app_model_config = MagicMock() - app_model_config.app_id = "app-1" - app_model_config.to_dict.return_value = result or { - "user_input_form": [{"name": "bar"}], - "annotation_reply": {"enabled": False}, - } - return app_model_config +class _DatabaseWithEngine: + def __init__(self, engine: Engine) -> None: + self.engine = engine + + +def _app( + *, + app_id: str = "app-1", + tenant_id: str = "tenant-1", + mode: AppMode = AppMode.WORKFLOW, + workflow_id: str | None = None, + app_model_config_id: str | None = None, +) -> App: + return App( + id=app_id, + tenant_id=tenant_id, + name="Plugin app", + description="", + mode=mode, + enable_site=False, + enable_api=False, + workflow_id=workflow_id, + app_model_config_id=app_model_config_id, + ) + + +def _workflow(*, workflow_id: str = "workflow-1", app_id: str = "app-1", tenant_id: str = "tenant-1") -> Workflow: + return Workflow( + id=workflow_id, + tenant_id=tenant_id, + app_id=app_id, + type=WorkflowType.WORKFLOW, + version=Workflow.VERSION_DRAFT, + graph="{}", + _features="{}", + created_by="account-1", + ) + + +def _end_user( + *, + user_id: str = "user-1", + tenant_id: str = "tenant-1", + app_id: str = "app-1", + session_id: str = "browser-session", +) -> EndUser: + return EndUser( + id=user_id, + tenant_id=tenant_id, + app_id=app_id, + type=EndUserType.BROWSER, + session_id=session_id, + name="Browser user", + is_anonymous=True, + ) class TestBaseBackwardsInvocation: @@ -53,23 +103,25 @@ class TestBaseBackwardsInvocation: class TestPluginAppBackwardsInvocation: - def patch_create_session(self, mocker: MockerFixture, *, return_value=None, side_effect=None): - session = MagicMock() - if side_effect is not None: - session.scalar.side_effect = side_effect - else: - session.scalar.return_value = return_value - session_ctx = MagicMock() - session_ctx.__enter__.return_value = session - session_ctx.__exit__.return_value = None - mocker.patch("core.plugin.backwards_invocation.app.create_session", return_value=session_ctx) - return session + @pytest.fixture(autouse=True) + def _real_sessions( + self, + mocker: MockerFixture, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + sqlite_engine: Engine, + ) -> None: + self.session = sqlite_session + self.session_factory = sqlite_session_factory + self.sqlite_engine = sqlite_engine + mocker.patch("core.plugin.backwards_invocation.app.create_session", side_effect=sqlite_session_factory) def test_fetch_app_info_workflow_path(self, mocker: MockerFixture): - workflow = MagicMock() - workflow.features_dict = {"feature": "v"} - workflow.user_input_form.return_value = [{"name": "foo"}] - app = MagicMock(mode=AppMode.WORKFLOW) + variable = {"type": "text-input", "variable": "foo", "label": "Foo", "required": False} + workflow = _workflow() + workflow.features = json.dumps({"feature": "v"}) + workflow.graph = json.dumps({"nodes": [{"data": {"type": "start", "variables": [variable]}}]}) + app = _app(mode=AppMode.WORKFLOW) mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow) mapper = mocker.patch( @@ -80,11 +132,11 @@ class TestPluginAppBackwardsInvocation: result = PluginAppBackwardsInvocation.fetch_app_info("app-1", "tenant-1") assert result == {"data": {"mapped": True}} - mapper.assert_called_once_with(features_dict={"feature": "v"}, user_input_form=[{"name": "foo"}]) + mapper.assert_called_once_with(features_dict={"feature": "v"}, user_input_form=[{"text-input": variable}]) def test_fetch_app_info_model_config_path(self, mocker: MockerFixture): model_config_dict = {"user_input_form": [{"name": "bar"}], "k": "v"} - app = MagicMock(mode=AppMode.COMPLETION) + app = _app(mode=AppMode.COMPLETION) mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app) mocker.patch.object(PluginAppBackwardsInvocation, "_get_app_model_config_dict", return_value=model_config_dict) mocker.patch( @@ -107,9 +159,9 @@ class TestPluginAppBackwardsInvocation: ], ) def test_invoke_app_routes_by_mode(self, mocker: MockerFixture, mode, route_method): - app = MagicMock(mode=mode) - user = MagicMock() - workflow = MagicMock() + app = _app(mode=mode) + user = _end_user() + workflow = _workflow() mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app) mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=user) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow) @@ -124,16 +176,16 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={"x": 1}, files=[], - session=MagicMock(), + session=self.session, ) assert result == {"routed": True} assert route.call_count == 1 def test_invoke_app_uses_end_user_when_user_id_missing(self, mocker: MockerFixture): - app = MagicMock(mode=AppMode.WORKFLOW) - end_user = MagicMock() - workflow = MagicMock() + app = _app(mode=AppMode.WORKFLOW) + end_user = _end_user() + workflow = _workflow() mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow) get_or_create = mocker.patch( @@ -151,7 +203,7 @@ class TestPluginAppBackwardsInvocation: stream=True, inputs={}, files=[], - session=MagicMock(), + session=self.session, ) assert result == {"ok": True} @@ -160,8 +212,8 @@ class TestPluginAppBackwardsInvocation: assert route.call_args.args[2] is end_user def test_invoke_app_missing_query_for_chat_raises(self, mocker: MockerFixture): - mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=MagicMock(mode=AppMode.CHAT)) - mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=MagicMock()) + mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=_app(mode=AppMode.CHAT)) + mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=_end_user()) with pytest.raises(ValueError, match="missing query"): PluginAppBackwardsInvocation.invoke_app( @@ -173,12 +225,16 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={}, files=[], - session=MagicMock(), + session=self.session, ) def test_invoke_app_unexpected_mode_raises(self, mocker: MockerFixture): - mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=MagicMock(mode="other")) - mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=MagicMock()) + mocker.patch.object( + PluginAppBackwardsInvocation, + "_get_app", + return_value=_app(mode=cast(AppMode, "other")), + ) + mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=_end_user()) with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_app( @@ -190,7 +246,7 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={}, files=[], - session=MagicMock(), + session=self.session, ) @pytest.mark.parametrize( @@ -201,44 +257,43 @@ class TestPluginAppBackwardsInvocation: ], ) def test_invoke_chat_app_agent_and_chat(self, mocker: MockerFixture, mode, generator_path): - app = MagicMock(mode=mode, workflow=None) + app = _app(mode=mode) spy = mocker.patch(generator_path, return_value={"result": "ok"}) result = PluginAppBackwardsInvocation.invoke_chat_app( app=app, - user=MagicMock(), + user=_end_user(), conversation_id="conv-1", query="hello", stream=False, inputs={"k": "v"}, files=[], - session=MagicMock(), + session=self.session, ) assert result == {"result": "ok"} assert spy.call_count == 1 def test_invoke_chat_app_advanced_chat_injects_pause_state_config(self, mocker: MockerFixture): - workflow = MagicMock() + workflow = _workflow() workflow.created_by = "owner-id" - app = MagicMock() - app.mode = AppMode.ADVANCED_CHAT + app = _app(mode=AppMode.ADVANCED_CHAT) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow) mocker.patch( "core.plugin.backwards_invocation.app.db", - SimpleNamespace(engine=MagicMock()), + _DatabaseWithEngine(self.sqlite_engine), ) generator_spy = mocker.patch( "core.plugin.backwards_invocation.app.AdvancedChatAppGenerator.generate", return_value={"result": "ok"}, ) - session = MagicMock() + session = self.session result = PluginAppBackwardsInvocation.invoke_chat_app( app=app, - user=MagicMock(), + user=_end_user(), conversation_id="conv-1", query="hello", stream=False, @@ -255,44 +310,43 @@ class TestPluginAppBackwardsInvocation: assert pause_state_config.state_owner_user_id == "owner-id" def test_invoke_chat_app_advanced_chat_without_workflow_raises(self, mocker: MockerFixture): - app = MagicMock(mode=AppMode.ADVANCED_CHAT) + app = _app(mode=AppMode.ADVANCED_CHAT) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=None) with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_chat_app( app=app, - user=MagicMock(), + user=_end_user(), conversation_id="conv-1", query="hello", stream=False, inputs={}, files=[], - session=MagicMock(), + session=self.session, ) def test_invoke_chat_app_unexpected_mode_raises(self): - app = MagicMock(mode="invalid") + app = _app(mode=cast(AppMode, "invalid")) with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_chat_app( app=app, - user=MagicMock(), + user=_end_user(), conversation_id="conv-1", query="hello", stream=False, inputs={}, files=[], - session=MagicMock(), + session=self.session, ) def test_invoke_workflow_app_injects_pause_state_config(self, mocker: MockerFixture): - workflow = MagicMock() + workflow = _workflow() workflow.created_by = "owner-id" - app = MagicMock() - app.mode = AppMode.WORKFLOW + app = _app(mode=AppMode.WORKFLOW) mocker.patch( "core.plugin.backwards_invocation.app.db", - SimpleNamespace(engine=MagicMock()), + _DatabaseWithEngine(self.sqlite_engine), ) generator_spy = mocker.patch( "core.plugin.backwards_invocation.app.WorkflowAppGenerator.generate", @@ -302,7 +356,7 @@ class TestPluginAppBackwardsInvocation: result = PluginAppBackwardsInvocation.invoke_workflow_app( app=app, workflow=workflow, - user=MagicMock(), + user=_end_user(), stream=False, inputs={"k": "v"}, files=[], @@ -315,9 +369,9 @@ class TestPluginAppBackwardsInvocation: assert pause_state_config.state_owner_user_id == "owner-id" def test_invoke_app_workflow_without_workflow_raises(self, mocker: MockerFixture): - app = MagicMock(mode=AppMode.WORKFLOW) + app = _app(mode=AppMode.WORKFLOW) mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app) - mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=MagicMock()) + mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", return_value=_end_user()) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=None) with pytest.raises(ValueError, match="unexpected app type"): PluginAppBackwardsInvocation.invoke_app( @@ -329,77 +383,140 @@ class TestPluginAppBackwardsInvocation: stream=False, inputs={}, files=[], - session=MagicMock(), + session=self.session, ) def test_invoke_completion_app(self, mocker: MockerFixture): spy = mocker.patch( "core.plugin.backwards_invocation.app.CompletionAppGenerator.generate", return_value={"ok": 1} ) - app = MagicMock(mode=AppMode.COMPLETION) + app = _app(mode=AppMode.COMPLETION) - result = PluginAppBackwardsInvocation.invoke_completion_app(app, MagicMock(), False, {"x": 1}, [], MagicMock()) + result = PluginAppBackwardsInvocation.invoke_completion_app(app, _end_user(), False, {"x": 1}, [], self.session) assert result == {"ok": 1} assert spy.call_count == 1 - def test_get_user_returns_end_user(self, mocker: MockerFixture): - session = self.patch_create_session(mocker, side_effect=[MagicMock(id="end-user")]) - app = SimpleNamespace(id="app-1", tenant_id="tenant-1") + def test_get_user_returns_end_user(self): + app = _app() + end_user = EndUser( + id="uid", + tenant_id=app.tenant_id, + app_id=app.id, + type=EndUserType.BROWSER, + session_id="browser-session", + name="Browser user", + is_anonymous=True, + ) + self.session.add(end_user) + self.session.commit() user = PluginAppBackwardsInvocation._get_user("uid", app) - assert user.id == "end-user" - stmt = session.scalar.call_args_list[0].args[0] - compiled = str(stmt.compile(dialect=postgresql.dialect())) - assert "end_users.id" in compiled - assert "end_users.tenant_id" in compiled - assert "end_users.app_id" in compiled - assert stmt.compile().params == {"id_1": "uid", "tenant_id_1": "tenant-1", "app_id_1": "app-1"} + assert user.id == "uid" + assert user.tenant_id == app.tenant_id + assert user.app_id == app.id - def test_get_user_returns_end_user_by_session_id(self, mocker: MockerFixture): - session = self.patch_create_session(mocker, side_effect=[None, MagicMock(id="session-user")]) - app = SimpleNamespace(id="app-1", tenant_id="tenant-1") + def test_get_user_returns_end_user_by_session_id(self): + app = _app() + end_user = EndUser( + id="session-user", + tenant_id=app.tenant_id, + app_id=app.id, + type=EndUserType.BROWSER, + session_id="wecom-sender-1", + name="External user", + is_anonymous=True, + ) + self.session.add(end_user) + self.session.commit() user = PluginAppBackwardsInvocation._get_user("wecom-sender-1", app) assert user.id == "session-user" - stmt = session.scalar.call_args_list[1].args[0] - compiled = str(stmt.compile(dialect=postgresql.dialect())) - assert "end_users.session_id" in compiled - assert "end_users.tenant_id" in compiled - assert "end_users.app_id" in compiled - assert stmt.compile().params == { - "session_id_1": "wecom-sender-1", - "tenant_id_1": "tenant-1", - "app_id_1": "app-1", - } - def test_get_user_falls_back_to_account_user(self, mocker: MockerFixture): - session = self.patch_create_session(mocker, side_effect=[None, None, MagicMock(id="account-user")]) - app = SimpleNamespace(id="app-1", tenant_id="tenant-1") + def test_get_user_rejects_end_user_from_another_app(self): + app = _app() + end_user = EndUser( + id="uid", + tenant_id=app.tenant_id, + app_id="other-app", + type=EndUserType.BROWSER, + session_id="browser-session", + name="Browser user", + is_anonymous=True, + ) + self.session.add(end_user) + self.session.commit() - user = PluginAppBackwardsInvocation._get_user("uid", app) + with pytest.raises(ValueError, match="user not found"): + PluginAppBackwardsInvocation._get_user("uid", app) + + def test_get_user_rejects_nonmatching_session_id(self): + app = _app() + end_user = EndUser( + id="session-user", + tenant_id=app.tenant_id, + app_id=app.id, + type=EndUserType.BROWSER, + session_id="other-session", + name="External user", + is_anonymous=True, + ) + self.session.add(end_user) + self.session.commit() + + with pytest.raises(ValueError, match="user not found"): + PluginAppBackwardsInvocation._get_user("wecom-sender-1", app) + + def test_get_user_falls_back_to_account_user(self): + app = _app() + tenant = Tenant(name="Plugin tenant") + tenant.id = app.tenant_id + account = Account(name="Account user", email="account-user@example.com") + account.id = "account-user" + membership = TenantAccountJoin(tenant_id=tenant.id, account_id=account.id) + self.session.add_all([tenant, account, membership]) + self.session.commit() + + user = PluginAppBackwardsInvocation._get_user(account.id, app) assert user.id == "account-user" - stmt = session.scalar.call_args_list[2].args[0] - compiled = str(stmt.compile(dialect=postgresql.dialect())) - assert "accounts.id" in compiled - assert "tenant_account_joins.account_id" in compiled - assert "tenant_account_joins.tenant_id" in compiled - assert stmt.compile().params == {"id_1": "uid", "tenant_id_1": "tenant-1"} - def test_get_user_raises_when_user_not_found(self, mocker: MockerFixture): - self.patch_create_session(mocker, side_effect=[None, None, None]) - app = SimpleNamespace(id="app-1", tenant_id="tenant-1") + def test_get_user_rejects_account_from_another_tenant(self): + app = _app() + tenant = Tenant(name="Plugin tenant") + tenant.id = "other-tenant" + account = Account(name="Account user", email="account-user@example.com") + account.id = "account-user" + membership = TenantAccountJoin(tenant_id=tenant.id, account_id=account.id) + self.session.add_all([tenant, account, membership]) + self.session.commit() + + with pytest.raises(ValueError, match="user not found"): + PluginAppBackwardsInvocation._get_user(account.id, app) + + def test_get_user_raises_when_user_not_found(self): + app = _app() + other_tenant_user = EndUser( + id="uid", + tenant_id="other-tenant", + app_id=app.id, + type=EndUserType.BROWSER, + session_id="uid", + name="Wrong tenant", + is_anonymous=True, + ) + self.session.add(other_tenant_user) + self.session.commit() with pytest.raises(ValueError, match="user not found"): PluginAppBackwardsInvocation._get_user("uid", app) def test_invoke_app_creates_end_user_for_unknown_external_user_id(self, mocker: MockerFixture): - app = MagicMock(mode=AppMode.WORKFLOW) - end_user = MagicMock() - workflow = MagicMock() + app = _app(mode=AppMode.WORKFLOW) + end_user = _end_user(session_id="wecom-sender-1") + workflow = _workflow() mocker.patch.object(PluginAppBackwardsInvocation, "_get_app", return_value=app) mocker.patch.object(PluginAppBackwardsInvocation, "_get_workflow", return_value=workflow) mocker.patch.object(PluginAppBackwardsInvocation, "_get_user", side_effect=ValueError("user not found")) @@ -418,70 +535,92 @@ class TestPluginAppBackwardsInvocation: stream=True, inputs={}, files=[], - session=MagicMock(), + session=self.session, ) assert result == {"ok": True} get_or_create.assert_called_once_with(app, user_id="wecom-sender-1") assert route.call_args.args[2] is end_user - def test_get_app_returns_app(self, mocker: MockerFixture): - app_obj = MagicMock(id="app") - self.patch_create_session(mocker, return_value=app_obj) + def test_get_app_returns_app(self): + app_obj = _app(app_id="app", tenant_id="tenant") + self.session.add(app_obj) + self.session.commit() - assert PluginAppBackwardsInvocation._get_app("app", "tenant") is app_obj + result = PluginAppBackwardsInvocation._get_app("app", "tenant") + assert result.id == app_obj.id + assert result.tenant_id == app_obj.tenant_id - def test_get_app_raises_when_missing(self, mocker: MockerFixture): - self.patch_create_session(mocker, return_value=None) + def test_get_app_raises_when_missing(self): + self.session.add(_app(app_id="app", tenant_id="other-tenant")) + self.session.commit() with pytest.raises(ValueError, match="app not found"): PluginAppBackwardsInvocation._get_app("app", "tenant") - def test_get_app_raises_when_query_fails(self, mocker: MockerFixture): - self.patch_create_session(mocker, side_effect=RuntimeError("db down")) + def test_get_app_raises_when_query_fails(self): + def fail_query(*_args, **_kwargs): + raise RuntimeError("db down") + + event.listen(self.sqlite_engine, "before_cursor_execute", fail_query, once=True) with pytest.raises(ValueError, match="app not found"): PluginAppBackwardsInvocation._get_app("app", "tenant") - def test_get_workflow_stays_inside_app_boundary(self, mocker: MockerFixture): - workflow = MagicMock(id="workflow") - session = self.patch_create_session(mocker, return_value=workflow) - app = SimpleNamespace(id="app-1", tenant_id="tenant-1", workflow_id="workflow-1") + def test_get_workflow_stays_inside_app_boundary(self): + workflow = _workflow() + other_workflow = _workflow(workflow_id="workflow-other", tenant_id="other-tenant") + self.session.add_all([workflow, other_workflow]) + self.session.commit() + app = _app(workflow_id="workflow-1") - assert PluginAppBackwardsInvocation._get_workflow(app) is workflow + result = PluginAppBackwardsInvocation._get_workflow(app) + assert result is not None + assert result.id == workflow.id + assert result.tenant_id == app.tenant_id - stmt = session.scalar.call_args.args[0] - compiled = str(stmt.compile(dialect=postgresql.dialect())) - assert "workflows.id" in compiled - assert "workflows.tenant_id" in compiled - assert "workflows.app_id" in compiled - assert stmt.compile().params == { - "id_1": "workflow-1", - "tenant_id_1": "tenant-1", - "app_id_1": "app-1", - "param_1": 1, - } + def test_get_workflow_rejects_workflow_from_another_tenant(self): + workflow = _workflow(tenant_id="other-tenant") + self.session.add(workflow) + self.session.commit() + app = _app(app_id=workflow.app_id, tenant_id="tenant-1", workflow_id=workflow.id) + + assert PluginAppBackwardsInvocation._get_workflow(app) is None + + def test_get_workflow_rejects_workflow_from_another_app(self): + workflow = _workflow(app_id="other-app") + self.session.add(workflow) + self.session.commit() + app = _app(app_id="app-1", tenant_id=workflow.tenant_id, workflow_id=workflow.id) + + assert PluginAppBackwardsInvocation._get_workflow(app) is None def test_get_app_model_config_dict_uses_explicit_session_for_annotation_reply(self, mocker: MockerFixture): annotation_reply = {"enabled": False} - app_model_config = _build_app_model_config() - session = self.patch_create_session(mocker, return_value=app_model_config) + app_model_config = AppModelConfig(app_id="app-1", user_input_form=json.dumps([{"name": "bar"}])) + app_model_config.id = "config-1" + self.session.add(app_model_config) + self.session.commit() load_annotation_reply_config = mocker.patch( "core.plugin.backwards_invocation.app.load_annotation_reply_config", return_value=annotation_reply, ) - app = SimpleNamespace(id="app-1", app_model_config_id="config-1") + app = _app(app_model_config_id="config-1") result = PluginAppBackwardsInvocation._get_app_model_config_dict(app) assert result is not None assert result["user_input_form"] == [{"name": "bar"}] assert result["annotation_reply"] == annotation_reply - load_annotation_reply_config.assert_called_once_with(session, "app-1") - app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + queried_session, queried_app_id = load_annotation_reply_config.call_args.args + assert isinstance(queried_session, Session) + assert queried_app_id == "app-1" - stmt = session.scalar.call_args.args[0] - compiled = str(stmt.compile(dialect=postgresql.dialect())) - assert "app_model_configs.id" in compiled - assert "app_model_configs.app_id" in compiled - assert stmt.compile().params == {"id_1": "config-1", "app_id_1": "app-1", "param_1": 1} + def test_get_app_model_config_dict_rejects_config_from_another_app(self): + app_model_config = AppModelConfig(app_id="other-app", user_input_form=json.dumps([{"name": "bar"}])) + app_model_config.id = "config-1" + self.session.add(app_model_config) + self.session.commit() + app = _app(app_model_config_id=app_model_config.id) + + assert PluginAppBackwardsInvocation._get_app_model_config_dict(app) is None diff --git a/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py b/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py index a6f80505e62..a436116f838 100644 --- a/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py +++ b/api/tests/unit_tests/core/plugin/test_model_runtime_adapter.py @@ -7,6 +7,7 @@ from unittest.mock import MagicMock, Mock, patch, sentinel import pytest +from core.plugin.entities.plugin import PluginInstallationSource from core.plugin.entities.plugin_daemon import PluginModelProviderEntity from core.plugin.impl import model_runtime as model_runtime_module from core.plugin.impl.model import PluginModelClient @@ -93,6 +94,7 @@ def _build_plugin_model_provider(*, tenant_id: str, provider: str = "openai") -> tenant_id=tenant_id, plugin_unique_identifier=f"langgenius/{provider}/{provider}", plugin_id=f"langgenius/{provider}", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider=provider, label=I18nObject(en_US=provider.title()), @@ -116,6 +118,7 @@ class TestPluginModelRuntime: tenant_id="tenant", plugin_unique_identifier="langgenius/openai/openai", plugin_id="langgenius/openai", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider="openai", label=I18nObject(en_US="OpenAI"), @@ -145,6 +148,7 @@ class TestPluginModelRuntime: tenant_id="tenant", plugin_unique_identifier="acme/openai/openai", plugin_id="acme/openai", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider="openai", label=I18nObject(en_US="Acme OpenAI"), @@ -160,6 +164,7 @@ class TestPluginModelRuntime: tenant_id="tenant", plugin_unique_identifier="langgenius/openai/openai", plugin_id="langgenius/openai", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider="openai", label=I18nObject(en_US="OpenAI"), @@ -187,6 +192,7 @@ class TestPluginModelRuntime: tenant_id="tenant", plugin_unique_identifier="langgenius/gemini/google", plugin_id="langgenius/gemini", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider="google", label=I18nObject(en_US="Google"), @@ -821,6 +827,7 @@ def test_get_provider_icon_reads_requested_variant_and_detects_svg_mime(monkeypa tenant_id="tenant", plugin_unique_identifier="langgenius/openai/openai", plugin_id="langgenius/openai", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider="openai", label=I18nObject(en_US="OpenAI"), @@ -857,6 +864,7 @@ def test_get_provider_icon_rejects_unsupported_types_and_missing_variants() -> N tenant_id="tenant", plugin_unique_identifier="langgenius/openai/openai", plugin_id="langgenius/openai", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider="openai", label=I18nObject(en_US="OpenAI"), @@ -989,6 +997,7 @@ def test_get_provider_schema_supports_short_alias_and_rejects_invalid_provider() tenant_id="tenant", plugin_unique_identifier="langgenius/openai/openai", plugin_id="langgenius/openai", + installation_source=PluginInstallationSource.Marketplace, declaration=ProviderEntity( provider="openai", label=I18nObject(en_US="OpenAI"), diff --git a/api/tests/unit_tests/core/plugin/test_plugin_entities.py b/api/tests/unit_tests/core/plugin/test_plugin_entities.py index 0a532646abb..acc0d0be6fa 100644 --- a/api/tests/unit_tests/core/plugin/test_plugin_entities.py +++ b/api/tests/unit_tests/core/plugin/test_plugin_entities.py @@ -265,6 +265,11 @@ class TestPluginParameterEntities: with pytest.raises(ValueError, match="not found in tool config"): init_frontend_parameter(required_rule, PluginParameterType.STRING, None) + tools_rule = PluginParameter(name="tools", label=self._label(), required=True, default=None) + assert init_frontend_parameter(tools_rule, PluginParameterType.TOOLS_SELECTOR, []) == [] + with pytest.raises(ValueError, match="not found in tool config"): + init_frontend_parameter(tools_rule, PluginParameterType.TOOLS_SELECTOR, None) + class TestPluginDaemonEntities: def test_credential_type_helpers(self): diff --git a/api/tests/unit_tests/core/plugin/test_plugin_manager.py b/api/tests/unit_tests/core/plugin/test_plugin_manager.py index 1aa019254ff..290c7301bbe 100644 --- a/api/tests/unit_tests/core/plugin/test_plugin_manager.py +++ b/api/tests/unit_tests/core/plugin/test_plugin_manager.py @@ -29,6 +29,7 @@ from core.plugin.entities.plugin import ( ) from core.plugin.entities.plugin_daemon import ( PluginDecodeResponse, + PluginInstalledIdsDaemonResponse, PluginInstallTask, PluginInstallTaskStartResponse, PluginInstallTaskStatus, @@ -132,7 +133,13 @@ class TestPluginDiscovery: plugin_installer, "_request_with_plugin_daemon_response", return_value=mock_response ) as mock_request: result = plugin_installer.list_plugins_by_category( - "test-tenant", category=PluginCategory.Tool, page=2, page_size=10 + "test-tenant", + category=PluginCategory.Tool, + page=2, + page_size=10, + query="weather", + tags=["search", "rag"], + language="zh_Hans", ) mock_request.assert_called_once() @@ -141,6 +148,9 @@ class TestPluginDiscovery: assert call_args.args[2] is PluginListWithoutTotalResponse assert call_args.kwargs["params"]["page"] == 2 assert call_args.kwargs["params"]["page_size"] == 10 + assert call_args.kwargs["params"]["query"] == "weather" + assert call_args.kwargs["params"]["tags"] == ["search", "rag"] + assert call_args.kwargs["params"]["language"] == "zh_Hans" assert result.list == [mock_plugin_entity] assert result.has_more is True @@ -156,6 +166,23 @@ class TestPluginDiscovery: # Assert: Verify empty list is returned assert len(result) == 0 + def test_list_installed_plugin_ids(self, plugin_installer): + """The lightweight ID endpoint is unpaginated and does not request plugin details.""" + mock_response = PluginInstalledIdsDaemonResponse(plugin_ids=["langgenius/openai", "langgenius/anthropic"]) + + with patch.object( + plugin_installer, "_request_with_plugin_daemon_response", return_value=mock_response + ) as mock_request: + result = plugin_installer.list_installed_plugin_ids("test-tenant", PluginCategory.Tool) + + mock_request.assert_called_once_with( + "GET", + "plugin/test-tenant/management/installation/ids", + PluginInstalledIdsDaemonResponse, + params={"category": "tool"}, + ) + assert result == ["langgenius/openai", "langgenius/anthropic"] + def test_fetch_plugin_by_identifier_found(self, plugin_installer): """Test fetching a plugin by its unique identifier when it exists.""" # Arrange: Mock successful fetch diff --git a/api/tests/unit_tests/core/prompt/test_extract_thread_messages.py b/api/tests/unit_tests/core/prompt/test_extract_thread_messages.py index 3b38a9af40a..a9438e6aa18 100644 --- a/api/tests/unit_tests/core/prompt/test_extract_thread_messages.py +++ b/api/tests/unit_tests/core/prompt/test_extract_thread_messages.py @@ -1,9 +1,42 @@ -from unittest.mock import MagicMock +from datetime import datetime, timedelta +from decimal import Decimal from uuid import uuid4 +import pytest +from sqlalchemy.orm import Session + from constants import UUID_NIL from core.prompt.utils.extract_thread_messages import extract_thread_messages from core.prompt.utils.get_thread_messages_length import get_thread_messages_length +from models.enums import ConversationFromSource +from models.model import Message + + +def _persisted_message( + *, + message_id: str, + conversation_id: str, + parent_message_id: str, + answer: str, + created_at: datetime, +) -> Message: + message = Message( + id=message_id, + app_id="app-id", + conversation_id=conversation_id, + query="question", + message={"role": "user", "content": "question"}, + answer=answer, + message_unit_price=Decimal("0.0001"), + answer_unit_price=Decimal("0.0001"), + currency="USD", + from_source=ConversationFromSource.API, + parent_message_id=parent_message_id, + created_at=created_at, + updated_at=created_at, + ) + message._inputs = {} + return message class MockMessage: @@ -104,33 +137,64 @@ def test_extract_thread_messages_breaks_when_parent_is_none(): assert result[0].id == id2 -def test_get_thread_messages_length_excludes_newly_created_empty_answer(): +@pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) +def test_get_thread_messages_length_excludes_newly_created_empty_answer(sqlite_session: Session): id1, id2 = str(uuid4()), str(uuid4()) + now = datetime.now() messages = [ - MockMessage(id2, id1, answer=""), # newest generated message should be excluded - MockMessage(id1, UUID_NIL, answer="ok"), + _persisted_message( + message_id=id2, + conversation_id="conversation-1", + parent_message_id=id1, + answer="", + created_at=now, + ), + _persisted_message( + message_id=id1, + conversation_id="conversation-1", + parent_message_id=UUID_NIL, + answer="ok", + created_at=now - timedelta(seconds=1), + ), + _persisted_message( + message_id=str(uuid4()), + conversation_id="other-conversation", + parent_message_id=UUID_NIL, + answer="unrelated", + created_at=now + timedelta(seconds=1), + ), ] + sqlite_session.add_all(messages) + sqlite_session.commit() - session = MagicMock() - session.scalars.return_value.all.return_value = messages - - length = get_thread_messages_length("conversation-1", session=session) + length = get_thread_messages_length("conversation-1", session=sqlite_session) assert length == 1 - session.scalars.assert_called_once() -def test_get_thread_messages_length_keeps_non_empty_latest_answer(): +@pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) +def test_get_thread_messages_length_keeps_non_empty_latest_answer(sqlite_session: Session): id1, id2 = str(uuid4()), str(uuid4()) + now = datetime.now() messages = [ - MockMessage(id2, id1, answer="latest-answer"), - MockMessage(id1, UUID_NIL, answer="older-answer"), + _persisted_message( + message_id=id2, + conversation_id="conversation-2", + parent_message_id=id1, + answer="latest-answer", + created_at=now, + ), + _persisted_message( + message_id=id1, + conversation_id="conversation-2", + parent_message_id=UUID_NIL, + answer="older-answer", + created_at=now - timedelta(seconds=1), + ), ] + sqlite_session.add_all(messages) + sqlite_session.commit() - session = MagicMock() - session.scalars.return_value.all.return_value = messages - - length = get_thread_messages_length("conversation-2", session=session) + length = get_thread_messages_length("conversation-2", session=sqlite_session) assert length == 2 - session.scalars.assert_called_once() diff --git a/api/tests/unit_tests/core/rag/cleaner/test_clean_processor.py b/api/tests/unit_tests/core/rag/cleaner/test_clean_processor.py index c7a4265a954..2155342ab98 100644 --- a/api/tests/unit_tests/core/rag/cleaner/test_clean_processor.py +++ b/api/tests/unit_tests/core/rag/cleaner/test_clean_processor.py @@ -22,6 +22,23 @@ class TestCleanProcessor: expected = "normalpadding" assert CleanProcessor.clean(text_with_ufffe, None) == expected + def test_clean_preserves_valid_extended_characters(self): + """Default cleaning must not strip valid printable characters. + + The invalid-symbol filter used to include the UTF-8 bytes of U+FFFE + (0xEF 0xBF 0xBE) inside a character class. On a decoded string those + bytes are the code points U+00EF, U+00BF and U+00BE, i.e. the valid + characters 'ï', '¿' and '¾', so words like "naïve" and Spanish + questions like "¿Cómo?" were being silently corrupted on ingest. + """ + assert CleanProcessor.clean("naïve", None) == "naïve" + assert CleanProcessor.clean("¿Cómo estás?", None) == "¿Cómo estás?" + assert CleanProcessor.clean("¾ cup sugar", None) == "¾ cup sugar" + assert CleanProcessor.clean("￾", None) == "￾" + + # The U+FFFE noncharacter is still stripped by its dedicated substitution. + assert CleanProcessor.clean("keep\ufffedrop", None) == "keepdrop" + def test_clean_with_none_process_rule(self): """Test cleaning with None process_rule - only default cleaning applied.""" text = "Hello<|World\x00" diff --git a/api/tests/unit_tests/core/rag/data_post_processor/test_data_post_processor.py b/api/tests/unit_tests/core/rag/data_post_processor/test_data_post_processor.py index 96bb8e8bfc7..6d7cb1e8a85 100644 --- a/api/tests/unit_tests/core/rag/data_post_processor/test_data_post_processor.py +++ b/api/tests/unit_tests/core/rag/data_post_processor/test_data_post_processor.py @@ -1,5 +1,7 @@ from unittest.mock import MagicMock, patch +from sqlalchemy.orm import Session + from core.rag.data_post_processor.data_post_processor import DataPostProcessor from core.rag.data_post_processor.reorder import ReorderRunner from core.rag.index_processor.constant.query_type import QueryType @@ -14,10 +16,9 @@ def _doc(content: str) -> Document: class TestDataPostProcessor: - def test_init_sets_rerank_and_reorder_runners(self): + def test_init_sets_rerank_and_reorder_runners(self, unbound_session: Session): rerank_runner = object() reorder_runner = object() - session = MagicMock() with patch.object(DataPostProcessor, "_get_rerank_runner", return_value=rerank_runner) as rerank_mock: with patch.object(DataPostProcessor, "_get_reorder_runner", return_value=reorder_runner) as reorder_mock: @@ -27,7 +28,7 @@ class TestDataPostProcessor: reranking_model={"config": "value"}, weights={"weight": "value"}, reorder_enabled=True, - session=session, + session=unbound_session, ) assert processor.rerank_runner is rerank_runner @@ -37,7 +38,7 @@ class TestDataPostProcessor: "tenant-1", {"config": "value"}, {"weight": "value"}, - session=session, + session=unbound_session, ) reorder_mock.assert_called_once_with(True) @@ -79,7 +80,7 @@ class TestDataPostProcessor: assert processor.invoke(query="query", documents=documents) == documents - def test_get_rerank_runner_for_weighted_score(self): + def test_get_rerank_runner_for_weighted_score(self, unbound_session: Session): weights_config = { "vector_setting": { "vector_weight": 0.7, @@ -90,7 +91,6 @@ class TestDataPostProcessor: } expected_runner = object() processor = DataPostProcessor.__new__(DataPostProcessor) - session = MagicMock() with patch( "core.rag.data_post_processor.data_post_processor.RerankRunnerFactory.create_rerank_runner", @@ -101,7 +101,7 @@ class TestDataPostProcessor: tenant_id="tenant-1", reranking_model=None, weights=weights_config, - session=session, + session=unbound_session, ) assert result is expected_runner @@ -113,13 +113,12 @@ class TestDataPostProcessor: assert kwargs["weights"].vector_setting.embedding_model_name == "embedding-y" assert kwargs["weights"].keyword_setting.keyword_weight == 0.3 - def test_get_rerank_runner_for_reranking_model_returns_none_without_model_instance(self): + def test_get_rerank_runner_for_reranking_model_returns_none_without_model_instance(self, unbound_session: Session): processor = DataPostProcessor.__new__(DataPostProcessor) reranking_model = { "reranking_provider_name": "provider-x", "reranking_model_name": "model-y", } - session = MagicMock() with patch.object(DataPostProcessor, "_get_rerank_model_instance", return_value=None) as model_mock: with patch( @@ -130,18 +129,17 @@ class TestDataPostProcessor: tenant_id="tenant-1", reranking_model=reranking_model, weights=None, - session=session, + session=unbound_session, ) assert result is None model_mock.assert_called_once_with("tenant-1", reranking_model) factory_mock.assert_not_called() - def test_get_rerank_runner_for_reranking_model_creates_runner_with_model_instance(self): + def test_get_rerank_runner_for_reranking_model_creates_runner_with_model_instance(self, unbound_session: Session): processor = DataPostProcessor.__new__(DataPostProcessor) model_instance = object() expected_runner = object() - session = MagicMock() with patch.object(DataPostProcessor, "_get_rerank_model_instance", return_value=model_instance): with patch( @@ -156,22 +154,24 @@ class TestDataPostProcessor: "reranking_model_name": "model-y", }, weights=None, - session=session, + session=unbound_session, ) assert result is expected_runner factory_mock.assert_called_once_with( runner_type=RerankMode.RERANKING_MODEL, rerank_model_instance=model_instance, - session=session, + session=unbound_session, ) - def test_get_rerank_runner_returns_none_for_unsupported_mode(self): + def test_get_rerank_runner_returns_none_for_unsupported_mode(self, unbound_session: Session): processor = DataPostProcessor.__new__(DataPostProcessor) - session = MagicMock() - assert processor._get_rerank_runner("unsupported", "tenant-1", None, None, session=session) is None - assert processor._get_rerank_runner(RerankMode.WEIGHTED_SCORE, "tenant-1", None, None, session=session) is None + assert processor._get_rerank_runner("unsupported", "tenant-1", None, None, session=unbound_session) is None + assert ( + processor._get_rerank_runner(RerankMode.WEIGHTED_SCORE, "tenant-1", None, None, session=unbound_session) + is None + ) def test_get_reorder_runner_by_flag(self): processor = DataPostProcessor.__new__(DataPostProcessor) diff --git a/api/tests/unit_tests/core/rag/datasource/keyword/jieba/test_jieba.py b/api/tests/unit_tests/core/rag/datasource/keyword/jieba/test_jieba.py index 8a576bc40a2..250a776299d 100644 --- a/api/tests/unit_tests/core/rag/datasource/keyword/jieba/test_jieba.py +++ b/api/tests/unit_tests/core/rag/datasource/keyword/jieba/test_jieba.py @@ -4,10 +4,13 @@ from typing import Any from unittest.mock import MagicMock import pytest +from sqlalchemy import event, select +from sqlalchemy.orm import Session import core.rag.datasource.keyword.jieba.jieba as jieba_module from core.rag.datasource.keyword.jieba.jieba import Jieba, dumps_with_sets, set_orjson_default from core.rag.models.document import Document +from models.dataset import DatasetKeywordTable, DocumentSegment class _DummyLock: @@ -18,37 +21,6 @@ class _DummyLock: return False -class _Field: - def __init__(self, name: str): - self._name = name - - def __eq__(self, other): - return ("eq", self._name, other) - - def in_(self, values): - return ("in", self._name, tuple(values)) - - -class _FakeExecuteResult: - def __init__(self, segments: list[SimpleNamespace]): - self._segments = segments - - def scalars(self): - return self - - def all(self): - return self._segments - - -class _FakeSelect: - def __init__(self): - self.where_conditions: tuple | None = None - - def where(self, *conditions): - self.where_conditions = conditions - return self - - def _dataset_keyword_table(data_source_type: str = "database", keyword_table_dict: dict[str, Any] | None = None): return SimpleNamespace( data_source_type=data_source_type, @@ -67,8 +39,7 @@ def _dataset(dataset_keyword_table=None, keyword_number=None): @pytest.fixture -def patched_runtime(monkeypatch: pytest.MonkeyPatch): - session = MagicMock() +def patched_runtime(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): storage = MagicMock() lock = MagicMock(return_value=_DummyLock()) redis_client = SimpleNamespace(lock=lock) @@ -76,7 +47,28 @@ def patched_runtime(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(jieba_module, "storage", storage) monkeypatch.setattr(jieba_module, "redis_client", redis_client) - return SimpleNamespace(session=session, storage=storage, lock=lock) + return SimpleNamespace(session=sqlite_session, storage=storage, lock=lock) + + +def _segment(*, index_node_id: str = "node-2") -> DocumentSegment: + segment = DocumentSegment( + tenant_id="tenant-1", + dataset_id="dataset-1", + document_id="doc-2", + position=1, + content="segment-content", + word_count=1, + tokens=1, + created_by="user-1", + enabled=True, + keywords=[], + answer=None, + index_node_id=index_node_id, + index_node_hash="hash-2", + status="completed", + ) + segment.id = "segment-1" + return segment def test_create_indexes_documents_and_returns_self(monkeypatch: pytest.MonkeyPatch, patched_runtime): @@ -156,9 +148,11 @@ def test_add_texts_without_keywords_list_always_uses_extractor(monkeypatch: pyte assert keyword._update_segment_keywords.call_args.args[3] is patched_runtime.session -def test_text_exists_handles_missing_and_existing_keyword_table(monkeypatch: pytest.MonkeyPatch): +def test_text_exists_handles_missing_and_existing_keyword_table( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +): keyword = Jieba(_dataset(_dataset_keyword_table(keyword_table_dict=None))) - session = MagicMock() + session = unbound_session assert keyword.text_exists("node-1", session=session) is False keyword = Jieba( @@ -205,24 +199,9 @@ def test_delete_by_ids_saves_none_when_keyword_table_is_missing(monkeypatch: pyt def test_search_returns_documents_in_rank_order_and_applies_filter(monkeypatch: pytest.MonkeyPatch, patched_runtime): - class _FakeDocumentSegment: - dataset_id = _Field("dataset_id") - index_node_id = _Field("index_node_id") - document_id = _Field("document_id") - keyword = Jieba(_dataset(_dataset_keyword_table())) - patched_runtime.session.scalars.return_value.all.return_value = [ - SimpleNamespace( - index_node_id="node-2", - content="segment-content", - index_node_hash="hash-2", - document_id="doc-2", - dataset_id="dataset-1", - ) - ] - - monkeypatch.setattr(jieba_module, "DocumentSegment", _FakeDocumentSegment) - monkeypatch.setattr(jieba_module, "select", lambda *_: _FakeSelect()) + patched_runtime.session.add(_segment()) + patched_runtime.session.flush() monkeypatch.setattr(keyword, "_retrieve_ids_by_query", MagicMock(return_value=["node-1", "node-2"])) documents = keyword.search("query", session=patched_runtime.session, top_k=2, document_ids_filter=["doc-2"]) @@ -233,39 +212,47 @@ def test_search_returns_documents_in_rank_order_and_applies_filter(monkeypatch: assert documents[0].metadata["doc_hash"] == "hash-2" -def test_delete_removes_keyword_table_and_optional_file(monkeypatch: pytest.MonkeyPatch, patched_runtime): - db_keyword = _dataset_keyword_table(data_source_type="database") - file_keyword = _dataset_keyword_table(data_source_type="object_storage") +def test_delete_removes_keyword_table_and_optional_file(patched_runtime): + db_keyword = DatasetKeywordTable(dataset_id="dataset-1", keyword_table="", data_source_type="database") + patched_runtime.session.add(db_keyword) + patched_runtime.session.commit() + commits: list[str] = [] + event.listen(patched_runtime.session, "after_commit", lambda _session: commits.append("commit")) keyword_db = Jieba(_dataset(db_keyword)) keyword_db.delete(session=patched_runtime.session) patched_runtime.storage.delete.assert_not_called() + assert patched_runtime.session.get(DatasetKeywordTable, db_keyword.id) is None + file_keyword = DatasetKeywordTable(dataset_id="dataset-1", keyword_table="", data_source_type="object_storage") + patched_runtime.session.add(file_keyword) + patched_runtime.session.commit() keyword_file = Jieba(_dataset(file_keyword)) keyword_file.delete(session=patched_runtime.session) patched_runtime.storage.delete.assert_called_once_with("keyword_files/tenant-1/dataset-1.txt") - assert patched_runtime.session.delete.call_count == 2 - assert patched_runtime.session.commit.call_count == 2 + assert patched_runtime.session.get(DatasetKeywordTable, file_keyword.id) is None + assert commits == ["commit", "commit", "commit"] -def test_save_dataset_keyword_table_to_database(monkeypatch: pytest.MonkeyPatch, patched_runtime): - dataset_keyword_table = _dataset_keyword_table(data_source_type="database") +def test_save_dataset_keyword_table_to_database(patched_runtime): + dataset_keyword_table = DatasetKeywordTable(dataset_id="dataset-1", keyword_table="", data_source_type="database") + patched_runtime.session.add(dataset_keyword_table) + patched_runtime.session.flush() keyword = Jieba(_dataset(dataset_keyword_table)) - patched_runtime.session.scalar.return_value = dataset_keyword_table keyword._save_dataset_keyword_table({"kw": {"node-1"}}, patched_runtime.session) assert '"__type__":"keyword_table"' in dataset_keyword_table.keyword_table assert '"index_id":"dataset-1"' in dataset_keyword_table.keyword_table - patched_runtime.session.flush.assert_called_once() -def test_save_dataset_keyword_table_to_file_storage(monkeypatch: pytest.MonkeyPatch, patched_runtime): - dataset_keyword_table = _dataset_keyword_table(data_source_type="file") +def test_save_dataset_keyword_table_to_file_storage(patched_runtime): + dataset_keyword_table = DatasetKeywordTable(dataset_id="dataset-1", keyword_table="", data_source_type="file") + patched_runtime.session.add(dataset_keyword_table) + patched_runtime.session.flush() keyword = Jieba(_dataset(dataset_keyword_table)) patched_runtime.storage.exists.return_value = True - patched_runtime.session.scalar.return_value = dataset_keyword_table keyword._save_dataset_keyword_table({"kw": {"node-1"}}, patched_runtime.session) @@ -276,33 +263,38 @@ def test_save_dataset_keyword_table_to_file_storage(monkeypatch: pytest.MonkeyPa assert isinstance(save_args[1], bytes) -def test_get_dataset_keyword_table_returns_existing_table_data(monkeypatch: pytest.MonkeyPatch, patched_runtime): - existing = _dataset_keyword_table( - keyword_table_dict={"__type__": "keyword_table", "__data__": {"table": {"kw": ["node-1"]}}} +def test_get_dataset_keyword_table_returns_existing_table_data(patched_runtime): + existing = DatasetKeywordTable( + dataset_id="dataset-1", + keyword_table="", + data_source_type="database", ) + existing.get_keyword_table_dict = MagicMock( + return_value={"__type__": "keyword_table", "__data__": {"table": {"kw": ["node-1"]}}} + ) + patched_runtime.session.add(existing) + patched_runtime.session.flush() keyword = Jieba(_dataset(existing)) - patched_runtime.session.scalar.return_value = existing assert keyword._get_dataset_keyword_table(patched_runtime.session) == {"kw": ["node-1"]} - missing_payload = _dataset_keyword_table(keyword_table_dict=None) - keyword_with_missing_payload = Jieba(_dataset(missing_payload)) - patched_runtime.session.scalar.return_value = missing_payload + existing.get_keyword_table_dict = MagicMock(return_value=None) + keyword_with_missing_payload = Jieba(_dataset(existing)) assert keyword_with_missing_payload._get_dataset_keyword_table(patched_runtime.session) == {} def test_get_dataset_keyword_table_creates_table_when_missing(monkeypatch: pytest.MonkeyPatch, patched_runtime): keyword = Jieba(_dataset(dataset_keyword_table=None)) monkeypatch.setattr(jieba_module.dify_config, "KEYWORD_DATA_SOURCE_TYPE", "database") - patched_runtime.session.scalar.return_value = None - result = keyword._get_dataset_keyword_table(patched_runtime.session) assert result == {} - created_table = patched_runtime.session.add.call_args.args[0] + created_table = patched_runtime.session.scalar( + select(DatasetKeywordTable).where(DatasetKeywordTable.dataset_id == "dataset-1") + ) + assert created_table is not None assert created_table.dataset_id == "dataset-1" assert created_table.data_source_type == "database" assert '"index_id":"dataset-1"' in created_table.keyword_table - patched_runtime.session.flush.assert_called_once() def test_add_and_delete_ids_from_keyword_table_helpers(): @@ -334,31 +326,18 @@ def test_retrieve_ids_by_query_ranks_by_keyword_frequency(monkeypatch: pytest.Mo assert ranked_ids == ["node-2"] -def test_update_segment_keywords_updates_when_segment_exists(monkeypatch: pytest.MonkeyPatch, patched_runtime): - class _FakeDocumentSegment: - dataset_id = _Field("dataset_id") - index_node_id = _Field("index_node_id") - - monkeypatch.setattr(jieba_module, "DocumentSegment", _FakeDocumentSegment) - monkeypatch.setattr(jieba_module, "select", lambda *_: _FakeSelect()) - +def test_update_segment_keywords_updates_when_segment_exists(patched_runtime): keyword = Jieba(_dataset(_dataset_keyword_table())) - segment = SimpleNamespace(keywords=[]) - patched_runtime.session.scalar.return_value = segment + segment = _segment(index_node_id="node-1") + patched_runtime.session.add(segment) + patched_runtime.session.flush() keyword._update_segment_keywords("dataset-1", "node-1", ["kw1", "kw2"], patched_runtime.session) assert segment.keywords == ["kw1", "kw2"] - patched_runtime.session.add.assert_called_once_with(segment) - patched_runtime.session.flush.assert_called_once() - - patched_runtime.session.reset_mock() - patched_runtime.session.scalar.return_value = None keyword._update_segment_keywords("dataset-1", "node-missing", ["kw3"], patched_runtime.session) - - patched_runtime.session.add.assert_not_called() - patched_runtime.session.flush.assert_not_called() + assert segment.keywords == ["kw1", "kw2"] def test_create_segment_keywords_and_update_segment_keywords_index(monkeypatch: pytest.MonkeyPatch, patched_runtime): diff --git a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_base.py b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_base.py index 5e60c7a1979..28d56961143 100644 --- a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_base.py +++ b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_base.py @@ -1,8 +1,8 @@ from types import SimpleNamespace from typing import override -from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from core.rag.datasource.keyword.keyword_base import BaseKeyword from core.rag.models.document import Document @@ -64,9 +64,9 @@ class _KeywordForHelpers(BaseKeyword): return [] -def test_abstract_methods_raise_not_implemented(): +def test_abstract_methods_raise_not_implemented(unbound_session: Session): keyword = _KeywordThatRaises(SimpleNamespace(id="dataset-1")) - session = MagicMock() + session = unbound_session with pytest.raises(NotImplementedError): keyword.create([], session) @@ -87,7 +87,7 @@ def test_abstract_methods_raise_not_implemented(): keyword.search("query", session=session) -def test_filter_duplicate_texts_removes_existing_doc_ids(): +def test_filter_duplicate_texts_removes_existing_doc_ids(unbound_session: Session): keyword = _KeywordForHelpers(SimpleNamespace(id="dataset-1"), existing_ids={"duplicate"}) texts = [ Document(page_content="keep", metadata={"doc_id": "keep"}), @@ -95,7 +95,7 @@ def test_filter_duplicate_texts_removes_existing_doc_ids(): SimpleNamespace(page_content="without-metadata", metadata=None), ] - filtered = keyword._filter_duplicate_texts(texts, session=MagicMock()) + filtered = keyword._filter_duplicate_texts(texts, session=unbound_session) assert [text.metadata["doc_id"] for text in filtered if text.metadata] == ["keep"] assert any(text.metadata is None for text in filtered) diff --git a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py index 57df6bfa18a..1642f8eb552 100644 --- a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py +++ b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py @@ -4,6 +4,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from core.rag.datasource.keyword.keyword_factory import Keyword from core.rag.datasource.keyword.keyword_type import KeyWordType @@ -39,7 +40,7 @@ def test_keyword_initialization_uses_configured_factory(monkeypatch: pytest.Monk assert keyword._keyword_processor is fake_processor -def test_keyword_methods_forward_to_processor(): +def test_keyword_methods_forward_to_processor(unbound_session: Session): processor = MagicMock() processor.text_exists.return_value = True processor.search.return_value = [Document(page_content="matched", metadata={"doc_id": "doc-1"})] @@ -48,7 +49,7 @@ def test_keyword_methods_forward_to_processor(): keyword._keyword_processor = processor docs = [Document(page_content="doc", metadata={"doc_id": "doc-1"})] - session = MagicMock() + session = unbound_session keyword.create(docs, session, foo="bar") keyword.add_texts(docs, session, batch=True, keywords_list=[["kw"]]) assert keyword.text_exists("doc-1", session=session) is True diff --git a/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py b/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py index 6a89e6ec5c8..589d29d8240 100644 --- a/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py +++ b/api/tests/unit_tests/core/rag/datasource/test_datasource_retrieval.py @@ -158,11 +158,12 @@ class _SimpleRetrievalSegment: class TestRetrievalServiceInternals: @pytest.fixture def internal_dataset(self) -> Dataset: - dataset = Mock(spec=Dataset) - dataset.id = "dataset-id" - dataset.tenant_id = "tenant-id" - dataset.is_multimodal = False - dataset.doc_form = IndexStructureType.PARENT_CHILD_INDEX + dataset = Dataset( + id="dataset-id", + tenant_id="tenant-id", + is_multimodal=False, + chunk_structure=IndexStructureType.PARENT_CHILD_INDEX, + ) return dataset @pytest.fixture @@ -264,7 +265,7 @@ class TestRetrievalServiceInternals: @patch("core.rag.datasource.retrieval_service.Session") def test_get_dataset_queries_by_id(self, mock_session_class): - expected_dataset = Mock(spec=Dataset) + expected_dataset = Dataset() mock_session = Mock() mock_session.scalar.return_value = expected_dataset mock_session_class.return_value.__enter__.return_value = mock_session diff --git a/api/tests/unit_tests/core/rag/datasource/vdb/test_vector_factory.py b/api/tests/unit_tests/core/rag/datasource/vdb/test_vector_factory.py index 673832fa2d4..47b917c86d5 100644 --- a/api/tests/unit_tests/core/rag/datasource/vdb/test_vector_factory.py +++ b/api/tests/unit_tests/core/rag/datasource/vdb/test_vector_factory.py @@ -1,12 +1,19 @@ import base64 import sys import types +from datetime import UTC, datetime from types import SimpleNamespace from unittest.mock import MagicMock, patch +from uuid import uuid4 import pytest +from sqlalchemy.orm import Session from core.rag.models.document import Document +from extensions.storage.storage_type import StorageType +from models.dataset import Whitelist +from models.enums import CreatorUserRole +from models.model import UploadFile def _register_fake_factory_module(monkeypatch: pytest.MonkeyPatch, module_path: str, class_name: str): @@ -145,13 +152,12 @@ def test_get_vector_factory_entry_point_overrides_builtin(vector_factory_module, assert result_cls is _PluginChromaFactory -def test_vector_init_uses_default_and_custom_attributes(vector_factory_module): +def test_vector_init_uses_default_and_custom_attributes(vector_factory_module, unbound_session: Session): dataset = SimpleNamespace(id="dataset-1") - session = MagicMock() with patch.object(vector_factory_module.Vector, "_init_vector", return_value="processor") as init_vector: - default_vector = vector_factory_module.Vector(dataset, session=session) - custom_vector = vector_factory_module.Vector(dataset, attributes=["doc_id"], session=session) + default_vector = vector_factory_module.Vector(dataset, session=unbound_session) + custom_vector = vector_factory_module.Vector(dataset, attributes=["doc_id"], session=unbound_session) # `is_summary` and `original_chunk_id` must be in the default return-properties # projection so summary index retrieval works on backends that honor the list @@ -171,10 +177,10 @@ def test_vector_init_uses_default_and_custom_attributes(vector_factory_module): # trigger billing/feature-service calls during ``Vector(dataset, session=...)`` # construction. See ``_LazyEmbeddings``. assert isinstance(default_vector._embeddings, vector_factory_module._LazyEmbeddings) - assert default_vector._session is session - assert custom_vector._session is session + assert default_vector._session is unbound_session + assert custom_vector._session is unbound_session assert default_vector._vector_processor == "processor" - assert [call.kwargs["session"] for call in init_vector.call_args_list] == [session, session] + assert [call.kwargs["session"] for call in init_vector.call_args_list] == [unbound_session, unbound_session] def test_lazy_embeddings_defer_real_load_until_first_embed_call(vector_factory_module, monkeypatch: pytest.MonkeyPatch): @@ -220,7 +226,9 @@ def test_lazy_embeddings_defer_real_load_until_first_embed_call(vector_factory_m inner_model.embed_documents.assert_called_once_with(["world"]) -def test_init_vector_prefers_dataset_index_struct(vector_factory_module, monkeypatch: pytest.MonkeyPatch): +def test_init_vector_prefers_dataset_index_struct( + vector_factory_module, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +): calls = {"vector_type": None, "init_args": None} class _Factory: @@ -241,28 +249,25 @@ def test_init_vector_prefers_dataset_index_struct(vector_factory_module, monkeyp vector._attributes = ["doc_id"] vector._embeddings = "embeddings" - result = vector._init_vector(session=MagicMock()) + result = vector._init_vector(session=unbound_session) assert result == "vector-processor" assert calls["vector_type"] == vector_factory_module.VectorType.UPSTASH assert calls["init_args"] == (vector._dataset, ["doc_id"], "embeddings") -def test_init_vector_uses_whitelist_override(vector_factory_module, monkeypatch: pytest.MonkeyPatch): - class _Expr: - def __eq__(self, _other): - return "expr" - +def test_init_vector_uses_whitelist_override( + vector_factory_module, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): calls = {"vector_type": None} class _Factory: def init_vector(self, dataset, attributes, embeddings): return "vector-processor" - monkeypatch.setattr(vector_factory_module, "Whitelist", SimpleNamespace(tenant_id=_Expr(), category=_Expr())) - monkeypatch.setattr(vector_factory_module, "select", lambda _model: SimpleNamespace(where=lambda *_args: "stmt")) - session = MagicMock() - session.scalars.return_value.one_or_none.return_value = object() + tenant_id = str(uuid4()) + sqlite_session.add(Whitelist(tenant_id=tenant_id, category="vector_db")) + sqlite_session.commit() monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE", vector_factory_module.VectorType.CHROMA) monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE_WHITELIST_ENABLE", True) monkeypatch.setattr( @@ -272,18 +277,19 @@ def test_init_vector_uses_whitelist_override(vector_factory_module, monkeypatch: ) vector = vector_factory_module.Vector.__new__(vector_factory_module.Vector) - vector._dataset = SimpleNamespace(index_struct_dict=None, tenant_id="tenant-1") + vector._dataset = SimpleNamespace(index_struct_dict=None, tenant_id=tenant_id) vector._attributes = ["doc_id"] vector._embeddings = "embeddings" - result = vector._init_vector(session=session) + result = vector._init_vector(session=sqlite_session) assert result == "vector-processor" assert calls["vector_type"] == vector_factory_module.VectorType.TIDB_ON_QDRANT - session.scalars.assert_called_once_with("stmt") -def test_init_vector_raises_when_vector_store_missing(vector_factory_module, monkeypatch: pytest.MonkeyPatch): +def test_init_vector_raises_when_vector_store_missing( + vector_factory_module, monkeypatch: pytest.MonkeyPatch, unbound_session: Session +): monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE", None) monkeypatch.setattr(vector_factory_module.dify_config, "VECTOR_STORE_WHITELIST_ENABLE", False) @@ -293,7 +299,7 @@ def test_init_vector_raises_when_vector_store_missing(vector_factory_module, mon vector._embeddings = "embeddings" with pytest.raises(ValueError, match="Vector store must be specified"): - vector._init_vector(session=MagicMock()) + vector._init_vector(session=unbound_session) def test_create_batches_texts_and_skips_empty_input(vector_factory_module): @@ -347,36 +353,41 @@ def test_create_skips_empty_text_documents_before_embedding(vector_factory_modul vector._vector_processor.create.assert_not_called() -def test_create_multimodal_filters_missing_uploads(vector_factory_module, monkeypatch: pytest.MonkeyPatch): - class _Field: - def in_(self, value): - return value - - def __eq__(self, value): - return value - +def test_create_multimodal_filters_missing_uploads( + vector_factory_module, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + upload_file = UploadFile( + tenant_id=str(uuid4()), + storage_type=StorageType.LOCAL, + key="k-1", + name="image.png", + size=3, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + created_at=datetime.now(UTC), + used=True, + ) + sqlite_session.add(upload_file) + sqlite_session.commit() vector = vector_factory_module.Vector.__new__(vector_factory_module.Vector) vector._embeddings = MagicMock() vector._embeddings.embed_multimodal_documents.return_value = [[0.1, 0.2]] vector._vector_processor = MagicMock() - session = MagicMock() - vector._session = session - session.scalars.return_value = SimpleNamespace(all=lambda: [SimpleNamespace(id="f-1", key="k-1")]) - - monkeypatch.setattr(vector_factory_module, "UploadFile", SimpleNamespace(id=_Field())) - monkeypatch.setattr(vector_factory_module, "select", lambda _model: SimpleNamespace(where=lambda *_args: "stmt")) + vector._session = sqlite_session monkeypatch.setattr(vector_factory_module.storage, "load_once", MagicMock(return_value=b"abc")) docs = [ - Document(page_content="file-1", metadata={"doc_id": "f-1", "doc_type": "image"}), - Document(page_content="file-2", metadata={"doc_id": "f-2", "doc_type": "image"}), + Document(page_content="file-1", metadata={"doc_id": upload_file.id, "doc_type": "image"}), + Document(page_content="file-2", metadata={"doc_id": str(uuid4()), "doc_type": "image"}), ] vector.create_multimodal(file_documents=docs, request_id="r-1") file_base64 = base64.b64encode(b"abc").decode() vector._embeddings.embed_multimodal_documents.assert_called_once_with( - [{"content": file_base64, "content_type": "image", "file_id": "f-1"}] + [{"content": file_base64, "content_type": "image", "file_id": upload_file.id}] ) vector._vector_processor.create.assert_called_once_with( texts=[docs[0]], @@ -482,30 +493,43 @@ def test_vector_delegation_methods(vector_factory_module): vector._vector_processor.delete_by_metadata_field.assert_called_once_with("doc_id", "doc-1") -def test_search_by_file_handles_missing_and_existing_upload(vector_factory_module, monkeypatch: pytest.MonkeyPatch): +def test_search_by_file_handles_missing_and_existing_upload( + vector_factory_module, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): vector = vector_factory_module.Vector.__new__(vector_factory_module.Vector) vector._embeddings = MagicMock() vector._vector_processor = MagicMock() - session = MagicMock() - session.get.return_value = None - vector._session = session + vector._session = sqlite_session + missing_id = str(uuid4()) - assert vector.search_by_file("file-1") == [] - session.get.assert_called_once_with(vector_factory_module.UploadFile, "file-1") + assert vector.search_by_file(missing_id) == [] - session.get.return_value = SimpleNamespace(key="blob-key") + upload_file = UploadFile( + tenant_id=str(uuid4()), + storage_type=StorageType.LOCAL, + key="blob-key", + name="query.png", + size=10, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + created_at=datetime.now(UTC), + used=True, + ) + sqlite_session.add(upload_file) + sqlite_session.commit() monkeypatch.setattr(vector_factory_module.storage, "load_once", MagicMock(return_value=b"file-bytes")) vector._embeddings.embed_multimodal_query.return_value = [0.3, 0.4] vector._vector_processor.search_by_vector.return_value = ["hit"] - result = vector.search_by_file("file-2", top_k=2) + result = vector.search_by_file(upload_file.id, top_k=2) assert result == ["hit"] - session.get.assert_called_with(vector_factory_module.UploadFile, "file-2") payload = vector._embeddings.embed_multimodal_query.call_args.args[0] assert payload["content_type"] == vector_factory_module.DocType.IMAGE - assert payload["file_id"] == "file-2" + assert payload["file_id"] == upload_file.id def test_delete_clears_redis_cache_when_collection_exists(vector_factory_module, monkeypatch: pytest.MonkeyPatch): diff --git a/api/tests/unit_tests/core/rag/docstore/test_dataset_docstore.py b/api/tests/unit_tests/core/rag/docstore/test_dataset_docstore.py index b7c51c77427..acc1af08438 100644 --- a/api/tests/unit_tests/core/rag/docstore/test_dataset_docstore.py +++ b/api/tests/unit_tests/core/rag/docstore/test_dataset_docstore.py @@ -5,7 +5,7 @@ Tests cover all public methods and error paths of the DatasetDocumentStore class which provides document storage and retrieval functionality for datasets in the RAG system. """ -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock import pytest from sqlalchemy import func, select @@ -19,12 +19,14 @@ TENANT_ID = "00000000-0000-0000-0000-000000000001" DATASET_ID = "00000000-0000-0000-0000-000000000002" DOCUMENT_ID = "00000000-0000-0000-0000-000000000003" USER_ID = "00000000-0000-0000-0000-000000000004" +ATTACHMENT_ID = "00000000-0000-0000-0000-000000000005" def _dataset() -> Dataset: - dataset = MagicMock(spec=Dataset) - dataset.id = DATASET_ID - dataset.tenant_id = TENANT_ID + dataset = Dataset( + id=DATASET_ID, + tenant_id=TENANT_ID, + ) return dataset @@ -36,7 +38,25 @@ def _persist_segment( content: str = "Test content", tokens: int = 5, ) -> DocumentSegment: - segment = DocumentSegment( + segment = _segment( + index_node_id=index_node_id, + index_node_hash=index_node_hash, + content=content, + tokens=tokens, + ) + session.add(segment) + session.flush() + return segment + + +def _segment( + *, + index_node_id: str = "doc-1", + index_node_hash: str = "hash-1", + content: str = "Test content", + tokens: int = 5, +) -> DocumentSegment: + return DocumentSegment( tenant_id=TENANT_ID, dataset_id=DATASET_ID, document_id=DOCUMENT_ID, @@ -48,19 +68,25 @@ def _persist_segment( index_node_id=index_node_id, index_node_hash=index_node_hash, ) - session.add(segment) - session.flush() - return segment -class TestDatasetDocumentStoreInit: +class _UsesSQLiteSession: + session: Session + + @pytest.fixture(autouse=True) + def _inject_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + + +class TestDatasetDocumentStoreInit(_UsesSQLiteSession): """Tests for DatasetDocumentStore initialization.""" def test_init_with_all_parameters(self): """Test initialization with dataset, user_id, and document_id.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + mock_dataset = Dataset( + id="test-dataset-id", + ) store = DatasetDocumentStore( dataset=mock_dataset, @@ -77,8 +103,9 @@ class TestDatasetDocumentStoreInit: def test_init_without_document_id(self): """Test initialization without document_id.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + mock_dataset = Dataset( + id="test-dataset-id", + ) store = DatasetDocumentStore( dataset=mock_dataset, @@ -89,14 +116,15 @@ class TestDatasetDocumentStoreInit: assert store.dataset_id == "test-dataset-id" -class TestDatasetDocumentStoreSerialization: +class TestDatasetDocumentStoreSerialization(_UsesSQLiteSession): """Tests for to_dict and from_dict methods.""" def test_to_dict(self): """Test serialization to dictionary.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + mock_dataset = Dataset( + id="test-dataset-id", + ) store = DatasetDocumentStore( dataset=mock_dataset, @@ -123,27 +151,20 @@ class TestDatasetDocumentStoreSerialization: assert store._document_id == "test-doc" -class TestDatasetDocumentStoreDocs: +class TestDatasetDocumentStoreDocs(_UsesSQLiteSession): """Tests for the docs property.""" def test_docs_returns_document_dict(self): """Test that docs property returns a dictionary of documents.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + mock_segment = _segment(index_node_id="node-1") - mock_segment = MagicMock(spec=DocumentSegment) - mock_segment.index_node_id = "node-1" - mock_segment.index_node_hash = "hash-1" - mock_segment.document_id = "doc-1" - mock_segment.dataset_id = "test-dataset-id" - mock_segment.content = "Test content" - - mock_session = MagicMock() - mock_session.scalars.return_value.all.return_value = [mock_segment] + mock_session = self.session + mock_session.add(mock_segment) + mock_session.flush() store = DatasetDocumentStore( - dataset=mock_dataset, + dataset=_dataset(), user_id="test-user-id", ) @@ -155,11 +176,11 @@ class TestDatasetDocumentStoreDocs: def test_docs_empty_dataset(self): """Test docs property with no segments.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + mock_dataset = Dataset( + id="test-dataset-id", + ) - mock_session = MagicMock() - mock_session.scalars.return_value.all.return_value = [] + mock_session = self.session store = DatasetDocumentStore( dataset=mock_dataset, @@ -171,12 +192,7 @@ class TestDatasetDocumentStoreDocs: assert result == {} -@pytest.mark.parametrize( - "sqlite_session", - [(DocumentSegment, ChildChunk, SegmentAttachmentBinding)], - indirect=True, -) -class TestDatasetDocumentStoreAddDocuments: +class TestDatasetDocumentStoreAddDocuments(_UsesSQLiteSession): """Tests for add_documents method.""" def test_add_documents_new_document_with_token_count(self, sqlite_session: Session): @@ -320,260 +336,147 @@ class TestDatasetDocumentStoreAddDocuments: assert sqlite_session.scalar(select(func.count()).select_from(DocumentSegment)) == 0 -class TestDatasetDocumentStoreExists: +class TestDatasetDocumentStoreExists(_UsesSQLiteSession): """Tests for document_exists method.""" def test_document_exists_returns_true(self): """Test document_exists returns True when segment exists.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + _persist_segment(self.session) + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - mock_segment = MagicMock() - mock_session = MagicMock() - - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) - - result = store.document_exists("doc-1", session=mock_session) - - assert result is True + assert store.document_exists("doc-1", session=self.session) is True def test_document_exists_returns_false(self): """Test document_exists returns False when segment doesn't exist.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_session = MagicMock() + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) - - result = store.document_exists("doc-1", session=mock_session) - - assert result is False + assert store.document_exists("doc-1", session=self.session) is False -class TestDatasetDocumentStoreGetDocument: +class TestDatasetDocumentStoreGetDocument(_UsesSQLiteSession): """Tests for get_document method.""" def test_get_document_success(self): """Test getting a document successfully.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + _persist_segment(self.session, index_node_id="node-1") + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - mock_segment = MagicMock(spec=DocumentSegment) - mock_segment.index_node_id = "node-1" - mock_segment.index_node_hash = "hash-1" - mock_segment.document_id = "doc-1" - mock_segment.dataset_id = "test-dataset-id" - mock_segment.content = "Test content" - mock_session = MagicMock() + result = store.get_document("node-1", session=self.session, raise_error=False) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) - - result = store.get_document("node-1", session=mock_session, raise_error=False) - - assert isinstance(result, Document) - assert result.page_content == "Test content" + assert isinstance(result, Document) + assert result.page_content == "Test content" def test_get_document_returns_none_when_not_found(self): """Test get_document returns None when not found and raise_error=False.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_session = MagicMock() + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + result = store.get_document("nonexistent", session=self.session, raise_error=False) - result = store.get_document("nonexistent", session=mock_session, raise_error=False) - - assert result is None + assert result is None def test_get_document_raises_when_not_found(self): """Test get_document raises ValueError when not found and raise_error=True.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_session = MagicMock() + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) - - with pytest.raises(ValueError, match="not found"): - store.get_document("nonexistent", session=mock_session, raise_error=True) + with pytest.raises(ValueError, match="not found"): + store.get_document("nonexistent", session=self.session, raise_error=True) -class TestDatasetDocumentStoreDeleteDocument: +class TestDatasetDocumentStoreDeleteDocument(_UsesSQLiteSession): """Tests for delete_document method.""" def test_delete_document_success(self): """Test deleting a document successfully.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + segment = _persist_segment(self.session) + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - mock_segment = MagicMock() - mock_session = MagicMock() + store.delete_document("doc-1", session=self.session) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) - - store.delete_document("doc-1", session=mock_session) - - mock_session.delete.assert_called_with(mock_segment) - mock_session.flush.assert_called() + assert self.session.get(DocumentSegment, segment.id) is None def test_delete_document_returns_none_when_not_found(self): """Test delete_document returns None when not found and raise_error=False.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_session = MagicMock() + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + result = store.delete_document("nonexistent", session=self.session, raise_error=False) - result = store.delete_document("nonexistent", session=mock_session, raise_error=False) - - assert result is None + assert result is None def test_delete_document_raises_when_not_found(self): """Test delete_document raises ValueError when not found and raise_error=True.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_session = MagicMock() + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) - - with pytest.raises(ValueError, match="not found"): - store.delete_document("nonexistent", session=mock_session, raise_error=True) + with pytest.raises(ValueError, match="not found"): + store.delete_document("nonexistent", session=self.session, raise_error=True) -class TestDatasetDocumentStoreHashOperations: +class TestDatasetDocumentStoreHashOperations(_UsesSQLiteSession): """Tests for set_document_hash and get_document_hash methods.""" def test_set_document_hash_success(self): """Test setting document hash successfully.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + segment = _persist_segment(self.session, index_node_hash="old-hash") + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - mock_segment = MagicMock() - mock_segment.index_node_hash = "old-hash" - mock_session = MagicMock() + store.set_document_hash("doc-1", "new-hash", session=self.session) + self.session.expire_all() - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) - - store.set_document_hash("doc-1", "new-hash", session=mock_session) - - assert mock_segment.index_node_hash == "new-hash" - mock_session.flush.assert_called() + updated_segment = self.session.get(DocumentSegment, segment.id) + assert updated_segment is not None + assert updated_segment.index_node_hash == "new-hash" def test_set_document_hash_returns_none_when_not_found(self): """Test set_document_hash returns None when segment not found.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_session = MagicMock() + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + result = store.set_document_hash("nonexistent", "new-hash", session=self.session) - result = store.set_document_hash("nonexistent", "new-hash", session=mock_session) - - assert result is None + assert result is None def test_get_document_hash_success(self): """Test getting document hash successfully.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + _persist_segment(self.session, index_node_hash="test-hash") + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - mock_segment = MagicMock() - mock_segment.index_node_hash = "test-hash" - mock_session = MagicMock() + result = store.get_document_hash("doc-1", session=self.session) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=mock_segment): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) - - result = store.get_document_hash("doc-1", session=mock_session) - - assert result == "test-hash" + assert result == "test-hash" def test_get_document_hash_returns_none_when_not_found(self): """Test get_document_hash returns None when segment not found.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_session = MagicMock() + store = DatasetDocumentStore(dataset=_dataset(), user_id=USER_ID) - with patch.object(DatasetDocumentStore, "get_document_segment", return_value=None): - store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - ) + result = store.get_document_hash("nonexistent", session=self.session) - result = store.get_document_hash("nonexistent", session=mock_session) - - assert result is None + assert result is None -class TestDatasetDocumentStoreSegment: +class TestDatasetDocumentStoreSegment(_UsesSQLiteSession): """Tests for get_document_segment method.""" def test_get_document_segment_returns_segment(self): """Test getting a document segment.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" + mock_segment = _segment() - mock_segment = MagicMock(spec=DocumentSegment) - - mock_session = MagicMock() - mock_session.scalar.return_value = mock_segment + mock_session = self.session + mock_session.add(mock_segment) + mock_session.flush() store = DatasetDocumentStore( - dataset=mock_dataset, + dataset=_dataset(), user_id="test-user-id", ) @@ -584,14 +487,10 @@ class TestDatasetDocumentStoreSegment: def test_get_document_segment_returns_none(self): """Test getting a non-existent document segment.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - - mock_session = MagicMock() - mock_session.scalar.return_value = None + mock_session = self.session store = DatasetDocumentStore( - dataset=mock_dataset, + dataset=_dataset(), user_id="test-user-id", ) @@ -600,98 +499,79 @@ class TestDatasetDocumentStoreSegment: assert result is None -class TestDatasetDocumentStoreMultimodelBinding: +class TestDatasetDocumentStoreMultimodelBinding(_UsesSQLiteSession): """Tests for add_multimodel_documents_binding method.""" def test_add_multimodel_documents_binding_with_attachments(self): """Test adding multimodel document bindings.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_dataset.tenant_id = "tenant-1" + attachment = AttachmentDocument(page_content="attachment", metadata={"doc_id": ATTACHMENT_ID}) - mock_attachment = MagicMock(spec=AttachmentDocument) - mock_attachment.metadata = {"doc_id": "attachment-1"} - - mock_session = MagicMock() + mock_session = self.session store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", + dataset=_dataset(), + user_id=USER_ID, + document_id=DOCUMENT_ID, ) - store.add_multimodel_documents_binding("seg-1", [mock_attachment], session=mock_session) + store.add_multimodel_documents_binding("seg-1", [attachment], session=mock_session) + mock_session.flush() - mock_session.add.assert_called() + binding = mock_session.scalar(select(SegmentAttachmentBinding)) + assert binding is not None + assert binding.segment_id == "seg-1" + assert binding.attachment_id == ATTACHMENT_ID def test_add_multimodel_documents_binding_without_attachments(self): """Test adding bindings with None attachments.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_dataset.tenant_id = "tenant-1" - - mock_session = MagicMock() + mock_session = self.session store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", + dataset=_dataset(), + user_id=USER_ID, + document_id=DOCUMENT_ID, ) store.add_multimodel_documents_binding("seg-1", None, session=mock_session) - mock_session.add.assert_not_called() + assert mock_session.scalar(select(func.count()).select_from(SegmentAttachmentBinding)) == 0 def test_add_multimodel_documents_binding_with_empty_list(self): """Test adding bindings with empty list.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_dataset.tenant_id = "tenant-1" - - mock_session = MagicMock() + mock_session = self.session store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", - document_id="test-doc-id", + dataset=_dataset(), + user_id=USER_ID, + document_id=DOCUMENT_ID, ) store.add_multimodel_documents_binding("seg-1", [], session=mock_session) - mock_session.add.assert_not_called() + assert mock_session.scalar(select(func.count()).select_from(SegmentAttachmentBinding)) == 0 def test_add_multimodel_documents_binding_with_none_document_id(self): """Test that no bindings are added when document_id is None.""" - mock_dataset = MagicMock(spec=Dataset) - mock_dataset.id = "test-dataset-id" - mock_dataset.tenant_id = "tenant-1" + attachment = AttachmentDocument(page_content="attachment", metadata={"doc_id": ATTACHMENT_ID}) - mock_attachment = MagicMock(spec=AttachmentDocument) - mock_attachment.metadata = {"doc_id": "attachment-1"} - - mock_session = MagicMock() + mock_session = self.session store = DatasetDocumentStore( - dataset=mock_dataset, - user_id="test-user-id", + dataset=_dataset(), + user_id=USER_ID, document_id=None, ) - store.add_multimodel_documents_binding("seg-1", [mock_attachment], session=mock_session) + store.add_multimodel_documents_binding("seg-1", [attachment], session=mock_session) - mock_session.add.assert_not_called() + assert mock_session.scalar(select(func.count()).select_from(SegmentAttachmentBinding)) == 0 -@pytest.mark.parametrize( - "sqlite_session", - [(DocumentSegment, ChildChunk, SegmentAttachmentBinding)], - indirect=True, -) -class TestDatasetDocumentStoreAddDocumentsUpdateChild: +class TestDatasetDocumentStoreAddDocumentsUpdateChild(_UsesSQLiteSession): """Tests for add_documents when updating existing documents with children.""" def test_add_documents_update_existing_with_children(self, sqlite_session: Session): @@ -739,12 +619,7 @@ class TestDatasetDocumentStoreAddDocumentsUpdateChild: assert children[0].index_node_id == "child-1" -@pytest.mark.parametrize( - "sqlite_session", - [(DocumentSegment, ChildChunk, SegmentAttachmentBinding)], - indirect=True, -) -class TestDatasetDocumentStoreAddDocumentsUpdateAnswer: +class TestDatasetDocumentStoreAddDocumentsUpdateAnswer(_UsesSQLiteSession): """Tests for add_documents when updating existing documents with answer metadata.""" def test_add_documents_update_existing_with_answer(self, sqlite_session: Session): diff --git a/api/tests/unit_tests/core/rag/embedding/test_cached_embedding.py b/api/tests/unit_tests/core/rag/embedding/test_cached_embedding.py index b33a7ba725c..97b1327bc88 100644 --- a/api/tests/unit_tests/core/rag/embedding/test_cached_embedding.py +++ b/api/tests/unit_tests/core/rag/embedding/test_cached_embedding.py @@ -8,19 +8,55 @@ This test file covers the methods not fully tested in test_embedding_service.py: import base64 import logging +from dataclasses import dataclass from decimal import Decimal from unittest.mock import Mock, patch import numpy as np import pytest -from sqlalchemy.exc import IntegrityError +from sqlalchemy import event, func, select +from sqlalchemy.orm import Session +from core.rag.embedding import cached_embedding as cached_embedding_module from core.rag.embedding.cached_embedding import CacheEmbedding from graphon.model_runtime.entities.model_entities import ModelPropertyKey from graphon.model_runtime.entities.text_embedding_entities import EmbeddingResult, EmbeddingUsage from models.dataset import Embedding +@dataclass(frozen=True) +class _DatabaseBinding: + session: Session + + +@pytest.fixture +def embedding_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Bind CacheEmbedding to the shared SQLite session.""" + + monkeypatch.setattr(cached_embedding_module, "db", _DatabaseBinding(session=sqlite_session)) + return sqlite_session + + +def _persist_embedding( + session: Session, + *, + cache_key: str, + vector: list[float], + model_name: str = "vision-embedding-model", + provider_name: str = "openai", +) -> Embedding: + embedding = Embedding( + model_name=model_name, + hash=cache_key, + provider_name=provider_name, + embedding=b"placeholder", + ) + embedding.set_embedding(vector) + session.add(embedding) + session.commit() + return embedding + + class TestCacheEmbeddingMultimodalDocuments: """Test suite for CacheEmbedding.embed_multimodal_documents method.""" @@ -65,27 +101,29 @@ class TestCacheEmbeddingMultimodalDocuments: ) def test_embed_single_multimodal_document_cache_miss( - self, mock_model_instance, sample_multimodal_result: EmbeddingResult + self, + mock_model_instance, + sample_multimodal_result: EmbeddingResult, + embedding_session: Session, ): """Test embedding a single multimodal document when cache is empty.""" cache_embedding = CacheEmbedding(mock_model_instance) documents = [{"file_id": "file123", "content": "test content"}] - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_multimodal_embedding.return_value = sample_multimodal_result + mock_model_instance.invoke_multimodal_embedding.return_value = sample_multimodal_result - result = cache_embedding.embed_multimodal_documents(documents) + result = cache_embedding.embed_multimodal_documents(documents) - assert len(result) == 1 - assert isinstance(result[0], list) - assert len(result[0]) == 1536 + assert len(result) == 1 + assert isinstance(result[0], list) + assert len(result[0]) == 1536 - mock_model_instance.invoke_multimodal_embedding.assert_called_once() - mock_session.add.assert_called_once() - mock_session.commit.assert_called_once() + mock_model_instance.invoke_multimodal_embedding.assert_called_once() + persisted = embedding_session.scalar(select(Embedding).where(Embedding.hash == "file123")) + assert persisted is not None + assert persisted.get_embedding() == result[0] - def test_embed_multiple_multimodal_documents_cache_miss(self, mock_model_instance): + def test_embed_multiple_multimodal_documents_cache_miss(self, mock_model_instance, embedding_session: Session): """Test embedding multiple multimodal documents when cache is empty.""" cache_embedding = CacheEmbedding(mock_model_instance) documents = [ @@ -116,16 +154,15 @@ class TestCacheEmbeddingMultimodalDocuments: usage=usage, ) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_multimodal_embedding.return_value = embedding_result + mock_model_instance.invoke_multimodal_embedding.return_value = embedding_result - result = cache_embedding.embed_multimodal_documents(documents) + result = cache_embedding.embed_multimodal_documents(documents) - assert len(result) == 3 - assert all(len(emb) == 1536 for emb in result) + assert len(result) == 3 + assert all(len(emb) == 1536 for emb in result) + assert embedding_session.scalar(select(func.count()).select_from(Embedding)) == 3 - def test_embed_multimodal_documents_cache_hit(self, mock_model_instance): + def test_embed_multimodal_documents_cache_hit(self, mock_model_instance, embedding_session: Session): """Test embedding multimodal documents when embeddings are cached.""" cache_embedding = CacheEmbedding(mock_model_instance) documents = [{"file_id": "file123"}] @@ -133,19 +170,15 @@ class TestCacheEmbeddingMultimodalDocuments: cached_vector = np.random.randn(1536) normalized_cached = (cached_vector / np.linalg.norm(cached_vector)).tolist() - mock_cached_embedding = Mock(spec=Embedding) - mock_cached_embedding.get_embedding.return_value = normalized_cached + _persist_embedding(embedding_session, cache_key="file123", vector=normalized_cached) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = mock_cached_embedding + result = cache_embedding.embed_multimodal_documents(documents) - result = cache_embedding.embed_multimodal_documents(documents) + assert len(result) == 1 + assert result[0] == normalized_cached + mock_model_instance.invoke_multimodal_embedding.assert_not_called() - assert len(result) == 1 - assert result[0] == normalized_cached - mock_model_instance.invoke_multimodal_embedding.assert_not_called() - - def test_embed_multimodal_documents_partial_cache_hit(self, mock_model_instance): + def test_embed_multimodal_documents_partial_cache_hit(self, mock_model_instance, embedding_session: Session): """Test embedding multimodal documents with mixed cache hits and misses.""" cache_embedding = CacheEmbedding(mock_model_instance) documents = [ @@ -157,8 +190,7 @@ class TestCacheEmbeddingMultimodalDocuments: cached_vector = np.random.randn(1536) normalized_cached = (cached_vector / np.linalg.norm(cached_vector)).tolist() - mock_cached_embedding = Mock(spec=Embedding) - mock_cached_embedding.get_embedding.return_value = normalized_cached + _persist_embedding(embedding_session, cache_key="cached_file", vector=normalized_cached) new_embeddings = [] for _ in range(2): @@ -182,16 +214,20 @@ class TestCacheEmbeddingMultimodalDocuments: usage=usage, ) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.side_effect = [mock_cached_embedding, None, None] - mock_model_instance.invoke_multimodal_embedding.return_value = embedding_result + mock_model_instance.invoke_multimodal_embedding.return_value = embedding_result - result = cache_embedding.embed_multimodal_documents(documents) + result = cache_embedding.embed_multimodal_documents(documents) - assert len(result) == 3 - assert result[0] == normalized_cached + assert len(result) == 3 + assert result[0] == normalized_cached + assert embedding_session.scalar(select(func.count()).select_from(Embedding)) == 3 - def test_embed_multimodal_documents_nan_handling(self, mock_model_instance, caplog: pytest.LogCaptureFixture): + def test_embed_multimodal_documents_nan_handling( + self, + mock_model_instance, + embedding_session: Session, + caplog: pytest.LogCaptureFixture, + ): """Test handling of NaN values in multimodal embeddings.""" cache_embedding = CacheEmbedding(mock_model_instance) documents = [{"file_id": "valid"}, {"file_id": "nan"}] @@ -215,20 +251,19 @@ class TestCacheEmbeddingMultimodalDocuments: usage=usage, ) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_multimodal_embedding.return_value = embedding_result + mock_model_instance.invoke_multimodal_embedding.return_value = embedding_result - with caplog.at_level(logging.WARNING, logger="core.rag.embedding.cached_embedding"): - result = cache_embedding.embed_multimodal_documents(documents) + with caplog.at_level(logging.WARNING, logger="core.rag.embedding.cached_embedding"): + result = cache_embedding.embed_multimodal_documents(documents) - assert len(result) == 2 - assert result[0] is not None - assert result[1] is None + assert len(result) == 2 + assert result[0] is not None + assert result[1] is None + assert embedding_session.scalar(select(func.count()).select_from(Embedding)) == 1 - assert any(record.levelno == logging.WARNING for record in caplog.records) + assert any(record.levelno == logging.WARNING for record in caplog.records) - def test_embed_multimodal_documents_large_batch(self, mock_model_instance): + def test_embed_multimodal_documents_large_batch(self, mock_model_instance, embedding_session: Session): """Test embedding large batch of multimodal documents respecting MAX_CHUNKS.""" cache_embedding = CacheEmbedding(mock_model_instance) documents = [{"file_id": f"file{i}"} for i in range(25)] @@ -256,49 +291,66 @@ class TestCacheEmbeddingMultimodalDocuments: usage=usage, ) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None + batch_results = [create_batch_result(10), create_batch_result(10), create_batch_result(5)] + mock_model_instance.invoke_multimodal_embedding.side_effect = batch_results - batch_results = [create_batch_result(10), create_batch_result(10), create_batch_result(5)] - mock_model_instance.invoke_multimodal_embedding.side_effect = batch_results + result = cache_embedding.embed_multimodal_documents(documents) - result = cache_embedding.embed_multimodal_documents(documents) + assert len(result) == 25 + assert mock_model_instance.invoke_multimodal_embedding.call_count == 3 + assert embedding_session.scalar(select(func.count()).select_from(Embedding)) == 25 - assert len(result) == 25 - assert mock_model_instance.invoke_multimodal_embedding.call_count == 3 - - def test_embed_multimodal_documents_api_error(self, mock_model_instance): + def test_embed_multimodal_documents_api_error(self, mock_model_instance, embedding_session: Session): """Test handling of API errors during multimodal embedding.""" cache_embedding = CacheEmbedding(mock_model_instance) documents = [{"file_id": "file123"}] - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_multimodal_embedding.side_effect = Exception("API Error") + mock_model_instance.invoke_multimodal_embedding.side_effect = Exception("API Error") - with pytest.raises(Exception) as exc_info: - cache_embedding.embed_multimodal_documents(documents) + with pytest.raises(Exception, match="API Error"): + cache_embedding.embed_multimodal_documents(documents) - assert "API Error" in str(exc_info.value) - mock_session.rollback.assert_called() + assert not embedding_session.in_transaction() + assert embedding_session.scalar(select(func.count()).select_from(Embedding)) == 0 def test_embed_multimodal_documents_integrity_error_during_transform( - self, mock_model_instance, sample_multimodal_result + self, + mock_model_instance, + sample_multimodal_result, + embedding_session: Session, ): """Test handling of IntegrityError during embedding transformation.""" cache_embedding = CacheEmbedding(mock_model_instance) documents = [{"file_id": "file123"}] - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_multimodal_embedding.return_value = sample_multimodal_result + mock_model_instance.invoke_multimodal_embedding.return_value = sample_multimodal_result - mock_session.commit.side_effect = IntegrityError("Duplicate key", None, None) + injected = False + def add_competing_row(session: Session, _flush_context: object, _instances: object) -> None: + nonlocal injected + if injected: + return + pending = next(item for item in session.new if isinstance(item, Embedding)) + session.add( + Embedding( + model_name=pending.model_name, + hash=pending.hash, + provider_name=pending.provider_name, + embedding=pending.embedding, + ) + ) + injected = True + + event.listen(embedding_session, "before_flush", add_competing_row) + try: result = cache_embedding.embed_multimodal_documents(documents) + finally: + event.remove(embedding_session, "before_flush", add_competing_row) - assert len(result) == 1 - mock_session.rollback.assert_called() + assert len(result) == 1 + assert not embedding_session.in_transaction() + assert embedding_session.scalar(select(func.count()).select_from(Embedding)) == 0 class TestCacheEmbeddingMultimodalQuery: diff --git a/api/tests/unit_tests/core/rag/embedding/test_embedding_service.py b/api/tests/unit_tests/core/rag/embedding/test_embedding_service.py index 168d0e466d4..2c395ea9807 100644 --- a/api/tests/unit_tests/core/rag/embedding/test_embedding_service.py +++ b/api/tests/unit_tests/core/rag/embedding/test_embedding_service.py @@ -1,58 +1,24 @@ -"""Comprehensive unit tests for embedding service (CacheEmbedding). +"""Behavior-focused tests for :class:`CacheEmbedding`. -This test module covers all aspects of the embedding service including: -- Batch embedding generation with proper batching logic -- Embedding model switching and configuration -- Embedding dimension validation -- Error handling for API failures -- Cache management (database and Redis) -- Normalization and NaN handling - -Test Coverage: -============== -1. **Batch Embedding Generation** - - Single text embedding - - Multiple texts in batches - - Large batch processing (respects MAX_CHUNKS) - - Empty text handling - -2. **Embedding Model Switching** - - Different providers (OpenAI, Cohere, etc.) - - Different models within same provider - - Model instance configuration - -3. **Embedding Dimension Validation** - - Correct dimensions for different models - - Vector normalization - - Dimension consistency across batches - -4. **Error Handling** - - API connection failures - - Rate limit errors - - Authorization errors - - Invalid input handling - - NaN value detection and handling - -5. **Cache Management** - - Database cache for document embeddings - - Redis cache for query embeddings - - Cache hit/miss scenarios - - Cache invalidation - -All tests use mocking to avoid external dependencies and ensure fast, reliable execution. -Tests follow the Arrange-Act-Assert pattern for clarity. +Document and multimodal caches are persisted in SQLite so cache hits, misses, +provider isolation, uniqueness races, and rollback behavior exercise real +SQLAlchemy statements. Model providers and Redis remain mocked external I/O. """ import base64 -import logging +import pickle +from dataclasses import dataclass from decimal import Decimal from unittest.mock import Mock, patch import numpy as np import pytest +from sqlalchemy import event, select from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session from core.entities.embedding_type import EmbeddingInputType +from core.rag.embedding import cached_embedding from core.rag.embedding.cached_embedding import CacheEmbedding from graphon.model_runtime.entities.model_entities import ModelPropertyKey from graphon.model_runtime.entities.text_embedding_entities import EmbeddingResult, EmbeddingUsage @@ -61,1833 +27,281 @@ from graphon.model_runtime.errors.invoke import ( InvokeConnectionError, InvokeRateLimitError, ) +from libs import helper from models.dataset import Embedding -class TestCacheEmbeddingDocuments: - """Test suite for CacheEmbedding.embed_documents method. +@dataclass(frozen=True) +class _Database: + """Expose a real session through the interface used by the cache.""" - This class tests the batch embedding generation functionality including: - - Single and multiple text processing - - Cache hit/miss scenarios - - Batch processing with MAX_CHUNKS - - Database cache management - - Error handling during embedding generation - """ + session: Session - @pytest.fixture - def mock_model_instance(self): - """Create a mock ModelInstance for testing. - Returns: - Mock: Configured ModelInstance with text embedding capabilities - """ - model_instance = Mock() - model_instance.model_name = "text-embedding-ada-002" - model_instance.provider = "openai" - model_instance.credentials = {"api_key": "test-key"} +@pytest.fixture +def sqlite_embedding_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Bind the shared SQLite session to the document-cache implementation.""" - # Mock the model type instance - model_type_instance = Mock() - model_instance.model_type_instance = model_type_instance + monkeypatch.setattr(cached_embedding, "db", _Database(sqlite_session)) + return sqlite_session - # Mock model schema with MAX_CHUNKS property - model_schema = Mock() - model_schema.model_properties = {ModelPropertyKey.MAX_CHUNKS: 10} - model_type_instance.get_model_schema.return_value = model_schema - return model_instance +@pytest.fixture +def model_instance() -> Mock: + instance = Mock() + instance.model_name = "text-embedding-ada-002" + instance.provider = "openai" + instance.credentials = {"api_key": "test-key"} + schema = Mock(model_properties={ModelPropertyKey.MAX_CHUNKS: 10}) + instance.model_type_instance.get_model_schema.return_value = schema + return instance - @pytest.fixture - def sample_embedding_result(self): - """Create a sample EmbeddingResult for testing. - Returns: - EmbeddingResult: Mock embedding result with proper structure - """ - # Create normalized embedding vectors (dimension 1536 for ada-002) - embedding_vector = np.random.randn(1536) - normalized_vector = (embedding_vector / np.linalg.norm(embedding_vector)).tolist() +def _vector(dimension: int = 8, *, offset: float = 0) -> np.ndarray: + return np.arange(1 + offset, dimension + 1 + offset, dtype=float) - usage = EmbeddingUsage( - tokens=10, - total_tokens=10, + +def _result(*vectors: np.ndarray) -> EmbeddingResult: + return EmbeddingResult( + model="text-embedding-ada-002", + embeddings=list(vectors), + usage=EmbeddingUsage( + tokens=len(vectors), + total_tokens=len(vectors), unit_price=Decimal("0.0001"), price_unit=Decimal(1000), total_price=Decimal("0.000001"), currency="USD", - latency=0.5, - ) - - return EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized_vector], - usage=usage, - ) - - def test_embed_single_document_cache_miss(self, mock_model_instance, sample_embedding_result): - """Test embedding a single document when cache is empty. - - Verifies: - - Model invocation with correct parameters - - Embedding normalization - - Database cache storage - - Correct return value - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Python is a programming language"] - - # Mock database query to return no cached embedding (cache miss) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - - # Mock model invocation - mock_model_instance.invoke_text_embedding.return_value = sample_embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 1 - assert isinstance(result[0], list) - assert len(result[0]) == 1536 # ada-002 dimension - assert all(isinstance(x, float) for x in result[0]) - - # Verify model was invoked with correct parameters - mock_model_instance.invoke_text_embedding.assert_called_once_with( - texts=texts, - input_type=EmbeddingInputType.DOCUMENT, - ) - - # Verify embedding was added to database cache - mock_session.add.assert_called_once() - mock_session.commit.assert_called_once() - - def test_embed_multiple_documents_cache_miss(self, mock_model_instance): - """Test embedding multiple documents when cache is empty. - - Verifies: - - Batch processing of multiple texts - - Multiple embeddings returned - - All embeddings are properly normalized - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = [ - "Python is a programming language", - "JavaScript is used for web development", - "Machine learning is a subset of AI", - ] - - # Create multiple embedding vectors - embeddings = [] - for _ in range(3): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) - - usage = EmbeddingUsage( - tokens=30, - total_tokens=30, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000003"), - currency="USD", - latency=0.8, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 3 - assert all(len(emb) == 1536 for emb in result) - assert all(isinstance(emb, list) for emb in result) - - # Verify all embeddings are normalized (L2 norm ≈ 1.0) - for emb in result: - norm = np.linalg.norm(emb) - assert abs(norm - 1.0) < 0.01 # Allow small floating point error - - def test_embed_documents_cache_hit(self, mock_model_instance): - """Test embedding documents when embeddings are already cached. - - Verifies: - - Cached embeddings are retrieved from database - - Model is not invoked for cached texts - - Correct embeddings are returned - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Python is a programming language"] - - # Create cached embedding - cached_vector = np.random.randn(1536) - normalized_cached = (cached_vector / np.linalg.norm(cached_vector)).tolist() - - mock_cached_embedding = Mock(spec=Embedding) - mock_cached_embedding.get_embedding.return_value = normalized_cached - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - # Mock database to return cached embedding (cache hit) - mock_session.scalar.return_value = mock_cached_embedding - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 1 - assert result[0] == normalized_cached - - # Verify model was NOT invoked (cache hit) - mock_model_instance.invoke_text_embedding.assert_not_called() - - # Verify no new cache entries were added - mock_session.add.assert_not_called() - - def test_embed_documents_partial_cache_hit(self, mock_model_instance): - """Test embedding documents with mixed cache hits and misses. - - Verifies: - - Cached embeddings are used when available - - Only non-cached texts are sent to model - - Results are properly merged - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = [ - "Cached text 1", - "New text 1", - "New text 2", - ] - - # Create cached embedding for first text - cached_vector = np.random.randn(1536) - normalized_cached = (cached_vector / np.linalg.norm(cached_vector)).tolist() - - mock_cached_embedding = Mock(spec=Embedding) - mock_cached_embedding.get_embedding.return_value = normalized_cached - - # Create new embeddings for non-cached texts - new_embeddings = [] - for _ in range(2): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - new_embeddings.append(normalized) - - usage = EmbeddingUsage( - tokens=20, - total_tokens=20, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000002"), - currency="USD", - latency=0.6, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=new_embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - with patch("core.rag.embedding.cached_embedding.helper.generate_text_hash") as mock_hash: - # Mock hash generation to return predictable values - hash_counter = [0] - - def generate_hash(text): - hash_counter[0] += 1 - return f"hash_{hash_counter[0]}" - - mock_hash.side_effect = generate_hash - - # Mock database to return cached embedding only for first text (hash_1) - mock_session.scalar.side_effect = [mock_cached_embedding, None, None] - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 3 - assert result[0] == normalized_cached # From cache - # The model returns already normalized embeddings, but the code normalizes again - # So we just verify the structure and dimensions - assert result[1] is not None - assert isinstance(result[1], list) - assert len(result[1]) == 1536 - assert result[2] is not None - assert isinstance(result[2], list) - assert len(result[2]) == 1536 - - # Verify all embeddings are normalized - for emb in result: - if emb is not None: - norm = np.linalg.norm(emb) - assert abs(norm - 1.0) < 0.01 - - # Verify model was invoked only for non-cached texts - mock_model_instance.invoke_text_embedding.assert_called_once() - call_args = mock_model_instance.invoke_text_embedding.call_args - assert len(call_args.kwargs["texts"]) == 2 # Only 2 non-cached texts - - def test_embed_documents_large_batch(self, mock_model_instance): - """Test embedding a large batch of documents respecting MAX_CHUNKS. - - Verifies: - - Large batches are split according to MAX_CHUNKS - - Multiple model invocations for large batches - - All embeddings are returned correctly - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - # Create 25 texts, MAX_CHUNKS is 10, so should be 3 batches (10, 10, 5) - texts = [f"Text number {i}" for i in range(25)] - - # Create embeddings for each batch - def create_batch_result(batch_size): - embeddings = [] - for _ in range(batch_size): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) - - usage = EmbeddingUsage( - tokens=batch_size * 10, - total_tokens=batch_size * 10, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal(str(batch_size * 0.000001)), - currency="USD", - latency=0.5, - ) - - return EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - - # Mock model to return appropriate batch results - batch_results = [ - create_batch_result(10), - create_batch_result(10), - create_batch_result(5), - ] - mock_model_instance.invoke_text_embedding.side_effect = batch_results - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 25 - assert all(len(emb) == 1536 for emb in result) - - # Verify model was invoked 3 times (for 3 batches) - assert mock_model_instance.invoke_text_embedding.call_count == 3 - - # Verify batch sizes - calls = mock_model_instance.invoke_text_embedding.call_args_list - assert len(calls[0].kwargs["texts"]) == 10 - assert len(calls[1].kwargs["texts"]) == 10 - assert len(calls[2].kwargs["texts"]) == 5 - - def test_embed_documents_nan_handling(self, mock_model_instance, caplog: pytest.LogCaptureFixture): - """Test handling of NaN values in embeddings. - - Verifies: - - NaN values are detected - - NaN embeddings are skipped - - Warning is logged - - Valid embeddings are still processed - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Valid text", "Text that produces NaN"] - - # Create one valid embedding and one with NaN - # Note: The code normalizes again, so we provide unnormalized vector - valid_vector = np.random.randn(1536) - - # Create NaN vector - nan_vector = [float("nan")] * 1536 - - usage = EmbeddingUsage( - tokens=20, - total_tokens=20, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000002"), - currency="USD", - latency=0.5, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[valid_vector.tolist(), nan_vector], - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - with caplog.at_level(logging.WARNING, logger="core.rag.embedding.cached_embedding"): - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - # NaN embedding is skipped, so only 1 embedding in result - # The first position gets the valid embedding, second is None - assert len(result) == 2 - assert result[0] is not None - assert isinstance(result[0], list) - assert len(result[0]) == 1536 - # Second embedding should be None since NaN was skipped - assert result[1] is None - - # Verify warning was logged - assert sum(1 for r in caplog.records if r.levelno == logging.WARNING) >= 1 - assert any("Normalized embedding is nan" in record.message for record in caplog.records) - - def test_embed_documents_api_connection_error(self, mock_model_instance): - """Test handling of API connection errors during embedding. - - Verifies: - - Connection errors are propagated - - Database transaction is rolled back - - Error message is preserved - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Test text"] - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - - # Mock model to raise connection error - mock_model_instance.invoke_text_embedding.side_effect = InvokeConnectionError("Failed to connect to API") - - # Act & Assert - with pytest.raises(InvokeConnectionError) as exc_info: - cache_embedding.embed_documents(texts) - - assert "Failed to connect to API" in str(exc_info.value) - - # Verify database rollback was called - mock_session.rollback.assert_called() - - def test_embed_documents_rate_limit_error(self, mock_model_instance): - """Test handling of rate limit errors during embedding. - - Verifies: - - Rate limit errors are propagated - - Database transaction is rolled back - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Test text"] - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - - # Mock model to raise rate limit error - mock_model_instance.invoke_text_embedding.side_effect = InvokeRateLimitError("Rate limit exceeded") - - # Act & Assert - with pytest.raises(InvokeRateLimitError) as exc_info: - cache_embedding.embed_documents(texts) - - assert "Rate limit exceeded" in str(exc_info.value) - mock_session.rollback.assert_called() - - def test_embed_documents_authorization_error(self, mock_model_instance): - """Test handling of authorization errors during embedding. - - Verifies: - - Authorization errors are propagated - - Database transaction is rolled back - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Test text"] - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - - # Mock model to raise authorization error - mock_model_instance.invoke_text_embedding.side_effect = InvokeAuthorizationError("Invalid API key") - - # Act & Assert - with pytest.raises(InvokeAuthorizationError) as exc_info: - cache_embedding.embed_documents(texts) - - assert "Invalid API key" in str(exc_info.value) - mock_session.rollback.assert_called() - - def test_embed_documents_database_integrity_error(self, mock_model_instance, sample_embedding_result): - """Test handling of database integrity errors during cache storage. - - Verifies: - - Integrity errors are caught (e.g., duplicate hash) - - Database transaction is rolled back - - Embeddings are still returned - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Test text"] - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = sample_embedding_result - - # Mock database commit to raise IntegrityError - mock_session.commit.side_effect = IntegrityError("Duplicate key", None, None) - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - # Embeddings should still be returned despite cache error - assert len(result) == 1 - assert isinstance(result[0], list) - - # Verify rollback was called - mock_session.rollback.assert_called() - - -class TestCacheEmbeddingQuery: - """Test suite for CacheEmbedding.embed_query method. - - This class tests the query embedding functionality including: - - Single query embedding - - Redis cache management - - Cache hit/miss scenarios - - Error handling - """ - - @pytest.fixture - def mock_model_instance(self): - """Create a mock ModelInstance for testing.""" - model_instance = Mock() - model_instance.model_name = "text-embedding-ada-002" - model_instance.provider = "openai" - model_instance.credentials = {"api_key": "test-key"} - return model_instance - - def test_embed_query_cache_miss(self, mock_model_instance): - """Test embedding a query when Redis cache is empty. - - Verifies: - - Model invocation with QUERY input type - - Embedding normalization - - Redis cache storage - - Correct return value - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - query = "What is Python?" - - # Create embedding result - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - - usage = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized], - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.redis_client") as mock_redis: - # Mock Redis cache miss - mock_redis.get.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_query(query) - - # Assert - assert isinstance(result, list) - assert len(result) == 1536 - assert all(isinstance(x, float) for x in result) - - # Verify model was invoked with QUERY input type - mock_model_instance.invoke_text_embedding.assert_called_once_with( - texts=[query], - input_type=EmbeddingInputType.QUERY, - ) - - # Verify Redis cache was set - mock_redis.setex.assert_called_once() - # Cache key format: {provider}_{model}_{hash} - cache_key = mock_redis.setex.call_args[0][0] - assert "openai" in cache_key - assert "text-embedding-ada-002" in cache_key - - # Verify cache TTL is 600 seconds - assert mock_redis.setex.call_args[0][1] == 600 - - def test_embed_query_cache_hit(self, mock_model_instance): - """Test embedding a query when Redis cache contains the result. - - Verifies: - - Cached embedding is retrieved from Redis - - Model is not invoked - - Cache TTL is extended - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - query = "What is Python?" - - # Create cached embedding - vector = np.random.randn(1536) - normalized = vector / np.linalg.norm(vector) - - # Encode to base64 (as stored in Redis) - vector_bytes = normalized.tobytes() - encoded_vector = base64.b64encode(vector_bytes).decode("utf-8") - - with patch("core.rag.embedding.cached_embedding.redis_client") as mock_redis: - # Mock Redis cache hit - mock_redis.get.return_value = encoded_vector - - # Act - result = cache_embedding.embed_query(query) - - # Assert - assert isinstance(result, list) - assert len(result) == 1536 - - # Verify model was NOT invoked (cache hit) - mock_model_instance.invoke_text_embedding.assert_not_called() - - # Verify cache TTL was extended - mock_redis.expire.assert_called_once() - assert mock_redis.expire.call_args[0][1] == 600 - - def test_embed_query_nan_handling(self, mock_model_instance): - """Test handling of NaN values in query embeddings. - - Verifies: - - NaN values are detected - - ValueError is raised - - Error message is descriptive - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - query = "Query that produces NaN" - - # Create NaN embedding - nan_vector = [float("nan")] * 1536 - - usage = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[nan_vector], - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.redis_client") as mock_redis: - mock_redis.get.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act & Assert - with pytest.raises(ValueError) as exc_info: - cache_embedding.embed_query(query) - - assert "Normalized embedding is nan" in str(exc_info.value) - - def test_embed_query_connection_error(self, mock_model_instance): - """Test handling of connection errors during query embedding. - - Verifies: - - Connection errors are propagated - - Error is logged in debug mode - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - query = "Test query" - - with patch("core.rag.embedding.cached_embedding.redis_client") as mock_redis: - mock_redis.get.return_value = None - - # Mock model to raise connection error - mock_model_instance.invoke_text_embedding.side_effect = InvokeConnectionError("Connection failed") - - # Act & Assert - with pytest.raises(InvokeConnectionError) as exc_info: - cache_embedding.embed_query(query) - - assert "Connection failed" in str(exc_info.value) - - def test_embed_query_redis_cache_error(self, mock_model_instance): - """Test handling of Redis cache errors during storage. - - Verifies: - - Redis errors are caught - - Embedding is still returned - - Error is logged in debug mode - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - query = "Test query" - - # Create valid embedding - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - - usage = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized], - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.redis_client") as mock_redis: - mock_redis.get.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Mock Redis setex to raise error - mock_redis.setex.side_effect = Exception("Redis connection failed") - - # Act & Assert - with pytest.raises(Exception) as exc_info: - cache_embedding.embed_query(query) - - assert "Redis connection failed" in str(exc_info.value) - - -class TestEmbeddingModelSwitching: - """Test suite for embedding model switching functionality. - - This class tests the ability to switch between different embedding models - and providers, ensuring proper configuration and dimension handling. - """ - - def test_switch_between_openai_models(self): - """Test switching between different OpenAI embedding models. - - Verifies: - - Different models produce different cache keys - - Model name is correctly used in cache lookup - - Embeddings are model-specific - """ - # Arrange - model_instance_ada = Mock() - model_instance_ada.model_name = "text-embedding-ada-002" - model_instance_ada.provider = "openai" - - # Mock model type instance for ada - model_type_instance_ada = Mock() - model_instance_ada.model_type_instance = model_type_instance_ada - model_schema_ada = Mock() - model_schema_ada.model_properties = {ModelPropertyKey.MAX_CHUNKS: 10} - model_type_instance_ada.get_model_schema.return_value = model_schema_ada - - model_instance_3_small = Mock() - model_instance_3_small.model_name = "text-embedding-3-small" - model_instance_3_small.provider = "openai" - - # Mock model type instance for 3-small - model_type_instance_3_small = Mock() - model_instance_3_small.model_type_instance = model_type_instance_3_small - model_schema_3_small = Mock() - model_schema_3_small.model_properties = {ModelPropertyKey.MAX_CHUNKS: 10} - model_type_instance_3_small.get_model_schema.return_value = model_schema_3_small - - cache_ada = CacheEmbedding(model_instance_ada) - cache_3_small = CacheEmbedding(model_instance_3_small) - - text = "Test text" - - # Create different embeddings for each model - vector_ada = np.random.randn(1536) - normalized_ada = (vector_ada / np.linalg.norm(vector_ada)).tolist() - - vector_3_small = np.random.randn(1536) - normalized_3_small = (vector_3_small / np.linalg.norm(vector_3_small)).tolist() - - usage = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, - ) - - result_ada = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized_ada], - usage=usage, - ) - - result_3_small = EmbeddingResult( - model="text-embedding-3-small", - embeddings=[normalized_3_small], - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - - model_instance_ada.invoke_text_embedding.return_value = result_ada - model_instance_3_small.invoke_text_embedding.return_value = result_3_small - - # Act - embedding_ada = cache_ada.embed_documents([text]) - embedding_3_small = cache_3_small.embed_documents([text]) - - # Assert - # Both should return embeddings but they should be different - assert len(embedding_ada) == 1 - assert len(embedding_3_small) == 1 - assert embedding_ada[0] != embedding_3_small[0] - - # Verify both models were invoked - model_instance_ada.invoke_text_embedding.assert_called_once() - model_instance_3_small.invoke_text_embedding.assert_called_once() - - def test_switch_between_providers(self): - """Test switching between different embedding providers. - - Verifies: - - Different providers use separate cache namespaces - - Provider name is correctly used in cache lookup - """ - # Arrange - model_instance_openai = Mock() - model_instance_openai.model_name = "text-embedding-ada-002" - model_instance_openai.provider = "openai" - - model_instance_cohere = Mock() - model_instance_cohere.model_name = "embed-english-v3.0" - model_instance_cohere.provider = "cohere" - - cache_openai = CacheEmbedding(model_instance_openai) - cache_cohere = CacheEmbedding(model_instance_cohere) - - query = "Test query" - - # Create embeddings - vector_openai = np.random.randn(1536) - normalized_openai = (vector_openai / np.linalg.norm(vector_openai)).tolist() - - vector_cohere = np.random.randn(1024) # Cohere uses different dimension - normalized_cohere = (vector_cohere / np.linalg.norm(vector_cohere)).tolist() - - usage_openai = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, - ) - - usage_cohere = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0002"), - price_unit=Decimal(1000), - total_price=Decimal("0.000001"), - currency="USD", - latency=0.4, - ) - - result_openai = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized_openai], - usage=usage_openai, - ) - - result_cohere = EmbeddingResult( - model="embed-english-v3.0", - embeddings=[normalized_cohere], - usage=usage_cohere, - ) - - with patch("core.rag.embedding.cached_embedding.redis_client") as mock_redis: - mock_redis.get.return_value = None - - model_instance_openai.invoke_text_embedding.return_value = result_openai - model_instance_cohere.invoke_text_embedding.return_value = result_cohere - - # Act - embedding_openai = cache_openai.embed_query(query) - embedding_cohere = cache_cohere.embed_query(query) - - # Assert - assert len(embedding_openai) == 1536 # OpenAI dimension - assert len(embedding_cohere) == 1024 # Cohere dimension - - # Verify different cache keys were used - calls = mock_redis.setex.call_args_list - assert len(calls) == 2 - cache_key_openai = calls[0][0][0] - cache_key_cohere = calls[1][0][0] - - assert "openai" in cache_key_openai - assert "cohere" in cache_key_cohere - assert cache_key_openai != cache_key_cohere - - -class TestEmbeddingDimensionValidation: - """Test suite for embedding dimension validation. - - This class tests that embeddings maintain correct dimensions - and are properly normalized across different scenarios. - """ - - @pytest.fixture - def mock_model_instance(self): - """Create a mock ModelInstance for testing.""" - model_instance = Mock() - model_instance.model_name = "text-embedding-ada-002" - model_instance.provider = "openai" - model_instance.credentials = {"api_key": "test-key"} - - model_type_instance = Mock() - model_instance.model_type_instance = model_type_instance - - model_schema = Mock() - model_schema.model_properties = {ModelPropertyKey.MAX_CHUNKS: 10} - model_type_instance.get_model_schema.return_value = model_schema - - return model_instance - - def test_embedding_dimension_consistency(self, mock_model_instance): - """Test that all embeddings have consistent dimensions. - - Verifies: - - All embeddings have the same dimension - - Dimension matches model specification (1536 for ada-002) - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = [f"Text {i}" for i in range(5)] - - # Create embeddings with consistent dimension - embeddings = [] - for _ in range(5): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) - - usage = EmbeddingUsage( - tokens=50, - total_tokens=50, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000005"), - currency="USD", - latency=0.7, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 5 - - # All embeddings should have same dimension - dimensions = [len(emb) for emb in result] - assert all(dim == 1536 for dim in dimensions) - - # All embeddings should be lists of floats - for emb in result: - assert isinstance(emb, list) - assert all(isinstance(x, float) for x in emb) - - def test_embedding_normalization(self, mock_model_instance): - """Test that embeddings are properly normalized (L2 norm ≈ 1.0). - - Verifies: - - All embeddings are L2 normalized - - Normalization is consistent across batches - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Text 1", "Text 2", "Text 3"] - - # Create unnormalized vectors (will be normalized by the service) - embeddings = [] - for _ in range(3): - vector = np.random.randn(1536) * 10 # Unnormalized - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) - - usage = EmbeddingUsage( - tokens=30, - total_tokens=30, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000003"), - currency="USD", - latency=0.5, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - for emb in result: - norm = np.linalg.norm(emb) - # L2 norm should be approximately 1.0 - assert abs(norm - 1.0) < 0.01, f"Embedding not normalized: norm={norm}" - - def test_different_model_dimensions(self): - """Test handling of different embedding dimensions for different models. - - Verifies: - - Different models can have different dimensions - - Dimensions are correctly preserved - """ - # Arrange - OpenAI ada-002 (1536 dimensions) - model_instance_ada = Mock() - model_instance_ada.model_name = "text-embedding-ada-002" - model_instance_ada.provider = "openai" - - # Mock model type instance for ada - model_type_instance_ada = Mock() - model_instance_ada.model_type_instance = model_type_instance_ada - model_schema_ada = Mock() - model_schema_ada.model_properties = {ModelPropertyKey.MAX_CHUNKS: 10} - model_type_instance_ada.get_model_schema.return_value = model_schema_ada - - cache_ada = CacheEmbedding(model_instance_ada) - - vector_ada = np.random.randn(1536) - normalized_ada = (vector_ada / np.linalg.norm(vector_ada)).tolist() - - usage_ada = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, - ) - - result_ada = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized_ada], - usage=usage_ada, - ) - - # Arrange - Cohere embed-english-v3.0 (1024 dimensions) - model_instance_cohere = Mock() - model_instance_cohere.model_name = "embed-english-v3.0" - model_instance_cohere.provider = "cohere" - - # Mock model type instance for cohere - model_type_instance_cohere = Mock() - model_instance_cohere.model_type_instance = model_type_instance_cohere - model_schema_cohere = Mock() - model_schema_cohere.model_properties = {ModelPropertyKey.MAX_CHUNKS: 10} - model_type_instance_cohere.get_model_schema.return_value = model_schema_cohere - - cache_cohere = CacheEmbedding(model_instance_cohere) - - vector_cohere = np.random.randn(1024) - normalized_cohere = (vector_cohere / np.linalg.norm(vector_cohere)).tolist() - - usage_cohere = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0002"), - price_unit=Decimal(1000), - total_price=Decimal("0.000001"), - currency="USD", - latency=0.4, - ) - - result_cohere = EmbeddingResult( - model="embed-english-v3.0", - embeddings=[normalized_cohere], - usage=usage_cohere, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - - model_instance_ada.invoke_text_embedding.return_value = result_ada - model_instance_cohere.invoke_text_embedding.return_value = result_cohere - - # Act - embedding_ada = cache_ada.embed_documents(["Test"]) - embedding_cohere = cache_cohere.embed_documents(["Test"]) - - # Assert - assert len(embedding_ada[0]) == 1536 # OpenAI dimension - assert len(embedding_cohere[0]) == 1024 # Cohere dimension - - -class TestEmbeddingEdgeCases: - """Test suite for edge cases and special scenarios. - - This class tests unusual inputs and boundary conditions including: - - Empty inputs (empty list, empty strings) - - Very long texts (exceeding typical limits) - - Special characters and Unicode - - Whitespace-only texts - - Duplicate texts in same batch - - Mixed valid and invalid inputs - """ - - @pytest.fixture - def mock_model_instance(self): - """Create a mock ModelInstance for testing. - - Returns: - Mock: Configured ModelInstance with standard settings - - Model: text-embedding-ada-002 - - Provider: openai - - MAX_CHUNKS: 10 - """ - model_instance = Mock() - model_instance.model_name = "text-embedding-ada-002" - model_instance.provider = "openai" - - model_type_instance = Mock() - model_instance.model_type_instance = model_type_instance - - model_schema = Mock() - model_schema.model_properties = {ModelPropertyKey.MAX_CHUNKS: 10} - model_type_instance.get_model_schema.return_value = model_schema - - return model_instance - - def test_embed_empty_list(self, mock_model_instance): - """Test embedding an empty list of documents. - - Verifies: - - Empty list returns empty result - - No model invocation occurs - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = [] - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert result == [] - mock_model_instance.invoke_text_embedding.assert_not_called() - - def test_embed_empty_string(self, mock_model_instance): - """Test embedding an empty string. - - Verifies: - - Empty string is handled correctly - - Model is invoked with empty string - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = [""] - - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - - usage = EmbeddingUsage( - tokens=0, - total_tokens=0, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal(0), - currency="USD", latency=0.1, + ), + ) + + +def _persist_embedding( + session: Session, + *, + text: str, + vector: list[float], + model_name: str = "text-embedding-ada-002", + provider_name: str = "openai", +) -> Embedding: + row = Embedding( + model_name=model_name, + hash=helper.generate_text_hash(text), + provider_name=provider_name, + embedding=pickle.dumps(vector, protocol=pickle.HIGHEST_PROTOCOL), + ) + session.add(row) + session.commit() + return row + + +class TestDocumentCache: + def test_cache_miss_normalizes_and_persists(self, sqlite_embedding_session: Session, model_instance: Mock) -> None: + model_instance.invoke_text_embedding.return_value = _result(_vector()) + + result = CacheEmbedding(model_instance).embed_documents(["new text"]) + + assert np.linalg.norm(result[0]) == pytest.approx(1.0) + model_instance.invoke_text_embedding.assert_called_once_with( + texts=["new text"], input_type=EmbeddingInputType.DOCUMENT ) + persisted = sqlite_embedding_session.scalar(select(Embedding)) + assert persisted is not None + assert persisted.get_embedding() == result[0] - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized], - usage=usage, - ) + def test_cache_hit_uses_persisted_vector(self, sqlite_embedding_session: Session, model_instance: Mock) -> None: + cached_vector = (_vector() / np.linalg.norm(_vector())).tolist() + row = _persist_embedding(sqlite_embedding_session, text="cached", vector=cached_vector) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result + result = CacheEmbedding(model_instance).embed_documents(["cached"]) - # Act - result = cache_embedding.embed_documents(texts) + assert result == [cached_vector] + model_instance.invoke_text_embedding.assert_not_called() + assert sqlite_embedding_session.scalars(select(Embedding)).all() == [row] - # Assert - assert len(result) == 1 - assert len(result[0]) == 1536 + def test_partial_hit_preserves_input_order(self, sqlite_embedding_session: Session, model_instance: Mock) -> None: + cached_vector = (_vector(offset=10) / np.linalg.norm(_vector(offset=10))).tolist() + _persist_embedding(sqlite_embedding_session, text="cached", vector=cached_vector) + model_instance.invoke_text_embedding.return_value = _result(_vector(), _vector(offset=20)) - def test_embed_very_long_text(self, mock_model_instance): - """Test embedding very long text. + result = CacheEmbedding(model_instance).embed_documents(["cached", "new one", "new two"]) - Verifies: - - Long texts are handled correctly - - No truncation errors occur - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - # Create a very long text (10000 characters) - long_text = "Python " * 2000 - texts = [long_text] + assert result[0] == cached_vector + assert all(np.linalg.norm(vector) == pytest.approx(1.0) for vector in result) + assert model_instance.invoke_text_embedding.call_args.kwargs["texts"] == ["new one", "new two"] + assert len(sqlite_embedding_session.scalars(select(Embedding)).all()) == 3 - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - - usage = EmbeddingUsage( - tokens=2000, - total_tokens=2000, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0002"), - currency="USD", - latency=1.5, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized], - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 1 - assert len(result[0]) == 1536 - - def test_embed_special_characters(self, mock_model_instance): - """Test embedding text with special characters. - - Verifies: - - Special characters are handled correctly - - Unicode characters work properly - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = [ - "Hello 世界! 🌍", - "Special chars: @#$%^&*()", - "Newlines\nand\ttabs", + def test_large_batch_respects_model_chunk_limit( + self, sqlite_embedding_session: Session, model_instance: Mock + ) -> None: + texts = [f"text-{index}" for index in range(25)] + model_instance.invoke_text_embedding.side_effect = [ + _result(*(_vector(offset=index) for index in range(10))), + _result(*(_vector(offset=index) for index in range(10, 20))), + _result(*(_vector(offset=index) for index in range(20, 25))), ] - embeddings = [] - for _ in range(3): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) + result = CacheEmbedding(model_instance).embed_documents(texts) - usage = EmbeddingUsage( - tokens=30, - total_tokens=30, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000003"), - currency="USD", - latency=0.5, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 3 - assert all(len(emb) == 1536 for emb in result) - - def test_embed_whitespace_only_text(self, mock_model_instance): - """Test embedding text containing only whitespace. - - Verifies: - - Whitespace-only texts are handled correctly - - Model is invoked with whitespace text - - Valid embedding is returned - - Context: - -------- - Whitespace-only texts can occur in real-world scenarios when - processing documents with formatting issues or empty sections. - The embedding model should handle these gracefully. - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = [" ", "\t\t", "\n\n\n"] - - # Create embeddings for whitespace texts - embeddings = [] - for _ in range(3): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) - - usage = EmbeddingUsage( - tokens=3, - total_tokens=3, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000003"), - currency="USD", - latency=0.2, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 3 - assert all(isinstance(emb, list) for emb in result) - assert all(len(emb) == 1536 for emb in result) - - def test_embed_duplicate_texts_in_batch(self, mock_model_instance): - """Test embedding when same text appears multiple times in batch. - - Verifies: - - Duplicate texts are handled correctly - - Each duplicate gets its own embedding - - All duplicates are processed - - Context: - -------- - In batch processing, the same text might appear multiple times. - The current implementation processes all texts individually, - even if they're duplicates. This ensures each position in the - input list gets a corresponding embedding in the output. - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - # Same text repeated 3 times - texts = ["Duplicate text", "Duplicate text", "Duplicate text"] - - # Create embeddings for all three (even though they're duplicates) - embeddings = [] - for _ in range(3): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) - - usage = EmbeddingUsage( - tokens=30, - total_tokens=30, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000003"), - currency="USD", - latency=0.3, - ) - - # Model returns embeddings for all texts - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - # All three should have embeddings - assert len(result) == 3 - # Model should be called once - mock_model_instance.invoke_text_embedding.assert_called_once() - # All three texts are sent to model (no deduplication) - call_args = mock_model_instance.invoke_text_embedding.call_args - assert len(call_args.kwargs["texts"]) == 3 - - def test_embed_mixed_languages(self, mock_model_instance): - """Test embedding texts in different languages. - - Verifies: - - Multi-language texts are handled correctly - - Unicode characters from various scripts work - - Embeddings are generated for all languages - - Context: - -------- - Modern embedding models support multiple languages. - This test ensures the service handles various scripts: - - Latin (English) - - CJK (Chinese, Japanese, Korean) - - Cyrillic (Russian) - - Arabic - - Emoji and symbols - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = [ - "Hello World", # English - "你好世界", # Chinese - "こんにちは世界", # Japanese - "Привет мир", # Russian - "مرحبا بالعالم", # Arabic - "🌍🌎🌏", # Emoji + assert len(result) == 25 + assert [len(call.kwargs["texts"]) for call in model_instance.invoke_text_embedding.call_args_list] == [ + 10, + 10, + 5, ] + assert len(sqlite_embedding_session.scalars(select(Embedding)).all()) == 25 - # Create embeddings for each language - embeddings = [] - for _ in range(6): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) + @pytest.mark.parametrize( + "provider_error", + [InvokeConnectionError("offline"), InvokeRateLimitError("limited"), InvokeAuthorizationError("denied")], + ) + def test_provider_error_rolls_back_real_transaction( + self, + sqlite_embedding_session: Session, + model_instance: Mock, + provider_error: Exception, + ) -> None: + model_instance.invoke_text_embedding.side_effect = provider_error - usage = EmbeddingUsage( - tokens=60, - total_tokens=60, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000006"), - currency="USD", - latency=0.8, + with pytest.raises(type(provider_error)): + CacheEmbedding(model_instance).embed_documents(["failure"]) + + assert not sqlite_embedding_session.in_transaction() + assert sqlite_embedding_session.scalar(select(Embedding)) is None + + def test_integrity_race_rolls_back_without_losing_result( + self, sqlite_embedding_session: Session, model_instance: Mock + ) -> None: + model_instance.invoke_text_embedding.return_value = _result(_vector()) + + def raise_integrity_error(_session: Session) -> None: + raise IntegrityError("duplicate embedding", {}, RuntimeError("unique constraint")) + + event.listen(sqlite_embedding_session, "before_commit", raise_integrity_error) + try: + result = CacheEmbedding(model_instance).embed_documents(["racing insert"]) + finally: + event.remove(sqlite_embedding_session, "before_commit", raise_integrity_error) + + assert np.linalg.norm(result[0]) == pytest.approx(1.0) + assert not sqlite_embedding_session.in_transaction() + assert sqlite_embedding_session.scalar(select(Embedding)) is None + + def test_nan_vector_is_skipped_and_logged( + self, + sqlite_embedding_session: Session, + model_instance: Mock, + caplog: pytest.LogCaptureFixture, + ) -> None: + model_instance.invoke_text_embedding.return_value = _result(np.array([np.nan, 1.0])) + + with caplog.at_level("WARNING", logger="core.rag.embedding.cached_embedding"): + result = CacheEmbedding(model_instance).embed_documents(["invalid"]) + + assert result == [None] + assert "Normalized embedding is nan" in caplog.text + assert sqlite_embedding_session.scalar(select(Embedding)) is None + + def test_duplicate_text_is_cached_once(self, sqlite_embedding_session: Session, model_instance: Mock) -> None: + model_instance.invoke_text_embedding.return_value = _result(_vector(), _vector()) + + result = CacheEmbedding(model_instance).embed_documents(["same", "same"]) + + assert result[0] == result[1] + assert len(sqlite_embedding_session.scalars(select(Embedding)).all()) == 1 + + def test_provider_and_model_are_part_of_cache_identity(self, sqlite_embedding_session: Session) -> None: + first = Mock(model_name="shared-model", provider="provider-a", credentials={}) + first.model_type_instance.get_model_schema.return_value = Mock( + model_properties={ModelPropertyKey.MAX_CHUNKS: 10} ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, + first.invoke_text_embedding.return_value = _result(_vector()) + second = Mock(model_name="shared-model", provider="provider-b", credentials={}) + second.model_type_instance.get_model_schema.return_value = Mock( + model_properties={ModelPropertyKey.MAX_CHUNKS: 10} ) + second.invoke_text_embedding.return_value = _result(_vector(offset=10)) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result + first_result = CacheEmbedding(first).embed_documents(["same text"]) + second_result = CacheEmbedding(second).embed_documents(["same text"]) - # Act - result = cache_embedding.embed_documents(texts) + assert first_result != second_result + assert len(sqlite_embedding_session.scalars(select(Embedding)).all()) == 2 - # Assert - assert len(result) == 6 - assert all(isinstance(emb, list) for emb in result) - assert all(len(emb) == 1536 for emb in result) - # Verify all embeddings are normalized - for emb in result: - norm = np.linalg.norm(emb) - assert abs(norm - 1.0) < 0.01 + def test_empty_input_does_not_touch_provider(self, model_instance: Mock) -> None: + assert CacheEmbedding(model_instance).embed_documents([]) == [] + model_instance.invoke_text_embedding.assert_not_called() - def test_embed_query_uses_bound_model_instance(self, mock_model_instance): - """Test query embedding using the provided model instance. - Verifies: - - Embedding generation works with the injected model instance - - Query input type is preserved - - No extra binding step is required at call time - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - query = "What is machine learning?" +class TestMultimodalDocumentCache: + def test_multimodal_miss_persists_and_hit_reuses( + self, sqlite_embedding_session: Session, model_instance: Mock + ) -> None: + document = {"file_id": "file-1"} + model_instance.invoke_multimodal_embedding.return_value = _result(_vector()) + cache = CacheEmbedding(model_instance) - # Create embedding - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() + generated = cache.embed_multimodal_documents([document]) + cached = cache.embed_multimodal_documents([document]) - usage = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, + assert cached == generated + model_instance.invoke_multimodal_embedding.assert_called_once_with( + multimodel_documents=[document], input_type=EmbeddingInputType.DOCUMENT ) + assert sqlite_embedding_session.scalar(select(Embedding).where(Embedding.hash == "file-1")) is not None - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized], - usage=usage, + +class TestQueryCache: + @patch("core.rag.embedding.cached_embedding.redis_client") + def test_query_cache_miss_normalizes_and_stores(self, redis: Mock, model_instance: Mock) -> None: + redis.get.return_value = None + model_instance.invoke_text_embedding.return_value = _result(_vector()) + + result = CacheEmbedding(model_instance).embed_query("query") + + assert np.linalg.norm(result) == pytest.approx(1.0) + model_instance.invoke_text_embedding.assert_called_once_with( + texts=["query"], input_type=EmbeddingInputType.QUERY ) + redis.setex.assert_called_once() - with patch("core.rag.embedding.cached_embedding.redis_client") as mock_redis: - mock_redis.get.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result + @patch("core.rag.embedding.cached_embedding.redis_client") + def test_query_cache_hit_refreshes_ttl(self, redis: Mock, model_instance: Mock) -> None: + vector = np.array([0.25, 0.75], dtype=float) + redis.get.return_value = base64.b64encode(vector.tobytes()) - # Act - result = cache_embedding.embed_query(query) + result = CacheEmbedding(model_instance).embed_query("query") - # Assert - assert isinstance(result, list) - assert len(result) == 1536 + assert result == vector.tolist() + redis.expire.assert_called_once_with(f"openai_text-embedding-ada-002_{helper.generate_text_hash('query')}", 600) + model_instance.invoke_text_embedding.assert_not_called() - mock_model_instance.invoke_text_embedding.assert_called_once_with( - texts=[query], - input_type=EmbeddingInputType.QUERY, - ) + @patch("core.rag.embedding.cached_embedding.redis_client") + def test_query_nan_is_rejected(self, redis: Mock, model_instance: Mock) -> None: + redis.get.return_value = None + model_instance.invoke_text_embedding.return_value = _result(np.array([np.nan, 1.0])) - def test_embed_documents_uses_bound_model_instance(self, mock_model_instance): - """Test document embedding using the provided model instance. + with pytest.raises(ValueError, match="Normalized embedding is nan"): + CacheEmbedding(model_instance).embed_query("query") - Verifies: - - Batch processing uses the injected model instance - - Document input type is preserved - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - texts = ["Document 1", "Document 2"] + redis.setex.assert_not_called() - # Create embeddings - embeddings = [] - for _ in range(2): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) + @patch("core.rag.embedding.cached_embedding.redis_client") + def test_redis_write_error_is_propagated(self, redis: Mock, model_instance: Mock) -> None: + redis.get.return_value = None + redis.setex.side_effect = ConnectionError("redis unavailable") + model_instance.invoke_text_embedding.return_value = _result(_vector()) - usage = EmbeddingUsage( - tokens=20, - total_tokens=20, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.000002"), - currency="USD", - latency=0.5, - ) + with pytest.raises(ConnectionError, match="redis unavailable"): + CacheEmbedding(model_instance).embed_query("query") - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) + @patch("core.rag.embedding.cached_embedding.redis_client") + def test_multimodal_query_uses_file_cache_key(self, redis: Mock, model_instance: Mock) -> None: + redis.get.return_value = None + model_instance.invoke_multimodal_embedding.return_value = _result(_vector()) - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result + result = CacheEmbedding(model_instance).embed_multimodal_query({"file_id": "file-1"}) - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 2 - - mock_model_instance.invoke_text_embedding.assert_called_once() - call_args = mock_model_instance.invoke_text_embedding.call_args - assert call_args.kwargs["input_type"] == EmbeddingInputType.DOCUMENT - - -class TestEmbeddingCachePerformance: - """Test suite for cache performance and optimization scenarios. - - This class tests cache-related performance optimizations: - - Cache hit rate improvements - - Batch processing efficiency - - Memory usage optimization - - Cache key generation - - TTL (Time To Live) management - """ - - @pytest.fixture - def mock_model_instance(self): - """Create a mock ModelInstance for testing. - - Returns: - Mock: Configured ModelInstance for performance testing - - Model: text-embedding-ada-002 - - Provider: openai - - MAX_CHUNKS: 10 - """ - model_instance = Mock() - model_instance.model_name = "text-embedding-ada-002" - model_instance.provider = "openai" - - model_type_instance = Mock() - model_instance.model_type_instance = model_type_instance - - model_schema = Mock() - model_schema.model_properties = {ModelPropertyKey.MAX_CHUNKS: 10} - model_type_instance.get_model_schema.return_value = model_schema - - return model_instance - - def test_cache_hit_reduces_api_calls(self, mock_model_instance): - """Test that cache hits prevent unnecessary API calls. - - Verifies: - - First call triggers API request - - Second call uses cache (no API call) - - Cache significantly reduces API usage - - Context: - -------- - Caching is critical for: - 1. Reducing API costs - 2. Improving response time - 3. Reducing rate limit pressure - 4. Better user experience - - This test demonstrates the cache working as expected. - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - text = "Frequently used text" - - # Create cached embedding - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - - mock_cached_embedding = Mock(spec=Embedding) - mock_cached_embedding.get_embedding.return_value = normalized - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - # First call: cache miss - mock_session.scalar.return_value = None - - usage = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized], - usage=usage, - ) - - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - First call (cache miss) - result1 = cache_embedding.embed_documents([text]) - - # Assert - Model was called - assert mock_model_instance.invoke_text_embedding.call_count == 1 - assert len(result1) == 1 - - # Arrange - Second call: cache hit - mock_session.scalar.return_value = mock_cached_embedding - - # Act - Second call (cache hit) - result2 = cache_embedding.embed_documents([text]) - - # Assert - Model was NOT called again (still 1 call total) - assert mock_model_instance.invoke_text_embedding.call_count == 1 - assert len(result2) == 1 - assert result2[0] == normalized # Same embedding from cache - - def test_batch_processing_efficiency(self, mock_model_instance): - """Test that batch processing is more efficient than individual calls. - - Verifies: - - Multiple texts are processed in single API call - - Batch size respects MAX_CHUNKS limit - - Batching reduces total API calls - - Context: - -------- - Batch processing is essential for: - 1. Reducing API overhead - 2. Better throughput - 3. Lower latency per text - 4. Cost optimization - - Example: 100 texts in batches of 10 = 10 API calls - vs 100 individual calls = 100 API calls - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - # 15 texts should be processed in 2 batches (10 + 5) - texts = [f"Text {i}" for i in range(15)] - - # Create embeddings for each batch - def create_batch_result(batch_size): - """Helper function to create batch embedding results.""" - embeddings = [] - for _ in range(batch_size): - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - embeddings.append(normalized) - - usage = EmbeddingUsage( - tokens=batch_size * 10, - total_tokens=batch_size * 10, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal(str(batch_size * 0.000001)), - currency="USD", - latency=0.5, - ) - - return EmbeddingResult( - model="text-embedding-ada-002", - embeddings=embeddings, - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.db.session") as mock_session: - mock_session.scalar.return_value = None - - # Mock model to return appropriate batch results - batch_results = [ - create_batch_result(10), # First batch - create_batch_result(5), # Second batch - ] - mock_model_instance.invoke_text_embedding.side_effect = batch_results - - # Act - result = cache_embedding.embed_documents(texts) - - # Assert - assert len(result) == 15 - # Only 2 API calls for 15 texts (batched) - assert mock_model_instance.invoke_text_embedding.call_count == 2 - - # Verify batch sizes - calls = mock_model_instance.invoke_text_embedding.call_args_list - assert len(calls[0].kwargs["texts"]) == 10 # First batch - assert len(calls[1].kwargs["texts"]) == 5 # Second batch - - def test_redis_cache_expiration(self, mock_model_instance): - """Test Redis cache TTL (Time To Live) management. - - Verifies: - - Cache entries have appropriate TTL (600 seconds) - - TTL is extended on cache hits - - Expired entries are regenerated - - Context: - -------- - Redis cache TTL ensures: - 1. Memory doesn't grow unbounded - 2. Stale embeddings are refreshed - 3. Frequently used queries stay cached longer - 4. Infrequently used queries expire naturally - """ - # Arrange - cache_embedding = CacheEmbedding(mock_model_instance) - query = "Test query" - - vector = np.random.randn(1536) - normalized = (vector / np.linalg.norm(vector)).tolist() - - usage = EmbeddingUsage( - tokens=5, - total_tokens=5, - unit_price=Decimal("0.0001"), - price_unit=Decimal(1000), - total_price=Decimal("0.0000005"), - currency="USD", - latency=0.3, - ) - - embedding_result = EmbeddingResult( - model="text-embedding-ada-002", - embeddings=[normalized], - usage=usage, - ) - - with patch("core.rag.embedding.cached_embedding.redis_client") as mock_redis: - # Test cache miss - sets TTL - mock_redis.get.return_value = None - mock_model_instance.invoke_text_embedding.return_value = embedding_result - - # Act - cache_embedding.embed_query(query) - - # Assert - TTL was set to 600 seconds - mock_redis.setex.assert_called_once() - call_args = mock_redis.setex.call_args - assert call_args[0][1] == 600 # TTL in seconds - - # Test cache hit - extends TTL - mock_redis.reset_mock() - vector_bytes = np.array(normalized).tobytes() - encoded_vector = base64.b64encode(vector_bytes).decode("utf-8") - mock_redis.get.return_value = encoded_vector - - # Act - cache_embedding.embed_query(query) - - # Assert - TTL was extended - mock_redis.expire.assert_called_once() - assert mock_redis.expire.call_args[0][1] == 600 + assert np.linalg.norm(result) == pytest.approx(1.0) + assert redis.setex.call_args.args[:2] == ("openai_text-embedding-ada-002_file-1", 600) diff --git a/api/tests/unit_tests/core/rag/embedding/test_token_counter.py b/api/tests/unit_tests/core/rag/embedding/test_token_counter.py index 0c134a94c90..733f55fa16c 100644 --- a/api/tests/unit_tests/core/rag/embedding/test_token_counter.py +++ b/api/tests/unit_tests/core/rag/embedding/test_token_counter.py @@ -1,4 +1,4 @@ -from unittest.mock import Mock, patch +from unittest.mock import patch from core.rag.embedding.token_counter import calculate_segment_token_counts from core.rag.index_processor.constant.index_type import IndexTechniqueType @@ -7,11 +7,12 @@ from models.dataset import Dataset def test_high_quality_counts_each_document_once() -> None: - dataset = Mock(spec=Dataset) - dataset.tenant_id = "tenant-1" - dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY - dataset.embedding_model_provider = "provider" - dataset.embedding_model = "model" + dataset = Dataset( + tenant_id="tenant-1", + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + embedding_model_provider="provider", + embedding_model="model", + ) documents = [ Document(page_content="first", metadata={}), Document(page_content="second", metadata={}), @@ -31,8 +32,9 @@ def test_high_quality_counts_each_document_once() -> None: def test_economy_returns_zero_without_loading_model() -> None: - dataset = Mock(spec=Dataset) - dataset.indexing_technique = IndexTechniqueType.ECONOMY + dataset = Dataset( + indexing_technique=IndexTechniqueType.ECONOMY, + ) documents = [ Document(page_content="first", metadata={}), Document(page_content="second", metadata={}), @@ -46,8 +48,9 @@ def test_economy_returns_zero_without_loading_model() -> None: def test_empty_documents_return_without_loading_model() -> None: - dataset = Mock(spec=Dataset) - dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY + dataset = Dataset( + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + ) with patch("core.rag.embedding.token_counter.ModelManager.for_tenant") as model_manager_factory: result = calculate_segment_token_counts(dataset=dataset, documents=[]) diff --git a/api/tests/unit_tests/core/rag/extractor/firecrawl/test_firecrawl.py b/api/tests/unit_tests/core/rag/extractor/firecrawl/test_firecrawl.py index db49221583f..3bf1441c0af 100644 --- a/api/tests/unit_tests/core/rag/extractor/firecrawl/test_firecrawl.py +++ b/api/tests/unit_tests/core/rag/extractor/firecrawl/test_firecrawl.py @@ -35,6 +35,18 @@ class TestFirecrawlApp: } assert app._build_url("/v2/crawl") == "https://custom.firecrawl.dev/v2/crawl" + def test_requests_use_bounded_timeout(self, mocker: MockerFixture): + """Outbound requests must carry a bounded timeout so a hanging endpoint cannot block extraction.""" + app = FirecrawlApp(api_key="fc-key", base_url="https://custom.firecrawl.dev") + mock_post = mocker.patch("httpx.post", return_value=_response(200, {"id": "job-1"})) + mock_get = mocker.patch("httpx.get", return_value=_response(200, {"status": "completed"})) + + app._post_request("https://custom.firecrawl.dev/v1/crawl", {}, app._prepare_headers()) + app._get_request("https://custom.firecrawl.dev/v1/crawl/job-1", app._prepare_headers()) + + assert mock_post.call_args.kwargs["timeout"] == firecrawl_module._REQUEST_TIMEOUT + assert mock_get.call_args.kwargs["timeout"] == firecrawl_module._REQUEST_TIMEOUT + def test_scrape_url_success(self, mocker: MockerFixture): app = FirecrawlApp(api_key="fc-key", base_url="https://custom.firecrawl.dev") mocker.patch( diff --git a/api/tests/unit_tests/core/rag/extractor/test_excel_extractor.py b/api/tests/unit_tests/core/rag/extractor/test_excel_extractor.py index ebe24c29007..af12b8780f6 100644 --- a/api/tests/unit_tests/core/rag/extractor/test_excel_extractor.py +++ b/api/tests/unit_tests/core/rag/extractor/test_excel_extractor.py @@ -2,9 +2,21 @@ from types import SimpleNamespace import pandas as pd import pytest +from sqlalchemy import Engine, select +from sqlalchemy.orm import Session, sessionmaker import core.rag.extractor.excel_extractor as excel_module from core.rag.extractor.excel_extractor import ExcelExtractor +from models.base import TypeBase +from models.model import UploadFile + + +@pytest.fixture +def database_session_maker(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> sessionmaker[Session]: + TypeBase.metadata.create_all(sqlite_engine, tables=[UploadFile.__table__]) + session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr(excel_module.session_factory, "create_session", session_maker) + return session_maker class _FakeCell: @@ -58,82 +70,22 @@ class _FakeImage: return self._raw_data -class _FieldExpression: - def __eq__(self, other): - return ("eq", other) - - def in_(self, values): - return ("in", tuple(values)) - - -class _SelectStub: - def where(self, *args, **kwargs): - return self - - -class _FakeUploadFile: - tenant_id = _FieldExpression() - key = _FieldExpression() - _i = 0 - - def __init__(self, **kwargs): - type(self)._i += 1 - self.id = f"u{self._i}" - self.key = kwargs["key"] - - -class _PersistentSession: - def __init__(self, persisted): - self._persisted = persisted - self.added = [] - self.commit_count = 0 - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb): - return False - - def scalars(self, _stmt): - return SimpleNamespace(all=lambda: list(self._persisted.values())) - - def add_all(self, objects) -> None: - self.added.extend(objects) - - def commit(self) -> None: - self.commit_count += 1 - for upload_file in self.added: - self._persisted[upload_file.key] = upload_file - self.added.clear() - - -class _PersistentSessionFactory: - def __init__(self): - self.persisted = {} - self.sessions = [] - - def create_session(self): - session = _PersistentSession(self.persisted) - self.sessions.append(session) - return session - - def _patch_image_persistence(monkeypatch: pytest.MonkeyPatch): saves: list[tuple[str, bytes]] = [] - session_factory = _PersistentSessionFactory() def save(key: str, data: bytes) -> None: saves.append((key, data)) - _FakeUploadFile._i = 0 - monkeypatch.setattr(excel_module, "storage", SimpleNamespace(save=save)) - monkeypatch.setattr(excel_module, "session_factory", session_factory) - monkeypatch.setattr(excel_module, "select", lambda *args, **kwargs: _SelectStub()) - monkeypatch.setattr(excel_module, "UploadFile", _FakeUploadFile) + monkeypatch.setattr(excel_module.storage, "save", save) monkeypatch.setattr(excel_module.dify_config, "FILES_URL", "http://files.local", raising=False) monkeypatch.setattr(excel_module.dify_config, "STORAGE_TYPE", "local", raising=False) - return saves, session_factory + return saves + + +def _get_upload_files(session_maker: sessionmaker[Session]) -> list[UploadFile]: + with session_maker() as session: + return list(session.scalars(select(UploadFile)).all()) class TestExcelExtractor: @@ -160,7 +112,11 @@ class TestExcelExtractor: assert docs[1].page_content == '"Name":"";"Link":"123"' assert all(doc.metadata["source"] == "/tmp/sample.xlsx" for doc in docs) - def test_extract_xlsx_turns_embedded_images_into_markdown_links(self, monkeypatch: pytest.MonkeyPatch): + def test_extract_xlsx_turns_embedded_images_into_markdown_links( + self, + monkeypatch: pytest.MonkeyPatch, + database_session_maker: sessionmaker[Session], + ): image_bytes = b"\x89PNG\r\n\x1a\nexcel-image" sheet = _FakeSheet( header_rows=[("Question", "Answer", "Image")], @@ -175,7 +131,7 @@ class TestExcelExtractor: ) workbook = _FakeWorkbook({"Data": sheet}) monkeypatch.setattr(excel_module, "load_workbook", lambda *args, **kwargs: workbook) - saves, session_factory = _patch_image_persistence(monkeypatch) + saves = _patch_image_persistence(monkeypatch) extractor = ExcelExtractor( "/tmp/sample.xlsx", @@ -184,23 +140,30 @@ class TestExcelExtractor: source_file_id="source-file-1", ) docs = extractor.extract() + upload_files = _get_upload_files(database_session_maker) assert workbook.closed is True assert len(docs) == 2 + assert len(upload_files) == 1 assert docs[0].page_content == ( '"Question":"Q1";"Answer":"A1";' - '"Image":"![image](http://files.local/files/u1/file-preview) ' - '![image](http://files.local/files/u1/file-preview)"' + f'"Image":"![image](http://files.local/files/{upload_files[0].id}/file-preview) ' + f'![image](http://files.local/files/{upload_files[0].id}/file-preview)"' ) assert docs[1].page_content == '"Question":"Q2";"Answer":"A2";"Image":""' assert len(saves) == 1 assert saves[0][0].startswith("image_files/tenant-1/source-file-1/") assert saves[0][0].endswith(".png") assert saves[0][1] == image_bytes - assert len(session_factory.persisted) == 1 - assert [session.commit_count for session in session_factory.sessions] == [1] + assert upload_files[0].tenant_id == "tenant-1" + assert upload_files[0].key == saves[0][0] + assert upload_files[0].used is True - def test_extract_xlsx_keeps_rows_with_only_embedded_images(self, monkeypatch: pytest.MonkeyPatch): + def test_extract_xlsx_keeps_rows_with_only_embedded_images( + self, + monkeypatch: pytest.MonkeyPatch, + database_session_maker: sessionmaker[Session], + ): image_bytes = b"\x89PNG\r\n\x1a\nimage-only-row" sheet = _FakeSheet( header_rows=[("Question", "Answer", "Image")], @@ -212,7 +175,7 @@ class TestExcelExtractor: ) workbook = _FakeWorkbook({"Data": sheet}) monkeypatch.setattr(excel_module, "load_workbook", lambda *args, **kwargs: workbook) - saves, session_factory = _patch_image_persistence(monkeypatch) + saves = _patch_image_persistence(monkeypatch) extractor = ExcelExtractor( "/tmp/sample.xlsx", @@ -221,17 +184,21 @@ class TestExcelExtractor: source_file_id="source-file-1", ) docs = extractor.extract() + upload_files = _get_upload_files(database_session_maker) assert workbook.closed is True assert len(docs) == 1 + assert len(upload_files) == 1 assert docs[0].page_content == ( - '"Question":"";"Answer":"";"Image":"![image](http://files.local/files/u1/file-preview)"' + f'"Question":"";"Answer":"";"Image":"![image](http://files.local/files/{upload_files[0].id}/file-preview)"' ) assert len(saves) == 1 - assert len(session_factory.persisted) == 1 - assert [session.commit_count for session in session_factory.sessions] == [1] - def test_extract_xlsx_reuses_existing_embedded_image_uploads_on_retry(self, monkeypatch: pytest.MonkeyPatch): + def test_extract_xlsx_reuses_existing_embedded_image_uploads_on_retry( + self, + monkeypatch: pytest.MonkeyPatch, + database_session_maker: sessionmaker[Session], + ): image_bytes = b"\x89PNG\r\n\x1a\nretry-safe-image" workbooks = [ _FakeWorkbook( @@ -254,7 +221,7 @@ class TestExcelExtractor: ), ] monkeypatch.setattr(excel_module, "load_workbook", lambda *args, **kwargs: workbooks.pop(0)) - saves, session_factory = _patch_image_persistence(monkeypatch) + saves = _patch_image_persistence(monkeypatch) extractor = ExcelExtractor( "/tmp/sample.xlsx", @@ -264,16 +231,17 @@ class TestExcelExtractor: ) first_docs = extractor.extract() second_docs = extractor.extract() + upload_files = _get_upload_files(database_session_maker) + assert len(upload_files) == 1 expected_page_content = ( - '"Question":"Q1";"Answer":"A1";"Image":"![image](http://files.local/files/u1/file-preview)"' + '"Question":"Q1";"Answer":"A1";' + f'"Image":"![image](http://files.local/files/{upload_files[0].id}/file-preview)"' ) assert first_docs[0].page_content == expected_page_content assert second_docs[0].page_content == expected_page_content assert len(saves) == 1 - assert len(session_factory.persisted) == 1 - assert [session.commit_count for session in session_factory.sessions] == [1, 0] def test_extract_xls_path(self, monkeypatch: pytest.MonkeyPatch): class FakeExcelFile: diff --git a/api/tests/unit_tests/core/rag/extractor/test_notion_extractor.py b/api/tests/unit_tests/core/rag/extractor/test_notion_extractor.py index f5d3187828d..8f49647fe7c 100644 --- a/api/tests/unit_tests/core/rag/extractor/test_notion_extractor.py +++ b/api/tests/unit_tests/core/rag/extractor/test_notion_extractor.py @@ -207,6 +207,26 @@ class TestNotionDatabase: mocker.patch("httpx.post", return_value=_mock_response({"results": None})) assert extractor._get_notion_database_data("db-1") == [] + def test_requests_use_bounded_timeout(self, mocker: MockerFixture): + """All outbound Notion API calls must carry a bounded timeout so a hanging endpoint cannot block extraction.""" + extractor = notion_extractor.NotionExtractor( + notion_workspace_id="ws", + notion_obj_id="obj", + notion_page_type="database", + tenant_id="tenant", + notion_access_token="token", + ) + + mock_post = mocker.patch("httpx.post", return_value=_mock_response({"results": None})) + extractor._get_notion_database_data("db-1") + assert mock_post.call_args.kwargs["timeout"] == notion_extractor._REQUEST_TIMEOUT + + mock_request = mocker.patch( + "httpx.request", return_value=_mock_response({"last_edited_time": "2024-01-01T00:00:00.000Z"}) + ) + extractor.get_notion_last_edited_time() + assert mock_request.call_args.kwargs["timeout"] == notion_extractor._REQUEST_TIMEOUT + def test_get_notion_database_data_requires_access_token(self): extractor = notion_extractor.NotionExtractor( notion_workspace_id="ws", diff --git a/api/tests/unit_tests/core/rag/extractor/test_pdf_extractor.py b/api/tests/unit_tests/core/rag/extractor/test_pdf_extractor.py index c41a0752b48..3f5cf0d37cb 100644 --- a/api/tests/unit_tests/core/rag/extractor/test_pdf_extractor.py +++ b/api/tests/unit_tests/core/rag/extractor/test_pdf_extractor.py @@ -1,83 +1,68 @@ -from types import SimpleNamespace +from dataclasses import dataclass from unittest.mock import MagicMock, patch +from uuid import uuid4 import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session import core.rag.extractor.pdf_extractor as pe +from models.model import UploadFile + +TENANT_ID = str(uuid4()) +USER_ID = str(uuid4()) + + +class _Storage: + saves: list[tuple[str, bytes]] + + def __init__(self) -> None: + self.saves = [] + + def save(self, key: str, data: bytes) -> None: + self.saves.append((key, data)) + + +class _DatabaseBinding: + session: Session + + def __init__(self, session: Session) -> None: + self.session = session + + +@dataclass(frozen=True) +class _Dependencies: + storage: _Storage + session: Session @pytest.fixture -def mock_dependencies(monkeypatch: pytest.MonkeyPatch): - # Mock storage - saves = [] - - def save(key, data): - saves.append((key, data)) - - monkeypatch.setattr(pe, "storage", SimpleNamespace(save=save)) - - # Mock db - class DummySession: - def __init__(self): - self.added = [] - self.committed = False - - def add(self, obj): - self.added.append(obj) - - def add_all(self, objs): - self.added.extend(objs) - - def commit(self): - self.committed = True - - db_stub = SimpleNamespace(session=DummySession()) - monkeypatch.setattr(pe, "db", db_stub) - - # Mock UploadFile - class FakeUploadFile: - DEFAULT_ID = "test_file_id" - - def __init__(self, **kwargs): - # Assign id from DEFAULT_ID, allow override via kwargs if needed - self.id = self.DEFAULT_ID - for k, v in kwargs.items(): - setattr(self, k, v) - - monkeypatch.setattr(pe, "UploadFile", FakeUploadFile) - - # Mock config +def mock_dependencies(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> _Dependencies: + storage = _Storage() + monkeypatch.setattr(pe, "storage", storage) + monkeypatch.setattr(pe, "db", _DatabaseBinding(sqlite_session)) monkeypatch.setattr(pe.dify_config, "FILES_URL", "http://files.local") monkeypatch.setattr(pe.dify_config, "INTERNAL_FILES_URL", None) monkeypatch.setattr(pe.dify_config, "STORAGE_TYPE", "local") - - return SimpleNamespace(saves=saves, db=db_stub, UploadFile=FakeUploadFile) + return _Dependencies(storage=storage, session=sqlite_session) @pytest.mark.parametrize( - ("image_bytes", "expected_mime", "expected_ext", "file_id"), + ("image_bytes", "expected_mime", "expected_ext"), [ - (b"\xff\xd8\xff some jpeg", "image/jpeg", "jpg", "test_file_id_jpeg"), - (b"\x89PNG\r\n\x1a\n some png", "image/png", "png", "test_file_id_png"), + (b"\xff\xd8\xff some jpeg", "image/jpeg", "jpg"), + (b"\x89PNG\r\n\x1a\n some png", "image/png", "png"), ], ) +@pytest.mark.parametrize("sqlite_session", [(UploadFile,)], indirect=True) @pytest.mark.parametrize("inject_session", [False, True]) def test_extract_images_formats( - mock_dependencies, - monkeypatch: pytest.MonkeyPatch, - image_bytes, - expected_mime, - expected_ext, - file_id, + mock_dependencies: _Dependencies, + image_bytes: bytes, + expected_mime: str, + expected_ext: str, inject_session: bool, ): - saves = mock_dependencies.saves - db_stub = mock_dependencies.db - - # Customize FakeUploadFile id for this test case. - # Using monkeypatch ensures the class attribute is reset between parameter sets. - monkeypatch.setattr(mock_dependencies.UploadFile, "DEFAULT_ID", file_id) - # Mock page and image objects mock_page = MagicMock() mock_image_obj = MagicMock() @@ -91,25 +76,32 @@ def test_extract_images_formats( extractor = pe.PdfExtractor( file_path="test.pdf", - tenant_id="t1", - user_id="u1", - session=db_stub.session if inject_session else None, + tenant_id=TENANT_ID, + user_id=USER_ID, + session=mock_dependencies.session if inject_session else None, ) # We need to handle the import inside _extract_images - with patch("pypdfium2.raw", autospec=True) as mock_raw: + with ( + patch("pypdfium2.raw", autospec=True) as mock_raw, + patch.object( + mock_dependencies.session, + "commit", + wraps=mock_dependencies.session.commit, + ) as commit, + ): mock_raw.FPDF_PAGEOBJ_IMAGE = 1 result = extractor._extract_images(mock_page) - assert f"![image](http://files.local/files/{file_id}/file-preview)" in result - assert len(saves) == 1 - assert saves[0][1] == image_bytes - assert len(db_stub.session.added) == 1 - assert db_stub.session.added[0].tenant_id == "t1" - assert db_stub.session.added[0].size == len(image_bytes) - assert db_stub.session.added[0].mime_type == expected_mime - assert db_stub.session.added[0].extension == expected_ext - assert db_stub.session.committed is not inject_session + assert commit.called is not inject_session + upload_file = mock_dependencies.session.scalar(select(UploadFile)) + assert upload_file is not None + assert f"![image](http://files.local/files/{upload_file.id}/file-preview)" in result + assert mock_dependencies.storage.saves == [(upload_file.key, image_bytes)] + assert upload_file.tenant_id == TENANT_ID + assert upload_file.size == len(image_bytes) + assert upload_file.mime_type == expected_mime + assert upload_file.extension == expected_ext @pytest.mark.parametrize( @@ -120,14 +112,17 @@ def test_extract_images_formats( (Exception("Failed to get objects"), None), # Exception raised ], ) -def test_extract_images_get_objects_scenarios(mock_dependencies, get_objects_side_effect, get_objects_return_value): +@pytest.mark.parametrize("sqlite_session", [(UploadFile,)], indirect=True) +def test_extract_images_get_objects_scenarios( + mock_dependencies: _Dependencies, get_objects_side_effect, get_objects_return_value +): mock_page = MagicMock() if get_objects_side_effect: mock_page.get_objects.side_effect = get_objects_side_effect else: mock_page.get_objects.return_value = get_objects_return_value - extractor = pe.PdfExtractor(file_path="test.pdf", tenant_id="t1", user_id="u1") + extractor = pe.PdfExtractor(file_path="test.pdf", tenant_id=TENANT_ID, user_id=USER_ID) with patch("pypdfium2.raw", autospec=True) as mock_raw: mock_raw.FPDF_PAGEOBJ_IMAGE = 1 @@ -136,7 +131,8 @@ def test_extract_images_get_objects_scenarios(mock_dependencies, get_objects_sid assert result == "" -def test_extract_calls_extract_images(mock_dependencies, monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite_session", [(UploadFile,)], indirect=True) +def test_extract_calls_extract_images(mock_dependencies: _Dependencies, monkeypatch: pytest.MonkeyPatch): # Mock pypdfium2 mock_pdf_doc = MagicMock() mock_page = MagicMock() @@ -152,7 +148,7 @@ def test_extract_calls_extract_images(mock_dependencies, monkeypatch: pytest.Mon mock_blob = MagicMock() mock_blob.source = "test.pdf" with patch("core.rag.extractor.pdf_extractor.Blob.from_path", return_value=mock_blob, autospec=True): - extractor = pe.PdfExtractor(file_path="test.pdf", tenant_id="t1", user_id="u1") + extractor = pe.PdfExtractor(file_path="test.pdf", tenant_id=TENANT_ID, user_id=USER_ID) # Mock _extract_images to return a known string monkeypatch.setattr(extractor, "_extract_images", lambda p: "![image](img_url)") @@ -165,10 +161,8 @@ def test_extract_calls_extract_images(mock_dependencies, monkeypatch: pytest.Mon assert documents[0].metadata["page"] == 0 -def test_extract_images_failures(mock_dependencies): - saves = mock_dependencies.saves - db_stub = mock_dependencies.db - +@pytest.mark.parametrize("sqlite_session", [(UploadFile,)], indirect=True) +def test_extract_images_failures(mock_dependencies: _Dependencies): # Mock page and image objects mock_page = MagicMock() mock_image_obj_fail = MagicMock() @@ -187,14 +181,14 @@ def test_extract_images_failures(mock_dependencies): mock_page.get_objects.return_value = [mock_image_obj_fail, mock_image_obj_ok] - extractor = pe.PdfExtractor(file_path="test.pdf", tenant_id="t1", user_id="u1") + extractor = pe.PdfExtractor(file_path="test.pdf", tenant_id=TENANT_ID, user_id=USER_ID) with patch("pypdfium2.raw", autospec=True) as mock_raw: mock_raw.FPDF_PAGEOBJ_IMAGE = 1 result = extractor._extract_images(mock_page) # Should have one success - assert "![image](http://files.local/files/test_file_id/file-preview)" in result - assert len(saves) == 1 - assert saves[0][1] == jpeg_bytes - assert db_stub.session.committed is True + upload_file = mock_dependencies.session.scalar(select(UploadFile)) + assert upload_file is not None + assert f"![image](http://files.local/files/{upload_file.id}/file-preview)" in result + assert mock_dependencies.storage.saves == [(upload_file.key, jpeg_bytes)] diff --git a/api/tests/unit_tests/core/rag/extractor/test_word_extractor.py b/api/tests/unit_tests/core/rag/extractor/test_word_extractor.py index 00c27504887..bd08a14c8cb 100644 --- a/api/tests/unit_tests/core/rag/extractor/test_word_extractor.py +++ b/api/tests/unit_tests/core/rag/extractor/test_word_extractor.py @@ -15,9 +15,12 @@ import pytest from docx import Document from docx.oxml import OxmlElement from docx.oxml.ns import qn +from sqlalchemy import event, select +from sqlalchemy.orm import Session import core.rag.extractor.word_extractor as we from core.rag.extractor.word_extractor import WordExtractor +from models.model import UploadFile class _TextOxmlElement(Protocol): @@ -112,7 +115,7 @@ def test_init_downloads_via_remote_fetcher(monkeypatch: pytest.MonkeyPatch): @pytest.mark.parametrize("inject_session", [False, True]) -def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_session: bool): +def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_session: bool, sqlite_session: Session): external_bytes = b"ext-bytes" internal_bytes = b"int-bytes" @@ -124,35 +127,13 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_sessio monkeypatch.setattr(we, "storage", SimpleNamespace(save=save)) - # Patch db.session to record adds/commit - class DummySession: - def __init__(self): - self.added = [] - self.committed = False - - def add_all(self, objects): - self.added.extend(objects) - - def commit(self): - self.committed = True - - db_stub = SimpleNamespace(session=DummySession()) + db_stub = SimpleNamespace(session=sqlite_session) monkeypatch.setattr(we, "db", db_stub) # Patch config values used for URL composition and storage type monkeypatch.setattr(we.dify_config, "FILES_URL", "http://files.local", raising=False) monkeypatch.setattr(we.dify_config, "STORAGE_TYPE", "local", raising=False) - # Patch UploadFile to avoid real DB models - class FakeUploadFile: - _i = 0 - - def __init__(self, **kwargs): # kwargs match the real signature fields - type(self)._i += 1 - self.id = f"u{self._i}" - - monkeypatch.setattr(we, "UploadFile", FakeUploadFile) - # Patch external image fetcher def fake_make_request(method: str, url: str, **kwargs): assert method == "GET" @@ -176,9 +157,11 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_sessio doc = SimpleNamespace(part=SimpleNamespace(rels={"rId1": rel_ext, "rId2": rel_int})) extractor = object.__new__(WordExtractor) - extractor.tenant_id = "t1" - extractor.user_id = "u1" + extractor.tenant_id = "00000000-0000-0000-0000-000000000001" + extractor.user_id = "00000000-0000-0000-0000-000000000002" extractor._session = db_stub.session if inject_session else None + transaction_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: transaction_events.append("commit")) image_map = extractor._extract_images_from_docx(doc) @@ -191,12 +174,13 @@ def test_extract_images_from_docx(monkeypatch: pytest.MonkeyPatch, inject_sessio assert external_bytes in payloads assert internal_bytes in payloads - # DB interactions should be recorded - assert len(db_stub.session.added) == 2 - assert db_stub.session.committed is not inject_session + assert len(sqlite_session.scalars(select(UploadFile)).all()) == 2 + assert transaction_events == ([] if inject_session else ["commit"]) -def test_extract_images_does_not_stage_partial_files_on_storage_failure(monkeypatch: pytest.MonkeyPatch): +def test_extract_images_does_not_stage_partial_files_on_storage_failure( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): class HashablePart: def __init__(self, blob: bytes): self.blob = blob @@ -222,22 +206,20 @@ def test_extract_images_does_not_stage_partial_files_on_storage_failure(monkeypa } ) ) - session = MagicMock() save = MagicMock(side_effect=[None, RuntimeError("storage failure")]) monkeypatch.setattr(we, "storage", SimpleNamespace(save=save)) monkeypatch.setattr(we.dify_config, "FILES_URL", "http://files.local", raising=False) monkeypatch.setattr(we.dify_config, "STORAGE_TYPE", "local", raising=False) extractor = object.__new__(WordExtractor) - extractor.tenant_id = "tenant" - extractor.user_id = "user" - extractor._session = session + extractor.tenant_id = "00000000-0000-0000-0000-000000000001" + extractor.user_id = "00000000-0000-0000-0000-000000000002" + extractor._session = sqlite_session with pytest.raises(RuntimeError, match="storage failure"): extractor._extract_images_from_docx(doc) - session.add_all.assert_not_called() - session.commit.assert_not_called() + assert sqlite_session.scalars(select(UploadFile)).all() == [] def test_extract_images_from_docx_uses_internal_files_url(): diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py index 6f9b0d0cc74..228aaf1a531 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_paragraph_index_processor.py @@ -1,21 +1,36 @@ import logging from contextlib import nullcontext +from datetime import datetime from types import SimpleNamespace from typing import Any from unittest.mock import Mock, patch import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker from core.entities.knowledge_entities import PreviewDetail from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.index_processor.processor.paragraph_index_processor import ParagraphIndexProcessor from core.rag.models.document import AttachmentDocument, Document +from extensions.storage.storage_type import StorageType from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, ImagePromptMessageContent from graphon.model_runtime.entities.model_entities import ModelFeature +from models.dataset import DocumentSegment, SegmentAttachmentBinding +from models.enums import CreatorUserRole +from models.model import UploadFile class TestParagraphIndexProcessor: + session: Session + session_factory: sessionmaker[Session] + + @pytest.fixture(autouse=True) + def _inject_sqlite_sessions(self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session]) -> None: + self.session = sqlite_session + self.session_factory = sqlite_session_factory + @pytest.fixture def processor(self) -> ParagraphIndexProcessor: return ParagraphIndexProcessor() @@ -54,9 +69,34 @@ class TestParagraphIndexProcessor: usage=LLMUsage.empty_usage(), ) + def _upload_file( + self, + *, + file_id: str, + tenant_id: str = "tenant-1", + name: str = "image.png", + extension: str = "png", + mime_type: str = "image/png", + ) -> UploadFile: + upload_file = UploadFile( + tenant_id=tenant_id, + storage_type=StorageType.LOCAL, + key=f"key-{file_id}", + name=name, + size=1, + extension=extension, + mime_type=mime_type, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + created_at=datetime(2025, 1, 1), + used=True, + ) + upload_file.id = file_id + return upload_file + def test_extract_forwards_automatic_flag(self, processor: ParagraphIndexProcessor) -> None: extract_setting = Mock() - session = Mock() + session = self.session expected_docs = [Document(page_content="chunk", metadata={})] with patch( @@ -69,7 +109,7 @@ class TestParagraphIndexProcessor: mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True, session=session) def test_transform_validates_process_rule(self, processor: ParagraphIndexProcessor) -> None: - session = Mock() + session = self.session with pytest.raises(ValueError, match="No process rule found"): processor.transform([Document(page_content="text", metadata={})], process_rule=None, session=session) @@ -82,7 +122,7 @@ class TestParagraphIndexProcessor: self, processor: ParagraphIndexProcessor, process_rule: dict[str, Any] ) -> None: rules_without_segmentation = SimpleNamespace(segmentation=None) - session = Mock() + session = self.session with patch( "core.rag.index_processor.processor.paragraph_index_processor.Rule.model_validate", @@ -99,7 +139,7 @@ class TestParagraphIndexProcessor: self, processor: ParagraphIndexProcessor, process_rule: dict[str, Any] ) -> None: source_document = Document(page_content="source", metadata={"dataset_id": "dataset-1", "document_id": "doc-1"}) - session = Mock() + session = self.session splitter = Mock() splitter.split_documents.return_value = [ Document(page_content=".first", metadata={}), @@ -138,7 +178,7 @@ class TestParagraphIndexProcessor: def test_transform_automatic_mode_uses_default_rules(self, processor: ParagraphIndexProcessor) -> None: splitter = Mock() splitter.split_documents.return_value = [Document(page_content="text", metadata={})] - session = Mock() + session = self.session with ( patch( @@ -173,7 +213,7 @@ class TestParagraphIndexProcessor: ) -> None: docs = [Document(page_content="chunk", metadata={})] multimodal_docs = [AttachmentDocument(page_content="image", metadata={})] - session = Mock() + session = self.session with ( patch("core.rag.index_processor.processor.paragraph_index_processor.Vector") as mock_vector_cls, @@ -191,7 +231,7 @@ class TestParagraphIndexProcessor: ) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY docs = [Document(page_content="chunk", metadata={})] - session = Mock() + session = self.session keywords_list = [["k1"], ["k2"]] with patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls: @@ -204,7 +244,7 @@ class TestParagraphIndexProcessor: ) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY docs = [Document(page_content="chunk", metadata={})] - session = Mock() + session = self.session with patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls: processor.load(dataset, docs, session=session) @@ -212,10 +252,20 @@ class TestParagraphIndexProcessor: mock_keyword_cls.return_value.add_texts.assert_called_once_with(docs, session) def test_clean_deletes_summaries_and_vector(self, processor: ParagraphIndexProcessor, dataset: Mock) -> None: - scalars_result = Mock() - scalars_result.all.return_value = [SimpleNamespace(id="seg-1")] - session = Mock() - session.scalars.return_value = scalars_result + session = self.session + segment = DocumentSegment( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id="doc-1", + position=1, + content="segment", + word_count=1, + tokens=1, + created_by="user-1", + index_node_id="node-1", + ) + session.add(segment) + session.flush() with ( patch( @@ -226,14 +276,14 @@ class TestParagraphIndexProcessor: vector = mock_vector_cls.return_value processor.clean(dataset, ["node-1"], delete_summaries=True, session=session) - mock_summary.assert_called_once_with(dataset, ["seg-1"], session=session) + mock_summary.assert_called_once_with(dataset, [segment.id], session=session) vector.delete_by_ids.assert_called_once_with(["node-1"]) def test_clean_economy_deletes_summaries_and_keywords( self, processor: ParagraphIndexProcessor, dataset: Mock ) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY - session = Mock() + session = self.session with ( patch( @@ -248,7 +298,7 @@ class TestParagraphIndexProcessor: def test_clean_deletes_keywords_by_ids(self, processor: ParagraphIndexProcessor, dataset: Mock) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY - session = Mock() + session = self.session with patch("core.rag.index_processor.processor.paragraph_index_processor.Keyword") as mock_keyword_cls: processor.clean(dataset, ["node-2"], with_keywords=True, session=session) @@ -257,9 +307,9 @@ class TestParagraphIndexProcessor: def test_index_list_chunks_high_quality( self, processor: ParagraphIndexProcessor, dataset: Mock, dataset_document: Mock ) -> None: - session = Mock() + session = self.session phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") + event.listen(session, "after_commit", lambda _session: phase_events.append("commit")) with ( patch( "core.rag.index_processor.processor.paragraph_index_processor.helper.generate_text_hash", @@ -299,9 +349,9 @@ class TestParagraphIndexProcessor: self, processor: ParagraphIndexProcessor, dataset: Mock, dataset_document: Mock ) -> None: dataset.indexing_technique = IndexTechniqueType.ECONOMY - session = Mock() + session = self.session phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") + event.listen(session, "after_commit", lambda _session: phase_events.append("commit")) with ( patch( "core.rag.index_processor.processor.paragraph_index_processor.helper.generate_text_hash", @@ -334,8 +384,8 @@ class TestParagraphIndexProcessor: ) chunk_without_files = SimpleNamespace(content="content-2", files=None) structure = SimpleNamespace(general_chunks=[chunk_with_files, chunk_without_files]) - session = Mock() - account_session = Mock() + session = self.session + account_session = self.session_factory() with ( patch( @@ -374,7 +424,7 @@ class TestParagraphIndexProcessor: self, processor: ParagraphIndexProcessor, dataset: Mock, dataset_document: Mock ) -> None: structure = SimpleNamespace(general_chunks=[SimpleNamespace(content="content", files=None)]) - session = Mock() + session = self.session with ( patch( @@ -403,8 +453,8 @@ class TestParagraphIndexProcessor: def test_generate_summary_preview_success_and_failure(self, processor: ParagraphIndexProcessor) -> None: preview_items = [PreviewDetail(content="chunk-1"), PreviewDetail(content="chunk-2")] - session = Mock() - worker_sessions = [Mock(), Mock()] + session = self.session + worker_sessions = [self.session_factory(), self.session_factory()] with ( patch( @@ -440,7 +490,9 @@ class TestParagraphIndexProcessor: patch("flask.current_app", fake_current_app), patch.object(processor, "generate_summary", return_value=("summary", LLMUsage.empty_usage())), ): - result = processor.generate_summary_preview("tenant-1", preview_items, {"enable": True}, session=Mock()) + result = processor.generate_summary_preview( + "tenant-1", preview_items, {"enable": True}, session=self.session + ) assert result[0].summary == "summary" @@ -456,16 +508,16 @@ class TestParagraphIndexProcessor: patch("concurrent.futures.wait", side_effect=[(set(), {future}), (set(), set())]), ): with pytest.raises(ValueError, match="timeout"): - processor.generate_summary_preview("tenant-1", preview_items, {"enable": True}, session=Mock()) + processor.generate_summary_preview("tenant-1", preview_items, {"enable": True}, session=self.session) future.cancel.assert_called_once() def test_generate_summary_validates_input(self) -> None: with pytest.raises(ValueError, match="must be enabled"): - ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": False}, session=Mock()) + ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": False}, session=self.session) with pytest.raises(ValueError, match="model_name and model_provider_name"): - ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": True}, session=Mock()) + ParagraphIndexProcessor.generate_summary("tenant-1", "text", {"enable": True}, session=self.session) def test_generate_summary_text_only_flow(self, caplog: pytest.LogCaptureFixture) -> None: model_instance = Mock() @@ -495,7 +547,7 @@ class TestParagraphIndexProcessor: "text content", {"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"}, document_language="English", - session=Mock(), + session=self.session, ) assert summary == "text summary" @@ -537,7 +589,7 @@ class TestParagraphIndexProcessor: "text content", {"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"}, segment_id="seg-1", - session=Mock(), + session=self.session, ) assert summary == "vision summary" @@ -580,7 +632,7 @@ class TestParagraphIndexProcessor: "tenant-1", "text content", {"enable": True, "model_name": "model-a", "model_provider_name": "provider-a"}, - session=Mock(), + session=self.session, ) assert sum(1 for r in caplog.records if r.levelno == logging.WARNING) == 1 @@ -594,30 +646,19 @@ class TestParagraphIndexProcessor: "![img2](/files/22222222-2222-2222-2222-222222222222/file-preview) " "![tool](/files/tools/33333333-3333-3333-3333-333333333333.png)" ) - image_upload = SimpleNamespace( - id="11111111-1111-1111-1111-111111111111", - tenant_id="tenant-1", - name="image.png", - mime_type="image/png", - extension="png", - source_url="", - size=1, - key="key", + session = self.session + session.add_all( + [ + self._upload_file(file_id="11111111-1111-1111-1111-111111111111"), + self._upload_file( + file_id="22222222-2222-2222-2222-222222222222", + name="file.txt", + extension="txt", + mime_type="text/plain", + ), + ] ) - non_image_upload = SimpleNamespace( - id="22222222-2222-2222-2222-222222222222", - tenant_id="tenant-1", - name="file.txt", - mime_type="text/plain", - extension="txt", - source_url="", - size=1, - key="key", - ) - scalars_result = Mock() - scalars_result.all.return_value = [image_upload, non_image_upload] - session = Mock() - session.scalars.return_value = scalars_result + session.flush() with ( patch( @@ -633,28 +674,14 @@ class TestParagraphIndexProcessor: assert not any(record.levelno == logging.WARNING for record in caplog.records) def test_extract_images_from_text_returns_empty_when_no_matches(self) -> None: - scalars_result = Mock() - scalars_result.all.return_value = [] - session = Mock() - session.scalars.return_value = scalars_result + session = self.session assert ParagraphIndexProcessor._extract_images_from_text("tenant-1", "no images here", session) == [] def test_extract_images_from_text_logs_when_build_fails(self, caplog: pytest.LogCaptureFixture) -> None: text = "![img](/files/11111111-1111-1111-1111-111111111111/image-preview)" - image_upload = SimpleNamespace( - id="11111111-1111-1111-1111-111111111111", - tenant_id="tenant-1", - name="image.png", - mime_type="image/png", - extension="png", - source_url="", - size=1, - key="key", - ) - scalars_result = Mock() - scalars_result.all.return_value = [image_upload] - session = Mock() - session.scalars.return_value = scalars_result + session = self.session + session.add(self._upload_file(file_id="11111111-1111-1111-1111-111111111111")) + session.flush() with ( patch( @@ -669,49 +696,39 @@ class TestParagraphIndexProcessor: assert sum(1 for r in caplog.records if r.levelno == logging.WARNING) == 1 def test_extract_images_from_segment_attachments(self, caplog: pytest.LogCaptureFixture) -> None: - image_upload = SimpleNamespace( - id="file-1", - name="image", - extension="png", - mime_type="image/png", - source_url="", - size=1, - key="k1", - ) - bad_upload = SimpleNamespace( - id="file-2", - name="broken", - extension=None, - mime_type="image/png", - source_url="", - size=1, - key="k2", - ) - non_image_upload = SimpleNamespace( - id="file-3", - name="text", - extension="txt", - mime_type="text/plain", - source_url="", - size=1, - key="k3", - ) - execute_result = Mock() - execute_result.all.return_value = [(None, image_upload), (None, bad_upload), (None, non_image_upload)] - session = Mock() - session.execute.return_value = execute_result + session = self.session + uploads = [ + self._upload_file(file_id="file-1", name="image"), + self._upload_file(file_id="file-2", name="broken"), + self._upload_file(file_id="file-3", name="text", extension="txt", mime_type="text/plain"), + ] + bindings = [ + SegmentAttachmentBinding( + tenant_id="tenant-1", + dataset_id="dataset-1", + document_id="doc-1", + segment_id="seg-1", + attachment_id=upload.id, + ) + for upload in uploads + ] + session.add_all([*uploads, *bindings]) + session.flush() - with caplog.at_level(logging.WARNING, logger="core.rag.index_processor.processor.paragraph_index_processor"): + with ( + patch( + "core.rag.index_processor.processor.paragraph_index_processor.File", + side_effect=[SimpleNamespace(id="file-1"), RuntimeError("bad file")], + ), + caplog.at_level(logging.WARNING, logger="core.rag.index_processor.processor.paragraph_index_processor"), + ): files = ParagraphIndexProcessor._extract_images_from_segment_attachments("tenant-1", "seg-1", session) assert len(files) == 1 assert sum(1 for r in caplog.records if r.levelno == logging.WARNING) == 1 def test_extract_images_from_segment_attachments_empty(self) -> None: - execute_result = Mock() - execute_result.all.return_value = [] - session = Mock() - session.execute.return_value = execute_result + session = self.session empty_files = ParagraphIndexProcessor._extract_images_from_segment_attachments("tenant-1", "seg-1", session) diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py index a2a5c67972d..1a7e54e35c1 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_parent_child_index_processor.py @@ -3,15 +3,26 @@ from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker from core.entities.knowledge_entities import PreviewDetail from core.rag.entities import ParentMode, Rule, Segmentation from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.index_processor.processor.parent_child_index_processor import ParentChildIndexProcessor from core.rag.models.document import AttachmentDocument, ChildDocument, Document +from models.dataset import ChildChunk, DatasetProcessRule, DocumentSegment class TestParentChildIndexProcessor: + session: Session + session_factory: sessionmaker[Session] + + @pytest.fixture(autouse=True) + def _inject_sqlite_sessions(self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session]) -> None: + self.session = sqlite_session + self.session_factory = sqlite_session_factory + @pytest.fixture def processor(self) -> ParentChildIndexProcessor: return ParentChildIndexProcessor() @@ -50,7 +61,7 @@ class TestParentChildIndexProcessor: def test_extract_forwards_automatic_flag(self, processor: ParentChildIndexProcessor) -> None: extract_setting = Mock() - session = Mock() + session = self.session expected = [Document(page_content="chunk", metadata={})] with patch( @@ -63,7 +74,7 @@ class TestParentChildIndexProcessor: mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True, session=session) def test_transform_validates_process_rule(self, processor: ParentChildIndexProcessor) -> None: - session = MagicMock() + session = self.session with pytest.raises(ValueError, match="No process rule found"): processor.transform([Document(page_content="text", metadata={})], process_rule=None, session=session) @@ -82,7 +93,7 @@ class TestParentChildIndexProcessor: processor.transform( [Document(page_content="text", metadata={})], process_rule={"mode": "custom", "rules": {"enabled": True}}, - session=MagicMock(), + session=self.session, ) def test_transform_paragraph_builds_parent_and_child_docs(self, processor: ParentChildIndexProcessor) -> None: @@ -117,7 +128,7 @@ class TestParentChildIndexProcessor: [parent_document], process_rule={"mode": "custom", "rules": {"enabled": True}}, preview=False, - session=MagicMock(), + session=self.session, ) assert len(result) == 1 @@ -154,7 +165,7 @@ class TestParentChildIndexProcessor: documents, process_rule={"mode": "custom", "rules": {"enabled": True}}, preview=True, - session=MagicMock(), + session=self.session, ) assert len(result) == 10 @@ -188,7 +199,7 @@ class TestParentChildIndexProcessor: docs, process_rule={"mode": "hierarchical", "rules": {"enabled": True}}, preview=True, - session=MagicMock(), + session=self.session, ) assert len(result) == 1 @@ -205,7 +216,7 @@ class TestParentChildIndexProcessor: ], ) multimodal_docs = [AttachmentDocument(page_content="image", metadata={})] - session = MagicMock() + session = self.session with patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls: vector = mock_vector_cls.return_value @@ -219,7 +230,7 @@ class TestParentChildIndexProcessor: vector.create_multimodal.assert_called_once_with(multimodal_docs) def test_clean_with_precomputed_child_ids(self, processor: ParentChildIndexProcessor, dataset: Mock) -> None: - session = MagicMock() + session = self.session with ( patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls, @@ -234,16 +245,42 @@ class TestParentChildIndexProcessor: ) vector.delete_by_ids.assert_called_once_with(["child-1", "child-2"]) - session.execute.assert_called() - session.flush.assert_called_once() + assert session.query(ChildChunk).count() == 0 def test_clean_queries_child_ids_when_not_precomputed( self, processor: ParentChildIndexProcessor, dataset: Mock ) -> None: - execute_result = Mock() - execute_result.all.return_value = [("child-1",), (None,), ("child-2",)] - session = MagicMock() - session.execute.return_value = execute_result + session = self.session + parent = DocumentSegment( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id="doc-1", + position=1, + content="parent", + word_count=1, + tokens=1, + created_by="user-1", + index_node_id="node-1", + ) + session.add(parent) + session.flush() + session.add_all( + [ + ChildChunk( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id="doc-1", + segment_id=parent.id, + position=index, + content=f"child-{index}", + word_count=1, + created_by="user-1", + index_node_id=f"child-{index}", + ) + for index in (1, 2) + ] + ) + session.flush() with ( patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls, @@ -254,7 +291,7 @@ class TestParentChildIndexProcessor: vector.delete_by_ids.assert_called_once_with(["child-1", "child-2"]) def test_clean_dataset_wide_cleanup(self, processor: ParentChildIndexProcessor, dataset: Mock) -> None: - session = MagicMock() + session = self.session with ( patch("core.rag.index_processor.processor.parent_child_index_processor.Vector") as mock_vector_cls, @@ -263,17 +300,23 @@ class TestParentChildIndexProcessor: processor.clean(dataset, None, delete_child_chunks=True, session=session) vector.delete.assert_called_once() - session.execute.assert_called() - session.flush.assert_called_once() + assert session.query(ChildChunk).count() == 0 def test_clean_deletes_summaries_when_requested(self, processor: ParentChildIndexProcessor, dataset: Mock) -> None: - scalars_result = Mock() - scalars_result.all.return_value = [SimpleNamespace(id="seg-1")] - session = MagicMock() - session.scalars.return_value = scalars_result - session_ctx = MagicMock() - session_ctx.__enter__.return_value = session - session_ctx.__exit__.return_value = False + session = self.session + segment = DocumentSegment( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id="doc-1", + position=1, + content="parent", + word_count=1, + tokens=1, + created_by="user-1", + index_node_id="node-1", + ) + session.add(segment) + session.flush() with ( patch( @@ -283,7 +326,7 @@ class TestParentChildIndexProcessor: ): processor.clean(dataset, ["node-1"], delete_summaries=True, precomputed_child_node_ids=[], session=session) - mock_summary.assert_called_once_with(dataset, ["seg-1"], session=session) + mock_summary.assert_called_once_with(dataset, [segment.id], session=session) def test_clean_deletes_all_summaries_when_node_ids_missing( self, processor: ParentChildIndexProcessor, dataset: Mock @@ -294,7 +337,7 @@ class TestParentChildIndexProcessor: ) as mock_summary, patch("core.rag.index_processor.processor.parent_child_index_processor.Vector"), ): - session = MagicMock() + session = self.session processor.clean(dataset, None, delete_summaries=True, session=session) mock_summary.assert_called_once_with(dataset, None, session=session) @@ -341,20 +384,15 @@ class TestParentChildIndexProcessor: ) ], ) - dataset_rule = SimpleNamespace(id="rule-1") - session = MagicMock() + session = self.session phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") + event.listen(session, "after_commit", lambda _session: phase_events.append("commit")) with ( patch( "core.rag.index_processor.processor.parent_child_index_processor.ParentChildStructureChunk.model_validate", return_value=parent_childs, ), - patch( - "core.rag.index_processor.processor.parent_child_index_processor.DatasetProcessRule", - return_value=dataset_rule, - ), patch( "core.rag.index_processor.processor.parent_child_index_processor.helper.generate_text_hash", side_effect=lambda text: f"hash-{text}", @@ -373,9 +411,8 @@ class TestParentChildIndexProcessor: processor.index(dataset, dataset_document, {"parent_child_chunks": []}, session) assert phase_events == ["count", "store", "commit", "vector"] - assert dataset_document.dataset_process_rule_id == "rule-1" - session.add.assert_called_once_with(dataset_rule) - session.flush.assert_called_once() + dataset_rule = session.get(DatasetProcessRule, dataset_document.dataset_process_rule_id) + assert dataset_rule is not None documents = mock_token_counter.call_args.kwargs["documents"] assert [document.page_content for document in documents] == ["parent text"] mock_token_counter.assert_called_once_with(dataset=dataset, documents=documents) @@ -396,19 +433,14 @@ class TestParentChildIndexProcessor: parent_mode=ParentMode.PARAGRAPH, parent_child_chunks=[SimpleNamespace(parent_content="parent", child_contents=["child"], files=None)], ) - dataset_rule = SimpleNamespace(id="rule-1") - session = MagicMock() - account_session = MagicMock() + session = self.session + account_session = self.session_factory() with ( patch( "core.rag.index_processor.processor.parent_child_index_processor.ParentChildStructureChunk.model_validate", return_value=parent_childs, ), - patch( - "core.rag.index_processor.processor.parent_child_index_processor.DatasetProcessRule", - return_value=dataset_rule, - ), patch( "core.rag.index_processor.processor.parent_child_index_processor.helper.generate_text_hash", return_value="hash", @@ -480,8 +512,8 @@ class TestParentChildIndexProcessor: def test_generate_summary_preview_sets_summaries(self, processor: ParentChildIndexProcessor) -> None: preview_texts = [PreviewDetail(content="chunk-1"), PreviewDetail(content="chunk-2")] - session = MagicMock() - worker_sessions = [MagicMock(), MagicMock()] + session = self.session + worker_sessions = [self.session_factory(), self.session_factory()] with ( patch( @@ -513,7 +545,7 @@ class TestParentChildIndexProcessor: side_effect=RuntimeError("summary failed"), ): with pytest.raises(ValueError, match="Failed to generate summaries"): - processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}, session=MagicMock()) + processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}, session=self.session) def test_generate_summary_preview_falls_back_without_flask_context( self, processor: ParentChildIndexProcessor @@ -529,7 +561,7 @@ class TestParentChildIndexProcessor: ), ): result = processor.generate_summary_preview( - "tenant-1", preview_texts, {"enable": True}, session=MagicMock() + "tenant-1", preview_texts, {"enable": True}, session=self.session ) assert result[0].summary == "summary" @@ -546,6 +578,6 @@ class TestParentChildIndexProcessor: patch("concurrent.futures.wait", side_effect=[(set(), {future}), (set(), set())]), ): with pytest.raises(ValueError, match="timeout"): - processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}, session=MagicMock()) + processor.generate_summary_preview("tenant-1", preview_texts, {"enable": True}, session=self.session) future.cancel.assert_called_once() diff --git a/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py b/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py index ec390fcc578..3bcc1c59a17 100644 --- a/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/processor/test_qa_index_processor.py @@ -1,16 +1,20 @@ import logging from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch import pandas as pd import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session from werkzeug.datastructures import FileStorage from core.entities.knowledge_entities import PreviewDetail from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.index_processor.processor.qa_index_processor import QAIndexProcessor from core.rag.models.document import AttachmentDocument, Document +from models.dataset import Dataset, DocumentSegment +from models.dataset import Document as DatasetDocument class _ImmediateThread: @@ -32,20 +36,19 @@ class TestQAIndexProcessor: return QAIndexProcessor() @pytest.fixture - def dataset(self) -> Mock: - dataset = Mock() - dataset.id = "dataset-1" - dataset.tenant_id = "tenant-1" - dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY - dataset.is_multimodal = True - return dataset + def dataset(self) -> Dataset: + return Dataset( + id="dataset-1", + tenant_id="tenant-1", + name="QA Dataset", + created_by="user-1", + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + is_multimodal=True, + ) @pytest.fixture - def dataset_document(self) -> Mock: - document = Mock() - document.id = "doc-1" - document.created_by = "user-1" - return document + def dataset_document(self) -> DatasetDocument: + return DatasetDocument(id="doc-1", created_by="user-1") @pytest.fixture def process_rule(self) -> dict: @@ -58,9 +61,9 @@ class TestQAIndexProcessor: segmentation = SimpleNamespace(max_tokens=256, chunk_overlap=10, separator="\n") return SimpleNamespace(segmentation=segmentation) - def test_extract_forwards_automatic_flag(self, processor: QAIndexProcessor) -> None: + def test_extract_forwards_automatic_flag(self, processor: QAIndexProcessor, sqlite_session: Session) -> None: extract_setting = Mock() - session = Mock() + session = sqlite_session expected_docs = [Document(page_content="chunk", metadata={})] with patch("core.rag.index_processor.processor.qa_index_processor.ExtractProcessor.extract") as mock_extract: @@ -71,22 +74,22 @@ class TestQAIndexProcessor: assert docs == expected_docs mock_extract.assert_called_once_with(extract_setting=extract_setting, is_automatic=True, session=session) - def test_transform_rejects_none_process_rule(self, processor: QAIndexProcessor) -> None: - session = MagicMock() + def test_transform_rejects_none_process_rule(self, processor: QAIndexProcessor, sqlite_session: Session) -> None: + session = sqlite_session with pytest.raises(ValueError, match="No process rule found"): processor.transform([Document(page_content="text", metadata={})], process_rule=None, session=session) - def test_transform_rejects_missing_rules_key(self, processor: QAIndexProcessor) -> None: - session = MagicMock() + def test_transform_rejects_missing_rules_key(self, processor: QAIndexProcessor, sqlite_session: Session) -> None: + session = sqlite_session with pytest.raises(ValueError, match="No rules found in process rule"): processor.transform( [Document(page_content="text", metadata={})], process_rule={"mode": "custom"}, session=session ) def test_transform_preview_calls_formatter_once( - self, processor: QAIndexProcessor, process_rule: dict[str, Any], fake_flask_app + self, processor: QAIndexProcessor, process_rule: dict[str, Any], fake_flask_app, sqlite_session: Session ) -> None: - session = MagicMock() + session = sqlite_session document = Document(page_content="raw text", metadata={"dataset_id": "dataset-1", "document_id": "doc-1"}) split_node = Document(page_content=".question", metadata={}) splitter = Mock() @@ -128,9 +131,9 @@ class TestQAIndexProcessor: mock_format.assert_called_once() def test_transform_non_preview_uses_thread_batches( - self, processor: QAIndexProcessor, process_rule: dict[str, Any], fake_flask_app + self, processor: QAIndexProcessor, process_rule: dict[str, Any], fake_flask_app, sqlite_session: Session ) -> None: - session = MagicMock() + session = sqlite_session documents = [ Document(page_content="doc-1", metadata={"document_id": "doc-1", "dataset_id": "dataset-1"}), Document(page_content="doc-2", metadata={"document_id": "doc-2", "dataset_id": "dataset-1"}), @@ -212,8 +215,10 @@ class TestQAIndexProcessor: with pytest.raises(ValueError, match="bad csv"): processor.format_by_template(csv_file) - def test_load_creates_vectors_for_high_quality_dataset(self, processor: QAIndexProcessor, dataset: Mock) -> None: - session = MagicMock() + def test_load_creates_vectors_for_high_quality_dataset( + self, processor: QAIndexProcessor, dataset: Dataset, sqlite_session: Session + ) -> None: + session = sqlite_session docs = [Document(page_content="Q1", metadata={"answer": "A1"})] multimodal_docs = [AttachmentDocument(page_content="image", metadata={})] @@ -225,8 +230,10 @@ class TestQAIndexProcessor: vector.create.assert_called_once_with(docs) vector.create_multimodal.assert_called_once_with(multimodal_docs) - def test_load_skips_vector_for_non_high_quality(self, processor: QAIndexProcessor, dataset: Mock) -> None: - session = MagicMock() + def test_load_skips_vector_for_non_high_quality( + self, processor: QAIndexProcessor, dataset: Dataset, sqlite_session: Session + ) -> None: + session = sqlite_session dataset.indexing_technique = IndexTechniqueType.ECONOMY docs = [Document(page_content="Q1", metadata={"answer": "A1"})] @@ -236,13 +243,46 @@ class TestQAIndexProcessor: mock_vector_cls.assert_not_called() def test_clean_handles_summary_deletion_and_vector_cleanup( - self, processor: QAIndexProcessor, dataset: Mock + self, + processor: QAIndexProcessor, + dataset: Dataset, + sqlite_session: Session, ) -> None: - mock_segment = SimpleNamespace(id="seg-1") - scalars_result = Mock() - scalars_result.all.return_value = [mock_segment] - mock_session = MagicMock() - mock_session.scalars.return_value = scalars_result + matching_segment = DocumentSegment( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id="doc-1", + position=1, + content="Q1", + word_count=1, + tokens=1, + created_by="user-1", + index_node_id="node-1", + ) + other_node_segment = DocumentSegment( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id="doc-1", + position=2, + content="Q2", + word_count=1, + tokens=1, + created_by="user-1", + index_node_id="node-2", + ) + other_dataset_segment = DocumentSegment( + tenant_id="tenant-2", + dataset_id="dataset-2", + document_id="doc-2", + position=1, + content="Q3", + word_count=1, + tokens=1, + created_by="user-2", + index_node_id="node-1", + ) + sqlite_session.add_all([matching_segment, other_node_segment, other_dataset_segment]) + sqlite_session.commit() with ( patch( @@ -251,13 +291,15 @@ class TestQAIndexProcessor: patch("core.rag.index_processor.processor.qa_index_processor.Vector") as mock_vector_cls, ): vector = mock_vector_cls.return_value - processor.clean(dataset, ["node-1"], delete_summaries=True, session=mock_session) + processor.clean(dataset, ["node-1"], delete_summaries=True, session=sqlite_session) - mock_summary.assert_called_once_with(dataset, ["seg-1"], session=mock_session) + mock_summary.assert_called_once_with(dataset, [matching_segment.id], session=sqlite_session) vector.delete_by_ids.assert_called_once_with(["node-1"]) - def test_clean_handles_dataset_wide_cleanup(self, processor: QAIndexProcessor, dataset: Mock) -> None: - session = MagicMock() + def test_clean_handles_dataset_wide_cleanup( + self, processor: QAIndexProcessor, dataset: Dataset, sqlite_session: Session + ) -> None: + session = sqlite_session with ( patch( "core.rag.index_processor.processor.qa_index_processor.SummaryIndexService.delete_summaries_for_segments" @@ -271,11 +313,15 @@ class TestQAIndexProcessor: vector.delete.assert_called_once() def test_index_adds_documents_and_vectors_for_high_quality( - self, processor: QAIndexProcessor, dataset: Mock, dataset_document: Mock + self, + processor: QAIndexProcessor, + dataset: Dataset, + dataset_document: DatasetDocument, + sqlite_session: Session, ) -> None: - session = MagicMock() + session = sqlite_session phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") + event.listen(session, "after_commit", lambda _session: phase_events.append("commit")) qa_chunks = SimpleNamespace( qa_chunks=[ SimpleNamespace(question="Q1", answer="A1"), @@ -315,9 +361,13 @@ class TestQAIndexProcessor: mock_vector_cls.return_value.create.assert_called_once() def test_index_requires_high_quality( - self, processor: QAIndexProcessor, dataset: Mock, dataset_document: Mock + self, + processor: QAIndexProcessor, + dataset: Dataset, + dataset_document: DatasetDocument, + sqlite_session: Session, ) -> None: - session = MagicMock() + session = sqlite_session dataset.indexing_technique = IndexTechniqueType.ECONOMY qa_chunks = SimpleNamespace(qa_chunks=[SimpleNamespace(question="Q1", answer="A1")]) @@ -351,10 +401,10 @@ class TestQAIndexProcessor: assert preview["total_segments"] == 1 assert preview["qa_preview"] == [{"question": "Q1", "answer": "A1"}] - def test_generate_summary_preview_returns_input(self, processor: QAIndexProcessor) -> None: + def test_generate_summary_preview_returns_input(self, processor: QAIndexProcessor, sqlite_session: Session) -> None: preview_items = [PreviewDetail(content="Q1")] assert ( - processor.generate_summary_preview("tenant-1", preview_items, {"enable": False}, session=MagicMock()) + processor.generate_summary_preview("tenant-1", preview_items, {"enable": False}, session=sqlite_session) is preview_items ) diff --git a/api/tests/unit_tests/core/rag/indexing/test_index_processor.py b/api/tests/unit_tests/core/rag/indexing/test_index_processor.py index ed943fe1129..51d817d8f53 100644 --- a/api/tests/unit_tests/core/rag/indexing/test_index_processor.py +++ b/api/tests/unit_tests/core/rag/indexing/test_index_processor.py @@ -1,12 +1,51 @@ -import datetime +import uuid from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import MagicMock, patch +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker + from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.index_processor.index_processor import IndexProcessor from core.workflow.nodes.knowledge_index.protocols import Preview, PreviewItem -from models.dataset import Dataset, Document +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, SegmentStatus + + +def _persist_dataset_and_document( + session: Session, + *, + indexing_technique: IndexTechniqueType = IndexTechniqueType.HIGH_QUALITY, + summary_index_setting: dict | None = None, +) -> tuple[Dataset, Document]: + tenant_id = str(uuid.uuid4()) + created_by = str(uuid.uuid4()) + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + indexing_technique=indexing_technique, + chunk_structure="text_model", + summary_index_setting=summary_index_setting, + created_by=created_by, + ) + document = Document( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + dataset_id=dataset.id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="Document", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, + doc_language="English", + ) + session.add_all([dataset, document]) + session.flush() + return dataset, document class TestIndexProcessor: @@ -22,161 +61,163 @@ class TestIndexProcessor: assert preview.qa_preview[0].question == "Q1" assert preview.qa_preview[0].answer == "A1" - def test_index_and_clean_ends_transactions_around_index_io(self) -> None: - document = SimpleNamespace( - id="document-1", - name="Document", - created_at=datetime.datetime(2026, 1, 1), - indexing_latency=None, - indexing_status=None, - completed_at=None, - word_count=0, - need_summary=False, - ) - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - name="Dataset", - chunk_structure="text_model", - summary_index_setting=None, - ) - session = MagicMock() - session.scalar.side_effect = [dataset, document, 3] + def test_index_and_clean_ends_transactions_around_index_io(self, sqlite_session: Session) -> None: + dataset, document = _persist_dataset_and_document(sqlite_session) phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") + event.listen(sqlite_session, "after_commit", lambda _session: phase_events.append("commit")) index_processor = MagicMock() index_processor.index.side_effect = lambda *args: phase_events.append("index") + processor = IndexProcessor() + admission_service = MagicMock() + chunks = {"general_chunks": ["content"]} - with patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory: + with ( + patch( + "core.rag.index_processor.index_processor.VectorSpaceAdmissionService", + return_value=admission_service, + ), + patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory, + ): index_processor_factory.return_value.init_index_processor.return_value = index_processor - IndexProcessor().index_and_clean( + processor.index_and_clean( dataset_id=dataset.id, document_id=document.id, original_document_id="", - chunks={"general_chunks": ["content"]}, + chunks=chunks, batch="batch-1", - session=session, + session=sqlite_session, ) assert phase_events == ["commit", "index", "commit"] - - def test_index_and_clean_scopes_replacement_queries_to_dataset_owner(self) -> None: - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - name="Dataset", - summary_index_setting=None, - chunk_structure="text_model", + admission_service.ensure_pipeline_can_be_indexed.assert_called_once_with( + dataset=dataset, + document_id=document.id, + chunk_structure=dataset.chunk_structure, + chunks=chunks, + include_summaries=False, + session=sqlite_session, ) - document = SimpleNamespace( - id="doc-1", - tenant_id="tenant-1", - dataset_id="dataset-1", - name="Document", - created_at=datetime.datetime(2026, 1, 1, tzinfo=datetime.UTC), - ) - segment = SimpleNamespace(index_node_id="node-1") - session = MagicMock() - def resolve_owner(statement): - entity = statement.column_descriptions[0]["entity"] - if entity is Dataset: - return dataset - if entity is Document: - return document - return 3 + def test_index_and_clean_skips_admission_for_replacement_without_existing_vector_points( + self, sqlite_session: Session + ) -> None: + dataset, document = _persist_dataset_and_document(sqlite_session) + index_processor = MagicMock() + processor = IndexProcessor() + chunks = {"general_chunks": ["content"]} - session.scalar.side_effect = resolve_owner - session.scalars.return_value.all.return_value = [segment] - - with patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory: - index_backend = index_processor_factory.return_value.init_index_processor.return_value - IndexProcessor().index_and_clean( - dataset_id="dataset-1", - document_id="doc-1", - original_document_id="original-doc", - chunks={}, + with ( + patch("core.rag.index_processor.index_processor.VectorSpaceAdmissionService") as admission_service_class, + patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory, + ): + index_processor_factory.return_value.init_index_processor.return_value = index_processor + processor.index_and_clean( + dataset_id=dataset.id, + document_id=document.id, + original_document_id=document.id, + chunks=chunks, batch="batch-1", - session=session, + session=sqlite_session, ) - document_statement = next( - call.args[0] - for call in session.scalar.call_args_list - if call.args[0].column_descriptions[0]["entity"] is Document - ) - segment_statement = session.scalars.call_args.args[0] - delete_statement = session.execute.call_args_list[0].args[0] - word_count_statement = session.scalar.call_args_list[-1].args[0] - segment_update_statement = session.execute.call_args_list[1].args[0] - document_owner = {"doc-1", "dataset-1", "tenant-1"} - original_document_owner = {"original-doc", "dataset-1", "tenant-1"} + admission_service_class.assert_not_called() - assert document_owner <= set(document_statement.compile().params.values()) - assert original_document_owner <= set(segment_statement.compile().params.values()) - assert original_document_owner <= set(delete_statement.compile().params.values()) - assert document_owner <= set(word_count_statement.compile().params.values()) - assert document_owner <= set(segment_update_statement.compile().params.values()) + def test_index_and_clean_scopes_replacement_queries_to_dataset_owner(self, sqlite_session: Session) -> None: + dataset, document = _persist_dataset_and_document(sqlite_session) + original_document = Document( + id=str(uuid.uuid4()), + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + position=2, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="Original document", + created_from=DocumentCreatedFrom.WEB, + created_by=dataset.created_by, + ) + segment = DocumentSegment( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id=original_document.id, + position=1, + content="old content", + word_count=3, + tokens=3, + created_by=dataset.created_by, + index_node_id="node-1", + status=SegmentStatus.COMPLETED, + ) + control_segment = DocumentSegment( + tenant_id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + document_id=original_document.id, + position=1, + content="other tenant content", + word_count=4, + tokens=4, + created_by=str(uuid.uuid4()), + index_node_id="other-node", + status=SegmentStatus.COMPLETED, + ) + sqlite_session.add_all([original_document, segment, control_segment]) + sqlite_session.flush() + + processor = IndexProcessor() + with ( + patch("core.rag.index_processor.index_processor.VectorSpaceAdmissionService") as admission_service_class, + patch("core.rag.index_processor.index_processor.IndexProcessorFactory") as index_processor_factory, + ): + index_backend = index_processor_factory.return_value.init_index_processor.return_value + processor.index_and_clean( + dataset_id=dataset.id, + document_id=document.id, + original_document_id=original_document.id, + chunks={}, + batch="batch-1", + session=sqlite_session, + ) + + assert sqlite_session.get(DocumentSegment, segment.id) is None + assert sqlite_session.get(DocumentSegment, control_segment.id) is control_segment + assert sqlite_session.get(Document, document.id).indexing_status == "completed" index_backend.clean.assert_called_once_with( dataset, ["node-1"], with_keywords=True, delete_child_chunks=True, - session=session, + session=sqlite_session, ) - index_backend.index.assert_called_once_with(dataset, document, {}, session) + index_backend.index.assert_called_once_with(dataset, document, {}, sqlite_session) + admission_service_class.assert_not_called() - def test_get_preview_output_scopes_document_to_dataset_owner(self) -> None: - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - indexing_technique=IndexTechniqueType.ECONOMY, - summary_index_setting=None, - ) - document = SimpleNamespace(doc_language="English") - session = MagicMock() - - def resolve_owner(statement): - entity = statement.column_descriptions[0]["entity"] - if entity is Dataset: - return dataset - if entity is Document: - return document - raise AssertionError(f"Unexpected entity: {entity}") - - session.scalar.side_effect = resolve_owner + def test_get_preview_output_scopes_document_to_dataset_owner(self, sqlite_session: Session) -> None: + dataset, document = _persist_dataset_and_document(sqlite_session, indexing_technique=IndexTechniqueType.ECONOMY) processor = IndexProcessor() expected_preview = MagicMock() with patch.object(processor, "format_preview", return_value=expected_preview): result = processor.get_preview_output( chunks={}, - dataset_id="dataset-1", - document_id="doc-1", + dataset_id=dataset.id, + document_id=document.id, chunk_structure="text_model", summary_index_setting=None, - session=session, + session=sqlite_session, ) - document_statement = next( - call.args[0] - for call in session.scalar.call_args_list - if call.args[0].column_descriptions[0]["entity"] is Document - ) - assert {"doc-1", "dataset-1", "tenant-1"} <= set(document_statement.compile().params.values()) assert result is expected_preview - def test_preview_summary_workers_use_independent_sessions(self) -> None: - caller_session = MagicMock() - phase_events: list[str] = [] - caller_session.commit.side_effect = lambda: phase_events.append("commit") - caller_session.scalar.return_value = SimpleNamespace( - indexing_technique=IndexTechniqueType.HIGH_QUALITY, + def test_preview_summary_workers_use_independent_sessions( + self, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] + ) -> None: + dataset, _ = _persist_dataset_and_document( + sqlite_session, summary_index_setting={"enable": True}, - tenant_id="tenant-1", ) - worker_sessions = [MagicMock(), MagicMock()] + phase_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: phase_events.append("commit")) + worker_sessions: list[Session] = [] preview = Preview( chunk_structure="text_model", total_segments=2, @@ -184,11 +225,11 @@ class TestIndexProcessor: ) flask_app = SimpleNamespace(app_context=lambda: nullcontext()) processor = IndexProcessor() - worker_contexts = iter(nullcontext(worker_session) for worker_session in worker_sessions) - def create_worker_session(): + def generate_summary(*_args, **kwargs): phase_events.append("worker") - return next(worker_contexts) + worker_sessions.append(kwargs["session"]) + return "summary", None with ( patch.object(processor, "format_preview", return_value=preview), @@ -197,27 +238,25 @@ class TestIndexProcessor: SimpleNamespace(_get_current_object=lambda: flask_app), ), patch( - "core.rag.index_processor.index_processor.session_factory.create_session", - side_effect=create_worker_session, + "core.rag.index_processor.index_processor.session_factory", + SimpleNamespace(create_session=sqlite_session_factory), ), patch( "core.rag.index_processor.index_processor.ParagraphIndexProcessor.generate_summary", - return_value=("summary", None), + side_effect=generate_summary, ) as generate_summary, ): result = processor.get_preview_output( chunks=[], - dataset_id="dataset-1", + dataset_id=dataset.id, document_id="", chunk_structure="text_model", summary_index_setting={"enable": True}, - session=caller_session, + session=sqlite_session, ) assert all(item.summary == "summary" for item in result.preview) assert phase_events == ["commit", "worker", "worker"] call_sessions = [call.kwargs["session"] for call in generate_summary.call_args_list] - assert all(call_session is not caller_session for call_session in call_sessions) - assert all( - any(call_session is worker_session for worker_session in worker_sessions) for call_session in call_sessions - ) + assert all(call_session is not sqlite_session for call_session in call_sessions) + assert call_sessions == worker_sessions diff --git a/api/tests/unit_tests/core/rag/indexing/test_index_processor_base.py b/api/tests/unit_tests/core/rag/indexing/test_index_processor_base.py index 93c039e2ff5..ad8fb37ea67 100644 --- a/api/tests/unit_tests/core/rag/indexing/test_index_processor_base.py +++ b/api/tests/unit_tests/core/rag/indexing/test_index_processor_base.py @@ -1,14 +1,40 @@ +from datetime import UTC, datetime from types import SimpleNamespace from typing import override from unittest.mock import Mock, patch +from uuid import uuid4 import httpx import pytest +from sqlalchemy.orm import Session from core.entities.knowledge_entities import PreviewDetail from core.rag.index_processor.constant.doc_type import DocType from core.rag.index_processor.index_processor_base import BaseIndexProcessor from core.rag.models.document import AttachmentDocument, Document +from extensions.storage.storage_type import StorageType +from models.enums import CreatorUserRole +from models.model import UploadFile +from models.tools import ToolFile + + +def _persist_upload(session: Session, *, upload_id: str, name: str) -> UploadFile: + upload = UploadFile( + tenant_id=str(uuid4()), + storage_type=StorageType.LOCAL, + key=f"uploads/{name}", + name=name, + size=4, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + created_at=datetime.now(UTC), + used=True, + ) + upload.id = upload_id + session.add(upload) + return upload class _ForwardingBaseIndexProcessor(BaseIndexProcessor): @@ -68,21 +94,23 @@ class TestBaseIndexProcessor: def processor(self) -> _ForwardingBaseIndexProcessor: return _ForwardingBaseIndexProcessor() - def test_abstract_methods_raise_not_implemented(self, processor: _ForwardingBaseIndexProcessor) -> None: + def test_abstract_methods_raise_not_implemented( + self, processor: _ForwardingBaseIndexProcessor, unbound_session: Session + ) -> None: with pytest.raises(NotImplementedError): - processor.extract(Mock(), session=Mock()) + processor.extract(Mock(), session=unbound_session) with pytest.raises(NotImplementedError): - processor.transform([], session=Mock()) + processor.transform([], session=unbound_session) with pytest.raises(NotImplementedError): processor.generate_summary_preview( - "tenant", [PreviewDetail(content="c")], {"enable": False}, session=Mock() + "tenant", [PreviewDetail(content="c")], {"enable": False}, session=unbound_session ) with pytest.raises(NotImplementedError): - processor.load(Mock(), [], session=Mock()) + processor.load(Mock(), [], session=unbound_session) with pytest.raises(NotImplementedError): - processor.clean(Mock(), None, session=Mock()) + processor.clean(Mock(), None, session=unbound_session) with pytest.raises(NotImplementedError): - processor.index(Mock(), Mock(), {}, Mock()) + processor.index(Mock(), Mock(), {}, unbound_session) with pytest.raises(NotImplementedError): processor.format_preview([]) @@ -123,12 +151,14 @@ class TestBaseIndexProcessor: images = processor._extract_markdown_images(markdown) assert images == ["https://a/img.png", "/files/123/file-preview"] - def test_get_content_files_without_images_returns_empty(self, processor: _ForwardingBaseIndexProcessor) -> None: + def test_get_content_files_without_images_returns_empty( + self, processor: _ForwardingBaseIndexProcessor, unbound_session: Session + ) -> None: document = Document(page_content="no image markdown", metadata={"document_id": "doc-1", "dataset_id": "ds-1"}) - assert processor._get_content_files(document, session=Mock()) == [] + assert processor._get_content_files(document, session=unbound_session) == [] def test_get_content_files_handles_all_sources_and_duplicates( - self, processor: _ForwardingBaseIndexProcessor + self, processor: _ForwardingBaseIndexProcessor, sqlite_session: Session ) -> None: document = Document(page_content="ignored", metadata={"document_id": "doc-1", "dataset_id": "ds-1"}) images = [ @@ -138,22 +168,19 @@ class TestBaseIndexProcessor: "/files/tools/cccccccc-cccc-cccc-cccc-cccccccccccc.png", "https://example.com/remote.png?x=1", ] - upload_a = SimpleNamespace(id="aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa", name="a.png") - upload_b = SimpleNamespace(id="bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb", name="b.png") - upload_tool = SimpleNamespace(id="tool-upload-id", name="tool.png") - upload_remote = SimpleNamespace(id="remote-upload-id", name="remote.png") - scalars_result = Mock() - scalars_result.all.return_value = [upload_a, upload_b, upload_tool, upload_remote] - db_session = Mock() - db_session.scalars.return_value = scalars_result + _persist_upload(sqlite_session, upload_id="aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa", name="a.png") + _persist_upload(sqlite_session, upload_id="bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb", name="b.png") + tool_upload = _persist_upload(sqlite_session, upload_id=str(uuid4()), name="tool.png") + remote_upload = _persist_upload(sqlite_session, upload_id=str(uuid4()), name="remote.png") + sqlite_session.commit() current_user = Mock() with ( patch.object(processor, "_extract_markdown_images", return_value=images), - patch.object(processor, "_download_tool_file", return_value="tool-upload-id") as mock_tool_download, - patch.object(processor, "_download_image", return_value="remote-upload-id") as mock_image_download, + patch.object(processor, "_download_tool_file", return_value=tool_upload.id) as mock_tool_download, + patch.object(processor, "_download_image", return_value=remote_upload.id) as mock_image_download, ): - files = processor._get_content_files(document, current_user=current_user, session=db_session) + files = processor._get_content_files(document, current_user=current_user, session=sqlite_session) assert len(files) == 5 assert all(isinstance(file, AttachmentDocument) for file in files) @@ -165,33 +192,30 @@ class TestBaseIndexProcessor: mock_tool_download.assert_called_once_with( "cccccccc-cccc-cccc-cccc-cccccccccccc", current_user, - session=db_session, + session=sqlite_session, ) mock_image_download.assert_called_once() def test_get_content_files_skips_tool_and_remote_download_without_user( - self, processor: _ForwardingBaseIndexProcessor + self, processor: _ForwardingBaseIndexProcessor, unbound_session: Session ) -> None: document = Document(page_content="ignored", metadata={"document_id": "doc-1", "dataset_id": "ds-1"}) images = ["/files/tools/cccccccc-cccc-cccc-cccc-cccccccccccc.png", "https://example.com/remote.png"] with patch.object(processor, "_extract_markdown_images", return_value=images): - files = processor._get_content_files(document, current_user=None, session=Mock()) + files = processor._get_content_files(document, current_user=None, session=unbound_session) assert files == [] - def test_get_content_files_ignores_missing_upload_records(self, processor: _ForwardingBaseIndexProcessor) -> None: + def test_get_content_files_ignores_missing_upload_records( + self, processor: _ForwardingBaseIndexProcessor, sqlite_session: Session + ) -> None: document = Document(page_content="ignored", metadata={"document_id": "doc-1", "dataset_id": "ds-1"}) images = ["/files/aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa/image-preview"] - scalars_result = Mock() - scalars_result.all.return_value = [] - db_session = Mock() - db_session.scalars.return_value = scalars_result - with ( patch.object(processor, "_extract_markdown_images", return_value=images), ): - files = processor._get_content_files(document, session=db_session) + files = processor._get_content_files(document, session=sqlite_session) assert files == [] @@ -270,16 +294,25 @@ class TestBaseIndexProcessor: ): assert processor._download_image("https://example.com/image.png", current_user=Mock()) is None - def test_download_tool_file_returns_none_when_not_found(self, processor: _ForwardingBaseIndexProcessor) -> None: - db_session = Mock() - db_session.get.return_value = None + def test_download_tool_file_returns_none_when_not_found( + self, processor: _ForwardingBaseIndexProcessor, sqlite_session: Session + ) -> None: + assert processor._download_tool_file(str(uuid4()), current_user=Mock(), session=sqlite_session) is None - assert processor._download_tool_file("tool-id", current_user=Mock(), session=db_session) is None - - def test_download_tool_file_uploads_file_when_found(self, processor: _ForwardingBaseIndexProcessor) -> None: - tool_file = SimpleNamespace(file_key="k1", name="tool.png", mimetype="image/png") - db_session = Mock() - db_session.get.return_value = tool_file + def test_download_tool_file_uploads_file_when_found( + self, processor: _ForwardingBaseIndexProcessor, sqlite_session: Session + ) -> None: + tool_file = ToolFile( + user_id=str(uuid4()), + tenant_id=str(uuid4()), + conversation_id=None, + file_key="k1", + mimetype="image/png", + name="tool.png", + size=4, + ) + sqlite_session.add(tool_file) + sqlite_session.commit() mock_db = Mock() mock_db.engine = Mock() upload_result = SimpleNamespace(id="upload-id") @@ -290,7 +323,7 @@ class TestBaseIndexProcessor: patch("services.file_service.FileService") as mock_file_service, ): mock_file_service.return_value.upload_file.return_value = upload_result - result = processor._download_tool_file("tool-id", current_user=Mock(), session=db_session) + result = processor._download_tool_file(tool_file.id, current_user=Mock(), session=sqlite_session) assert result == "upload-id" mock_load.assert_called_once_with("k1") diff --git a/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py b/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py index 5307da6d343..ecb2985dc4b 100644 --- a/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py +++ b/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py @@ -65,12 +65,14 @@ from core.indexing_runner import ( ) from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType from core.rag.models.document import ChildDocument, Document +from enums import DeploymentEdition from graphon.model_runtime.entities.model_entities import ModelType from libs.datetime_utils import naive_utc_now from models.dataset import Dataset, DatasetProcessRule, DocumentSegment from models.dataset import Document as DatasetDocument from models.enums import SegmentStatus from models.model import Account +from services.vector_space_admission_service import VectorSpaceAdmissionError # ============================================================================ # Helper Functions @@ -83,10 +85,10 @@ def create_mock_dataset( indexing_technique: str = IndexTechniqueType.HIGH_QUALITY, embedding_provider: str = "openai", embedding_model: str = "text-embedding-ada-002", -) -> Mock: - """Create a mock Dataset object with configurable parameters. +) -> Dataset: + """Create a Dataset object with configurable parameters. - This helper function creates a properly configured mock Dataset object that can be + This helper function creates a properly configured Dataset object that can be used across multiple tests, ensuring consistency in test data. Args: @@ -97,18 +99,19 @@ def create_mock_dataset( embedding_model: The embedding model name. Returns: - Mock: A configured mock Dataset object with all required attributes. + Dataset: A configured Dataset object with all required attributes. Example: >>> dataset = create_mock_dataset(indexing_technique="economy") >>> assert dataset.indexing_technique == "economy" """ - dataset = Mock(spec=Dataset) - dataset.id = dataset_id or str(uuid.uuid4()) - dataset.tenant_id = tenant_id or str(uuid.uuid4()) - dataset.indexing_technique = indexing_technique - dataset.embedding_model_provider = embedding_provider - dataset.embedding_model = embedding_model + dataset = Dataset( + id=dataset_id or str(uuid.uuid4()), + tenant_id=tenant_id or str(uuid.uuid4()), + indexing_technique=indexing_technique, + embedding_model_provider=embedding_provider, + embedding_model=embedding_model, + ) return dataset @@ -119,10 +122,10 @@ def create_mock_dataset_document( doc_form: str = IndexStructureType.PARAGRAPH_INDEX, data_source_type: str = "upload_file", doc_language: str = "English", -) -> Mock: - """Create a mock DatasetDocument object with configurable parameters. +) -> DatasetDocument: + """Create a DatasetDocument object with configurable parameters. - This helper function creates a properly configured mock DatasetDocument object, + This helper function creates a properly configured DatasetDocument object, reducing boilerplate code in individual tests. Args: @@ -134,23 +137,23 @@ def create_mock_dataset_document( doc_language: The document language. Returns: - Mock: A configured mock DatasetDocument object with all required attributes. + DatasetDocument: A configured DatasetDocument object with all required attributes. Example: >>> doc = create_mock_dataset_document(doc_form=IndexStructureType.QA_INDEX) >>> assert doc.doc_form == IndexStructureType.QA_INDEX """ - doc = Mock(spec=DatasetDocument) - doc.id = document_id or str(uuid.uuid4()) - doc.dataset_id = dataset_id or str(uuid.uuid4()) - doc.tenant_id = tenant_id or str(uuid.uuid4()) - doc.doc_form = doc_form - doc.doc_language = doc_language - doc.data_source_type = data_source_type - doc.data_source_info_dict = {"upload_file_id": str(uuid.uuid4())} - doc.dataset_process_rule_id = str(uuid.uuid4()) - doc.created_by = str(uuid.uuid4()) - return doc + return DatasetDocument( + id=document_id or str(uuid.uuid4()), + dataset_id=dataset_id or str(uuid.uuid4()), + tenant_id=tenant_id or str(uuid.uuid4()), + doc_form=doc_form, + doc_language=doc_language, + data_source_type=data_source_type, + data_source_info=json.dumps({"upload_file_id": str(uuid.uuid4())}), + dataset_process_rule_id=str(uuid.uuid4()), + created_by=str(uuid.uuid4()), + ) def create_sample_documents( @@ -275,14 +278,14 @@ class TestIndexingRunnerExtract: @pytest.fixture def sample_dataset_document(self): """Create a sample dataset document for testing.""" - doc = Mock(spec=DatasetDocument) - doc.id = str(uuid.uuid4()) - doc.dataset_id = str(uuid.uuid4()) - doc.tenant_id = str(uuid.uuid4()) - doc.doc_form = IndexStructureType.PARAGRAPH_INDEX - doc.data_source_type = "upload_file" - doc.data_source_info_dict = {"upload_file_id": str(uuid.uuid4())} - return doc + return DatasetDocument( + id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + doc_form=IndexStructureType.PARAGRAPH_INDEX, + data_source_type="upload_file", + data_source_info=json.dumps({"upload_file_id": str(uuid.uuid4())}), + ) @pytest.fixture def sample_process_rule(self): @@ -353,12 +356,14 @@ class TestIndexingRunnerExtract: # Arrange runner = IndexingRunner() sample_dataset_document.data_source_type = "notion_import" - sample_dataset_document.data_source_info_dict = { - "credential_id": str(uuid.uuid4()), - "notion_workspace_id": "workspace123", - "notion_page_id": "page123", - "type": "page", - } + sample_dataset_document.data_source_info = json.dumps( + { + "credential_id": str(uuid.uuid4()), + "notion_workspace_id": "workspace123", + "notion_page_id": "page123", + "type": "page", + } + ) mock_processor = MagicMock() mock_dependencies["factory"].return_value.init_index_processor.return_value = mock_processor @@ -383,13 +388,15 @@ class TestIndexingRunnerExtract: # Arrange runner = IndexingRunner() sample_dataset_document.data_source_type = "website_crawl" - sample_dataset_document.data_source_info_dict = { - "provider": "firecrawl", - "url": "https://example.com", - "job_id": "job123", - "mode": "crawl", - "only_main_content": True, - } + sample_dataset_document.data_source_info = json.dumps( + { + "provider": "firecrawl", + "url": "https://example.com", + "job_id": "job123", + "mode": "crawl", + "only_main_content": True, + } + ) mock_processor = MagicMock() mock_dependencies["factory"].return_value.init_index_processor.return_value = mock_processor @@ -415,7 +422,7 @@ class TestIndexingRunnerExtract: """Test extraction fails when upload file is missing.""" # Arrange runner = IndexingRunner() - sample_dataset_document.data_source_info_dict = {} + sample_dataset_document.data_source_info = "{}" mock_processor = MagicMock() mock_dependencies["factory"].return_value.init_index_processor.return_value = mock_processor @@ -466,12 +473,13 @@ class TestIndexingRunnerTransform: @pytest.fixture def sample_dataset(self): """Create a sample dataset for testing.""" - dataset = Mock(spec=Dataset) - dataset.id = str(uuid.uuid4()) - dataset.tenant_id = str(uuid.uuid4()) - dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY - dataset.embedding_model_provider = "openai" - dataset.embedding_model = "text-embedding-ada-002" + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + embedding_model_provider="openai", + embedding_model="text-embedding-ada-002", + ) return dataset @pytest.fixture @@ -637,22 +645,23 @@ class TestIndexingRunnerLoad: @pytest.fixture def sample_dataset(self): """Create a sample dataset for testing.""" - dataset = Mock(spec=Dataset) - dataset.id = str(uuid.uuid4()) - dataset.tenant_id = str(uuid.uuid4()) - dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY - dataset.embedding_model_provider = "openai" - dataset.embedding_model = "text-embedding-ada-002" + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + embedding_model_provider="openai", + embedding_model="text-embedding-ada-002", + ) return dataset @pytest.fixture def sample_dataset_document(self): """Create a sample dataset document for testing.""" - doc = Mock(spec=DatasetDocument) - doc.id = str(uuid.uuid4()) - doc.dataset_id = str(uuid.uuid4()) - doc.doc_form = IndexStructureType.PARAGRAPH_INDEX - return doc + return DatasetDocument( + id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) @pytest.fixture def sample_documents(self): @@ -842,16 +851,18 @@ class TestIndexingRunnerRun: """Create sample dataset documents for testing.""" docs = [] for i in range(2): - doc = Mock(spec=DatasetDocument) - doc.id = str(uuid.uuid4()) - doc.dataset_id = str(uuid.uuid4()) - doc.tenant_id = str(uuid.uuid4()) - doc.doc_form = IndexStructureType.PARAGRAPH_INDEX - doc.doc_language = "English" - doc.data_source_type = "upload_file" - doc.data_source_info_dict = {"upload_file_id": str(uuid.uuid4())} - doc.dataset_process_rule_id = str(uuid.uuid4()) - docs.append(doc) + docs.append( + DatasetDocument( + id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + doc_form=IndexStructureType.PARAGRAPH_INDEX, + doc_language="English", + data_source_type="upload_file", + data_source_info=json.dumps({"upload_file_id": str(uuid.uuid4())}), + dataset_process_rule_id=str(uuid.uuid4()), + ) + ) return docs def test_run_in_indexing_status_loads_child_chunks_with_caller_session( @@ -860,44 +871,63 @@ class TestIndexingRunnerRun: runner = IndexingRunner() dataset_document = sample_dataset_documents[0] dataset_document.doc_form = IndexStructureType.PARENT_CHILD_INDEX - dataset = Mock(spec=Dataset) - segment = Mock(spec=DocumentSegment) - segment.status = "waiting" - segment.content = "parent" - segment.index_node_id = "parent-node" - segment.index_node_hash = "parent-hash" - segment.document_id = dataset_document.id - segment.dataset_id = dataset_document.dataset_id - segment.tokens = 12 - segment.get_child_chunks.return_value = [ - SimpleNamespace(content="child", index_node_id="child-node", index_node_hash="child-hash") - ] + dataset = Dataset() + segment = DocumentSegment( + tenant_id="tenant-id", + dataset_id=dataset_document.dataset_id, + document_id=dataset_document.id, + position=1, + content="parent", + word_count=0, + tokens=12, + created_by="account-id", + status="waiting", + index_node_id="parent-node", + index_node_hash="parent-hash", + ) + child_chunks = [SimpleNamespace(content="child", index_node_id="child-node", index_node_hash="child-hash")] session = mock_dependencies["session"] session.get.side_effect = lambda model, _: dataset_document if model is DatasetDocument else dataset session.scalars.return_value.all.return_value = [segment] - with patch.object(runner, "_load") as load: + with ( + patch.object(DocumentSegment, "get_child_chunks", return_value=child_chunks) as get_child_chunks, + patch.object(runner, "_load") as load, + ): runner.run_in_indexing_status(dataset_document, session) - segment.get_child_chunks.assert_called_once_with(session=session) + get_child_chunks.assert_called_once_with(session=session) assert load.call_args.kwargs["documents"][0].children[0].page_content == "child" assert load.call_args.kwargs["total_tokens"] == 12 def test_run_in_indexing_status_uses_tokens_from_all_segments(self, mock_dependencies, sample_dataset_documents): runner = IndexingRunner() dataset_document = sample_dataset_documents[0] - dataset = Mock(spec=Dataset) - completed_segment = Mock(spec=DocumentSegment) - completed_segment.status = SegmentStatus.COMPLETED - completed_segment.tokens = 10 - incomplete_segment = Mock(spec=DocumentSegment) - incomplete_segment.status = SegmentStatus.WAITING - incomplete_segment.tokens = 20 - incomplete_segment.content = "pending" - incomplete_segment.index_node_id = "pending-node" - incomplete_segment.index_node_hash = "pending-hash" - incomplete_segment.document_id = dataset_document.id - incomplete_segment.dataset_id = dataset_document.dataset_id + dataset = Dataset() + completed_segment = DocumentSegment( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + position=1, + content="", + word_count=0, + tokens=10, + created_by="account-id", + status=SegmentStatus.COMPLETED, + ) + incomplete_segment = DocumentSegment( + tenant_id="tenant-id", + dataset_id=dataset_document.dataset_id, + document_id=dataset_document.id, + position=1, + content="pending", + word_count=0, + tokens=20, + created_by="account-id", + status=SegmentStatus.WAITING, + index_node_id="pending-node", + index_node_hash="pending-hash", + ) session = mock_dependencies["session"] session.get.side_effect = lambda model, _: dataset_document if model is DatasetDocument else dataset session.scalars.return_value.all.return_value = [completed_segment, incomplete_segment] @@ -908,25 +938,28 @@ class TestIndexingRunnerRun: assert load.call_args.kwargs["documents"][0].page_content == "pending" assert load.call_args.kwargs["total_tokens"] == 30 - def test_run_success_single_document(self, mock_dependencies, sample_dataset_documents): + @patch.object(Account, "set_tenant_id_with_session", autospec=True) + def test_run_success_single_document(self, set_tenant_id, mock_dependencies, sample_dataset_documents): """Test successful run with single document.""" # Arrange runner = IndexingRunner() doc = sample_dataset_documents[0] # Mock database queries - mock_dataset = Mock(spec=Dataset) - mock_dataset.id = doc.dataset_id - mock_dataset.tenant_id = doc.tenant_id - mock_dataset.indexing_technique = IndexTechniqueType.ECONOMY + mock_dataset = Dataset( + id=doc.dataset_id, + tenant_id=doc.tenant_id, + indexing_technique=IndexTechniqueType.ECONOMY, + ) - mock_current_user = Mock(spec=Account) + mock_current_user = Account(name="Test Account", email="test@example.com") get_dispatch = {"Document": doc, "Dataset": mock_dataset, "Account": mock_current_user} mock_dependencies["session"].get.side_effect = lambda model, id: get_dispatch.get(model.__name__) - mock_process_rule = Mock(spec=DatasetProcessRule) - mock_process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} + mock_process_rule = DatasetProcessRule( + dataset_id="dataset-id", mode="automatic", rules="{}", created_by="account-id" + ) mock_dependencies["session"].scalar.return_value = mock_process_rule # Mock processor @@ -963,7 +996,8 @@ class TestIndexingRunnerRun: # Assert - verify the methods were called # Since we're mocking the internal methods, we just verify no exceptions were raised - mock_current_user.set_tenant_id_with_session.assert_called_once_with( + set_tenant_id.assert_called_once_with( + mock_current_user, mock_dataset.tenant_id, session=mock_dependencies["session"], ) @@ -998,14 +1032,18 @@ class TestIndexingRunnerRun: with pytest.raises(DocumentIsPausedError): runner.run([doc], mock_dependencies["session"]) - def test_run_counts_each_transformed_document_once(self, mock_dependencies, sample_dataset_documents): + @patch.object(Account, "set_tenant_id_with_session", autospec=True) + def test_run_counts_each_transformed_document_once( + self, set_tenant_id, mock_dependencies, sample_dataset_documents + ): runner = IndexingRunner() dataset_document = sample_dataset_documents[0] - dataset = Mock(spec=Dataset) - dataset.id = dataset_document.dataset_id - dataset.tenant_id = dataset_document.tenant_id - dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY - current_user = Mock(spec=Account) + dataset = Dataset( + id=dataset_document.dataset_id, + tenant_id=dataset_document.tenant_id, + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + ) + current_user = Account(name="Test Account", email="test@example.com") transformed_documents = [ Document(page_content="first", metadata={"doc_id": "first", "doc_hash": "hash-first"}), Document(page_content="second", metadata={"doc_id": "second", "doc_hash": "hash-second"}), @@ -1016,8 +1054,9 @@ class TestIndexingRunnerRun: Account: current_user, } mock_dependencies["session"].get.side_effect = lambda model, _: model_dispatch.get(model) - process_rule = Mock(spec=DatasetProcessRule) - process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} + process_rule = DatasetProcessRule( + dataset_id="dataset-id", mode="automatic", rules="{}", created_by="account-id" + ) mock_dependencies["session"].scalar.return_value = process_rule with ( @@ -1041,18 +1080,84 @@ class TestIndexingRunnerRun: token_counts=[11, 22], ) assert load.call_args.kwargs["total_tokens"] == 33 + set_tenant_id.assert_called_once_with( + current_user, + dataset.tenant_id, + session=mock_dependencies["session"], + ) + @patch.object(Account, "set_tenant_id_with_session", autospec=True) + def test_run_rejects_before_segment_or_vector_writes( + self, set_tenant_id, mock_dependencies, sample_dataset_documents + ): + runner = IndexingRunner(enforce_vector_space_admission=True) + dataset_document = sample_dataset_documents[0] + dataset_document.need_summary = False + dataset = Dataset( + id=dataset_document.dataset_id, + tenant_id=dataset_document.tenant_id, + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + ) + current_user = Account(name="Test Account", email="test@example.com") + model_dispatch = { + DatasetDocument: dataset_document, + Dataset: dataset, + Account: current_user, + } + mock_dependencies["session"].get.side_effect = lambda model, _: model_dispatch.get(model) + process_rule = DatasetProcessRule( + dataset_id="dataset-id", mode="automatic", rules="{}", created_by="account-id" + ) + mock_dependencies["session"].scalar.return_value = process_rule + transformed_documents = [Document(page_content="Chunk", metadata={"doc_id": "c1", "doc_hash": "h1"})] + admission_error = VectorSpaceAdmissionError("estimated storage exceeds capacity") + admission_service = Mock() + admission_service.ensure_document_can_be_indexed.side_effect = admission_error + + with ( + patch("core.indexing_runner.VectorSpaceAdmissionService", return_value=admission_service), + patch.object(runner, "_extract", return_value=[Document(page_content="source", metadata={})]), + patch.object( + runner, + "_transform", + return_value=transformed_documents, + ), + patch.object(runner, "_load_segments") as load_segments, + patch.object(runner, "_load") as load, + patch.object(runner, "_handle_indexing_error") as handle_error, + ): + runner.run([dataset_document], mock_dependencies["session"]) + + load_segments.assert_not_called() + load.assert_not_called() + admission_service.ensure_document_can_be_indexed.assert_called_once_with( + dataset=dataset, + document_id=dataset_document.id, + doc_form=dataset_document.doc_form, + documents=transformed_documents, + include_summaries=False, + session=mock_dependencies["session"], + ) + handle_error.assert_called_once_with(dataset_document.id, admission_error, mock_dependencies["session"]) + set_tenant_id.assert_called_once_with( + current_user, + dataset.tenant_id, + session=mock_dependencies["session"], + ) + + @patch.object(Account, "set_tenant_id_with_session", autospec=True) def test_run_in_splitting_status_counts_each_transformed_document_once( - self, mock_dependencies, sample_dataset_documents + self, set_tenant_id, mock_dependencies, sample_dataset_documents ): runner = IndexingRunner() dataset_document = sample_dataset_documents[0] dataset_document.created_by = "user-1" - dataset = Mock(spec=Dataset) - dataset.id = dataset_document.dataset_id - dataset.tenant_id = dataset_document.tenant_id - dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY - current_user = Mock(spec=Account) + dataset = Dataset( + id=dataset_document.dataset_id, + tenant_id=dataset_document.tenant_id, + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + ) + current_user = Account(name="Test Account", email="test@example.com") transformed_documents = [ Document(page_content="first", metadata={"doc_id": "first", "doc_hash": "hash-first"}), Document(page_content="second", metadata={"doc_id": "second", "doc_hash": "hash-second"}), @@ -1064,8 +1169,9 @@ class TestIndexingRunnerRun: } mock_dependencies["session"].get.side_effect = lambda model, _: model_dispatch.get(model) mock_dependencies["session"].scalars.return_value.all.return_value = [] - process_rule = Mock(spec=DatasetProcessRule) - process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} + process_rule = DatasetProcessRule( + dataset_id="dataset-id", mode="automatic", rules="{}", created_by="account-id" + ) mock_dependencies["session"].scalar.return_value = process_rule with ( @@ -1089,6 +1195,11 @@ class TestIndexingRunnerRun: token_counts=[11, 22], ) assert load.call_args.kwargs["total_tokens"] == 33 + set_tenant_id.assert_called_once_with( + current_user, + dataset.tenant_id, + session=mock_dependencies["session"], + ) def test_run_handles_provider_token_error(self, mock_dependencies, sample_dataset_documents): """Test run handles ProviderTokenNotInitError and updates document status.""" @@ -1097,14 +1208,16 @@ class TestIndexingRunnerRun: doc = sample_dataset_documents[0] # Mock database - mock_dataset = Mock(spec=Dataset) - mock_dataset.tenant_id = doc.tenant_id + mock_dataset = Dataset( + tenant_id=doc.tenant_id, + ) get_dispatch = {"Document": doc, "Dataset": mock_dataset} mock_dependencies["session"].get.side_effect = lambda model, id: get_dispatch.get(model.__name__) - mock_process_rule = Mock(spec=DatasetProcessRule) - mock_process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} + mock_process_rule = DatasetProcessRule( + dataset_id="dataset-id", mode="automatic", rules="{}", created_by="account-id" + ) mock_dependencies["session"].scalar.return_value = mock_process_rule mock_processor = MagicMock() @@ -1126,14 +1239,16 @@ class TestIndexingRunnerRun: doc = sample_dataset_documents[0] # Mock database - mock_dataset = Mock(spec=Dataset) - mock_dataset.tenant_id = doc.tenant_id + mock_dataset = Dataset( + tenant_id=doc.tenant_id, + ) get_dispatch = {"Document": doc, "Dataset": mock_dataset} mock_dependencies["session"].get.side_effect = lambda model, id: get_dispatch.get(model.__name__) - mock_process_rule = Mock(spec=DatasetProcessRule) - mock_process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} + mock_process_rule = DatasetProcessRule( + dataset_id="dataset-id", mode="automatic", rules="{}", created_by="account-id" + ) mock_dependencies["session"].scalar.return_value = mock_process_rule mock_processor = MagicMock() @@ -1147,16 +1262,18 @@ class TestIndexingRunnerRun: # Assert - should not raise, just log warning # No exception should be raised - def test_run_processes_multiple_documents(self, mock_dependencies, sample_dataset_documents): + @patch.object(Account, "set_tenant_id_with_session", autospec=True) + def test_run_processes_multiple_documents(self, set_tenant_id, mock_dependencies, sample_dataset_documents): """Test run processes multiple documents sequentially.""" # Arrange runner = IndexingRunner() docs = sample_dataset_documents # Mock database - mock_dataset = Mock(spec=Dataset) - mock_dataset.indexing_technique = IndexTechniqueType.ECONOMY - mock_current_user = Mock(spec=Account) + mock_dataset = Dataset( + indexing_technique=IndexTechniqueType.ECONOMY, + ) + mock_current_user = Account(name="Test Account", email="test@example.com") doc_map = {doc.id: doc for doc in docs} model_dispatch = {"Dataset": mock_dataset, "Account": mock_current_user} @@ -1169,8 +1286,9 @@ class TestIndexingRunnerRun: mock_dependencies["session"].get.side_effect = get_side_effect - mock_process_rule = Mock(spec=DatasetProcessRule) - mock_process_rule.to_dict.return_value = {"mode": "automatic", "rules": {}} + mock_process_rule = DatasetProcessRule( + dataset_id="dataset-id", mode="automatic", rules="{}", created_by="account-id" + ) mock_dependencies["session"].scalar.return_value = mock_process_rule mock_processor = MagicMock() @@ -1197,7 +1315,7 @@ class TestIndexingRunnerRun: # Assert # Verify extract was called for each document assert mock_extract.call_count == len(docs) - assert mock_current_user.set_tenant_id_with_session.call_count == len(docs) + assert set_tenant_id.call_count == len(docs) class TestIndexingRunnerRetryLogic: @@ -1244,8 +1362,7 @@ class TestIndexingRunnerRetryLogic: """Test successful document status update.""" # Arrange document_id = str(uuid.uuid4()) - mock_document = Mock(spec=DatasetDocument) - mock_document.id = document_id + mock_document = DatasetDocument(id=document_id) mock_dependencies["session"].scalar.return_value = 0 mock_dependencies["session"].get.return_value = mock_document @@ -1296,23 +1413,29 @@ class TestIndexingRunnerDocumentCleaning: @pytest.fixture def sample_process_rule_automatic(self): """Create automatic processing rule.""" - rule = Mock(spec=DatasetProcessRule) - rule.mode = "automatic" - rule.rules = None + rule = DatasetProcessRule( + dataset_id="dataset-id", + mode="automatic", + rules=None, + created_by="account-id", + ) return rule @pytest.fixture def sample_process_rule_custom(self): """Create custom processing rule.""" - rule = Mock(spec=DatasetProcessRule) - rule.mode = "custom" - rule.rules = json.dumps( - { - "pre_processing_rules": [ - {"id": "remove_extra_spaces", "enabled": True}, - {"id": "remove_urls_emails", "enabled": True}, - ] - } + rule = DatasetProcessRule( + dataset_id="dataset-id", + mode="custom", + rules=json.dumps( + { + "pre_processing_rules": [ + {"id": "remove_extra_spaces", "enabled": True}, + {"id": "remove_urls_emails", "enabled": True}, + ] + } + ), + created_by="account-id", ) return rule @@ -1372,6 +1495,19 @@ class TestIndexingRunnerDocumentCleaning: assert "\ufffe" not in result assert "Text with" in result + def test_filter_string_preserves_valid_extended_characters(self): + """filter_string must keep valid printable characters like 'ï', '¿', '¾'.""" + # Arrange + text = "naïve ¿Cómo? ¾ done" + + # Act + result = IndexingRunner.filter_string(text) + + # Assert + assert result == text + # The U+FFFE noncharacter is still stripped. + assert IndexingRunner.filter_string("keep\ufffedrop") == "keepdrop" + class TestIndexingRunnerSplitter: """Unit tests for text splitter configuration. @@ -1487,20 +1623,21 @@ class TestIndexingRunnerLoadSegments: @pytest.fixture def sample_dataset(self): """Create sample dataset.""" - dataset = Mock(spec=Dataset) - dataset.id = str(uuid.uuid4()) - dataset.tenant_id = str(uuid.uuid4()) + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=str(uuid.uuid4()), + ) return dataset @pytest.fixture def sample_dataset_document(self): """Create sample dataset document.""" - doc = Mock(spec=DatasetDocument) - doc.id = str(uuid.uuid4()) - doc.dataset_id = str(uuid.uuid4()) - doc.created_by = str(uuid.uuid4()) - doc.doc_form = IndexStructureType.PARAGRAPH_INDEX - return doc + return DatasetDocument( + id=str(uuid.uuid4()), + dataset_id=str(uuid.uuid4()), + created_by=str(uuid.uuid4()), + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) @pytest.fixture def sample_documents(self): @@ -1655,7 +1792,7 @@ class TestIndexingRunnerEstimate: # Create too many extract settings with patch("core.indexing_runner.dify_config") as mock_config: - mock_config.BILLING_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.BATCH_UPLOAD_LIMIT = 10 extract_settings = [MagicMock() for _ in range(15)] @@ -1696,7 +1833,7 @@ class TestIndexingRunnerEstimate: patch("core.indexing_runner.storage") as mock_storage, patch("core.indexing_runner.dify_config") as mock_config, ): - mock_config.BILLING_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY result = runner.indexing_estimate( tenant_id=tenant_id, @@ -1756,11 +1893,11 @@ class TestIndexingRunnerProcessChunk: Document(page_content="Chunk 2", metadata={"doc_id": "c2"}), ] - mock_dataset = Mock(spec=Dataset) - mock_dataset.id = str(uuid.uuid4()) + mock_dataset = Dataset( + id=str(uuid.uuid4()), + ) - mock_dataset_document = Mock(spec=DatasetDocument) - mock_dataset_document.id = str(uuid.uuid4()) + mock_dataset_document = DatasetDocument(id=str(uuid.uuid4())) mock_dependencies["redis"].get.return_value = None @@ -1810,9 +1947,8 @@ class TestIndexingRunnerProcessChunk: runner = IndexingRunner() chunk_documents = [Document(page_content="Chunk", metadata={"doc_id": "c1"})] - mock_dataset = Mock(spec=Dataset) - mock_dataset_document = Mock(spec=DatasetDocument) - mock_dataset_document.id = str(uuid.uuid4()) + mock_dataset = Dataset() + mock_dataset_document = DatasetDocument(id=str(uuid.uuid4())) # Mock Redis to return paused status mock_dependencies["redis"].get.return_value = "1" diff --git a/api/tests/unit_tests/core/rag/rerank/test_reranker.py b/api/tests/unit_tests/core/rag/rerank/test_reranker.py index 3ed676bd442..b1fecffd9b5 100644 --- a/api/tests/unit_tests/core/rag/rerank/test_reranker.py +++ b/api/tests/unit_tests/core/rag/rerank/test_reranker.py @@ -12,11 +12,13 @@ All tests use mocking to avoid external dependencies and ensure fast, reliable e Tests follow the Arrange-Act-Assert pattern for clarity. """ +from datetime import UTC, datetime from operator import itemgetter from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch import pytest +from sqlalchemy.orm import Session from core.model_manager import ModelInstance from core.rag.index_processor.constant.doc_type import DocType @@ -28,7 +30,10 @@ from core.rag.rerank.rerank_factory import RerankRunnerFactory from core.rag.rerank.rerank_model import RerankModelRunner from core.rag.rerank.rerank_type import RerankMode from core.rag.rerank.weight_rerank import WeightRerankRunner +from extensions.storage.storage_type import StorageType from graphon.model_runtime.entities.rerank_entities import RerankDocument, RerankResult +from models.enums import CreatorUserRole +from models.model import UploadFile def create_mock_model_instance() -> ModelInstance: @@ -43,7 +48,15 @@ def create_mock_model_instance() -> ModelInstance: return mock_instance -class TestRerankModelRunner: +class _UsesSQLiteSession: + session: Session + + @pytest.fixture(autouse=True) + def _inject_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + + +class TestRerankModelRunner(_UsesSQLiteSession): """Unit tests for RerankModelRunner. Tests cover: @@ -69,7 +82,7 @@ class TestRerankModelRunner: @pytest.fixture def rerank_runner(self, mock_model_instance): """Create a RerankModelRunner with mocked model instance.""" - return RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + return RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) @pytest.fixture def sample_documents(self): @@ -420,14 +433,14 @@ class TestBaseRerankRunner: runner.run(query="python", documents=[]) -class TestRerankModelRunnerMultimodal: +class TestRerankModelRunnerMultimodal(_UsesSQLiteSession): @pytest.fixture def mock_model_instance(self): return create_mock_model_instance() @pytest.fixture def rerank_runner(self, mock_model_instance): - return RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + return RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) def test_run_returns_original_documents_for_non_text_query_without_vision_support( self, rerank_runner, mock_model_instance @@ -534,7 +547,7 @@ class TestRerankModelRunnerMultimodal: assert len(docs_arg) == 1 def test_fetch_multimodal_rerank_image_query_invokes_multimodal_model( - self, rerank_runner: RerankModelRunner, mock_model_instance + self, rerank_runner: RerankModelRunner, mock_model_instance, sqlite_session: Session ): text_doc = Document( page_content="text-content", @@ -547,14 +560,28 @@ class TestRerankModelRunnerMultimodal: ) mock_model_instance.invoke_multimodal_rerank.return_value = rerank_result - session = MagicMock() - session.get.return_value = SimpleNamespace(key="query-image-key") + upload_file = UploadFile( + tenant_id="00000000-0000-0000-0000-000000000001", + storage_type=StorageType.LOCAL, + key="query-image-key", + name="query.png", + size=10, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="00000000-0000-0000-0000-000000000002", + created_at=datetime.now(UTC), + used=True, + ) + upload_file.id = "00000000-0000-0000-0000-000000000003" + sqlite_session.add(upload_file) + sqlite_session.commit() with ( - patch.object(rerank_runner, "_session", session), + patch.object(rerank_runner, "_session", sqlite_session), patch("core.rag.rerank.rerank_model.storage.load_once", return_value=b"query-image-bytes"), ): result, unique_documents = rerank_runner.fetch_multimodal_rerank( - query="query-upload-id", + query=upload_file.id, documents=[text_doc], score_threshold=0.2, top_n=2, @@ -1047,7 +1074,7 @@ class TestWeightRerankRunner: assert result[0].metadata["score"] == pytest.approx(expected_score, rel=1e-6) -class TestRerankRunnerFactory: +class TestRerankRunnerFactory(_UsesSQLiteSession): """Unit tests for RerankRunnerFactory. Tests cover: @@ -1071,7 +1098,7 @@ class TestRerankRunnerFactory: runner = RerankRunnerFactory.create_rerank_runner( runner_type=RerankMode.RERANKING_MODEL, rerank_model_instance=mock_model_instance, - session=MagicMock(), + session=self.session, ) # Assert: Correct runner type is created @@ -1134,14 +1161,14 @@ class TestRerankRunnerFactory: runner = RerankRunnerFactory.create_rerank_runner( runner_type=RerankMode.RERANKING_MODEL.value, rerank_model_instance=mock_model_instance, - session=MagicMock(), + session=self.session, ) # Assert: Runner is created successfully assert isinstance(runner, RerankModelRunner) -class TestRerankIntegration: +class TestRerankIntegration(_UsesSQLiteSession): """Integration tests for reranker components. Tests cover: @@ -1199,7 +1226,7 @@ class TestRerankIntegration: runner = RerankRunnerFactory.create_rerank_runner( runner_type=RerankMode.RERANKING_MODEL, rerank_model_instance=mock_model_instance, - session=MagicMock(), + session=self.session, ) result = runner.run( query="best programming language", @@ -1240,7 +1267,7 @@ class TestRerankIntegration: Document(page_content="Low relevance", metadata={"doc_id": "doc3"}, provider="dify"), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1252,7 +1279,7 @@ class TestRerankIntegration: assert 0.0 <= result[2].metadata["score"] <= 1.0 -class TestRerankEdgeCases: +class TestRerankEdgeCases(_UsesSQLiteSession): """Edge case tests for reranker components. Tests cover: @@ -1302,7 +1329,7 @@ class TestRerankEdgeCases: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1342,7 +1369,7 @@ class TestRerankEdgeCases: Document(page_content="Negative score", metadata={"doc_id": "doc3"}, provider="dify"), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking with zero threshold result = runner.run(query="test", documents=documents, score_threshold=0.0) @@ -1378,7 +1405,7 @@ class TestRerankEdgeCases: Document(page_content="Perfect 3", metadata={"doc_id": "doc3"}, provider="dify"), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1419,7 +1446,7 @@ class TestRerankEdgeCases: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking result = runner.run(query="test 测试", documents=documents) @@ -1457,7 +1484,7 @@ class TestRerankEdgeCases: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1493,7 +1520,7 @@ class TestRerankEdgeCases: for i in range(num_docs) ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking with top_n result = runner.run(query="test", documents=documents, top_n=10) @@ -1583,7 +1610,7 @@ class TestRerankEdgeCases: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking with empty query result = runner.run(query="", documents=documents) @@ -1594,7 +1621,7 @@ class TestRerankEdgeCases: assert mock_model_instance.invoke_rerank.call_args.kwargs["query"] == "" -class TestRerankPerformance: +class TestRerankPerformance(_UsesSQLiteSession): """Performance and optimization tests for reranker. Tests cover: @@ -1636,7 +1663,7 @@ class TestRerankPerformance: for i in range(5) ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act: Run reranking result = runner.run(query="test", documents=documents) @@ -1711,7 +1738,7 @@ class TestRerankPerformance: assert "keywords" in result[1].metadata -class TestRerankErrorHandling: +class TestRerankErrorHandling(_UsesSQLiteSession): """Error handling tests for reranker components. Tests cover: @@ -1748,7 +1775,7 @@ class TestRerankErrorHandling: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act & Assert: Exception is raised with pytest.raises(RuntimeError, match="Model invocation failed"): @@ -1781,7 +1808,7 @@ class TestRerankErrorHandling: ), ] - runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock()) + runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session) # Act & Assert: Should raise IndexError or handle gracefully with pytest.raises(IndexError): diff --git a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py index 9e504dd1b9b..c677d64c59c 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py @@ -8,6 +8,7 @@ from uuid import uuid4 import pytest from flask import Flask, current_app +from sqlalchemy.orm import Session from core.app.app_config.entities import ( DatasetEntity, @@ -252,24 +253,25 @@ class TestRetrievalService: @pytest.fixture def mock_dataset(self) -> Dataset: """ - Create a mock Dataset object for testing. + Create a Dataset object for testing. Returns: - Dataset: Mock dataset with standard configuration + Dataset with standard configuration. """ - dataset = Mock(spec=Dataset) - dataset.id = str(uuid4()) - dataset.tenant_id = str(uuid4()) - dataset.name = "test_dataset" - dataset.indexing_technique = "high_quality" - dataset.embedding_model = "text-embedding-ada-002" - dataset.embedding_model_provider = "openai" - dataset.retrieval_model = { - "search_method": RetrievalMethod.SEMANTIC_SEARCH, - "reranking_enable": False, - "top_k": 4, - "score_threshold_enabled": False, - } + dataset = Dataset( + id=str(uuid4()), + tenant_id=str(uuid4()), + name="test_dataset", + indexing_technique="high_quality", + embedding_model="text-embedding-ada-002", + embedding_model_provider="openai", + retrieval_model={ + "search_method": RetrievalMethod.SEMANTIC_SEARCH, + "reranking_enable": False, + "top_k": 4, + "score_threshold_enabled": False, + }, + ) return dataset @pytest.fixture @@ -1834,10 +1836,11 @@ class TestRetrievalService: all_documents = [] # Create second dataset - mock_dataset2 = Mock(spec=Dataset) - mock_dataset2.id = str(uuid4()) - mock_dataset2.indexing_technique = "high_quality" - mock_dataset2.provider = "dify" + mock_dataset2 = Dataset( + id=str(uuid4()), + indexing_technique="high_quality", + provider="dify", + ) # Act - Call with dataset_count = 2 with _patched_retriever_session() as rerank_session: @@ -2207,36 +2210,33 @@ def create_mock_dataset_methods( tenant_id: str | None = None, provider: str = "dify", indexing_technique: str = "high_quality", - available_document_count: int = 10, -) -> Mock: +) -> Dataset: """ - Create a mock Dataset object for testing. + Create a Dataset object for testing. Args: dataset_id: Unique identifier for the dataset tenant_id: Tenant ID for the dataset provider: Provider type ("dify" or "external") indexing_technique: Indexing technique ("high_quality" or "economy") - available_document_count: Number of available documents - Returns: - Mock: A properly configured Dataset mock + A configured Dataset model. """ - dataset = Mock(spec=Dataset) - dataset.id = dataset_id or str(uuid4()) - dataset.tenant_id = tenant_id or str(uuid4()) - dataset.name = "test_dataset" - dataset.provider = provider - dataset.indexing_technique = indexing_technique - dataset.available_document_count = available_document_count - dataset.embedding_model = "text-embedding-ada-002" - dataset.embedding_model_provider = "openai" - dataset.retrieval_model = { - "search_method": "semantic_search", - "reranking_enable": False, - "top_k": 4, - "score_threshold_enabled": False, - } + dataset = Dataset( + id=dataset_id or str(uuid4()), + tenant_id=tenant_id or str(uuid4()), + name="test_dataset", + provider=provider, + indexing_technique=indexing_technique, + embedding_model="text-embedding-ada-002", + embedding_model_provider="openai", + retrieval_model={ + "search_method": "semantic_search", + "reranking_enable": False, + "top_k": 4, + "score_threshold_enabled": False, + }, + ) return dataset @@ -3740,12 +3740,13 @@ class TestProcessMetadataFilterFunc: class TestKnowledgeRetrievalRegression: @pytest.fixture def mock_dataset(self) -> Dataset: - dataset = Mock(spec=Dataset) - dataset.id = str(uuid4()) - dataset.tenant_id = str(uuid4()) - dataset.name = "test_dataset" - dataset.indexing_technique = "high_quality" - dataset.provider = "dify" + dataset = Dataset( + id=str(uuid4()), + tenant_id=str(uuid4()), + name="test_dataset", + indexing_technique="high_quality", + provider="dify", + ) return dataset def test_multiple_retrieve_reranking_with_app_context(self, mock_dataset): @@ -3760,10 +3761,11 @@ class TestKnowledgeRetrievalRegression: tenant_id = str(uuid4()) # second dataset to ensure dataset_count > 1 reranking branch - secondary_dataset = Mock(spec=Dataset) - secondary_dataset.id = str(uuid4()) - secondary_dataset.provider = "dify" - secondary_dataset.indexing_technique = "high_quality" + secondary_dataset = Dataset( + id=str(uuid4()), + provider="dify", + indexing_technique="high_quality", + ) # retriever returns 1 doc into internal list (all_documents_item) document = Document( @@ -3893,6 +3895,96 @@ class TestKnowledgeRetrievalRegression: assert cancel_event.is_set() assert thread_exceptions == [expected_error] + def test_run_retriever_thread_safely_skips_failed_dataset_when_requested(self, caplog): + dataset_retrieval = DatasetRetrieval() + all_documents: list[Document] = [] + cancel_event = threading.Event() + thread_exceptions: list[Exception] = [] + expected_error = RuntimeError("retrieval failed") + + with _patched_retriever_session(): + with patch.object(dataset_retrieval, "_retriever", side_effect=expected_error): + dataset_retrieval._run_retriever_thread_safely( + flask_app=_FakeFlaskApp(), + dataset_id="dataset-1", + query="test query", + top_k=3, + all_documents=all_documents, + document_ids_filter=None, + metadata_condition=None, + attachment_ids=None, + cancel_event=cancel_event, + thread_exceptions=thread_exceptions, + skip_on_error=True, + ) + + assert not cancel_event.is_set() + assert thread_exceptions == [] + assert "dataset_id=dataset-1" in caplog.text + assert "Skipping dataset retrieval because retriever failed" in caplog.text + + def test_multiple_retrieve_thread_skips_failed_dataset(self, mock_dataset, caplog): + dataset_retrieval = DatasetRetrieval() + flask_app = Flask(__name__) + successful_dataset = Dataset( + id=str(uuid4()), + provider="dify", + indexing_technique="high_quality", + ) + document = Document( + page_content="successful doc", + metadata={ + "doc_id": "doc1", + "score": 0.95, + "document_id": str(uuid4()), + "dataset_id": successful_dataset.id, + }, + provider="dify", + ) + + def fake_retriever( + flask_app, + session, + dataset_id, + query, + top_k, + all_documents, + document_ids_filter, + metadata_condition, + attachment_ids, + ): + if dataset_id == mock_dataset.id: + raise RuntimeError("dataset unavailable") + all_documents.append(document) + + all_documents: list[Document] = [] + + with ( + patch.object(dataset_retrieval, "_retriever", side_effect=fake_retriever), + _patched_retriever_session(), + ): + dataset_retrieval._multiple_retrieve_thread( + flask_app=flask_app, + available_datasets=[mock_dataset, successful_dataset], + metadata_condition=None, + metadata_filter_document_ids=None, + all_documents=all_documents, + tenant_id=str(uuid4()), + reranking_enable=False, + reranking_mode="reranking_model", + reranking_model=None, + weights=None, + top_k=3, + score_threshold=0.0, + query="test query", + attachment_id=None, + dataset_count=2, + ) + + assert all_documents == [document] + assert f"dataset_id={mock_dataset.id}" in caplog.text + assert "Skipping dataset retrieval because retriever failed" in caplog.text + class _FakeFlaskApp: def app_context(self): @@ -4870,6 +4962,7 @@ class TestSingleAndMultipleRetrieveCoverage: assert len(result) == 1 assert result[0].provider == "external" + session.scalar.assert_called_once() mock_end.assert_called_once() assert retrieval.llm_usage.total_tokens == 2 @@ -4946,6 +5039,95 @@ class TestSingleAndMultipleRetrieveCoverage: ) assert results == [] + def test_single_retrieve_rejects_dataset_outside_available_datasets(self, retrieval: DatasetRetrieval) -> None: + available_dataset = _dataset(id="ds-1", name="Available DS", description=None) + session = MagicMock() + session.scalar.return_value = _dataset( + id="ds-2", + name="Foreign DS", + provider="external", + tenant_id="tenant-2", + retrieval_model={}, + ) + + with ( + patch("core.rag.retrieval.dataset_retrieval.ReactMultiDatasetRouter") as mock_router_cls, + patch( + "core.rag.retrieval.dataset_retrieval.ExternalDatasetService.fetch_external_knowledge_retrieval", + return_value=[], + ) as mock_external_retrieve, + patch.object(retrieval, "_on_query") as mock_on_query, + ): + mock_router_cls.return_value.invoke.return_value = ("ds-2", LLMUsage.empty_usage()) + results = retrieval.single_retrieve( + session, + app_id="app-1", + tenant_id="tenant-1", + user_id="user-1", + user_from="workflow", + query="python", + available_datasets=[available_dataset], + model_instance=Mock(), + model_config=Mock(), + planning_strategy=PlanningStrategy.REACT_ROUTER, + ) + + assert results == [] + session.scalar.assert_not_called() + mock_external_retrieve.assert_not_called() + mock_on_query.assert_not_called() + + def test_single_retrieve_rejects_allowlisted_dataset_owned_by_another_tenant( + self, retrieval: DatasetRetrieval, sqlite_session: Session + ) -> None: + dataset_id = str(uuid4()) + caller_tenant_id = str(uuid4()) + foreign_dataset = Dataset( + id=dataset_id, + tenant_id=str(uuid4()), + name="Foreign DS", + provider="external", + indexing_technique="high_quality", + retrieval_model={}, + created_by=str(uuid4()), + ) + sqlite_session.add(foreign_dataset) + available_dataset = _dataset( + id=dataset_id, + tenant_id=caller_tenant_id, + name="Available DS", + description=None, + ) + + with ( + patch("core.rag.retrieval.dataset_retrieval.ReactMultiDatasetRouter") as mock_router_cls, + patch( + "core.rag.retrieval.dataset_retrieval.ExternalDatasetService.fetch_external_knowledge_retrieval", + ) as mock_external_retrieve, + patch( + "core.rag.retrieval.dataset_retrieval.RetrievalService.retrieve", + ) as mock_internal_retrieve, + patch.object(retrieval, "_on_query") as mock_on_query, + ): + mock_router_cls.return_value.invoke.return_value = (dataset_id, LLMUsage.empty_usage()) + results = retrieval.single_retrieve( + sqlite_session, + app_id="app-1", + tenant_id=caller_tenant_id, + user_id="user-1", + user_from="workflow", + query="python", + available_datasets=[available_dataset], + model_instance=Mock(), + model_config=Mock(), + planning_strategy=PlanningStrategy.REACT_ROUTER, + ) + + assert results == [] + mock_internal_retrieve.assert_not_called() + mock_external_retrieve.assert_not_called() + mock_on_query.assert_not_called() + def test_single_retrieve_respects_metadata_filter_shortcuts(self, retrieval: DatasetRetrieval) -> None: dataset = _dataset( id="ds-1", diff --git a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_methods.py b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_methods.py index 98413840d00..6768651d6df 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_methods.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval_methods.py @@ -1,348 +1,302 @@ -from typing import Any -from unittest.mock import MagicMock, Mock, patch -from uuid import uuid4 +"""SQLite-backed tests for dataset availability, rate limiting, and retrieval orchestration.""" + +from dataclasses import dataclass +from types import SimpleNamespace +from unittest.mock import MagicMock import pytest +from sqlalchemy import Engine, select +from sqlalchemy.orm import Session, sessionmaker +from core.rag.index_processor.constant.index_type import IndexTechniqueType from core.rag.models.document import Document +from core.rag.retrieval import dataset_retrieval as retrieval_module from core.rag.retrieval.dataset_retrieval import DatasetRetrieval from core.workflow.nodes.knowledge_retrieval import exc from core.workflow.nodes.knowledge_retrieval.retrieval import KnowledgeRetrievalRequest -from models.dataset import Dataset - -# ==================== Helper Functions ==================== +from models.dataset import Dataset, DocumentSegment, RateLimitLog +from models.dataset import Document as DatasetDocument +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus -def create_mock_dataset( - dataset_id: str | None = None, - tenant_id: str | None = None, - provider: str = "dify", - indexing_technique: str = "high_quality", - available_document_count: int = 10, -) -> Mock: - """ - Create a mock Dataset object for testing. +@dataclass(frozen=True) +class RetrievalDatabase: + session_maker: sessionmaker[Session] - Args: - dataset_id: Unique identifier for the dataset - tenant_id: Tenant ID for the dataset - provider: Provider type ("dify" or "external") - indexing_technique: Indexing technique ("high_quality" or "economy") - available_document_count: Number of available documents - Returns: - Mock: A properly configured Dataset mock - """ - dataset = Mock(spec=Dataset) - dataset.id = dataset_id or str(uuid4()) - dataset.tenant_id = tenant_id or str(uuid4()) - dataset.name = "test_dataset" - dataset.provider = provider - dataset.indexing_technique = indexing_technique - dataset.available_document_count = available_document_count - dataset.embedding_model = "text-embedding-ada-002" - dataset.embedding_model_provider = "openai" - dataset.retrieval_model = { - "search_method": "semantic_search", - "reranking_enable": False, - "top_k": 4, - "score_threshold_enabled": False, - } +@pytest.fixture +def retrieval_database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> RetrievalDatabase: + """Bind every retrieval-owned session to a disposable SQLite database.""" + Dataset.metadata.create_all( + sqlite_engine, + tables=[Dataset.__table__, DatasetDocument.__table__, DocumentSegment.__table__, RateLimitLog.__table__], + ) + session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr(retrieval_module.session_factory, "create_session", session_maker) + return RetrievalDatabase(session_maker=session_maker) + + +def _persist_dataset( + database: RetrievalDatabase, + *, + dataset_id: str, + tenant_id: str = "tenant-1", + provider: str = "vendor", +) -> Dataset: + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name=f"Dataset {dataset_id}", + created_by="user-1", + provider=provider, + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + embedding_model="text-embedding-ada-002", + embedding_model_provider="openai", + retrieval_model={ + "search_method": "semantic_search", + "reranking_enable": False, + "top_k": 4, + "score_threshold_enabled": False, + }, + ) + with database.session_maker.begin() as session: + session.add(dataset) return dataset -def create_mock_document( - content: str, - doc_id: str, - score: float = 0.8, - provider: str = "dify", - additional_metadata: dict[str, Any] | None = None, -) -> Document: - """ - Create a mock Document object for testing. +def _persist_document( + database: RetrievalDatabase, + *, + dataset_id: str, + document_id: str, + tenant_id: str = "tenant-1", + indexing_status: IndexingStatus = IndexingStatus.COMPLETED, + enabled: bool = True, + archived: bool = False, +) -> DatasetDocument: + document = DatasetDocument( + id=document_id, + tenant_id=tenant_id, + dataset_id=dataset_id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name=f"Document {document_id}", + created_from=DocumentCreatedFrom.API, + created_by="user-1", + indexing_status=indexing_status, + enabled=enabled, + archived=archived, + ) + with database.session_maker.begin() as session: + session.add(document) + return document - Args: - content: The text content of the document - doc_id: Unique identifier for the document chunk - score: Relevance score (0.0 to 1.0) - provider: Document provider ("dify" or "external") - additional_metadata: Optional extra metadata fields - Returns: - Document: A properly structured Document object - """ - metadata = { - "doc_id": doc_id, - "document_id": str(uuid4()), - "dataset_id": str(uuid4()), - "score": score, - } +def _persist_segment( + database: RetrievalDatabase, + *, + dataset_id: str, + document_id: str, +) -> DocumentSegment: + segment = DocumentSegment( + tenant_id="tenant-1", + dataset_id=dataset_id, + document_id=document_id, + position=1, + content="Python is great", + word_count=3, + tokens=3, + created_by="user-1", + index_node_id="node-1", + index_node_hash="hash-1", + hit_count=5, + ) + with database.session_maker.begin() as session: + session.add(segment) + return segment - if additional_metadata: - metadata.update(additional_metadata) - return Document( - page_content=content, - metadata=metadata, - provider=provider, +def _persist_available_dataset( + database: RetrievalDatabase, + *, + dataset_id: str = "dataset-1", + document_id: str = "document-1", +) -> tuple[Dataset, DatasetDocument]: + return ( + _persist_dataset(database, dataset_id=dataset_id), + _persist_document(database, dataset_id=dataset_id, document_id=document_id), ) -# ==================== Test _check_knowledge_rate_limit ==================== +def _request( + *, + dataset_ids: list[str], + retrieval_mode: str = "multiple", + metadata_filtering_mode: str = "disabled", +) -> KnowledgeRetrievalRequest: + return KnowledgeRetrievalRequest( + tenant_id="tenant-1", + user_id="user-1", + app_id="app-1", + user_from="web", + dataset_ids=dataset_ids, + query="What is Python?", + retrieval_mode=retrieval_mode, + metadata_filtering_mode=metadata_filtering_mode, + top_k=5, + score_threshold=0.7, + reranking_enable=True, + reranking_mode="reranking_model", + reranking_model={"reranking_provider_name": "cohere", "reranking_model_name": "rerank-v2"}, + ) + + +def _rag_document( + content: str, + doc_id: str, + *, + score: float = 0.8, + provider: str = "dify", + additional_metadata: dict[str, object] | None = None, +) -> Document: + metadata: dict[str, object] = { + "doc_id": doc_id, + "document_id": "document-1", + "dataset_id": "dataset-1", + "score": score, + } + if additional_metadata: + metadata.update(additional_metadata) + return Document(page_content=content, metadata=metadata, provider=provider) + + +def _patch_rate_limit( + monkeypatch: pytest.MonkeyPatch, + *, + enabled: bool, + request_count: int = 0, +) -> MagicMock: + limit = SimpleNamespace(enabled=enabled, limit=100, subscription_plan="professional") + monkeypatch.setattr( + retrieval_module.FeatureService, + "get_knowledge_rate_limit", + MagicMock(return_value=limit), + ) + redis = MagicMock() + redis.zcard.return_value = request_count + monkeypatch.setattr(retrieval_module, "redis_client", redis) + monkeypatch.setattr(retrieval_module.time, "time", lambda: 1234567890) + return redis class TestCheckKnowledgeRateLimit: - """ - Test suite for _check_knowledge_rate_limit method. + def test_rate_limit_disabled_performs_no_redis_or_database_work( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + redis = _patch_rate_limit(monkeypatch, enabled=False) - The _check_knowledge_rate_limit method validates whether a tenant has - exceeded their knowledge retrieval rate limit. This is important for: - - Preventing abuse of the knowledge retrieval system - - Enforcing subscription plan limits - - Tracking usage for billing purposes + DatasetRetrieval()._check_knowledge_rate_limit("tenant-1") - Test Cases: - ============ - 1. Rate limit disabled - no exception raised - 2. Rate limit enabled but not exceeded - no exception raised - 3. Rate limit enabled and exceeded - RateLimitExceededError raised - 4. Redis operations are performed correctly - 5. RateLimitLog is created when limit is exceeded - """ + redis.zadd.assert_not_called() + with retrieval_database.session_maker() as session: + assert session.scalar(select(RateLimitLog)) is None - @patch("core.rag.retrieval.dataset_retrieval.FeatureService") - @patch("core.rag.retrieval.dataset_retrieval.redis_client") - def test_rate_limit_disabled_no_exception(self, mock_redis, mock_feature_service): - """ - Test that when rate limit is disabled, no exception is raised. + def test_rate_limit_enabled_not_exceeded_tracks_request_without_log( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + redis = _patch_rate_limit(monkeypatch, enabled=True, request_count=50) - This test verifies the behavior when the tenant's subscription - does not have rate limiting enabled. + DatasetRetrieval()._check_knowledge_rate_limit("tenant-1") - Verifies: - - FeatureService.get_knowledge_rate_limit is called - - No Redis operations are performed - - No exception is raised - - Retrieval proceeds normally - """ - # Arrange - tenant_id = str(uuid4()) - dataset_retrieval = DatasetRetrieval() - - # Mock rate limit disabled - mock_limit = Mock() - mock_limit.enabled = False - mock_feature_service.get_knowledge_rate_limit.return_value = mock_limit - - # Act & Assert - should not raise any exception - dataset_retrieval._check_knowledge_rate_limit(tenant_id) - - # Verify FeatureService was called - mock_feature_service.get_knowledge_rate_limit.assert_called_once_with(tenant_id) - - # Verify no Redis operations were performed - assert not mock_redis.zadd.called - assert not mock_redis.zremrangebyscore.called - assert not mock_redis.zcard.called - - @patch("core.rag.retrieval.dataset_retrieval.session_factory") - @patch("core.rag.retrieval.dataset_retrieval.FeatureService") - @patch("core.rag.retrieval.dataset_retrieval.redis_client") - @patch("core.rag.retrieval.dataset_retrieval.time") - def test_rate_limit_enabled_not_exceeded(self, mock_time, mock_redis, mock_feature_service, mock_session_factory): - """ - Test that when rate limit is enabled but not exceeded, no exception is raised. - - This test simulates a tenant making requests within their rate limit. - The Redis sorted set stores timestamps of recent requests, and old - requests (older than 60 seconds) are removed. - - Verifies: - - Redis zadd is called to track the request - - Redis zremrangebyscore removes old entries - - Redis zcard returns count within limit - - No exception is raised - """ - # Arrange - tenant_id = str(uuid4()) - dataset_retrieval = DatasetRetrieval() - - # Mock rate limit enabled with limit of 100 requests per minute - mock_limit = Mock() - mock_limit.enabled = True - mock_limit.limit = 100 - mock_limit.subscription_plan = "professional" - mock_feature_service.get_knowledge_rate_limit.return_value = mock_limit - - # Mock time - current_time = 1234567890000 # Current time in milliseconds - mock_time.time.return_value = current_time / 1000 # Return seconds - mock_time.time.__mul__ = lambda self, x: int(self * x) # Multiply to get milliseconds - - # Mock Redis operations - # zcard returns 50 (within limit of 100) - mock_redis.zcard.return_value = 50 - - # Mock session_factory.create_session - mock_session = MagicMock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session - mock_session_factory.create_session.return_value.__exit__.return_value = None - - # Act & Assert - should not raise any exception - dataset_retrieval._check_knowledge_rate_limit(tenant_id) - - # Verify Redis operations - expected_key = f"rate_limit_{tenant_id}" - mock_redis.zadd.assert_called_once_with(expected_key, {current_time: current_time}) - mock_redis.zremrangebyscore.assert_called_once_with(expected_key, 0, current_time - 60000) - mock_redis.zcard.assert_called_once_with(expected_key) - - @patch("core.rag.retrieval.dataset_retrieval.session_factory") - @patch("core.rag.retrieval.dataset_retrieval.FeatureService") - @patch("core.rag.retrieval.dataset_retrieval.redis_client") - @patch("core.rag.retrieval.dataset_retrieval.time") - def test_rate_limit_enabled_exceeded_raises_exception( - self, mock_time, mock_redis, mock_feature_service, mock_session_factory - ): - """ - Test that when rate limit is enabled and exceeded, RateLimitExceededError is raised. - - This test simulates a tenant exceeding their rate limit. When the count - of recent requests exceeds the limit, an exception should be raised and - a RateLimitLog should be created. - - Verifies: - - Redis zcard returns count exceeding limit - - RateLimitExceededError is raised with correct message - - RateLimitLog is created in database - - Session operations are performed correctly - """ - # Arrange - tenant_id = str(uuid4()) - dataset_retrieval = DatasetRetrieval() - - # Mock rate limit enabled with limit of 100 requests per minute - mock_limit = Mock() - mock_limit.enabled = True - mock_limit.limit = 100 - mock_limit.subscription_plan = "professional" - mock_feature_service.get_knowledge_rate_limit.return_value = mock_limit - - # Mock time current_time = 1234567890000 - mock_time.time.return_value = current_time / 1000 + redis.zadd.assert_called_once_with("rate_limit_tenant-1", {current_time: current_time}) + redis.zremrangebyscore.assert_called_once_with("rate_limit_tenant-1", 0, current_time - 60000) + with retrieval_database.session_maker() as session: + assert session.scalar(select(RateLimitLog)) is None - # Mock Redis operations - return count exceeding limit - mock_redis.zcard.return_value = 150 # Exceeds limit of 100 + def test_rate_limit_exceeded_commits_audit_log_before_raising( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + _patch_rate_limit(monkeypatch, enabled=True, request_count=150) - # Mock session_factory.create_session - mock_session = MagicMock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session - mock_session_factory.create_session.return_value.__exit__.return_value = None + with pytest.raises(exc.RateLimitExceededError, match="knowledge base request rate limit"): + DatasetRetrieval()._check_knowledge_rate_limit("tenant-1") - # Act & Assert - with pytest.raises(exc.RateLimitExceededError) as exc_info: - dataset_retrieval._check_knowledge_rate_limit(tenant_id) - - # Verify exception message - assert "knowledge base request rate limit" in str(exc_info.value) - - # Verify RateLimitLog was created - mock_session.add.assert_called_once() - added_log = mock_session.add.call_args[0][0] - assert added_log.tenant_id == tenant_id - assert added_log.subscription_plan == "professional" - assert added_log.operation == "knowledge" - - -# ==================== Test _get_available_datasets ==================== + with retrieval_database.session_maker() as session: + logs = session.scalars(select(RateLimitLog)).all() + assert len(logs) == 1 + assert logs[0].tenant_id == "tenant-1" + assert logs[0].subscription_plan == "professional" + assert logs[0].operation == "knowledge" class TestGetAvailableDatasets: - """ - Test suite for _get_available_datasets method. + def test_returns_completed_or_external_datasets_with_tenant_scope( + self, + retrieval_database: RetrievalDatabase, + ) -> None: + _persist_available_dataset(retrieval_database, dataset_id="available", document_id="available-doc") + _persist_dataset(retrieval_database, dataset_id="disabled") + _persist_document( + retrieval_database, + dataset_id="disabled", + document_id="disabled-doc", + enabled=False, + ) + _persist_dataset(retrieval_database, dataset_id="archived") + _persist_document( + retrieval_database, + dataset_id="archived", + document_id="archived-doc", + archived=True, + ) + _persist_dataset(retrieval_database, dataset_id="waiting") + _persist_document( + retrieval_database, + dataset_id="waiting", + document_id="waiting-doc", + indexing_status=IndexingStatus.WAITING, + ) + _persist_dataset(retrieval_database, dataset_id="external", provider="external") + _persist_dataset(retrieval_database, dataset_id="other-tenant", tenant_id="tenant-2") + _persist_document( + retrieval_database, + dataset_id="other-tenant", + document_id="other-doc", + tenant_id="tenant-2", + ) - The _get_available_datasets method retrieves datasets that are available - for retrieval. A dataset is considered available if: - - It belongs to the specified tenant - - It's in the list of requested dataset_ids - - It has at least one completed, enabled, non-archived document OR - - It's an external provider dataset + datasets = DatasetRetrieval()._get_available_datasets( + "tenant-1", + ["available", "disabled", "archived", "waiting", "external", "other-tenant"], + ) - Note: Due to SQLAlchemy subquery complexity, full testing is done in - integration tests. Unit tests here verify basic behavior. - """ + assert {dataset.id for dataset in datasets} == {"available", "external"} - def test_method_exists_and_has_correct_signature(self): - """ - Test that the method exists and has the correct signature. + def test_returns_empty_for_vendor_dataset_without_documents( + self, + retrieval_database: RetrievalDatabase, + ) -> None: + _persist_dataset(retrieval_database, dataset_id="empty") - Verifies: - - Method exists on DatasetRetrieval class - - Accepts tenant_id and dataset_ids parameters - """ - # Arrange - dataset_retrieval = DatasetRetrieval() - - # Assert - method exists - assert hasattr(dataset_retrieval, "_get_available_datasets") - # Assert - method is callable - assert callable(dataset_retrieval._get_available_datasets) - - -# ==================== Test knowledge_retrieval ==================== + assert DatasetRetrieval()._get_available_datasets("tenant-1", ["empty"]) == [] class TestDatasetRetrievalKnowledgeRetrieval: - """ - Test suite for knowledge_retrieval method. - - The knowledge_retrieval method is the main entry point for retrieving - knowledge from datasets. It orchestrates the entire retrieval process: - 1. Checks rate limits - 2. Gets available datasets - 3. Applies metadata filtering if enabled - 4. Performs retrieval (single or multiple mode) - 5. Formats and returns results - - Test Cases: - ============ - 1. Single mode retrieval - 2. Multiple mode retrieval - 3. Metadata filtering disabled - 4. Metadata filtering automatic - 5. Metadata filtering manual - 6. External documents handling - 7. Dify documents handling - 8. Empty results handling - 9. Rate limit exceeded - 10. No available datasets - """ - - def test_knowledge_retrieval_single_mode_basic(self): - """ - Test knowledge_retrieval in single retrieval mode - basic check. - - Note: Full single mode testing requires complex model mocking and - is better suited for integration tests. This test verifies the - method accepts single mode requests. - - Verifies: - - Method can accept single mode request - - Request parameters are correctly structured - """ - # Arrange - tenant_id = str(uuid4()) - user_id = str(uuid4()) - app_id = str(uuid4()) - dataset_id = str(uuid4()) - + def test_single_mode_request_shape(self) -> None: request = KnowledgeRetrievalRequest( - tenant_id=tenant_id, - user_id=user_id, - app_id=app_id, + tenant_id="tenant-1", + user_id="user-1", + app_id="app-1", user_from="web", - dataset_ids=[dataset_id], + dataset_ids=["dataset-1"], query="What is Python?", retrieval_mode="single", model_provider="openai", @@ -351,365 +305,140 @@ class TestDatasetRetrievalKnowledgeRetrieval: completion_params={"temperature": 0.7}, ) - # Assert - request is properly structured assert request.retrieval_mode == "single" assert request.model_provider == "openai" assert request.model_name == "gpt-4" - assert request.model_mode == "chat" - @patch("core.rag.retrieval.dataset_retrieval.DataPostProcessor") - @patch("core.rag.retrieval.dataset_retrieval.session_factory") - def test_knowledge_retrieval_multiple_mode(self, mock_session_factory, mock_data_processor): - """ - Test knowledge_retrieval in multiple retrieval mode. - - In multiple mode, retrieval is performed across all datasets and - results are combined and reranked. - - Verifies: - - Rate limit is checked - - Available datasets are retrieved - - Multiple retrieval is performed - - Results are combined and reranked - - Results are formatted correctly - """ - # Arrange - tenant_id = str(uuid4()) - user_id = str(uuid4()) - app_id = str(uuid4()) - dataset_id1 = str(uuid4()) - dataset_id2 = str(uuid4()) - - request = KnowledgeRetrievalRequest( - tenant_id=tenant_id, - user_id=user_id, - app_id=app_id, - user_from="web", - dataset_ids=[dataset_id1, dataset_id2], - query="What is Python?", - retrieval_mode="multiple", - top_k=5, - score_threshold=0.7, - reranking_enable=True, - reranking_mode="reranking_model", - reranking_model={"reranking_provider_name": "cohere", "reranking_model_name": "rerank-v2"}, + def test_multiple_mode_formats_persisted_dataset_and_document( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + dataset, document = _persist_available_dataset(retrieval_database) + segment = _persist_segment(retrieval_database, dataset_id=dataset.id, document_id=document.id) + retrieval = DatasetRetrieval() + monkeypatch.setattr(retrieval, "_check_knowledge_rate_limit", MagicMock()) + monkeypatch.setattr(retrieval, "multiple_retrieve", MagicMock(return_value=[_rag_document("Python", "node-1")])) + record = SimpleNamespace(segment=segment, score=0.9, child_chunks=[], summary=None, files=None) + monkeypatch.setattr( + retrieval_module.RetrievalService, + "format_retrieval_documents", + MagicMock(return_value=[record]), ) + grant_access = MagicMock() + monkeypatch.setattr(retrieval_module, "grant_retriever_segment_access", grant_access) - dataset_retrieval = DatasetRetrieval() + with retrieval_database.session_maker() as caller_session: + result = retrieval.knowledge_retrieval(caller_session, _request(dataset_ids=[dataset.id])) - # Mock _check_knowledge_rate_limit - with patch.object(dataset_retrieval, "_check_knowledge_rate_limit"): - # Mock _get_available_datasets - mock_dataset1 = create_mock_dataset(dataset_id=dataset_id1, tenant_id=tenant_id) - mock_dataset2 = create_mock_dataset(dataset_id=dataset_id2, tenant_id=tenant_id) - with patch.object( - dataset_retrieval, "_get_available_datasets", return_value=[mock_dataset1, mock_dataset2] - ): - # Mock get_metadata_filter_condition - with patch.object(dataset_retrieval, "get_metadata_filter_condition", return_value=(None, None)): - # Mock multiple_retrieve to return documents - doc1 = create_mock_document("Python is great", "doc1", score=0.9) - doc2 = create_mock_document("Python is awesome", "doc2", score=0.8) - with patch.object( - dataset_retrieval, "multiple_retrieve", return_value=[doc1, doc2] - ) as mock_multiple_retrieve: - # Mock format_retrieval_documents - mock_record = Mock() - mock_record.segment = Mock() - mock_record.segment.dataset_id = dataset_id1 - mock_record.segment.document_id = str(uuid4()) - mock_record.segment.index_node_hash = "hash123" - mock_record.segment.hit_count = 5 - mock_record.segment.word_count = 100 - mock_record.segment.position = 1 - mock_record.segment.get_sign_content.return_value = "Python is great" - mock_record.segment.answer = None - mock_record.score = 0.9 - mock_record.child_chunks = [] - mock_record.summary = None - mock_record.files = None + assert len(result) == 1 + assert result[0].metadata.dataset_id == dataset.id + assert result[0].metadata.document_id == document.id + assert result[0].title == document.name + grant_access.assert_called_once_with([segment.id]) - mock_retrieval_service = Mock() - mock_retrieval_service.format_retrieval_documents.return_value = [mock_record] + def test_metadata_filtering_disabled_skips_filter_builder( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + _persist_available_dataset(retrieval_database) + retrieval = DatasetRetrieval() + monkeypatch.setattr(retrieval, "_check_knowledge_rate_limit", MagicMock()) + metadata_filter = MagicMock(return_value=(None, None)) + monkeypatch.setattr(retrieval, "get_metadata_filter_condition", metadata_filter) + monkeypatch.setattr(retrieval, "multiple_retrieve", MagicMock(return_value=[])) - with patch( - "core.rag.retrieval.dataset_retrieval.RetrievalService", - return_value=mock_retrieval_service, - ): - # Mock database queries - mock_session = MagicMock() - mock_session_factory.create_session.return_value.__enter__.return_value = mock_session - mock_session_factory.create_session.return_value.__exit__.return_value = None + with retrieval_database.session_maker() as caller_session: + result = retrieval.knowledge_retrieval(caller_session, _request(dataset_ids=["dataset-1"])) - mock_dataset_from_db = Mock() - mock_dataset_from_db.id = dataset_id1 - mock_dataset_from_db.name = "test_dataset" + assert result == [] + metadata_filter.assert_not_called() - mock_document = Mock() - mock_document.id = str(uuid4()) - mock_document.name = "test_doc" - mock_document.data_source_type = "upload_file" - mock_document.doc_metadata = {} - - mock_datasets = MagicMock() - mock_datasets.all.return_value = [mock_dataset_from_db] - mock_documents = MagicMock() - mock_documents.all.return_value = [mock_document] - mock_session.scalars.side_effect = [mock_datasets, mock_documents] - - # Act - result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) - - # Assert - assert isinstance(result, list) - mock_multiple_retrieve.assert_called_once() - - def test_knowledge_retrieval_metadata_filtering_disabled(self): - """ - Test knowledge_retrieval with metadata filtering disabled. - - When metadata filtering is disabled, get_metadata_filter_condition is - NOT called (the method checks metadata_filtering_mode != "disabled"). - - Verifies: - - get_metadata_filter_condition is NOT called when mode is "disabled" - - Retrieval proceeds without metadata filters - """ - # Arrange - tenant_id = str(uuid4()) - user_id = str(uuid4()) - app_id = str(uuid4()) - dataset_id = str(uuid4()) - - request = KnowledgeRetrievalRequest( - tenant_id=tenant_id, - user_id=user_id, - app_id=app_id, - user_from="web", - dataset_ids=[dataset_id], - query="What is Python?", - retrieval_mode="multiple", - metadata_filtering_mode="disabled", - top_k=5, + def test_external_documents_are_formatted_without_database_document( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + _persist_dataset(retrieval_database, dataset_id="external", provider="external") + retrieval = DatasetRetrieval() + monkeypatch.setattr(retrieval, "_check_knowledge_rate_limit", MagicMock()) + external_document = _rag_document( + "External knowledge", + "external-node", + score=0.9, + provider="external", + additional_metadata={ + "dataset_id": "external", + "dataset_name": "External Dataset", + "document_id": "external-document", + "title": "External Document", + }, ) + monkeypatch.setattr(retrieval, "multiple_retrieve", MagicMock(return_value=[external_document])) - dataset_retrieval = DatasetRetrieval() + with retrieval_database.session_maker() as caller_session: + result = retrieval.knowledge_retrieval(caller_session, _request(dataset_ids=["external"])) - # Mock dependencies - with patch.object(dataset_retrieval, "_check_knowledge_rate_limit"): - mock_dataset = create_mock_dataset(dataset_id=dataset_id, tenant_id=tenant_id) - with patch.object(dataset_retrieval, "_get_available_datasets", return_value=[mock_dataset]): - # Mock get_metadata_filter_condition - should NOT be called when disabled - with patch.object( - dataset_retrieval, - "get_metadata_filter_condition", - return_value=(None, None), - ) as mock_get_metadata: - with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[]): - # Act - result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) + assert len(result) == 1 + assert result[0].metadata.data_source_type == "external" + assert result[0].metadata.dataset_id == "external" - # Assert - assert isinstance(result, list) - # get_metadata_filter_condition should NOT be called when mode is "disabled" - mock_get_metadata.assert_not_called() + def test_empty_retrieval_results_return_empty_list( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + _persist_available_dataset(retrieval_database) + retrieval = DatasetRetrieval() + monkeypatch.setattr(retrieval, "_check_knowledge_rate_limit", MagicMock()) + monkeypatch.setattr(retrieval, "multiple_retrieve", MagicMock(return_value=[])) - def test_knowledge_retrieval_with_external_documents(self): - """ - Test knowledge_retrieval with external documents. + with retrieval_database.session_maker() as caller_session: + result = retrieval.knowledge_retrieval(caller_session, _request(dataset_ids=["dataset-1"])) - External documents come from external knowledge bases and should - be formatted differently than Dify documents. + assert result == [] - Verifies: - - External documents are handled correctly - - Provider is set to "external" - - Metadata includes external-specific fields - """ - # Arrange - tenant_id = str(uuid4()) - user_id = str(uuid4()) - app_id = str(uuid4()) - dataset_id = str(uuid4()) - - request = KnowledgeRetrievalRequest( - tenant_id=tenant_id, - user_id=user_id, - app_id=app_id, - user_from="web", - dataset_ids=[dataset_id], - query="What is Python?", - retrieval_mode="multiple", - top_k=5, - ) - - dataset_retrieval = DatasetRetrieval() - - # Mock dependencies - with patch.object(dataset_retrieval, "_check_knowledge_rate_limit"): - mock_dataset = create_mock_dataset(dataset_id=dataset_id, tenant_id=tenant_id, provider="external") - with patch.object(dataset_retrieval, "_get_available_datasets", return_value=[mock_dataset]): - with patch.object(dataset_retrieval, "get_metadata_filter_condition", return_value=(None, None)): - # Create external document - external_doc = create_mock_document( - "External knowledge", - "doc1", - score=0.9, - provider="external", - additional_metadata={ - "dataset_id": dataset_id, - "dataset_name": "external_kb", - "document_id": "ext_doc1", - "title": "External Document", - }, - ) - with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[external_doc]): - # Act - result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) - - # Assert - assert isinstance(result, list) - if result: - assert result[0].metadata.data_source_type == "external" - - def test_knowledge_retrieval_empty_results(self): - """ - Test knowledge_retrieval when no documents are found. - - Verifies: - - Empty list is returned - - No errors are raised - - All dependencies are still called - """ - # Arrange - tenant_id = str(uuid4()) - user_id = str(uuid4()) - app_id = str(uuid4()) - dataset_id = str(uuid4()) - - request = KnowledgeRetrievalRequest( - tenant_id=tenant_id, - user_id=user_id, - app_id=app_id, - user_from="web", - dataset_ids=[dataset_id], - query="What is Python?", - retrieval_mode="multiple", - top_k=5, - ) - - dataset_retrieval = DatasetRetrieval() - - # Mock dependencies - with patch.object(dataset_retrieval, "_check_knowledge_rate_limit"): - mock_dataset = create_mock_dataset(dataset_id=dataset_id, tenant_id=tenant_id) - with patch.object(dataset_retrieval, "_get_available_datasets", return_value=[mock_dataset]): - with patch.object(dataset_retrieval, "get_metadata_filter_condition", return_value=(None, None)): - # Mock multiple_retrieve to return empty list - with patch.object(dataset_retrieval, "multiple_retrieve", return_value=[]): - # Act - result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) - - # Assert - assert result == [] - - def test_knowledge_retrieval_rate_limit_exceeded(self): - """ - Test knowledge_retrieval when rate limit is exceeded. - - Verifies: - - RateLimitExceededError is raised - - No further processing occurs - """ - # Arrange - tenant_id = str(uuid4()) - user_id = str(uuid4()) - app_id = str(uuid4()) - dataset_id = str(uuid4()) - - request = KnowledgeRetrievalRequest( - tenant_id=tenant_id, - user_id=user_id, - app_id=app_id, - user_from="web", - dataset_ids=[dataset_id], - query="What is Python?", - retrieval_mode="multiple", - top_k=5, - ) - - dataset_retrieval = DatasetRetrieval() - - # Mock _check_knowledge_rate_limit to raise exception - with patch.object( - dataset_retrieval, + def test_rate_limit_exception_stops_retrieval( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + retrieval = DatasetRetrieval() + monkeypatch.setattr( + retrieval, "_check_knowledge_rate_limit", - side_effect=exc.RateLimitExceededError("Rate limit exceeded"), - ): - # Act & Assert - with pytest.raises(exc.RateLimitExceededError): - dataset_retrieval.knowledge_retrieval(MagicMock(), request) - - def test_knowledge_retrieval_no_available_datasets(self): - """ - Test knowledge_retrieval when no datasets are available. - - Verifies: - - Empty list is returned - - No retrieval is attempted - """ - # Arrange - tenant_id = str(uuid4()) - user_id = str(uuid4()) - app_id = str(uuid4()) - dataset_id = str(uuid4()) - - request = KnowledgeRetrievalRequest( - tenant_id=tenant_id, - user_id=user_id, - app_id=app_id, - user_from="web", - dataset_ids=[dataset_id], - query="What is Python?", - retrieval_mode="multiple", - top_k=5, + MagicMock(side_effect=exc.RateLimitExceededError("Rate limit exceeded")), ) - dataset_retrieval = DatasetRetrieval() + with retrieval_database.session_maker() as caller_session: + with pytest.raises(exc.RateLimitExceededError): + retrieval.knowledge_retrieval(caller_session, _request(dataset_ids=["dataset-1"])) - # Mock dependencies - with patch.object(dataset_retrieval, "_check_knowledge_rate_limit"): - # Mock _get_available_datasets to return empty list - with patch.object(dataset_retrieval, "_get_available_datasets", return_value=[]): - # Act - result = dataset_retrieval.knowledge_retrieval(MagicMock(), request) + def test_no_available_datasets_skips_retrieval( + self, + monkeypatch: pytest.MonkeyPatch, + retrieval_database: RetrievalDatabase, + ) -> None: + _persist_dataset(retrieval_database, dataset_id="empty") + retrieval = DatasetRetrieval() + monkeypatch.setattr(retrieval, "_check_knowledge_rate_limit", MagicMock()) + multiple_retrieve = MagicMock() + monkeypatch.setattr(retrieval, "multiple_retrieve", multiple_retrieve) - # Assert - assert result == [] + with retrieval_database.session_maker() as caller_session: + result = retrieval.knowledge_retrieval(caller_session, _request(dataset_ids=["empty"])) - def test_knowledge_retrieval_handles_multiple_documents_with_different_scores(self): - """ - Test that knowledge_retrieval processes multiple documents with different scores. + assert result == [] + multiple_retrieve.assert_not_called() - Note: Full sorting and position testing requires complex SQLAlchemy mocking - which is better suited for integration tests. This test verifies documents - with different scores can be created and have their metadata. + def test_document_scores_sort_descending(self) -> None: + documents = [ + _rag_document("Low", "doc1", score=0.6), + _rag_document("High", "doc2", score=0.95), + _rag_document("Medium", "doc3", score=0.8), + ] - Verifies: - - Documents can be created with different scores - - Score metadata is properly set - """ - # Create documents with different scores - doc1 = create_mock_document("Low score", "doc1", score=0.6) - doc2 = create_mock_document("High score", "doc2", score=0.95) - doc3 = create_mock_document("Medium score", "doc3", score=0.8) + sorted_documents = sorted(documents, key=lambda document: document.metadata["score"], reverse=True) - # Assert - each document has the correct score - assert doc1.metadata["score"] == 0.6 - assert doc2.metadata["score"] == 0.95 - assert doc3.metadata["score"] == 0.8 - - # Assert - documents are correctly sorted (not the retrieval result, just the list) - unsorted = [doc1, doc2, doc3] - sorted_docs = sorted(unsorted, key=lambda d: d.metadata["score"], reverse=True) - assert [d.metadata["score"] for d in sorted_docs] == [0.95, 0.8, 0.6] + assert [document.metadata["score"] for document in sorted_documents] == [0.95, 0.8, 0.6] diff --git a/api/tests/unit_tests/core/rag/retrieval/test_multi_dataset_react_route.py b/api/tests/unit_tests/core/rag/retrieval/test_multi_dataset_react_route.py index c56528cf55e..bcc8ee8fcb5 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_multi_dataset_react_route.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_multi_dataset_react_route.py @@ -165,10 +165,7 @@ class TestReactMultiDatasetRouter: model_instance = Mock() model_instance.invoke_llm.return_value = iter([chunk]) - with ( - patch("core.rag.retrieval.router.multi_dataset_react_route.ModelManager.for_tenant") as mock_manager, - patch("core.rag.retrieval.router.multi_dataset_react_route.deduct_llm_quota") as mock_deduct, - ): + with patch("core.rag.retrieval.router.multi_dataset_react_route.ModelManager.for_tenant") as mock_manager: mock_manager.return_value.get_model_instance.return_value = model_instance text, returned_usage = router._invoke_llm( completion_param={"temperature": 0.1}, @@ -188,7 +185,6 @@ class TestReactMultiDatasetRouter: model_type=ModelType.LLM, model=model_instance.model_name, ) - mock_deduct.assert_called_once() def test_handle_invoke_result_with_empty_usage(self) -> None: router = ReactMultiDatasetRouter() diff --git a/api/tests/unit_tests/core/rag/splitter/test_text_splitter.py b/api/tests/unit_tests/core/rag/splitter/test_text_splitter.py index 12117241b5d..7a0726d3a9c 100644 --- a/api/tests/unit_tests/core/rag/splitter/test_text_splitter.py +++ b/api/tests/unit_tests/core/rag/splitter/test_text_splitter.py @@ -952,6 +952,21 @@ class TestFixedRecursiveCharacterTextSplitter: assert "word1" in combined assert "word2" in combined + def test_preserves_spaces_when_recursively_splitting_long_paragraph(self): + """Ensure recursive space splitting preserves word boundaries.""" + text = "여름철에는 항상 기상상황에 주목하며 주변 사람들과 함께 정보를 공유합니다." + splitter = FixedRecursiveCharacterTextSplitter( + fixed_separator="\n\n", + chunk_size=20, + chunk_overlap=0, + keep_separator=True, + ) + + result = splitter.split_text(text) + + assert len(result) > 1 + assert " ".join(result) == text + def test_character_level_splitting(self): """Test character-level splitting when no separator works.""" text = "verylongwordwithoutspaces" diff --git a/api/tests/unit_tests/core/repositories/test_celery_workflow_execution_repository.py b/api/tests/unit_tests/core/repositories/test_celery_workflow_execution_repository.py index b158061dbe6..64a1b1a33df 100644 --- a/api/tests/unit_tests/core/repositories/test_celery_workflow_execution_repository.py +++ b/api/tests/unit_tests/core/repositories/test_celery_workflow_execution_repository.py @@ -5,7 +5,7 @@ These tests verify the Celery-based asynchronous storage functionality for workflow execution data. """ -from unittest.mock import Mock, patch +from unittest.mock import patch from uuid import uuid4 import pytest @@ -14,7 +14,7 @@ from core.repositories.celery_workflow_execution_repository import CeleryWorkflo from graphon.entities import WorkflowExecution from graphon.enums import WorkflowType from libs.datetime_utils import naive_utc_now -from models import Account, EndUser +from models import Account, EndUser, Tenant from models.enums import WorkflowRunTriggeredFrom RESOURCE_TENANT_ID = "resource-tenant-id" @@ -34,18 +34,20 @@ def mock_session_factory(): @pytest.fixture def mock_account(): """Mock Account user.""" - account = Mock(spec=Account) + account = Account(name="Test Account", email="test@example.com") account.id = str(uuid4()) - account.current_tenant_id = str(uuid4()) + account._current_tenant = Tenant(name="Test Tenant") + account._current_tenant.id = str(uuid4()) return account @pytest.fixture def mock_end_user(): """Mock EndUser.""" - user = Mock(spec=EndUser) - user.id = str(uuid4()) - user.tenant_id = str(uuid4()) + user = EndUser( + id=str(uuid4()), + tenant_id=str(uuid4()), + ) return user @@ -114,9 +116,8 @@ class TestCeleryWorkflowExecutionRepository: def test_init_without_tenant_id_raises_error(self, mock_session_factory): """Test that initialization fails without tenant_id.""" - # Create a mock Account with no tenant_id - user = Mock(spec=Account) - user.current_tenant_id = None + # Create an Account with no tenant_id. + user = Account(name="Test Account", email="test@example.com") user.id = str(uuid4()) with pytest.raises(ValueError, match="tenant_id is required"): @@ -129,8 +130,7 @@ class TestCeleryWorkflowExecutionRepository: ) def test_init_uses_resource_tenant_when_account_has_no_current_tenant(self, mock_session_factory): - user = Mock(spec=Account) - user.current_tenant_id = None + user = Account(name="Test Account", email="test@example.com") user.id = str(uuid4()) repo = CeleryWorkflowExecutionRepository( diff --git a/api/tests/unit_tests/core/repositories/test_celery_workflow_node_execution_repository.py b/api/tests/unit_tests/core/repositories/test_celery_workflow_node_execution_repository.py index 018626c450c..2ca60937353 100644 --- a/api/tests/unit_tests/core/repositories/test_celery_workflow_node_execution_repository.py +++ b/api/tests/unit_tests/core/repositories/test_celery_workflow_node_execution_repository.py @@ -18,7 +18,7 @@ from graphon.entities.workflow_node_execution import ( ) from graphon.enums import BuiltinNodeTypes from libs.datetime_utils import naive_utc_now -from models import Account, EndUser +from models import Account, EndUser, Tenant from models.workflow import WorkflowNodeExecutionTriggeredFrom RESOURCE_TENANT_ID = "resource-tenant-id" @@ -38,18 +38,20 @@ def mock_session_factory(): @pytest.fixture def mock_account(): """Mock Account user.""" - account = Mock(spec=Account) + account = Account(name="Test Account", email="test@example.com") account.id = str(uuid4()) - account.current_tenant_id = str(uuid4()) + account._current_tenant = Tenant(name="Test Tenant") + account._current_tenant.id = str(uuid4()) return account @pytest.fixture def mock_end_user(): """Mock EndUser.""" - user = Mock(spec=EndUser) - user.id = str(uuid4()) - user.tenant_id = str(uuid4()) + user = EndUser( + id=str(uuid4()), + tenant_id=str(uuid4()), + ) return user @@ -120,9 +122,8 @@ class TestCeleryWorkflowNodeExecutionRepository: def test_init_without_tenant_id_raises_error(self, mock_session_factory): """Test that initialization fails without tenant_id.""" - # Create a mock Account with no tenant_id - user = Mock(spec=Account) - user.current_tenant_id = None + # Create an Account with no tenant_id. + user = Account(name="Test Account", email="test@example.com") user.id = str(uuid4()) with pytest.raises(ValueError, match="tenant_id is required"): @@ -135,8 +136,7 @@ class TestCeleryWorkflowNodeExecutionRepository: ) def test_init_uses_resource_tenant_when_account_has_no_current_tenant(self, mock_session_factory): - user = Mock(spec=Account) - user.current_tenant_id = None + user = Account(name="Test Account", email="test@example.com") user.id = str(uuid4()) repo = CeleryWorkflowNodeExecutionRepository( @@ -186,6 +186,23 @@ class TestCeleryWorkflowNodeExecutionRepository: in repo._workflow_execution_mapping[sample_workflow_node_execution.workflow_execution_id] ) + def test_save_synchronously_uses_sql_repository_without_queueing( + self, mock_session_factory, mock_account, sample_workflow_node_execution + ): + repo = CeleryWorkflowNodeExecutionRepository( + session_factory=mock_session_factory, + tenant_id=RESOURCE_TENANT_ID, + user=mock_account, + app_id="test-app", + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + ) + repo._sql_repository.save_synchronously = Mock() + + repo.save_synchronously(sample_workflow_node_execution) + + repo._sql_repository.save_synchronously.assert_called_once_with(sample_workflow_node_execution) + assert repo._execution_cache[sample_workflow_node_execution.id] is sample_workflow_node_execution + @patch("core.repositories.celery_workflow_node_execution_repository.save_workflow_node_execution_task") def test_save_handles_celery_failure( self, mock_task, mock_session_factory, mock_account, sample_workflow_node_execution @@ -245,6 +262,60 @@ class TestCeleryWorkflowNodeExecutionRepository: # Should return empty list since nothing in cache assert len(result) == 0 + def test_get_by_workflow_execution_loads_persisted_executions_on_cache_miss( + self, mock_session_factory, mock_account, sample_workflow_node_execution + ): + repo = CeleryWorkflowNodeExecutionRepository( + session_factory=mock_session_factory, + tenant_id=RESOURCE_TENANT_ID, + user=mock_account, + app_id="test-app", + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + ) + repo._sql_repository = Mock() + repo._sql_repository.get_by_workflow_execution.return_value = [sample_workflow_node_execution] + + result = repo.get_by_workflow_execution(sample_workflow_node_execution.workflow_execution_id) + + assert result == [sample_workflow_node_execution] + assert repo._execution_cache[sample_workflow_node_execution.id] is sample_workflow_node_execution + assert repo._workflow_execution_mapping[sample_workflow_node_execution.workflow_execution_id] == [ + sample_workflow_node_execution.id + ] + + @patch("core.repositories.celery_workflow_node_execution_repository.save_workflow_node_execution_task") + def test_get_by_workflow_execution_merges_database_and_newer_cache( + self, mock_task, mock_session_factory, mock_account, sample_workflow_node_execution + ): + repo = CeleryWorkflowNodeExecutionRepository( + session_factory=mock_session_factory, + tenant_id=RESOURCE_TENANT_ID, + user=mock_account, + app_id="test-app", + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + ) + persisted_current = sample_workflow_node_execution.model_copy(deep=True) + historical = sample_workflow_node_execution.model_copy( + update={ + "id": str(uuid4()), + "node_execution_id": str(uuid4()), + "index": 0, + "node_id": "start", + } + ) + sample_workflow_node_execution.status = WorkflowNodeExecutionStatus.SUCCEEDED + repo.save(sample_workflow_node_execution) + repo._sql_repository = Mock() + repo._sql_repository.get_by_workflow_execution.return_value = [persisted_current, historical] + + result = repo.get_by_workflow_execution( + sample_workflow_node_execution.workflow_execution_id, + OrderConfig(order_by=["index"], order_direction="asc"), + ) + + assert [execution.id for execution in result] == [historical.id, sample_workflow_node_execution.id] + assert result[1] is sample_workflow_node_execution + @patch("core.repositories.celery_workflow_node_execution_repository.save_workflow_node_execution_task") def test_cache_operations(self, mock_task, mock_session_factory, mock_account, sample_workflow_node_execution): """Test cache operations work correctly.""" diff --git a/api/tests/unit_tests/core/repositories/test_factory.py b/api/tests/unit_tests/core/repositories/test_factory.py index 1b47cb64a70..700a986a310 100644 --- a/api/tests/unit_tests/core/repositories/test_factory.py +++ b/api/tests/unit_tests/core/repositories/test_factory.py @@ -9,7 +9,7 @@ from unittest.mock import MagicMock, patch import pytest from sqlalchemy.engine import Engine -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.repositories.factory import ( DifyCoreRepositoryFactory, @@ -25,6 +25,15 @@ from models.workflow import WorkflowNodeExecutionTriggeredFrom RESOURCE_TENANT_ID = "resource-tenant-id" +@pytest.fixture +def sqlite_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]: + """Return a real session factory bound to the test's isolated SQLite engine.""" + factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + with factory() as session: + assert session.get_bind() is sqlite_engine + return factory + + class TestRepositoryFactory: """Test cases for RepositoryFactory.""" @@ -54,14 +63,13 @@ class TestRepositoryFactory: assert "doesn't look like a module path" in str(exc_info.value) @patch("core.repositories.factory.dify_config") - def test_create_workflow_execution_repository_success(self, mock_config): + def test_create_workflow_execution_repository_success(self, mock_config, sqlite_session_factory): """Test successful WorkflowExecutionRepository creation.""" # Setup mock configuration mock_config.CORE_WORKFLOW_EXECUTION_REPOSITORY = "unittest.mock.MagicMock" - # Create mock dependencies - mock_session_factory = MagicMock(spec=sessionmaker) - mock_user = MagicMock(spec=Account) + # Create non-database dependencies + mock_user = Account(name="Test Account", email="test@example.com") app_id = "test-app-id" triggered_from = WorkflowRunTriggeredFrom.APP_RUN @@ -73,7 +81,7 @@ class TestRepositoryFactory: # Mock import_string with patch("core.repositories.factory.import_string", return_value=mock_repository_class, autospec=True): result = DifyCoreRepositoryFactory.create_workflow_execution_repository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id=app_id, @@ -82,7 +90,7 @@ class TestRepositoryFactory: # Verify the repository was created with correct parameters mock_repository_class.assert_called_once_with( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id=app_id, @@ -91,17 +99,16 @@ class TestRepositoryFactory: assert result is mock_repository_instance @patch("core.repositories.factory.dify_config") - def test_create_workflow_execution_repository_import_error(self, mock_config): + def test_create_workflow_execution_repository_import_error(self, mock_config, sqlite_session_factory): """Test WorkflowExecutionRepository creation with import error.""" # Setup mock configuration with invalid class path mock_config.CORE_WORKFLOW_EXECUTION_REPOSITORY = "invalid.module.InvalidClass" - mock_session_factory = MagicMock(spec=sessionmaker) - mock_user = MagicMock(spec=Account) + mock_user = Account(name="Test Account", email="test@example.com") with pytest.raises(RepositoryImportError) as exc_info: DifyCoreRepositoryFactory.create_workflow_execution_repository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id="test-app-id", @@ -110,13 +117,12 @@ class TestRepositoryFactory: assert "Failed to create WorkflowExecutionRepository" in str(exc_info.value) @patch("core.repositories.factory.dify_config") - def test_create_workflow_execution_repository_instantiation_error(self, mock_config): + def test_create_workflow_execution_repository_instantiation_error(self, mock_config, sqlite_session_factory): """Test WorkflowExecutionRepository creation with instantiation error.""" # Setup mock configuration mock_config.CORE_WORKFLOW_EXECUTION_REPOSITORY = "unittest.mock.MagicMock" - mock_session_factory = MagicMock(spec=sessionmaker) - mock_user = MagicMock(spec=Account) + mock_user = Account(name="Test Account", email="test@example.com") # Create a mock repository class that raises exception on instantiation mock_repository_class = MagicMock() @@ -126,7 +132,7 @@ class TestRepositoryFactory: with patch("core.repositories.factory.import_string", return_value=mock_repository_class, autospec=True): with pytest.raises(RepositoryImportError) as exc_info: DifyCoreRepositoryFactory.create_workflow_execution_repository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id="test-app-id", @@ -135,14 +141,13 @@ class TestRepositoryFactory: assert "Failed to create WorkflowExecutionRepository" in str(exc_info.value) @patch("core.repositories.factory.dify_config") - def test_create_workflow_node_execution_repository_success(self, mock_config): + def test_create_workflow_node_execution_repository_success(self, mock_config, sqlite_session_factory): """Test successful WorkflowNodeExecutionRepository creation.""" # Setup mock configuration mock_config.CORE_WORKFLOW_NODE_EXECUTION_REPOSITORY = "unittest.mock.MagicMock" - # Create mock dependencies - mock_session_factory = MagicMock(spec=sessionmaker) - mock_user = MagicMock(spec=EndUser) + # Create non-database dependencies + mock_user = EndUser() app_id = "test-app-id" triggered_from = WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP @@ -154,7 +159,7 @@ class TestRepositoryFactory: # Mock import_string with patch("core.repositories.factory.import_string", return_value=mock_repository_class, autospec=True): result = DifyCoreRepositoryFactory.create_workflow_node_execution_repository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id=app_id, @@ -163,7 +168,7 @@ class TestRepositoryFactory: # Verify the repository was created with correct parameters mock_repository_class.assert_called_once_with( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id=app_id, @@ -172,17 +177,16 @@ class TestRepositoryFactory: assert result is mock_repository_instance @patch("core.repositories.factory.dify_config") - def test_create_workflow_node_execution_repository_import_error(self, mock_config): + def test_create_workflow_node_execution_repository_import_error(self, mock_config, sqlite_session_factory): """Test WorkflowNodeExecutionRepository creation with import error.""" # Setup mock configuration with invalid class path mock_config.CORE_WORKFLOW_NODE_EXECUTION_REPOSITORY = "invalid.module.InvalidClass" - mock_session_factory = MagicMock(spec=sessionmaker) - mock_user = MagicMock(spec=EndUser) + mock_user = EndUser() with pytest.raises(RepositoryImportError) as exc_info: DifyCoreRepositoryFactory.create_workflow_node_execution_repository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id="test-app-id", @@ -191,13 +195,12 @@ class TestRepositoryFactory: assert "Failed to create WorkflowNodeExecutionRepository" in str(exc_info.value) @patch("core.repositories.factory.dify_config") - def test_create_workflow_node_execution_repository_instantiation_error(self, mock_config): + def test_create_workflow_node_execution_repository_instantiation_error(self, mock_config, sqlite_session_factory): """Test WorkflowNodeExecutionRepository creation with instantiation error.""" # Setup mock configuration mock_config.CORE_WORKFLOW_NODE_EXECUTION_REPOSITORY = "unittest.mock.MagicMock" - mock_session_factory = MagicMock(spec=sessionmaker) - mock_user = MagicMock(spec=EndUser) + mock_user = EndUser() # Create a mock repository class that raises exception on instantiation mock_repository_class = MagicMock() @@ -207,7 +210,7 @@ class TestRepositoryFactory: with patch("core.repositories.factory.import_string", return_value=mock_repository_class, autospec=True): with pytest.raises(RepositoryImportError) as exc_info: DifyCoreRepositoryFactory.create_workflow_node_execution_repository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id="test-app-id", @@ -222,14 +225,13 @@ class TestRepositoryFactory: assert str(error) == error_message @patch("core.repositories.factory.dify_config") - def test_create_with_engine_instead_of_sessionmaker(self, mock_config): + def test_create_with_engine_instead_of_sessionmaker(self, mock_config, sqlite_engine: Engine): """Test repository creation with Engine instead of sessionmaker.""" # Setup mock configuration mock_config.CORE_WORKFLOW_EXECUTION_REPOSITORY = "unittest.mock.MagicMock" - # Create mock dependencies using Engine instead of sessionmaker - mock_engine = MagicMock(spec=Engine) - mock_user = MagicMock(spec=Account) + # Pass the real Engine directly instead of wrapping it in sessionmaker + mock_user = Account(name="Test Account", email="test@example.com") app_id = "test-app-id" triggered_from = WorkflowRunTriggeredFrom.APP_RUN @@ -241,7 +243,7 @@ class TestRepositoryFactory: # Mock import_string with patch("core.repositories.factory.import_string", return_value=mock_repository_class, autospec=True): result = DifyCoreRepositoryFactory.create_workflow_execution_repository( - session_factory=mock_engine, # Using Engine instead of sessionmaker + session_factory=sqlite_engine, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id=app_id, @@ -250,7 +252,7 @@ class TestRepositoryFactory: # Verify the repository was created with correct parameters mock_repository_class.assert_called_once_with( - session_factory=mock_engine, + session_factory=sqlite_engine, tenant_id=RESOURCE_TENANT_ID, user=mock_user, app_id=app_id, diff --git a/api/tests/unit_tests/core/repositories/test_human_input_form_repository_impl.py b/api/tests/unit_tests/core/repositories/test_human_input_form_repository_impl.py index f8c6f6d81de..15925ccdc14 100644 --- a/api/tests/unit_tests/core/repositories/test_human_input_form_repository_impl.py +++ b/api/tests/unit_tests/core/repositories/test_human_input_form_repository_impl.py @@ -2,11 +2,12 @@ from __future__ import annotations -import dataclasses from datetime import datetime -from types import SimpleNamespace import pytest +from sqlalchemy import Engine, event +from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy.orm import Session from core.repositories.human_input_repository import ( HumanInputFormRecord, @@ -27,9 +28,11 @@ from core.workflow.nodes.human_input.entities import ( ) from core.workflow.nodes.human_input.enums import HumanInputFormKind, HumanInputFormStatus from libs.datetime_utils import naive_utc_now +from models import Account, TenantAccountJoin from models.human_input import ( EmailExternalRecipientPayload, EmailMemberRecipientPayload, + HumanInputForm, HumanInputFormRecipient, RecipientType, StandaloneWebAppRecipientPayload, @@ -40,55 +43,40 @@ def _build_repository() -> HumanInputFormRepositoryImpl: return HumanInputFormRepositoryImpl(tenant_id="tenant-id") -def _patch_recipient_factory(monkeypatch: pytest.MonkeyPatch) -> list[SimpleNamespace]: - created: list[SimpleNamespace] = [] - - def fake_new(cls, form_id: str, delivery_id: str, payload): # type: ignore[no-untyped-def] - recipient = SimpleNamespace( - form_id=form_id, - delivery_id=delivery_id, - recipient_type=payload.TYPE, - recipient_payload=payload.model_dump_json(), - ) - created.append(recipient) - return recipient - - monkeypatch.setattr(HumanInputFormRecipient, "new", classmethod(fake_new)) - return created - - -@pytest.fixture(autouse=True) -def _stub_selectinload(monkeypatch: pytest.MonkeyPatch) -> None: - """Avoid SQLAlchemy mapper configuration in tests using fake sessions.""" - - class _FakeSelect: - def options(self, *_args, **_kwargs): # type: ignore[no-untyped-def] - return self - - def where(self, *_args, **_kwargs): # type: ignore[no-untyped-def] - return self - - monkeypatch.setattr( - "core.repositories.human_input_repository.selectinload", lambda *args, **kwargs: "_loader_option" - ) - monkeypatch.setattr("core.repositories.human_input_repository.select", lambda *args, **kwargs: _FakeSelect()) +def _add_workspace_member( + session: Session, + *, + user_id: str, + email: str, + tenant_id: str = "tenant-id", +) -> None: + account = Account(name=user_id, email=email) + account.id = user_id + session.add_all([account, TenantAccountJoin(tenant_id=tenant_id, account_id=user_id)]) + session.commit() class TestHumanInputFormRepositoryImplHelpers: - def test_build_email_recipients_with_member_and_external(self, monkeypatch: pytest.MonkeyPatch) -> None: + def test_build_email_recipients_with_member_and_external( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ) -> None: repo = _build_repository() - session_stub = object() - _patch_recipient_factory(monkeypatch) + _add_workspace_member(sqlite_session, user_id="member-1", email="member@example.com") + query_workspace_members_by_ids = repo._query_workspace_members_by_ids - def fake_query(self, session, restrict_to_user_ids): # type: ignore[no-untyped-def] - assert session is session_stub - assert restrict_to_user_ids == ["member-1"] - return [_WorkspaceMemberInfo(user_id="member-1", email="member@example.com")] + def query_members(session: Session, restrict_to_user_ids: list[str]) -> list[_WorkspaceMemberInfo]: + assert session is sqlite_session + return query_workspace_members_by_ids( + session=session, + restrict_to_user_ids=restrict_to_user_ids, + ) - monkeypatch.setattr(HumanInputFormRepositoryImpl, "_query_workspace_members_by_ids", fake_query) + monkeypatch.setattr(repo, "_query_workspace_members_by_ids", query_members) recipients = repo._build_email_recipients( - session=session_stub, + session=sqlite_session, form_id="form-id", delivery_id="delivery-id", recipients_config=EmailRecipients( @@ -111,20 +99,11 @@ class TestHumanInputFormRepositoryImplHelpers: external_payload = EmailExternalRecipientPayload.model_validate_json(external_recipient.recipient_payload) assert external_payload.email == "external@example.com" - def test_build_email_recipients_skips_unknown_members(self, monkeypatch: pytest.MonkeyPatch) -> None: + def test_build_email_recipients_skips_unknown_members(self, sqlite_session: Session) -> None: repo = _build_repository() - session_stub = object() - created = _patch_recipient_factory(monkeypatch) - - def fake_query(self, session, restrict_to_user_ids): # type: ignore[no-untyped-def] - assert session is session_stub - assert restrict_to_user_ids == ["missing-member"] - return [] - - monkeypatch.setattr(HumanInputFormRepositoryImpl, "_query_workspace_members_by_ids", fake_query) recipients = repo._build_email_recipients( - session=session_stub, + session=sqlite_session, form_id="form-id", delivery_id="delivery-id", recipients_config=EmailRecipients( @@ -138,24 +117,14 @@ class TestHumanInputFormRepositoryImplHelpers: assert len(recipients) == 1 assert recipients[0].recipient_type == RecipientType.EMAIL_EXTERNAL - assert len(created) == 1 # only external recipient created via factory - def test_build_email_recipients_whole_workspace_uses_all_members(self, monkeypatch: pytest.MonkeyPatch) -> None: + def test_build_email_recipients_whole_workspace_uses_all_members(self, sqlite_session: Session) -> None: repo = _build_repository() - session_stub = object() - _patch_recipient_factory(monkeypatch) - - def fake_query(self, session): # type: ignore[no-untyped-def] - assert session is session_stub - return [ - _WorkspaceMemberInfo(user_id="member-1", email="member1@example.com"), - _WorkspaceMemberInfo(user_id="member-2", email="member2@example.com"), - ] - - monkeypatch.setattr(HumanInputFormRepositoryImpl, "_query_all_workspace_members", fake_query) + _add_workspace_member(sqlite_session, user_id="member-1", email="member1@example.com") + _add_workspace_member(sqlite_session, user_id="member-2", email="member2@example.com") recipients = repo._build_email_recipients( - session=session_stub, + session=sqlite_session, form_id="form-id", delivery_id="delivery-id", recipients_config=EmailRecipients( @@ -168,20 +137,11 @@ class TestHumanInputFormRepositoryImplHelpers: emails = {EmailMemberRecipientPayload.model_validate_json(r.recipient_payload).email for r in recipients} assert emails == {"member1@example.com", "member2@example.com"} - def test_build_email_recipients_dedupes_external_by_email(self, monkeypatch: pytest.MonkeyPatch) -> None: + def test_build_email_recipients_dedupes_external_by_email(self, sqlite_session: Session) -> None: repo = _build_repository() - session_stub = object() - created = _patch_recipient_factory(monkeypatch) - - def fake_query(self, session, restrict_to_user_ids): # type: ignore[no-untyped-def] - assert session is session_stub - assert restrict_to_user_ids == [] - return [] - - monkeypatch.setattr(HumanInputFormRepositoryImpl, "_query_workspace_members_by_ids", fake_query) recipients = repo._build_email_recipients( - session=session_stub, + session=sqlite_session, form_id="form-id", delivery_id="delivery-id", recipients_config=EmailRecipients( @@ -194,24 +154,13 @@ class TestHumanInputFormRepositoryImplHelpers: ) assert len(recipients) == 1 - assert len(created) == 1 - def test_build_email_recipients_prefers_member_over_external_by_email( - self, monkeypatch: pytest.MonkeyPatch - ) -> None: + def test_build_email_recipients_prefers_member_over_external_by_email(self, sqlite_session: Session) -> None: repo = _build_repository() - session_stub = object() - _patch_recipient_factory(monkeypatch) - - def fake_query(self, session, restrict_to_user_ids): # type: ignore[no-untyped-def] - assert session is session_stub - assert restrict_to_user_ids == ["member-1"] - return [_WorkspaceMemberInfo(user_id="member-1", email="shared@example.com")] - - monkeypatch.setattr(HumanInputFormRepositoryImpl, "_query_workspace_members_by_ids", fake_query) + _add_workspace_member(sqlite_session, user_id="member-1", email="shared@example.com") recipients = repo._build_email_recipients( - session=session_stub, + session=sqlite_session, form_id="form-id", delivery_id="delivery-id", recipients_config=EmailRecipients( @@ -228,20 +177,11 @@ class TestHumanInputFormRepositoryImplHelpers: def test_delivery_method_to_model_includes_external_recipients_with_whole_workspace( self, - monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: repo = _build_repository() - session_stub = object() - _patch_recipient_factory(monkeypatch) - - def fake_query(self, session): # type: ignore[no-untyped-def] - assert session is session_stub - return [ - _WorkspaceMemberInfo(user_id="member-1", email="member1@example.com"), - _WorkspaceMemberInfo(user_id="member-2", email="member2@example.com"), - ] - - monkeypatch.setattr(HumanInputFormRepositoryImpl, "_query_all_workspace_members", fake_query) + _add_workspace_member(sqlite_session, user_id="member-1", email="member1@example.com") + _add_workspace_member(sqlite_session, user_id="member-2", email="member2@example.com") method = EmailDeliveryMethod( config=EmailDeliveryConfig( @@ -254,7 +194,7 @@ class TestHumanInputFormRepositoryImplHelpers: ) ) - result = repo._delivery_method_to_model(session=session_stub, form_id="form-id", delivery_method=method) + result = repo._delivery_method_to_model(session=sqlite_session, form_id="form-id", delivery_method=method) assert len(result.recipients) == 3 member_emails = { @@ -279,155 +219,69 @@ def _make_form_definition() -> str: ).model_dump_json() -@dataclasses.dataclass -class _DummyForm: - id: str - workflow_run_id: str - node_id: str - tenant_id: str - app_id: str - form_definition: str - rendered_content: str - expiration_time: datetime - conversation_id: str | None = None - form_kind: HumanInputFormKind = HumanInputFormKind.RUNTIME - created_at: datetime = dataclasses.field(default_factory=naive_utc_now) - selected_action_id: str | None = None - submitted_data: str | None = None - submitted_at: datetime | None = None - submission_user_id: str | None = None - submission_end_user_id: str | None = None - completed_by_recipient_id: str | None = None - status: HumanInputFormStatus = HumanInputFormStatus.WAITING - - -@dataclasses.dataclass -class _DummyRecipient: - id: str - form_id: str - recipient_type: RecipientType - access_token: str - form: _DummyForm | None = None - recipient_payload: str = dataclasses.field( - default_factory=lambda: StandaloneWebAppRecipientPayload().model_dump_json() +def _make_form( + *, + form_id: str = "form-1", + workflow_run_id: str = "run-1", + node_id: str = "node-1", + tenant_id: str = "tenant-id", + selected_action_id: str | None = None, + submitted_data: str | None = None, + submitted_at: datetime | None = None, + expiration_time: datetime | None = None, +) -> HumanInputForm: + return HumanInputForm( + id=form_id, + workflow_run_id=workflow_run_id, + conversation_id=None, + node_id=node_id, + tenant_id=tenant_id, + app_id="app-id", + form_kind=HumanInputFormKind.RUNTIME, + form_definition=_make_form_definition(), + rendered_content="

hello

", + expiration_time=expiration_time or naive_utc_now(), + status=HumanInputFormStatus.WAITING, + selected_action_id=selected_action_id, + submitted_data=submitted_data, + submitted_at=submitted_at, ) -class _FakeScalarResult: - def __init__(self, obj): - self._obj = obj - - def first(self): - if isinstance(self._obj, list): - return self._obj[0] if self._obj else None - return self._obj - - def all(self): - if isinstance(self._obj, list): - return list(self._obj) - if self._obj is None: - return [] - return [self._obj] - - -class _FakeSession: - def __init__( - self, - *, - scalars_result=None, - scalars_results: list[object] | None = None, - forms: dict[str, _DummyForm] | None = None, - recipients: dict[str, _DummyRecipient] | None = None, - ): - if scalars_results is not None: - self._scalars_queue = list(scalars_results) - elif scalars_result is not None: - self._scalars_queue = [scalars_result] - else: - self._scalars_queue = [] - self.forms = forms or {} - self.recipients = recipients or {} - - def scalars(self, _query): - if self._scalars_queue: - result = self._scalars_queue.pop(0) - else: - result = None - return _FakeScalarResult(result) - - def get(self, model_cls, obj_id): # type: ignore[no-untyped-def] - if getattr(model_cls, "__name__", None) == "HumanInputForm": - return self.forms.get(obj_id) - if getattr(model_cls, "__name__", None) == "HumanInputFormRecipient": - return self.recipients.get(obj_id) - return None - - def add(self, _obj): - return None - - def flush(self): - return None - - def refresh(self, _obj): - return None - - def begin(self): - return self - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb): - return None - - -def _session_factory(session: _FakeSession): - class _SessionContext: - def __enter__(self): - return session - - def __exit__(self, exc_type, exc, tb): - return None - - def _factory(*_args, **_kwargs): - return _SessionContext() - - return _factory - - -def _patch_repo_session_factory(monkeypatch: pytest.MonkeyPatch, session: _FakeSession) -> None: - """Patch repository's global session factory to return our fake session. - - The repositories under test now use a global session factory; patch its - create_session method so unit tests don't hit a real database. - """ - monkeypatch.setattr( - "core.repositories.human_input_repository.session_factory.create_session", - _session_factory(session), - raising=True, +def _make_recipient( + form_id: str, + *, + recipient_id: str = "recipient-1", + recipient_type: RecipientType = RecipientType.STANDALONE_WEB_APP, + access_token: str = "token-123", +) -> HumanInputFormRecipient: + return HumanInputFormRecipient( + id=recipient_id, + form_id=form_id, + delivery_id="delivery-1", + recipient_type=recipient_type, + access_token=access_token, + recipient_payload=StandaloneWebAppRecipientPayload().model_dump_json(), ) +def _persist_form( + session: Session, + form: HumanInputForm, + recipients: list[HumanInputFormRecipient] | None = None, +) -> None: + session.add(form) + session.add_all(recipients or []) + session.commit() + + class TestHumanInputFormRepositoryImplPublicMethods: - def test_get_form_returns_entity_and_recipients(self, monkeypatch: pytest.MonkeyPatch): - form = _DummyForm( - id="form-1", - workflow_run_id="run-1", - node_id="node-1", - tenant_id="tenant-id", - app_id="app-id", - form_definition=_make_form_definition(), - rendered_content="

hello

", - expiration_time=naive_utc_now(), - ) - recipient = _DummyRecipient( - id="recipient-1", - form_id=form.id, - recipient_type=RecipientType.STANDALONE_WEB_APP, - access_token="token-123", - ) - session = _FakeSession(scalars_results=[form, [recipient]]) - _patch_repo_session_factory(monkeypatch, session) + def test_get_form_returns_entity_and_recipients(self, sqlite_session: Session): + form = _make_form() + recipient = _make_recipient(form.id) + other_tenant_form = _make_form(form_id="form-2", tenant_id="other-tenant") + sqlite_session.add(other_tenant_form) + _persist_form(sqlite_session, form, [recipient]) repo = HumanInputFormRepositoryImpl(tenant_id="tenant-id", workflow_execution_id=form.workflow_run_id) entity = repo.get_form(form.node_id) @@ -438,26 +292,14 @@ class TestHumanInputFormRepositoryImplPublicMethods: assert len(entity.recipients) == 1 assert entity.recipients[0].token == "token-123" - def test_get_form_returns_none_when_missing(self, monkeypatch: pytest.MonkeyPatch): - session = _FakeSession(scalars_results=[None]) - _patch_repo_session_factory(monkeypatch, session) + def test_get_form_returns_none_when_missing(self): repo = HumanInputFormRepositoryImpl(tenant_id="tenant-id", workflow_execution_id="run-1") assert repo.get_form("node-1") is None - def test_get_form_returns_unsubmitted_state(self, monkeypatch: pytest.MonkeyPatch): - form = _DummyForm( - id="form-1", - workflow_run_id="run-1", - node_id="node-1", - tenant_id="tenant-id", - app_id="app-id", - form_definition=_make_form_definition(), - rendered_content="

hello

", - expiration_time=naive_utc_now(), - ) - session = _FakeSession(scalars_results=[form, []]) - _patch_repo_session_factory(monkeypatch, session) + def test_get_form_returns_unsubmitted_state(self, sqlite_session: Session): + form = _make_form() + _persist_form(sqlite_session, form) repo = HumanInputFormRepositoryImpl(tenant_id="tenant-id", workflow_execution_id=form.workflow_run_id) entity = repo.get_form(form.node_id) @@ -467,22 +309,13 @@ class TestHumanInputFormRepositoryImplPublicMethods: assert entity.selected_action_id is None assert entity.submitted_data is None - def test_get_form_returns_submission_when_completed(self, monkeypatch: pytest.MonkeyPatch): - form = _DummyForm( - id="form-1", - workflow_run_id="run-1", - node_id="node-1", - tenant_id="tenant-id", - app_id="app-id", - form_definition=_make_form_definition(), - rendered_content="

hello

", - expiration_time=naive_utc_now(), + def test_get_form_returns_submission_when_completed(self, sqlite_session: Session): + form = _make_form( selected_action_id="approve", submitted_data='{"field": "value"}', submitted_at=naive_utc_now(), ) - session = _FakeSession(scalars_results=[form, []]) - _patch_repo_session_factory(monkeypatch, session) + _persist_form(sqlite_session, form) repo = HumanInputFormRepositoryImpl(tenant_id="tenant-id", workflow_execution_id=form.workflow_run_id) entity = repo.get_form(form.node_id) @@ -494,26 +327,10 @@ class TestHumanInputFormRepositoryImplPublicMethods: class TestHumanInputFormSubmissionRepository: - def test_get_by_token_returns_record(self, monkeypatch: pytest.MonkeyPatch): - form = _DummyForm( - id="form-1", - workflow_run_id="run-1", - node_id="node-1", - tenant_id="tenant-1", - app_id="app-1", - form_definition=_make_form_definition(), - rendered_content="

hello

", - expiration_time=naive_utc_now(), - ) - recipient = _DummyRecipient( - id="recipient-1", - form_id=form.id, - recipient_type=RecipientType.STANDALONE_WEB_APP, - access_token="token-123", - form=form, - ) - session = _FakeSession(scalars_result=recipient) - _patch_repo_session_factory(monkeypatch, session) + def test_get_by_token_returns_record(self, sqlite_session: Session): + form = _make_form(tenant_id="tenant-1") + recipient = _make_recipient(form.id) + _persist_form(sqlite_session, form, [recipient]) repo = HumanInputFormSubmissionRepository() record = repo.get_by_token("token-123") @@ -523,26 +340,10 @@ class TestHumanInputFormSubmissionRepository: assert record.recipient_type == RecipientType.STANDALONE_WEB_APP assert record.submitted is False - def test_get_by_form_id_and_recipient_type_uses_recipient(self, monkeypatch: pytest.MonkeyPatch): - form = _DummyForm( - id="form-1", - workflow_run_id="run-1", - node_id="node-1", - tenant_id="tenant-1", - app_id="app-1", - form_definition=_make_form_definition(), - rendered_content="

hello

", - expiration_time=naive_utc_now(), - ) - recipient = _DummyRecipient( - id="recipient-1", - form_id=form.id, - recipient_type=RecipientType.STANDALONE_WEB_APP, - access_token="token-123", - form=form, - ) - session = _FakeSession(scalars_result=recipient) - _patch_repo_session_factory(monkeypatch, session) + def test_get_by_form_id_and_recipient_type_uses_recipient(self, sqlite_session: Session): + form = _make_form(tenant_id="tenant-1") + recipient = _make_recipient(form.id) + _persist_form(sqlite_session, form, [recipient]) repo = HumanInputFormSubmissionRepository() record = repo.get_by_form_id_and_recipient_type( @@ -554,31 +355,17 @@ class TestHumanInputFormSubmissionRepository: assert record.recipient_id == recipient.id assert record.access_token == recipient.access_token - def test_mark_submitted_updates_fields(self, monkeypatch: pytest.MonkeyPatch): + def test_mark_submitted_updates_fields( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ): fixed_now = datetime(2024, 1, 1, 0, 0, 0) monkeypatch.setattr("core.repositories.human_input_repository.naive_utc_now", lambda: fixed_now) - form = _DummyForm( - id="form-1", - workflow_run_id="run-1", - node_id="node-1", - tenant_id="tenant-1", - app_id="app-1", - form_definition=_make_form_definition(), - rendered_content="

hello

", - expiration_time=fixed_now, - ) - recipient = _DummyRecipient( - id="recipient-1", - form_id="form-1", - recipient_type=RecipientType.STANDALONE_WEB_APP, - access_token="token-123", - ) - session = _FakeSession( - forms={form.id: form}, - recipients={recipient.id: recipient}, - ) - _patch_repo_session_factory(monkeypatch, session) + form = _make_form(tenant_id="tenant-1", expiration_time=fixed_now) + recipient = _make_recipient(form.id) + _persist_form(sqlite_session, form, [recipient]) repo = HumanInputFormSubmissionRepository() record: HumanInputFormRecord = repo.mark_submitted( @@ -590,11 +377,50 @@ class TestHumanInputFormSubmissionRepository: submission_end_user_id="end-user-1", ) - assert form.selected_action_id == "approve" - assert form.completed_by_recipient_id == recipient.id - assert form.submission_user_id == "user-1" - assert form.submission_end_user_id == "end-user-1" - assert form.submitted_at == fixed_now + sqlite_session.expire_all() + persisted_form = sqlite_session.get(HumanInputForm, form.id) + assert persisted_form is not None + assert persisted_form.selected_action_id == "approve" + assert persisted_form.completed_by_recipient_id == recipient.id + assert persisted_form.submission_user_id == "user-1" + assert persisted_form.submission_end_user_id == "end-user-1" + assert persisted_form.submitted_at == fixed_now assert record.submitted is True assert record.selected_action_id == "approve" assert record.submitted_data == {"field": "value"} + + def test_mark_submitted_rolls_back_on_database_failure( + self, + sqlite_session: Session, + sqlite_engine: Engine, + ) -> None: + form = _make_form(tenant_id="tenant-1") + recipient = _make_recipient(form.id) + _persist_form(sqlite_session, form, [recipient]) + repo = HumanInputFormSubmissionRepository() + + def fail_form_update(_connection, _cursor, statement, _parameters, _context, _executemany) -> None: + if statement.lstrip().upper().startswith("UPDATE HUMAN_INPUT_FORMS"): + raise SQLAlchemyError("forced update failure") + + event.listen(sqlite_engine, "before_cursor_execute", fail_form_update) + try: + with pytest.raises(SQLAlchemyError, match="forced update failure"): + repo.mark_submitted( + form_id=form.id, + recipient_id=recipient.id, + selected_action_id="approve", + form_data={"field": "value"}, + submission_user_id="user-1", + submission_end_user_id="end-user-1", + ) + finally: + event.remove(sqlite_engine, "before_cursor_execute", fail_form_update) + + sqlite_session.expire_all() + persisted_form = sqlite_session.get(HumanInputForm, form.id) + assert persisted_form is not None + assert persisted_form.status == HumanInputFormStatus.WAITING + assert persisted_form.selected_action_id is None + assert persisted_form.submitted_data is None + assert persisted_form.submitted_at is None diff --git a/api/tests/unit_tests/core/repositories/test_human_input_repository.py b/api/tests/unit_tests/core/repositories/test_human_input_repository.py index 5cb17a53e7b..6e149fdb697 100644 --- a/api/tests/unit_tests/core/repositories/test_human_input_repository.py +++ b/api/tests/unit_tests/core/repositories/test_human_input_repository.py @@ -1,14 +1,13 @@ from __future__ import annotations -import dataclasses import json -from collections.abc import Sequence +from collections.abc import Iterator, Sequence from datetime import datetime, timedelta -from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock import pytest +from sqlalchemy import Engine, select +from sqlalchemy.orm import Session, sessionmaker from core.repositories.human_input_repository import ( FormCreateParams, @@ -22,6 +21,7 @@ from core.repositories.human_input_repository import ( _WorkspaceMemberInfo, ) from core.workflow.human_input_adapter import ( + DeliveryMethodType, EmailDeliveryConfig, EmailDeliveryMethod, EmailRecipients, @@ -32,23 +32,35 @@ from core.workflow.human_input_adapter import ( from core.workflow.nodes.human_input.entities import HumanInputNodeData, UserActionConfig from core.workflow.nodes.human_input.enums import HumanInputFormKind, HumanInputFormStatus from libs.datetime_utils import naive_utc_now -from models.human_input import HumanInputFormRecipient, RecipientType +from models.account import Account, TenantAccountJoin, TenantAccountRole +from models.base import TypeBase +from models.human_input import ( + HumanInputDelivery, + HumanInputForm, + HumanInputFormRecipient, + RecipientType, +) -@pytest.fixture(autouse=True) -def _stub_select(monkeypatch: pytest.MonkeyPatch) -> None: - class _FakeSelect: - def join(self, *_args: Any, **_kwargs: Any) -> _FakeSelect: - return self +@pytest.fixture +def repository_session(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]: + """Bind repository-owned sessions to an isolated SQLite database.""" - def where(self, *_args: Any, **_kwargs: Any) -> _FakeSelect: - return self - - def options(self, *_args: Any, **_kwargs: Any) -> _FakeSelect: - return self - - monkeypatch.setattr("core.repositories.human_input_repository.select", lambda *_args, **_kwargs: _FakeSelect()) - monkeypatch.setattr("core.repositories.human_input_repository.selectinload", lambda *_args, **_kwargs: "_loader") + tables = [ + HumanInputForm.__table__, + HumanInputDelivery.__table__, + HumanInputFormRecipient.__table__, + Account.__table__, + TenantAccountJoin.__table__, + ] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + repository_session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr( + "core.repositories.human_input_repository.session_factory.create_session", + repository_session_factory, + ) + with repository_session_factory() as session: + yield session def _make_form_definition_json(*, include_expiration_time: bool) -> str: @@ -63,196 +75,116 @@ def _make_form_definition_json(*, include_expiration_time: bool) -> str: return json.dumps(payload, default=str) -@dataclasses.dataclass -class _DummyForm: - id: str - workflow_run_id: str | None - node_id: str - tenant_id: str - app_id: str - form_definition: str - rendered_content: str - expiration_time: datetime - conversation_id: str | None = None - form_kind: HumanInputFormKind = HumanInputFormKind.RUNTIME - created_at: datetime = dataclasses.field(default_factory=naive_utc_now) - selected_action_id: str | None = None - submitted_data: str | None = None - submitted_at: datetime | None = None - submission_user_id: str | None = None - submission_end_user_id: str | None = None - completed_by_recipient_id: str | None = None - status: HumanInputFormStatus = HumanInputFormStatus.WAITING +def _persist_form( + session: Session, + *, + form_id: str = "form-1", + tenant_id: str = "tenant", + workflow_run_id: str | None = "run", + node_id: str = "node", + status: HumanInputFormStatus = HumanInputFormStatus.WAITING, +) -> HumanInputForm: + form = HumanInputForm( + id=form_id, + tenant_id=tenant_id, + app_id="app", + workflow_run_id=workflow_run_id, + conversation_id=None, + form_kind=HumanInputFormKind.RUNTIME, + node_id=node_id, + form_definition=_make_form_definition_json(include_expiration_time=True), + rendered_content="

x

", + expiration_time=naive_utc_now() + timedelta(hours=1), + status=status, + ) + session.add(form) + session.commit() + return form -@dataclasses.dataclass -class _DummyRecipient: - id: str - form_id: str - recipient_type: RecipientType - access_token: str | None - - -class _FakeScalarResult: - def __init__(self, obj: Any): - self._obj = obj - - def first(self) -> Any: - if isinstance(self._obj, list): - return self._obj[0] if self._obj else None - return self._obj - - def all(self) -> list[Any]: - if self._obj is None: - return [] - if isinstance(self._obj, list): - return list(self._obj) - return [self._obj] - - -class _FakeExecuteResult: - def __init__(self, rows: Sequence[tuple[Any, ...]]): - self._rows = list(rows) - - def all(self) -> list[tuple[Any, ...]]: - return list(self._rows) - - -class _FakeSession: - def __init__( - self, - *, - scalars_result: Any = None, - scalars_results: list[Any] | None = None, - forms: dict[str, _DummyForm] | None = None, - recipients: dict[str, _DummyRecipient] | None = None, - execute_rows: Sequence[tuple[Any, ...]] = (), - ): - if scalars_results is not None: - self._scalars_queue = list(scalars_results) - else: - self._scalars_queue = [scalars_result] - self._forms = forms or {} - self._recipients = recipients or {} - self._execute_rows = list(execute_rows) - self.added: list[Any] = [] - - def scalars(self, _query: Any) -> _FakeScalarResult: - if self._scalars_queue: - value = self._scalars_queue.pop(0) - else: - value = None - return _FakeScalarResult(value) - - def execute(self, _stmt: Any) -> _FakeExecuteResult: - return _FakeExecuteResult(self._execute_rows) - - def get(self, model_cls: Any, obj_id: str) -> Any: - name = getattr(model_cls, "__name__", "") - if name == "HumanInputForm": - return self._forms.get(obj_id) - if name == "HumanInputFormRecipient": - return self._recipients.get(obj_id) - return None - - def add(self, obj: Any) -> None: - self.added.append(obj) - - def add_all(self, objs: Sequence[Any]) -> None: - self.added.extend(list(objs)) - - def flush(self) -> None: - # Simulate DB default population for attributes referenced in entity wrappers. - for obj in self.added: - if hasattr(obj, "id") and obj.id in (None, ""): - obj.id = f"gen-{len(str(self.added))}" - if isinstance(obj, HumanInputFormRecipient) and obj.access_token is None: - if obj.recipient_type == RecipientType.CONSOLE: - obj.access_token = "token-console" - elif obj.recipient_type == RecipientType.BACKSTAGE: - obj.access_token = "token-backstage" - else: - obj.access_token = "token-webapp" - - def refresh(self, _obj: Any) -> None: - return None - - def begin(self) -> _FakeSession: - return self - - def __enter__(self) -> _FakeSession: - return self - - def __exit__(self, exc_type, exc, tb) -> None: - return None - - -class _SessionFactoryStub: - def __init__(self, session: _FakeSession): - self._session = session - - def create_session(self) -> _FakeSession: - return self._session - - -def _patch_session_factory(monkeypatch: pytest.MonkeyPatch, session: _FakeSession) -> None: - monkeypatch.setattr("core.repositories.human_input_repository.session_factory", _SessionFactoryStub(session)) +def _persist_recipient( + session: Session, + *, + form_id: str, + recipient_id: str = "recipient-1", + recipient_type: RecipientType = RecipientType.STANDALONE_WEB_APP, + access_token: str = "token-1", +) -> HumanInputFormRecipient: + delivery = HumanInputDelivery( + id=f"delivery-{recipient_id}", + form_id=form_id, + delivery_method_type=DeliveryMethodType.WEBAPP, + delivery_config_id=None, + channel_payload="{}", + ) + recipient = HumanInputFormRecipient( + id=recipient_id, + form_id=form_id, + delivery_id=delivery.id, + recipient_type=recipient_type, + recipient_payload="{}", + access_token=access_token, + ) + session.add_all([delivery, recipient]) + session.commit() + return recipient def test_recipient_entity_token_raises_when_missing() -> None: - recipient = SimpleNamespace(id="r1", access_token=None) - entity = _HumanInputFormRecipientEntityImpl(recipient) # type: ignore[arg-type] + recipient = HumanInputFormRecipient( + id="r1", + form_id="f1", + delivery_id="d1", + recipient_type=RecipientType.CONSOLE, + recipient_payload="{}", + access_token=None, + ) + entity = _HumanInputFormRecipientEntityImpl(recipient) with pytest.raises(AssertionError, match="access_token should not be None"): _ = entity.token -def test_recipient_entity_id_and_token_success() -> None: - recipient = SimpleNamespace(id="r1", access_token="tok") - entity = _HumanInputFormRecipientEntityImpl(recipient) # type: ignore[arg-type] +def test_recipient_entity_id_and_token_success(repository_session: Session) -> None: + form = _persist_form(repository_session) + recipient = _persist_recipient(repository_session, form_id=form.id, recipient_id="r1", access_token="tok") + entity = _HumanInputFormRecipientEntityImpl(recipient) assert entity.id == "r1" assert entity.token == "tok" -def test_form_entity_submission_token_prefers_console_then_webapp_then_none() -> None: - form = _DummyForm( - id="f1", - workflow_run_id="run", - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=naive_utc_now(), +def test_form_entity_submission_token_prefers_console_then_webapp_then_none(repository_session: Session) -> None: + form = _persist_form(repository_session, form_id="f1") + console = _persist_recipient( + repository_session, + form_id=form.id, + recipient_id="c1", + recipient_type=RecipientType.CONSOLE, + access_token="ctok", ) - console = _DummyRecipient(id="c1", form_id=form.id, recipient_type=RecipientType.CONSOLE, access_token="ctok") - webapp = _DummyRecipient( - id="w1", form_id=form.id, recipient_type=RecipientType.STANDALONE_WEB_APP, access_token="wtok" + webapp = _persist_recipient( + repository_session, + form_id=form.id, + recipient_id="w1", + recipient_type=RecipientType.STANDALONE_WEB_APP, + access_token="wtok", ) - entity = _HumanInputFormEntityImpl(form_model=form, recipient_models=[webapp, console]) # type: ignore[arg-type] + entity = _HumanInputFormEntityImpl(form_model=form, recipient_models=[webapp, console]) assert entity.submission_token == "ctok" - entity = _HumanInputFormEntityImpl(form_model=form, recipient_models=[webapp]) # type: ignore[arg-type] + entity = _HumanInputFormEntityImpl(form_model=form, recipient_models=[webapp]) assert entity.submission_token == "wtok" - entity = _HumanInputFormEntityImpl(form_model=form, recipient_models=[]) # type: ignore[arg-type] + entity = _HumanInputFormEntityImpl(form_model=form, recipient_models=[]) assert entity.submission_token is None -def test_form_entity_submitted_data_parsed() -> None: - form = _DummyForm( - id="f1", - workflow_run_id="run", - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=naive_utc_now(), - submitted_data='{"a": 1}', - submitted_at=naive_utc_now(), - ) - entity = _HumanInputFormEntityImpl(form_model=form, recipient_models=[]) # type: ignore[arg-type] +def test_form_entity_submitted_data_parsed(repository_session: Session) -> None: + form = _persist_form(repository_session, form_id="f1") + form.submitted_data = '{"a": 1}' + form.submitted_at = naive_utc_now() + repository_session.commit() + entity = _HumanInputFormEntityImpl(form_model=form, recipient_models=[]) assert entity.submitted is True assert entity.submitted_data == {"a": 1} assert entity.rendered_content == "

x

" @@ -260,42 +192,20 @@ def test_form_entity_submitted_data_parsed() -> None: assert entity.status == HumanInputFormStatus.WAITING -def test_form_record_from_models_injects_expiration_time_when_missing() -> None: +def test_form_record_from_models_injects_expiration_time_when_missing(repository_session: Session) -> None: expiration = naive_utc_now() - form = _DummyForm( - id="f1", - workflow_run_id=None, - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=False), - rendered_content="

x

", - expiration_time=expiration, - submitted_data='{"k": "v"}', - ) - record = HumanInputFormRecord.from_models(form, None) # type: ignore[arg-type] + form = _persist_form(repository_session, form_id="f1", workflow_run_id=None) + form.form_definition = _make_form_definition_json(include_expiration_time=False) + form.expiration_time = expiration + form.submitted_data = '{"k": "v"}' + repository_session.commit() + record = HumanInputFormRecord.from_models(form, None) assert record.definition.expiration_time == expiration assert record.submitted_data == {"k": "v"} assert record.submitted is False -def test_create_email_recipients_from_resolved_dedupes_and_skips_blank(monkeypatch: pytest.MonkeyPatch) -> None: - created: list[SimpleNamespace] = [] - - def fake_new(cls, form_id: str, delivery_id: str, payload: Any): # type: ignore[no-untyped-def] - recipient = SimpleNamespace( - id=f"{payload.TYPE}-{len(created)}", - form_id=form_id, - delivery_id=delivery_id, - recipient_type=payload.TYPE, - recipient_payload=payload.model_dump_json(), - access_token="tok", - ) - created.append(recipient) - return recipient - - monkeypatch.setattr("core.repositories.human_input_repository.HumanInputFormRecipient.new", classmethod(fake_new)) - +def test_create_email_recipients_from_resolved_dedupes_and_skips_blank() -> None: repo = HumanInputFormRepositoryImpl(tenant_id="tenant") recipients = repo._create_email_recipients_from_resolved( # type: ignore[attr-defined] form_id="f", @@ -310,26 +220,47 @@ def test_create_email_recipients_from_resolved_dedupes_and_skips_blank(monkeypat assert [r.recipient_type for r in recipients] == [RecipientType.EMAIL_MEMBER, RecipientType.EMAIL_EXTERNAL] -def test_query_workspace_members_by_ids_empty_returns_empty() -> None: +def test_query_workspace_members_by_ids_empty_returns_empty(repository_session: Session) -> None: repo = HumanInputFormRepositoryImpl(tenant_id="tenant") - assert repo._query_workspace_members_by_ids(session=MagicMock(), restrict_to_user_ids=["", ""]) == [] + assert repo._query_workspace_members_by_ids(session=repository_session, restrict_to_user_ids=["", ""]) == [] -def test_query_workspace_members_by_ids_maps_rows() -> None: - session = _FakeSession(execute_rows=[("u1", "a@example.com"), ("u2", "b@example.com")]) - repo = HumanInputFormRepositoryImpl(tenant_id="tenant") - rows = repo._query_workspace_members_by_ids(session=session, restrict_to_user_ids=["u1", "u2"]) - assert rows == [ - _WorkspaceMemberInfo(user_id="u1", email="a@example.com"), - _WorkspaceMemberInfo(user_id="u2", email="b@example.com"), +def test_query_workspace_members_by_ids_maps_rows_and_scopes_tenant(repository_session: Session) -> None: + accounts = [ + Account(name="One", email="a@example.com"), + Account(name="Two", email="b@example.com"), + Account(name="Other", email="other@example.com"), ] - - -def test_query_all_workspace_members_maps_rows() -> None: - session = _FakeSession(execute_rows=[("u1", "a@example.com")]) + repository_session.add_all(accounts) + repository_session.flush() + repository_session.add_all( + [ + TenantAccountJoin(tenant_id="tenant", account_id=accounts[0].id, role=TenantAccountRole.NORMAL), + TenantAccountJoin(tenant_id="tenant", account_id=accounts[1].id, role=TenantAccountRole.NORMAL), + TenantAccountJoin(tenant_id="other-tenant", account_id=accounts[2].id, role=TenantAccountRole.NORMAL), + ] + ) + repository_session.commit() repo = HumanInputFormRepositoryImpl(tenant_id="tenant") - rows = repo._query_all_workspace_members(session=session) - assert rows == [_WorkspaceMemberInfo(user_id="u1", email="a@example.com")] + rows = repo._query_workspace_members_by_ids( + session=repository_session, + restrict_to_user_ids=[accounts[0].id, accounts[1].id, accounts[2].id], + ) + assert set(rows) == { + _WorkspaceMemberInfo(user_id=accounts[0].id, email="a@example.com"), + _WorkspaceMemberInfo(user_id=accounts[1].id, email="b@example.com"), + } + + +def test_query_all_workspace_members_maps_rows(repository_session: Session) -> None: + account = Account(name="One", email="a@example.com") + repository_session.add(account) + repository_session.flush() + repository_session.add(TenantAccountJoin(tenant_id="tenant", account_id=account.id, role=TenantAccountRole.NORMAL)) + repository_session.commit() + repo = HumanInputFormRepositoryImpl(tenant_id="tenant") + rows = repo._query_all_workspace_members(session=repository_session) + assert rows == [_WorkspaceMemberInfo(user_id=account.id, email="a@example.com")] def test_repository_init_sets_tenant_id() -> None: @@ -337,11 +268,13 @@ def test_repository_init_sets_tenant_id() -> None: assert repo._tenant_id == "tenant" -def test_delivery_method_to_model_webapp_creates_delivery_and_recipient(monkeypatch: pytest.MonkeyPatch) -> None: +def test_delivery_method_to_model_webapp_creates_delivery_and_recipient( + repository_session: Session, monkeypatch: pytest.MonkeyPatch +) -> None: repo = HumanInputFormRepositoryImpl(tenant_id="tenant") monkeypatch.setattr("core.repositories.human_input_repository.uuidv7", lambda: "del-1") result = repo._delivery_method_to_model( - session=MagicMock(), form_id="form-1", delivery_method=WebAppDeliveryMethod() + session=repository_session, form_id="form-1", delivery_method=WebAppDeliveryMethod() ) assert result.delivery.id == "del-1" assert result.delivery.form_id == "form-1" @@ -349,12 +282,14 @@ def test_delivery_method_to_model_webapp_creates_delivery_and_recipient(monkeypa assert result.recipients[0].recipient_type == RecipientType.STANDALONE_WEB_APP -def test_delivery_method_to_model_email_uses_build_email_recipients(monkeypatch: pytest.MonkeyPatch) -> None: +def test_delivery_method_to_model_email_uses_build_email_recipients( + repository_session: Session, monkeypatch: pytest.MonkeyPatch +) -> None: repo = HumanInputFormRepositoryImpl(tenant_id="tenant") monkeypatch.setattr("core.repositories.human_input_repository.uuidv7", lambda: "del-1") called: dict[str, Any] = {} - def fake_build(*, session: Any, form_id: str, delivery_id: str, recipients_config: Any) -> list[Any]: + def fake_build(*, session: Session, form_id: str, delivery_id: str, recipients_config: Any) -> list[Any]: called.update( {"session": session, "form_id": form_id, "delivery_id": delivery_id, "recipients_config": recipients_config} ) @@ -372,12 +307,15 @@ def test_delivery_method_to_model_email_uses_build_email_recipients(monkeypatch: body="b", ) ) - result = repo._delivery_method_to_model(session="sess", form_id="form-1", delivery_method=method) + result = repo._delivery_method_to_model(session=repository_session, form_id="form-1", delivery_method=method) assert result.recipients == ["r"] + assert called["session"] is repository_session assert called["delivery_id"] == "del-1" -def test_build_email_recipients_uses_all_members_when_whole_workspace(monkeypatch: pytest.MonkeyPatch) -> None: +def test_build_email_recipients_uses_all_members_when_whole_workspace( + repository_session: Session, monkeypatch: pytest.MonkeyPatch +) -> None: repo = HumanInputFormRepositoryImpl(tenant_id="tenant") monkeypatch.setattr( repo, @@ -386,7 +324,7 @@ def test_build_email_recipients_uses_all_members_when_whole_workspace(monkeypatc ) monkeypatch.setattr(repo, "_create_email_recipients_from_resolved", lambda **_: ["ok"]) recipients = repo._build_email_recipients( - session=MagicMock(), + session=repository_session, form_id="f", delivery_id="d", recipients_config=EmailRecipients(include_bound_group=True, items=[ExternalRecipient(email="e@example.com")]), @@ -394,7 +332,9 @@ def test_build_email_recipients_uses_all_members_when_whole_workspace(monkeypatc assert recipients == ["ok"] -def test_build_email_recipients_uses_selected_members_when_not_whole_workspace(monkeypatch: pytest.MonkeyPatch) -> None: +def test_build_email_recipients_uses_selected_members_when_not_whole_workspace( + repository_session: Session, monkeypatch: pytest.MonkeyPatch +) -> None: repo = HumanInputFormRepositoryImpl(tenant_id="tenant") def fake_query(*, session: Any, restrict_to_user_ids: Sequence[str]) -> list[_WorkspaceMemberInfo]: @@ -404,7 +344,7 @@ def test_build_email_recipients_uses_selected_members_when_not_whole_workspace(m monkeypatch.setattr(repo, "_query_workspace_members_by_ids", fake_query) monkeypatch.setattr(repo, "_create_email_recipients_from_resolved", lambda **_: ["ok"]) recipients = repo._build_email_recipients( - session=MagicMock(), + session=repository_session, form_id="f", delivery_id="d", recipients_config=EmailRecipients( @@ -415,29 +355,13 @@ def test_build_email_recipients_uses_selected_members_when_not_whole_workspace(m assert recipients == ["ok"] -def test_get_form_returns_entity_and_none_when_missing(monkeypatch: pytest.MonkeyPatch) -> None: - _patch_session_factory(monkeypatch, _FakeSession(scalars_results=[None])) +def test_get_form_returns_entity_and_none_when_missing(repository_session: Session) -> None: repo = HumanInputFormRepositoryImpl(tenant_id="tenant", workflow_execution_id="run") assert repo.get_form("node") is None - form = _DummyForm( - id="f1", - workflow_run_id="run", - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=naive_utc_now(), - ) - recipient = _DummyRecipient( - id="r1", - form_id=form.id, - recipient_type=RecipientType.STANDALONE_WEB_APP, - access_token="tok", - ) - session = _FakeSession(scalars_results=[form, [recipient]]) - _patch_session_factory(monkeypatch, session) + form = _persist_form(repository_session, form_id="f1") + _persist_form(repository_session, form_id="other-tenant-form", tenant_id="other-tenant") + recipient = _persist_recipient(repository_session, form_id=form.id, recipient_id="r1", access_token="tok") repo = HumanInputFormRepositoryImpl(tenant_id="tenant", workflow_execution_id="run") entity = repo.get_form("node") assert entity is not None @@ -446,15 +370,15 @@ def test_get_form_returns_entity_and_none_when_missing(monkeypatch: pytest.Monke assert entity.recipients[0].token == "tok" -def test_create_form_adds_console_and_backstage_recipients(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_form_adds_console_and_backstage_recipients( + repository_session: Session, monkeypatch: pytest.MonkeyPatch +) -> None: fixed_now = datetime(2024, 1, 1, 0, 0, 0) monkeypatch.setattr("core.repositories.human_input_repository.naive_utc_now", lambda: fixed_now) ids = iter(["form-id", "del-web", "del-console", "del-backstage"]) monkeypatch.setattr("core.repositories.human_input_repository.uuidv7", lambda: next(ids)) - session = _FakeSession() - _patch_session_factory(monkeypatch, session) repo = HumanInputFormRepositoryImpl( tenant_id="tenant", app_id="app", @@ -485,19 +409,28 @@ def test_create_form_adds_console_and_backstage_recipients(monkeypatch: pytest.M assert entity.id == "form-id" assert entity.expiration_time == fixed_now + timedelta(hours=form_config.timeout) # Console token should take precedence when console recipient is present. - assert entity.submission_token == "token-console" + assert entity.submission_token is not None assert len(entity.recipients) == 3 + repository_session.expire_all() + persisted_form = repository_session.get(HumanInputForm, "form-id") + assert persisted_form is not None + assert persisted_form.status == HumanInputFormStatus.WAITING + recipients = repository_session.scalars( + select(HumanInputFormRecipient).where(HumanInputFormRecipient.form_id == "form-id") + ).all() + assert {recipient.recipient_type for recipient in recipients} == { + RecipientType.STANDALONE_WEB_APP, + RecipientType.CONSOLE, + RecipientType.BACKSTAGE, + } -def test_submission_get_by_token_returns_none_when_missing_or_form_missing(monkeypatch: pytest.MonkeyPatch) -> None: - _patch_session_factory(monkeypatch, _FakeSession(scalars_result=None)) +def test_submission_get_by_token_returns_none_when_missing_or_form_missing(repository_session: Session) -> None: repo = HumanInputFormSubmissionRepository() assert repo.get_by_token("tok") is None - recipient = SimpleNamespace(form=None) - _patch_session_factory(monkeypatch, _FakeSession(scalars_result=recipient)) - repo = HumanInputFormSubmissionRepository() - assert repo.get_by_token("tok") is None + _persist_recipient(repository_session, form_id="missing-form", access_token="orphan-token") + assert repo.get_by_token("orphan-token") is None def test_submission_repository_init_no_args() -> None: @@ -505,50 +438,36 @@ def test_submission_repository_init_no_args() -> None: assert isinstance(repo, HumanInputFormSubmissionRepository) -def test_submission_get_by_token_and_get_by_form_id_success_paths(monkeypatch: pytest.MonkeyPatch) -> None: - form = _DummyForm( - id="f1", - workflow_run_id=None, - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=naive_utc_now(), - ) - recipient = SimpleNamespace( - id="r1", - form_id=form.id, - recipient_type=RecipientType.STANDALONE_WEB_APP, - access_token="tok", - form=form, - ) - - _patch_session_factory(monkeypatch, _FakeSession(scalars_result=recipient)) +def test_submission_get_by_token_and_get_by_form_id_success_paths(repository_session: Session) -> None: + form = _persist_form(repository_session, form_id="f1", workflow_run_id=None) + recipient = _persist_recipient(repository_session, form_id=form.id, recipient_id="r1", access_token="tok") repo = HumanInputFormSubmissionRepository() record = repo.get_by_token("tok") assert record is not None assert record.access_token == "tok" - _patch_session_factory(monkeypatch, _FakeSession(scalars_result=recipient)) - repo = HumanInputFormSubmissionRepository() record = repo.get_by_form_id_and_recipient_type(form_id=form.id, recipient_type=RecipientType.STANDALONE_WEB_APP) assert record is not None assert record.recipient_id == "r1" + record = repo.get_by_form_id(form.id) + assert record is not None + assert record.form_id == form.id + assert record.recipient_id is None -def test_submission_get_by_form_id_returns_none_on_missing(monkeypatch: pytest.MonkeyPatch) -> None: - _patch_session_factory(monkeypatch, _FakeSession(scalars_result=None)) + +def test_submission_get_by_form_id_returns_none_on_missing(repository_session: Session) -> None: repo = HumanInputFormSubmissionRepository() assert repo.get_by_form_id_and_recipient_type(form_id="f", recipient_type=RecipientType.CONSOLE) is None + assert repo.get_by_form_id("f") is None -def test_mark_submitted_updates_and_raises_when_missing(monkeypatch: pytest.MonkeyPatch) -> None: +def test_mark_submitted_updates_and_raises_when_missing( + repository_session: Session, monkeypatch: pytest.MonkeyPatch +) -> None: fixed_now = datetime(2024, 1, 1, 0, 0, 0) monkeypatch.setattr("core.repositories.human_input_repository.naive_utc_now", lambda: fixed_now) - missing_session = _FakeSession(forms={}) - _patch_session_factory(monkeypatch, missing_session) repo = HumanInputFormSubmissionRepository() with pytest.raises(FormNotFoundError, match="form not found"): repo.mark_submitted( @@ -560,20 +479,14 @@ def test_mark_submitted_updates_and_raises_when_missing(monkeypatch: pytest.Monk submission_end_user_id=None, ) - form = _DummyForm( - id="f", - workflow_run_id=None, - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=fixed_now, + form = _persist_form(repository_session, form_id="f", workflow_run_id=None) + recipient = _persist_recipient( + repository_session, + form_id=form.id, + recipient_id="r", + recipient_type=RecipientType.CONSOLE, + access_token="tok", ) - recipient = _DummyRecipient(id="r", form_id=form.id, recipient_type=RecipientType.CONSOLE, access_token="tok") - session = _FakeSession(forms={form.id: form}, recipients={recipient.id: recipient}) - _patch_session_factory(monkeypatch, session) - repo = HumanInputFormSubmissionRepository() record = repo.mark_submitted( form_id=form.id, recipient_id=recipient.id, @@ -582,33 +495,29 @@ def test_mark_submitted_updates_and_raises_when_missing(monkeypatch: pytest.Monk submission_user_id="u", submission_end_user_id="eu", ) - assert form.status == HumanInputFormStatus.SUBMITTED - assert form.submitted_at == fixed_now + repository_session.expire_all() + persisted = repository_session.get(HumanInputForm, form.id) + assert persisted is not None + assert persisted.status == HumanInputFormStatus.SUBMITTED + assert persisted.submitted_at == fixed_now + assert persisted.completed_by_recipient_id == recipient.id assert record.submitted_data == {"k": "v"} -def test_mark_submitted_serializes_select_and_file_payloads(monkeypatch: pytest.MonkeyPatch) -> None: +def test_mark_submitted_serializes_select_and_file_payloads( + repository_session: Session, monkeypatch: pytest.MonkeyPatch +) -> None: fixed_now = datetime(2024, 1, 1, 0, 0, 0) monkeypatch.setattr("core.repositories.human_input_repository.naive_utc_now", lambda: fixed_now) - form = _DummyForm( - id="f-complex", - workflow_run_id=None, - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=fixed_now, - ) - recipient = _DummyRecipient( - id="r-complex", + form = _persist_form(repository_session, form_id="f-complex", workflow_run_id=None) + recipient = _persist_recipient( + repository_session, form_id=form.id, + recipient_id="r-complex", recipient_type=RecipientType.CONSOLE, access_token="tok", ) - session = _FakeSession(forms={form.id: form}, recipients={recipient.id: recipient}) - _patch_session_factory(monkeypatch, session) payload = { "decision": "approve", @@ -650,98 +559,71 @@ def test_mark_submitted_serializes_select_and_file_payloads(monkeypatch: pytest. submission_end_user_id="end-user-1", ) - assert json.loads(form.submitted_data or "") == payload + repository_session.expire_all() + persisted = repository_session.get(HumanInputForm, form.id) + assert persisted is not None + assert json.loads(persisted.submitted_data or "") == payload assert record.submitted_data == payload -def test_mark_timeout_invalid_status_raises(monkeypatch: pytest.MonkeyPatch) -> None: - form = _DummyForm( - id="f", - workflow_run_id=None, - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=naive_utc_now(), - ) - session = _FakeSession(forms={form.id: form}) - _patch_session_factory(monkeypatch, session) +def test_mark_timeout_invalid_status_rolls_back(repository_session: Session) -> None: + form = _persist_form(repository_session, form_id="f", workflow_run_id=None) repo = HumanInputFormSubmissionRepository() with pytest.raises(_InvalidTimeoutStatusError, match="invalid timeout status"): repo.mark_timeout(form_id=form.id, timeout_status=HumanInputFormStatus.SUBMITTED) # type: ignore[arg-type] + repository_session.expire_all() + persisted = repository_session.get(HumanInputForm, form.id) + assert persisted is not None + assert persisted.status == HumanInputFormStatus.WAITING -def test_mark_timeout_already_timed_out_returns_record(monkeypatch: pytest.MonkeyPatch) -> None: - form = _DummyForm( - id="f", +def test_mark_timeout_already_timed_out_returns_record(repository_session: Session) -> None: + form = _persist_form( + repository_session, + form_id="f", workflow_run_id=None, - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=naive_utc_now(), status=HumanInputFormStatus.TIMEOUT, ) - session = _FakeSession(forms={form.id: form}) - _patch_session_factory(monkeypatch, session) repo = HumanInputFormSubmissionRepository() record = repo.mark_timeout(form_id=form.id, timeout_status=HumanInputFormStatus.TIMEOUT, reason="r") assert record.status == HumanInputFormStatus.TIMEOUT -def test_mark_timeout_submitted_raises_form_not_found(monkeypatch: pytest.MonkeyPatch) -> None: - form = _DummyForm( - id="f", +def test_mark_timeout_submitted_raises_form_not_found(repository_session: Session) -> None: + form = _persist_form( + repository_session, + form_id="f", workflow_run_id=None, - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=naive_utc_now(), status=HumanInputFormStatus.SUBMITTED, ) - session = _FakeSession(forms={form.id: form}) - _patch_session_factory(monkeypatch, session) repo = HumanInputFormSubmissionRepository() with pytest.raises(FormNotFoundError, match="form already submitted"): repo.mark_timeout(form_id=form.id, timeout_status=HumanInputFormStatus.EXPIRED) -def test_mark_timeout_updates_fields(monkeypatch: pytest.MonkeyPatch) -> None: - form = _DummyForm( - id="f", - workflow_run_id=None, - node_id="node", - tenant_id="tenant", - app_id="app", - form_definition=_make_form_definition_json(include_expiration_time=True), - rendered_content="

x

", - expiration_time=naive_utc_now(), - selected_action_id="a", - submitted_data="{}", - submission_user_id="u", - submission_end_user_id="eu", - completed_by_recipient_id="r", - status=HumanInputFormStatus.WAITING, - ) - session = _FakeSession(forms={form.id: form}) - _patch_session_factory(monkeypatch, session) +def test_mark_timeout_updates_fields(repository_session: Session) -> None: + form = _persist_form(repository_session, form_id="f", workflow_run_id=None) + form.selected_action_id = "a" + form.submitted_data = "{}" + form.submission_user_id = "u" + form.submission_end_user_id = "eu" + form.completed_by_recipient_id = "r" + repository_session.commit() repo = HumanInputFormSubmissionRepository() record = repo.mark_timeout(form_id=form.id, timeout_status=HumanInputFormStatus.EXPIRED) - assert form.status == HumanInputFormStatus.EXPIRED - assert form.selected_action_id is None - assert form.submitted_data is None - assert form.submission_user_id is None - assert form.submission_end_user_id is None - assert form.completed_by_recipient_id is None + repository_session.expire_all() + persisted = repository_session.get(HumanInputForm, form.id) + assert persisted is not None + assert persisted.status == HumanInputFormStatus.EXPIRED + assert persisted.selected_action_id is None + assert persisted.submitted_data is None + assert persisted.submission_user_id is None + assert persisted.submission_end_user_id is None + assert persisted.completed_by_recipient_id is None assert record.status == HumanInputFormStatus.EXPIRED -def test_mark_timeout_raises_when_form_missing(monkeypatch: pytest.MonkeyPatch) -> None: - _patch_session_factory(monkeypatch, _FakeSession(forms={})) +def test_mark_timeout_raises_when_form_missing(repository_session: Session) -> None: repo = HumanInputFormSubmissionRepository() with pytest.raises(FormNotFoundError, match="form not found"): repo.mark_timeout(form_id="missing", timeout_status=HumanInputFormStatus.TIMEOUT) diff --git a/api/tests/unit_tests/core/repositories/test_sqlalchemy_workflow_execution_repository.py b/api/tests/unit_tests/core/repositories/test_sqlalchemy_workflow_execution_repository.py index f247525aedd..01c4905e4cb 100644 --- a/api/tests/unit_tests/core/repositories/test_sqlalchemy_workflow_execution_repository.py +++ b/api/tests/unit_tests/core/repositories/test_sqlalchemy_workflow_execution_repository.py @@ -1,56 +1,59 @@ +import json from datetime import UTC, datetime -from unittest.mock import MagicMock from uuid import uuid4 import pytest from sqlalchemy.engine import Engine -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.repositories.sqlalchemy_workflow_execution_repository import SQLAlchemyWorkflowExecutionRepository from graphon.entities import WorkflowExecution from graphon.enums import WorkflowExecutionStatus, WorkflowType -from models import Account, CreatorUserRole, EndUser, WorkflowRun -from models.enums import WorkflowRunTriggeredFrom +from models import Account, CreatorUserRole, EndUser, Tenant, WorkflowRun +from models.enums import EndUserType, WorkflowRunTriggeredFrom +from models.workflow import WorkflowType as ModelWorkflowType + +TABLES = (WorkflowRun,) RESOURCE_TENANT_ID = "resource-tenant-id" -@pytest.fixture -def mock_session_factory(): - """Mock SQLAlchemy session factory.""" - session_factory = MagicMock(spec=sessionmaker) - session = MagicMock() - session.get.return_value = None - session_factory.return_value.__enter__.return_value = session - return session_factory - - -@pytest.fixture -def mock_engine(): - """Mock SQLAlchemy Engine.""" - return MagicMock(spec=Engine) - - -@pytest.fixture -def mock_account(): - """Mock Account user.""" - account = MagicMock(spec=Account) +def _make_account(*, tenant_id: str | None = None) -> Account: + account = Account(name="Repository User", email=f"{uuid4()}@example.com") account.id = str(uuid4()) - account.current_tenant_id = str(uuid4()) + if tenant_id is not None: + tenant = Tenant(name="Repository Tenant") + tenant.id = tenant_id + account._current_tenant = tenant return account @pytest.fixture -def mock_end_user(): - """Mock EndUser.""" - user = MagicMock(spec=EndUser) - user.id = str(uuid4()) - user.tenant_id = str(uuid4()) - return user +def sqlite_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]: + """Create repository-owned sessions bound to the isolated SQLite engine.""" + return sessionmaker(bind=sqlite_engine, expire_on_commit=False) @pytest.fixture -def sample_workflow_execution(): +def account() -> Account: + return _make_account(tenant_id=str(uuid4())) + + +@pytest.fixture +def end_user() -> EndUser: + return EndUser( + id=str(uuid4()), + tenant_id=str(uuid4()), + app_id=None, + type=EndUserType.SERVICE_API, + external_user_id=None, + name="Repository End User", + session_id=str(uuid4()), + ) + + +@pytest.fixture +def sample_workflow_execution() -> WorkflowExecution: """Sample WorkflowExecution for testing.""" return WorkflowExecution( id_=str(uuid4()), @@ -71,125 +74,147 @@ def sample_workflow_execution(): class TestSQLAlchemyWorkflowExecutionRepository: - def test_init_with_sessionmaker(self, mock_session_factory, mock_account): + def test_init_with_sessionmaker(self, sqlite_session_factory: sessionmaker[Session], account: Account): app_id = "test_app_id" triggered_from = WorkflowRunTriggeredFrom.APP_RUN repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id=app_id, triggered_from=triggered_from, ) - assert repo._session_factory == mock_session_factory + assert repo._session_factory is sqlite_session_factory assert repo._tenant_id == RESOURCE_TENANT_ID assert repo._app_id == app_id assert repo._triggered_from == triggered_from - assert repo._creator_user_id == mock_account.id + assert repo._creator_user_id == account.id assert repo._creator_user_role == CreatorUserRole.ACCOUNT - def test_init_with_engine(self, mock_engine, mock_account): + def test_init_with_engine(self, sqlite_engine: Engine, account: Account): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_engine, + session_factory=sqlite_engine, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id="test_app_id", triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) assert isinstance(repo._session_factory, sessionmaker) - assert repo._session_factory.kw["bind"] == mock_engine + assert repo._session_factory.kw["bind"] is sqlite_engine - def test_init_invalid_session_factory(self, mock_account): + def test_init_invalid_session_factory(self, account: Account): with pytest.raises(ValueError, match="Invalid session_factory type"): SQLAlchemyWorkflowExecutionRepository( session_factory="invalid", tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id=None, triggered_from=None, ) - def test_init_no_tenant_id(self, mock_session_factory): - user = MagicMock(spec=Account) - user.current_tenant_id = None + def test_init_no_tenant_id(self, sqlite_session_factory: sessionmaker[Session]): + user = _make_account() with pytest.raises(ValueError, match="tenant_id is required"): SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id="", user=user, app_id=None, triggered_from=None, ) - def test_init_uses_resource_tenant_when_account_has_no_current_tenant(self, mock_session_factory): - user = MagicMock(spec=Account) - user.current_tenant_id = None - user.id = str(uuid4()) + def test_init_uses_resource_tenant_when_account_has_no_current_tenant( + self, sqlite_session_factory: sessionmaker[Session] + ): + user = _make_account() repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, - tenant_id="resource-tenant-id", + session_factory=sqlite_session_factory, + tenant_id=RESOURCE_TENANT_ID, user=user, app_id="test-app", triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) - assert repo._tenant_id == "resource-tenant-id" + assert repo._tenant_id == RESOURCE_TENANT_ID assert repo._creator_user_id == user.id - def test_init_with_end_user(self, mock_session_factory, mock_end_user): + def test_init_with_end_user(self, sqlite_session_factory: sessionmaker[Session], end_user: EndUser): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_end_user, + user=end_user, app_id=None, triggered_from=None, ) assert repo._tenant_id == RESOURCE_TENANT_ID assert repo._creator_user_role == CreatorUserRole.END_USER - def test_to_domain_model(self, mock_session_factory, mock_account): + @pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True) + def test_to_domain_model( + self, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, + account: Account, + ): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id=None, triggered_from=None, ) - db_model = MagicMock(spec=WorkflowRun) - db_model.id = str(uuid4()) - db_model.workflow_id = str(uuid4()) - db_model.type = "workflow" - db_model.version = "1.0" - db_model.inputs_dict = {"in": "val"} - db_model.outputs_dict = {"out": "val"} - db_model.graph_dict = {"nodes": []} - db_model.status = "succeeded" - db_model.error = "some error" - db_model.total_tokens = 50 - db_model.total_steps = 3 - db_model.exceptions_count = 1 - db_model.created_at = datetime.now(UTC) - db_model.finished_at = datetime.now(UTC) + db_model = WorkflowRun( + id=str(uuid4()), + tenant_id=account.current_tenant_id, + app_id=str(uuid4()), + workflow_id=str(uuid4()), + type=ModelWorkflowType.WORKFLOW, + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, + version="1.0", + inputs=json.dumps({"in": "val"}), + outputs=json.dumps({"out": "val"}), + graph=json.dumps({"nodes": []}), + status=WorkflowExecutionStatus.SUCCEEDED, + error="some error", + elapsed_time=1.0, + total_tokens=50, + total_steps=3, + exceptions_count=1, + created_by_role=CreatorUserRole.ACCOUNT, + created_by=account.id, + created_at=datetime.now(UTC), + finished_at=datetime.now(UTC), + ) + sqlite_session.add(db_model) + sqlite_session.commit() + sqlite_session.expunge_all() + persisted_model = sqlite_session.get(WorkflowRun, db_model.id) + assert persisted_model is not None - domain_model = repo._to_domain_model(db_model) + domain_model = repo._to_domain_model(persisted_model) - assert domain_model.id_ == db_model.id - assert domain_model.workflow_id == db_model.workflow_id + assert domain_model.id_ == persisted_model.id + assert domain_model.workflow_id == persisted_model.workflow_id assert domain_model.status == WorkflowExecutionStatus.SUCCEEDED - assert domain_model.inputs == db_model.inputs_dict + assert domain_model.inputs == {"in": "val"} assert domain_model.error_message == "some error" - def test_to_db_model(self, mock_session_factory, mock_account, sample_workflow_execution): + def test_to_db_model( + self, + sqlite_session_factory: sessionmaker[Session], + account: Account, + sample_workflow_execution: WorkflowExecution, + ): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id="test_app", triggered_from=WorkflowRunTriggeredFrom.DEBUGGING, ) @@ -208,11 +233,16 @@ class TestSQLAlchemyWorkflowExecutionRepository: assert db_model.total_tokens == sample_workflow_execution.total_tokens assert db_model.elapsed_time == 10.0 - def test_to_db_model_edge_cases(self, mock_session_factory, mock_account, sample_workflow_execution): + def test_to_db_model_edge_cases( + self, + sqlite_session_factory: sessionmaker[Session], + account: Account, + sample_workflow_execution: WorkflowExecution, + ): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id="test_app", triggered_from=WorkflowRunTriggeredFrom.DEBUGGING, ) @@ -231,11 +261,16 @@ class TestSQLAlchemyWorkflowExecutionRepository: assert db_model.error is None assert db_model.elapsed_time == 0 - def test_to_db_model_app_id_none(self, mock_session_factory, mock_account, sample_workflow_execution): + def test_to_db_model_app_id_none( + self, + sqlite_session_factory: sessionmaker[Session], + account: Account, + sample_workflow_execution: WorkflowExecution, + ): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id=None, triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) @@ -244,11 +279,16 @@ class TestSQLAlchemyWorkflowExecutionRepository: assert not hasattr(db_model, "app_id") or db_model.app_id is None assert db_model.tenant_id == repo._tenant_id - def test_to_db_model_missing_context(self, mock_session_factory, mock_account, sample_workflow_execution): + def test_to_db_model_missing_context( + self, + sqlite_session_factory: sessionmaker[Session], + account: Account, + sample_workflow_execution: WorkflowExecution, + ): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id=None, triggered_from=None, ) @@ -267,33 +307,47 @@ class TestSQLAlchemyWorkflowExecutionRepository: with pytest.raises(ValueError, match="created_by_role is required"): repo._to_db_model(sample_workflow_execution) - def test_save(self, mock_session_factory, mock_account, sample_workflow_execution): + @pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True) + def test_save( + self, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, + account: Account, + sample_workflow_execution: WorkflowExecution, + ): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id="test_app", triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) repo.save(sample_workflow_execution) - session = mock_session_factory.return_value.__enter__.return_value - session.merge.assert_called_once() - session.commit.assert_called_once() + persisted_model = sqlite_session.get(WorkflowRun, sample_workflow_execution.id_) + assert persisted_model is not None + assert persisted_model.tenant_id == RESOURCE_TENANT_ID + assert persisted_model.inputs_dict == sample_workflow_execution.inputs + assert persisted_model.outputs_dict == sample_workflow_execution.outputs # Check cache assert sample_workflow_execution.id_ in repo._execution_cache cached_model = repo._execution_cache[sample_workflow_execution.id_] assert cached_model.id == sample_workflow_execution.id_ + @pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True) def test_save_uses_execution_started_at_when_record_does_not_exist( - self, mock_session_factory, mock_account, sample_workflow_execution + self, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, + account: Account, + sample_workflow_execution: WorkflowExecution, ): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id="test_app", triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) @@ -301,41 +355,71 @@ class TestSQLAlchemyWorkflowExecutionRepository: started_at = datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC) sample_workflow_execution.started_at = started_at - session = mock_session_factory.return_value.__enter__.return_value - session.get.return_value = None - repo.save(sample_workflow_execution) - saved_model = session.merge.call_args.args[0] - assert saved_model.created_at == started_at - session.commit.assert_called_once() + persisted_model = sqlite_session.get(WorkflowRun, sample_workflow_execution.id_) + assert persisted_model is not None + assert persisted_model.created_at == started_at.replace(tzinfo=None) + @pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True) def test_save_preserves_existing_created_at_when_record_already_exists( - self, mock_session_factory, mock_account, sample_workflow_execution + self, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, + account: Account, + sample_workflow_execution: WorkflowExecution, ): repo = SQLAlchemyWorkflowExecutionRepository( - session_factory=mock_session_factory, + session_factory=sqlite_session_factory, tenant_id=RESOURCE_TENANT_ID, - user=mock_account, + user=account, app_id="test_app", triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) execution_id = sample_workflow_execution.id_ existing_created_at = datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC) - - existing_run = WorkflowRun() - existing_run.id = execution_id - existing_run.tenant_id = repo._tenant_id - existing_run.created_at = existing_created_at - - session = mock_session_factory.return_value.__enter__.return_value - session.get.return_value = existing_run + sample_workflow_execution.started_at = existing_created_at + repo.save(sample_workflow_execution) sample_workflow_execution.started_at = datetime(2026, 1, 1, 12, 30, 0, tzinfo=UTC) repo.save(sample_workflow_execution) - saved_model = session.merge.call_args.args[0] - assert saved_model.created_at == existing_created_at - session.commit.assert_called_once() + persisted_model = sqlite_session.get(WorkflowRun, execution_id) + assert persisted_model is not None + assert persisted_model.created_at == existing_created_at.replace(tzinfo=None) + + @pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True) + def test_save_rejects_execution_owned_by_another_tenant( + self, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, + account: Account, + sample_workflow_execution: WorkflowExecution, + ): + other_tenant_id = str(uuid4()) + other_account = _make_account(tenant_id=str(uuid4())) + other_repo = SQLAlchemyWorkflowExecutionRepository( + session_factory=sqlite_session_factory, + tenant_id=other_tenant_id, + user=other_account, + app_id="test_app", + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, + ) + other_repo.save(sample_workflow_execution) + + repo = SQLAlchemyWorkflowExecutionRepository( + session_factory=sqlite_session_factory, + tenant_id=RESOURCE_TENANT_ID, + user=account, + app_id="test_app", + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, + ) + with pytest.raises(ValueError, match="Unauthorized access to workflow run"): + repo.save(sample_workflow_execution) + + sqlite_session.expire_all() + persisted_model = sqlite_session.get(WorkflowRun, sample_workflow_execution.id_) + assert persisted_model is not None + assert persisted_model.tenant_id == other_tenant_id diff --git a/api/tests/unit_tests/core/repositories/test_sqlalchemy_workflow_node_execution_repository.py b/api/tests/unit_tests/core/repositories/test_sqlalchemy_workflow_node_execution_repository.py index 3289e2b50e8..b437cdcf205 100644 --- a/api/tests/unit_tests/core/repositories/test_sqlalchemy_workflow_node_execution_repository.py +++ b/api/tests/unit_tests/core/repositories/test_sqlalchemy_workflow_node_execution_repository.py @@ -1,18 +1,21 @@ +"""SQLite-backed tests for the workflow node execution repository.""" + from __future__ import annotations import json import logging -from collections.abc import Mapping +from collections.abc import Callable, Generator, Mapping +from contextlib import contextmanager from datetime import UTC, datetime from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock, Mock +from unittest.mock import Mock import psycopg2.errors import pytest -from sqlalchemy import Engine, create_engine +from sqlalchemy import Engine, event, select from sqlalchemy.exc import IntegrityError -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from configs import dify_config from core.repositories.factory import OrderConfig @@ -23,135 +26,145 @@ from core.repositories.sqlalchemy_workflow_node_execution_repository import ( _find_first, _replace_or_append_offload, ) +from extensions.storage.storage_type import StorageType from graphon.entities import WorkflowNodeExecution -from graphon.enums import ( - BuiltinNodeTypes, - WorkflowNodeExecutionMetadataKey, - WorkflowNodeExecutionStatus, -) -from models import Account, EndUser -from models.enums import ExecutionOffLoadType +from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus +from models import Account, EndUser, Tenant +from models.enums import CreatorUserRole, ExecutionOffLoadType +from models.model import UploadFile from models.workflow import WorkflowNodeExecutionModel, WorkflowNodeExecutionOffload, WorkflowNodeExecutionTriggeredFrom -RESOURCE_TENANT_ID = "tenant" - -def _mock_account(*, tenant_id: str = "tenant", user_id: str = "user") -> Account: - user = Mock(spec=Account) +def _account(*, tenant_id: str = "tenant-1", user_id: str = "user-1") -> Account: + user = Account(name="Test Account", email="test@example.com") user.id = user_id - user.current_tenant_id = tenant_id + user._current_tenant = Tenant(name="Test Tenant") + user._current_tenant.id = tenant_id return user -def _mock_end_user(*, tenant_id: str = "tenant", user_id: str = "user") -> EndUser: - user = Mock(spec=EndUser) - user.id = user_id - user.tenant_id = tenant_id - return user +def _end_user(*, tenant_id: str = "tenant-1", user_id: str = "end-user-1") -> EndUser: + return EndUser(id=user_id, tenant_id=tenant_id) + + +def _upload_file(*, key: str = "storage-key") -> UploadFile: + return UploadFile( + tenant_id="tenant-1", + storage_type=StorageType.LOCAL, + key=key, + name="offload.json", + size=1, + extension="json", + mime_type="application/json", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + created_at=datetime.now(UTC), + used=False, + ) def _execution( *, - execution_id: str = "exec-id", - node_execution_id: str = "node-exec-id", - workflow_run_id: str = "run-id", + execution_id: str = "execution-1", + node_execution_id: str = "node-execution-1", + run_id: str = "run-1", + index: int = 1, status: WorkflowNodeExecutionStatus = WorkflowNodeExecutionStatus.SUCCEEDED, inputs: Mapping[str, Any] | None = None, outputs: Mapping[str, Any] | None = None, process_data: Mapping[str, Any] | None = None, - metadata: Mapping[WorkflowNodeExecutionMetadataKey, Any] | None = None, ) -> WorkflowNodeExecution: return WorkflowNodeExecution( id=execution_id, node_execution_id=node_execution_id, - workflow_id="workflow-id", - workflow_execution_id=workflow_run_id, - index=1, + workflow_id="workflow-1", + workflow_execution_id=run_id, + index=index, predecessor_node_id=None, - node_id="node-id", + node_id=f"node-{index}", node_type=BuiltinNodeTypes.LLM, - title="Title", + title=f"Node {index}", inputs=inputs, outputs=outputs, process_data=process_data, status=status, error=None, elapsed_time=1.0, - metadata=metadata, + metadata={WorkflowNodeExecutionMetadataKey.TOTAL_TOKENS: index}, created_at=datetime.now(UTC), finished_at=None, ) -class _SessionCtx: - def __init__(self, session: Any): - self._session = session - - def __enter__(self) -> Any: - return self._session - - def __exit__(self, exc_type, exc, tb) -> None: - return None - - -def _session_factory(session: Any) -> sessionmaker: - factory = Mock(spec=sessionmaker) - factory.return_value = _SessionCtx(session) - return factory - - -def test_init_accepts_engine_and_sessionmaker_and_sets_role(monkeypatch: pytest.MonkeyPatch) -> None: +def _repository( + monkeypatch: pytest.MonkeyPatch, + factory: sessionmaker[Session] | Engine, + *, + tenant_id: str = "tenant-1", + app_id: str | None = "app-1", + user: Account | EndUser | None = None, + triggered_from: WorkflowNodeExecutionTriggeredFrom | None = WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, +) -> SQLAlchemyWorkflowNodeExecutionRepository: monkeypatch.setattr( "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), + lambda *_args: SimpleNamespace(upload_file=Mock()), + ) + return SQLAlchemyWorkflowNodeExecutionRepository( + session_factory=factory, + tenant_id=tenant_id, + user=user or _account(tenant_id=tenant_id), + app_id=app_id, + triggered_from=triggered_from, ) - engine: Engine = create_engine("sqlite:///:memory:") - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=engine, - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - assert isinstance(repo._session_factory, sessionmaker) - sm = Mock(spec=sessionmaker) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=sm, - tenant_id=RESOURCE_TENANT_ID, - user=_mock_end_user(), - app_id="app", - triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP, - ) - assert repo._creator_user_role.value == "end_user" +@contextmanager +def _raise_on_execution_insert(engine: Engine) -> Generator[None]: + def raise_error( + _conn: Any, + _cursor: Any, + statement: str, + _parameters: Any, + _context: Any, + _executemany: Any, + ) -> None: + if statement.lstrip().upper().startswith("INSERT") and "workflow_node_executions" in statement: + raise RuntimeError("forced execution INSERT") + + event.listen(engine, "before_cursor_execute", raise_error) + try: + yield + finally: + event.remove(engine, "before_cursor_execute", raise_error) -def test_init_rejects_invalid_session_factory_type(monkeypatch: pytest.MonkeyPatch) -> None: +def test_init_accepts_real_engine_and_sessionmaker_and_sets_role( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session_factory: sessionmaker[Session] +) -> None: + engine_repo = _repository(monkeypatch, sqlite_engine) + assert isinstance(engine_repo._session_factory, sessionmaker) + end_user_repo = _repository(monkeypatch, sqlite_session_factory, user=_end_user()) + assert end_user_repo._creator_user_role.value == "end_user" + + +def test_init_rejects_invalid_factory_and_missing_tenant(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), + lambda *_args: SimpleNamespace(upload_file=Mock()), ) with pytest.raises(ValueError, match="Invalid session_factory type"): - SQLAlchemyWorkflowNodeExecutionRepository( # type: ignore[arg-type] - session_factory=object(), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), + SQLAlchemyWorkflowNodeExecutionRepository( + session_factory=object(), # type: ignore[arg-type] + tenant_id="tenant-1", + user=_account(), app_id=None, triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, ) - - -def test_init_requires_tenant_id(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - user = _mock_account() - user.current_tenant_id = None - with pytest.raises(ValueError, match="tenant_id is required"): + user = _account() + user._current_tenant = None + with pytest.raises(ValueError, match="tenant_id"): SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), + session_factory=sessionmaker(), tenant_id="", user=user, app_id=None, @@ -159,663 +172,411 @@ def test_init_requires_tenant_id(monkeypatch: pytest.MonkeyPatch) -> None: ) -def test_init_uses_resource_tenant_when_account_has_no_current_tenant(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - user = _mock_account() - user.current_tenant_id = None +def test_init_uses_resource_tenant_when_account_has_no_current_tenant( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + user = _account() + user._current_tenant = None - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id="resource-tenant-id", + repo = _repository( + monkeypatch, + sqlite_session_factory, + tenant_id="resource-tenant", user=user, - app_id="app-id", - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, ) - assert repo._tenant_id == "resource-tenant-id" + assert repo._tenant_id == "resource-tenant" assert repo._creator_user_id == user.id -def test_create_truncator_uses_config(monkeypatch: pytest.MonkeyPatch) -> None: - created: dict[str, Any] = {} +def test_helper_functions_and_truncator_configuration( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + assert _deterministic_json_dump({"b": 1, "a": 2}) == '{"a": 2, "b": 1}' + assert _find_first([], lambda _value: True) is None + assert _find_first([1, 2, 3], lambda value: value > 1) == 2 + inputs = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.INPUTS) + outputs = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.OUTPUTS) + assert _find_first([inputs, outputs], _filter_by_offload_type(ExecutionOffLoadType.OUTPUTS)) is outputs + replaced = _replace_or_append_offload( + [inputs, outputs], WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.INPUTS) + ) + assert [item.type_ for item in replaced] == [ExecutionOffLoadType.OUTPUTS, ExecutionOffLoadType.INPUTS] - class FakeTruncator: - def __init__(self, *, max_size_bytes: int, array_element_limit: int, string_length_limit: int): + created: dict[str, int] = {} + + class Truncator: + def __init__(self, *, max_size_bytes: int, array_element_limit: int, string_length_limit: int) -> None: created.update( - { - "max_size_bytes": max_size_bytes, - "array_element_limit": array_element_limit, - "string_length_limit": string_length_limit, - } + max_size_bytes=max_size_bytes, + array_element_limit=array_element_limit, + string_length_limit=string_length_limit, ) - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.VariableTruncator", - FakeTruncator, - ) - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - _ = repo._create_truncator() + monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.VariableTruncator", Truncator) + _repository(monkeypatch, sqlite_session_factory)._create_truncator() assert created["max_size_bytes"] == dify_config.WORKFLOW_VARIABLE_TRUNCATION_MAX_SIZE -def test_helpers_find_first_and_replace_or_append_and_filter() -> None: - assert _deterministic_json_dump({"b": 1, "a": 2}) == '{"a": 2, "b": 1}' - assert _find_first([], lambda _: True) is None - assert _find_first([1, 2, 3], lambda x: x > 1) == 2 - - off1 = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.INPUTS) - off2 = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.OUTPUTS) - assert _find_first([off1, off2], _filter_by_offload_type(ExecutionOffLoadType.OUTPUTS)) is off2 - - replaced = _replace_or_append_offload([off1, off2], WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.INPUTS)) - assert len(replaced) == 2 - assert [o.type_ for o in replaced] == [ExecutionOffLoadType.OUTPUTS, ExecutionOffLoadType.INPUTS] - - -def test_to_db_model_requires_constructor_context(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), +def test_to_db_model_uses_context_and_deterministic_json( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + db_model = repo._to_db_model( + _execution( + inputs={"b": 1, "a": 2}, + process_data={"agent_workspace_binding_id": "participant-1"}, + ) ) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - execution = _execution(inputs={"b": 1, "a": 2}, metadata={WorkflowNodeExecutionMetadataKey.TOTAL_TOKENS: 1}) - - # Happy path: deterministic json dump should be sorted - db_model = repo._to_db_model(execution) - assert db_model.tenant_id == RESOURCE_TENANT_ID - assert db_model.created_by == "user" - assert db_model.created_by_role.value == "account" assert json.loads(db_model.inputs or "{}") == {"a": 2, "b": 1} - assert json.loads(db_model.execution_metadata or "{}")["total_tokens"] == 1 - + assert db_model.tenant_id == "tenant-1" + assert db_model.app_id == "app-1" + assert db_model.created_by == "user-1" + assert db_model.created_by_role == CreatorUserRole.ACCOUNT + assert json.loads(db_model.execution_metadata or "{}") == {"total_tokens": 1} + assert db_model.agent_workspace_binding_id is None + assert _repository(monkeypatch, sqlite_session_factory, app_id=None)._to_db_model(_execution()).app_id is None repo._triggered_from = None with pytest.raises(ValueError, match="triggered_from is required"): - repo._to_db_model(execution) + repo._to_db_model(_execution()) -def test_to_db_model_requires_creator_user_id_and_role(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id="app", - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) +def test_to_db_model_requires_creator_context( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) execution = _execution() - db_model = repo._to_db_model(execution) - assert db_model.app_id == "app" - repo._creator_user_id = None + monkeypatch.setattr(repo, "_creator_user_id", None) with pytest.raises(ValueError, match="created_by is required"): repo._to_db_model(execution) - repo._creator_user_id = "user" - repo._creator_user_role = None + monkeypatch.setattr(repo, "_creator_user_id", "user-1") + monkeypatch.setattr(repo, "_creator_user_role", None) with pytest.raises(ValueError, match="created_by_role is required"): repo._to_db_model(execution) -def test_is_duplicate_key_error_and_regenerate_id( - monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture -) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - unique = Mock(spec=psycopg2.errors.UniqueViolation) - duplicate_error = IntegrityError("dup", params=None, orig=unique) - assert repo._is_duplicate_key_error(duplicate_error) is True - assert repo._is_duplicate_key_error(IntegrityError("other", params=None, orig=None)) is False - - execution = _execution(execution_id="old-id") - db_model = WorkflowNodeExecutionModel() - db_model.id = "old-id" - monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.uuidv7", lambda: "new-id") - caplog.set_level(logging.WARNING) - repo._regenerate_id_on_duplicate(execution, db_model) - assert execution.id == "new-id" - assert db_model.id == "new-id" - assert any("Duplicate key conflict" in r.message for r in caplog.records) - - -def test_persist_to_database_updates_existing_and_inserts_new(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - session = MagicMock() - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=_session_factory(session), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - db_model = WorkflowNodeExecutionModel() - db_model.id = "id1" - db_model.node_execution_id = "node1" - db_model.foo = "bar" # type: ignore[attr-defined] - db_model.__dict__["_private"] = "x" - - existing = SimpleNamespace() - session.get.return_value = existing - repo._persist_to_database(db_model) - assert existing.foo == "bar" - session.add.assert_not_called() - assert repo._node_execution_cache["node1"] is db_model - - session.reset_mock() - session.get.return_value = None - repo._node_execution_cache.clear() - repo._persist_to_database(db_model) - session.add.assert_called_once_with(db_model) - assert repo._node_execution_cache["node1"] is db_model - - -def test_truncate_and_upload_returns_none_when_no_values_or_not_truncated(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id="app", - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - assert repo._truncate_and_upload(None, "e", ExecutionOffLoadType.INPUTS) is None - - class FakeTruncator: - def truncate_variable_mapping(self, value: Any): # type: ignore[no-untyped-def] - return value, False - - monkeypatch.setattr(repo, "_create_truncator", lambda: FakeTruncator()) - assert repo._truncate_and_upload({"a": 1}, "e", ExecutionOffLoadType.INPUTS) is None - - -def test_truncate_and_upload_uploads_and_builds_offload(monkeypatch: pytest.MonkeyPatch) -> None: - uploaded: dict[str, Any] = {} - - class FakeFileService: - def upload_file(self, *, filename: str, content: bytes, mimetype: str, user: Any, tenant_id: str): # type: ignore[no-untyped-def] - uploaded.update( - {"filename": filename, "content": content, "mimetype": mimetype, "user": user, "tenant_id": tenant_id} - ) - return SimpleNamespace(id="file-id", key="file-key") - - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", lambda *_: FakeFileService() - ) - monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.uuidv7", lambda: "offload-id") - - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id="app", - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - class FakeTruncator: - def truncate_variable_mapping(self, value: Any): # type: ignore[no-untyped-def] - return {"truncated": True}, True - - monkeypatch.setattr(repo, "_create_truncator", lambda: FakeTruncator()) - - result = repo._truncate_and_upload({"a": 1}, "exec", ExecutionOffLoadType.INPUTS) - assert result is not None - assert result.truncated_value == {"truncated": True} - assert uploaded["filename"].startswith("node_execution_exec_inputs.json") - assert uploaded["tenant_id"] == RESOURCE_TENANT_ID - assert result.offload.file_id == "file-id" - assert result.offload.type_ == ExecutionOffLoadType.INPUTS - - -def test_to_domain_model_loads_offloaded_files(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - db_model = WorkflowNodeExecutionModel() - db_model.id = "id" - db_model.node_execution_id = "node-exec" - db_model.workflow_id = "wf" - db_model.workflow_run_id = "run" - db_model.index = 1 - db_model.predecessor_node_id = None - db_model.node_id = "node" - db_model.node_type = BuiltinNodeTypes.LLM - db_model.title = "t" - db_model.inputs = json.dumps({"trunc": "i"}) - db_model.process_data = json.dumps({"trunc": "p"}) - db_model.outputs = json.dumps({"trunc": "o"}) - db_model.status = WorkflowNodeExecutionStatus.SUCCEEDED - db_model.error = None - db_model.elapsed_time = 0.1 - db_model.execution_metadata = json.dumps({"total_tokens": 3}) - db_model.created_at = datetime.now(UTC) - db_model.finished_at = None - - off_in = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.INPUTS) - off_out = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.OUTPUTS) - off_proc = WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.PROCESS_DATA) - off_in.file = SimpleNamespace(key="k-in") - off_out.file = SimpleNamespace(key="k-out") - off_proc.file = SimpleNamespace(key="k-proc") - db_model.offload_data = [off_out, off_in, off_proc] - - def fake_load(key: str) -> bytes: - return json.dumps({"full": key}).encode() - - monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.storage.load", fake_load) - - domain = repo._to_domain_model(db_model) - assert domain.inputs == {"full": "k-in"} - assert domain.outputs == {"full": "k-out"} - assert domain.process_data == {"full": "k-proc"} - assert domain.get_truncated_inputs() == {"trunc": "i"} - assert domain.get_truncated_outputs() == {"trunc": "o"} - assert domain.get_truncated_process_data() == {"trunc": "p"} - - -def test_to_domain_model_returns_early_when_no_offload_data(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - db_model = WorkflowNodeExecutionModel() - db_model.id = "id" - db_model.node_execution_id = "node-exec" - db_model.workflow_id = "wf" - db_model.workflow_run_id = "run" - db_model.index = 1 - db_model.predecessor_node_id = None - db_model.node_id = "node" - db_model.node_type = BuiltinNodeTypes.LLM - db_model.title = "t" - db_model.inputs = json.dumps({"i": 1}) - db_model.process_data = json.dumps({"p": 2}) - db_model.outputs = json.dumps({"o": 3}) - db_model.status = WorkflowNodeExecutionStatus.SUCCEEDED - db_model.error = None - db_model.elapsed_time = 0.1 - db_model.execution_metadata = "{}" - db_model.created_at = datetime.now(UTC) - db_model.finished_at = None - db_model.offload_data = [] - - domain = repo._to_domain_model(db_model) - assert domain.inputs == {"i": 1} - assert domain.outputs == {"o": 3} - - def test_json_encode_uses_runtime_converter(monkeypatch: pytest.MonkeyPatch) -> None: - class FakeConverter: + class Converter: def to_json_encodable(self, values: Mapping[str, Any]) -> Mapping[str, Any]: - return {"wrapped": values["a"]} + return {"wrapped": values["value"]} monkeypatch.setattr( "core.repositories.sqlalchemy_workflow_node_execution_repository.WorkflowRuntimeTypeConverter", - FakeConverter, - ) - assert SQLAlchemyWorkflowNodeExecutionRepository._json_encode({"a": 1}) == '{"wrapped": 1}' - - -def test_save_execution_data_handles_existing_db_model_and_truncation(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - session = MagicMock() - session.execute.return_value.scalars.return_value.first.return_value = SimpleNamespace( - id="id", - offload_data=[WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.INPUTS)], - inputs=None, - outputs=None, - process_data=None, - ) - session.merge = Mock() - session.flush = Mock() - session.begin.return_value.__enter__ = Mock(return_value=session) - session.begin.return_value.__exit__ = Mock(return_value=None) - - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=_session_factory(session), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id="app", - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + Converter, ) - execution = _execution(inputs={"a": 1}, outputs={"b": 2}, process_data={"c": 3}) - - trunc_result = SimpleNamespace( - truncated_value={"trunc": True}, - offload=WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.INPUTS, file_id="f1"), - ) - monkeypatch.setattr( - repo, "_truncate_and_upload", lambda values, *_args, **_kwargs: trunc_result if values == {"a": 1} else None - ) - monkeypatch.setattr(repo, "_json_encode", lambda values: json.dumps(values, sort_keys=True)) - - repo.save_execution_data(execution) - # Inputs should be truncated, outputs/process_data encoded directly - db_model = session.merge.call_args.args[0] - assert json.loads(db_model.inputs) == {"trunc": True} - assert json.loads(db_model.outputs) == {"b": 2} - assert json.loads(db_model.process_data) == {"c": 3} - assert any(off.type_ == ExecutionOffLoadType.INPUTS for off in db_model.offload_data) - assert execution.get_truncated_inputs() == {"trunc": True} + assert SQLAlchemyWorkflowNodeExecutionRepository._json_encode({"value": 1}) == '{"wrapped": 1}' -def test_save_execution_data_truncates_outputs_and_process_data(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - existing = SimpleNamespace( - id="id", - offload_data=[], - inputs=None, - outputs=None, - process_data=None, - ) - session = MagicMock() - session.execute.return_value.scalars.return_value.first.return_value = existing - session.merge = Mock() - session.flush = Mock() - session.begin.return_value.__enter__ = Mock(return_value=session) - session.begin.return_value.__exit__ = Mock(return_value=None) - - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=_session_factory(session), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id="app", - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - execution = _execution(inputs={"a": 1}, outputs={"b": 2}, process_data={"c": 3}) - - def trunc(values: Mapping[str, Any], *_args: Any, **_kwargs: Any) -> Any: - if values == {"b": 2}: - return SimpleNamespace( - truncated_value={"b": "trunc"}, - offload=WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.OUTPUTS, file_id="f2"), - ) - if values == {"c": 3}: - return SimpleNamespace( - truncated_value={"c": "trunc"}, - offload=WorkflowNodeExecutionOffload(type_=ExecutionOffLoadType.PROCESS_DATA, file_id="f3"), - ) - return None - - monkeypatch.setattr(repo, "_truncate_and_upload", trunc) - monkeypatch.setattr(repo, "_json_encode", lambda values: json.dumps(values, sort_keys=True)) - - repo.save_execution_data(execution) - db_model = session.merge.call_args.args[0] - assert json.loads(db_model.outputs) == {"b": "trunc"} - assert json.loads(db_model.process_data) == {"c": "trunc"} - assert execution.get_truncated_outputs() == {"b": "trunc"} - assert execution.get_truncated_process_data() == {"c": "trunc"} - - -def test_save_execution_data_handles_missing_db_model(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - session = MagicMock() - session.execute.return_value.scalars.return_value.first.return_value = None - session.merge = Mock() - session.flush = Mock() - session.begin.return_value.__enter__ = Mock(return_value=session) - session.begin.return_value.__exit__ = Mock(return_value=None) - - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=_session_factory(session), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - execution = _execution(inputs={"a": 1}) - fake_db_model = SimpleNamespace(id=execution.id, offload_data=[], inputs=None, outputs=None, process_data=None) - monkeypatch.setattr(repo, "_to_db_model", lambda *_: fake_db_model) - monkeypatch.setattr(repo, "_truncate_and_upload", lambda *_args, **_kwargs: None) - monkeypatch.setattr(repo, "_json_encode", lambda values: json.dumps(values)) - - repo.save_execution_data(execution) - merged = session.merge.call_args.args[0] - assert merged.inputs == '{"a": 1}' - - -def test_save_retries_duplicate_and_logs_non_duplicate( - monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +def test_save_inserts_and_updates_persisted_execution( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] ) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - execution = _execution(execution_id="id") - unique = Mock(spec=psycopg2.errors.UniqueViolation) - duplicate_error = IntegrityError("dup", params=None, orig=unique) - other_error = IntegrityError("other", params=None, orig=None) - - calls = {"n": 0} - - def persist(_db_model: Any) -> None: - calls["n"] += 1 - if calls["n"] == 1: - raise duplicate_error - - monkeypatch.setattr(repo, "_persist_to_database", persist) - monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.uuidv7", lambda: "new-id") + repo = _repository(monkeypatch, sqlite_session_factory) + execution = _execution(inputs={"value": 1}, outputs={"result": "first"}) repo.save(execution) - assert execution.id == "new-id" - assert repo._node_execution_cache[execution.node_execution_id] is not None - - caplog.set_level(logging.ERROR) - monkeypatch.setattr(repo, "_persist_to_database", lambda _db: (_ for _ in ()).throw(other_error)) - with pytest.raises(IntegrityError): - repo.save(_execution(execution_id="id2", node_execution_id="node2")) - assert any("Non-duplicate key integrity error" in r.message for r in caplog.records) + with sqlite_session_factory() as session: + persisted = session.get(WorkflowNodeExecutionModel, execution.id) + assert persisted is not None + assert persisted.outputs_dict == {"result": "first"} + execution.title = "Updated" + execution.outputs = {"result": "second"} + repo.save(execution) + with sqlite_session_factory() as session: + persisted = session.get(WorkflowNodeExecutionModel, execution.id) + assert persisted is not None + assert persisted.title == "Updated" + assert persisted.outputs_dict == {"result": "second"} + assert execution.node_execution_id is not None + assert repo._node_execution_cache[execution.node_execution_id].id == execution.id -def test_save_logs_and_reraises_on_unexpected_error( - monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +def test_save_owned_session_rolls_back_failed_insert( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], ) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + with _raise_on_execution_insert(sqlite_engine), pytest.raises(RuntimeError, match="forced execution INSERT"): + repo.save(_execution()) + with sqlite_session_factory() as session: + assert session.scalar(select(WorkflowNodeExecutionModel)) is None + + +def test_save_execution_data_updates_existing_and_creates_missing( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + existing = _execution( + inputs={"initial": True}, + process_data={"workflow_agent_binding_id": "binding-1"}, + ) + repo.save(existing) + existing.inputs = {"updated": True} + existing.outputs = {"result": 2} + existing.process_data = {"step": 3} + monkeypatch.setattr(repo, "_truncate_and_upload", lambda *_args, **_kwargs: None) + repo.save_execution_data(existing) + with sqlite_session_factory() as session: + persisted = session.get(WorkflowNodeExecutionModel, existing.id) + assert persisted is not None + assert persisted.inputs_dict == {"updated": True} + assert persisted.outputs_dict == {"result": 2} + assert persisted.process_data_dict == { + "step": 3, + "workflow_agent_binding_id": "binding-1", + } + + missing = _execution(execution_id="missing", node_execution_id="missing-node", inputs={"new": True}) + repo.save_execution_data(missing) + with sqlite_session_factory() as session: + persisted = session.get(WorkflowNodeExecutionModel, missing.id) + assert persisted is not None + assert persisted.inputs_dict == {"new": True} + + +@pytest.mark.parametrize( + ("execution_factory", "offload_type", "read_persisted", "read_truncated"), + [ + ( + lambda: _execution(inputs={"large": "value"}), + ExecutionOffLoadType.INPUTS, + lambda model: model.inputs_dict, + lambda execution: execution.get_truncated_inputs(), + ), + ( + lambda: _execution(outputs={"large": "value"}), + ExecutionOffLoadType.OUTPUTS, + lambda model: model.outputs_dict, + lambda execution: execution.get_truncated_outputs(), + ), + ( + lambda: _execution(process_data={"large": "value"}), + ExecutionOffLoadType.PROCESS_DATA, + lambda model: model.process_data_dict, + lambda execution: execution.get_truncated_process_data(), + ), + ], +) +def test_save_execution_data_persists_each_truncation_offload( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], + execution_factory: Callable[[], WorkflowNodeExecution], + offload_type: ExecutionOffLoadType, + read_persisted: Callable[[WorkflowNodeExecutionModel], Mapping[str, Any] | None], + read_truncated: Callable[[WorkflowNodeExecution], Mapping[str, Any] | None], +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + execution = execution_factory() + repo.save(execution) + offload = WorkflowNodeExecutionOffload( + tenant_id="tenant-1", + app_id="app-1", + node_execution_id=execution.id, + type_=offload_type, + file_id="file-1", + ) + result = SimpleNamespace(truncated_value={"large": "truncated"}, offload=offload) + monkeypatch.setattr(repo, "_truncate_and_upload", lambda values, *_args: result if values else None) + repo.save_execution_data(execution) + with sqlite_session_factory() as session: + persisted = session.get(WorkflowNodeExecutionModel, execution.id) + assert persisted is not None + assert read_persisted(persisted) == {"large": "truncated"} + offloads = session.scalars( + select(WorkflowNodeExecutionOffload).where(WorkflowNodeExecutionOffload.node_execution_id == execution.id) + ).all() + assert [item.type_ for item in offloads] == [offload_type] + assert read_truncated(execution) == {"large": "truncated"} + + +def test_get_by_workflow_run_filters_tenant_app_trigger_and_paused_and_orders( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + repo.save(_execution(execution_id="two", node_execution_id="node-two", index=2)) + repo.save(_execution(execution_id="one", node_execution_id="node-one", index=1)) + repo.save( + _execution( + execution_id="paused", + node_execution_id="node-paused", + index=3, + status=WorkflowNodeExecutionStatus.PAUSED, + ) + ) + _repository(monkeypatch, sqlite_session_factory, tenant_id="tenant-2").save( + _execution(execution_id="foreign-tenant", node_execution_id="foreign-tenant") + ) + _repository(monkeypatch, sqlite_session_factory, app_id="app-2").save( + _execution(execution_id="foreign-app", node_execution_id="foreign-app") + ) + _repository( + monkeypatch, + sqlite_session_factory, + triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP, + ).save(_execution(execution_id="single-step", node_execution_id="single-step")) + + models = repo.get_db_models_by_workflow_run( + "run-1", + OrderConfig(order_by=["missing", "index"], order_direction="desc"), + ) + assert [model.id for model in models] == ["two", "one"] + assert set(repo._node_execution_cache) >= {"node-one", "node-two"} + assert repo.get_db_models_by_workflow_run("missing-run") == [] + no_app_repo = _repository(monkeypatch, sqlite_session_factory, app_id=None) + assert ( + no_app_repo.get_db_models_by_workflow_run( + "missing-run", + OrderConfig(order_by=["missing"], order_direction="asc"), + ) + == [] + ) + + +def test_get_by_workflow_execution_maps_real_rows_to_domain( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + repo.save(_execution(inputs={"input": 1}, outputs={"output": 2})) + domains = repo.get_by_workflow_execution("run-1", OrderConfig(order_by=["index"], order_direction="asc")) + assert len(domains) == 1 + assert domains[0].inputs == {"input": 1} + assert domains[0].outputs == {"output": 2} + + +def test_to_domain_model_loads_offloaded_storage( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + db_model = repo._to_db_model( + _execution( + inputs={"truncated": "inputs"}, + outputs={"truncated": "outputs"}, + process_data={"truncated": "process_data"}, + ) + ) + offloads = [] + for offload_type in ExecutionOffLoadType: + offload = WorkflowNodeExecutionOffload(type_=offload_type) + offload.file = _upload_file(key=offload_type.value) + offloads.append(offload) + db_model.offload_data = offloads monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), + "core.repositories.sqlalchemy_workflow_node_execution_repository.storage.load", + lambda key: json.dumps({"full": key}).encode(), ) - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + domain = repo._to_domain_model(db_model) + assert domain.inputs == {"full": "inputs"} + assert domain.outputs == {"full": "outputs"} + assert domain.process_data == {"full": "process_data"} + assert domain.get_truncated_inputs() == {"truncated": "inputs"} + assert domain.get_truncated_outputs() == {"truncated": "outputs"} + assert domain.get_truncated_process_data() == {"truncated": "process_data"} + + +def test_truncate_and_upload_keeps_file_boundary_mocked( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + uploaded = _upload_file(key="file-key") + uploaded.id = "file-1" + repo = _repository(monkeypatch, sqlite_session_factory) + upload_file = Mock(return_value=uploaded) + monkeypatch.setattr(repo._file_service, "upload_file", upload_file) + + class Truncator: + def truncate_variable_mapping(self, _value: Any) -> tuple[dict[str, bool], bool]: + return {"truncated": True}, True + + monkeypatch.setattr(repo, "_create_truncator", lambda: Truncator()) + result = repo._truncate_and_upload({"value": 1}, "execution-1", ExecutionOffLoadType.INPUTS) + assert result is not None + assert result.truncated_value == {"truncated": True} + upload_file.assert_called_once_with( + filename="node_execution_execution-1_inputs.json", + content=b'{"value": 1}', + mimetype="application/json", + user=repo._user, + tenant_id="tenant-1", ) + assert result.offload.file_id == "file-1" + assert result.offload.type_ == ExecutionOffLoadType.INPUTS + assert result.offload.tenant_id == "tenant-1" + assert result.offload.app_id == "app-1" + assert result.offload.node_execution_id == "execution-1" + + +def test_truncate_and_upload_returns_none_for_missing_or_small_values( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + assert repo._truncate_and_upload(None, "execution-1", ExecutionOffLoadType.INPUTS) is None + + class Truncator: + def truncate_variable_mapping(self, value: Mapping[str, Any]) -> tuple[Mapping[str, Any], bool]: + return value, False + + monkeypatch.setattr(repo, "_create_truncator", lambda: Truncator()) + assert repo._truncate_and_upload({"value": 1}, "execution-1", ExecutionOffLoadType.INPUTS) is None + + +def test_duplicate_detection_and_id_regeneration( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + sqlite_session_factory: sessionmaker[Session], +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + duplicate = IntegrityError("duplicate", params=None, orig=Mock(spec=psycopg2.errors.UniqueViolation)) + assert repo._is_duplicate_key_error(duplicate) + assert not repo._is_duplicate_key_error(IntegrityError("other", params=None, orig=Exception("other"))) + execution = _execution(execution_id="old") + db_model = repo._to_db_model(execution) + monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.uuidv7", lambda: "new") + caplog.set_level(logging.WARNING) + repo._regenerate_id_on_duplicate(execution, db_model) + assert execution.id == db_model.id == "new" + assert "Duplicate key conflict" in caplog.text + + +def test_save_retries_postgres_duplicate_key( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + execution = _execution(execution_id="old") + duplicate = IntegrityError( + "duplicate", + params=None, + orig=Mock(spec=psycopg2.errors.UniqueViolation), + ) + persist = Mock(side_effect=[duplicate, None]) + monkeypatch.setattr(repo, "_persist_to_database", persist) + monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.uuidv7", lambda: "new") + + repo.save(execution) + + assert persist.call_count == 2 + assert execution.id == "new" + assert execution.node_execution_id is not None + assert repo._node_execution_cache[execution.node_execution_id].id == "new" + + +def test_save_logs_and_reraises_non_duplicate_and_unexpected_errors( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + sqlite_session_factory: sessionmaker[Session], +) -> None: + repo = _repository(monkeypatch, sqlite_session_factory) + non_duplicate = IntegrityError("other", params=None, orig=Exception("constraint")) + monkeypatch.setattr(repo, "_persist_to_database", Mock(side_effect=non_duplicate)) caplog.set_level(logging.ERROR) - monkeypatch.setattr(repo, "_persist_to_database", lambda _db: (_ for _ in ()).throw(RuntimeError("boom"))) + + with pytest.raises(IntegrityError): + repo.save(_execution()) + assert "Non-duplicate key integrity error" in caplog.text + + caplog.clear() + monkeypatch.setattr(repo, "_persist_to_database", Mock(side_effect=RuntimeError("boom"))) with pytest.raises(RuntimeError, match="boom"): - repo.save(_execution(execution_id="id3", node_execution_id="node3")) - assert any("Failed to save workflow node execution" in r.message for r in caplog.records) - - -def test_get_db_models_by_workflow_run_orders_and_caches(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - - class FakeStmt: - def __init__(self) -> None: - self.where_calls = 0 - self.order_by_args: tuple[Any, ...] | None = None - - def where(self, *_args: Any) -> FakeStmt: - self.where_calls += 1 - return self - - def order_by(self, *args: Any) -> FakeStmt: - self.order_by_args = args - return self - - stmt = FakeStmt() - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.WorkflowNodeExecutionModel.preload_offload_data_and_files", - lambda _q: stmt, - ) - monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.select", lambda *_: "select") - - model1 = SimpleNamespace(node_execution_id="n1") - model2 = SimpleNamespace(node_execution_id=None) - session = MagicMock() - session.scalars.return_value.all.return_value = [model1, model2] - - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=_session_factory(session), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id="app", - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - order = OrderConfig(order_by=["index", "missing"], order_direction="desc") - db_models = repo.get_db_models_by_workflow_run("run", order) - assert db_models == [model1, model2] - assert repo._node_execution_cache["n1"] is model1 - assert stmt.order_by_args is not None - - -def test_get_db_models_by_workflow_run_uses_asc_order(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - - class FakeStmt: - def where(self, *_args: Any) -> FakeStmt: - return self - - def order_by(self, *args: Any) -> FakeStmt: - self.args = args # type: ignore[attr-defined] - return self - - stmt = FakeStmt() - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.WorkflowNodeExecutionModel.preload_offload_data_and_files", - lambda _q: stmt, - ) - monkeypatch.setattr("core.repositories.sqlalchemy_workflow_node_execution_repository.select", lambda *_: "select") - - session = MagicMock() - session.scalars.return_value.all.return_value = [] - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=_session_factory(session), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - repo.get_db_models_by_workflow_run("run", OrderConfig(order_by=["index"], order_direction="asc")) - - -def test_get_by_workflow_run_maps_to_domain(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.FileService", - lambda *_: SimpleNamespace(upload_file=Mock()), - ) - - repo = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=Mock(spec=sessionmaker), - tenant_id=RESOURCE_TENANT_ID, - user=_mock_account(), - app_id=None, - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - db_models = [SimpleNamespace(id="db1"), SimpleNamespace(id="db2")] - monkeypatch.setattr(repo, "get_db_models_by_workflow_run", lambda *_args, **_kwargs: db_models) - monkeypatch.setattr(repo, "_to_domain_model", lambda m: f"domain:{m.id}") - - class FakeExecutor: - def __enter__(self) -> FakeExecutor: - return self - - def __exit__(self, exc_type, exc, tb) -> None: - return None - - def map(self, func, items, timeout: int): # type: ignore[no-untyped-def] - assert timeout == 30 - return list(map(func, items)) - - monkeypatch.setattr( - "core.repositories.sqlalchemy_workflow_node_execution_repository.ThreadPoolExecutor", - lambda max_workers: FakeExecutor(), - ) - - result = repo.get_by_workflow_execution("run", order_config=None) - assert result == ["domain:db1", "domain:db2"] + repo.save(_execution(execution_id="unexpected")) + assert "Failed to save workflow node execution" in caplog.text diff --git a/api/tests/unit_tests/core/repositories/test_workflow_node_execution_conflict_handling.py b/api/tests/unit_tests/core/repositories/test_workflow_node_execution_conflict_handling.py index e70745f70b3..0335bff634b 100644 --- a/api/tests/unit_tests/core/repositories/test_workflow_node_execution_conflict_handling.py +++ b/api/tests/unit_tests/core/repositories/test_workflow_node_execution_conflict_handling.py @@ -1,11 +1,14 @@ """Unit tests for workflow node execution conflict handling.""" -from unittest.mock import MagicMock, Mock +from collections.abc import Iterator +from contextlib import contextmanager +from dataclasses import dataclass import psycopg2.errors import pytest +from sqlalchemy import Engine, event from sqlalchemy.exc import IntegrityError -from sqlalchemy.orm import sessionmaker +from sqlalchemy.orm import Session, sessionmaker from core.repositories.sqlalchemy_workflow_node_execution_repository import ( SQLAlchemyWorkflowNodeExecutionRepository, @@ -16,196 +19,176 @@ from graphon.entities.workflow_node_execution import ( ) from graphon.enums import BuiltinNodeTypes from libs.datetime_utils import naive_utc_now -from models import Account, WorkflowNodeExecutionTriggeredFrom +from models import Account, Tenant, WorkflowNodeExecutionModel, WorkflowNodeExecutionTriggeredFrom + + +@dataclass(frozen=True) +class ConflictDatabase: + engine: Engine + session: Session + session_factory: sessionmaker[Session] + + +@dataclass(frozen=True) +class ConflictEvents: + insert_attempts: list[str] + rollbacks: list[bool] + + +@pytest.fixture +def conflict_database( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> ConflictDatabase: + """Use the shared SQLite fixtures for repository-owned sessions and assertions.""" + engine = sqlite_session.get_bind() + assert isinstance(engine, Engine) + return ConflictDatabase( + engine=engine, + session=sqlite_session, + session_factory=sqlite_session_factory, + ) + + +def _account() -> Account: + tenant = Tenant(name="Conflict Tenant") + tenant.id = "test-tenant-id" + account = Account(name="Conflict User", email="conflict@example.com") + account.id = "test-user-id" + account._current_tenant = tenant + return account + + +@pytest.fixture +def repository(conflict_database: ConflictDatabase) -> SQLAlchemyWorkflowNodeExecutionRepository: + return SQLAlchemyWorkflowNodeExecutionRepository( + session_factory=conflict_database.session_factory, + tenant_id="test-tenant-id", + user=_account(), + app_id="test-app-id", + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + ) + + +def _execution( + *, + execution_id: str, + status: WorkflowNodeExecutionStatus = WorkflowNodeExecutionStatus.RUNNING, +) -> WorkflowNodeExecution: + return WorkflowNodeExecution( + id=execution_id, + workflow_id="test-workflow-id", + workflow_execution_id="test-workflow-execution-id", + node_execution_id="test-node-execution-id", + node_id="test-node-id", + node_type=BuiltinNodeTypes.START, + title="Test Node", + index=1, + status=status, + created_at=naive_utc_now(), + ) + + +@contextmanager +def _fail_inserts( + engine: Engine, + *, + failure_count: int, + duplicate: bool, +) -> Iterator[ConflictEvents]: + insert_attempts: list[str] = [] + rollbacks: list[bool] = [] + + def fail_insert(_connection, _cursor, statement, parameters, _context, _executemany) -> None: + if not statement.lstrip().upper().startswith("INSERT INTO WORKFLOW_NODE_EXECUTIONS"): + return + insert_attempts.append(statement) + if len(insert_attempts) > failure_count: + return + original_error = ( + psycopg2.errors.UniqueViolation("forced duplicate key") + if duplicate + else RuntimeError("forced non-duplicate constraint failure") + ) + raise IntegrityError(statement, parameters, original_error) + + def record_rollback(_connection) -> None: + rollbacks.append(True) + + event.listen(engine, "before_cursor_execute", fail_insert) + event.listen(engine, "rollback", record_rollback) + try: + yield ConflictEvents(insert_attempts=insert_attempts, rollbacks=rollbacks) + finally: + event.remove(engine, "before_cursor_execute", fail_insert) + event.remove(engine, "rollback", record_rollback) class TestWorkflowNodeExecutionConflictHandling: """Test cases for handling duplicate key conflicts in workflow node execution.""" - def setup_method(self): - """Set up test fixtures.""" - # Create a mock user with tenant_id - self.mock_user = Mock(spec=Account) - self.mock_user.id = "test-user-id" - self.mock_user.current_tenant_id = "test-tenant-id" - - # Create mock session factory - self.mock_session_factory = Mock(spec=sessionmaker) - - # Create repository instance - self.repository = SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=self.mock_session_factory, - tenant_id="test-tenant-id", - user=self.mock_user, - app_id="test-app-id", - triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, - ) - - def test_save_with_duplicate_key_retries_with_new_uuid(self): - """Test that save retries with a new UUID v7 when encountering duplicate key error.""" - # Create a mock session - mock_session = MagicMock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - self.mock_session_factory.return_value = mock_session - - # Mock session.get to return None (no existing record) - mock_session.get.return_value = None - - # Create IntegrityError for duplicate key with proper psycopg2.errors.UniqueViolation - mock_unique_violation = Mock(spec=psycopg2.errors.UniqueViolation) - duplicate_error = IntegrityError( - "duplicate key value violates unique constraint", - params=None, - orig=mock_unique_violation, - ) - - # First call to session.add raises IntegrityError, second succeeds - mock_session.add.side_effect = [duplicate_error, None] - mock_session.commit.side_effect = [None, None] - - # Create test execution - execution = WorkflowNodeExecution( - id="original-id", - workflow_id="test-workflow-id", - workflow_execution_id="test-workflow-execution-id", - node_execution_id="test-node-execution-id", - node_id="test-node-id", - node_type=BuiltinNodeTypes.START, - title="Test Node", - index=1, - status=WorkflowNodeExecutionStatus.RUNNING, - created_at=naive_utc_now(), - ) - + def test_save_with_duplicate_key_retries_with_new_uuid( + self, + repository: SQLAlchemyWorkflowNodeExecutionRepository, + conflict_database: ConflictDatabase, + ) -> None: + execution = _execution(execution_id="original-id") original_id = execution.id - # Save should succeed after retry - self.repository.save(execution) + with _fail_inserts(conflict_database.engine, failure_count=1, duplicate=True) as conflicts: + repository.save(execution) - # Verify that session.add was called twice (initial attempt + retry) - assert mock_session.add.call_count == 2 - - # Verify that the ID was changed (new UUID v7 generated) + assert len(conflicts.insert_attempts) == 2 + assert len(conflicts.rollbacks) == 1 assert execution.id != original_id + assert conflict_database.session.get(WorkflowNodeExecutionModel, original_id) is None + persisted = conflict_database.session.get(WorkflowNodeExecutionModel, execution.id) + assert persisted is not None + assert persisted.node_execution_id == execution.node_execution_id - def test_save_with_existing_record_updates_instead_of_insert(self): - """Test that save updates existing record instead of inserting duplicate.""" - # Create a mock session - mock_session = MagicMock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - self.mock_session_factory.return_value = mock_session + def test_save_with_existing_record_updates_instead_of_insert( + self, + repository: SQLAlchemyWorkflowNodeExecutionRepository, + conflict_database: ConflictDatabase, + ) -> None: + execution = _execution(execution_id="existing-id") + repository.save(execution) - # Mock existing record - mock_existing = MagicMock() - mock_session.get.return_value = mock_existing - mock_session.commit.return_value = None + execution.status = WorkflowNodeExecutionStatus.SUCCEEDED + repository.save(execution) - # Create test execution - execution = WorkflowNodeExecution( - id="existing-id", - workflow_id="test-workflow-id", - workflow_execution_id="test-workflow-execution-id", - node_execution_id="test-node-execution-id", - node_id="test-node-id", - node_type=BuiltinNodeTypes.START, - title="Test Node", - index=1, - status=WorkflowNodeExecutionStatus.SUCCEEDED, - created_at=naive_utc_now(), - ) + conflict_database.session.expire_all() + persisted = conflict_database.session.get(WorkflowNodeExecutionModel, execution.id) + assert persisted is not None + assert persisted.status == WorkflowNodeExecutionStatus.SUCCEEDED + assert conflict_database.session.query(WorkflowNodeExecutionModel).count() == 1 - # Save should update existing record - self.repository.save(execution) + def test_save_exceeds_max_retries_raises_error( + self, + repository: SQLAlchemyWorkflowNodeExecutionRepository, + conflict_database: ConflictDatabase, + ) -> None: + execution = _execution(execution_id="test-id") - # Verify that session.add was not called (update path) - mock_session.add.assert_not_called() + with _fail_inserts(conflict_database.engine, failure_count=3, duplicate=True) as conflicts: + with pytest.raises(IntegrityError): + repository.save(execution) - # Verify that session.commit was called - mock_session.commit.assert_called_once() + assert len(conflicts.insert_attempts) == 3 + assert len(conflicts.rollbacks) == 3 + assert conflict_database.session.query(WorkflowNodeExecutionModel).count() == 0 - def test_save_exceeds_max_retries_raises_error(self): - """Test that save raises error after exceeding max retries.""" - # Create a mock session - mock_session = MagicMock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - self.mock_session_factory.return_value = mock_session + def test_save_non_duplicate_integrity_error_raises_immediately( + self, + repository: SQLAlchemyWorkflowNodeExecutionRepository, + conflict_database: ConflictDatabase, + ) -> None: + execution = _execution(execution_id="test-id") - # Mock session.get to return None (no existing record) - mock_session.get.return_value = None + with _fail_inserts(conflict_database.engine, failure_count=1, duplicate=False) as conflicts: + with pytest.raises(IntegrityError): + repository.save(execution) - # Create IntegrityError for duplicate key with proper psycopg2.errors.UniqueViolation - mock_unique_violation = Mock(spec=psycopg2.errors.UniqueViolation) - duplicate_error = IntegrityError( - "duplicate key value violates unique constraint", - params=None, - orig=mock_unique_violation, - ) - - # All attempts fail with duplicate error - mock_session.add.side_effect = duplicate_error - - # Create test execution - execution = WorkflowNodeExecution( - id="test-id", - workflow_id="test-workflow-id", - workflow_execution_id="test-workflow-execution-id", - node_execution_id="test-node-execution-id", - node_id="test-node-id", - node_type=BuiltinNodeTypes.START, - title="Test Node", - index=1, - status=WorkflowNodeExecutionStatus.RUNNING, - created_at=naive_utc_now(), - ) - - # Save should raise IntegrityError after max retries - with pytest.raises(IntegrityError): - self.repository.save(execution) - - # Verify that session.add was called 3 times (max_retries) - assert mock_session.add.call_count == 3 - - def test_save_non_duplicate_integrity_error_raises_immediately(self): - """Test that non-duplicate IntegrityErrors are raised immediately without retry.""" - # Create a mock session - mock_session = MagicMock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - self.mock_session_factory.return_value = mock_session - - # Mock session.get to return None (no existing record) - mock_session.get.return_value = None - - # Create IntegrityError for non-duplicate constraint - other_error = IntegrityError( - "null value in column violates not-null constraint", - params=None, - orig=None, - ) - - # First call raises non-duplicate error - mock_session.add.side_effect = other_error - - # Create test execution - execution = WorkflowNodeExecution( - id="test-id", - workflow_id="test-workflow-id", - workflow_execution_id="test-workflow-execution-id", - node_execution_id="test-node-execution-id", - node_id="test-node-id", - node_type=BuiltinNodeTypes.START, - title="Test Node", - index=1, - status=WorkflowNodeExecutionStatus.RUNNING, - created_at=naive_utc_now(), - ) - - # Save should raise error immediately - with pytest.raises(IntegrityError): - self.repository.save(execution) - - # Verify that session.add was called only once (no retry) - assert mock_session.add.call_count == 1 + assert len(conflicts.insert_attempts) == 1 + assert len(conflicts.rollbacks) == 1 + assert conflict_database.session.get(WorkflowNodeExecutionModel, execution.id) is None diff --git a/api/tests/unit_tests/core/repositories/test_workflow_node_execution_truncation.py b/api/tests/unit_tests/core/repositories/test_workflow_node_execution_truncation.py index 93ea4271093..b444c93e5f8 100644 --- a/api/tests/unit_tests/core/repositories/test_workflow_node_execution_truncation.py +++ b/api/tests/unit_tests/core/repositories/test_workflow_node_execution_truncation.py @@ -9,9 +9,9 @@ import json from dataclasses import dataclass from datetime import UTC, datetime from typing import Any -from unittest.mock import MagicMock from sqlalchemy import Engine +from sqlalchemy.orm import Session from configs import dify_config from core.repositories.sqlalchemy_workflow_node_execution_repository import ( @@ -22,7 +22,7 @@ from graphon.entities.workflow_node_execution import ( WorkflowNodeExecutionStatus, ) from graphon.enums import BuiltinNodeTypes -from models import Account, WorkflowNodeExecutionTriggeredFrom +from models import Account, Tenant, WorkflowNodeExecutionTriggeredFrom from models.enums import ExecutionOffLoadType from models.workflow import WorkflowNodeExecutionModel, WorkflowNodeExecutionOffload @@ -112,53 +112,58 @@ def create_workflow_node_execution( def mock_user() -> Account: - """Create a mock Account user for testing.""" - from unittest.mock import MagicMock + """Create an Account user for testing.""" - user = MagicMock(spec=Account) + user = Account(name="Test Account", email="test@example.com") user.id = "test-user-id" - user.current_tenant_id = "test-tenant-id" + user._current_tenant = Tenant(name="Test Tenant") + user._current_tenant.id = "test-tenant-id" return user class TestSQLAlchemyWorkflowNodeExecutionRepositoryTruncation: """Test class for truncation functionality in SQLAlchemyWorkflowNodeExecutionRepository.""" - def create_repository(self) -> SQLAlchemyWorkflowNodeExecutionRepository: - """Create a repository instance for testing.""" - return SQLAlchemyWorkflowNodeExecutionRepository( - session_factory=MagicMock(spec=Engine), + def create_repository(self, sqlite_engine: Engine) -> SQLAlchemyWorkflowNodeExecutionRepository: + """Create a repository backed by the test's isolated SQLite engine.""" + repository = SQLAlchemyWorkflowNodeExecutionRepository( + session_factory=sqlite_engine, tenant_id="test-tenant-id", user=mock_user(), app_id="test-app-id", triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, ) + with repository._session_factory() as session: + assert isinstance(session, Session) + assert session.get_bind() is sqlite_engine + return repository - def test_to_domain_model_without_offload_data(self): + def test_to_domain_model_without_offload_data(self, sqlite_engine: Engine): """Test _to_domain_model correctly handles models without offload data.""" - repo = self.create_repository() + repo = self.create_repository(sqlite_engine) - # Create a mock database model without offload data - db_model = WorkflowNodeExecutionModel() - db_model.id = "test-id" - db_model.node_execution_id = "node-exec-id" - db_model.workflow_id = "workflow-id" - db_model.workflow_run_id = "run-id" - db_model.index = 1 - db_model.predecessor_node_id = None - db_model.node_id = "node-id" - db_model.node_type = BuiltinNodeTypes.LLM - db_model.title = "Test Node" - db_model.inputs = json.dumps({"value": "inputs"}) - db_model.process_data = json.dumps({"value": "process_data"}) - db_model.outputs = json.dumps({"value": "outputs"}) - db_model.status = WorkflowNodeExecutionStatus.SUCCEEDED - db_model.error = None - db_model.elapsed_time = 1.0 - db_model.execution_metadata = "{}" - db_model.created_at = datetime.now(UTC) - db_model.finished_at = None - db_model.offload_data = [] + # Create a database model without offload data + db_model = WorkflowNodeExecutionModel( + id="test-id", + node_execution_id="node-exec-id", + workflow_id="workflow-id", + workflow_run_id="run-id", + index=1, + predecessor_node_id=None, + node_id="node-id", + node_type=BuiltinNodeTypes.LLM, + title="Test Node", + inputs=json.dumps({"value": "inputs"}), + process_data=json.dumps({"value": "process_data"}), + outputs=json.dumps({"value": "outputs"}), + status=WorkflowNodeExecutionStatus.SUCCEEDED, + error=None, + elapsed_time=1.0, + execution_metadata="{}", + created_at=datetime.now(UTC), + finished_at=None, + offload_data=[], + ) domain_model = repo._to_domain_model(db_model) @@ -202,8 +207,9 @@ class TestWorkflowNodeExecutionModelTruncatedProperties: def test_truncated_properties_without_offload_data(self): """Test truncated properties when no offload data exists.""" - model = WorkflowNodeExecutionModel() - model.offload_data = [] + model = WorkflowNodeExecutionModel( + offload_data=[], + ) assert model.inputs_truncated is False assert model.outputs_truncated is False diff --git a/api/tests/unit_tests/core/test_model_manager.py b/api/tests/unit_tests/core/test_model_manager.py index 5a7e7e30a50..fa873e8d865 100644 --- a/api/tests/unit_tests/core/test_model_manager.py +++ b/api/tests/unit_tests/core/test_model_manager.py @@ -4,10 +4,20 @@ import pytest import redis from pytest_mock import MockerFixture -from core.entities.provider_entities import ModelLoadBalancingConfiguration -from core.model_manager import LBModelManager, ModelManager +from core.entities.provider_entities import ( + ModelLoadBalancingConfiguration, + ProviderQuotaType, + QuotaConfiguration, + QuotaUnit, + RestrictModel, +) +from core.errors.error import ModelCurrentlyNotSupportError +from core.model_manager import LBModelManager, ModelInstance, ModelManager, QuotaManagedModelInstance from extensions.ext_redis import redis_client +from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage +from graphon.model_runtime.entities.message_entities import AssistantPromptMessage from graphon.model_runtime.entities.model_entities import ModelType +from models.provider import ProviderType @pytest.fixture @@ -63,6 +73,208 @@ def test_model_manager_with_cache_enabled_reuses_stored_credentials(): get_creds.assert_called_once() +def _build_model_manager_bundle( + *, + provider_type: ProviderType, + restrict_models: list[RestrictModel], +) -> tuple[ModelManager, MagicMock]: + provider_manager = MagicMock() + bundle = MagicMock() + bundle.configuration.provider.provider = "openai" + bundle.configuration.tenant_id = "tenant-1" + bundle.configuration.model_settings = [] + bundle.configuration.using_provider_type = provider_type + bundle.configuration.system_configuration.current_quota_type = ProviderQuotaType.TRIAL + bundle.configuration.system_configuration.quota_configurations = [ + QuotaConfiguration( + quota_type=ProviderQuotaType.TRIAL, + quota_unit=QuotaUnit.CREDITS, + quota_limit=200, + quota_used=0, + is_valid=True, + restrict_models=restrict_models, + ) + ] + bundle.configuration.get_current_credentials.return_value = {"api_key": "hosted"} + bundle.model_type_instance.model_type = ModelType.LLM + provider_manager.get_provider_model_bundle.return_value = bundle + return ModelManager(provider_manager), bundle + + +def test_model_manager_wraps_allowlisted_system_llm() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + + assert isinstance(model_instance, QuotaManagedModelInstance) + + +def test_model_manager_rejects_system_model_by_exact_name() -> None: + manager, bundle = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + + with pytest.raises(ModelCurrentlyNotSupportError, match="llm/gpt-4o is not allowed"): + manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4o") + + bundle.configuration.get_current_credentials.assert_not_called() + + +def test_model_manager_matches_allowlist_name_across_model_types() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="shared-model", model_type=ModelType.TEXT_EMBEDDING)], + ) + + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "shared-model") + + assert isinstance(model_instance, QuotaManagedModelInstance) + + +def test_quota_managed_non_streaming_invocation_finalizes_reservation() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + usage = LLMUsage.empty_usage().model_copy(update={"total_tokens": 12}) + result = MagicMock(spec=LLMResult, usage=usage) + reservation = MagicMock(commit_before_delivery=True) + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=result) as invoke, + ): + response = model_instance.invoke_llm(prompt_messages=[], stream=False) + + assert response is result + invoke.assert_called_once() + reservation.commit.assert_called_once_with(usage) + reservation.release.assert_called_once_with() + + +def test_quota_managed_stream_commits_before_first_chunk() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + chunk = LLMResultChunk( + model="gpt-4", + prompt_messages=[], + delta=LLMResultChunkDelta(index=0, message=AssistantPromptMessage(content="hello")), + ) + reservation = MagicMock(commit_before_delivery=True) + events: list[str] = [] + reservation.commit.side_effect = lambda _usage: events.append("commit") + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=(item for item in [chunk])), + ): + response = model_instance.invoke_llm(prompt_messages=[], stream=True) + assert next(response) is chunk + events.append("delivered") + with pytest.raises(StopIteration): + next(response) + + assert events == ["commit", "delivered"] + reservation.release.assert_called_once_with() + + +def test_quota_managed_stream_releases_when_provider_fails_before_first_chunk() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + reservation = MagicMock(commit_before_delivery=True) + + def failing_stream(): + raise RuntimeError("provider failed") + yield + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=failing_stream()), + pytest.raises(RuntimeError, match="provider failed"), + ): + list(model_instance.invoke_llm(prompt_messages=[], stream=True)) + + reservation.commit.assert_not_called() + reservation.release.assert_called_once_with() + + +def test_quota_managed_usage_stream_commits_before_delivering_buffered_chunks() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + usage = LLMUsage.empty_usage().model_copy(update={"total_tokens": 12}) + chunks = [ + LLMResultChunk( + model="gpt-4", + prompt_messages=[], + delta=LLMResultChunkDelta(index=0, message=AssistantPromptMessage(content="hello")), + ), + LLMResultChunk( + model="gpt-4", + prompt_messages=[], + delta=LLMResultChunkDelta(index=1, message=AssistantPromptMessage(content=" world"), usage=usage), + ), + ] + reservation = MagicMock(commit_before_delivery=False) + events: list[str] = [] + reservation.commit.side_effect = lambda _usage: events.append("commit") + + def provider_stream(): + for index, chunk in enumerate(chunks): + events.append(f"provider-{index}") + yield chunk + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=provider_stream()), + ): + response = model_instance.invoke_llm(prompt_messages=[], stream=True) + assert next(response) is chunks[0] + events.append("delivered") + assert list(response) == [chunks[1]] + + assert events == ["provider-0", "provider-1", "commit", "delivered"] + reservation.commit.assert_called_once_with(usage) + reservation.release.assert_called_once_with() + + +def test_quota_managed_usage_stream_does_not_deliver_when_settlement_fails() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + chunk = LLMResultChunk( + model="gpt-4", + prompt_messages=[], + delta=LLMResultChunkDelta(index=0, message=AssistantPromptMessage(content="hello")), + ) + reservation = MagicMock(commit_before_delivery=False) + reservation.commit.side_effect = ValueError("terminal usage is required") + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=(item for item in [chunk])), + pytest.raises(ValueError, match="terminal usage is required"), + ): + next(model_instance.invoke_llm(prompt_messages=[], stream=True)) + + reservation.release.assert_called_once_with() + + def test_lb_model_manager_fetch_next(mocker: MockerFixture, lb_model_manager: LBModelManager): # initialize redis client redis_client.initialize(redis.Redis()) diff --git a/api/tests/unit_tests/core/test_provider_manager.py b/api/tests/unit_tests/core/test_provider_manager.py index 128eebdd5af..bf805e761ab 100644 --- a/api/tests/unit_tests/core/test_provider_manager.py +++ b/api/tests/unit_tests/core/test_provider_manager.py @@ -1,33 +1,119 @@ +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor from types import SimpleNamespace from unittest.mock import MagicMock, Mock, PropertyMock, patch import pytest -from pytest_mock import MockerFixture +from flask import has_app_context +from sqlalchemy import Engine, event, select +from sqlalchemy.orm import Session, sessionmaker -from core.entities.provider_entities import ModelSettings +from core import provider_manager as provider_manager_module +from core.entities.provider_entities import ( + CustomConfiguration, + CustomProviderConfiguration, + ModelSettings, + ProviderQuotaType, +) +from core.hosting_configuration import HostingProvider, TrialHostingQuota +from core.plugin.entities.plugin import PluginInstallationSource +from core.plugin.entities.plugin_daemon import PluginModelProviderDeclaration from core.provider_manager import ProviderConfigurationCacheSource, ProviderManager +from enums import DeploymentEdition from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import ModelType +from graphon.model_runtime.entities.provider_entities import ConfigurateMethod +from models.base import TypeBase from models.provider import ( LoadBalancingModelConfig, Provider, ProviderCredential, + ProviderModel, + ProviderModelCredential, ProviderModelSetting, ProviderType, TenantDefaultModel, + TenantPreferredModelProvider, ) from models.provider_ids import ModelProviderID -def _build_provider_manager(mocker: MockerFixture) -> ProviderManager: - return ProviderManager(model_runtime=mocker.Mock()) +def _build_provider_manager() -> ProviderManager: + return ProviderManager(model_runtime=Mock()) -def _build_session_context(session: Mock) -> MagicMock: - session_cm = MagicMock() - session_cm.__enter__.return_value = session - session_cm.__exit__.return_value = False - return session_cm +def _persist_model_configuration( + session: Session, + setting: ProviderModelSetting, + load_balancing_configs: list[LoadBalancingModelConfig], +) -> tuple[list[ProviderModelSetting], list[LoadBalancingModelConfig]]: + session.add_all([setting, *load_balancing_configs]) + session.commit() + session.expire_all() + settings = list( + session.scalars(select(ProviderModelSetting).where(ProviderModelSetting.tenant_id == setting.tenant_id)).all() + ) + configs = list( + session.scalars( + select(LoadBalancingModelConfig).where(LoadBalancingModelConfig.tenant_id == setting.tenant_id) + ).all() + ) + return settings, configs + + +@pytest.fixture +def provider_db(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Session]: + """Bind request-owned and provider-owned sessions to one isolated SQLite database.""" + models = ( + Provider, + ProviderCredential, + ProviderModel, + ProviderModelCredential, + ProviderModelSetting, + LoadBalancingModelConfig, + TenantDefaultModel, + TenantPreferredModelProvider, + ) + tables = [model.metadata.tables[model.__tablename__] for model in models] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + owned_session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + with owned_session_factory() as request_session: + monkeypatch.setattr(provider_manager_module.db, "session", request_session) + monkeypatch.setattr(provider_manager_module.session_factory, "create_session", owned_session_factory) + yield request_session + + +def _build_plugin_provider_declaration( + installation_source: PluginInstallationSource | None, +) -> PluginModelProviderDeclaration: + return PluginModelProviderDeclaration( + provider="langgenius/openai/openai", + plugin_unique_identifier="langgenius/openai:1.0.0@checksum", + installation_source=installation_source, + label=I18nObject(en_US="OpenAI"), + supported_model_types=[ModelType.LLM], + configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL], + ) + + +def _build_hosting_provider() -> HostingProvider: + return HostingProvider( + enabled=True, + credentials={"api_key": "system-secret"}, + quotas=[TrialHostingQuota(quota_limit=100)], + ) + + +def _build_trial_provider_record() -> Provider: + return Provider( + tenant_id="tenant-id", + provider_name="openai", + provider_type=ProviderType.SYSTEM, + quota_type=ProviderQuotaType.TRIAL, + quota_limit=100, + quota_used=0, + is_valid=True, + ) class _FakeRedis: @@ -56,17 +142,6 @@ class _FakeRedis: self.expirations[key] = time -class _FakeScalarResult: - def __init__(self, values: list[object]) -> None: - self._values = values - - def __iter__(self): - return iter(self._values) - - def all(self) -> list[object]: - return self._values - - @pytest.fixture def mock_provider_entity(): mock_entity = Mock() @@ -86,7 +161,7 @@ def mock_provider_entity(): return mock_entity -def test__to_model_settings(mocker: MockerFixture, mock_provider_entity): +def test__to_model_settings(mock_provider_entity, provider_db: Session): # Mocking the inputs ps = ProviderModelSetting( tenant_id="tenant_id", @@ -122,12 +197,15 @@ def test__to_model_settings(mocker: MockerFixture, mock_provider_entity): ] load_balancing_model_configs[0].id = "id1" load_balancing_model_configs[1].id = "id2" + provider_model_settings, load_balancing_model_configs = _persist_model_configuration( + provider_db, ps, load_balancing_model_configs + ) with patch( "core.helper.model_provider_cache.ProviderCredentialsCache.get", return_value={"openai_api_key": "fake_key"}, ): - provider_manager = _build_provider_manager(mocker) + provider_manager = _build_provider_manager() # Running the method result = provider_manager._to_model_settings( @@ -147,7 +225,181 @@ def test__to_model_settings(mocker: MockerFixture, mock_provider_entity): assert result[0].load_balancing_configs[1].name == "first" -def test__to_model_settings_only_one_lb(mocker: MockerFixture, mock_provider_entity): +@pytest.mark.parametrize( + "installation_source", + [ + None, + PluginInstallationSource.Github, + PluginInstallationSource.Package, + PluginInstallationSource.Remote, + ], +) +def test_to_system_configuration_rejects_non_marketplace_provider( + installation_source: PluginInstallationSource | None, +) -> None: + provider_entity = _build_plugin_provider_declaration(installation_source) + manager = _build_provider_manager() + + with ( + patch.object(manager, "_choice_current_using_quota_type") as choose_quota, + patch( + "core.provider_manager.ext_hosting_provider.hosting_configuration.provider_map", + {provider_entity.provider: _build_hosting_provider()}, + ), + ): + configuration = manager._to_system_configuration("tenant-id", provider_entity, []) + + assert configuration.enabled is False + assert configuration.credentials is None + choose_quota.assert_not_called() + + +def test_to_system_configuration_rejects_unverified_marketplace_provider() -> None: + provider_entity = _build_plugin_provider_declaration(PluginInstallationSource.Marketplace) + manager = _build_provider_manager() + + with ( + patch.object(manager, "_choice_current_using_quota_type") as choose_quota, + patch( + "core.provider_manager.ext_hosting_provider.hosting_configuration.provider_map", + {provider_entity.provider: _build_hosting_provider()}, + ), + patch( + "core.plugin.plugin_service.PluginService.is_plugin_verified", + return_value=False, + ) as is_plugin_verified, + ): + configuration = manager._to_system_configuration("tenant-id", provider_entity, []) + + assert configuration.enabled is False + assert configuration.credentials is None + is_plugin_verified.assert_called_once_with("tenant-id", provider_entity.plugin_unique_identifier) + choose_quota.assert_not_called() + + +def test_to_system_configuration_never_returns_hosting_credentials_for_package_with_valid_quota() -> None: + provider_entity = _build_plugin_provider_declaration(PluginInstallationSource.Package) + manager = _build_provider_manager() + + with patch( + "core.provider_manager.ext_hosting_provider.hosting_configuration.provider_map", + {provider_entity.provider: _build_hosting_provider()}, + ): + configuration = manager._to_system_configuration( + "tenant-id", + provider_entity, + [_build_trial_provider_record()], + ) + + assert configuration.enabled is False + assert configuration.credentials is None + + +def test_to_system_configuration_uses_owned_session_for_cloud_credit_pools() -> None: + provider_entity = _build_plugin_provider_declaration(PluginInstallationSource.Marketplace) + manager = _build_provider_manager() + trial_pool = SimpleNamespace(quota_used=0, quota_limit=100) + paid_pool = SimpleNamespace(quota_used=0, quota_limit=0) + + with ( + patch.object(provider_manager_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + patch( + "core.provider_manager.ext_hosting_provider.hosting_configuration.provider_map", + {provider_entity.provider: _build_hosting_provider()}, + ), + patch( + "core.plugin.plugin_service.PluginService.is_plugin_verified", + return_value=True, + ), + patch( + "services.credit_pool_service.CreditPoolService.get_pool", + side_effect=[trial_pool, paid_pool], + ) as get_pool, + ): + with ThreadPoolExecutor(max_workers=1) as executor: + assert executor.submit(has_app_context).result() is False + configuration = executor.submit( + manager._to_system_configuration, + "tenant-id", + provider_entity, + [_build_trial_provider_record()], + ).result() + + assert configuration.enabled is True + assert [call.kwargs["pool_type"] for call in get_pool.call_args_list] == [ + ProviderQuotaType.TRIAL, + ProviderQuotaType.PAID, + ] + owned_sessions = [call.kwargs["session"] for call in get_pool.call_args_list] + assert all(isinstance(session, Session) for session in owned_sessions) + assert owned_sessions[0] is owned_sessions[1] + + +def test_to_system_configuration_preserves_marketplace_behavior() -> None: + provider_entity = _build_plugin_provider_declaration(PluginInstallationSource.Marketplace) + manager = _build_provider_manager() + + with ( + patch( + "core.provider_manager.ext_hosting_provider.hosting_configuration.provider_map", + {provider_entity.provider: _build_hosting_provider()}, + ), + patch( + "core.plugin.plugin_service.PluginService.is_plugin_verified", + return_value=True, + ) as is_plugin_verified, + ): + configuration = manager._to_system_configuration( + "tenant-id", + provider_entity, + [_build_trial_provider_record()], + ) + + assert configuration.enabled is True + assert configuration.credentials == {"api_key": "system-secret"} + assert configuration.current_quota_type == ProviderQuotaType.TRIAL + is_plugin_verified.assert_called_once_with("tenant-id", provider_entity.plugin_unique_identifier) + + +def test_package_provider_keeps_custom_configuration() -> None: + provider_entity = _build_plugin_provider_declaration(PluginInstallationSource.Package) + manager = _build_provider_manager() + provider_factory = Mock() + provider_factory.get_providers.return_value = [provider_entity] + custom_configuration = CustomConfiguration( + provider=CustomProviderConfiguration(credentials={"api_key": "user-secret"}) + ) + + with ( + patch.object(manager, "_get_all_providers", return_value={provider_entity.provider: []}), + patch.object( + manager, + "_init_trial_provider_records", + return_value={provider_entity.provider: []}, + ), + patch.object(manager, "_get_all_provider_models", return_value={}), + patch.object(manager, "_get_all_preferred_model_providers", return_value={}), + patch.object(manager, "_get_all_provider_model_settings", return_value={}), + patch.object(manager, "_get_all_provider_load_balancing_configs", return_value={}), + patch.object(manager, "_get_all_provider_model_credentials", return_value={}), + patch.object(manager, "_get_all_provider_credentials", return_value={}), + patch.object(manager, "_to_custom_configuration", return_value=custom_configuration), + patch.object(manager, "_to_model_settings", return_value=[]), + patch("core.provider_manager.ModelProviderFactory", return_value=provider_factory), + patch( + "core.provider_manager.ext_hosting_provider.hosting_configuration.provider_map", + {provider_entity.provider: _build_hosting_provider()}, + ), + ): + configuration = manager.get_configurations("tenant-id").get(provider_entity.provider) + + assert configuration is not None + assert configuration.system_configuration.enabled is False + assert configuration.custom_configuration.provider is not None + assert configuration.custom_configuration.provider.credentials == {"api_key": "user-secret"} + + +def test__to_model_settings_only_one_lb(mock_provider_entity, provider_db: Session): # Mocking the inputs ps = ProviderModelSetting( @@ -172,12 +424,15 @@ def test__to_model_settings_only_one_lb(mocker: MockerFixture, mock_provider_ent ) ] load_balancing_model_configs[0].id = "id1" + provider_model_settings, load_balancing_model_configs = _persist_model_configuration( + provider_db, ps, load_balancing_model_configs + ) with patch( "core.helper.model_provider_cache.ProviderCredentialsCache.get", return_value={"openai_api_key": "fake_key"}, ): - provider_manager = _build_provider_manager(mocker) + provider_manager = _build_provider_manager() # Running the method result = provider_manager._to_model_settings( @@ -195,7 +450,7 @@ def test__to_model_settings_only_one_lb(mocker: MockerFixture, mock_provider_ent assert len(result[0].load_balancing_configs) == 0 -def test__to_model_settings_lb_disabled(mocker: MockerFixture, mock_provider_entity): +def test__to_model_settings_lb_disabled(mock_provider_entity, provider_db: Session): # Mocking the inputs ps = ProviderModelSetting( tenant_id="tenant_id", @@ -229,12 +484,15 @@ def test__to_model_settings_lb_disabled(mocker: MockerFixture, mock_provider_ent ] load_balancing_model_configs[0].id = "id1" load_balancing_model_configs[1].id = "id2" + provider_model_settings, load_balancing_model_configs = _persist_model_configuration( + provider_db, ps, load_balancing_model_configs + ) with patch( "core.helper.model_provider_cache.ProviderCredentialsCache.get", return_value={"openai_api_key": "fake_key"}, ): - provider_manager = _build_provider_manager(mocker) + provider_manager = _build_provider_manager() # Running the method result = provider_manager._to_model_settings( @@ -252,19 +510,23 @@ def test__to_model_settings_lb_disabled(mocker: MockerFixture, mock_provider_ent assert len(result[0].load_balancing_configs) == 0 -def test_get_default_model_uses_first_available_active_model(mocker: MockerFixture): - mock_session = Mock() - mock_session.scalar.return_value = None - +def test_get_default_model_uses_first_available_active_model(provider_db: Session): + other_tenant_default = TenantDefaultModel( + tenant_id="other-tenant", + provider_name="anthropic", + model_name="claude", + model_type=ModelType.LLM, + ) + provider_db.add(other_tenant_default) + provider_db.commit() provider_configurations = Mock() provider_configurations.get_models.return_value = [ Mock(model="gpt-3.5-turbo", provider=Mock(provider="openai")), Mock(model="gpt-4", provider=Mock(provider="openai")), ] - manager = _build_provider_manager(mocker) + manager = _build_provider_manager() with ( - patch("core.provider_manager.db.session", mock_session), patch.object(manager, "get_configurations", return_value=provider_configurations), patch("core.provider_manager.ModelProviderFactory") as mock_factory_cls, ): @@ -281,22 +543,33 @@ def test_get_default_model_uses_first_available_active_model(mocker: MockerFixtu assert result.model == "gpt-3.5-turbo" assert result.provider.provider == "openai" provider_configurations.get_models.assert_called_once_with(model_type=ModelType.LLM, only_active=True) - mock_session.add.assert_called_once() - saved_default_model = mock_session.add.call_args.args[0] + saved_default_model = provider_db.scalar( + select(TenantDefaultModel).where(TenantDefaultModel.tenant_id == "tenant-id") + ) + assert saved_default_model is not None assert saved_default_model.model_name == "gpt-3.5-turbo" assert saved_default_model.provider_name == "openai" - mock_session.commit.assert_called_once() + assert ( + provider_db.scalar(select(TenantDefaultModel).where(TenantDefaultModel.tenant_id == "other-tenant")) + is other_tenant_default + ) -def test_get_default_model_returns_none_when_no_default_or_active_models(mocker: MockerFixture): - mock_session = Mock() - mock_session.scalar.return_value = None +def test_get_default_model_returns_none_when_no_default_or_active_models(provider_db: Session): + provider_db.add( + TenantDefaultModel( + tenant_id="other-tenant", + provider_name="anthropic", + model_name="claude", + model_type=ModelType.LLM, + ) + ) + provider_db.commit() provider_configurations = Mock() provider_configurations.get_models.return_value = [] - manager = _build_provider_manager(mocker) + manager = _build_provider_manager() with ( - patch("core.provider_manager.db.session", mock_session), patch.object(manager, "get_configurations", return_value=provider_configurations), patch("core.provider_manager.ModelProviderFactory") as mock_factory_cls, ): @@ -305,23 +578,24 @@ def test_get_default_model_returns_none_when_no_default_or_active_models(mocker: assert result is None provider_configurations.get_models.assert_called_once_with(model_type=ModelType.LLM, only_active=True) mock_factory_cls.assert_not_called() - mock_session.add.assert_not_called() - mock_session.commit.assert_not_called() + assert provider_db.scalar(select(TenantDefaultModel).where(TenantDefaultModel.tenant_id == "tenant-id")) is None + assert ( + provider_db.scalar(select(TenantDefaultModel).where(TenantDefaultModel.tenant_id == "other-tenant")) is not None + ) -def test_get_default_model_uses_injected_runtime_for_existing_default_record(mocker: MockerFixture): +def test_get_default_model_uses_injected_runtime_for_existing_default_record(provider_db: Session): existing_default_model = TenantDefaultModel( tenant_id="tenant-id", provider_name="openai", model_name="gpt-4", model_type=ModelType.LLM, ) - mock_session = Mock() - mock_session.scalar.return_value = existing_default_model - manager = _build_provider_manager(mocker) + provider_db.add(existing_default_model) + provider_db.commit() + manager = _build_provider_manager() with ( - patch("core.provider_manager.db.session", mock_session), patch("core.provider_manager.ModelProviderFactory") as mock_factory_cls, ): mock_factory_cls.return_value.get_provider_schema.return_value = Mock( @@ -339,11 +613,32 @@ def test_get_default_model_uses_injected_runtime_for_existing_default_record(moc assert result.provider.provider == "openai" -def test_get_configurations_uses_injected_runtime_and_adds_provider_aliases(mocker: MockerFixture): - manager = _build_provider_manager(mocker) - provider_records = {"openai": [SimpleNamespace(provider_name="openai")]} - provider_model_records = {"openai": [SimpleNamespace(provider_name="openai")]} - preferred_provider_records = {"openai": SimpleNamespace(preferred_provider_type="system")} +def test_get_configurations_uses_injected_runtime_and_adds_provider_aliases(provider_db: Session): + manager = _build_provider_manager() + provider = Provider( + tenant_id="tenant-id", + provider_name="openai", + provider_type=ProviderType.CUSTOM, + is_valid=True, + ) + provider_model = ProviderModel( + tenant_id="tenant-id", + provider_name="openai", + model_name="gpt-4", + model_type=ModelType.LLM, + is_valid=True, + ) + preferred_provider = TenantPreferredModelProvider( + tenant_id="tenant-id", + provider_name="openai", + preferred_provider_type=ProviderType.SYSTEM, + ) + provider_db.add_all([provider, provider_model, preferred_provider]) + provider_db.commit() + with patch("core.provider_manager.redis_client", _FakeRedis()): + provider_model_records = ProviderManager._get_all_provider_models("tenant-id") + preferred_provider_records = ProviderManager._get_all_preferred_model_providers("tenant-id") + provider_records = {"openai": [provider]} with ( patch.object(manager, "_get_all_providers", return_value=provider_records), @@ -380,18 +675,16 @@ def test_get_provider_names_returns_short_and_full_aliases(provider_name: str, e assert ProviderManager._get_provider_names(provider_name) == expected_provider_names -def test_get_provider_model_bundle_raises_for_unknown_provider(mocker: MockerFixture): - manager = _build_provider_manager(mocker) +def test_get_provider_model_bundle_raises_for_unknown_provider(): + manager = _build_provider_manager() with patch.object(manager, "get_configurations", return_value={}): with pytest.raises(ValueError, match="Provider openai does not exist."): manager.get_provider_model_bundle("tenant-id", "openai", ModelType.LLM) -def test_get_configurations_binds_manager_runtime_to_provider_configuration( - mocker: MockerFixture, mock_provider_entity -): - manager = _build_provider_manager(mocker) +def test_get_configurations_binds_manager_runtime_to_provider_configuration(mock_provider_entity): + manager = _build_provider_manager() provider_configuration = Mock() provider_factory = Mock() provider_factory.get_providers.return_value = [mock_provider_entity] @@ -418,8 +711,65 @@ def test_get_configurations_binds_manager_runtime_to_provider_configuration( provider_configuration.bind_model_runtime.assert_called_once_with(manager._model_runtime) -def test_get_configurations_reuses_cached_result_for_same_tenant(mocker: MockerFixture, mock_provider_entity): - manager = _build_provider_manager(mocker) +@pytest.mark.parametrize( + ("quota_is_valid", "has_custom_provider", "has_custom_model", "expected_using_provider_type"), + [ + (True, False, False, ProviderType.SYSTEM), + (False, False, False, ProviderType.SYSTEM), + (False, True, False, ProviderType.CUSTOM), + (False, False, True, ProviderType.CUSTOM), + (None, False, False, ProviderType.CUSTOM), + ], +) +def test_get_configurations_resolves_system_preference_with_quota_and_custom_fallback( + mock_provider_entity, + quota_is_valid: bool | None, + has_custom_provider: bool, + has_custom_model: bool, + expected_using_provider_type: ProviderType, +): + manager = _build_provider_manager() + provider_configuration = Mock() + provider_factory = Mock() + provider_factory.get_providers.return_value = [mock_provider_entity] + custom_configuration = SimpleNamespace( + provider=Mock() if has_custom_provider else None, + models=[Mock()] if has_custom_model else [], + ) + system_configuration = SimpleNamespace( + enabled=True, + quota_configurations=[] if quota_is_valid is None else [SimpleNamespace(is_valid=quota_is_valid)], + current_quota_type=None, + ) + preferred_provider_records = { + "openai": SimpleNamespace(preferred_provider_type=ProviderType.SYSTEM), + } + + with ( + patch.object(manager, "_get_all_providers", return_value={"openai": []}), + patch.object(manager, "_init_trial_provider_records", return_value={"openai": []}), + patch.object(manager, "_get_all_provider_models", return_value={"openai": []}), + patch.object(manager, "_get_all_preferred_model_providers", return_value=preferred_provider_records), + patch.object(manager, "_get_all_provider_model_settings", return_value={}), + patch.object(manager, "_get_all_provider_load_balancing_configs", return_value={}), + patch.object(manager, "_get_all_provider_model_credentials", return_value={}), + patch.object(manager, "_get_all_provider_credentials", return_value={}), + patch.object(manager, "_to_custom_configuration", return_value=custom_configuration), + patch.object(manager, "_to_system_configuration", return_value=system_configuration), + patch.object(manager, "_to_model_settings", return_value=[]), + patch("core.provider_manager.ModelProviderFactory", return_value=provider_factory), + patch( + "core.provider_manager.ProviderConfiguration", + return_value=provider_configuration, + ) as mock_provider_configuration, + ): + manager.get_configurations("tenant-id") + + assert mock_provider_configuration.call_args.kwargs["using_provider_type"] == expected_using_provider_type + + +def test_get_configurations_reuses_cached_result_for_same_tenant(mock_provider_entity): + manager = _build_provider_manager() provider_configuration = Mock() provider_factory = Mock() provider_factory.get_providers.return_value = [mock_provider_entity] @@ -454,8 +804,8 @@ def test_get_configurations_reuses_cached_result_for_same_tenant(mocker: MockerF provider_configuration.bind_model_runtime.assert_called_once_with(manager._model_runtime) -def test_clear_configurations_cache_rebuilds_requested_tenant(mocker: MockerFixture, mock_provider_entity): - manager = _build_provider_manager(mocker) +def test_clear_configurations_cache_rebuilds_requested_tenant(mock_provider_entity): + manager = _build_provider_manager() provider_factory = Mock() provider_factory.get_providers.return_value = [mock_provider_entity] custom_configuration = SimpleNamespace(provider=None, models=[]) @@ -492,8 +842,8 @@ def test_clear_configurations_cache_rebuilds_requested_tenant(mocker: MockerFixt provider_configuration_second.bind_model_runtime.assert_called_once_with(manager._model_runtime) -def test_get_provider_model_bundle_returns_selected_model_type_instance(mocker: MockerFixture): - manager = _build_provider_manager(mocker) +def test_get_provider_model_bundle_returns_selected_model_type_instance(): + manager = _build_provider_manager() provider_configuration = Mock() model_type_instance = Mock() provider_configuration.get_model_type_instance.return_value = model_type_instance @@ -513,8 +863,8 @@ def test_get_provider_model_bundle_returns_selected_model_type_instance(mocker: assert result is expected_bundle -def test_get_first_provider_first_model_returns_none_when_no_models(mocker: MockerFixture): - manager = _build_provider_manager(mocker) +def test_get_first_provider_first_model_returns_none_when_no_models(): + manager = _build_provider_manager() provider_configurations = Mock() provider_configurations.get_models.return_value = [] @@ -525,8 +875,8 @@ def test_get_first_provider_first_model_returns_none_when_no_models(mocker: Mock provider_configurations.get_models.assert_called_once_with(model_type=ModelType.LLM, only_active=False) -def test_get_first_provider_first_model_returns_first_model_and_provider(mocker: MockerFixture): - manager = _build_provider_manager(mocker) +def test_get_first_provider_first_model_returns_first_model_and_provider(): + manager = _build_provider_manager() provider_configurations = Mock() provider_configurations.get_models.return_value = [ Mock(model="gpt-4", provider=Mock(provider="openai")), @@ -539,16 +889,16 @@ def test_get_first_provider_first_model_returns_first_model_and_provider(mocker: assert result == ("openai", "gpt-4") -def test_update_default_model_record_raises_for_unknown_provider(mocker: MockerFixture): - manager = _build_provider_manager(mocker) +def test_update_default_model_record_raises_for_unknown_provider(): + manager = _build_provider_manager() with patch.object(manager, "get_configurations", return_value={}): with pytest.raises(ValueError, match="Provider openai does not exist."): manager.update_default_model_record("tenant-id", ModelType.LLM, "openai", "gpt-4") -def test_update_default_model_record_raises_for_unknown_model(mocker: MockerFixture): - manager = _build_provider_manager(mocker) +def test_update_default_model_record_raises_for_unknown_model(): + manager = _build_provider_manager() provider_configurations = MagicMock() provider_configurations.__contains__.return_value = True provider_configurations.get_models.return_value = [Mock(model="gpt-4")] @@ -560,8 +910,8 @@ def test_update_default_model_record_raises_for_unknown_model(mocker: MockerFixt provider_configurations.get_models.assert_called_once_with(model_type=ModelType.LLM, only_active=True) -def test_update_default_model_record_updates_existing_record(mocker: MockerFixture): - manager = _build_provider_manager(mocker) +def test_update_default_model_record_updates_existing_record(provider_db: Session): + manager = _build_provider_manager() provider_configurations = MagicMock() provider_configurations.__contains__.return_value = True provider_configurations.get_models.return_value = [Mock(model="gpt-3.5-turbo")] @@ -571,62 +921,87 @@ def test_update_default_model_record_updates_existing_record(mocker: MockerFixtu model_name="claude-3-sonnet", model_type=ModelType.LLM, ) - mock_session = Mock() - mock_session.scalar.return_value = existing_default_model + other_tenant_default = TenantDefaultModel( + tenant_id="other-tenant", + provider_name="cohere", + model_name="command-r", + model_type=ModelType.LLM, + ) + provider_db.add_all([existing_default_model, other_tenant_default]) + provider_db.commit() - with ( - patch.object(manager, "get_configurations", return_value=provider_configurations), - patch("core.provider_manager.db.session", mock_session), - ): + with patch.object(manager, "get_configurations", return_value=provider_configurations): result = manager.update_default_model_record("tenant-id", ModelType.LLM, "openai", "gpt-3.5-turbo") assert result is existing_default_model assert existing_default_model.provider_name == "openai" assert existing_default_model.model_name == "gpt-3.5-turbo" - mock_session.commit.assert_called_once() - mock_session.add.assert_not_called() + provider_db.expire_all() + persisted_default = provider_db.get(TenantDefaultModel, existing_default_model.id) + assert persisted_default is not None + assert persisted_default.provider_name == "openai" + persisted_other_default = provider_db.get(TenantDefaultModel, other_tenant_default.id) + assert persisted_other_default is not None + assert persisted_other_default.provider_name == "cohere" + assert persisted_other_default.model_name == "command-r" -def test_update_default_model_record_creates_record_with_origin_model_type(mocker: MockerFixture): - manager = _build_provider_manager(mocker) +def test_update_default_model_record_creates_record_with_origin_model_type(provider_db: Session): + manager = _build_provider_manager() provider_configurations = MagicMock() provider_configurations.__contains__.return_value = True provider_configurations.get_models.return_value = [Mock(model="gpt-4")] - mock_session = Mock() - mock_session.scalar.return_value = None + provider_db.add( + TenantDefaultModel( + tenant_id="other-tenant", + provider_name="anthropic", + model_name="claude", + model_type=ModelType.LLM, + ) + ) + provider_db.commit() - with ( - patch.object(manager, "get_configurations", return_value=provider_configurations), - patch("core.provider_manager.db.session", mock_session), - ): + with patch.object(manager, "get_configurations", return_value=provider_configurations): result = manager.update_default_model_record("tenant-id", ModelType.LLM, "openai", "gpt-4") - mock_session.add.assert_called_once() - created_default_model = mock_session.add.call_args.args[0] - assert result is created_default_model - assert created_default_model.tenant_id == "tenant-id" - assert created_default_model.provider_name == "openai" - assert created_default_model.model_name == "gpt-4" - assert created_default_model.model_type == ModelType.LLM - mock_session.commit.assert_called_once() + assert result.tenant_id == "tenant-id" + assert result.provider_name == "openai" + assert result.model_name == "gpt-4" + assert result.model_type == ModelType.LLM + persisted_defaults = list(provider_db.scalars(select(TenantDefaultModel)).all()) + assert {record.tenant_id for record in persisted_defaults} == {"tenant-id", "other-tenant"} -def test_get_all_providers_normalizes_provider_names_with_model_provider_id() -> None: - session = Mock() - openai_provider = SimpleNamespace(provider_name="openai") - gemini_provider = SimpleNamespace(provider_name="langgenius/gemini/google") - session.scalars.return_value = [openai_provider, gemini_provider] +def test_get_all_providers_normalizes_provider_names_with_model_provider_id(provider_db: Session) -> None: + openai_provider = Provider( + tenant_id="tenant-id", + provider_name="openai", + provider_type=ProviderType.CUSTOM, + is_valid=True, + ) + gemini_provider = Provider( + tenant_id="tenant-id", + provider_name="langgenius/gemini/google", + provider_type=ProviderType.CUSTOM, + is_valid=True, + ) + other_tenant_provider = Provider( + tenant_id="other-tenant", + provider_name="anthropic", + provider_type=ProviderType.CUSTOM, + is_valid=True, + ) + provider_db.add_all([openai_provider, gemini_provider, other_tenant_provider]) + provider_db.commit() - with ( - patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), - ): - result = ProviderManager._get_all_providers("tenant-id") + result = ProviderManager._get_all_providers("tenant-id") assert list(result[str(ModelProviderID("openai"))]) == [openai_provider] assert list(result[str(ModelProviderID("langgenius/gemini/google"))]) == [gemini_provider] + assert str(ModelProviderID("anthropic")) not in result -def test_get_all_providers_attaches_active_credentials() -> None: +def test_get_all_providers_attaches_active_credentials(provider_db: Session) -> None: provider = Provider( tenant_id="tenant-id", provider_name="openai", @@ -634,7 +1009,6 @@ def test_get_all_providers_attaches_active_credentials() -> None: is_valid=True, credential_id="credential-id", ) - provider.id = "provider-id" credential = ProviderCredential( tenant_id="tenant-id", provider_name="openai", @@ -642,20 +1016,25 @@ def test_get_all_providers_attaches_active_credentials() -> None: encrypted_config='{"api_key": "secret"}', ) credential.id = "credential-id" - session = Mock() - session.scalars.side_effect = [ - _FakeScalarResult([provider]), - _FakeScalarResult([credential]), - ] + provider_db.add_all( + [ + provider, + credential, + Provider( + tenant_id="other-tenant", + provider_name="anthropic", + provider_type=ProviderType.CUSTOM, + is_valid=True, + ), + ] + ) + provider_db.commit() - with ( - patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), - ): - result = ProviderManager._get_all_providers("tenant-id") + result = ProviderManager._get_all_providers("tenant-id") - assert session.scalars.call_count == 2 assert result[str(ModelProviderID("openai"))][0].credential_name == "primary" assert result[str(ModelProviderID("openai"))][0].encrypted_config == '{"api_key": "secret"}' + assert str(ModelProviderID("anthropic")) not in result def test_invalidate_configurations_cache_bumps_selected_source_version() -> None: @@ -676,57 +1055,93 @@ def test_invalidate_configurations_cache_bumps_selected_source_version() -> None assert "provider_configurations:tenant:tenant-id:source:provider_models:version" not in fake_redis.store -def test_provider_model_credentials_cache_returns_cache_entries() -> None: +def test_provider_model_credentials_cache_returns_cache_entries(provider_db: Session, sqlite_engine: Engine) -> None: fake_redis = _FakeRedis() - credential_record = SimpleNamespace( - id="credential-id", + credential_record = ProviderModelCredential( + tenant_id="tenant-id", provider_name="openai", model_name="gpt-4", model_type=ModelType.LLM, credential_name="primary", + encrypted_config='{"api_key": "secret"}', ) - session = Mock() - session.scalars.return_value = [credential_record] + credential_record.id = "credential-id" + provider_db.add_all( + [ + credential_record, + ProviderModelCredential( + tenant_id="other-tenant", + provider_name="anthropic", + model_name="claude", + model_type=ModelType.LLM, + credential_name="other", + encrypted_config='{"api_key": "other"}', + ), + ] + ) + provider_db.commit() - with ( - patch("core.provider_manager.redis_client", fake_redis), - patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), - ): - first = ProviderManager._get_all_provider_model_credentials("tenant-id") - second = ProviderManager._get_all_provider_model_credentials("tenant-id") + credential_query_count = 0 + + def count_credential_queries(_conn, _cursor, statement, _parameters, _context, _executemany) -> None: + nonlocal credential_query_count + if statement.lstrip().upper().startswith("SELECT") and "provider_model_credentials" in statement: + credential_query_count += 1 + + # Count real SQL statements now that this test uses SQLite instead of a mocked session. + event.listen(sqlite_engine, "before_cursor_execute", count_credential_queries) + try: + with patch("core.provider_manager.redis_client", fake_redis): + first = ProviderManager._get_all_provider_model_credentials("tenant-id") + second = ProviderManager._get_all_provider_model_credentials("tenant-id") + finally: + event.remove(sqlite_engine, "before_cursor_execute", count_credential_queries) - assert session.scalars.call_count == 1 version_key = "provider_configurations:tenant:tenant-id:source:provider_model_credentials:version" + assert credential_query_count == 1 assert fake_redis.expirations[version_key] == 360 assert first["openai"][0] is not credential_record assert second["openai"][0].credential_name == "primary" assert second["openai"][0].model_type == ModelType.LLM + assert "anthropic" not in second -def test_provider_configuration_cache_skips_write_when_version_changes_during_load() -> None: +def test_provider_configuration_cache_skips_write_when_version_changes_during_load( + provider_db: Session, sqlite_engine: Engine +) -> None: fake_redis = _FakeRedis() version_key = "provider_configurations:tenant:tenant-id:source:provider_model_credentials:version" - credential_record = SimpleNamespace( - id="credential-id", + credential_record = ProviderModelCredential( + tenant_id="tenant-id", provider_name="openai", model_name="gpt-4", model_type=ModelType.LLM, credential_name="primary", + encrypted_config='{"api_key": "secret"}', ) - session = Mock() + credential_record.id = "credential-id" + provider_db.add(credential_record) + provider_db.commit() + version_bumped = False - def load_records(_stmt): - fake_redis.incr(version_key) - return [credential_record] + def bump_version_on_credential_query(_conn, _cursor, statement, _parameters, _context, _executemany) -> None: + nonlocal version_bumped + if ( + not version_bumped + and statement.lstrip().upper().startswith("SELECT") + and "provider_model_credentials" in statement + ): + version_bumped = True + fake_redis.incr(version_key) - session.scalars.side_effect = load_records - - with ( - patch("core.provider_manager.redis_client", fake_redis), - patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), - ): - result = ProviderManager._get_all_provider_model_credentials("tenant-id") + event.listen(sqlite_engine, "before_cursor_execute", bump_version_on_credential_query) + try: + with patch("core.provider_manager.redis_client", fake_redis): + result = ProviderManager._get_all_provider_model_credentials("tenant-id") + finally: + event.remove(sqlite_engine, "before_cursor_execute", bump_version_on_credential_query) + assert version_bumped is True assert fake_redis.store[version_key] == "1" assert "provider_configurations:tenant:tenant-id:source:provider_model_credentials:v:0" not in fake_redis.store assert "provider_configurations:tenant:tenant-id:source:provider_model_credentials:v:1" not in fake_redis.store @@ -742,19 +1157,21 @@ def test_provider_configuration_cache_skips_write_when_version_changes_during_lo "_get_all_provider_credentials", ], ) -def test_provider_grouping_helpers_group_records_by_provider_name(method_name: str) -> None: - def build_record(provider_name: str, index: int): +def test_provider_grouping_helpers_group_records_by_provider_name(method_name: str, provider_db: Session) -> None: + def build_record(provider_name: str, index: int, *, tenant_id: str = "tenant-id"): match method_name: case "_get_all_provider_models": - return SimpleNamespace( - id=f"model-{index}", + record = ProviderModel( + tenant_id=tenant_id, provider_name=provider_name, model_name=f"model-{index}", model_type=ModelType.LLM, credential_id=None, + is_valid=True, ) case "_get_all_provider_model_settings": - return SimpleNamespace( + record = ProviderModelSetting( + tenant_id=tenant_id, provider_name=provider_name, model_name=f"model-{index}", model_type=ModelType.LLM, @@ -762,69 +1179,95 @@ def test_provider_grouping_helpers_group_records_by_provider_name(method_name: s load_balancing_enabled=False, ) case "_get_all_provider_model_credentials": - return SimpleNamespace( - id=f"model-credential-{index}", + record = ProviderModelCredential( + tenant_id=tenant_id, provider_name=provider_name, model_name=f"model-{index}", model_type=ModelType.LLM, credential_name=f"credential-{index}", + encrypted_config='{"api_key": "secret"}', ) case "_get_all_provider_credentials": - return SimpleNamespace( - id=f"credential-{index}", + record = ProviderCredential( + tenant_id=tenant_id, provider_name=provider_name, credential_name=f"credential-{index}", + encrypted_config='{"api_key": "secret"}', ) case _: raise AssertionError(f"Unexpected method: {method_name}") + record.id = f"record-{index}" + return record - session = Mock() openai_primary = build_record("openai", 1) openai_secondary = build_record("openai", 2) anthropic_record = build_record("anthropic", 3) - session.scalars.return_value = [openai_primary, openai_secondary, anthropic_record] + other_tenant_record = build_record("other-provider", 4, tenant_id="other-tenant") + provider_db.add_all([openai_primary, openai_secondary, anthropic_record, other_tenant_record]) + provider_db.commit() - with ( - patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), - ): + with patch("core.provider_manager.redis_client", _FakeRedis()): result = getattr(ProviderManager, method_name)("tenant-id") assert [record.provider_name for record in result["openai"]] == ["openai", "openai"] assert [record.provider_name for record in result["anthropic"]] == ["anthropic"] + assert "other-provider" not in result -def test_get_all_preferred_model_providers_returns_mapping_by_provider_name() -> None: - session = Mock() - openai_preference = SimpleNamespace(provider_name="openai", preferred_provider_type=ProviderType.SYSTEM) - anthropic_preference = SimpleNamespace(provider_name="anthropic", preferred_provider_type=ProviderType.CUSTOM) - session.scalars.return_value = [openai_preference, anthropic_preference] +def test_get_all_preferred_model_providers_returns_mapping_by_provider_name(provider_db: Session) -> None: + openai_preference = TenantPreferredModelProvider( + tenant_id="tenant-id", provider_name="openai", preferred_provider_type=ProviderType.SYSTEM + ) + anthropic_preference = TenantPreferredModelProvider( + tenant_id="tenant-id", provider_name="anthropic", preferred_provider_type=ProviderType.CUSTOM + ) + provider_db.add_all( + [ + openai_preference, + anthropic_preference, + TenantPreferredModelProvider( + tenant_id="other-tenant", + provider_name="other-provider", + preferred_provider_type=ProviderType.SYSTEM, + ), + ] + ) + provider_db.commit() - with ( - patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), - ): + with patch("core.provider_manager.redis_client", _FakeRedis()): result = ProviderManager._get_all_preferred_model_providers("tenant-id") assert result["openai"].preferred_provider_type == ProviderType.SYSTEM assert result["anthropic"].preferred_provider_type == ProviderType.CUSTOM + assert "other-provider" not in result -def test_get_all_provider_load_balancing_configs_returns_empty_when_cached_flag_is_disabled() -> None: - with ( - patch("core.provider_manager.redis_client.get", return_value=b"False"), - patch("core.provider_manager.FeatureService.get_features") as mock_get_features, - patch("core.provider_manager.session_factory.create_session") as mock_create_session, - ): - result = ProviderManager._get_all_provider_load_balancing_configs("tenant-id") +@pytest.mark.usefixtures("provider_db") +def test_get_all_provider_load_balancing_configs_returns_empty_when_cached_flag_is_disabled( + sqlite_engine: Engine, +) -> None: + statements: list[str] = [] + + def record_statement(_conn, _cursor, statement, _parameters, _context, _executemany) -> None: + statements.append(statement) + + event.listen(sqlite_engine, "before_cursor_execute", record_statement) + try: + with ( + patch("core.provider_manager.redis_client.get", return_value=b"False"), + patch("core.provider_manager.FeatureService.get_features") as mock_get_features, + ): + result = ProviderManager._get_all_provider_load_balancing_configs("tenant-id") + finally: + event.remove(sqlite_engine, "before_cursor_execute", record_statement) assert result == {} mock_get_features.assert_not_called() - mock_create_session.assert_not_called() + assert statements == [] -def test_get_all_provider_load_balancing_configs_populates_cache_and_groups_configs() -> None: - session = Mock() - openai_config = SimpleNamespace( - id="lb-1", +def test_get_all_provider_load_balancing_configs_populates_cache_and_groups_configs(provider_db: Session) -> None: + openai_config = LoadBalancingModelConfig( tenant_id="tenant-id", provider_name="openai", model_name="gpt-4", @@ -835,8 +1278,8 @@ def test_get_all_provider_load_balancing_configs_populates_cache_and_groups_conf credential_source_type=None, enabled=True, ) - anthropic_config = SimpleNamespace( - id="lb-2", + openai_config.id = "lb-1" + anthropic_config = LoadBalancingModelConfig( tenant_id="tenant-id", provider_name="anthropic", model_name="claude", @@ -847,7 +1290,20 @@ def test_get_all_provider_load_balancing_configs_populates_cache_and_groups_conf credential_source_type=None, enabled=True, ) - session.scalars.return_value = [openai_config, anthropic_config] + anthropic_config.id = "lb-2" + other_tenant_config = LoadBalancingModelConfig( + tenant_id="other-tenant", + provider_name="other-provider", + model_name="other-model", + model_type=ModelType.LLM, + name="primary", + encrypted_config=None, + credential_id=None, + credential_source_type=None, + enabled=True, + ) + provider_db.add_all([openai_config, anthropic_config, other_tenant_config]) + provider_db.commit() with ( patch("core.provider_manager.redis_client.get", return_value=None), @@ -856,10 +1312,10 @@ def test_get_all_provider_load_balancing_configs_populates_cache_and_groups_conf "core.provider_manager.FeatureService.get_features", return_value=SimpleNamespace(model_load_balancing_enabled=True), ), - patch("core.provider_manager.session_factory.create_session", return_value=_build_session_context(session)), ): result = ProviderManager._get_all_provider_load_balancing_configs("tenant-id") mock_setex.assert_any_call("tenant:tenant-id:model_load_balancing_enabled", 120, "True") assert [record.provider_name for record in result["openai"]] == ["openai"] assert [record.provider_name for record in result["anthropic"]] == ["anthropic"] + assert "other-provider" not in result diff --git a/api/tests/unit_tests/core/tools/test_base_tool.py b/api/tests/unit_tests/core/tools/test_base_tool.py index f164e3fddea..8cbd538b66e 100644 --- a/api/tests/unit_tests/core/tools/test_base_tool.py +++ b/api/tests/unit_tests/core/tools/test_base_tool.py @@ -3,9 +3,9 @@ from __future__ import annotations from collections.abc import Generator from dataclasses import dataclass from typing import Any, cast -from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool import Tool @@ -94,7 +94,7 @@ def _build_tool(runtime: ToolRuntime | None = None) -> DummyTool: return DummyTool(entity=entity, runtime=runtime) -def test_invoke_supports_single_message_and_parameter_casting(): +def test_invoke_supports_single_message_and_parameter_casting(sqlite_session: Session): runtime = ToolRuntime( tenant_id="tenant-1", invoke_from=InvokeFrom.DEBUGGER, @@ -112,7 +112,7 @@ def test_invoke_supports_single_message_and_parameter_casting(): messages = list( tool.invoke( - session=MagicMock(), + session=sqlite_session, user_id="user-1", tool_parameters={"age": "18", "raw": "keep"}, conversation_id="conv-1", @@ -132,7 +132,7 @@ def test_invoke_supports_single_message_and_parameter_casting(): } -def test_invoke_preserves_multiple_select_values(): +def test_invoke_preserves_multiple_select_values(sqlite_session: Session): tool = _build_tool() parameter = ToolParameter.get_simple_instance( name="choice", @@ -144,18 +144,18 @@ def test_invoke_preserves_multiple_select_values(): parameter.multiple = True tool.entity.parameters = [parameter] - list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"choice": ["a", "b"]})) + list(tool.invoke(session=sqlite_session, 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"}) + tool.invoke(session=sqlite_session, user_id="user-1", tool_parameters={"choice": "a"}) -def test_invoke_supports_list_and_generator_results(): +def test_invoke_supports_list_and_generator_results(sqlite_session: Session): tool = _build_tool() tool.result = [tool.create_text_message("a"), tool.create_text_message("b")] - list_messages = list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={})) + list_messages = list(tool.invoke(session=sqlite_session, user_id="user-1", tool_parameters={})) assert [msg.message.text for msg in list_messages] == ["a", "b"] def _message_generator() -> Generator[ToolInvokeMessage, None, None]: @@ -163,7 +163,7 @@ def test_invoke_supports_list_and_generator_results(): yield tool.create_text_message("g2") tool.result = _message_generator() - generated_messages = list(tool.invoke(session=MagicMock(), user_id="user-2", tool_parameters={})) + generated_messages = list(tool.invoke(session=sqlite_session, user_id="user-2", tool_parameters={})) assert [msg.message.text for msg in generated_messages] == ["g1", "g2"] @@ -372,6 +372,6 @@ def test_message_factory_helpers(): assert variable_message.message.stream is False -def test_base_abstract_invoke_placeholder_returns_none(): +def test_base_abstract_invoke_placeholder_returns_none(sqlite_session: Session): tool = _build_tool() - assert Tool._invoke(tool, session=MagicMock(), user_id="u", tool_parameters={}) is None + assert Tool._invoke(tool, session=sqlite_session, user_id="u", tool_parameters={}) is None diff --git a/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py b/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py index 4dac9b7260d..d6bb7503793 100644 --- a/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py +++ b/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py @@ -4,10 +4,10 @@ import calendar import math from datetime import date from types import SimpleNamespace -from unittest.mock import MagicMock from zoneinfo import ZoneInfo import pytest +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool_runtime import ToolRuntime @@ -51,24 +51,26 @@ def _raise_runtime_error(*_args: object, **_kwargs: object) -> None: raise RuntimeError("boom") -def test_current_time_tool(): +def test_current_time_tool(sqlite_session: Session): current_tool = _build_builtin_tool(CurrentTimeTool) - utc_text = list(current_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"timezone": "UTC"}))[ + utc_text = list(current_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"timezone": "UTC"}))[ 0 ].message.text assert utc_text invalid_tz = list( - current_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"timezone": "Invalid/TZ"}) + current_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"timezone": "Invalid/TZ"}) )[0].message.text assert "Invalid timezone" in invalid_tz -def test_localtime_to_timestamp_tool(): +def test_localtime_to_timestamp_tool(sqlite_session: Session): localtime_tool = _build_builtin_tool(LocaltimeToTimestampTool) ts_message = list( localtime_tool.invoke( - session=MagicMock(), user_id="u", tool_parameters={"localtime": "2024-01-01 10:00:00", "timezone": "UTC"} + session=sqlite_session, + user_id="u", + tool_parameters={"localtime": "2024-01-01 10:00:00", "timezone": "UTC"}, ) )[0].message.text ts_value = float(ts_message.strip()) @@ -92,11 +94,11 @@ def test_localtime_to_timestamp_tool(): LocaltimeToTimestampTool.localtime_to_timestamp("bad", "%Y-%m-%d %H:%M:%S", "UTC") -def test_timestamp_to_localtime_tool(): +def test_timestamp_to_localtime_tool(sqlite_session: Session): to_local_tool = _build_builtin_tool(TimestampToLocaltimeTool) local_text = list( to_local_tool.invoke( - session=MagicMock(), user_id="u", tool_parameters={"timestamp": 1704067200, "timezone": "UTC"} + session=sqlite_session, user_id="u", tool_parameters={"timestamp": 1704067200, "timezone": "UTC"} ) )[0].message.text assert "2024" in local_text @@ -104,11 +106,11 @@ def test_timestamp_to_localtime_tool(): TimestampToLocaltimeTool.timestamp_to_localtime("bad", "UTC") # type: ignore[arg-type] -def test_timezone_conversion_tool(): +def test_timezone_conversion_tool(sqlite_session: Session): timezone_tool = _build_builtin_tool(TimezoneConversionTool) converted = list( timezone_tool.invoke( - session=MagicMock(), + session=sqlite_session, user_id="u", tool_parameters={ "current_time": "2024-01-01 08:00:00", @@ -122,10 +124,10 @@ def test_timezone_conversion_tool(): TimezoneConversionTool.timezone_convert("bad", "UTC", "Asia/Tokyo") -def test_weekday_tool(): +def test_weekday_tool(sqlite_session: Session): weekday_tool = _build_builtin_tool(WeekdayTool) valid = list( - weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "month": 1, "day": 1}) + weekday_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"year": 2024, "month": 1, "day": 1}) )[0].message.text expected_date = date(2024, 1, 1) expected_message = ( @@ -135,14 +137,14 @@ def test_weekday_tool(): ) assert valid == expected_message invalid = list( - weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "month": 2, "day": 31}) + weekday_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"year": 2024, "month": 2, "day": 31}) )[0].message.text assert "Invalid date" in invalid with pytest.raises(ValueError, match="Month is required"): - list(weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "day": 1})) + list(weekday_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"year": 2024, "day": 1})) -def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch): +def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): simple_code = _build_builtin_tool(SimpleCode) monkeypatch.setattr( @@ -151,7 +153,7 @@ def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch): ) result = list( simple_code.invoke( - session=MagicMock(), + session=sqlite_session, user_id="u", tool_parameters={"language": "python3", "code": "print(1)"}, ) @@ -159,18 +161,18 @@ def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch): assert result == "ok" -def test_simple_code_invalid_language(): +def test_simple_code_invalid_language(sqlite_session: Session): simple_code = _build_builtin_tool(SimpleCode) with pytest.raises(ValueError, match="Only python3 and javascript"): list( simple_code.invoke( - session=MagicMock(), user_id="u", tool_parameters={"language": "go", "code": "fmt.Println(1)"} + session=sqlite_session, user_id="u", tool_parameters={"language": "go", "code": "fmt.Println(1)"} ) ) -def test_simple_code_execution_error(monkeypatch: pytest.MonkeyPatch): +def test_simple_code_execution_error(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): simple_code = _build_builtin_tool(SimpleCode) monkeypatch.setattr( @@ -180,33 +182,35 @@ def test_simple_code_execution_error(monkeypatch: pytest.MonkeyPatch): with pytest.raises(ToolInvokeError, match="boom"): list( simple_code.invoke( - session=MagicMock(), user_id="u", tool_parameters={"language": "python3", "code": "print(1)"} + session=sqlite_session, + user_id="u", + tool_parameters={"language": "python3", "code": "print(1)"}, ) ) -def test_webscraper_empty_url(): +def test_webscraper_empty_url(sqlite_session: Session): webscraper = _build_builtin_tool(WebscraperTool) - empty = list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": ""}))[0].message.text + empty = list(webscraper.invoke(session=sqlite_session, user_id="u", tool_parameters={"url": ""}))[0].message.text assert empty == "Please input url" -def test_webscraper_fetch(monkeypatch: pytest.MonkeyPatch): +def test_webscraper_fetch(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): webscraper = _build_builtin_tool(WebscraperTool) monkeypatch.setattr("core.tools.builtin_tool.providers.webscraper.tools.webscraper.get_url", lambda *a, **k: "page") - full = list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": "https://example.com"}))[ + full = list(webscraper.invoke(session=sqlite_session, user_id="u", tool_parameters={"url": "https://example.com"}))[ 0 ].message.text assert full == "page" -def test_webscraper_summary(monkeypatch: pytest.MonkeyPatch): +def test_webscraper_summary(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): webscraper = _build_builtin_tool(WebscraperTool) monkeypatch.setattr("core.tools.builtin_tool.providers.webscraper.tools.webscraper.get_url", lambda *a, **k: "page") monkeypatch.setattr(webscraper, "summary", lambda user_id, content: "summary") summarized = list( webscraper.invoke( - session=MagicMock(), + session=sqlite_session, user_id="u", tool_parameters={"url": "https://example.com", "generate_summary": True}, ) @@ -214,26 +218,26 @@ def test_webscraper_summary(monkeypatch: pytest.MonkeyPatch): assert summarized == "summary" -def test_webscraper_fetch_error(monkeypatch: pytest.MonkeyPatch): +def test_webscraper_fetch_error(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): webscraper = _build_builtin_tool(WebscraperTool) monkeypatch.setattr( "core.tools.builtin_tool.providers.webscraper.tools.webscraper.get_url", _raise_runtime_error, ) with pytest.raises(ToolInvokeError, match="boom"): - list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": "https://example.com"})) + list(webscraper.invoke(session=sqlite_session, user_id="u", tool_parameters={"url": "https://example.com"})) -def test_asr_invalid_file(): +def test_asr_invalid_file(sqlite_session: Session): asr = _build_builtin_tool(ASRTool) file_obj = SimpleNamespace(type=FileType.DOCUMENT) - invalid_file = list(asr.invoke(session=MagicMock(), user_id="u", tool_parameters={"audio_file": file_obj}))[ + invalid_file = list(asr.invoke(session=sqlite_session, user_id="u", tool_parameters={"audio_file": file_obj}))[ 0 ].message.text assert "not a valid audio file" in invalid_file -def test_asr_valid_file_invocation(monkeypatch: pytest.MonkeyPatch): +def test_asr_valid_file_invocation(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): asr = _build_builtin_tool(ASRTool) model_instance = type("M", (), {"invoke_speech2text": lambda self, file: "transcript"})() model_manager = type("Mgr", (), {"get_model_instance": lambda *a, **k: model_instance})() @@ -245,9 +249,9 @@ def test_asr_valid_file_invocation(monkeypatch: pytest.MonkeyPatch): lambda **kwargs: captured_manager_kwargs.update(kwargs) or model_manager, ) audio_file = SimpleNamespace(type=FileType.AUDIO) - ok = list(asr.invoke(session=MagicMock(), user_id="u", tool_parameters={"audio_file": audio_file, "model": "p#m"}))[ - 0 - ].message.text + ok = list( + asr.invoke(session=sqlite_session, user_id="u", tool_parameters={"audio_file": audio_file, "model": "p#m"}) + )[0].message.text assert ok == "transcript" assert captured_manager_kwargs == {"tenant_id": "tenant-1", "user_id": "u"} @@ -263,7 +267,7 @@ def test_asr_available_models_and_runtime_parameters(monkeypatch: pytest.MonkeyP assert asr.get_runtime_parameters()[0].name == "model" -def test_tts_invoke_returns_messages(monkeypatch: pytest.MonkeyPatch): +def test_tts_invoke_returns_messages(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): tts = _build_builtin_tool(TTSTool) captured_manager_kwargs = {} voices_model_instance = type( @@ -281,7 +285,7 @@ def test_tts_invoke_returns_messages(monkeypatch: pytest.MonkeyPatch): or type("M", (), {"get_model_instance": lambda *a, **k: voices_model_instance})() ), ) - messages = list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + messages = list(tts.invoke(session=sqlite_session, user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) assert [m.type for m in messages] == [ToolInvokeMessage.MessageType.TEXT, ToolInvokeMessage.MessageType.BLOB] assert captured_manager_kwargs == {"tenant_id": "tenant-1", "user_id": "u"} @@ -293,18 +297,18 @@ def test_tts_get_available_models_requires_runtime(): tts.get_available_models() -def test_tts_tool_raises_when_runtime_missing(): +def test_tts_tool_raises_when_runtime_missing(sqlite_session: Session): tts = _build_builtin_tool(TTSTool) tts.runtime = None with pytest.raises(ValueError, match="Runtime is required"): - list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + list(tts.invoke(session=sqlite_session, user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) @pytest.mark.parametrize( "voices", [[{"value": None}], []], ) -def test_tts_tool_raises_when_voice_unavailable(monkeypatch, voices): +def test_tts_tool_raises_when_voice_unavailable(monkeypatch, voices, sqlite_session: Session): tts = _build_builtin_tool(TTSTool) tts.runtime = ToolRuntime(tenant_id="tenant-1", invoke_from=InvokeFrom.DEBUGGER) model_without_voice = type( @@ -320,7 +324,7 @@ def test_tts_tool_raises_when_voice_unavailable(monkeypatch, voices): lambda **_: type("Manager", (), {"get_model_instance": lambda *args, **kwargs: model_without_voice})(), ) with pytest.raises(ValueError, match="no voice available"): - list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + list(tts.invoke(session=sqlite_session, user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) def test_tts_tool_get_available_models_and_runtime_parameters(monkeypatch: pytest.MonkeyPatch): diff --git a/api/tests/unit_tests/core/tools/test_custom_tool.py b/api/tests/unit_tests/core/tools/test_custom_tool.py index 64cd128f18e..9dd68334b6c 100644 --- a/api/tests/unit_tests/core/tools/test_custom_tool.py +++ b/api/tests/unit_tests/core/tools/test_custom_tool.py @@ -2,10 +2,10 @@ from __future__ import annotations from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock import httpx import pytest +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool_runtime import ToolRuntime @@ -241,7 +241,7 @@ def test_do_http_request_builds_arguments_and_handles_invalid_method(monkeypatch invalid_method_tool.do_http_request("https://api.example.com", "TRACE", headers={}, parameters={}) -def test_do_http_request_handles_file_upload_and_invoke_paths(monkeypatch: pytest.MonkeyPatch): +def test_do_http_request_handles_file_upload_and_invoke_paths(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): openapi = { "parameters": [], "requestBody": { @@ -281,11 +281,11 @@ def test_do_http_request_handles_file_upload_and_invoke_paths(monkeypatch: pytes monkeypatch.setattr(tool, "assembling_request", lambda parameters: {}) monkeypatch.setattr(tool, "do_http_request", lambda *args, **kwargs: httpx.Response(200, text='{"a":1}')) monkeypatch.setattr(tool, "validate_and_parse_response", lambda _: ParsedResponse({"a": 1}, True)) - messages = list(tool.invoke(session=MagicMock(), user_id="u1", tool_parameters={})) + messages = list(tool.invoke(session=sqlite_session, user_id="u1", tool_parameters={})) assert [m.type for m in messages] == [ToolInvokeMessage.MessageType.JSON, ToolInvokeMessage.MessageType.TEXT] # _invoke text path monkeypatch.setattr(tool, "validate_and_parse_response", lambda _: ParsedResponse("plain", False)) - messages = list(tool.invoke(session=MagicMock(), user_id="u1", tool_parameters={})) + messages = list(tool.invoke(session=sqlite_session, user_id="u1", tool_parameters={})) assert len(messages) == 1 assert messages[0].message.text == "plain" diff --git a/api/tests/unit_tests/core/tools/test_custom_tool_provider.py b/api/tests/unit_tests/core/tools/test_custom_tool_provider.py index 93ae217e24e..760a6b49657 100644 --- a/api/tests/unit_tests/core/tools/test_custom_tool_provider.py +++ b/api/tests/unit_tests/core/tools/test_custom_tool_provider.py @@ -1,17 +1,31 @@ +"""Tests for custom API tool providers with persisted provider lookup state.""" + from __future__ import annotations +import json +from dataclasses import dataclass from types import SimpleNamespace -from unittest.mock import Mock, patch +from typing import cast +from uuid import uuid4 import pytest +from sqlalchemy import delete +from sqlalchemy.orm import Session +from core.tools.custom_tool import provider as provider_module from core.tools.custom_tool.provider import ApiToolProviderController from core.tools.custom_tool.tool import ApiTool from core.tools.entities.tool_bundle import ApiToolBundle -from core.tools.entities.tool_entities import ApiProviderAuthType, ToolProviderType +from core.tools.entities.tool_entities import ApiProviderAuthType, ApiProviderSchemaType, ToolProviderType +from models.tools import ApiToolProvider -def _db_provider() -> SimpleNamespace: +@dataclass(frozen=True) +class _Database: + session: Session + + +def _db_provider() -> ApiToolProvider: bundle = ApiToolBundle( server_url="https://api.example.com/items", method="GET", @@ -21,18 +35,39 @@ def _db_provider() -> SimpleNamespace: author="author", openapi={"parameters": []}, ) - return SimpleNamespace( - id="provider-id", - tenant_id="tenant-1", - name="provider-a", - description="desc", - icon="icon.svg", - user=SimpleNamespace(name="Alice"), - tools=[bundle], + return cast( + ApiToolProvider, + SimpleNamespace( + id="provider-id", + tenant_id="tenant-1", + name="provider-a", + description="desc", + icon="icon.svg", + user=SimpleNamespace(name="Alice"), + tools=[bundle], + ), ) -def test_api_tool_provider_from_db_and_parse_tool_bundle(): +def _persist_provider(session: Session, *, tenant_id: str, name: str = "provider-a") -> ApiToolProvider: + bundle = _db_provider().tools[0] + provider = ApiToolProvider( + name=name, + icon="icon.svg", + schema="{}", + schema_type_str=ApiProviderSchemaType.OPENAPI, + user_id=str(uuid4()), + tenant_id=tenant_id, + description="desc", + tools_str=json.dumps([bundle.model_dump(mode="json")]), + credentials_str='{"auth_type":"none"}', + ) + session.add(provider) + session.commit() + return provider + + +def test_api_tool_provider_from_db_and_parse_tool_bundle() -> None: controller = ApiToolProviderController.from_db(_db_provider(), ApiProviderAuthType.API_KEY_HEADER) assert controller.provider_type == ToolProviderType.API assert any(c.name == "api_key_value" for c in controller.entity.credentials_schema) @@ -42,7 +77,7 @@ def test_api_tool_provider_from_db_and_parse_tool_bundle(): assert tool.entity.identity.provider == "provider-id" -def test_api_tool_provider_from_db_query_auth_and_none_auth(): +def test_api_tool_provider_from_db_query_auth_and_none_auth() -> None: query_controller = ApiToolProviderController.from_db(_db_provider(), ApiProviderAuthType.API_KEY_QUERY) assert any(c.name == "api_key_query_param" for c in query_controller.entity.credentials_schema) @@ -50,7 +85,9 @@ def test_api_tool_provider_from_db_query_auth_and_none_auth(): assert [c.name for c in none_controller.entity.credentials_schema] == ["auth_type"] -def test_api_tool_provider_load_get_tools_and_get_tool(): +def test_api_tool_provider_load_get_tools_and_get_tool( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: controller = ApiToolProviderController.from_db(_db_provider(), ApiProviderAuthType.NONE) loaded = controller.load_bundled_tools(_db_provider().tools) assert len(loaded) == 1 @@ -66,10 +103,17 @@ def test_api_tool_provider_load_get_tools_and_get_tool(): # Force DB fetch branch. controller.tools = [] - provider_with_tools = _db_provider() - with patch("core.tools.custom_tool.provider.db") as mock_db: - scalars_result = Mock() - scalars_result.all.return_value = [provider_with_tools] - mock_db.session.scalars.return_value = scalars_result - tools = controller.get_tools("tenant-1") + tenant_id = str(uuid4()) + provider_with_tools = _persist_provider(sqlite_session, tenant_id=tenant_id) + _persist_provider(sqlite_session, tenant_id=str(uuid4())) + controller.tenant_id = tenant_id + monkeypatch.setattr(provider_module, "db", _Database(session=sqlite_session)) + + tools = controller.get_tools(tenant_id) assert len(tools) == 1 + assert tools[0].entity.identity.provider == controller.provider_id + + sqlite_session.execute(delete(ApiToolProvider).where(ApiToolProvider.id == provider_with_tools.id)) + sqlite_session.commit() + controller.tools = [] + assert controller.get_tools(tenant_id) == [] diff --git a/api/tests/unit_tests/core/tools/test_mcp_tool.py b/api/tests/unit_tests/core/tools/test_mcp_tool.py index be0ce20ad60..0c6ab408a37 100644 --- a/api/tests/unit_tests/core/tools/test_mcp_tool.py +++ b/api/tests/unit_tests/core/tools/test_mcp_tool.py @@ -22,6 +22,7 @@ from core.tools.entities.common_entities import I18nObject from core.tools.entities.tool_entities import ToolEntity, ToolIdentity, ToolInvokeMessage, ToolProviderType from core.tools.errors import ToolInvokeError from core.tools.mcp_tool.tool import MCPTool +from enums import DeploymentEdition def _build_mcp_tool(*, with_output_schema: bool = True) -> MCPTool: @@ -262,12 +263,12 @@ def test_invoke_remote_mcp_tool_fails_closed_when_user_id_missing(): tool = _build_forwarding_tool() with patch("core.tools.mcp_tool.tool.dify_config") as cfg: - cfg.ENTERPRISE_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE with pytest.raises(ToolInvokeError, match="no end-user context"): tool.invoke_remote_mcp_tool({}, user_id=None, app_id=None) -def test_invoke_skips_forwarding_when_enterprise_disabled(): +def test_invoke_skips_forwarding_outside_enterprise_edition(): """Non-enterprise deployments treat the DB selector as a no-op: a stale `identity_mode="idp_token"` row must NOT raise (fail-closed) AND must NOT call the enterprise inner API. The runtime falls through to the @@ -275,7 +276,7 @@ def test_invoke_skips_forwarding_when_enterprise_disabled(): tool = _build_forwarding_tool() with patch("core.tools.mcp_tool.tool.dify_config") as cfg: - cfg.ENTERPRISE_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY # The fail-closed branch must NOT fire (no enterprise → no forwarding). # The function will still try the legacy DB-load path; we patch that # to keep the test unit-scoped. diff --git a/api/tests/unit_tests/core/tools/test_signature.py b/api/tests/unit_tests/core/tools/test_signature.py index ee80cedd499..142b82902bf 100644 --- a/api/tests/unit_tests/core/tools/test_signature.py +++ b/api/tests/unit_tests/core/tools/test_signature.py @@ -7,14 +7,40 @@ from urllib.parse import parse_qs, urlparse import pytest from core.tools.signature import ( - get_signed_file_url_for_plugin, + bind_file_uri, + get_signed_file_uri_for_plugin, sign_tool_file, + sign_tool_file_uri, sign_upload_file_preview_url, verify_plugin_file_signature, verify_tool_file_signature, ) +def test_bind_file_uri_uses_selected_base_and_preserves_remote_url() -> None: + uri = "/files/tools/tool-file-id.png?sign=1" + + assert bind_file_uri(uri, "https://files.example.com") == f"https://files.example.com{uri}" + assert bind_file_uri(uri, "") == uri + assert bind_file_uri("https://remote.example.com/report.pdf", "https://files.example.com") == ( + "https://remote.example.com/report.pdf" + ) + + +def test_sign_tool_file_uri_has_no_origin(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("core.tools.signature.time.time", lambda: 1700000000) + monkeypatch.setattr("core.tools.signature.os.urandom", lambda _: b"\x08" * 16) + monkeypatch.setattr("core.tools.signature.dify_config.SECRET_KEY", "unit-secret") + + uri = sign_tool_file_uri("tool-file-id", ".png") + parsed = urlparse(uri) + + assert parsed.scheme == "" + assert parsed.netloc == "" + assert parsed.path == "/files/tools/tool-file-id.png" + assert parse_qs(parsed.query)["timestamp"] == ["1700000000"] + + def test_sign_tool_file_and_verify_roundtrip(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr("core.tools.signature.time.time", lambda: 1700000000) monkeypatch.setattr("core.tools.signature.os.urandom", lambda _: b"\x01" * 16) @@ -125,25 +151,23 @@ def test_sign_upload_file_preview_url_ignores_internal_files_url(monkeypatch: py assert query["sign"][0] -def test_get_signed_file_url_for_plugin_and_verify_roundtrip(monkeypatch: pytest.MonkeyPatch) -> None: +def test_get_signed_file_uri_for_plugin_and_verify_roundtrip(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr("core.tools.signature.time.time", lambda: 1700000000) monkeypatch.setattr("core.tools.signature.os.urandom", lambda _: b"\x06" * 16) monkeypatch.setattr("core.tools.signature.dify_config.SECRET_KEY", "unit-secret") - monkeypatch.setattr("core.tools.signature.dify_config.FILES_URL", "https://files.example.com") - monkeypatch.setattr("core.tools.signature.dify_config.INTERNAL_FILES_URL", "https://internal.example.com") monkeypatch.setattr("core.tools.signature.dify_config.FILES_ACCESS_TIMEOUT", 60) - url = get_signed_file_url_for_plugin( + uri = get_signed_file_uri_for_plugin( filename="report.pdf", mimetype="application/pdf", tenant_id="tenant-id", user_id="user-id", conversation_id="conversation-id", ) - parsed = urlparse(url) + parsed = urlparse(uri) query = parse_qs(parsed.query) - assert parsed.netloc == "internal.example.com" + assert parsed.netloc == "" assert parsed.path == "/files/upload/for-plugin" assert query["tenant_id"] == ["tenant-id"] assert query["user_id"] == ["user-id"] @@ -163,21 +187,49 @@ def test_get_signed_file_url_for_plugin_and_verify_roundtrip(monkeypatch: pytest ) +def test_plugin_upload_signature_binds_account_user_from(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("core.tools.signature.time.time", lambda: 1700000000) + monkeypatch.setattr("core.tools.signature.os.urandom", lambda _: b"\x09" * 16) + monkeypatch.setattr("core.tools.signature.dify_config.SECRET_KEY", "unit-secret") + monkeypatch.setattr("core.tools.signature.dify_config.FILES_ACCESS_TIMEOUT", 60) + + uri = get_signed_file_uri_for_plugin( + filename="report.pdf", + mimetype="application/pdf", + tenant_id="tenant-id", + user_id="account-id", + user_from="account", + ) + query = parse_qs(urlparse(uri).query) + + assert query["user_from"] == ["account"] + signed = { + "filename": "report.pdf", + "mimetype": "application/pdf", + "tenant_id": "tenant-id", + "user_id": "account-id", + "timestamp": query["timestamp"][0], + "nonce": query["nonce"][0], + "sign": query["sign"][0], + } + assert verify_plugin_file_signature(**signed, user_from="account") is True + assert verify_plugin_file_signature(**signed, user_from="end-user") is False + assert verify_plugin_file_signature(**signed) is False + + def test_verify_plugin_file_signature_rejects_invalid_signatures(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr("core.tools.signature.time.time", lambda: 1700000000) monkeypatch.setattr("core.tools.signature.os.urandom", lambda _: b"\x07" * 16) monkeypatch.setattr("core.tools.signature.dify_config.SECRET_KEY", "unit-secret") - monkeypatch.setattr("core.tools.signature.dify_config.FILES_URL", "https://files.example.com") - monkeypatch.setattr("core.tools.signature.dify_config.INTERNAL_FILES_URL", "") monkeypatch.setattr("core.tools.signature.dify_config.FILES_ACCESS_TIMEOUT", 30) - url = get_signed_file_url_for_plugin( + uri = get_signed_file_uri_for_plugin( filename="report.pdf", mimetype="application/pdf", tenant_id="tenant-id", user_id="user-id", ) - query = parse_qs(urlparse(url).query) + query = parse_qs(urlparse(uri).query) assert ( verify_plugin_file_signature( diff --git a/api/tests/unit_tests/core/tools/test_tool_engine.py b/api/tests/unit_tests/core/tools/test_tool_engine.py index f38ab2a2fab..f688c68cfa6 100644 --- a/api/tests/unit_tests/core/tools/test_tool_engine.py +++ b/api/tests/unit_tests/core/tools/test_tool_engine.py @@ -1,11 +1,14 @@ from __future__ import annotations from collections.abc import Generator -from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch +from uuid import uuid4 import pytest +from sqlalchemy import select +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool import Tool @@ -26,6 +29,45 @@ from core.tools.errors import ( ToolParameterValidationError, ) from core.tools.tool_engine import ToolEngine +from models.model import AppMode, Message, MessageFile + + +class _DatabaseBinding: + engine: Engine + + def __init__(self, engine: Engine) -> None: + self.engine = engine + + +def _message() -> Message: + message = Message( + app_id=str(uuid4()), + model_provider="provider", + model_id="model", + override_model_configs=None, + conversation_id=str(uuid4()), + inputs={}, + query="query", + message="", + message_tokens=0, + message_unit_price=0, + message_price_unit=0, + 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="debugger", + from_source="console", + from_end_user_id=None, + from_account_id=str(uuid4()), + app_mode=AppMode.CHAT, + ) + message.id = str(uuid4()) + return message class _DummyTool(Tool): @@ -120,52 +162,41 @@ def test_convert_tool_response_to_str_and_extract_binary_messages(): ) -def test_create_message_files_and_invoke_generator(): +@pytest.mark.parametrize("sqlite_session", [(MessageFile,)], indirect=True) +def test_create_message_files_and_invoke_generator(sqlite_engine: Engine, sqlite_session: Session): binaries = [ ToolInvokeMessageBinary(mimetype="image/png", url="https://example.com/abc.png"), ToolInvokeMessageBinary(mimetype="audio/wav", url="https://example.com/def.wav"), ] - created = [] - - def _message_file_factory(**kwargs): - obj = SimpleNamespace(id=f"mf-{len(created) + 1}", **kwargs) - created.append(obj) - return obj - - file_session = MagicMock() - session_factory = MagicMock() - session_factory.begin.return_value.__enter__.return_value = file_session - with ( - patch("core.tools.tool_engine.MessageFile", side_effect=_message_file_factory), - patch("core.tools.tool_engine.db") as mock_db, - patch("core.tools.tool_engine.sessionmaker", return_value=session_factory) as mock_sessionmaker, - ): + agent_message = _message() + with patch("core.tools.tool_engine.db", _DatabaseBinding(sqlite_engine)): ids = ToolEngine._create_message_files( tool_messages=binaries, - agent_message=SimpleNamespace(id="msg-1"), + agent_message=agent_message, invoke_from=InvokeFrom.DEBUGGER, - user_id="user-1", + user_id=str(uuid4()), ) - assert ids == ["mf-1", "mf-2"] - mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) - assert file_session.add.call_count == 2 - mock_db.session.close.assert_not_called() + message_files = list(sqlite_session.scalars(select(MessageFile).order_by(MessageFile.created_at)).all()) + assert ids == [message_file.id for message_file in message_files] + assert len(message_files) == 2 + assert {message_file.message_id for message_file in message_files} == {agent_message.id} tool = _build_tool() - invoked = list(ToolEngine._invoke(MagicMock(), tool, {"a": 1}, user_id="u")) + invoked = list(ToolEngine._invoke(sqlite_session, tool, {"a": 1}, user_id="u")) assert invoked[0].type == ToolInvokeMessage.MessageType.TEXT assert isinstance(invoked[-1], ToolInvokeMeta) assert invoked[-1].error is None -def test_generic_invoke_success_and_error_paths(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_generic_invoke_success_and_error_paths(sqlite_session: Session): tool = _build_tool() callback = Mock() callback.on_tool_execution.side_effect = lambda **kwargs: kwargs["tool_outputs"] response = list( ToolEngine.generic_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"x": 1}, user_id="u1", @@ -186,7 +217,7 @@ def test_generic_invoke_success_and_error_paths(): with pytest.raises(RuntimeError, match="boom"): list( ToolEngine.generic_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"x": 1}, user_id="u1", @@ -197,10 +228,11 @@ def test_generic_invoke_success_and_error_paths(): error_callback.on_tool_error.assert_called_once() -def test_agent_invoke_success(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_agent_invoke_success(sqlite_session: Session): tool = _build_tool(with_llm_parameter=True) callback = Mock() - message = SimpleNamespace(id="m1", conversation_id="c1") + message = _message() meta = ToolInvokeMeta.empty() with patch.object(ToolEngine, "_invoke", return_value=iter([tool.create_text_message("ok"), meta])): @@ -211,7 +243,7 @@ def test_agent_invoke_success(): with patch.object(ToolEngine, "_extract_tool_response_binary_and_text", return_value=iter([])): with patch.object(ToolEngine, "_create_message_files", return_value=[]): result_text, message_files, result_meta = ToolEngine.agent_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters="hello", user_id="u1", @@ -228,14 +260,15 @@ def test_agent_invoke_success(): callback.on_tool_end.assert_called_once() -def test_agent_invoke_param_validation_error(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_agent_invoke_param_validation_error(sqlite_session: Session): tool = _build_tool(with_llm_parameter=True) callback = Mock() - message = SimpleNamespace(id="m1", conversation_id="c1") + message = _message() with patch.object(ToolEngine, "_invoke", side_effect=ToolParameterValidationError("bad-param")): error_text, files, error_meta = ToolEngine.agent_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"a": 1}, user_id="u1", @@ -250,15 +283,16 @@ def test_agent_invoke_param_validation_error(): assert error_meta.error -def test_agent_invoke_engine_meta_error(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_agent_invoke_engine_meta_error(sqlite_session: Session): tool = _build_tool(with_llm_parameter=True) callback = Mock() - message = SimpleNamespace(id="m1", conversation_id="c1") + message = _message() engine_error = ToolEngineInvokeError(ToolInvokeMeta.error_instance("meta failure")) with patch.object(ToolEngine, "_invoke", side_effect=engine_error): error_text, files, error_meta = ToolEngine.agent_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"a": 1}, user_id="u1", @@ -295,14 +329,15 @@ def test_convert_tool_response_excludes_variable_messages(): assert "variable_name" not in result -def test_agent_invoke_tool_invoke_error(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_agent_invoke_tool_invoke_error(sqlite_session: Session): tool = _build_tool(with_llm_parameter=True) callback = Mock() - message = SimpleNamespace(id="m1", conversation_id="c1") + message = _message() with patch.object(ToolEngine, "_invoke", side_effect=ToolInvokeError("invoke boom")): error_text, files, _ = ToolEngine.agent_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"a": 1}, user_id="u1", diff --git a/api/tests/unit_tests/core/tools/test_tool_manager.py b/api/tests/unit_tests/core/tools/test_tool_manager.py index 89f50b207db..a2947dbd94a 100644 --- a/api/tests/unit_tests/core/tools/test_tool_manager.py +++ b/api/tests/unit_tests/core/tools/test_tool_manager.py @@ -1,14 +1,20 @@ -from __future__ import annotations +"""Unit tests for ToolManager with persisted providers and isolated external collaborators.""" -"""Unit tests for ToolManager behavior with mocked providers and collaborators.""" +from __future__ import annotations import json import threading +from collections.abc import Iterator +from dataclasses import dataclass +from datetime import datetime from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch import pytest +from sqlalchemy import event +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.plugin.entities.plugin_daemon import CredentialType @@ -22,6 +28,94 @@ from core.tools.entities.tool_entities import ( from core.tools.errors import ToolProviderNotFoundError from core.tools.plugin_tool.provider import PluginToolProviderController from core.tools.tool_manager import ToolManager +from models.base import TypeBase +from models.tools import ApiToolProvider, BuiltinToolProvider, WorkflowToolProvider + + +@dataclass(frozen=True) +class _ToolDatabase: + engine: Engine + session: Session + + +@pytest.fixture +def tool_database(sqlite_engine: Engine) -> Iterator[_ToolDatabase]: + """Provide isolated provider tables through the same engine/session split used in production.""" + + tables = [ + TypeBase.metadata.tables[model.__tablename__] + for model in (BuiltinToolProvider, ApiToolProvider, WorkflowToolProvider) + ] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + with Session(sqlite_engine, expire_on_commit=False) as session: + yield _ToolDatabase(engine=sqlite_engine, session=session) + + +def _builtin_provider( + *, + provider_id: str, + tenant_id: str, + provider: str = "time", + name: str = "Time credentials", + encrypted_credentials: str = '{"encrypted":"value"}', + is_default: bool = True, + credential_type: CredentialType = CredentialType.API_KEY, + expires_at: int = -1, +) -> BuiltinToolProvider: + record = BuiltinToolProvider( + tenant_id=tenant_id, + user_id="00000000-0000-0000-0000-000000000099", + provider=provider, + name=name, + encrypted_credentials=encrypted_credentials, + is_default=is_default, + credential_type=credential_type, + expires_at=expires_at, + ) + record.id = provider_id + return record + + +def _api_provider( + *, provider_id: str, tenant_id: str, name: str = "api-provider", icon: str = '{"background":"#000","content":"A"}' +) -> ApiToolProvider: + record = ApiToolProvider( + name=name, + icon=icon, + schema="{}", + schema_type_str="openapi", + user_id="00000000-0000-0000-0000-000000000099", + tenant_id=tenant_id, + description="desc", + tools_str="[]", + credentials_str='{"auth_type":"api_key_query","api_key_value":"secret"}', + privacy_policy="privacy", + custom_disclaimer="disclaimer", + ) + record.id = provider_id + return record + + +def _workflow_provider( + *, + provider_id: str, + tenant_id: str, + name: str = "workflow-provider", + icon: str = '{"background":"#222","content":"W"}', +) -> WorkflowToolProvider: + record = WorkflowToolProvider( + name=name, + label=name, + icon=icon, + app_id=provider_id, + version="1", + user_id="00000000-0000-0000-0000-000000000099", + tenant_id=tenant_id, + description="desc", + parameter_configuration="[]", + ) + record.id = provider_id + return record class _SimpleContextVar: @@ -39,26 +133,24 @@ class _SimpleContextVar: self._is_set = True -def _cm(session: Any): - context = Mock() - context.__enter__ = Mock(return_value=session) - context.__exit__ = Mock(return_value=False) - return context - - def _setup_list_providers_from_api_mocks( monkeypatch, *, - session: Mock, + tool_database: _ToolDatabase, hardcoded_controller: SimpleNamespace, plugin_controller: PluginToolProviderController, api_controller: SimpleNamespace, workflow_controller: SimpleNamespace, ): - mock_db = Mock() - mock_db.engine = object() - monkeypatch.setattr("core.tools.tool_manager.db", mock_db) - monkeypatch.setattr("core.tools.tool_manager.Session", lambda *args, **kwargs: _cm(session)) + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) + monkeypatch.setattr( + "core.tools.tool_manager.dify_config", + SimpleNamespace( + SQLALCHEMY_DATABASE_URI_SCHEME="mysql", + POSITION_TOOL_INCLUDES_SET=None, + POSITION_TOOL_EXCLUDES_SET=None, + ), + ) monkeypatch.setattr( ToolManager, "list_builtin_providers", @@ -92,7 +184,12 @@ def _setup_list_providers_from_api_mocks( ) monkeypatch.setattr( "core.tools.tool_manager.ToolLabelManager.get_tools_labels", - Mock(side_effect=[{"api-1": ["search"]}, {"wf-1": ["utility"]}]), + Mock( + side_effect=[ + {api_controller.provider_id: ["search"]}, + {workflow_controller.provider_id: ["utility"]}, + ] + ), ) mock_mcp_service = Mock() mock_mcp_service.list_providers.return_value = [SimpleNamespace(name="mcp-provider")] @@ -201,7 +298,9 @@ def test_get_tool_runtime_builtin_missing_tool_raises(): ) -def test_get_tool_runtime_builtin_with_credentials_decrypts_and_forks(): +def test_get_tool_runtime_builtin_with_credentials_decrypts_and_forks( + monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase +): tool = Mock() tool.fork_tool_runtime.return_value = "runtime-tool" controller = SimpleNamespace( @@ -209,28 +308,27 @@ def test_get_tool_runtime_builtin_with_credentials_decrypts_and_forks(): need_credentials=True, get_credentials_schema_by_type=Mock(return_value=[]), ) - builtin_provider = SimpleNamespace( - id="cred-1", - credential_type=CredentialType.API_KEY.value, - credentials={"encrypted": "value"}, - expires_at=-1, - user_id="user-1", + tenant_id = "00000000-0000-0000-0000-000000000001" + builtin_provider = _builtin_provider( + provider_id="00000000-0000-0000-0000-000000000002", + tenant_id=tenant_id, ) + tool_database.session.add(builtin_provider) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) with patch.object(ToolManager, "get_builtin_provider", return_value=controller): with patch("core.helper.credential_utils.check_credential_policy_compliance"): - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.session.scalar.return_value = builtin_provider - encrypter = Mock() - encrypter.decrypt.return_value = {"api_key": "secret"} - cache = Mock() - with patch("core.tools.tool_manager.create_provider_encrypter", return_value=(encrypter, cache)): - result = ToolManager.get_tool_runtime( - provider_type=ToolProviderType.BUILT_IN, - provider_id="time", - tool_name="weekday", - tenant_id="tenant-1", - ) + encrypter = Mock() + encrypter.decrypt.return_value = {"api_key": "secret"} + cache = Mock() + with patch("core.tools.tool_manager.create_provider_encrypter", return_value=(encrypter, cache)): + result = ToolManager.get_tool_runtime( + provider_type=ToolProviderType.BUILT_IN, + provider_id="time", + tool_name="weekday", + tenant_id=tenant_id, + ) assert result == "runtime-tool" runtime = tool.fork_tool_runtime.call_args.kwargs["runtime"] @@ -244,16 +342,16 @@ def test_get_tool_runtime_builtin_with_credentials_decrypts_and_forks(): "services.tools.builtin_tools_manage_service.BuiltinToolManageService.get_oauth_client", return_value={"client_id": "id"}, ) -@patch("core.tools.tool_manager.db") @patch("core.tools.tool_manager.time.time", return_value=1000) @patch("core.helper.credential_utils.check_credential_policy_compliance") def test_get_tool_runtime_builtin_refreshes_expired_oauth_credentials( mock_check, mock_time, - mock_db, mock_get_oauth_client, mock_oauth_handler_cls, mock_create_provider_encrypter, + monkeypatch: pytest.MonkeyPatch, + tool_database: _ToolDatabase, ): tool = Mock() tool.fork_tool_runtime.return_value = "runtime-tool" @@ -262,17 +360,19 @@ def test_get_tool_runtime_builtin_refreshes_expired_oauth_credentials( need_credentials=True, get_credentials_schema_by_type=Mock(return_value=[]), ) - builtin_provider = SimpleNamespace( - id="cred-1", - credential_type=CredentialType.OAUTH2.value, - credentials={"encrypted": "value"}, - encrypted_credentials=None, + tenant_id = "00000000-0000-0000-0000-000000000001" + provider_id = "00000000-0000-0000-0000-000000000002" + builtin_provider = _builtin_provider( + provider_id=provider_id, + tenant_id=tenant_id, + credential_type=CredentialType.OAUTH2, expires_at=1, - user_id="user-1", ) refreshed = SimpleNamespace(credentials={"token": "new"}, expires_at=123456) - mock_db.session.scalar.return_value = builtin_provider + tool_database.session.add(builtin_provider) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) encrypter = Mock() encrypter.decrypt.return_value = {"token": "old"} encrypter.encrypt.return_value = {"token": "encrypted"} @@ -285,34 +385,38 @@ def test_get_tool_runtime_builtin_refreshes_expired_oauth_credentials( provider_type=ToolProviderType.BUILT_IN, provider_id="time", tool_name="weekday", - tenant_id="tenant-1", + tenant_id=tenant_id, ) assert result == "runtime-tool" assert builtin_provider.expires_at == refreshed.expires_at assert builtin_provider.encrypted_credentials == json.dumps({"token": "encrypted"}) - mock_db.session.commit.assert_called_once() + tool_database.session.expire_all() + persisted = tool_database.session.get(BuiltinToolProvider, provider_id) + assert persisted is not None + assert persisted.expires_at == refreshed.expires_at + assert persisted.encrypted_credentials == json.dumps({"token": "encrypted"}) cache.delete.assert_called_once() -def test_get_tool_runtime_builtin_plugin_provider_deleted_raises(): +def test_get_tool_runtime_builtin_plugin_provider_deleted_raises( + monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase +): plugin_controller = object.__new__(PluginToolProviderController) plugin_controller.entity = SimpleNamespace(credentials_schema=[{"name": "k"}], oauth_schema=None) plugin_controller.get_tool = Mock(return_value=Mock()) plugin_controller.get_credentials_schema_by_type = Mock(return_value=[]) + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) with patch.object(ToolManager, "get_builtin_provider", return_value=plugin_controller): - with patch("core.tools.tool_manager.is_valid_uuid", return_value=True): - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.session.scalar.return_value = None - with pytest.raises(ToolProviderNotFoundError, match="provider has been deleted"): - ToolManager.get_tool_runtime( - provider_type=ToolProviderType.BUILT_IN, - provider_id="time", - tool_name="weekday", - tenant_id="tenant-1", - credential_id="uuid-id", - ) + with pytest.raises(ToolProviderNotFoundError, match="provider has been deleted"): + ToolManager.get_tool_runtime( + provider_type=ToolProviderType.BUILT_IN, + provider_id="time", + tool_name="weekday", + tenant_id="00000000-0000-0000-0000-000000000001", + credential_id="00000000-0000-0000-0000-000000000002", + ) def test_get_tool_runtime_api_path(): @@ -336,32 +440,30 @@ def test_get_tool_runtime_api_path(): ) -def test_get_tool_runtime_workflow_path(): - workflow_provider = SimpleNamespace(tenant_id="tenant-1") +def test_get_tool_runtime_workflow_path(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase): + tenant_id = "00000000-0000-0000-0000-000000000001" + provider_id = "00000000-0000-0000-0000-000000000002" + workflow_provider = _workflow_provider(provider_id=provider_id, tenant_id=tenant_id) + tool_database.session.add(workflow_provider) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) workflow_tool = Mock() workflow_tool.fork_tool_runtime.return_value = "wf-runtime" workflow_controller = Mock() workflow_controller.get_tools.return_value = [workflow_tool] - session = Mock() - session.begin.return_value = _cm(None) - session.scalar.return_value = workflow_provider - - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.engine = object() - with patch("core.tools.tool_manager.Session", return_value=_cm(session)): - with patch( - "core.tools.tool_manager.ToolTransformService.workflow_provider_to_controller", - return_value=workflow_controller, - ): - assert ( - ToolManager.get_tool_runtime( - provider_type=ToolProviderType.WORKFLOW, - provider_id="wf-1", - tool_name="wf", - tenant_id="tenant-1", - ) - == "wf-runtime" - ) + with patch( + "core.tools.tool_manager.ToolTransformService.workflow_provider_to_controller", + return_value=workflow_controller, + ): + assert ( + ToolManager.get_tool_runtime( + provider_type=ToolProviderType.WORKFLOW, + provider_id=provider_id, + tool_name="wf", + tenant_id=tenant_id, + ) + == "wf-runtime" + ) def test_get_tool_runtime_plugin_path(): @@ -631,79 +733,152 @@ def test_get_tool_label_loads_cache_and_handles_missing(): assert ToolManager.get_tool_label("missing") is None -def test_list_default_builtin_providers_for_postgres_and_mysql(): - provider_records = [SimpleNamespace(id="id-1"), SimpleNamespace(id="id-2")] +@pytest.mark.parametrize("database_scheme", ["mysql", "postgresql"]) +def test_list_default_builtin_providers_uses_persisted_defaults( + monkeypatch: pytest.MonkeyPatch, + tool_database: _ToolDatabase, + database_scheme: str, +): + tenant_id = "00000000-0000-0000-0000-000000000001" + default_provider = _builtin_provider( + provider_id="00000000-0000-0000-0000-000000000002", + tenant_id=tenant_id, + name="default", + is_default=True, + ) + older_provider = _builtin_provider( + provider_id="00000000-0000-0000-0000-000000000003", + tenant_id=tenant_id, + name="older", + is_default=False, + ) + other_tenant_provider = _builtin_provider( + provider_id="00000000-0000-0000-0000-000000000004", + tenant_id="00000000-0000-0000-0000-000000000005", + name="foreign", + ) + default_provider.created_at = datetime(2026, 1, 2) + older_provider.created_at = datetime(2026, 1, 1) + tool_database.session.add_all([default_provider, older_provider, other_tenant_provider]) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) + monkeypatch.setattr( + "core.tools.tool_manager.dify_config", + SimpleNamespace(SQLALCHEMY_DATABASE_URI_SCHEME=database_scheme), + ) - for scheme in ("postgresql", "mysql"): - session = Mock() - session.execute.return_value.all.return_value = [SimpleNamespace(id="id-1"), SimpleNamespace(id="id-2")] - session.scalars.return_value = iter(provider_records) + postgresql_statements = [] - with patch("core.tools.tool_manager.dify_config", SimpleNamespace(SQLALCHEMY_DATABASE_URI_SCHEME=scheme)): - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.engine = object() - with patch("core.tools.tool_manager.Session", return_value=_cm(session)): - providers = ToolManager.list_default_builtin_providers("tenant-1") + def translate_postgresql_distinct_on(_connection, _cursor, statement, parameters, _context, _executemany): + if "SELECT DISTINCT ON (tenant_id, provider) id" not in statement: + return statement, parameters - assert providers == provider_records + postgresql_statements.append(statement) + sqlite_statement = """ + SELECT id FROM ( + SELECT id, + ROW_NUMBER() OVER ( + PARTITION BY tenant_id, provider + ORDER BY is_default DESC, created_at DESC + ) AS rn + FROM tool_builtin_providers + WHERE tenant_id = ? + ) ranked WHERE rn = 1 + """ + return sqlite_statement, parameters + + if database_scheme == "postgresql": + event.listen( + tool_database.engine, + "before_cursor_execute", + translate_postgresql_distinct_on, + retval=True, + ) + + try: + providers = ToolManager.list_default_builtin_providers(tenant_id) + finally: + if database_scheme == "postgresql": + event.remove( + tool_database.engine, + "before_cursor_execute", + translate_postgresql_distinct_on, + ) + + assert [provider.id for provider in providers] == [default_provider.id] + if database_scheme == "postgresql": + assert postgresql_statements -def test_list_providers_from_api_covers_builtin_api_workflow_and_mcp(monkeypatch: pytest.MonkeyPatch): +def test_list_providers_from_api_covers_builtin_api_workflow_and_mcp( + monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase +): + tenant_id = "00000000-0000-0000-0000-000000000001" hardcoded_controller = SimpleNamespace(entity=SimpleNamespace(identity=SimpleNamespace(name="hardcoded"))) plugin_controller = object.__new__(PluginToolProviderController) plugin_controller.entity = SimpleNamespace(identity=SimpleNamespace(name="plugin-provider")) - api_db_provider_good = SimpleNamespace(id="api-1") - api_db_provider_bad = SimpleNamespace(id="api-2") - api_controller = SimpleNamespace(provider_id="api-1") + api_db_provider_good = _api_provider( + provider_id="00000000-0000-0000-0000-000000000002", tenant_id=tenant_id, name="api-good" + ) + api_db_provider_bad = _api_provider( + provider_id="00000000-0000-0000-0000-000000000003", tenant_id=tenant_id, name="api-bad" + ) + api_controller = SimpleNamespace(provider_id=api_db_provider_good.id) - workflow_db_provider_good = SimpleNamespace(id="wf-1") - workflow_db_provider_bad = SimpleNamespace(id="wf-2") - workflow_controller = SimpleNamespace(provider_id="wf-1") - - session = Mock() - session.scalars.side_effect = [ - SimpleNamespace(all=lambda: [api_db_provider_good, api_db_provider_bad]), - SimpleNamespace(all=lambda: [workflow_db_provider_good, workflow_db_provider_bad]), - ] + workflow_db_provider_good = _workflow_provider( + provider_id="00000000-0000-0000-0000-000000000004", tenant_id=tenant_id, name="workflow-good" + ) + workflow_db_provider_bad = _workflow_provider( + provider_id="00000000-0000-0000-0000-000000000005", tenant_id=tenant_id, name="workflow-bad" + ) + workflow_controller = SimpleNamespace(provider_id=workflow_db_provider_good.id) + tool_database.session.add_all( + [ + api_db_provider_good, + api_db_provider_bad, + _api_provider( + provider_id="00000000-0000-0000-0000-000000000006", + tenant_id="00000000-0000-0000-0000-000000000099", + name="foreign-api", + ), + workflow_db_provider_good, + workflow_db_provider_bad, + _workflow_provider( + provider_id="00000000-0000-0000-0000-000000000007", + tenant_id="00000000-0000-0000-0000-000000000099", + name="foreign-workflow", + ), + ] + ) + tool_database.session.commit() _setup_list_providers_from_api_mocks( monkeypatch, - session=session, + tool_database=tool_database, hardcoded_controller=hardcoded_controller, plugin_controller=plugin_controller, api_controller=api_controller, workflow_controller=workflow_controller, ) - providers = ToolManager.list_providers_from_api(user_id="user-1", tenant_id="tenant-1", typ="") + providers = ToolManager.list_providers_from_api(user_id="user-1", tenant_id=tenant_id, typ="") names = {provider.name for provider in providers} assert {"hardcoded", "plugin-provider", "api-provider", "workflow-provider", "mcp-provider"} <= names -def test_get_api_provider_controller_returns_controller_and_credentials(): - provider = SimpleNamespace( - id="api-1", - tenant_id="tenant-1", - name="api-provider", - description="desc", - credentials={"auth_type": "api_key_query"}, - credentials_str='{"auth_type": "api_key_query", "api_key_value": "secret"}', - schema_type="openapi", - schema="schema", - tools=[], - icon='{"background": "#000", "content": "A"}', - privacy_policy="privacy", - custom_disclaimer="disclaimer", - ) +def test_get_api_provider_controller_returns_controller_and_credentials( + monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase +): + tenant_id = "00000000-0000-0000-0000-000000000001" + provider = _api_provider(provider_id="00000000-0000-0000-0000-000000000002", tenant_id=tenant_id) + tool_database.session.add(provider) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) controller = Mock() - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.session.scalar.return_value = provider - with patch( - "core.tools.tool_manager.ApiToolProviderController.from_db", return_value=controller - ) as mock_from_db: - built_controller, credentials = ToolManager.get_api_provider_controller("tenant-1", "api-1") + with patch("core.tools.tool_manager.ApiToolProviderController.from_db", return_value=controller) as mock_from_db: + built_controller, credentials = ToolManager.get_api_provider_controller(tenant_id, provider.id) assert built_controller is controller assert credentials == provider.credentials @@ -711,83 +886,74 @@ def test_get_api_provider_controller_returns_controller_and_credentials(): controller.load_bundled_tools.assert_called_once_with(provider.tools) -def test_user_get_api_provider_masks_credentials_and_adds_labels(): - provider = SimpleNamespace( - id="api-1", - tenant_id="tenant-1", - name="api-provider", - description="desc", - credentials={"auth_type": "api_key_query"}, - credentials_str='{"auth_type": "api_key_query", "api_key_value": "secret"}', - schema_type="openapi", - schema="schema", - tools=[], - icon='{"background": "#000", "content": "A"}', - privacy_policy="privacy", - custom_disclaimer="disclaimer", - ) +def test_user_get_api_provider_masks_credentials_and_adds_labels( + monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase +): + tenant_id = "00000000-0000-0000-0000-000000000001" + provider = _api_provider(provider_id="00000000-0000-0000-0000-000000000002", tenant_id=tenant_id) + tool_database.session.add(provider) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) controller = Mock() - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.session.scalar.return_value = provider - with patch("core.tools.tool_manager.ApiToolProviderController.from_db", return_value=controller): - encrypter = Mock() - encrypter.decrypt.return_value = {"api_key_value": "secret"} - encrypter.mask_plugin_credentials.return_value = {"api_key_value": "***"} - with patch("core.tools.tool_manager.create_tool_provider_encrypter", return_value=(encrypter, Mock())): - with patch("core.tools.tool_manager.ToolLabelManager.get_tool_labels", return_value=["search"]): - user_payload = ToolManager.user_get_api_provider("api-provider", "tenant-1") + with patch("core.tools.tool_manager.ApiToolProviderController.from_db", return_value=controller): + encrypter = Mock() + encrypter.decrypt.return_value = {"api_key_value": "secret"} + encrypter.mask_plugin_credentials.return_value = {"api_key_value": "***"} + with patch("core.tools.tool_manager.create_tool_provider_encrypter", return_value=(encrypter, Mock())): + with patch("core.tools.tool_manager.ToolLabelManager.get_tool_labels", return_value=["search"]): + user_payload = ToolManager.user_get_api_provider(provider.name, tenant_id) assert user_payload["credentials"]["api_key_value"] == "***" assert user_payload["labels"] == ["search"] -def test_get_api_provider_controller_not_found_raises(): - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.session.scalar.return_value = None - with pytest.raises(ToolProviderNotFoundError, match="api provider missing not found"): - ToolManager.get_api_provider_controller("tenant-1", "missing") +def test_get_api_provider_controller_not_found_raises(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase): + provider_id = "00000000-0000-0000-0000-000000000002" + tool_database.session.add( + _api_provider( + provider_id=provider_id, + tenant_id="00000000-0000-0000-0000-000000000099", + ) + ) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) + + with pytest.raises(ToolProviderNotFoundError, match=f"api provider {provider_id} not found"): + ToolManager.get_api_provider_controller("00000000-0000-0000-0000-000000000001", provider_id) -def test_get_mcp_provider_controller_returns_controller(): +def test_get_mcp_provider_controller_returns_controller(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase): provider_entity = SimpleNamespace(provider_icon={"background": "#111", "content": "M"}) controller = Mock() - session = Mock() - - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.engine = object() - with patch("core.tools.tool_manager.Session", return_value=_cm(session)): - with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls: - mock_service = mock_service_cls.return_value - mock_service.get_provider.return_value = provider_entity - with patch("core.tools.tool_manager.MCPToolProviderController.from_db", return_value=controller): - built = ToolManager.get_mcp_provider_controller("tenant-1", "mcp-1") - assert built is controller + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) + with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls: + mock_service = mock_service_cls.return_value + mock_service.get_provider.return_value = provider_entity + with patch("core.tools.tool_manager.MCPToolProviderController.from_db", return_value=controller): + built = ToolManager.get_mcp_provider_controller("tenant-1", "mcp-1") + assert built is controller + assert isinstance(mock_service_cls.call_args.kwargs["session"], Session) -def test_generate_mcp_tool_icon_url_returns_provider_icon(): +def test_generate_mcp_tool_icon_url_returns_provider_icon( + monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase +): provider_entity = SimpleNamespace(provider_icon={"background": "#111", "content": "M"}) - session = Mock() - - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.engine = object() - with patch("core.tools.tool_manager.Session", return_value=_cm(session)): - with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls: - mock_service = mock_service_cls.return_value - mock_service.get_provider_entity.return_value = provider_entity - assert ToolManager.generate_mcp_tool_icon_url("tenant-1", "mcp-1") == provider_entity.provider_icon + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) + with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls: + mock_service = mock_service_cls.return_value + mock_service.get_provider_entity.return_value = provider_entity + assert ToolManager.generate_mcp_tool_icon_url("tenant-1", "mcp-1") == provider_entity.provider_icon + assert isinstance(mock_service_cls.call_args.kwargs["session"], Session) -def test_get_mcp_provider_controller_missing_raises(): - session = Mock() - - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.engine = object() - with patch("core.tools.tool_manager.Session", return_value=_cm(session)): - with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls: - mock_service_cls.return_value.get_provider.side_effect = ValueError("missing") - with pytest.raises(ToolProviderNotFoundError, match="mcp provider mcp-1 not found"): - ToolManager.get_mcp_provider_controller("tenant-1", "mcp-1") +def test_get_mcp_provider_controller_missing_raises(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase): + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) + with patch("core.tools.tool_manager.MCPToolManageService") as mock_service_cls: + mock_service_cls.return_value.get_provider.side_effect = ValueError("missing") + with pytest.raises(ToolProviderNotFoundError, match="mcp provider mcp-1 not found"): + ToolManager.get_mcp_provider_controller("tenant-1", "mcp-1") def test_generate_tool_icon_urls_for_builtin_and_plugin(): @@ -799,39 +965,49 @@ def test_generate_tool_icon_urls_for_builtin_and_plugin(): assert "/plugin/icon" in plugin_url -def test_generate_tool_icon_urls_for_workflow_and_api(): - workflow_provider = SimpleNamespace(icon='{"background": "#222", "content": "W"}') - api_provider = SimpleNamespace(icon='{"background": "#333", "content": "A"}') - mock_engine = object() - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.engine = mock_engine - with patch("core.tools.tool_manager.Session") as mock_session_cls: - mock_session = MagicMock() - mock_session.scalar.side_effect = [workflow_provider, api_provider] - mock_session_cls.return_value.__enter__ = MagicMock(return_value=mock_session) - mock_session_cls.return_value.__exit__ = MagicMock(return_value=False) - assert ToolManager.generate_workflow_tool_icon_url("tenant-1", "wf-1") == { - "background": "#222", - "content": "W", - } - assert ToolManager.generate_api_tool_icon_url("tenant-1", "api-1") == {"background": "#333", "content": "A"} - # Verify sessions are created with the engine - assert mock_session_cls.call_count == 2 - mock_session_cls.assert_called_with(mock_engine, expire_on_commit=False) +def test_generate_tool_icon_urls_for_workflow_and_api(monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase): + tenant_id = "00000000-0000-0000-0000-000000000001" + workflow_provider = _workflow_provider( + provider_id="00000000-0000-0000-0000-000000000002", + tenant_id=tenant_id, + ) + api_provider = _api_provider( + provider_id="00000000-0000-0000-0000-000000000003", + tenant_id=tenant_id, + icon='{"background":"#333","content":"A"}', + ) + tool_database.session.add_all([workflow_provider, api_provider]) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) + + assert ToolManager.generate_workflow_tool_icon_url(tenant_id, workflow_provider.id) == { + "background": "#222", + "content": "W", + } + assert ToolManager.generate_api_tool_icon_url(tenant_id, api_provider.id) == { + "background": "#333", + "content": "A", + } -def test_generate_tool_icon_urls_missing_workflow_and_api_use_default(): - mock_engine = object() - with patch("core.tools.tool_manager.db") as mock_db: - mock_db.engine = mock_engine - with patch("core.tools.tool_manager.Session") as mock_session_cls: - mock_session = MagicMock() - mock_session.scalar.return_value = None - mock_session_cls.return_value.__enter__ = MagicMock(return_value=mock_session) - mock_session_cls.return_value.__exit__ = MagicMock(return_value=False) - assert ToolManager.generate_workflow_tool_icon_url("tenant-1", "missing")["background"] == "#252525" - assert ToolManager.generate_api_tool_icon_url("tenant-1", "missing")["background"] == "#252525" - assert mock_session_cls.call_count == 2 +def test_generate_tool_icon_urls_missing_workflow_and_api_use_default( + monkeypatch: pytest.MonkeyPatch, tool_database: _ToolDatabase +): + tenant_id = "00000000-0000-0000-0000-000000000001" + foreign_tenant_id = "00000000-0000-0000-0000-000000000099" + workflow_provider_id = "00000000-0000-0000-0000-000000000002" + api_provider_id = "00000000-0000-0000-0000-000000000003" + tool_database.session.add_all( + [ + _workflow_provider(provider_id=workflow_provider_id, tenant_id=foreign_tenant_id), + _api_provider(provider_id=api_provider_id, tenant_id=foreign_tenant_id), + ] + ) + tool_database.session.commit() + monkeypatch.setattr("core.tools.tool_manager.db", tool_database) + + assert ToolManager.generate_workflow_tool_icon_url(tenant_id, workflow_provider_id)["background"] == "#252525" + assert ToolManager.generate_api_tool_icon_url(tenant_id, api_provider_id)["background"] == "#252525" def test_get_tool_icon_for_builtin_provider_variants(): @@ -907,14 +1083,14 @@ def test_convert_tool_parameters_type_agent_and_workflow_branches(): variable_pool = Mock() variable_pool.get.return_value = SimpleNamespace(value="from-variable") - variable_pool.convert_template.return_value = SimpleNamespace(text="from-template") - mixed = ToolManager._convert_tool_parameters_type( - parameters=[text_param], - variable_pool=variable_pool, - tool_configurations={"text": {"type": "mixed", "value": "Hello {{name}}"}}, - typ="workflow", - ) + with patch("core.tools.tool_manager.convert_template", return_value=SimpleNamespace(text="from-template")): + mixed = ToolManager._convert_tool_parameters_type( + parameters=[text_param], + variable_pool=variable_pool, + tool_configurations={"text": {"type": "mixed", "value": "Hello {{name}}"}}, + typ="workflow", + ) assert mixed == {"text": "from-template"} variable = ToolManager._convert_tool_parameters_type( diff --git a/api/tests/unit_tests/core/tools/utils/test_misc_utils_extra.py b/api/tests/unit_tests/core/tools/utils/test_misc_utils_extra.py index 007ab09aabc..afd2a1b8272 100644 --- a/api/tests/unit_tests/core/tools/utils/test_misc_utils_extra.py +++ b/api/tests/unit_tests/core/tools/utils/test_misc_utils_extra.py @@ -2,15 +2,20 @@ from __future__ import annotations import uuid from contextlib import nullcontext -from types import SimpleNamespace -from unittest.mock import Mock, patch +from datetime import datetime +from unittest.mock import patch import pytest +from sqlalchemy.orm import Session, scoped_session, sessionmaker from yaml import YAMLError from core.app.app_config.entities import DatasetRetrieveConfigEntity from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler +from core.model_manager import ModelInstance, ModelManager +from core.provider_manager import ProviderManager +from core.rag.embedding.retrieval import RetrievalSegments from core.rag.models.document import Document as RagDocument +from core.rag.rerank.rerank_model import RerankModelRunner from core.tools.utils.dataset_retriever import dataset_multi_retriever_tool as multi_retriever_module from core.tools.utils.dataset_retriever import dataset_retriever_tool as single_retriever_module from core.tools.utils.dataset_retriever.dataset_multi_retriever_tool import DatasetMultiRetrieverTool @@ -18,17 +23,88 @@ from core.tools.utils.dataset_retriever.dataset_retriever_tool import DatasetRet from core.tools.utils.text_processing_utils import remove_leading_symbols from core.tools.utils.uuid_utils import is_valid_uuid from core.tools.utils.yaml_utils import _load_yaml_file, load_yaml_file_cached +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, SegmentStatus def _retrieve_config() -> DatasetRetrieveConfigEntity: return DatasetRetrieveConfigEntity(retrieve_strategy=DatasetRetrieveConfigEntity.RetrieveStrategy.SINGLE) +def _persist_dataset( + session: Session, + *, + dataset_id: str | None = None, + tenant_id: str | None = None, + name: str = "Knowledge Base", + provider: str = "vendor", + retrieval_model: dict | None = None, +) -> Dataset: + dataset = Dataset( + id=dataset_id or str(uuid.uuid4()), + tenant_id=tenant_id or str(uuid.uuid4()), + name=name, + provider=provider, + data_source_type=DataSourceType.UPLOAD_FILE, + indexing_technique="high_quality", + retrieval_model=retrieval_model, + created_by=str(uuid.uuid4()), + ) + session.add(dataset) + session.commit() + return dataset + + +def _persist_document( + session: Session, + *, + dataset: Dataset, + name: str, + data_source_type: DataSourceType = DataSourceType.UPLOAD_FILE, + doc_metadata: dict | None = None, +) -> Document: + document = Document( + id=str(uuid.uuid4()), + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + position=1, + data_source_type=data_source_type, + batch="batch-1", + name=name, + created_from=DocumentCreatedFrom.WEB, + created_by=dataset.created_by, + doc_metadata=doc_metadata, + ) + session.add(document) + session.commit() + return document + + class _FakeFlaskApp: def app_context(self): return nullcontext() +class _FakeCurrentApp: + def _get_current_object(self) -> _FakeFlaskApp: + return _FakeFlaskApp() + + +class _DatabaseWithSession: + def __init__(self, session: scoped_session[Session]) -> None: + self.session = session + + +class _UnusedProviderManager(ProviderManager): + def __init__(self) -> None: + pass + + +class _UnusedModelInstance(ModelInstance): + def __init__(self) -> None: + pass + + class _ImmediateThread: def __init__(self, target=None, kwargs=None, **_kwargs): self._target = target @@ -105,8 +181,13 @@ def test_load_yaml_file_cached_hits(tmp_path): assert load_yaml_file_cached.cache_info().hits == 1 -def test_single_dataset_retriever_from_dataset_builds_name_and_description(): - dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1", name="Knowledge", description=None) +def test_single_dataset_retriever_from_dataset_builds_name_and_description(sqlite_session: Session): + dataset = _persist_dataset( + sqlite_session, + dataset_id="dataset-1", + tenant_id="tenant-1", + name="Knowledge", + ) tool = SingleDatasetRetrieverTool.from_dataset( dataset=dataset, @@ -120,31 +201,21 @@ def test_single_dataset_retriever_from_dataset_builds_name_and_description(): assert tool.description == "useful for when you want to answer queries about the Knowledge" -def test_single_dataset_retriever_external_run_returns_content_and_resources(): - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - name="Knowledge Base", - provider="external", - indexing_technique="high_quality", - retrieval_model={}, - ) +def test_single_dataset_retriever_external_run_returns_content_and_resources(sqlite_session: Session): + dataset = _persist_dataset(sqlite_session, provider="external", retrieval_model={}) callback = _TestHitCallback() - dataset_retrieval = Mock() - dataset_retrieval.get_metadata_filter_condition.return_value = ( - {"dataset-1": ["doc-a"]}, + metadata_filter_result = ( + {dataset.id: ["doc-a"]}, {"logical_operator": "and"}, ) - session = Mock() - session.scalar.return_value = dataset external_documents = [ {"content": "first", "metadata": {"document_id": "doc-a"}, "score": 0.9, "title": "Doc A"}, {"content": "second", "metadata": {"document_id": "doc-b"}, "score": 0.8, "title": "Doc B"}, ] tool = SingleDatasetRetrieverTool( - tenant_id="tenant-1", - dataset_id="dataset-1", + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, retrieve_config=_retrieve_config(), return_resource=True, retriever_from="dev", @@ -152,41 +223,35 @@ def test_single_dataset_retriever_external_run_returns_content_and_resources(): inputs={"x": 1}, ) - with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): + with patch.object( + single_retriever_module.DatasetRetrieval, + "get_metadata_filter_condition", + return_value=metadata_filter_result, + ): with patch.object( single_retriever_module.ExternalDatasetService, "fetch_external_knowledge_retrieval", return_value=external_documents, ) as fetch_mock: - result = tool.run(session=session, query="hello") + result = tool.run(session=sqlite_session, query="hello") assert result == "first\nsecond" - assert callback.queries == [("hello", "dataset-1")] + assert callback.queries == [("hello", dataset.id)] assert callback.resources is not None resource_info = callback.resources assert [item.position for item in resource_info] == [1, 2] - assert resource_info[0].dataset_id == "dataset-1" + assert resource_info[0].dataset_id == dataset.id fetch_mock.assert_called_once() - assert fetch_mock.call_args.kwargs["session"] is session + assert fetch_mock.call_args.kwargs["session"] is sqlite_session -def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_documents(): - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - name="Knowledge Base", - provider="internal", - indexing_technique="high_quality", - retrieval_model=None, - ) - dataset_retrieval = Mock() - dataset_retrieval.get_metadata_filter_condition.return_value = ({"dataset-1": []}, {"logical_operator": "and"}) - session = Mock() - session.scalar.return_value = dataset - +def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_documents( + sqlite_session: Session, +): + dataset = _persist_dataset(sqlite_session) tool = SingleDatasetRetrieverTool( - tenant_id="tenant-1", - dataset_id="dataset-1", + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, retrieve_config=_retrieve_config(), return_resource=False, retriever_from="prod", @@ -194,21 +259,21 @@ def test_single_dataset_retriever_returns_empty_when_metadata_filter_finds_no_do inputs={}, ) - with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): + with patch.object( + single_retriever_module.DatasetRetrieval, + "get_metadata_filter_condition", + return_value=({dataset.id: []}, {"logical_operator": "and"}), + ): with patch.object(single_retriever_module.RetrievalService, "retrieve") as retrieve_mock: - result = tool.run(session=session, query="hello") + result = tool.run(session=sqlite_session, query="hello") assert result == "" retrieve_mock.assert_not_called() -def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources(): - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - name="Knowledge Base", - provider="internal", - indexing_technique="high_quality", +def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources(sqlite_session: Session): + dataset = _persist_dataset( + sqlite_session, retrieval_model={ "search_method": "semantic_search", "score_threshold_enabled": True, @@ -219,54 +284,64 @@ def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources(): "weights": {"vector_setting": {"vector_weight": 0.6}}, }, ) + document_low = _persist_document( + sqlite_session, + dataset=dataset, + name="Document Low", + doc_metadata={"lang": "en"}, + ) + document_high = _persist_document( + sqlite_session, + dataset=dataset, + name="Document High", + data_source_type=DataSourceType.NOTION_IMPORT, + doc_metadata={"lang": "fr"}, + ) callback = _TestHitCallback() - dataset_retrieval = Mock() - dataset_retrieval.get_metadata_filter_condition.return_value = (None, None) - low_segment = SimpleNamespace( - id="seg-low", - dataset_id="dataset-1", - document_id="doc-low", + low_segment = DocumentSegment( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id=document_low.id, + index_node_id="node-low", content="raw low", - answer="low answer", hit_count=1, word_count=10, position=3, index_node_hash="hash-low", - get_sign_content=lambda: "signed low", + tokens=10, + created_by=dataset.created_by, + answer="low answer", + status=SegmentStatus.COMPLETED, + completed_at=datetime.now(), ) - high_segment = SimpleNamespace( - id="seg-high", - dataset_id="dataset-1", - document_id="doc-high", + high_segment = DocumentSegment( + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id=document_high.id, + index_node_id="node-high", content="raw high", - answer=None, hit_count=9, word_count=25, position=1, index_node_hash="hash-high", - get_sign_content=lambda: "signed high", + tokens=25, + created_by=dataset.created_by, + status=SegmentStatus.COMPLETED, + completed_at=datetime.now(), ) + sqlite_session.add_all([low_segment, high_segment]) + sqlite_session.commit() records = [ - SimpleNamespace(segment=low_segment, score=0.2, summary="summary low"), - SimpleNamespace(segment=high_segment, score=0.9, summary=None), + RetrievalSegments(segment=low_segment, score=0.2, summary="summary low"), + RetrievalSegments(segment=high_segment, score=0.9), ] documents = [ RagDocument(page_content="first", metadata={"doc_id": "node-low", "score": 0.2}), RagDocument(page_content="second", metadata={"doc_id": "node-high", "score": 0.9}), ] - lookup_doc_low = SimpleNamespace( - id="doc-low", name="Document Low", data_source_type="upload_file", doc_metadata={"lang": "en"} - ) - lookup_doc_high = SimpleNamespace( - id="doc-high", name="Document High", data_source_type="notion", doc_metadata={"lang": "fr"} - ) - session = Mock() - session.scalar.side_effect = [dataset, lookup_doc_low, lookup_doc_high] - session.get.return_value = dataset - tool = SingleDatasetRetrieverTool( - tenant_id="tenant-1", - dataset_id="dataset-1", + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, retrieve_config=_retrieve_config(), return_resource=True, retriever_from="dev", @@ -275,21 +350,28 @@ def test_single_dataset_retriever_non_economy_run_sorts_context_and_resources(): top_k=2, ) - with patch.object(single_retriever_module, "DatasetRetrieval", return_value=dataset_retrieval): - with patch.object(single_retriever_module.RetrievalService, "retrieve", return_value=documents): - with patch.object( - single_retriever_module.RetrievalService, - "format_retrieval_documents", - return_value=records, - ): - result = tool.run(session=session, query="hello") + with ( + patch.object( + single_retriever_module.DatasetRetrieval, + "get_metadata_filter_condition", + return_value=(None, None), + ), + patch.object(single_retriever_module.RetrievalService, "retrieve", return_value=documents), + patch.object( + single_retriever_module.RetrievalService, + "format_retrieval_documents", + return_value=records, + ), + patch.object(DocumentSegment, "get_sign_content", lambda segment: segment.content.replace("raw", "signed")), + ): + result = tool.run(session=sqlite_session, query="hello") assert result == "signed high\nsummary low\nquestion:signed low answer:low answer" assert callback.documents == documents assert callback.resources is not None resource_info = callback.resources assert [item.position for item in resource_info] == [1, 2] - assert resource_info[0].segment_id == "seg-high" + assert resource_info[0].segment_id == high_segment.id assert resource_info[0].hit_count == 9 assert resource_info[1].summary == "summary low" assert resource_info[1].content == "question:raw low \nanswer:low answer" @@ -308,11 +390,12 @@ def test_multi_dataset_retriever_from_dataset_sets_tool_name(): assert tool.name == "dataset_tenant_1" -def test_multi_dataset_retriever_retriever_returns_early_when_dataset_is_missing(): +def test_multi_dataset_retriever_retriever_returns_early_when_dataset_is_missing( + sqlite_session_factory: sessionmaker[Session], +): callback = _TestHitCallback() all_documents: list[RagDocument] = [] - db_session = Mock() - db_session.scalar.return_value = None + db_session = scoped_session(sqlite_session_factory) tool = DatasetMultiRetrieverTool( tenant_id="tenant-1", dataset_ids=["dataset-1"], @@ -322,15 +405,18 @@ def test_multi_dataset_retriever_retriever_returns_early_when_dataset_is_missing retriever_from="prod", ) - with patch.object(multi_retriever_module, "db", SimpleNamespace(session=db_session)): - with patch.object(multi_retriever_module.RetrievalService, "retrieve") as retrieve_mock: - result = tool._retriever( - flask_app=_FakeFlaskApp(), - dataset_id="dataset-1", - query="hello", - all_documents=all_documents, - hit_callbacks=[callback], - ) + try: + with patch.object(multi_retriever_module, "db", _DatabaseWithSession(db_session)): + with patch.object(multi_retriever_module.RetrievalService, "retrieve") as retrieve_mock: + result = tool._retriever( + flask_app=_FakeFlaskApp(), + dataset_id=str(uuid.uuid4()), + query="hello", + all_documents=all_documents, + hit_callbacks=[callback], + ) + finally: + db_session.remove() assert result == [] assert all_documents == [] @@ -338,11 +424,12 @@ def test_multi_dataset_retriever_retriever_returns_early_when_dataset_is_missing retrieve_mock.assert_not_called() -def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model(): - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - indexing_technique="high_quality", +def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +): + dataset = _persist_dataset( + sqlite_session, retrieval_model={ "search_method": "semantic_search", "top_k": 6, @@ -356,11 +443,10 @@ def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model(): callback = _TestHitCallback() documents = [RagDocument(page_content="retrieved", metadata={"doc_id": "node-1", "score": 0.4})] all_documents: list[RagDocument] = [] - db_session = Mock() - db_session.scalar.return_value = dataset + db_session = scoped_session(sqlite_session_factory) tool = DatasetMultiRetrieverTool( - tenant_id="tenant-1", - dataset_ids=["dataset-1"], + tenant_id=dataset.tenant_id, + dataset_ids=[dataset.id], reranking_provider_name="provider", reranking_model_name="model", return_resource=False, @@ -368,21 +454,26 @@ def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model(): top_k=2, ) - with patch.object(multi_retriever_module, "db", SimpleNamespace(session=db_session)): - with patch.object(multi_retriever_module.RetrievalService, "retrieve", return_value=documents) as retrieve_mock: - tool._retriever( - flask_app=_FakeFlaskApp(), - dataset_id="dataset-1", - query="hello", - all_documents=all_documents, - hit_callbacks=[callback], - ) + try: + with patch.object(multi_retriever_module, "db", _DatabaseWithSession(db_session)): + with patch.object( + multi_retriever_module.RetrievalService, "retrieve", return_value=documents + ) as retrieve_mock: + tool._retriever( + flask_app=_FakeFlaskApp(), + dataset_id=dataset.id, + query="hello", + all_documents=all_documents, + hit_callbacks=[callback], + ) + finally: + db_session.remove() assert all_documents == documents - assert callback.queries == [("hello", "dataset-1")] + assert callback.queries == [("hello", dataset.id)] retrieve_mock.assert_called_once_with( retrieval_method="semantic_search", - dataset_id="dataset-1", + dataset_id=dataset.id, query="hello", top_k=6, score_threshold=0.4, @@ -392,11 +483,26 @@ def test_multi_dataset_retriever_retriever_non_economy_uses_retrieval_model(): ) -def test_multi_dataset_retriever_run_orders_segments_and_returns_resources(): +def test_multi_dataset_retriever_run_orders_segments_and_returns_resources(sqlite_session: Session): + dataset_one = _persist_dataset(sqlite_session, name="Dataset One") + dataset_two = _persist_dataset(sqlite_session, tenant_id=dataset_one.tenant_id, name="Dataset Two") + document_two = _persist_document( + sqlite_session, + dataset=dataset_one, + name="Doc Two", + data_source_type=DataSourceType.NOTION_IMPORT, + doc_metadata={"p": 2}, + ) + document_one = _persist_document( + sqlite_session, + dataset=dataset_two, + name="Doc One", + doc_metadata={"p": 1}, + ) callback = _TestHitCallback() tool = DatasetMultiRetrieverTool( - tenant_id="tenant-1", - dataset_ids=["dataset-1", "dataset-2"], + tenant_id=dataset_one.tenant_id, + dataset_ids=[dataset_one.id, dataset_two.id], reranking_provider_name="provider", reranking_model_name="model", return_resource=True, @@ -409,64 +515,63 @@ def test_multi_dataset_retriever_run_orders_segments_and_returns_resources(): second_doc = RagDocument(page_content="second", metadata={"doc_id": "node-1", "score": 0.9}) def fake_retriever(**kwargs): - if kwargs["dataset_id"] == "dataset-1": + if kwargs["dataset_id"] == dataset_one.id: kwargs["all_documents"].append(first_doc) else: kwargs["all_documents"].append(second_doc) - segment_for_node_2 = SimpleNamespace( - id="seg-2", - dataset_id="dataset-1", - document_id="doc-2", + segment_for_node_2 = DocumentSegment( + tenant_id=dataset_one.tenant_id, + dataset_id=dataset_one.id, + document_id=document_two.id, index_node_id="node-2", content="raw two", - answer="answer two", hit_count=2, word_count=20, position=2, index_node_hash="hash-2", - get_sign_content=lambda: "signed two", + tokens=20, + created_by=dataset_one.created_by, + answer="answer two", + status=SegmentStatus.COMPLETED, + completed_at=datetime.now(), ) - segment_for_node_1 = SimpleNamespace( - id="seg-1", - dataset_id="dataset-2", - document_id="doc-1", + segment_for_node_1 = DocumentSegment( + tenant_id=dataset_two.tenant_id, + dataset_id=dataset_two.id, + document_id=document_one.id, index_node_id="node-1", content="raw one", - answer=None, hit_count=7, word_count=30, position=1, index_node_hash="hash-1", - get_sign_content=lambda: "signed one", + tokens=30, + created_by=dataset_two.created_by, + status=SegmentStatus.COMPLETED, + completed_at=datetime.now(), ) - db_session = Mock() - db_session.scalars.return_value.all.return_value = [segment_for_node_2, segment_for_node_1] - db_session.get.side_effect = [ - SimpleNamespace(id="dataset-2", name="Dataset Two"), - SimpleNamespace(id="dataset-1", name="Dataset One"), - ] - db_session.scalar.side_effect = [ - SimpleNamespace(id="doc-1", name="Doc One", data_source_type="upload_file", doc_metadata={"p": 1}), - SimpleNamespace(id="doc-2", name="Doc Two", data_source_type="notion", doc_metadata={"p": 2}), - ] - model_manager = Mock() - model_manager.get_model_instance.return_value = Mock() - rerank_runner = Mock() - rerank_runner.run.return_value = [second_doc, first_doc] - fake_current_app = SimpleNamespace(_get_current_object=lambda: _FakeFlaskApp()) + sqlite_session.add_all([segment_for_node_2, segment_for_node_1]) + sqlite_session.commit() + model_manager = ModelManager(provider_manager=_UnusedProviderManager()) + model_instance = _UnusedModelInstance() + rerank_runner = RerankModelRunner(model_instance, session=sqlite_session) + fake_current_app = _FakeCurrentApp() - with patch.object(tool, "_retriever", side_effect=fake_retriever) as retriever_mock: - with patch.object(multi_retriever_module, "current_app", fake_current_app): - with patch.object(multi_retriever_module.threading, "Thread", _ImmediateThread): - with patch.object(multi_retriever_module.ModelManager, "for_tenant", return_value=model_manager): - with patch.object( - multi_retriever_module, "RerankModelRunner", return_value=rerank_runner - ) as rerank_runner_class: - result = tool.run(session=db_session, query="hello") + with ( + patch.object(DocumentSegment, "get_sign_content", lambda segment: segment.content.replace("raw", "signed")), + patch.object(tool, "_retriever", side_effect=fake_retriever) as retriever_mock, + patch.object(multi_retriever_module, "current_app", fake_current_app), + patch.object(multi_retriever_module.threading, "Thread", _ImmediateThread), + patch.object(multi_retriever_module.ModelManager, "for_tenant", return_value=model_manager), + patch.object(model_manager, "get_model_instance", return_value=model_instance), + patch.object(multi_retriever_module, "RerankModelRunner", return_value=rerank_runner) as rerank_runner_class, + patch.object(rerank_runner, "run", return_value=[second_doc, first_doc]), + ): + result = tool.run(session=sqlite_session, query="hello") assert result == "signed one\nquestion:signed two answer:answer two" - rerank_runner_class.assert_called_once_with(model_manager.get_model_instance.return_value, session=db_session) + rerank_runner_class.assert_called_once_with(model_instance, session=sqlite_session) assert retriever_mock.call_count == 2 assert callback.documents == [second_doc, first_doc] assert callback.resources is not None diff --git a/api/tests/unit_tests/core/tools/utils/test_model_invocation_utils.py b/api/tests/unit_tests/core/tools/utils/test_model_invocation_utils.py index 44785f939ca..357a62e9b11 100644 --- a/api/tests/unit_tests/core/tools/utils/test_model_invocation_utils.py +++ b/api/tests/unit_tests/core/tools/utils/test_model_invocation_utils.py @@ -4,6 +4,7 @@ Covers success and error branches for ModelInvocationUtils, including InvokeModelError and invoke error mappings for InvokeAuthorizationError, InvokeBadRequestError, InvokeConnectionError, InvokeRateLimitError, and InvokeServerUnavailableError. Assumes mocked model instances and managers. +Invocation logging uses a real SQLite-backed SQLAlchemy session. """ from __future__ import annotations @@ -14,7 +15,10 @@ from typing import Any from unittest.mock import Mock, patch import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session +from core.tools.entities.tool_entities import ToolProviderType from core.tools.utils.model_invocation_utils import InvokeModelError, ModelInvocationUtils from graphon.model_runtime.entities.model_entities import ModelPropertyKey from graphon.model_runtime.errors.invoke import ( @@ -24,6 +28,11 @@ from graphon.model_runtime.errors.invoke import ( InvokeRateLimitError, InvokeServerUnavailableError, ) +from models.tools import ToolModelInvoke + +TENANT_ID = "11111111-1111-1111-1111-111111111111" +USER_ID = "22222222-2222-2222-2222-222222222222" +CALLER_ID = "33333333-3333-3333-3333-333333333333" def _mock_model_instance(*, schema: dict[str, Any] | None = None) -> SimpleNamespace: @@ -80,7 +89,7 @@ def test_calculate_tokens_handles_missing_model(): mock_factory.assert_called_once_with(tenant_id="tenant", user_id=None) -def test_invoke_success_and_error_mappings(): +def test_invoke_success_and_error_mappings(sqlite_session: Session): model_instance = _mock_model_instance(schema={ModelPropertyKey.CONTEXT_SIZE: 2048}) model_instance.invoke_llm.return_value = SimpleNamespace( message=SimpleNamespace(content="ok"), @@ -96,28 +105,37 @@ def test_invoke_success_and_error_mappings(): manager = Mock() manager.get_default_model_instance.return_value = model_instance - class _ToolModelInvoke: - def __init__(self, **kwargs): - self.__dict__.update(kwargs) + database = SimpleNamespace(session=sqlite_session) - db_mock = SimpleNamespace(session=Mock()) - - with patch("core.tools.utils.model_invocation_utils.ModelManager.for_tenant", return_value=manager) as mock_factory: - with patch("core.tools.utils.model_invocation_utils.ToolModelInvoke", _ToolModelInvoke): - with patch("core.tools.utils.model_invocation_utils.db", db_mock): - response = ModelInvocationUtils.invoke( - user_id="u1", - tenant_id="tenant", - tool_type="builtin", - tool_name="tool-a", - prompt_messages=[], - caller_user_id="caller-1", - ) + with ( + patch("core.tools.utils.model_invocation_utils.ModelManager.for_tenant", return_value=manager) as mock_factory, + patch("core.tools.utils.model_invocation_utils.db", database), + patch.object(sqlite_session, "add", wraps=sqlite_session.add) as mock_session_add, + patch.object(sqlite_session, "commit", wraps=sqlite_session.commit) as mock_session_commit, + ): + response = ModelInvocationUtils.invoke( + user_id=USER_ID, + tenant_id=TENANT_ID, + tool_type=ToolProviderType.BUILT_IN, + tool_name="tool-a", + prompt_messages=[], + caller_user_id=CALLER_ID, + ) assert response.message.content == "ok" - assert db_mock.session.add.call_count == 1 - assert db_mock.session.commit.call_count == 2 - mock_factory.assert_called_once_with(tenant_id="tenant", user_id="caller-1") + assert mock_session_add.call_count == 1 + assert mock_session_commit.call_count == 2 + assert not sqlite_session.in_transaction() + persisted = sqlite_session.scalar(select(ToolModelInvoke)) + assert persisted is not None + assert persisted.user_id == USER_ID + assert persisted.tenant_id == TENANT_ID + assert persisted.tool_type == ToolProviderType.BUILT_IN + assert persisted.model_response == "ok" + assert persisted.prompt_tokens == 5 + assert persisted.answer_tokens == 7 + assert persisted.total_price == Decimal("0.7000000") + mock_factory.assert_called_once_with(tenant_id=TENANT_ID, user_id=CALLER_ID) @pytest.mark.parametrize( @@ -139,27 +157,30 @@ def test_invoke_success_and_error_mappings(): "generic-error", ], ) -def test_invoke_error_mappings(exc, expected): +def test_invoke_error_mappings(exc, expected, sqlite_session: Session): model_instance = _mock_model_instance(schema={ModelPropertyKey.CONTEXT_SIZE: 2048}) model_instance.invoke_llm.side_effect = exc manager = Mock() manager.get_default_model_instance.return_value = model_instance - class _ToolModelInvoke: - def __init__(self, **kwargs): - self.__dict__.update(kwargs) + database = SimpleNamespace(session=sqlite_session) - db_mock = SimpleNamespace(session=Mock()) - - with patch("core.tools.utils.model_invocation_utils.ModelManager.for_tenant", return_value=manager) as mock_factory: - with patch("core.tools.utils.model_invocation_utils.ToolModelInvoke", _ToolModelInvoke): - with patch("core.tools.utils.model_invocation_utils.db", db_mock): - with pytest.raises(InvokeModelError, match=expected): - ModelInvocationUtils.invoke( - user_id="u1", - tenant_id="tenant", - tool_type="builtin", - tool_name="tool-a", - prompt_messages=[], - ) - mock_factory.assert_called_once_with(tenant_id="tenant", user_id="u1") + with ( + patch("core.tools.utils.model_invocation_utils.ModelManager.for_tenant", return_value=manager) as mock_factory, + patch("core.tools.utils.model_invocation_utils.db", database), + ): + with pytest.raises(InvokeModelError, match=expected): + ModelInvocationUtils.invoke( + user_id=USER_ID, + tenant_id=TENANT_ID, + tool_type=ToolProviderType.BUILT_IN, + tool_name="tool-a", + prompt_messages=[], + ) + assert not sqlite_session.in_transaction() + persisted = sqlite_session.scalar(select(ToolModelInvoke)) + assert persisted is not None + assert persisted.model_response == "" + assert persisted.prompt_tokens == 5 + assert persisted.answer_tokens == 0 + mock_factory.assert_called_once_with(tenant_id=TENANT_ID, user_id=USER_ID) diff --git a/api/tests/unit_tests/core/tools/workflow_as_tool/test_provider.py b/api/tests/unit_tests/core/tools/workflow_as_tool/test_provider.py index b876fa64b9c..45d9bdbf770 100644 --- a/api/tests/unit_tests/core/tools/workflow_as_tool/test_provider.py +++ b/api/tests/unit_tests/core/tools/workflow_as_tool/test_provider.py @@ -1,12 +1,17 @@ from __future__ import annotations import json +import uuid +from collections.abc import Iterator from types import SimpleNamespace from typing import Any, cast -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import PropertyMock, patch import pytest +from sqlalchemy import Engine +from sqlalchemy.orm import Session, sessionmaker +from core.db.session_factory import session_factory from core.tools.__base.tool_runtime import ToolRuntime from core.tools.entities.common_entities import I18nObject from core.tools.entities.tool_entities import ( @@ -20,14 +25,31 @@ from core.tools.entities.tool_entities import ( ) from core.tools.workflow_as_tool.provider import WorkflowToolProviderController from core.tools.workflow_as_tool.tool import WorkflowTool +from extensions.ext_database import db from graphon.variables.input_entities import VariableEntity, VariableEntityType from models.account import Account -from models.model import App +from models.base import TypeBase +from models.model import App, AppMode, IconType from models.tools import WorkflowToolProvider from models.workflow import Workflow, WorkflowType -def _controller() -> WorkflowToolProviderController: +@pytest.fixture +def database_session(sqlite_engine: Engine) -> Iterator[Session]: + models = (Account, App, Workflow, WorkflowToolProvider) + tables = [model.metadata.tables[model.__tablename__] for model in models] + TypeBase.metadata.create_all(sqlite_engine, tables=tables) + session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + + with ( + patch.object(session_factory, "create_session", session_maker), + patch.object(type(db), "engine", new_callable=PropertyMock, return_value=sqlite_engine), + ): + with session_maker() as session: + yield session + + +def _controller(provider_id: str = "provider-1") -> WorkflowToolProviderController: entity = ToolProviderEntity( identity=ToolProviderIdentity( author="author", @@ -38,47 +60,64 @@ def _controller() -> WorkflowToolProviderController: ), credentials_schema=[], ) - return WorkflowToolProviderController(entity=entity, provider_id="provider-1") + return WorkflowToolProviderController(entity=entity, provider_id=provider_id) -def _app() -> App: - return App(id="app-1") +def _app(*, tenant_id: str | None = None) -> App: + return App( + id=str(uuid.uuid4()), + tenant_id=tenant_id or str(uuid.uuid4()), + name="Workflow App", + mode=AppMode.WORKFLOW, + icon_type=IconType.EMOJI, + icon="workflow", + icon_background="#FFFFFF", + enable_site=True, + enable_api=False, + ) def _account() -> Account: return Account(name="Alice", email="alice@example.com") -def _workflow() -> Workflow: +def _workflow(app: App, account: Account | None = None) -> Workflow: return Workflow.new( - tenant_id="tenant-1", - app_id="app-1", + tenant_id=app.tenant_id, + app_id=app.id, type=WorkflowType.WORKFLOW.value, version="1", graph=json.dumps({"nodes": []}), features="{}", - created_by="user-1", + created_by=account.id if account else str(uuid.uuid4()), environment_variables=[], conversation_variables=[], rag_pipeline_variables=[], ) -def _db_provider(*, parameter_configuration: str = "[]") -> WorkflowToolProvider: +def _db_provider( + app: App, + account: Account, + *, + parameter_configuration: str = "[]", +) -> WorkflowToolProvider: return WorkflowToolProvider( name="workflow_tool", label="WF Provider", icon="icon.svg", - app_id="app-1", + app_id=app.id, version="1", - user_id="user-1", - tenant_id="tenant-1", + user_id=account.id, + tenant_id=app.tenant_id, description="desc", parameter_configuration=parameter_configuration, ) -def _workflow_tool(name: str = "workflow_tool") -> WorkflowTool: +def _workflow_tool(name: str = "workflow_tool", *, tenant_id: str | None = None) -> WorkflowTool: + app = _app(tenant_id=tenant_id) + workflow = _workflow(app) return WorkflowTool( workflow_as_tool_id="provider-1", entity=ToolEntity( @@ -91,38 +130,46 @@ def _workflow_tool(name: str = "workflow_tool") -> WorkflowTool: description=ToolDescription(human=I18nObject(en_US="desc"), llm="desc"), parameters=[], ), - runtime=ToolRuntime(tenant_id="tenant-1"), - workflow_app_id="app-1", - workflow_entities={"app": _app(), "workflow": _workflow()}, + runtime=ToolRuntime(tenant_id=app.tenant_id), + workflow_app_id=app.id, + workflow_entities={"app": app, "workflow": workflow}, version="1", workflow_call_depth=0, ) -def _mock_session_with_begin() -> Mock: - session = Mock() - begin_cm = Mock() - begin_cm.__enter__ = Mock(return_value=None) - begin_cm.__exit__ = Mock(return_value=False) - session.begin.return_value = begin_cm - return session - - -def test_get_db_provider_tool_builds_entity(): - controller = _controller() - session = Mock() - workflow = _workflow() - session.scalar.return_value = workflow +def _persist_provider_graph( + session: Session, + *, + parameter_configuration: str = "[]", + include_app: bool = True, + include_workflow: bool = True, +) -> tuple[WorkflowToolProvider, App, Account, Workflow]: + account = _account() app = _app() - db_provider = _db_provider( + workflow = _workflow(app, account) + db_provider = _db_provider(app, account, parameter_configuration=parameter_configuration) + + session.add_all([account, db_provider]) + if include_app: + session.add(app) + if include_workflow: + session.add(workflow) + session.commit() + return db_provider, app, account, workflow + + +def test_get_db_provider_tool_builds_entity(database_session: Session): + db_provider, app, user, _ = _persist_provider_graph( + database_session, parameter_configuration=json.dumps( [ {"name": "country", "description": "Country", "form": ToolParameter.ToolParameterForm.FORM.value}, {"name": "files", "description": "files", "form": ToolParameter.ToolParameterForm.FORM.value}, ] - ) + ), ) - user = _account() + controller = _controller(db_provider.id) variables = [ VariableEntity( variable="country", @@ -152,7 +199,7 @@ def test_get_db_provider_tool_builds_entity(): return_value=outputs, ), ): - tool = controller._get_db_provider_tool(db_provider, app, session=session, user=user) + tool = controller._get_db_provider_tool(db_provider, app, session=database_session, user=user) assert tool.entity.identity.name == "workflow_tool" # "json" output is reserved for ToolInvokeMessage.VariableMessage and filtered out. @@ -175,61 +222,55 @@ def test_get_tool_returns_hit_or_none(): def test_get_tools_returns_cached(): controller = _controller() - cached_tools = [_workflow_tool("wf-cached")] + cached_tools = [_workflow_tool("wf-cached", tenant_id="tenant-1")] controller.tools = cached_tools assert controller.get_tools("tenant-1") == cached_tools -def test_from_db_builds_controller(): - app = _app() - user = _account() - db_provider = _db_provider() - session = _mock_session_with_begin() - session.scalar.return_value = db_provider - session.get.side_effect = [app, user] - fake_cm = MagicMock() - fake_cm.__enter__.return_value = session - fake_cm.__exit__.return_value = False - fake_session_factory = Mock() - fake_session_factory.create_session.return_value = fake_cm +def test_from_db_builds_controller(database_session: Session): + db_provider, app, user, workflow = _persist_provider_graph(database_session) + + with ( + patch( + "core.tools.workflow_as_tool.provider.WorkflowAppConfigManager.convert_features", + return_value=SimpleNamespace(file_upload=False), + ), + patch( + "core.tools.workflow_as_tool.provider.WorkflowToolConfigurationUtils.get_workflow_graph_variables", + return_value=[], + ), + patch( + "core.tools.workflow_as_tool.provider.WorkflowToolConfigurationUtils.get_workflow_graph_output", + return_value=[], + ), + ): + built = WorkflowToolProviderController.from_db(db_provider) - with patch("core.tools.workflow_as_tool.provider.session_factory", fake_session_factory): - with patch.object( - WorkflowToolProviderController, - "_get_db_provider_tool", - return_value=_workflow_tool("wf"), - ): - built = WorkflowToolProviderController.from_db(db_provider) assert isinstance(built, WorkflowToolProviderController) - assert built.tools + assert built.entity.identity.author == user.name + assert built.provider_id == db_provider.id + assert built.tools is not None + assert built.tools[0].workflow_app_id == app.id + assert built.tools[0].workflow_entities["workflow"].id == workflow.id -def test_get_tools_returns_empty_when_provider_missing(): - controller = _controller() +def test_get_tools_returns_empty_when_provider_missing(database_session: Session): + db_provider, _, _, _ = _persist_provider_graph(database_session) + controller = _controller(db_provider.id) controller.tools = None - with patch("core.tools.workflow_as_tool.provider.db") as mock_db: - mock_db.engine = object() - with patch("core.tools.workflow_as_tool.provider.Session") as session_cls: - session = _mock_session_with_begin() - session.scalar.return_value = None - session_cls.return_value.__enter__.return_value = session - - assert controller.get_tools("tenant-1") == [] + assert controller.get_tools(str(uuid.uuid4())) == [] -def test_get_tools_raises_when_app_missing(): - controller = _controller() +def test_get_tools_raises_when_app_missing(database_session: Session): + db_provider, _, _, _ = _persist_provider_graph( + database_session, + include_app=False, + include_workflow=False, + ) + controller = _controller(db_provider.id) controller.tools = None - db_provider = _db_provider() - with patch("core.tools.workflow_as_tool.provider.db") as mock_db: - mock_db.engine = object() - with patch("core.tools.workflow_as_tool.provider.Session") as session_cls: - session = _mock_session_with_begin() - session.scalar.return_value = db_provider - session.get.return_value = None - session_cls.return_value.__enter__.return_value = session - with pytest.raises(ValueError, match="app not found"): - controller.get_tools("tenant-1") + with pytest.raises(ValueError, match="app not found"): + controller.get_tools(db_provider.tenant_id) diff --git a/api/tests/unit_tests/core/tools/workflow_as_tool/test_tool.py b/api/tests/unit_tests/core/tools/workflow_as_tool/test_tool.py index 8df8e28bda4..34fe819e0d8 100644 --- a/api/tests/unit_tests/core/tools/workflow_as_tool/test_tool.py +++ b/api/tests/unit_tests/core/tools/workflow_as_tool/test_tool.py @@ -1,16 +1,16 @@ -"""Unit tests for workflow-as-tool behavior. - -StubSession/StubScalars emulate SQLAlchemy session/scalars with minimal methods -(`scalar`, `scalars`, `expunge`, `commit`, `refresh`, context manager) to keep -database access mocked and predictable in tests. -""" +"""Unit tests for workflow-as-tool behavior with real SQLite ORM boundaries.""" import json +import uuid +from collections.abc import Iterator +from dataclasses import dataclass from types import SimpleNamespace from typing import Any from unittest.mock import MagicMock, Mock, patch import pytest +from sqlalchemy import Engine, inspect +from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool_runtime import ToolRuntime @@ -23,74 +23,142 @@ from core.tools.entities.tool_entities import ( ToolProviderType, ) from core.tools.errors import ToolInvokeError +from core.tools.workflow_as_tool import tool as workflow_tool_module from core.tools.workflow_as_tool.tool import WorkflowTool from graphon.file import FILE_MODEL_IDENTITY, FileTransferMethod, FileType +from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole +from models.base import TypeBase +from models.enums import EndUserType +from models.model import App, AppMode, EndUser +from models.workflow import Workflow, WorkflowType + +TENANT_ID = "00000000-0000-0000-0000-000000000001" +OTHER_TENANT_ID = "00000000-0000-0000-0000-000000000002" +APP_ID = "00000000-0000-0000-0000-000000000003" +ACCOUNT_ID = "00000000-0000-0000-0000-000000000004" +END_USER_ID = "00000000-0000-0000-0000-000000000005" +CREATOR_ID = "00000000-0000-0000-0000-000000000006" -class StubScalars: - """Minimal stub for SQLAlchemy scalar results.""" - - _value: Any - - def __init__(self, value: Any) -> None: - self._value = value - - def first(self) -> Any: - return self._value +@dataclass(frozen=True) +class SqliteToolDb: + engine: Engine + session_maker: sessionmaker[Session] + caller_session: Session -class StubSession: - """Minimal stub for session_factory-created sessions.""" +@pytest.fixture +def sqlite_tool_db( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, +) -> Iterator[SqliteToolDb]: + """Bind service-owned sessions and Account tenant reloads to SQLite.""" + models = (App, Workflow, EndUser, Account, Tenant, TenantAccountJoin) + TypeBase.metadata.create_all(sqlite_engine, tables=[model.__table__ for model in models]) + session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr(workflow_tool_module.session_factory, "create_session", session_maker) - scalar_results: list[Any] - scalars_results: list[Any] - expunge_calls: list[object] + from models import account as account_module - def __init__(self, *, scalar_results: list[Any] | None = None, scalars_results: list[Any] | None = None) -> None: - self.scalar_results = list(scalar_results or []) - self.scalars_results = list(scalars_results or []) - self.expunge_calls: list[object] = [] - - def scalar(self, _stmt: Any) -> Any: - return self.scalar_results.pop(0) - - def scalars(self, _stmt: Any) -> StubScalars: - return StubScalars(self.scalars_results.pop(0)) - - def expunge(self, value: Any) -> None: - self.expunge_calls.append(value) - - def begin(self) -> "StubSession": - return self - - def commit(self) -> None: - pass - - def refresh(self, _value: Any) -> None: - pass - - def close(self) -> None: - pass - - def __enter__(self) -> "StubSession": - return self - - def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> bool: - return False + monkeypatch.setattr(account_module, "db", SimpleNamespace(engine=sqlite_engine)) + with session_maker() as caller_session: + yield SqliteToolDb(engine=sqlite_engine, session_maker=session_maker, caller_session=caller_session) -def _build_tool() -> WorkflowTool: +def _persist_tenant(db: SqliteToolDb, *, tenant_id: str = TENANT_ID) -> Tenant: + tenant = Tenant(name="Tenant") + tenant.id = tenant_id + db.caller_session.add(tenant) + db.caller_session.commit() + return tenant + + +def _persist_account(db: SqliteToolDb, *, tenant_id: str = TENANT_ID) -> Account: + account = Account(name="Account", email="account@example.com") + account.id = ACCOUNT_ID + join = TenantAccountJoin( + tenant_id=tenant_id, + account_id=account.id, + current=True, + role=TenantAccountRole.NORMAL, + ) + db.caller_session.add_all([account, join]) + db.caller_session.commit() + return account + + +def _persist_end_user( + db: SqliteToolDb, + *, + end_user_id: str = END_USER_ID, + tenant_id: str = TENANT_ID, +) -> EndUser: + end_user = EndUser( + id=end_user_id, + tenant_id=tenant_id, + app_id=APP_ID, + type=EndUserType.SERVICE_API, + name="End user", + session_id="end-user-session", + ) + db.caller_session.add(end_user) + db.caller_session.commit() + return end_user + + +def _persist_app(db: SqliteToolDb) -> App: + app = App( + id=APP_ID, + tenant_id=TENANT_ID, + name="Workflow app", + description="", + mode=AppMode.WORKFLOW, + icon_type=None, + icon="", + icon_background=None, + app_model_config_id=None, + workflow_id=None, + enable_site=False, + enable_api=True, + max_active_requests=None, + created_by=CREATOR_ID, + ) + db.caller_session.add(app) + db.caller_session.commit() + return app + + +def _persist_workflow(db: SqliteToolDb, *, version: str, workflow_id: str | None = None) -> Workflow: + workflow = Workflow.new( + tenant_id=TENANT_ID, + app_id=APP_ID, + type=WorkflowType.WORKFLOW.value, + version=version, + graph=json.dumps({"nodes": [], "edges": []}), + features="{}", + created_by=CREATOR_ID, + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + workflow.id = workflow_id or str(uuid.uuid4()) + db.caller_session.add(workflow) + db.caller_session.commit() + return workflow + + +def _build_tool(*, tenant_id: str = "test_tool", workflow_app_id: str = "app-1", version: str = "1") -> WorkflowTool: entity = ToolEntity( identity=ToolIdentity(author="test", name="test tool", label=I18nObject(en_US="test tool"), provider="test"), parameters=[], description=None, has_runtime_parameters=False, ) - runtime = ToolRuntime(tenant_id="test_tool", invoke_from=InvokeFrom.EXPLORE) + runtime = ToolRuntime(tenant_id=tenant_id, invoke_from=InvokeFrom.EXPLORE) return WorkflowTool( - workflow_app_id="app-1", + workflow_app_id=workflow_app_id, workflow_as_tool_id="wf-tool-1", - version="1", + version=version, workflow_entities={}, workflow_call_depth=1, entity=entity, @@ -98,7 +166,10 @@ def _build_tool() -> WorkflowTool: ) -def test_workflow_tool_should_raise_tool_invoke_error_when_result_has_error_field(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_should_raise_tool_invoke_error_when_result_has_error_field( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Ensure that WorkflowTool will throw a `ToolInvokeError` exception when `WorkflowAppGenerator.generate` returns a result with `error` key inside the `data` element. @@ -122,11 +193,14 @@ def test_workflow_tool_should_raise_tool_invoke_error_when_result_has_error_fiel with pytest.raises(ToolInvokeError) as exc_info: # WorkflowTool always returns a generator, so we need to iterate to # actually `run` the tool. - list(tool.invoke(MagicMock(), "test_user", {})) + list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {})) assert exc_info.value.args == ("oops",) -def test_workflow_tool_does_not_use_pause_state_config(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_does_not_use_pause_state_config( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Ensure pause_state_config is passed as None.""" tool = _build_tool() @@ -140,14 +214,17 @@ def test_workflow_tool_does_not_use_pause_state_config(monkeypatch: pytest.Monke monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke(MagicMock(), "test_user", {})) + list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert "pause_state_config" in call_kwargs assert call_kwargs["pause_state_config"] is None -def test_workflow_tool_passes_parent_trace_context_from_runtime(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_passes_parent_trace_context_from_runtime( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Ensure nested workflow runtime metadata is forwarded as parent trace context.""" tool = _build_tool() tool.set_parent_trace_context( @@ -165,7 +242,7 @@ def test_workflow_tool_passes_parent_trace_context_from_runtime(monkeypatch: pyt monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke(MagicMock(), "test_user", {})) + list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert call_kwargs["args"]["parent_trace_context"].model_dump() == { @@ -174,7 +251,10 @@ def test_workflow_tool_passes_parent_trace_context_from_runtime(monkeypatch: pyt } -def test_workflow_tool_passes_parent_trace_session_id(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_passes_parent_trace_session_id( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Ensure nested workflows inherit the parent observability session ID.""" tool = _build_tool() tool.entity.parameters = [ @@ -197,14 +277,17 @@ def test_workflow_tool_passes_parent_trace_session_id(monkeypatch: pytest.Monkey monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke(MagicMock(), "test_user", {"trace_session_id": "user-input-session"})) + list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {"trace_session_id": "user-input-session"})) call_kwargs = generate_mock.call_args.kwargs assert call_kwargs["args"]["inputs"]["trace_session_id"] == "user-input-session" assert call_kwargs["args"]["trace_session_id"] == "session-1" -def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Ensure private trace context does not overwrite same-named workflow inputs.""" tool = _build_tool() tool.entity.parameters = [ @@ -238,7 +321,7 @@ def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(monkeypat list( tool.invoke( - MagicMock(), + sqlite_tool_db.caller_session, "test_user", { "outer_workflow_run_id": "user-workflow-input", @@ -256,7 +339,10 @@ def test_workflow_tool_keeps_user_inputs_named_like_trace_runtime_keys(monkeypat } -def test_workflow_tool_can_clear_parent_trace_context(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_can_clear_parent_trace_context( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Ensure reused WorkflowTool instances do not keep stale parent trace context.""" tool = _build_tool() tool.set_parent_trace_context( @@ -275,13 +361,16 @@ def test_workflow_tool_can_clear_parent_trace_context(monkeypatch: pytest.Monkey monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke(MagicMock(), "test_user", {})) + list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert "parent_trace_context" not in call_kwargs["args"] -def test_workflow_tool_can_clear_trace_session_id(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_can_clear_trace_session_id( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Ensure reused WorkflowTool instances do not keep stale trace session IDs.""" tool = _build_tool() tool.set_trace_session_id("session-1") @@ -297,7 +386,7 @@ def test_workflow_tool_can_clear_trace_session_id(monkeypatch: pytest.MonkeyPatc monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke(MagicMock(), "test_user", {})) + list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert "trace_session_id" not in call_kwargs["args"] @@ -315,6 +404,7 @@ def test_workflow_tool_can_clear_trace_session_id(monkeypatch: pytest.MonkeyPatc def test_workflow_tool_omits_parent_trace_context_when_runtime_is_incomplete( monkeypatch: pytest.MonkeyPatch, runtime_parameters: dict[str, Any], + sqlite_tool_db: SqliteToolDb, ): """Ensure incomplete runtime metadata does not leak parent trace context into generator args.""" tool = _build_tool() @@ -330,13 +420,16 @@ def test_workflow_tool_omits_parent_trace_context_when_runtime_is_incomplete( monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) - list(tool.invoke(MagicMock(), "test_user", {})) + list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {})) call_kwargs = generate_mock.call_args.kwargs assert "parent_trace_context" not in call_kwargs["args"] -def test_workflow_tool_should_generate_variable_messages_for_outputs(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_should_generate_variable_messages_for_outputs( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Test that WorkflowTool should generate variable messages when there are outputs""" tool = _build_tool() @@ -359,7 +452,7 @@ def test_workflow_tool_should_generate_variable_messages_for_outputs(monkeypatch monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) # Execute tool invocation - messages = list(tool.invoke(MagicMock(), "test_user", {})) + messages = list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {})) # Verify variable messages variable_messages = [msg for msg in messages if msg.type == ToolInvokeMessage.MessageType.VARIABLE] @@ -382,7 +475,10 @@ def test_workflow_tool_should_generate_variable_messages_for_outputs(monkeypatch assert json_messages[0].message.json_object == mock_outputs -def test_workflow_tool_should_handle_empty_outputs(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_should_handle_empty_outputs( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Test that WorkflowTool should handle empty outputs correctly""" tool = _build_tool() @@ -402,7 +498,7 @@ def test_workflow_tool_should_handle_empty_outputs(monkeypatch: pytest.MonkeyPat monkeypatch.setattr("libs.login.current_user", lambda *args, **kwargs: None) # Execute tool invocation - messages = list(tool.invoke(MagicMock(), "test_user", {})) + messages = list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {})) # Verify generated messages # Should contain: 0 variable messages + 1 text message + 1 JSON message = 2 messages @@ -458,41 +554,32 @@ def test_create_file_message_should_include_file_marker(): assert message.meta == {"file": file_obj} -def test_resolve_user_from_database_falls_back_to_end_user(monkeypatch: pytest.MonkeyPatch): +def test_resolve_user_from_database_falls_back_to_end_user(sqlite_tool_db: SqliteToolDb): """Ensure worker context can resolve EndUser when Account is missing.""" - - tenant = SimpleNamespace(id="tenant_id") - end_user = SimpleNamespace(id="end_user_id", tenant_id="tenant_id") - - # Monkeypatch session factory to return our stub session - stub_session = StubSession(scalar_results=[tenant, None, end_user]) - monkeypatch.setattr( - "core.tools.workflow_as_tool.tool.session_factory.create_session", - lambda: stub_session, + _persist_tenant(sqlite_tool_db) + end_user = _persist_end_user(sqlite_tool_db) + other_tenant_end_user = _persist_end_user( + sqlite_tool_db, + end_user_id="00000000-0000-0000-0000-000000000007", + tenant_id=OTHER_TENANT_ID, ) - tool = _build_tool() + tool = _build_tool(tenant_id=TENANT_ID) tool.runtime.invoke_from = InvokeFrom.SERVICE_API - tool.runtime.tenant_id = "tenant_id" resolved_user = tool._resolve_user_from_database(user_id=end_user.id) - assert resolved_user is end_user - assert stub_session.expunge_calls == [end_user] + assert isinstance(resolved_user, EndUser) + assert resolved_user.id == end_user.id + assert resolved_user.tenant_id == TENANT_ID + assert inspect(resolved_user).detached is True + assert tool._resolve_user_from_database(user_id=other_tenant_end_user.id) is None -def test_resolve_user_from_database_returns_none_when_no_tenant(monkeypatch: pytest.MonkeyPatch): +def test_resolve_user_from_database_returns_none_when_no_tenant(sqlite_tool_db: SqliteToolDb): """Return None if tenant cannot be found in worker context.""" - - # Monkeypatch session factory to return our stub session with no tenant - monkeypatch.setattr( - "core.tools.workflow_as_tool.tool.session_factory.create_session", - lambda: StubSession(scalar_results=[None]), - ) - - tool = _build_tool() + tool = _build_tool(tenant_id=OTHER_TENANT_ID) tool.runtime.invoke_from = InvokeFrom.SERVICE_API - tool.runtime.tenant_id = "missing_tenant" resolved_user = tool._resolve_user_from_database(user_id="any") @@ -544,7 +631,10 @@ def test_extract_usage_from_nested(): assert nested == {"total_tokens": 3} -def test_invoke_raises_when_user_not_found(monkeypatch: pytest.MonkeyPatch): +def test_invoke_raises_when_user_not_found( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Raise ToolInvokeError when user resolution fails.""" tool = _build_tool() monkeypatch.setattr(tool, "_get_app", lambda *args, **kwargs: None) @@ -552,58 +642,45 @@ def test_invoke_raises_when_user_not_found(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(tool, "_resolve_user", lambda *args, **kwargs: None) with pytest.raises(ToolInvokeError, match="User not found"): - list(tool.invoke(MagicMock(), "missing", {})) + list(tool.invoke(sqlite_tool_db.caller_session, "missing", {})) -def test_resolve_user_from_database_returns_account(monkeypatch: pytest.MonkeyPatch): +def test_resolve_user_from_database_returns_account(sqlite_tool_db: SqliteToolDb): """Resolve Account and set tenant in worker context.""" - tenant = SimpleNamespace(id="tenant_id") - account = SimpleNamespace(id="account_id", current_tenant=None) - set_current_tenant = Mock(side_effect=lambda tenant, *, session: setattr(account, "current_tenant", tenant)) - account.set_current_tenant_with_session = set_current_tenant - session = StubSession(scalar_results=[tenant, account]) + tenant = _persist_tenant(sqlite_tool_db) + account = _persist_account(sqlite_tool_db) + tool = _build_tool(tenant_id=TENANT_ID) - monkeypatch.setattr("core.tools.workflow_as_tool.tool.session_factory.create_session", lambda: session) - tool = _build_tool() - tool.runtime.tenant_id = "tenant_id" - - resolved = tool._resolve_user_from_database(user_id="account_id") - assert resolved is account - assert account.current_tenant is tenant - set_current_tenant.assert_called_once_with(tenant, session=session) - assert session.expunge_calls == [account] + resolved = tool._resolve_user_from_database(user_id=account.id) + assert isinstance(resolved, Account) + assert resolved.id == account.id + assert resolved.current_tenant_id == tenant.id + assert inspect(resolved).detached is True -def test_get_workflow_and_get_app_db_branches(monkeypatch: pytest.MonkeyPatch): +def test_get_workflow_and_get_app_db_branches(sqlite_tool_db: SqliteToolDb): """Cover workflow/app retrieval branches and error cases.""" - tool = _build_tool() - latest_workflow = SimpleNamespace(id="wf-latest") - specific_workflow = SimpleNamespace(id="wf-v1") - app = SimpleNamespace(id="app-1") - sessions = iter( - [ - StubSession(scalar_results=[], scalars_results=[latest_workflow]), - StubSession(scalar_results=[specific_workflow], scalars_results=[]), - StubSession(scalar_results=[app], scalars_results=[]), - ] - ) - monkeypatch.setattr( - "core.tools.workflow_as_tool.tool.session_factory.create_session", - lambda: next(sessions), - ) + app = _persist_app(sqlite_tool_db) + specific_workflow = _persist_workflow(sqlite_tool_db, version="1") + latest_workflow = _persist_workflow(sqlite_tool_db, version="2") + _persist_workflow(sqlite_tool_db, version=Workflow.VERSION_DRAFT) + tool = _build_tool(tenant_id=TENANT_ID, workflow_app_id=APP_ID) - assert tool._get_workflow("app-1", "") is latest_workflow - assert tool._get_workflow("app-1", "1") is specific_workflow - assert tool._get_app("app-1") is app + latest = tool._get_workflow(APP_ID, "") + specific = tool._get_workflow(APP_ID, "1") + resolved_app = tool._get_app(APP_ID) + + assert latest.id == latest_workflow.id + assert specific.id == specific_workflow.id + assert resolved_app.id == app.id + assert inspect(latest).detached is True + assert inspect(specific).detached is True + assert inspect(resolved_app).detached is True - monkeypatch.setattr( - "core.tools.workflow_as_tool.tool.session_factory.create_session", - lambda: StubSession(scalar_results=[None, None], scalars_results=[None]), - ) with pytest.raises(ValueError, match="workflow not found"): - tool._get_workflow("app-1", "1") + tool._get_workflow(APP_ID, "missing") with pytest.raises(ValueError, match="app not found"): - tool._get_app("app-1") + tool._get_app("00000000-0000-0000-0000-000000000099") def _setup_transform_args_tool(monkeypatch: pytest.MonkeyPatch) -> WorkflowTool: @@ -722,7 +799,10 @@ def test_transform_args_normalizes_optional_files_parameter( assert files == [] -def test_workflow_tool_invocation_normalizes_optional_files_parameter(monkeypatch: pytest.MonkeyPatch): +def test_workflow_tool_invocation_normalizes_optional_files_parameter( + monkeypatch: pytest.MonkeyPatch, + sqlite_tool_db: SqliteToolDb, +): """Ensure casted empty FILES values do not reach workflow input validation as [None].""" tool = _build_tool() images_param = ToolParameter.get_simple_instance( @@ -741,7 +821,7 @@ def test_workflow_tool_invocation_normalizes_optional_files_parameter(monkeypatc generate_mock = MagicMock(return_value={"data": {}}) monkeypatch.setattr("core.app.apps.workflow.app_generator.WorkflowAppGenerator.generate", generate_mock) - list(tool.invoke(MagicMock(), "test_user", {"images": None})) + list(tool.invoke(sqlite_tool_db.caller_session, "test_user", {"images": None})) call_kwargs = generate_mock.call_args.kwargs assert call_kwargs["args"]["inputs"]["images"] == [] diff --git a/api/tests/unit_tests/core/variables/test_segment.py b/api/tests/unit_tests/core/variables/test_segment.py index 9e07ea1b6db..f65e5bbde75 100644 --- a/api/tests/unit_tests/core/variables/test_segment.py +++ b/api/tests/unit_tests/core/variables/test_segment.py @@ -28,6 +28,7 @@ from graphon.variables.segments import ( StringSegment, get_segment_discriminator, ) +from graphon.variables.template_resolution import convert_template from graphon.variables.types import SegmentType from graphon.variables.utils import ( dumps_with_segments, @@ -98,7 +99,7 @@ def test_segment_group_to_text(): template = ( "Hello, {{#sys.user_id#}}! Your query is {{#node_id.custom_query#}}. And your key is {{#env.secret_key#}}." ) - segments_group = variable_pool.convert_template(template) + segments_group = convert_template(variable_pool, template) assert segments_group.text == "Hello, fake-user-id! Your query is fake-user-query. And your key is fake-secret-key." assert segments_group.log == ( @@ -112,7 +113,7 @@ def test_convert_constant_to_segment_group(): system_variables=build_system_variables(user_id="1", app_id="1", workflow_id="1"), ) template = "Hello, world!" - segments_group = variable_pool.convert_template(template) + segments_group = convert_template(variable_pool, template) assert segments_group.text == "Hello, world!" assert segments_group.log == "Hello, world!" @@ -120,7 +121,7 @@ def test_convert_constant_to_segment_group(): def test_convert_variable_to_segment_group(): variable_pool = _build_variable_pool(system_variables=build_system_variables(user_id="fake-user-id")) template = "{{#sys.user_id#}}" - segments_group = variable_pool.convert_template(template) + segments_group = convert_template(variable_pool, template) assert segments_group.text == "fake-user-id" assert segments_group.log == "fake-user-id" assert isinstance(segments_group.value[0], StringVariable) diff --git a/api/tests/unit_tests/core/workflow/entities/test_private_workflow_pause.py b/api/tests/unit_tests/core/workflow/entities/test_private_workflow_pause.py index 3f476103127..5dde9985d78 100644 --- a/api/tests/unit_tests/core/workflow/entities/test_private_workflow_pause.py +++ b/api/tests/unit_tests/core/workflow/entities/test_private_workflow_pause.py @@ -1,45 +1,56 @@ """Tests for _PrivateWorkflowPauseEntity implementation.""" from datetime import datetime -from unittest.mock import MagicMock, patch +from unittest.mock import patch from models.workflow import WorkflowPause as WorkflowPauseModel from repositories.sqlalchemy_api_workflow_run_repository import _PrivateWorkflowPauseEntity +def _make_workflow_pause( + *, + pause_id: str = "pause-123", + workflow_run_id: str = "execution-456", + state_object_key: str = "test-state-key", + resumed_at: datetime | None = None, +) -> WorkflowPauseModel: + pause = WorkflowPauseModel( + workflow_id="workflow-789", + workflow_run_id=workflow_run_id, + state_object_key=state_object_key, + resumed_at=resumed_at, + ) + pause.id = pause_id + return pause + + class TestPrivateWorkflowPauseEntity: """Test _PrivateWorkflowPauseEntity implementation.""" def test_entity_initialization(self): """Test entity initialization with required parameters.""" - # Create mock models - mock_pause_model = MagicMock(spec=WorkflowPauseModel) - mock_pause_model.id = "pause-123" - mock_pause_model.workflow_run_id = "execution-456" - mock_pause_model.resumed_at = None + pause_model = _make_workflow_pause() # Create entity - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) # Verify initialization - assert entity._pause_model is mock_pause_model + assert entity._pause_model is pause_model assert entity._cached_state is None def test_id_property(self): """Test id property returns pause model ID.""" - mock_pause_model = MagicMock(spec=WorkflowPauseModel) - mock_pause_model.id = "pause-123" + pause_model = _make_workflow_pause() - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) assert entity.id == "pause-123" def test_workflow_execution_id_property(self): """Test workflow_execution_id property returns workflow run ID.""" - mock_pause_model = MagicMock(spec=WorkflowPauseModel) - mock_pause_model.workflow_run_id = "execution-456" + pause_model = _make_workflow_pause() - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) assert entity.workflow_execution_id == "execution-456" @@ -47,19 +58,17 @@ class TestPrivateWorkflowPauseEntity: """Test resumed_at property returns pause model resumed_at.""" resumed_at = datetime(2023, 12, 25, 15, 30, 45) - mock_pause_model = MagicMock(spec=WorkflowPauseModel) - mock_pause_model.resumed_at = resumed_at + pause_model = _make_workflow_pause(resumed_at=resumed_at) - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) assert entity.resumed_at == resumed_at def test_resumed_at_property_none(self): """Test resumed_at property returns None when not set.""" - mock_pause_model = MagicMock(spec=WorkflowPauseModel) - mock_pause_model.resumed_at = None + pause_model = _make_workflow_pause() - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) assert entity.resumed_at is None @@ -69,10 +78,9 @@ class TestPrivateWorkflowPauseEntity: state_data = b'{"test": "data", "step": 5}' mock_storage.load.return_value = state_data - mock_pause_model = MagicMock(spec=WorkflowPauseModel) - mock_pause_model.state_object_key = "test-state-key" + pause_model = _make_workflow_pause() - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) # First call should load from storage result = entity.get_state() @@ -87,10 +95,9 @@ class TestPrivateWorkflowPauseEntity: state_data = b'{"test": "data", "step": 5}' mock_storage.load.return_value = state_data - mock_pause_model = MagicMock(spec=WorkflowPauseModel) - mock_pause_model.state_object_key = "test-state-key" + pause_model = _make_workflow_pause() - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) # First call result1 = entity.get_state() @@ -107,9 +114,9 @@ class TestPrivateWorkflowPauseEntity: """Test get_state returns pre-cached data.""" state_data = b'{"test": "data", "step": 5}' - mock_pause_model = MagicMock(spec=WorkflowPauseModel) + pause_model = _make_workflow_pause() - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) # Pre-cache data entity._cached_state = state_data @@ -128,9 +135,9 @@ class TestPrivateWorkflowPauseEntity: with patch("repositories.sqlalchemy_api_workflow_run_repository.storage", autospec=True) as mock_storage: mock_storage.load.return_value = binary_data - mock_pause_model = MagicMock(spec=WorkflowPauseModel) + pause_model = _make_workflow_pause() - entity = _PrivateWorkflowPauseEntity(pause_model=mock_pause_model, reason_models=[], human_input_form=[]) + entity = _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) result = entity.get_state() diff --git a/api/tests/unit_tests/core/workflow/graph_engine/layers/test_llm_quota.py b/api/tests/unit_tests/core/workflow/graph_engine/layers/test_llm_quota.py deleted file mode 100644 index 97d7e4a937b..00000000000 --- a/api/tests/unit_tests/core/workflow/graph_engine/layers/test_llm_quota.py +++ /dev/null @@ -1,353 +0,0 @@ -import logging -import threading -from datetime import datetime -from types import SimpleNamespace -from unittest.mock import MagicMock, patch - -import pytest - -from core.app.workflow.layers.llm_quota import LLMQuotaLayer -from core.errors.error import QuotaExceededError -from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus -from graphon.graph_engine.entities.commands import CommandType -from graphon.graph_events import NodeRunSucceededEvent -from graphon.model_runtime.entities.llm_entities import LLMUsage -from graphon.node_events import NodeRunResult - - -def _build_succeeded_event(*, provider: str = "openai", model_name: str = "gpt-4o") -> NodeRunSucceededEvent: - return NodeRunSucceededEvent( - id="execution-id", - node_id="llm-node-id", - node_type=BuiltinNodeTypes.LLM, - start_at=datetime.now(), - node_run_result=NodeRunResult( - status=WorkflowNodeExecutionStatus.SUCCEEDED, - inputs={ - "question": "hello", - "model_provider": provider, - "model_name": model_name, - }, - llm_usage=LLMUsage.empty_usage(), - ), - ) - - -def _build_public_model_identity(*, provider: str = "openai", model_name: str = "gpt-4o") -> SimpleNamespace: - return SimpleNamespace(provider=provider, name=model_name) - - -def _build_node_data(*, model: SimpleNamespace | None = None) -> SimpleNamespace: - return SimpleNamespace( - error_strategy=None, - retry_config=SimpleNamespace(retry_enabled=False), - model=model, - ) - - -def _build_node(*, node_type: BuiltinNodeTypes = BuiltinNodeTypes.LLM) -> MagicMock: - node = MagicMock() - node.id = "node-id" - node.execution_id = "execution-id" - node.node_type = node_type - node.node_data = _build_node_data(model=_build_public_model_identity()) - node.model_instance = SimpleNamespace(provider="stale-provider", model_name="stale-model") - return node - - -class _RunnableQuotaNode: - id = "node-id" - execution_id = "execution-id" - node_type = BuiltinNodeTypes.LLM - title = "LLM node" - - def __init__(self, *, stop_event: threading.Event, node_data: SimpleNamespace | None = None) -> None: - self.node_data = node_data or _build_node_data(model=_build_public_model_identity()) - self.graph_runtime_state = SimpleNamespace(stop_event=stop_event) - self.original_run_called = False - - def _run(self) -> NodeRunResult: - self.original_run_called = True - return NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED) - - -def test_deduct_quota_called_for_successful_llm_node() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = _build_node(node_type=BuiltinNodeTypes.LLM) - result_event = _build_succeeded_event() - - with patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", autospec=True) as mock_deduct: - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - mock_deduct.assert_called_once_with( - tenant_id="tenant-id", - provider="openai", - model="gpt-4o", - usage=result_event.node_run_result.llm_usage, - ) - - -def test_deduct_quota_called_for_question_classifier_node() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = _build_node(node_type=BuiltinNodeTypes.QUESTION_CLASSIFIER) - result_event = _build_succeeded_event(provider="anthropic", model_name="claude-3-7-sonnet") - - with patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", autospec=True) as mock_deduct: - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - mock_deduct.assert_called_once_with( - tenant_id="tenant-id", - provider="anthropic", - model="claude-3-7-sonnet", - usage=result_event.node_run_result.llm_usage, - ) - - -def test_non_llm_node_is_ignored() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = _build_node(node_type=BuiltinNodeTypes.START) - result_event = _build_succeeded_event() - - with patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", autospec=True) as mock_deduct: - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - mock_deduct.assert_not_called() - - -def test_precheck_ignores_non_quota_node() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = _build_node(node_type=BuiltinNodeTypes.START) - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - mock_check.assert_not_called() - - -def test_quota_error_is_handled_in_layer(caplog: pytest.LogCaptureFixture) -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - result_event = _build_succeeded_event() - - with ( - caplog.at_level(logging.ERROR, logger="core.app.workflow.layers.llm_quota"), - patch( - "core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", - autospec=True, - side_effect=ValueError("quota exceeded"), - ) as mock_deduct, - ): - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - mock_deduct.assert_called_once_with( - tenant_id="tenant-id", - provider="openai", - model="gpt-4o", - usage=result_event.node_run_result.llm_usage, - ) - assert "LLM quota deduction failed, node_id=node-id" in caplog.text - assert not stop_event.is_set() - layer.command_channel.send_command.assert_not_called() - - -def test_send_abort_command_is_noop_without_channel_or_after_abort() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - - layer._send_abort_command(reason="no channel") - - layer.command_channel = MagicMock() - layer._abort_sent = True - layer._send_abort_command(reason="already aborted") - - layer.command_channel.send_command.assert_not_called() - - -def test_quota_deduction_exceeded_aborts_workflow_immediately() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - result_event = _build_succeeded_event() - with patch( - "core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", - autospec=True, - side_effect=QuotaExceededError("No credits remaining"), - ): - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - assert stop_event.is_set() - layer.command_channel.send_command.assert_called_once() - abort_command = layer.command_channel.send_command.call_args.args[0] - assert abort_command.command_type == CommandType.ABORT - assert abort_command.reason == "No credits remaining" - - -def test_quota_precheck_failure_aborts_workflow_immediately() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - with patch( - "core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", - autospec=True, - side_effect=QuotaExceededError("Model provider openai quota exceeded."), - ): - layer.on_node_run_start(node) - - assert stop_event.is_set() - layer.command_channel.send_command.assert_called_once() - abort_command = layer.command_channel.send_command.call_args.args[0] - assert abort_command.command_type == CommandType.ABORT - assert abort_command.reason == "Model provider openai quota exceeded." - - -def test_quota_precheck_failure_blocks_current_node_run() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _RunnableQuotaNode(stop_event=stop_event) - - with patch( - "core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", - autospec=True, - side_effect=QuotaExceededError("Model provider openai quota exceeded."), - ): - layer.on_node_run_start(node) - - result = node._run() - assert not node.original_run_called - assert result.status == WorkflowNodeExecutionStatus.FAILED - assert result.error == "Model provider openai quota exceeded." - assert result.error_type == QuotaExceededError.__name__ - - -def test_missing_model_identity_blocks_current_node_run() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _RunnableQuotaNode(stop_event=stop_event, node_data=_build_node_data()) - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - result = node._run() - assert not node.original_run_called - assert result.status == WorkflowNodeExecutionStatus.FAILED - assert result.error == "LLM quota check requires public node model identity before execution." - assert result.error_type == "LLMQuotaIdentityError" - mock_check.assert_not_called() - - -def test_quota_precheck_passes_without_abort() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - assert not stop_event.is_set() - mock_check.assert_called_once_with( - tenant_id="tenant-id", - provider="openai", - model="gpt-4o", - ) - layer.command_channel.send_command.assert_not_called() - - -def test_precheck_reads_model_identity_from_data_when_node_data_is_absent() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = SimpleNamespace( - id="node-id", - node_type=BuiltinNodeTypes.LLM, - data=_build_node_data(model=_build_public_model_identity(provider="anthropic", model_name="claude")), - ) - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - mock_check.assert_called_once_with( - tenant_id="tenant-id", - provider="anthropic", - model="claude", - ) - - -def test_precheck_rejects_invalid_public_model_identity() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.node_data = _build_node_data(model=_build_public_model_identity(provider="", model_name="gpt-4o")) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - assert stop_event.is_set() - mock_check.assert_not_called() - layer.command_channel.send_command.assert_called_once() - - -def test_precheck_requires_public_node_model_config() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.node_data = _build_node_data() - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - assert stop_event.is_set() - mock_check.assert_not_called() - layer.command_channel.send_command.assert_called_once() - abort_command = layer.command_channel.send_command.call_args.args[0] - assert abort_command.command_type == CommandType.ABORT - assert abort_command.reason == "LLM quota check requires public node model identity before execution." - - -def test_deduction_requires_public_event_model_identity() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - result_event = _build_succeeded_event() - result_event.node_run_result.inputs = {"question": "hello"} - - with patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", autospec=True) as mock_deduct: - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - assert stop_event.is_set() - mock_deduct.assert_not_called() - layer.command_channel.send_command.assert_called_once() - abort_command = layer.command_channel.send_command.call_args.args[0] - assert abort_command.command_type == CommandType.ABORT - assert abort_command.reason == "LLM quota deduction requires model identity in the node result event." diff --git a/api/tests/unit_tests/core/workflow/graph_engine/test_mock_factory.py b/api/tests/unit_tests/core/workflow/graph_engine/test_mock_factory.py index c721c7b0ebd..9c507bb1619 100644 --- a/api/tests/unit_tests/core/workflow/graph_engine/test_mock_factory.py +++ b/api/tests/unit_tests/core/workflow/graph_engine/test_mock_factory.py @@ -5,7 +5,7 @@ The factory follows the same config adaptation path as production implementations before instantiation. """ -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, override from core.workflow.human_input_adapter import adapt_node_config_for_graph from core.workflow.node_factory import DifyNodeFactory @@ -76,6 +76,14 @@ class MockNodeFactory(DifyNodeFactory): BuiltinNodeTypes.CODE: MockCodeNode, } + @override + def with_runtime_state(self, graph_runtime_state: "GraphRuntimeState") -> "MockNodeFactory": + return MockNodeFactory( + graph_init_params=self.graph_init_params, + graph_runtime_state=graph_runtime_state, + mock_config=self.mock_config, + ) + def create_node(self, node_config: dict[str, Any] | NodeConfigDict) -> Node: """ Create a node instance, using mock implementations for third-party service nodes. diff --git a/api/tests/unit_tests/core/workflow/graph_engine/test_mock_nodes.py b/api/tests/unit_tests/core/workflow/graph_engine/test_mock_nodes.py index 55dcbdb7a11..67b94e46399 100644 --- a/api/tests/unit_tests/core/workflow/graph_engine/test_mock_nodes.py +++ b/api/tests/unit_tests/core/workflow/graph_engine/test_mock_nodes.py @@ -615,69 +615,6 @@ class MockIterationNode(MockNodeMixin, IterationNode): """Return the version of this mock node.""" return "1" - def _create_graph_engine(self, index: int, item: Any): - """Create a graph engine with MockNodeFactory instead of DifyNodeFactory.""" - # Import dependencies - from graphon.entities import GraphInitParams - from graphon.graph import Graph - from graphon.graph_engine import GraphEngine, GraphEngineConfig - from graphon.graph_engine.command_channels import InMemoryChannel - from graphon.runtime import GraphRuntimeState - - # Import our MockNodeFactory instead of DifyNodeFactory - from .test_mock_factory import MockNodeFactory - - # Create GraphInitParams from node attributes - graph_init_params = GraphInitParams( - workflow_id=self.workflow_id, - graph_config=self.graph_config, - run_context=self.run_context, - call_depth=self.workflow_call_depth, - ) - - # Create a deep copy of the variable pool for each iteration - variable_pool_copy = self.graph_runtime_state.variable_pool.model_copy(deep=True) - - # append iteration variable (item, index) to variable pool - variable_pool_copy.add([self._node_id, "index"], index) - variable_pool_copy.add([self._node_id, "item"], item) - - # Create a new GraphRuntimeState for this iteration - graph_runtime_state_copy = GraphRuntimeState( - variable_pool=variable_pool_copy, - start_at=self.graph_runtime_state.start_at, - total_tokens=0, - node_run_steps=0, - ) - - # Create a MockNodeFactory with the same mock_config - node_factory = MockNodeFactory( - graph_init_params=graph_init_params, - graph_runtime_state=graph_runtime_state_copy, - mock_config=self.mock_config, # Pass the mock configuration - ) - - # Initialize the iteration graph with the mock node factory - iteration_graph = Graph.init( - graph_config=self.graph_config, node_factory=node_factory, root_node_id=self._node_data.start_node_id - ) - - if not iteration_graph: - from graphon.nodes.iteration.exc import IterationGraphNotFoundError - - raise IterationGraphNotFoundError("iteration graph not found") - - # Create a new GraphEngine for this iteration - graph_engine = GraphEngine( - workflow_id=self.workflow_id, - graph=iteration_graph, - graph_runtime_state=graph_runtime_state_copy, - command_channel=InMemoryChannel(), # Use InMemoryChannel for sub-graphs - config=GraphEngineConfig(), - ) - - return graph_engine - class MockLoopNode(MockNodeMixin, LoopNode): """Mock implementation of LoopNode that preserves mock configuration.""" @@ -687,56 +624,6 @@ class MockLoopNode(MockNodeMixin, LoopNode): """Return the version of this mock node.""" return "1" - def _create_graph_engine(self, start_at, root_node_id: str): - """Create a graph engine with MockNodeFactory instead of DifyNodeFactory.""" - # Import dependencies - from graphon.entities import GraphInitParams - from graphon.graph import Graph - from graphon.graph_engine import GraphEngine, GraphEngineConfig - from graphon.graph_engine.command_channels import InMemoryChannel - from graphon.runtime import GraphRuntimeState - - # Import our MockNodeFactory instead of DifyNodeFactory - from .test_mock_factory import MockNodeFactory - - # Create GraphInitParams from node attributes - graph_init_params = GraphInitParams( - workflow_id=self.workflow_id, - graph_config=self.graph_config, - run_context=self.run_context, - call_depth=self.workflow_call_depth, - ) - - # Create a new GraphRuntimeState for this iteration - graph_runtime_state_copy = GraphRuntimeState( - variable_pool=self.graph_runtime_state.variable_pool, - start_at=start_at.timestamp(), - ) - - # Create a MockNodeFactory with the same mock_config - node_factory = MockNodeFactory( - graph_init_params=graph_init_params, - graph_runtime_state=graph_runtime_state_copy, - mock_config=self.mock_config, # Pass the mock configuration - ) - - # Initialize the loop graph with the mock node factory - loop_graph = Graph.init(graph_config=self.graph_config, node_factory=node_factory, root_node_id=root_node_id) - - if not loop_graph: - raise ValueError("loop graph not found") - - # Create a new GraphEngine for this iteration - graph_engine = GraphEngine( - workflow_id=self.workflow_id, - graph=loop_graph, - graph_runtime_state=graph_runtime_state_copy, - command_channel=InMemoryChannel(), # Use InMemoryChannel for sub-graphs - config=GraphEngineConfig(), - ) - - return graph_engine - class MockTemplateTransformNode(MockNodeMixin, TemplateTransformNode): """Mock implementation of TemplateTransformNode for testing.""" diff --git a/api/tests/unit_tests/core/workflow/graph_engine/test_table_runner.py b/api/tests/unit_tests/core/workflow/graph_engine/test_table_runner.py index e9c9e04e17b..175b5da8a1b 100644 --- a/api/tests/unit_tests/core/workflow/graph_engine/test_table_runner.py +++ b/api/tests/unit_tests/core/workflow/graph_engine/test_table_runner.py @@ -51,53 +51,6 @@ from .test_mock_factory import MockNodeFactory logger = logging.getLogger(__name__) -class _TableTestChildEngineBuilder: - def __init__(self, *, use_mock_factory: bool, mock_config: MockConfig | None) -> None: - self._use_mock_factory = use_mock_factory - self._mock_config = mock_config - - def build_child_engine( - self, - *, - workflow_id: str, - graph_init_params: GraphInitParams, - parent_graph_runtime_state: GraphRuntimeState, - root_node_id: str, - variable_pool: VariablePool | None = None, - ) -> GraphEngine: - child_graph_runtime_state = GraphRuntimeState( - variable_pool=variable_pool if variable_pool is not None else parent_graph_runtime_state.variable_pool, - start_at=time.perf_counter(), - execution_context=parent_graph_runtime_state.execution_context, - ) - if self._use_mock_factory: - node_factory = MockNodeFactory( - graph_init_params=graph_init_params, - graph_runtime_state=child_graph_runtime_state, - mock_config=self._mock_config, - ) - else: - node_factory = DifyNodeFactory( - graph_init_params=graph_init_params, - graph_runtime_state=child_graph_runtime_state, - ) - - graph_config = graph_init_params.graph_config - child_graph = Graph.init(graph_config=graph_config, node_factory=node_factory, root_node_id=root_node_id) - if not child_graph: - raise ValueError("child graph not found") - - child_engine = GraphEngine( - workflow_id=workflow_id, - graph=child_graph, - graph_runtime_state=child_graph_runtime_state, - command_channel=InMemoryChannel(), - config=GraphEngineConfig(), - child_engine_builder=self, - ) - return child_engine - - @dataclass class WorkflowTestCase: """Represents a single test case for table-driven testing.""" @@ -379,10 +332,6 @@ class TableTestRunner: scale_up_threshold=self.graph_engine_scale_up_threshold, scale_down_idle_time=self.graph_engine_scale_down_idle_time, ), - child_engine_builder=_TableTestChildEngineBuilder( - use_mock_factory=test_case.use_auto_mock, - mock_config=test_case.mock_config, - ), ) # Execute and collect events diff --git a/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_entities.py b/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_entities.py new file mode 100644 index 00000000000..023fccbefd3 --- /dev/null +++ b/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_entities.py @@ -0,0 +1,58 @@ +from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom, build_dify_run_context +from core.workflow.node_factory import DifyNodeFactory +from core.workflow.nodes.agent.entities import AgentNodeData +from graphon.entities import GraphInitParams +from graphon.graph import Graph +from graphon.graph_engine import GraphEngine, GraphEngineConfig +from graphon.graph_engine.command_channels import InMemoryChannel +from graphon.graph_events import GraphNodeEventBase, GraphRunSucceededEvent +from graphon.runtime import GraphRuntimeState, VariablePool + + +def test_agent_node_data_unconfigured_defaults() -> None: + data = AgentNodeData.model_validate({"title": "Agent"}) + + assert data.agent_strategy_provider_name == "" + assert data.agent_strategy_name == "" + assert data.agent_strategy_label == "" + assert not data.agent_parameters + + +def test_unconfigured_disconnected_agent_does_not_block_workflow() -> None: + graph_config: dict[str, object] = { + "nodes": [ + {"id": "start", "data": {"type": "start", "title": "Start", "variables": []}}, + {"id": "agent", "data": {"type": "agent", "title": "Agent", "tool_node_version": "2"}}, + ], + "edges": [], + } + graph_runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0) + graph_init_params = GraphInitParams( + workflow_id="workflow", + graph_config=graph_config, + run_context=build_dify_run_context( + tenant_id="tenant", + app_id="app", + user_id="user", + user_from=UserFrom.ACCOUNT, + invoke_from=InvokeFrom.DEBUGGER, + ), + call_depth=0, + ) + graph = Graph.init( + graph_config=graph_config, + node_factory=DifyNodeFactory(graph_init_params, graph_runtime_state), + root_node_id="start", + ) + engine = GraphEngine( + workflow_id="workflow", + graph=graph, + graph_runtime_state=graph_runtime_state, + command_channel=InMemoryChannel(), + config=GraphEngineConfig(min_workers=1, max_workers=1), + ) + + events = list(engine.run()) + + assert isinstance(events[-1], GraphRunSucceededEvent) + assert not any(isinstance(event, GraphNodeEventBase) and event.node_id == "agent" for event in events) diff --git a/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_node.py b/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_node.py new file mode 100644 index 00000000000..b1596fecd1d --- /dev/null +++ b/api/tests/unit_tests/core/workflow/nodes/agent/test_agent_node.py @@ -0,0 +1,61 @@ +from unittest.mock import MagicMock + +from core.workflow.nodes.agent.agent_node import AgentNode +from core.workflow.nodes.agent.entities import AgentNodeData +from core.workflow.nodes.agent.events import AgentLogEvent, NodeRunAgentLogEvent +from graphon.entities import GraphInitParams +from graphon.enums import BuiltinNodeTypes +from graphon.graph_events import NodeRunStreamChunkEvent +from graphon.node_events import StreamChunkEvent +from graphon.runtime import GraphRuntimeState, VariablePool + + +def test_dispatch_converts_agent_events_and_delegates_other_events() -> None: + node = AgentNode( + node_id="node-id", + data=AgentNodeData(title="Agent"), + graph_init_params=GraphInitParams( + workflow_id="workflow-id", + graph_config={}, + run_context={}, + call_depth=0, + ), + graph_runtime_state=GraphRuntimeState(variable_pool=VariablePool(), start_at=0), + strategy_resolver=MagicMock(), + presentation_provider=MagicMock(), + runtime_support=MagicMock(), + message_transformer=MagicMock(), + ) + node._node_execution_id = "execution-id" + + graph_event = node._dispatch( + AgentLogEvent( + message_id="message-id", + label="label", + node_execution_id="agent-execution-id", + parent_id="parent-id", + error=None, + status="succeeded", + data={"output": "done"}, + metadata={"provider": "test"}, + node_id="source-node-id", + ) + ) + + assert graph_event == NodeRunAgentLogEvent( + id="execution-id", + node_id="node-id", + node_type=BuiltinNodeTypes.AGENT, + message_id="message-id", + label="label", + node_execution_id="agent-execution-id", + parent_id="parent-id", + error=None, + status="succeeded", + data={"output": "done"}, + metadata={"provider": "test"}, + ) + assert isinstance( + node._dispatch(StreamChunkEvent(selector=["node-id", "text"], chunk="hello")), + NodeRunStreamChunkEvent, + ) diff --git a/api/tests/unit_tests/core/workflow/nodes/agent/test_runtime_support.py b/api/tests/unit_tests/core/workflow/nodes/agent/test_runtime_support.py index c86de7f6e63..f79f4282649 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent/test_runtime_support.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent/test_runtime_support.py @@ -2,7 +2,16 @@ from types import SimpleNamespace from unittest.mock import Mock, patch from core.workflow.nodes.agent.runtime_support import AgentRuntimeSupport -from graphon.model_runtime.entities.model_entities import ModelType +from graphon.model_runtime.entities.common_entities import I18nObject +from graphon.model_runtime.entities.model_entities import ( + AIModelEntity, + FetchFrom, + ModelFeature, + ModelPropertyKey, + ModelType, + ParameterRule, + ParameterType, +) def test_fetch_model_reuses_single_model_assembly(): @@ -47,3 +56,98 @@ def test_fetch_model_reuses_single_model_assembly(): model_type=ModelType.LLM, model="gpt-4o-mini", ) + + +def _make_model_schema_with_defaults() -> AIModelEntity: + """Return a minimal AIModelEntity whose parameter_rules carry defaults.""" + return AIModelEntity( + model="qwen-max", + label=I18nObject(en_US="Qwen Max"), + model_type=ModelType.LLM, + features=[ModelFeature.AGENT_THOUGHT, ModelFeature.MULTI_TOOL_CALL], + fetch_from=FetchFrom.PREDEFINED_MODEL, + model_properties={ + ModelPropertyKey.MODE: "chat", + ModelPropertyKey.CONTEXT_SIZE: 32768, + }, + parameter_rules=[ + ParameterRule( + name="temperature", + use_template="temperature", + label=I18nObject(en_US="Temperature"), + type=ParameterType.FLOAT, + required=False, + default=0.7, + min=0.0, + max=2.0, + precision=2, + ), + ParameterRule( + name="max_tokens", + use_template="max_tokens", + label=I18nObject(en_US="Max Tokens"), + type=ParameterType.INT, + required=False, + default=2048, + min=1, + max=32768, + ), + ParameterRule( + name="top_p", + use_template="top_p", + label=I18nObject(en_US="Top P"), + type=ParameterType.FLOAT, + required=False, + default=1.0, + ), + ], + ) + + +def test_extract_default_completion_params_collects_rule_defaults(): + """_extract_default_completion_params should gather every rule.default.""" + schema = _make_model_schema_with_defaults() + params = AgentRuntimeSupport._extract_default_completion_params(schema) + assert params == {"temperature": 0.7, "max_tokens": 2048, "top_p": 1.0} + + +def test_extract_default_completion_params_skips_rules_without_default(): + """Rules whose default is None must not appear in the result.""" + schema = AIModelEntity( + model="test-model", + label=I18nObject(en_US="Test"), + model_type=ModelType.LLM, + fetch_from=FetchFrom.PREDEFINED_MODEL, + model_properties={ModelPropertyKey.MODE: "chat"}, + parameter_rules=[ + ParameterRule( + name="seed", + label=I18nObject(en_US="Seed"), + type=ParameterType.INT, + required=False, + default=None, + ), + ParameterRule( + name="temperature", + label=I18nObject(en_US="Temperature"), + type=ParameterType.FLOAT, + required=False, + default=0.5, + ), + ], + ) + params = AgentRuntimeSupport._extract_default_completion_params(schema) + assert params == {"temperature": 0.5} + + +def test_extract_default_completion_params_empty_when_no_defaults(): + """An empty parameter_rules list yields an empty dict.""" + schema = AIModelEntity( + model="test-model", + label=I18nObject(en_US="Test"), + model_type=ModelType.LLM, + fetch_from=FetchFrom.PREDEFINED_MODEL, + model_properties={ModelPropertyKey.MODE: "chat"}, + parameter_rules=[], + ) + assert AgentRuntimeSupport._extract_default_completion_params(schema) == {} diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_agent_node.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_agent_node.py index d142591c3fb..1ee031d9462 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_agent_node.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_agent_node.py @@ -3,7 +3,9 @@ from datetime import UTC, datetime from types import SimpleNamespace from typing import cast from unittest.mock import MagicMock, patch +from uuid import UUID, uuid4 +import pytest from agenton.compositor import CompositorSessionSnapshot from dify_agent.layers.ask_human import AskHumanToolResult from dify_agent.protocol import ( @@ -24,7 +26,6 @@ from clients.agent_backend import ( AgentBackendStreamInternalEvent, FakeAgentBackendRunClient, FakeAgentBackendScenario, - RuntimeLayerSpec, ) from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext, InvokeFrom, UserFrom from core.workflow.file_reference import build_file_reference @@ -33,18 +34,22 @@ from core.workflow.nodes.agent_v2.ask_human_resume import AskHumanResumeOutcome from core.workflow.nodes.agent_v2.binding_resolver import WorkflowAgentBindingBundle, WorkflowAgentBindingResolver from core.workflow.nodes.agent_v2.entities import DifyAgentNodeData from core.workflow.nodes.agent_v2.output_adapter import WorkflowAgentOutputAdapter -from core.workflow.nodes.agent_v2.runtime_request_builder import WorkflowAgentRuntimeRequestBuilder +from core.workflow.nodes.agent_v2.runtime_request_builder import ( + WorkflowAgentRuntimeBuildContext, + WorkflowAgentRuntimeRequestBuilder, +) from core.workflow.nodes.agent_v2.session_store import ( StoredWorkflowAgentSession, - WorkflowAgentRuntimeSessionStore, WorkflowAgentSessionScope, + WorkflowAgentWorkspaceStore, ) from core.workflow.nodes.human_input.pause_reason import HumanInputRequired from graphon.entities import GraphInitParams from graphon.entities.pause_reason import HitlRequired from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus from graphon.file import File, FileTransferMethod, FileType -from graphon.node_events import PauseRequestedEvent, StreamCompletedEvent +from graphon.graph_events import NodeRunPauseRequestedEvent +from graphon.node_events import StreamCompletedEvent from graphon.runtime import GraphRuntimeState from graphon.variables.segments import ArrayFileSegment, FileSegment, StringSegment from models.agent import Agent, AgentConfigSnapshot, WorkflowAgentNodeBinding @@ -54,6 +59,7 @@ from models.agent_config_entities import ( DeclaredOutputType, WorkflowNodeJobConfig, ) +from services.agent.workspace_service import AgentWorkspaceNotFoundError class FakeCredentialsProvider: @@ -95,12 +101,14 @@ class FakeVariablePool: class FakeBindingResolver(WorkflowAgentBindingResolver): def __init__(self): + self.calls: list[dict[str, object]] = [] self.agent = Agent(id="agent-1", tenant_id="tenant-1", name="Agent") self.snapshot = AgentConfigSnapshot( id="snapshot-1", tenant_id="tenant-1", agent_id="agent-1", version=1, + home_snapshot_id="home-1", config_snapshot=AgentSoulConfig( prompt={"system_prompt": "You are careful."}, model=AgentSoulModelConfig( @@ -127,13 +135,29 @@ class FakeBindingResolver(WorkflowAgentBindingResolver): ), ) - def resolve(self, **_kwargs): + def resolve(self, **kwargs): + self.calls.append(kwargs) + snapshot_id = kwargs.get("snapshot_id") + if isinstance(snapshot_id, str): + self.snapshot.id = snapshot_id return WorkflowAgentBindingBundle(binding=self.binding, agent=self.agent, snapshot=self.snapshot) class FakeSessionStore: - def __init__(self, snapshot: CompositorSessionSnapshot | None = None) -> None: + def __init__( + self, + snapshot: CompositorSessionSnapshot | None = None, + *, + binding_id: str = "binding-1", + workspace_id: str = "workspace-1", + backend_binding_ref: str = "backend-binding-1", + ) -> None: self.loaded_snapshot = snapshot + self.binding_id = binding_id + self.workspace_id = workspace_id + self.backend_binding_ref = backend_binding_ref + self.resolved_scopes: list[WorkflowAgentSessionScope] = [] + self.existing_scope_lookups: list[dict[str, object]] = [] # ENG-638: set to simulate resume after a submitted/timed-out form. self.loaded_session: StoredWorkflowAgentSession | None = None self.saved: list[ @@ -141,40 +165,43 @@ class FakeSessionStore: WorkflowAgentSessionScope, str, CompositorSessionSnapshot | None, - list[RuntimeLayerSpec], str | None, str | None, ] ] = [] - self.cleaned: list[tuple[WorkflowAgentSessionScope, str | None]] = [] - def load_active_snapshot(self, scope: WorkflowAgentSessionScope) -> CompositorSessionSnapshot | None: - return self.loaded_snapshot + def load_existing_node_execution_scope(self, **kwargs: object) -> WorkflowAgentSessionScope | None: + self.existing_scope_lookups.append(kwargs) + return self.loaded_session.scope if self.loaded_session is not None else None - def load_active_session(self, scope: WorkflowAgentSessionScope) -> StoredWorkflowAgentSession | None: - return self.loaded_session + def load_or_create_node_execution_session( + self, + scope: WorkflowAgentSessionScope, + *, + home_snapshot_id: str, + ) -> StoredWorkflowAgentSession: + assert home_snapshot_id == "home-1" + self.resolved_scopes.append(scope) + if self.loaded_session is not None: + return self.loaded_session + return StoredWorkflowAgentSession( + scope=scope, + binding_id=self.binding_id, + workspace_id=self.workspace_id, + backend_binding_ref=self.backend_binding_ref, + session_snapshot=self.loaded_snapshot, + ) def save_active_snapshot( self, *, scope: WorkflowAgentSessionScope, - backend_run_id: str, + binding_id: str, snapshot: CompositorSessionSnapshot | None, - runtime_layer_specs: list[RuntimeLayerSpec], pending_form_id: str | None = None, pending_tool_call_id: str | None = None, ) -> None: - self.saved.append( - (scope, backend_run_id, snapshot, list(runtime_layer_specs), pending_form_id, pending_tool_call_id) - ) - - def mark_cleaned( - self, - *, - scope: WorkflowAgentSessionScope, - backend_run_id: str | None = None, - ) -> None: - self.cleaned.append((scope, backend_run_id)) + self.saved.append((scope, binding_id, snapshot, pending_form_id, pending_tool_call_id)) class FileOutputBackendClient(FakeAgentBackendRunClient): @@ -286,6 +313,8 @@ def _node( session_store: FakeSessionStore | None = None, declared_outputs: list[dict[str, object]] | None = None, agent_backend_client: FakeAgentBackendRunClient | None = None, + binding_resolver: FakeBindingResolver | None = None, + runtime_request_builder: WorkflowAgentRuntimeRequestBuilder | None = None, ) -> DifyAgentNode: graph_init_params = GraphInitParams( workflow_id="workflow-1", @@ -309,7 +338,7 @@ def _node( return True client = agent_backend_client or FakeAgentBackendRunClient(scenario=scenario) - binding_resolver = FakeBindingResolver() + binding_resolver = binding_resolver or FakeBindingResolver() if declared_outputs is not None: binding_resolver.binding.node_job_config = WorkflowNodeJobConfig.model_validate( { @@ -319,20 +348,29 @@ def _node( } ) - return DifyAgentNode( + node = DifyAgentNode( node_id="agent-node", data=DifyAgentNodeData.model_validate({"type": BuiltinNodeTypes.AGENT, "version": "2"}), graph_init_params=graph_init_params, - graph_runtime_state=cast(GraphRuntimeState, SimpleNamespace(variable_pool=FakeVariablePool())), + graph_runtime_state=cast( + GraphRuntimeState, + SimpleNamespace( + variable_pool=FakeVariablePool(), + graph_execution=SimpleNamespace(aborted=False), + ), + ), binding_resolver=binding_resolver, - runtime_request_builder=WorkflowAgentRuntimeRequestBuilder(credentials_provider=FakeCredentialsProvider()), + runtime_request_builder=runtime_request_builder + or WorkflowAgentRuntimeRequestBuilder(credentials_provider=FakeCredentialsProvider()), agent_backend_client=client, event_adapter=AgentBackendRunEventAdapter(), output_adapter=WorkflowAgentOutputAdapter(), type_checker=PerOutputTypeChecker(file_validator=_AlwaysAllowFileValidator()), failure_orchestrator=OutputFailureOrchestrator(), - session_store=cast(WorkflowAgentRuntimeSessionStore | None, session_store), + session_store=cast(WorkflowAgentWorkspaceStore, session_store or FakeSessionStore()), ) + node.bind_execution_id(str(uuid4())) + return node def test_extract_variable_selector_to_variable_mapping_uses_frontend_agent_task_markers(): @@ -367,6 +405,85 @@ def test_agent_node_run_maps_successful_agent_backend_run_to_node_result(): assert layers["llm"]["config"]["credentials"] == "[REDACTED]" +def test_agent_node_uses_resolved_backend_binding_before_backend_invocation() -> None: + client = FakeAgentBackendRunClient() + store = FakeSessionStore(binding_id="binding-2", backend_binding_ref="backend-binding-2") + + events = list(_node(agent_backend_client=client, session_store=store)._run()) + + assert len(events) == 1 + assert client.request is not None + layers = {layer["name"]: layer for layer in client.request.model_dump(mode="json")["composition"]["layers"]} + assert layers["runtime"]["config"]["backend_binding_ref"] == "backend-binding-2" + assert store.saved[0][1] == "binding-2" + assert len(store.resolved_scopes) == 1 + + +def test_agent_node_resume_resolves_the_generation_from_the_persisted_execution() -> None: + binding_resolver = FakeBindingResolver() + store = FakeSessionStore() + store.loaded_session = StoredWorkflowAgentSession( + scope=WorkflowAgentSessionScope( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_run_id="workflow-run-1", + node_id="agent-node", + node_execution_id="exec-1", + workflow_agent_binding_id="binding-1", + agent_id="agent-1", + agent_config_snapshot_id="snapshot-pinned", + ), + binding_id="workspace-binding-1", + workspace_id="workspace-1", + backend_binding_ref="backend-binding-1", + session_snapshot=None, + ) + node = _node(binding_resolver=binding_resolver, session_store=store) + + events = list(node._run()) + + assert len(events) == 1 + assert store.existing_scope_lookups[0]["node_execution_id"] == node.execution_id + assert binding_resolver.calls[0]["binding_id"] == "binding-1" + assert binding_resolver.calls[0]["snapshot_id"] == "snapshot-pinned" + + +def test_agent_node_maps_persisted_participant_lookup_error_to_node_failure() -> None: + class _UnavailableParticipantStore(FakeSessionStore): + def load_existing_node_execution_scope(self, **kwargs: object) -> WorkflowAgentSessionScope | None: + del kwargs + raise AgentWorkspaceNotFoundError("Workflow node participant Binding is unavailable") + + events = list(_node(session_store=_UnavailableParticipantStore())._run()) + + assert len(events) == 1 + result = cast(StreamCompletedEvent, events[0]).node_run_result + assert result.status == WorkflowNodeExecutionStatus.FAILED + assert result.error == "Workflow node participant Binding is unavailable" + assert result.error_type == "agent_workflow_node_runtime_error" + + +def test_agent_node_passes_execution_id_to_session_store_and_runtime_request_builder() -> None: + store = FakeSessionStore() + request_builder = WorkflowAgentRuntimeRequestBuilder(credentials_provider=FakeCredentialsProvider()) + node = _node(session_store=store, runtime_request_builder=request_builder) + execution_id = node.execution_id + + with patch.object(request_builder, "build", wraps=request_builder.build) as build: + list(node._run()) + + assert str(UUID(execution_id)) == execution_id + assert execution_id != node.id + scope = store.resolved_scopes[0] + build.assert_called_once() + context = cast(WorkflowAgentRuntimeBuildContext, build.call_args.args[0]) + assert scope.node_id == node.id + assert scope.node_execution_id == execution_id + assert context.node_id == node.id + assert context.node_execution_id == execution_id + + def test_agent_node_run_ignores_agent_message_delta_until_terminal_result(): events = list(_node(agent_backend_client=AgentMessageDeltaBackendClient())._run()) @@ -491,23 +608,6 @@ def test_agent_node_run_maps_failed_agent_backend_run_to_node_result(): assert result.error_type == "unit_test" -def test_agent_node_failed_run_marks_session_cleaned_to_prevent_stale_reuse(): - """A failed agent run must retire the local ACTIVE session row so a workflow - loop back into the same Agent node does not resume from a stale snapshot.""" - existing_snapshot = CompositorSessionSnapshot(layers=[]) - store = FakeSessionStore(snapshot=existing_snapshot) - - events = list(_node(scenario=FakeAgentBackendScenario.FAILED, session_store=store)._run()) - - assert len(events) == 1 - assert store.cleaned, "failed agent run should mark the session cleaned" - cleaned_scope, cleaned_backend_run_id = store.cleaned[0] - assert cleaned_scope.workflow_run_id == "workflow-run-1" - assert cleaned_backend_run_id == "fake-run-1" - # A failed run does not produce a fresh snapshot to persist. - assert store.saved == [] - - def test_agent_node_saves_success_snapshot_and_reuses_existing_snapshot(): existing_snapshot = CompositorSessionSnapshot(layers=[]) store = FakeSessionStore(snapshot=existing_snapshot) @@ -518,21 +618,15 @@ def test_agent_node_saves_success_snapshot_and_reuses_existing_snapshot(): assert len(events) == 1 assert store.saved - scope, backend_run_id, saved_snapshot, saved_specs, pending_form_id, pending_tool_call_id = store.saved[0] + scope, binding_id, saved_snapshot, pending_form_id, pending_tool_call_id = store.saved[0] assert scope.workflow_run_id == "workflow-run-1" - assert backend_run_id == "fake-run-1" + assert binding_id == "binding-1" assert saved_snapshot is not None # A successful terminal carries no ask_human pause correlation. assert pending_form_id is None assert pending_tool_call_id is None assert client.request is not None assert client.request.session_snapshot is existing_snapshot - # Persist enough composition shape to replay a cleanup run; plugin layers - # (which would carry credentials) are intentionally absent. - saved_layer_names = [spec.name for spec in saved_specs] - assert saved_layer_names, "cleanup specs must persist at least the non-plugin layers" - plugin_types = {"dify.plugin.llm", "dify.plugin.tools"} - assert not {spec.type for spec in saved_specs} & plugin_types def test_agent_node_run_when_session_store_save_raises_records_persist_error_in_metadata(): @@ -553,83 +647,7 @@ def test_agent_node_run_when_session_store_save_raises_records_persist_error_in_ assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED agent_backend = result.metadata[WorkflowNodeExecutionMetadataKey.AGENT_LOG]["agent_backend"] assert agent_backend["session_snapshot_persisted"] is False - assert agent_backend["session_snapshot_persist_error"] == "workflow_agent_runtime_session_store_error" - - -def test_agent_node_failed_run_when_mark_cleaned_raises_records_cleanup_error_in_metadata(): - """Same defensive pattern: a DB-side mark_cleaned failure must surface as - a ``session_snapshot_cleanup_error`` in metadata, not as a node crash.""" - - class _ExplodingMarkCleanedStore(FakeSessionStore): - def mark_cleaned(self, **kwargs): # type: ignore[override] - del kwargs - raise RuntimeError("simulated DB failure") - - store = _ExplodingMarkCleanedStore() - events = list(_node(scenario=FakeAgentBackendScenario.FAILED, session_store=store)._run()) - - assert len(events) == 1 - result = cast(StreamCompletedEvent, events[0]).node_run_result - assert result.status == WorkflowNodeExecutionStatus.FAILED - agent_backend = result.metadata[WorkflowNodeExecutionMetadataKey.AGENT_LOG]["agent_backend"] - assert agent_backend["session_snapshot_cleaned_on_failure"] is False - assert agent_backend["session_snapshot_cleanup_error"] == "workflow_agent_runtime_session_store_error" - - -def test_agent_node_success_run_without_session_store_skips_persistence(): - """When ``session_store`` is None the node still completes successfully — - the lifecycle branch is a no-op and the run result is unaffected.""" - events = list(_node(session_store=None)._run()) - - assert len(events) == 1 - result = cast(StreamCompletedEvent, events[0]).node_run_result - assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED - agent_backend = result.metadata[WorkflowNodeExecutionMetadataKey.AGENT_LOG]["agent_backend"] - # No persistence metadata is attached when the store is missing. - assert "session_snapshot_persisted" not in agent_backend - - -def test_agent_node_failed_run_without_session_store_skips_mark_cleaned(): - """``session_store=None`` + failed terminal must remain a no-op for - the cleanup branch — the node failure path still surfaces correctly.""" - events = list(_node(scenario=FakeAgentBackendScenario.FAILED, session_store=None)._run()) - - assert len(events) == 1 - result = cast(StreamCompletedEvent, events[0]).node_run_result - assert result.status == WorkflowNodeExecutionStatus.FAILED - agent_backend = result.metadata[WorkflowNodeExecutionMetadataKey.AGENT_LOG]["agent_backend"] - assert "session_snapshot_cleaned_on_failure" not in agent_backend - - -def test_agent_node_failed_run_enqueues_backend_cleanup_before_local_retirement(monkeypatch): - store = FakeSessionStore() - store.loaded_session = StoredWorkflowAgentSession( - scope=_pending_session(CompositorSessionSnapshot(layers=[])).scope, - session_snapshot=CompositorSessionSnapshot(layers=[]), - backend_run_id="stored-run-1", - runtime_layer_specs=[RuntimeLayerSpec(name="history", type="pydantic_ai.history")], - ) - queued_payloads: list[dict[str, object]] = [] - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.agent_node.cleanup_workflow_agent_runtime_session.delay", - lambda payload: queued_payloads.append(payload), - ) - - events = list(_node(scenario=FakeAgentBackendScenario.FAILED, session_store=store)._run()) - - assert len(events) == 1 - result = cast(StreamCompletedEvent, events[0]).node_run_result - assert result.status == WorkflowNodeExecutionStatus.FAILED - assert store.cleaned[0][1] == "fake-run-1" - assert store.cleaned[0][0].workflow_run_id == "workflow-run-1" - assert store.cleaned[0][0].node_id == "agent-node" - assert len(queued_payloads) == 1 - assert ( - queued_payloads[0]["idempotency_key"] - == "tenant-1:workflow-run-1:agent-node:binding-1:workflow-agent-failure-cleanup:stored-run-1:fake-run-1" - ) - assert queued_payloads[0]["metadata"]["previous_agent_backend_run_id"] == "stored-run-1" - assert queued_payloads[0]["metadata"]["failed_agent_backend_run_id"] == "fake-run-1" + assert agent_backend["session_snapshot_persist_error"] == "workflow_agent_workspace_store_error" def test_agent_node_paused_run_requests_workflow_pause_and_persists_snapshot(): @@ -646,17 +664,21 @@ def test_agent_node_paused_run_requests_workflow_pause_and_persists_snapshot(): events = list(node._run()) assert len(events) == 1 - assert isinstance(events[0], PauseRequestedEvent) + assert isinstance(events[0], NodeRunPauseRequestedEvent) assert isinstance(events[0].reason, HitlRequired) assert events[0].reason.session_id == "form-1" assert events[0].reason.node_id == "agent-node" + assert events[0].node_run_result.process_data == { + "agent_id": "agent-1", + "agent_config_snapshot_id": "snapshot-1", + "workflow_agent_binding_id": "binding-1", + } fake_repo.create_form.assert_called_once() assert store.saved - assert store.saved[0][1] == "fake-run-1" - assert store.saved[0][3], "paused agent run should still persist replayable layer specs" + assert store.saved[0][1] == "binding-1" # ENG-637: the awaiting form + deferred tool_call correlation is persisted. - assert store.saved[0][4] == "form-1" - assert store.saved[0][5] == "fake-ask-human-1" + assert store.saved[0][3] == "form-1" + assert store.saved[0][4] == "fake-ask-human-1" def _pending_session(snapshot: CompositorSessionSnapshot) -> StoredWorkflowAgentSession: @@ -668,18 +690,20 @@ def _pending_session(snapshot: CompositorSessionSnapshot) -> StoredWorkflowAgent workflow_run_id="workflow-run-1", node_id="agent-node", node_execution_id="exec-1", - binding_id="binding-1", + workflow_agent_binding_id="binding-1", agent_id="agent-1", agent_config_snapshot_id="snapshot-1", ), + binding_id="binding-1", + workspace_id="workspace-1", + backend_binding_ref="backend-binding-1", session_snapshot=snapshot, - backend_run_id="run-0", pending_form_id="form-1", pending_tool_call_id="call-1", ) -def test_agent_node_resumes_with_deferred_tool_results_after_submitted_form(monkeypatch): +def test_agent_node_resumes_with_deferred_tool_results_after_submitted_form(monkeypatch: pytest.MonkeyPatch): # ENG-638: a submitted form re-enters _run; the human's answer is threaded # into the second Agent run as deferred_tool_results. snapshot = CompositorSessionSnapshot(layers=[]) @@ -703,7 +727,7 @@ def test_agent_node_resumes_with_deferred_tool_results_after_submitted_form(monk assert any(isinstance(event, StreamCompletedEvent) for event in events) -def test_agent_node_repauses_when_resumed_form_still_waiting(monkeypatch): +def test_agent_node_repauses_when_resumed_form_still_waiting(monkeypatch: pytest.MonkeyPatch): snapshot = CompositorSessionSnapshot(layers=[]) store = FakeSessionStore(snapshot=snapshot) store.loaded_session = _pending_session(snapshot) @@ -727,19 +751,75 @@ def test_agent_node_repauses_when_resumed_form_still_waiting(monkeypatch): events = list(node._run()) assert len(events) == 1 - assert isinstance(events[0], PauseRequestedEvent) + assert isinstance(events[0], NodeRunPauseRequestedEvent) assert isinstance(events[0].reason, HitlRequired) + assert events[0].node_run_result.process_data["workflow_agent_binding_id"] == "binding-1" assert client.request is None # no second Agent run was created +def test_agent_node_expired_ask_human_failure_keeps_binding_identity(monkeypatch: pytest.MonkeyPatch): + snapshot = CompositorSessionSnapshot(layers=[]) + store = FakeSessionStore(snapshot=snapshot) + store.loaded_session = _pending_session(snapshot) + + def _raise_expired_form(**_kwargs): + raise AssertionError("cannot resume globally expired ask_human form, form_id=form-1") + + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.agent_node.resolve_ask_human_form", + _raise_expired_form, + ) + + events = list(_node(session_store=store)._run()) + + assert len(events) == 1 + result = cast(StreamCompletedEvent, events[0]).node_run_result + assert result.status == WorkflowNodeExecutionStatus.FAILED + assert result.error == "cannot resume globally expired ask_human form, form_id=form-1" + assert result.error_type == "agent_workflow_node_runtime_error" + assert result.process_data["workflow_agent_binding_id"] == "binding-1" + assert "agent_workspace_binding_id" not in result.process_data + + +def test_agent_node_unexpected_post_resolution_failure_keeps_binding_identity(): + class _FailingSessionStore(FakeSessionStore): + def load_or_create_node_execution_session( + self, + scope: WorkflowAgentSessionScope, + *, + home_snapshot_id: str, + ) -> StoredWorkflowAgentSession: + del scope, home_snapshot_id + raise RuntimeError("session store failed") + + events = list(_node(session_store=_FailingSessionStore())._run()) + + assert len(events) == 1 + result = cast(StreamCompletedEvent, events[0]).node_run_result + assert result.status == WorkflowNodeExecutionStatus.FAILED + assert result.error == "session store failed" + assert result.error_type == "agent_workflow_node_runtime_error" + assert result.process_data == { + "agent_id": "agent-1", + "agent_config_snapshot_id": "snapshot-1", + "workflow_agent_binding_id": "binding-1", + } + + def test_agent_node_cancels_backend_run_when_stream_fails(): client = FailingStreamBackendClient() node = _node(agent_backend_client=client) - terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + terminal, failure = node._consume_event_stream( + "run-1", + inputs={}, + process_data={"workflow_agent_binding_id": "binding-1"}, + metadata={"agent_backend": {}}, + ) assert terminal is None assert failure is not None + assert failure.node_run_result.process_data == {"workflow_agent_binding_id": "binding-1"} assert len(client.cancel_requests) == 1 assert client.cancel_requests[0] is not None assert client.cancel_requests[0].reason == "event_stream_failed" @@ -749,7 +829,12 @@ def test_agent_node_cancels_backend_run_when_stream_ends_without_terminal_event( client = EmptyStreamBackendClient() node = _node(agent_backend_client=client) - terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + terminal, failure = node._consume_event_stream( + "run-1", + inputs={}, + process_data={"workflow_agent_binding_id": "binding-1"}, + metadata={"agent_backend": {}}, + ) assert terminal is None assert failure is None @@ -761,11 +846,17 @@ def test_agent_node_cancels_backend_run_when_stream_raises_unexpected_error(): client = GenericFailingStreamBackendClient() node = _node(agent_backend_client=client) - terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + terminal, failure = node._consume_event_stream( + "run-1", + inputs={}, + process_data={"workflow_agent_binding_id": "binding-1"}, + metadata={"agent_backend": {}}, + ) assert terminal is None assert failure is not None assert failure.node_run_result.error == "unexpected stream failure" + assert failure.node_run_result.process_data == {"workflow_agent_binding_id": "binding-1"} assert client.cancel_requests[0] is not None assert client.cancel_requests[0].reason == "event_stream_failed" @@ -775,7 +866,12 @@ def test_agent_node_uses_graph_abort_reason_when_cancel_request_fails(caplog): node = _node(agent_backend_client=client) node.graph_runtime_state.graph_execution = SimpleNamespace(aborted=True) - terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + terminal, failure = node._consume_event_stream( + "run-1", + inputs={}, + process_data={"workflow_agent_binding_id": "binding-1"}, + metadata={"agent_backend": {}}, + ) assert terminal is None assert failure is not None @@ -794,13 +890,19 @@ def test_agent_node_cancels_backend_run_for_unexpected_internal_event(): return_value=[SimpleNamespace(type=AgentBackendInternalEventType.RUN_FAILED)] ) - terminal, failure = node._consume_event_stream("run-1", {"agent_backend": {}}) + terminal, failure = node._consume_event_stream( + "run-1", + inputs={}, + process_data={"workflow_agent_binding_id": "binding-1"}, + metadata={"agent_backend": {}}, + ) assert terminal is None assert failure is not None assert failure.node_run_result.error == ( "Unexpected internal event type " ) + assert failure.node_run_result.process_data == {"workflow_agent_binding_id": "binding-1"} node._agent_backend_client.cancel_run.assert_called_once() diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_binding_resolver.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_binding_resolver.py index 86f3046d47f..c7774e6836f 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_binding_resolver.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_binding_resolver.py @@ -1,10 +1,13 @@ -import pytest -from sqlalchemy.orm import Session +from uuid import uuid4 -from core.workflow.nodes.agent_v2.binding_resolver import ( - WorkflowAgentBindingError, - WorkflowAgentBindingResolver, -) +import pytest +from sqlalchemy import event, inspect +from sqlalchemy.engine import Engine +from sqlalchemy.orm import ORMExecuteState, Session, sessionmaker +from sqlalchemy.sql import Executable + +import core.workflow.nodes.agent_v2.binding_resolver as resolver_module +from core.workflow.nodes.agent_v2.binding_resolver import WorkflowAgentBindingError, WorkflowAgentBindingResolver from models.agent import ( Agent, AgentConfigRevision, @@ -18,49 +21,57 @@ from models.agent import ( ) from models.agent_config_entities import AgentSoulConfig, AgentSoulModelConfig, WorkflowNodeJobConfig - -class FakeSession: - def __init__(self, scalar_results): - self._scalar_results = list(scalar_results) - self.expunge_calls = [] - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb): - return False - - def scalar(self, _stmt): - if not self._scalar_results: - return None - return self._scalar_results.pop(0) - - def expunge(self, value): - self.expunge_calls.append(value) +RESOLVER_MODELS = (WorkflowAgentNodeBinding, Agent, AgentConfigSnapshot, AgentConfigRevision) -def _binding() -> WorkflowAgentNodeBinding: - return WorkflowAgentNodeBinding( - id="binding-1", - tenant_id="tenant-1", - app_id="app-1", - workflow_id="workflow-1", - node_id="agent-node", - agent_id="agent-1", - current_snapshot_id="snapshot-1", - node_job_config=WorkflowNodeJobConfig(), +def _resolve_ids() -> dict[str, str]: + return { + "tenant_id": str(uuid4()), + "app_id": str(uuid4()), + "workflow_id": str(uuid4()), + "node_id": "agent-node", + } + + +def _agent( + *, + tenant_id: str, + status: AgentStatus = AgentStatus.ACTIVE, + scope: AgentScope = AgentScope.WORKFLOW_ONLY, + source: AgentSource = AgentSource.WORKFLOW, + app_id: str | None = None, + active_config_snapshot_id: str | None = None, + active_config_is_published: bool = True, +) -> Agent: + return Agent( + tenant_id=tenant_id, + name=f"Agent {uuid4()}", + description="", + role="", + icon_type=None, + icon=None, + icon_background=None, + scope=scope, + source=source, + app_id=app_id, + backing_app_id=None, + workflow_id=None, + workflow_node_id=None, + active_config_snapshot_id=active_config_snapshot_id, + active_config_has_model=True, + active_config_is_published=active_config_is_published, + status=status, + created_by=None, + updated_by=None, + archived_by=None, + archived_at=None, ) -def _agent(*, status: AgentStatus = AgentStatus.ACTIVE) -> Agent: - return Agent(id="agent-1", tenant_id="tenant-1", name="Agent", status=status) - - -def _snapshot() -> AgentConfigSnapshot: +def _snapshot(*, tenant_id: str, agent_id: str) -> AgentConfigSnapshot: return AgentConfigSnapshot( - id="snapshot-1", - tenant_id="tenant-1", - agent_id="agent-1", + tenant_id=tenant_id, + agent_id=agent_id, version=1, config_snapshot=AgentSoulConfig( model=AgentSoulModelConfig( @@ -69,166 +80,353 @@ def _snapshot() -> AgentConfigSnapshot: model="gpt-test", ) ), + summary=None, + version_note=None, + created_by=None, ) -def _resolve() -> dict[str, str]: - return { - "tenant_id": "tenant-1", - "app_id": "app-1", - "workflow_id": "workflow-1", - "node_id": "agent-node", - } - - -def test_binding_resolver_returns_detached_binding_bundle(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession([_binding(), _agent(), _snapshot()]) - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session", - lambda: fake_session, +def _binding( + *, ids: dict[str, str], agent_id: str, snapshot_id: str, binding_type: WorkflowAgentBindingType +) -> WorkflowAgentNodeBinding: + return WorkflowAgentNodeBinding( + tenant_id=ids["tenant_id"], + app_id=ids["app_id"], + workflow_id=ids["workflow_id"], + workflow_version="draft", + node_id=ids["node_id"], + binding_type=binding_type, + agent_id=agent_id, + current_snapshot_id=snapshot_id, + node_job_config=WorkflowNodeJobConfig(), + created_by=None, + updated_by=None, ) - bundle = WorkflowAgentBindingResolver().resolve(**_resolve()) - assert bundle.binding.id == "binding-1" - assert bundle.agent.id == "agent-1" - assert bundle.snapshot.id == "snapshot-1" - assert fake_session.expunge_calls == [bundle.binding, bundle.agent, bundle.snapshot] +def _bind_factory(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> list[Executable]: + scalar_statements: list[Executable] = [] + + class RecordingSession(Session): + pass + + def record_statement(execute_state: ORMExecuteState) -> None: + scalar_statements.append(execute_state.statement) + + event.listen(RecordingSession, "do_orm_execute", record_statement) + factory = sessionmaker(bind=sqlite_engine, class_=RecordingSession, expire_on_commit=False) + monkeypatch.setattr(resolver_module.session_factory, "create_session", factory) + return scalar_statements -def test_binding_resolver_uses_active_snapshot_for_roster_agent(monkeypatch: pytest.MonkeyPatch): - binding = _binding() - binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT - binding.current_snapshot_id = "old-snapshot" - agent = _agent() - agent.active_config_snapshot_id = "active-snapshot" - snapshot = _snapshot() - snapshot.id = "active-snapshot" - fake_session = FakeSession([binding, agent, snapshot]) - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session", - lambda: fake_session, +@pytest.mark.parametrize("sqlite_session", [RESOLVER_MODELS], indirect=True) +def test_binding_resolver_returns_detached_binding_bundle( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session +) -> None: + ids = _resolve_ids() + agent = _agent(tenant_id=ids["tenant_id"]) + sqlite_session.add(agent) + sqlite_session.flush() + snapshot = _snapshot(tenant_id=ids["tenant_id"], agent_id=agent.id) + sqlite_session.add(snapshot) + sqlite_session.flush() + binding = _binding( + ids=ids, + agent_id=agent.id, + snapshot_id=snapshot.id, + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + ) + sqlite_session.add(binding) + sqlite_session.commit() + _bind_factory(monkeypatch, sqlite_engine) + + bundle = WorkflowAgentBindingResolver().resolve(**ids) + + assert bundle.binding.id == binding.id + assert bundle.agent.id == agent.id + assert bundle.snapshot.id == snapshot.id + assert inspect(bundle.binding).detached + assert inspect(bundle.agent).detached + assert inspect(bundle.snapshot).detached + + +@pytest.mark.parametrize("sqlite_session", [RESOLVER_MODELS], indirect=True) +def test_binding_resolver_uses_active_snapshot_for_roster_agent( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session +) -> None: + ids = _resolve_ids() + agent = _agent( + tenant_id=ids["tenant_id"], + scope=AgentScope.ROSTER, + source=AgentSource.ROSTER, + ) + sqlite_session.add(agent) + sqlite_session.flush() + active_snapshot = _snapshot(tenant_id=ids["tenant_id"], agent_id=agent.id) + sqlite_session.add(active_snapshot) + sqlite_session.flush() + agent.active_config_snapshot_id = active_snapshot.id + binding = _binding( + ids=ids, + agent_id=agent.id, + snapshot_id=str(uuid4()), + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + ) + sqlite_session.add(binding) + sqlite_session.commit() + _bind_factory(monkeypatch, sqlite_engine) + + bundle = WorkflowAgentBindingResolver().resolve(**ids) + + assert bundle.snapshot.id == active_snapshot.id + + +@pytest.mark.parametrize( + ("binding_type", "scope", "source"), + [ + (WorkflowAgentBindingType.ROSTER_AGENT, AgentScope.ROSTER, AgentSource.ROSTER), + (WorkflowAgentBindingType.INLINE_AGENT, AgentScope.WORKFLOW_ONLY, AgentSource.WORKFLOW), + ], +) +@pytest.mark.parametrize("sqlite_session", [RESOLVER_MODELS], indirect=True) +def test_binding_resolver_uses_pinned_snapshot_for_existing_node_execution( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, + binding_type: WorkflowAgentBindingType, + scope: AgentScope, + source: AgentSource, +) -> None: + ids = _resolve_ids() + agent = _agent( + tenant_id=ids["tenant_id"], + scope=scope, + source=source, + active_config_snapshot_id=str(uuid4()), + ) + sqlite_session.add(agent) + sqlite_session.flush() + pinned_snapshot = _snapshot(tenant_id=ids["tenant_id"], agent_id=agent.id) + sqlite_session.add(pinned_snapshot) + sqlite_session.flush() + binding = _binding( + ids=ids, + agent_id=agent.id, + snapshot_id=str(uuid4()), + binding_type=binding_type, + ) + sqlite_session.add(binding) + sqlite_session.commit() + scalar_statements = _bind_factory(monkeypatch, sqlite_engine) + + bundle = WorkflowAgentBindingResolver().resolve( + **ids, + binding_id=binding.id, + snapshot_id=pinned_snapshot.id, ) - bundle = WorkflowAgentBindingResolver().resolve(**_resolve()) - - assert bundle.snapshot.id == "active-snapshot" + assert bundle.snapshot.id == pinned_snapshot.id + assert binding.id in scalar_statements[0].compile().params.values() + assert pinned_snapshot.id in scalar_statements[-1].compile().params.values() -def test_binding_resolver_rejects_unpublished_roster_agent(monkeypatch: pytest.MonkeyPatch): - binding = _binding() - binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session", - lambda: FakeSession([binding, None]), +@pytest.mark.parametrize("sqlite_session", [RESOLVER_MODELS], indirect=True) +def test_binding_resolver_does_not_fallback_from_an_explicit_empty_snapshot( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session +) -> None: + ids = _resolve_ids() + agent = _agent( + tenant_id=ids["tenant_id"], + scope=AgentScope.ROSTER, + source=AgentSource.ROSTER, ) + sqlite_session.add(agent) + sqlite_session.flush() + active_snapshot = _snapshot(tenant_id=ids["tenant_id"], agent_id=agent.id) + sqlite_session.add(active_snapshot) + sqlite_session.flush() + agent.active_config_snapshot_id = active_snapshot.id + binding = _binding( + ids=ids, + agent_id=agent.id, + snapshot_id=str(uuid4()), + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + ) + sqlite_session.add(binding) + sqlite_session.commit() + _bind_factory(monkeypatch, sqlite_engine) with pytest.raises(WorkflowAgentBindingError) as exc_info: - WorkflowAgentBindingResolver().resolve(**_resolve()) + WorkflowAgentBindingResolver().resolve(**ids, binding_id=binding.id, snapshot_id="") + + assert exc_info.value.error_code == "agent_config_snapshot_not_found" + + +@pytest.mark.parametrize( + ("binding_id", "snapshot_id"), + [("binding-1", None), (None, "snapshot-1")], +) +def test_binding_resolver_rejects_half_pinned_generation( + binding_id: str | None, + snapshot_id: str | None, +) -> None: + with pytest.raises(WorkflowAgentBindingError) as exc_info: + WorkflowAgentBindingResolver().resolve( + **_resolve_ids(), + binding_id=binding_id, + snapshot_id=snapshot_id, + ) + + assert exc_info.value.error_code == "agent_binding_generation_invalid" + + +@pytest.mark.parametrize("sqlite_session", [RESOLVER_MODELS], indirect=True) +def test_binding_resolver_rejects_unpublished_roster_agent( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session +) -> None: + ids = _resolve_ids() + snapshot_id = str(uuid4()) + agent = _agent( + tenant_id=ids["tenant_id"], + scope=AgentScope.ROSTER, + source=AgentSource.IMPORTED, + app_id=str(uuid4()), + active_config_snapshot_id=snapshot_id, + active_config_is_published=False, + ) + sqlite_session.add(agent) + sqlite_session.flush() + binding = _binding( + ids=ids, + agent_id=agent.id, + snapshot_id=snapshot_id, + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + ) + sqlite_session.add(binding) + sqlite_session.commit() + _bind_factory(monkeypatch, sqlite_engine) + + with pytest.raises(WorkflowAgentBindingError) as exc_info: + WorkflowAgentBindingResolver().resolve(**ids) assert exc_info.value.error_code == "agent_not_available" assert "not been published" in str(exc_info.value) -@pytest.mark.parametrize( - "sqlite_session", - [(Agent, AgentConfigSnapshot, AgentConfigRevision, WorkflowAgentNodeBinding)], - indirect=True, -) +@pytest.mark.parametrize("sqlite_session", [RESOLVER_MODELS], indirect=True) def test_binding_resolver_requires_publish_provenance_for_active_roster_snapshot( monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, sqlite_session: Session, ) -> None: - binding = _binding() - binding.binding_type = WorkflowAgentBindingType.ROSTER_AGENT - binding.workflow_version = "draft" - agent = Agent( - id="agent-1", - tenant_id="tenant-1", - name="Imported Agent", + ids = _resolve_ids() + agent = _agent( + tenant_id=ids["tenant_id"], scope=AgentScope.ROSTER, source=AgentSource.IMPORTED, - app_id="agent-app-1", - status=AgentStatus.ACTIVE, - active_config_snapshot_id="snapshot-1", - active_config_has_model=True, - # Dirty draft state must not hide a snapshot after it has publish provenance. + app_id=str(uuid4()), active_config_is_published=False, ) + sqlite_session.add(agent) + sqlite_session.flush() + snapshot = _snapshot(tenant_id=ids["tenant_id"], agent_id=agent.id) + sqlite_session.add(snapshot) + sqlite_session.flush() + agent.active_config_snapshot_id = snapshot.id + binding = _binding( + ids=ids, + agent_id=agent.id, + snapshot_id=snapshot.id, + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + ) sqlite_session.add_all( [ binding, - agent, - _snapshot(), AgentConfigRevision( - id="revision-import", - tenant_id="tenant-1", - agent_id="agent-1", - current_snapshot_id="snapshot-1", + tenant_id=ids["tenant_id"], + agent_id=agent.id, + current_snapshot_id=snapshot.id, revision=1, operation=AgentConfigRevisionOperation.IMPORT_PACKAGE, ), ] ) sqlite_session.commit() - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session", - lambda: sqlite_session, - ) + _bind_factory(monkeypatch, sqlite_engine) with pytest.raises(WorkflowAgentBindingError) as exc_info: - WorkflowAgentBindingResolver().resolve(**_resolve()) + WorkflowAgentBindingResolver().resolve(**ids) assert exc_info.value.error_code == "agent_not_available" sqlite_session.add( AgentConfigRevision( - id="revision-publish", - tenant_id="tenant-1", - agent_id="agent-1", - current_snapshot_id="snapshot-1", + tenant_id=ids["tenant_id"], + agent_id=agent.id, + current_snapshot_id=snapshot.id, revision=2, operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, ) ) sqlite_session.commit() - bundle = WorkflowAgentBindingResolver().resolve(**_resolve()) + bundle = WorkflowAgentBindingResolver().resolve(**ids) assert bundle.agent.id == agent.id - assert bundle.snapshot.id == "snapshot-1" + assert bundle.snapshot.id == snapshot.id -def test_binding_resolver_raises_when_binding_missing(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session", - lambda: FakeSession([None]), - ) +def test_binding_resolver_raises_when_binding_missing(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None: + _bind_factory(monkeypatch, sqlite_engine) with pytest.raises(WorkflowAgentBindingError) as exc_info: - WorkflowAgentBindingResolver().resolve(**_resolve()) + WorkflowAgentBindingResolver().resolve(**_resolve_ids()) assert exc_info.value.error_code == "agent_binding_not_found" -def test_binding_resolver_raises_when_agent_archived(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session", - lambda: FakeSession([_binding(), _agent(status=AgentStatus.ARCHIVED)]), +@pytest.mark.parametrize("sqlite_session", [RESOLVER_MODELS], indirect=True) +def test_binding_resolver_raises_when_agent_archived( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session +) -> None: + ids = _resolve_ids() + agent = _agent(tenant_id=ids["tenant_id"], status=AgentStatus.ARCHIVED) + sqlite_session.add(agent) + sqlite_session.flush() + binding = _binding( + ids=ids, + agent_id=agent.id, + snapshot_id=str(uuid4()), + binding_type=WorkflowAgentBindingType.INLINE_AGENT, ) + sqlite_session.add(binding) + sqlite_session.commit() + _bind_factory(monkeypatch, sqlite_engine) with pytest.raises(WorkflowAgentBindingError) as exc_info: - WorkflowAgentBindingResolver().resolve(**_resolve()) + WorkflowAgentBindingResolver().resolve(**ids) assert exc_info.value.error_code == "agent_not_available" -def test_binding_resolver_raises_when_snapshot_missing(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.binding_resolver.session_factory.create_session", - lambda: FakeSession([_binding(), _agent(), None]), +@pytest.mark.parametrize("sqlite_session", [RESOLVER_MODELS], indirect=True) +def test_binding_resolver_raises_when_snapshot_missing( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session +) -> None: + ids = _resolve_ids() + agent = _agent(tenant_id=ids["tenant_id"]) + sqlite_session.add(agent) + sqlite_session.flush() + binding = _binding( + ids=ids, + agent_id=agent.id, + snapshot_id=str(uuid4()), + binding_type=WorkflowAgentBindingType.INLINE_AGENT, ) + sqlite_session.add(binding) + sqlite_session.commit() + _bind_factory(monkeypatch, sqlite_engine) with pytest.raises(WorkflowAgentBindingError) as exc_info: - WorkflowAgentBindingResolver().resolve(**_resolve()) + WorkflowAgentBindingResolver().resolve(**ids) assert exc_info.value.error_code == "agent_config_snapshot_not_found" diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_dify_tools_builder.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_dify_tools_builder.py index b2abfbf222d..cf169ca6edf 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_dify_tools_builder.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_dify_tools_builder.py @@ -1007,7 +1007,7 @@ def test_provider_level_entry_unknown_provider_maps_to_declaration_not_found(): assert exc_info.value.error_code == "agent_tool_declaration_not_found" -def test_list_provider_tool_names_reads_builtin_provider(monkeypatch): +def test_list_provider_tool_names_reads_builtin_provider(monkeypatch: pytest.MonkeyPatch): """The default provider-tools lister maps ToolManager's provider controller to the plain name list the expansion step consumes.""" from types import SimpleNamespace diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_output_adapter.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_output_adapter.py index 760786c5b75..33f41a5ea3d 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_output_adapter.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_output_adapter.py @@ -3,6 +3,7 @@ from unittest.mock import patch import pytest from agenton.compositor import CompositorSessionSnapshot +from dify_agent.protocol import RunFailureType from clients.agent_backend import ( AgentBackendRunCancelledInternalEvent, @@ -162,6 +163,38 @@ def test_failure_output_adapter_preserves_backend_failed_reason(): assert result.error_type == "validation" +def test_failure_output_adapter_prefers_run_failure_type_over_reason(): + result = WorkflowAgentOutputAdapter().build_failure_result( + event=AgentBackendRunFailedInternalEvent( + run_id="run-1", + error="run limit reached", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + reason="runtime", + ), + inputs={}, + process_data={}, + metadata={}, + ) + + assert result.error_type == "agent_run_limit_exceeded" + + +def test_failure_output_adapter_uses_default_error_type_without_backend_classification(): + result = WorkflowAgentOutputAdapter().build_failure_result( + event=AgentBackendRunFailedInternalEvent( + run_id="run-1", + error="backend failed", + error_type=None, + reason=None, + ), + inputs={}, + process_data={}, + metadata={}, + ) + + assert result.error_type == "agent_backend_run_failed" + + def test_success_output_adapter_normalizes_string_and_scalar_outputs(): adapter = WorkflowAgentOutputAdapter() string_result = adapter.build_success_result( diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_runtime_request_builder.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_runtime_request_builder.py index c66bcf5b6e6..c969b296efb 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_runtime_request_builder.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_runtime_request_builder.py @@ -203,6 +203,8 @@ def _context() -> WorkflowAgentRuntimeBuildContext: binding=binding, agent=agent, snapshot=snapshot, + binding_id="binding-1", + backend_binding_ref="binding-ref-1", ) @@ -430,7 +432,7 @@ def test_builds_workflow_run_request_with_file_output_schema_and_reserved_metada assert "never invent the `reference` value" in output_description assert "Do not call `final_output` before the upload command succeeds" in output_description assert "accepted file-mapping shape and the returned `reference`" in output_description - assert "include the returned `download_url` in that reply" in output_description + assert "include the returned `public_download_url` in that reply" in output_description assert output_schema["properties"]["confidence"]["type"] == "number" assert output_schema["required"] == ["report"] assert layers[DIFY_AGENT_MODEL_LAYER_ID]["config"]["model_settings"] == {"temperature": 0.2} @@ -467,7 +469,7 @@ def test_build_maps_agent_soul_shell_settings_to_shell_layer(monkeypatch: pytest shell_config = {layer["name"]: layer for layer in dumped["composition"]["layers"]}[DIFY_SHELL_LAYER_ID]["config"] assert shell_config["cli_tools"][0]["install_commands"] == ["apt-get install -y ripgrep"] assert shell_config["env"][0] == {"name": "PROJECT_NAME", "value": "demo"} - assert shell_config["sandbox"] == {"provider": "independent", "config": {"cpu": 2}} + assert "sandbox" not in shell_config assert result.metadata["agent_tools"] == { "dify_tool_count": 0, "dify_tool_names": [], @@ -520,7 +522,7 @@ def test_build_shell_layer_config_accepts_legacy_fallback_keys(): {"name": "API_KEY", "ref": "credential-2"}, {"name": "LEGACY_SECRET_REF", "ref": "credential-3"}, ] - assert config["sandbox"] is None + assert "sandbox" not in config def test_build_shell_layer_config_maps_typed_command_field(): @@ -742,7 +744,13 @@ def test_build_maps_agent_soul_knowledge_to_knowledge_layer_config(): "conditions": { "logical_operator": "and", "conditions": [ - {"name": "category", "comparison_operator": "contains", "value": "auth"} + { + "id": "cond-1", + "metadata_id": "meta-1", + "name": "category", + "comparison_operator": "contains", + "value": "auth", + } ], }, }, @@ -1459,7 +1467,10 @@ def test_workflow_run_request_has_config_layer_with_empty_agent_soul(monkeypatch "mentioned_skill_names": [], "mentioned_file_names": [], } - assert layers[DIFY_SHELL_LAYER_ID]["deps"] == {"execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID} + assert layers[DIFY_SHELL_LAYER_ID]["deps"] == { + "execution_context": DIFY_EXECUTION_CONTEXT_LAYER_ID, + "runtime": "runtime", + } assert layers[DIFY_SHELL_LAYER_ID]["config"]["agent_stub_drive_ref"] is None @@ -1474,7 +1485,7 @@ def test_workflow_run_request_contains_config_layer(): layer_names = [layer["name"] for layer in dumped["composition"]["layers"]] assert DIFY_CONFIG_LAYER_ID in layer_names # shell enters first; config uses that shell to materialize mentioned targets. - assert layer_names.index(DIFY_SHELL_LAYER_ID) == layer_names.index("execution_context") + 1 + assert layer_names.index(DIFY_SHELL_LAYER_ID) == layer_names.index("execution_context") + 2 assert layer_names.index(DIFY_CONFIG_LAYER_ID) == layer_names.index(DIFY_SHELL_LAYER_ID) + 1 config = next(layer for layer in dumped["composition"]["layers"] if layer["name"] == DIFY_CONFIG_LAYER_ID) assert config["type"] == "dify.config" @@ -1498,11 +1509,6 @@ def test_workflow_run_request_contains_config_layer(): } warnings = result.metadata["runtime_support"]["unsupported_runtime_warnings"] assert warnings == [] - # the config layer is non-sensitive and must survive into persistable specs - from dify_agent.protocol import extract_runtime_layer_specs - - specs = extract_runtime_layer_specs(result.request.composition) - assert any(spec.name == DIFY_CONFIG_LAYER_ID and spec.type == "dify.config" for spec in specs) def test_workflow_runtime_expands_config_mentions_in_agent_soul_prompt(): diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_session_cleanup_layer.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_session_cleanup_layer.py deleted file mode 100644 index eda2962198e..00000000000 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_session_cleanup_layer.py +++ /dev/null @@ -1,256 +0,0 @@ -from typing import cast - -import pytest -from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot -from agenton.layers.base import LifecycleState -from dify_agent.protocol import RuntimeLayerSpec - -from core.workflow.nodes.agent_v2.session_cleanup_layer import ( - WorkflowAgentSessionCleanupLayer, - build_workflow_agent_session_cleanup_layer, -) -from core.workflow.nodes.agent_v2.session_store import ( - StoredWorkflowAgentSession, - WorkflowAgentRuntimeSessionStore, - WorkflowAgentSessionScope, -) -from core.workflow.system_variables import build_system_variables -from graphon.entities.pause_reason import SchedulingPause -from graphon.graph_engine.command_channels import CommandChannel -from graphon.graph_events import ( - GraphRunAbortedEvent, - GraphRunFailedEvent, - GraphRunPartialSucceededEvent, - GraphRunPausedEvent, - GraphRunStartedEvent, - GraphRunSucceededEvent, -) -from graphon.runtime import GraphRuntimeState, ReadOnlyGraphRuntimeStateWrapper, VariablePool - - -def _layer_snapshot(name: str) -> LayerSessionSnapshot: - return LayerSessionSnapshot( - name=name, - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ) - - -def _default_scope() -> WorkflowAgentSessionScope: - return WorkflowAgentSessionScope( - tenant_id="tenant-1", - app_id="app-1", - workflow_id="workflow-1", - workflow_run_id="workflow-run-1", - node_id="agent-node", - node_execution_id="node-exec-1", - binding_id="binding-1", - agent_id="agent-1", - agent_config_snapshot_id="snapshot-1", - ) - - -def _stored_session(scope: WorkflowAgentSessionScope, *, index: int = 1) -> StoredWorkflowAgentSession: - return StoredWorkflowAgentSession( - scope=scope, - session_snapshot=CompositorSessionSnapshot( - layers=[ - _layer_snapshot("workflow_node_job_prompt"), - _layer_snapshot("execution_context"), - _layer_snapshot("history"), - _layer_snapshot("llm"), - ] - ), - backend_run_id=f"agent-run-{index}", - runtime_layer_specs=[ - RuntimeLayerSpec(name="workflow_node_job_prompt", type="plain.prompt", config={"prefix": "ok"}), - RuntimeLayerSpec(name="execution_context", type="dify.execution_context", config={"tenant_id": "t"}), - RuntimeLayerSpec(name="history", type="pydantic_ai.history"), - ], - ) - - -class FakeSessionStore: - def __init__(self, *, stored: list[StoredWorkflowAgentSession] | None = None) -> None: - self._stored = stored if stored is not None else [_stored_session(_default_scope())] - self.list_calls: list[str] = [] - self.cleaned: list[tuple[WorkflowAgentSessionScope, str | None]] = [] - - def list_active_sessions(self, *, workflow_run_id: str) -> list[StoredWorkflowAgentSession]: - self.list_calls.append(workflow_run_id) - return list(self._stored) - - def mark_cleaned(self, *, scope: WorkflowAgentSessionScope, backend_run_id: str | None = None) -> None: - self.cleaned.append((scope, backend_run_id)) - - -def _build_layer(*, session_store: FakeSessionStore) -> WorkflowAgentSessionCleanupLayer: - variable_pool = VariablePool.from_bootstrap( - system_variables=build_system_variables(workflow_execution_id="workflow-run-1"), - user_inputs={}, - conversation_variables=[], - ) - runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=0.0) - layer = WorkflowAgentSessionCleanupLayer( - session_store=cast(WorkflowAgentRuntimeSessionStore, session_store), - ) - layer.initialize(ReadOnlyGraphRuntimeStateWrapper(runtime_state), cast(CommandChannel, object())) - return layer - - -@pytest.mark.parametrize( - "terminal_event", - [ - GraphRunSucceededEvent(outputs={}), - GraphRunPartialSucceededEvent(exceptions_count=1, outputs={}), - GraphRunFailedEvent(error="boom"), - GraphRunAbortedEvent(reason="user cancelled", outputs={}), - ], - ids=["succeeded", "partial_succeeded", "failed", "aborted"], -) -def test_cleanup_layer_enqueues_cleanup_and_marks_cleaned_on_terminal_events(monkeypatch, terminal_event): - session_store = FakeSessionStore() - queued_payloads: list[dict[str, object]] = [] - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.session_cleanup_layer.cleanup_workflow_agent_runtime_session.delay", - lambda payload: queued_payloads.append(payload), - ) - layer = _build_layer(session_store=session_store) - - layer.on_event(terminal_event) - - assert session_store.list_calls == ["workflow-run-1"] - assert len(queued_payloads) == 1 - assert queued_payloads[0]["metadata"]["workflow_run_id"] == "workflow-run-1" - assert queued_payloads[0]["metadata"]["previous_agent_backend_run_id"] == "agent-run-1" - assert session_store.cleaned == [(_default_scope(), "agent-run-1")] - - -@pytest.mark.parametrize( - "non_terminal_event", - [ - GraphRunStartedEvent(), - GraphRunPausedEvent(reasons=[SchedulingPause(message="awaiting human input")], outputs={}), - ], - ids=["started", "paused"], -) -def test_cleanup_layer_ignores_non_terminal_events(monkeypatch, non_terminal_event): - session_store = FakeSessionStore() - queued_payloads: list[dict[str, object]] = [] - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.session_cleanup_layer.cleanup_workflow_agent_runtime_session.delay", - lambda payload: queued_payloads.append(payload), - ) - layer = _build_layer(session_store=session_store) - - layer.on_event(non_terminal_event) - - assert session_store.list_calls == [] - assert queued_payloads == [] - assert session_store.cleaned == [] - - -def test_cleanup_layer_marks_cleaned_even_when_specs_are_missing(monkeypatch, caplog: pytest.LogCaptureFixture): - scope = _default_scope() - session_store = FakeSessionStore( - stored=[ - StoredWorkflowAgentSession( - scope=scope, - session_snapshot=CompositorSessionSnapshot(layers=[_layer_snapshot("history")]), - backend_run_id="legacy-run", - runtime_layer_specs=[], - ) - ] - ) - queued_payloads: list[dict[str, object]] = [] - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.session_cleanup_layer.cleanup_workflow_agent_runtime_session.delay", - lambda payload: queued_payloads.append(payload), - ) - layer = _build_layer(session_store=session_store) - - layer.on_event(GraphRunSucceededEvent(outputs={})) - - assert queued_payloads == [] - assert session_store.cleaned == [(scope, "legacy-run")] - assert any("no runtime_layer_specs persisted" in record.message for record in caplog.records) - - -def test_cleanup_layer_marks_cleaned_even_when_enqueue_fails(monkeypatch): - session_store = FakeSessionStore() - - def _explode(_payload: dict[str, object]) -> None: - raise RuntimeError("queue down") - - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.session_cleanup_layer.cleanup_workflow_agent_runtime_session.delay", - _explode, - ) - layer = _build_layer(session_store=session_store) - - layer.on_event(GraphRunSucceededEvent(outputs={})) - - assert session_store.cleaned == [(_default_scope(), "agent-run-1")] - - -def test_cleanup_layer_does_not_raise_when_mark_cleaned_fails(monkeypatch): - session_store = FakeSessionStore() - - def _explode(*, scope: WorkflowAgentSessionScope, backend_run_id: str | None = None) -> None: - del scope, backend_run_id - raise RuntimeError("cleanup bookkeeping failed") - - monkeypatch.setattr(session_store, "mark_cleaned", _explode) - layer = _build_layer(session_store=session_store) - - layer.on_event(GraphRunSucceededEvent(outputs={})) - - -def test_cleanup_layer_fans_out_to_every_active_session(monkeypatch): - scopes = [ - WorkflowAgentSessionScope( - tenant_id="tenant-1", - app_id="app-1", - workflow_id="workflow-1", - workflow_run_id="workflow-run-1", - node_id=f"agent-node-{i}", - node_execution_id=f"node-exec-{i}", - binding_id=f"binding-{i}", - agent_id=f"agent-{i}", - agent_config_snapshot_id=f"snapshot-{i}", - ) - for i in range(3) - ] - session_store = FakeSessionStore(stored=[_stored_session(scope, index=i) for i, scope in enumerate(scopes, 1)]) - queued_payloads: list[dict[str, object]] = [] - monkeypatch.setattr( - "core.workflow.nodes.agent_v2.session_cleanup_layer.cleanup_workflow_agent_runtime_session.delay", - lambda payload: queued_payloads.append(payload), - ) - layer = _build_layer(session_store=session_store) - - layer.on_event(GraphRunSucceededEvent(outputs={})) - - assert len(queued_payloads) == 3 - assert [entry[0] for entry in session_store.cleaned] == scopes - - -def test_cleanup_layer_skips_when_workflow_run_id_missing(caplog: pytest.LogCaptureFixture): - session_store = FakeSessionStore() - variable_pool = VariablePool.from_bootstrap(system_variables={}, user_inputs={}, conversation_variables=[]) - runtime_state = GraphRuntimeState(variable_pool=variable_pool, start_at=0.0) - layer = WorkflowAgentSessionCleanupLayer(session_store=cast(WorkflowAgentRuntimeSessionStore, session_store)) - layer.initialize(ReadOnlyGraphRuntimeStateWrapper(runtime_state), cast(CommandChannel, object())) - - layer.on_event(GraphRunSucceededEvent(outputs={})) - - assert session_store.list_calls == [] - assert session_store.cleaned == [] - assert any("workflow_run_id is missing" in record.message for record in caplog.records) - - -def test_build_workflow_agent_session_cleanup_layer_returns_layer() -> None: - layer = build_workflow_agent_session_cleanup_layer() - - assert isinstance(layer, WorkflowAgentSessionCleanupLayer) diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_session_store.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_session_store.py index a57b2adc168..193049e5fb9 100644 --- a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_session_store.py +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_session_store.py @@ -1,288 +1,434 @@ -"""Unit tests for :mod:`core.workflow.nodes.agent_v2.session_store`. - -Uses the in-memory SQLite engine configured by the project conftest plus a -per-test ``CREATE TABLE`` so the real ORM round-trip exercises every store -method. Keeps the suite self-contained — no Postgres / Docker required — while -still hitting the actual ``session_factory`` code path that production uses. -""" - -from __future__ import annotations - -from collections.abc import Generator +import json +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import MagicMock import pytest from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot -from agenton.layers.base import LifecycleState -from dify_agent.protocol import RuntimeLayerSpec -from sqlalchemy import delete +from sqlalchemy.orm import Session -from core.db.session_factory import session_factory -from core.workflow.nodes.agent_v2.session_store import ( - StoredWorkflowAgentSession, - WorkflowAgentRuntimeSessionStore, - WorkflowAgentSessionScope, +from core.workflow.nodes.agent_v2.session_store import WorkflowAgentSessionScope, WorkflowAgentWorkspaceStore +from models.agent import ( + AgentConfigVersionKind, + AgentWorkingResourceStatus, + AgentWorkspace, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, ) -from models.agent import WorkflowAgentRuntimeSession, WorkflowAgentRuntimeSessionStatus +from models.workflow import WorkflowNodeExecutionModel +from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService +from services.agent_app_sandbox_service import WorkflowAgentSandboxService -def _scope(workflow_run_id: str | None = "wfr-1", binding_id: str = "binding-1") -> WorkflowAgentSessionScope: +def _scope() -> WorkflowAgentSessionScope: return WorkflowAgentSessionScope( tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", - workflow_run_id=workflow_run_id, - node_id="agent-node", - node_execution_id="node-exec-1", - binding_id=binding_id, + workflow_run_id="run-1", + node_id="node-1", + node_execution_id="execution-1", + workflow_agent_binding_id="workflow-binding-1", agent_id="agent-1", - agent_config_snapshot_id="snapshot-1", + agent_config_snapshot_id="config-1", ) -def _snapshot(messages: int = 1) -> CompositorSessionSnapshot: - return CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot( - name="history", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={"messages": [{"role": "user", "content": f"m{i}"} for i in range(messages)]}, - ) - ] +def _binding() -> SimpleNamespace: + return SimpleNamespace( + id="binding-1", + workspace_id="workspace-1", + agent_id="agent-1", + base_home_snapshot_id="home-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + backend_binding_ref="backend-binding-1", + session_snapshot=None, + pending_form_id=None, + pending_tool_call_id=None, ) -def _specs() -> list[RuntimeLayerSpec]: - return [ - RuntimeLayerSpec(name="workflow_node_job_prompt", type="plain.prompt", config={"prefix": "ok"}), - RuntimeLayerSpec(name="history", type="pydantic_ai.history"), - ] - - -@pytest.fixture(autouse=True) -def _create_table() -> Generator[None, None, None]: - """Create the lifecycle table on the in-memory SQLite engine, drop after.""" - engine = session_factory.get_session_maker().kw["bind"] - WorkflowAgentRuntimeSession.__table__.create(bind=engine, checkfirst=True) - yield - with session_factory.create_session() as session: - session.execute(delete(WorkflowAgentRuntimeSession)) - session.commit() - WorkflowAgentRuntimeSession.__table__.drop(bind=engine, checkfirst=True) - - -def test_load_active_snapshot_returns_none_when_scope_has_no_workflow_run_id(): - """``workflow_run_id`` is the keying column; no row can match without it.""" - store = WorkflowAgentRuntimeSessionStore() - assert store.load_active_snapshot(_scope(workflow_run_id=None)) is None - - -def test_load_active_snapshot_returns_none_when_no_row_matches(): - store = WorkflowAgentRuntimeSessionStore() - assert store.load_active_snapshot(_scope()) is None - - -def test_save_active_snapshot_creates_row_and_load_round_trips(): - store = WorkflowAgentRuntimeSessionStore() - snapshot = _snapshot(messages=2) - store.save_active_snapshot(scope=_scope(), backend_run_id="run-1", snapshot=snapshot, runtime_layer_specs=_specs()) - - loaded = store.load_active_snapshot(_scope()) - assert loaded is not None - assert len(loaded.layers) == 1 - assert loaded.layers[0].name == "history" - assert loaded.layers[0].runtime_state["messages"] == snapshot.layers[0].runtime_state["messages"] - with session_factory.create_session() as session: - row = session.query(WorkflowAgentRuntimeSession).one() - assert "workflow_node_job_prompt" in row.composition_layer_specs - assert "history" in row.composition_layer_specs - - -def test_save_active_snapshot_skips_when_workflow_run_id_missing(): - """Without a workflow_run_id the row cannot be keyed; save is a no-op.""" - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(workflow_run_id=None), - backend_run_id="run-skipped", - snapshot=_snapshot(), - runtime_layer_specs=_specs(), - ) - with session_factory.create_session() as session: - assert session.query(WorkflowAgentRuntimeSession).count() == 0 - - -def test_save_active_snapshot_skips_when_snapshot_missing(): - """A run that produced no snapshot (e.g. failed agent run) does not write.""" - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-empty", - snapshot=None, - runtime_layer_specs=_specs(), - ) - with session_factory.create_session() as session: - assert session.query(WorkflowAgentRuntimeSession).count() == 0 - - -def test_save_active_snapshot_updates_existing_row_on_re_entry(): - """A second save under the same scope must update in place, not insert.""" - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(messages=1), - runtime_layer_specs=_specs(), - ) - # Second call with new snapshot + backend_run_id. - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-2", - snapshot=_snapshot(messages=2), - runtime_layer_specs=_specs(), +def _workspace_row( + *, + workspace_id: str = "workspace-1", + tenant_id: str = "tenant-1", + app_id: str = "app-1", + owner_scope_key: str = "node-1:workflow-binding-1", + status: AgentWorkingResourceStatus = AgentWorkingResourceStatus.ACTIVE, +) -> AgentWorkspace: + return AgentWorkspace( + id=workspace_id, + tenant_id=tenant_id, + app_id=app_id, + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id="run-1", + owner_scope_key=owner_scope_key, + backend_workspace_ref=f"{workspace_id}-ref", + status=status, + active_guard=1 if status is AgentWorkingResourceStatus.ACTIVE else None, ) - with session_factory.create_session() as session: - rows = session.query(WorkflowAgentRuntimeSession).all() - assert len(rows) == 1 - assert rows[0].backend_run_id == "run-2" - assert rows[0].status == WorkflowAgentRuntimeSessionStatus.ACTIVE - assert rows[0].cleaned_at is None - -def test_save_active_snapshot_resurrects_cleaned_row(): - """If a prior cleanup retired the row, a re-entry flips it back to ACTIVE.""" - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(), - runtime_layer_specs=_specs(), - ) - store.mark_cleaned(scope=_scope(), backend_run_id="cleanup-1") - # Save again — the existing row was CLEANED; should be revived. - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-2", - snapshot=_snapshot(messages=3), - runtime_layer_specs=_specs(), +def _binding_row( + *, + binding_id: str = "binding-1", + workspace_id: str = "workspace-1", + status: AgentWorkingResourceStatus = AgentWorkingResourceStatus.ACTIVE, +) -> AgentWorkspaceBinding: + return AgentWorkspaceBinding( + id=binding_id, + tenant_id="tenant-1", + app_id="app-1", + workspace_id=workspace_id, + agent_id="agent-1", + base_home_snapshot_id="home-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + backend_binding_ref=f"{binding_id}-ref", + status=status, ) - with session_factory.create_session() as session: - rows = session.query(WorkflowAgentRuntimeSession).all() - assert len(rows) == 1 - assert rows[0].status == WorkflowAgentRuntimeSessionStatus.ACTIVE - assert rows[0].cleaned_at is None - assert rows[0].backend_run_id == "run-2" + +def test_scope_uses_node_and_workflow_binding_as_workspace_subscope() -> None: + owner = _scope().workspace_owner + assert owner.owner_type is AgentWorkspaceOwnerType.WORKFLOW_RUN + assert owner.owner_id == "run-1" + assert owner.owner_scope_key == "node-1:workflow-binding-1" -def test_list_active_sessions_returns_specs_and_snapshot(): - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(binding_id="binding-A"), - backend_run_id="run-A", - snapshot=_snapshot(), - runtime_layer_specs=_specs(), +def test_load_existing_scope_reads_the_generation_from_the_persisted_binding( + monkeypatch: pytest.MonkeyPatch, +) -> None: + execution = SimpleNamespace( + agent_workspace_binding_id="binding-1", + process_data_dict={"workflow_agent_binding_id": "workflow-binding-1"}, ) - store.save_active_snapshot( - scope=_scope(binding_id="binding-B"), - backend_run_id="run-B", - snapshot=_snapshot(messages=2), - runtime_layer_specs=_specs(), + context = MagicMock() + store = WorkflowAgentWorkspaceStore() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr(store, "_load_execution_by_identity", MagicMock(return_value=execution)) + get_active = MagicMock(return_value=_binding()) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active) + + scope = store.load_existing_node_execution_scope( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_run_id="run-1", + node_id="node-1", + node_execution_id="execution-1", ) - listed = store.list_active_sessions(workflow_run_id="wfr-1") - assert {s.backend_run_id for s in listed} == {"run-A", "run-B"} - by_run = {s.backend_run_id: s for s in listed} - assert isinstance(by_run["run-A"], StoredWorkflowAgentSession) - # Specs round-trip through pydantic TypeAdapter — ensure deserialize works. - assert by_run["run-A"].runtime_layer_specs[0].name == "workflow_node_job_prompt" - assert by_run["run-A"].runtime_layer_specs[1].type == "pydantic_ai.history" - # node_execution_id default-replaces NULL with "" when the DB column is None. - assert by_run["run-A"].scope.node_execution_id == "node-exec-1" + assert scope is not None + assert scope.workflow_agent_binding_id == "workflow-binding-1" + assert scope.agent_id == "agent-1" + assert scope.agent_config_snapshot_id == "config-1" + assert get_active.call_args.kwargs["binding_id"] == "binding-1" -def test_list_active_sessions_skips_cleaned_rows(): - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(binding_id="binding-A"), - backend_run_id="run-A", - snapshot=_snapshot(), - runtime_layer_specs=_specs(), +def test_load_existing_scope_rejects_unavailable_persisted_binding(monkeypatch: pytest.MonkeyPatch) -> None: + execution = SimpleNamespace( + agent_workspace_binding_id="binding-missing", + process_data_dict={"workflow_agent_binding_id": "workflow-binding-1"}, ) - store.save_active_snapshot( - scope=_scope(binding_id="binding-B"), - backend_run_id="run-B", - snapshot=_snapshot(), - runtime_layer_specs=_specs(), + context = MagicMock() + store = WorkflowAgentWorkspaceStore() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, ) - store.mark_cleaned(scope=_scope(binding_id="binding-A"), backend_run_id="cleanup-A") + monkeypatch.setattr(store, "_load_execution_by_identity", MagicMock(return_value=execution)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=None)) - listed = store.list_active_sessions(workflow_run_id="wfr-1") - assert {s.backend_run_id for s in listed} == {"run-B"} + with pytest.raises(AgentWorkspaceNotFoundError, match="participant Binding is unavailable"): + store.load_existing_node_execution_scope( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_run_id="run-1", + node_id="node-1", + node_execution_id="execution-1", + ) -def test_list_active_sessions_handles_legacy_rows_without_specs(): - """Rows persisted before runtime_layer_specs landed have an empty string.""" - # Insert a legacy-shape row directly: empty specs payload simulates a row - # written before the spec persistence feature landed in A.1. - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-legacy", - snapshot=_snapshot(), - runtime_layer_specs=[], +@pytest.mark.parametrize("home_snapshot_id", ["home-1", None]) +def test_load_or_create_persists_binding_on_node_execution(monkeypatch, home_snapshot_id: str | None) -> None: + execution = WorkflowNodeExecutionModel( + agent_workspace_binding_id=None, + process_data=json.dumps({"existing": "value"}), ) - listed = store.list_active_sessions(workflow_run_id="wfr-1") - assert len(listed) == 1 - assert listed[0].runtime_layer_specs == [] - - -def test_mark_cleaned_sets_status_and_cleaned_at_with_backend_run_id(): - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(), - runtime_layer_specs=_specs(), + context = MagicMock() + session = context.__enter__.return_value + create = MagicMock(return_value=_binding()) + store = WorkflowAgentWorkspaceStore() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, ) - store.mark_cleaned(scope=_scope(), backend_run_id="cleanup-1") + monkeypatch.setattr(store, "_load_execution", MagicMock(return_value=execution)) + monkeypatch.setattr(AgentWorkspaceService, "create_binding", create) - with session_factory.create_session() as session: - row = session.query(WorkflowAgentRuntimeSession).one() - assert row.status == WorkflowAgentRuntimeSessionStatus.CLEANED - assert row.cleaned_at is not None - assert row.backend_run_id == "cleanup-1" + stored = store.load_or_create_node_execution_session(_scope(), home_snapshot_id=home_snapshot_id) + assert stored.binding_id == "binding-1" + assert stored.workspace_id == "workspace-1" + assert stored.backend_binding_ref == "backend-binding-1" + assert execution.agent_workspace_binding_id == "binding-1" + assert execution.process_data_dict == { + "existing": "value", + "workflow_agent_binding_id": "workflow-binding-1", + } + assert "agent_workspace_binding_id" not in execution.process_data_dict + assert create.call_args.kwargs["session"] is session + assert create.call_args.kwargs["base_home_snapshot_id"] == home_snapshot_id + session.commit.assert_called_once_with() -def test_mark_cleaned_preserves_existing_backend_run_id_when_none_given(): - """``backend_run_id=None`` means "leave the previous one in place".""" - store = WorkflowAgentRuntimeSessionStore() - store.save_active_snapshot( - scope=_scope(), - backend_run_id="run-1", - snapshot=_snapshot(), - runtime_layer_specs=_specs(), + get_active = MagicMock(return_value=_binding()) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active) + session.scalar.return_value = execution + resolved = WorkflowAgentSandboxService._resolve_binding( + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + node_id="node-1", + node_execution_id="execution-1", + session=session, ) - store.mark_cleaned(scope=_scope(), backend_run_id=None) - with session_factory.create_session() as session: - row = session.query(WorkflowAgentRuntimeSession).one() - assert row.status == WorkflowAgentRuntimeSessionStatus.CLEANED - assert row.backend_run_id == "run-1" + assert resolved.backend_binding_ref == "backend-binding-1" + assert resolved.agent_id == "agent-1" + assert resolved.agent_config_version_id == "config-1" + assert resolved.agent_config_version_kind == "snapshot" + owner_scope = get_active.call_args.kwargs["expected_owner_scope"] + assert owner_scope.owner_scope_key == "node-1:workflow-binding-1" + session.rollback.assert_called_once_with() -def test_mark_cleaned_is_a_noop_when_no_active_row(): - """No matching ACTIVE row → no-op (already-cleaned rows are not re-touched).""" - store = WorkflowAgentRuntimeSessionStore() - store.mark_cleaned(scope=_scope(), backend_run_id="cleanup-1") - with session_factory.create_session() as session: - assert session.query(WorkflowAgentRuntimeSession).count() == 0 +def test_load_existing_pointer_rejects_missing_workflow_identity(monkeypatch: pytest.MonkeyPatch) -> None: + execution = SimpleNamespace( + agent_workspace_binding_id="binding-1", + process_data=json.dumps({"existing": "value"}), + process_data_dict={"existing": "value"}, + ) + context = MagicMock() + session = context.__enter__.return_value + store = WorkflowAgentWorkspaceStore() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr(store, "_load_execution", MagicMock(return_value=execution)) + get_active = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active) + + with pytest.raises(AgentWorkspaceNotFoundError, match="caller identity is missing"): + store.load_or_create_node_execution_session(_scope(), home_snapshot_id="home-1") + + assert json.loads(execution.process_data) == {"existing": "value"} + get_active.assert_not_called() + session.commit.assert_not_called() -def test_mark_cleaned_is_a_noop_when_workflow_run_id_missing(): - """Without a workflow_run_id we cannot key the row; ignore the call.""" - store = WorkflowAgentRuntimeSessionStore() - store.mark_cleaned(scope=_scope(workflow_run_id=None), backend_run_id="cleanup-1") - # Sanity — no rows created or touched. - with session_factory.create_session() as session: - assert session.query(WorkflowAgentRuntimeSession).count() == 0 +def test_load_existing_pointer_reuses_matching_workflow_identity(monkeypatch: pytest.MonkeyPatch) -> None: + original_process_data = json.dumps( + { + "existing": "value", + "workflow_agent_binding_id": "workflow-binding-1", + } + ) + execution = SimpleNamespace( + agent_workspace_binding_id="binding-1", + process_data=original_process_data, + process_data_dict=json.loads(original_process_data), + ) + context = MagicMock() + session = context.__enter__.return_value + store = WorkflowAgentWorkspaceStore() + create = MagicMock() + binding = _binding() + validate_generation = MagicMock() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr(store, "_load_execution", MagicMock(return_value=execution)) + monkeypatch.setattr(AgentWorkspaceService, "create_binding", create) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=binding)) + monkeypatch.setattr(AgentWorkspaceService, "validate_binding_generation", validate_generation) + + stored = store.load_or_create_node_execution_session(_scope(), home_snapshot_id="home-1") + + assert stored.binding_id == "binding-1" + assert execution.process_data == original_process_data + assert "agent_workspace_binding_id" not in execution.process_data_dict + create.assert_not_called() + session.commit.assert_not_called() + validate_generation.assert_called_once_with( + binding, + base_home_snapshot_id="home-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + + +def test_load_existing_pointer_rejects_conflicting_workflow_identity(monkeypatch: pytest.MonkeyPatch) -> None: + execution = SimpleNamespace( + agent_workspace_binding_id="binding-1", + process_data=json.dumps({"workflow_agent_binding_id": "workflow-binding-other"}), + process_data_dict={"workflow_agent_binding_id": "workflow-binding-other"}, + ) + context = MagicMock() + session = context.__enter__.return_value + store = WorkflowAgentWorkspaceStore() + get_active = MagicMock() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr(store, "_load_execution", MagicMock(return_value=execution)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active) + + with pytest.raises(AgentWorkspaceNotFoundError, match="caller identity does not match"): + store.load_or_create_node_execution_session(_scope(), home_snapshot_id="home-1") + + get_active.assert_not_called() + session.commit.assert_not_called() + + +def test_load_or_create_fails_before_binding_create_when_caller_row_is_missing(monkeypatch: pytest.MonkeyPatch) -> None: + context = MagicMock() + session = context.__enter__.return_value + create = MagicMock() + sleep = MagicMock() + store = WorkflowAgentWorkspaceStore() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr("core.workflow.nodes.agent_v2.session_store.time.sleep", sleep) + monkeypatch.setattr(session, "scalar", MagicMock(return_value=None)) + monkeypatch.setattr(AgentWorkspaceService, "create_binding", create) + + with pytest.raises(AgentWorkspaceNotFoundError, match="Workflow node execution caller is unavailable"): + store.load_or_create_node_execution_session(_scope(), home_snapshot_id="home-1") + + assert session.scalar.call_count == 60 + assert sleep.call_count == 59 + create.assert_not_called() + session.commit.assert_not_called() + + +def test_load_existing_scope_waits_for_caller_row_to_become_visible(monkeypatch: pytest.MonkeyPatch) -> None: + execution = SimpleNamespace(agent_workspace_binding_id=None) + context = MagicMock() + session = context.__enter__.return_value + session.scalar.side_effect = [None, None, execution] + sleep = MagicMock() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr("core.workflow.nodes.agent_v2.session_store.time.sleep", sleep) + + scope = WorkflowAgentWorkspaceStore().load_existing_node_execution_scope( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_run_id="run-1", + node_id="node-1", + node_execution_id="execution-1", + ) + + assert scope is None + assert session.scalar.call_count == 3 + assert sleep.call_count == 2 + + +def test_save_snapshot_targets_binding(monkeypatch: pytest.MonkeyPatch) -> None: + save = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "save_binding_session_snapshot", save) + snapshot = CompositorSessionSnapshot(layers=[]) + + WorkflowAgentWorkspaceStore().save_active_snapshot(scope=_scope(), binding_id="binding-1", snapshot=snapshot) + + assert save.call_args.kwargs["binding_id"] == "binding-1" + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_retire_workflow_run_only_retires_matching_tenant_and_app( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + matching = _workspace_row() + other_tenant = _workspace_row(workspace_id="workspace-other-tenant", tenant_id="tenant-2") + other_app = _workspace_row( + workspace_id="workspace-other-app", + app_id="app-2", + owner_scope_key="node-2:workflow-binding-2", + ) + sqlite_session.add_all([matching, other_tenant, other_app]) + sqlite_session.commit() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: nullcontext(sqlite_session), + ) + + workspace_ids = WorkflowAgentWorkspaceStore().retire_workflow_run( + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + ) + + assert matching.status is AgentWorkingResourceStatus.RETIRED + assert other_tenant.status is AgentWorkingResourceStatus.ACTIVE + assert other_app.status is AgentWorkingResourceStatus.ACTIVE + assert workspace_ids == [matching.id] + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_retire_workflow_run_transitions_active_workspace( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + workspace = _workspace_row() + binding = _binding_row() + sqlite_session.add_all([workspace, binding]) + sqlite_session.commit() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: nullcontext(sqlite_session), + ) + + workspace_ids = WorkflowAgentWorkspaceStore().retire_workflow_run( + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + ) + + assert workspace.status is AgentWorkingResourceStatus.RETIRED + assert binding.status is AgentWorkingResourceStatus.RETIRED + assert workspace_ids == [workspace.id] + + +def test_retire_workflow_run_returns_existing_retired_workspace(monkeypatch: pytest.MonkeyPatch) -> None: + workspace = _workspace_row(status=AgentWorkingResourceStatus.RETIRED) + context = MagicMock() + session = context.__enter__.return_value + session.scalars.return_value.all.return_value = [workspace] + retire = MagicMock() + monkeypatch.setattr( + "core.workflow.nodes.agent_v2.session_store.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr(AgentWorkspaceService, "retire_workspace", retire) + + workspace_ids = WorkflowAgentWorkspaceStore().retire_workflow_run( + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + ) + + retire.assert_not_called() + assert workspace_ids == [workspace.id] diff --git a/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_workspace_retirement_layer.py b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_workspace_retirement_layer.py new file mode 100644 index 00000000000..b5298c23235 --- /dev/null +++ b/api/tests/unit_tests/core/workflow/nodes/agent_v2/test_workspace_retirement_layer.py @@ -0,0 +1,75 @@ +from unittest.mock import MagicMock + +import pytest + +from core.app.entities.app_invoke_entities import DifyRunContext, InvokeFrom, UserFrom +from core.workflow.nodes.agent_v2 import workspace_retirement_layer as layer_module +from core.workflow.nodes.agent_v2.workspace_retirement_layer import WorkflowAgentWorkspaceRetirementLayer +from graphon.graph_events import GraphRunSucceededEvent, NodeRunStartedEvent + + +def _run_context() -> DifyRunContext: + return DifyRunContext( + tenant_id="tenant-1", + app_id="app-1", + user_id="account-1", + user_from=UserFrom.ACCOUNT, + invoke_from=InvokeFrom.DEBUGGER, + ) + + +def test_terminal_event_retires_workflow_workspace(monkeypatch: pytest.MonkeyPatch) -> None: + store = MagicMock() + events: list[str] = [] + store.retire_workflow_run.side_effect = lambda **_kwargs: events.append("retire") or ["workspace-1"] + enqueue = MagicMock(side_effect=lambda **_kwargs: events.append("enqueue")) + monkeypatch.setattr(layer_module, "WorkflowAgentWorkspaceStore", MagicMock(return_value=store)) + monkeypatch.setattr(layer_module, "enqueue_agent_resource_collection", enqueue) + layer = WorkflowAgentWorkspaceRetirementLayer(dify_run_context=_run_context()) + layer.initialize(MagicMock(), MagicMock()) + monkeypatch.setattr(layer_module, "get_system_text", lambda *_: "workflow-run-1") + + layer.on_event(MagicMock(spec=GraphRunSucceededEvent)) + + store.retire_workflow_run.assert_called_once_with( + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="workflow-run-1", + ) + enqueue.assert_called_once_with(tenant_id="tenant-1", workspace_ids=["workspace-1"]) + assert events == ["retire", "enqueue"] + + +def test_non_terminal_event_does_not_retire_workspace(monkeypatch: pytest.MonkeyPatch) -> None: + store = MagicMock() + monkeypatch.setattr(layer_module, "WorkflowAgentWorkspaceStore", MagicMock(return_value=store)) + layer = WorkflowAgentWorkspaceRetirementLayer(dify_run_context=_run_context()) + + layer.on_event(MagicMock(spec=NodeRunStartedEvent)) + + store.retire_workflow_run.assert_not_called() + + +def test_terminal_retirement_failure_does_not_replace_terminal_event(monkeypatch: pytest.MonkeyPatch) -> None: + store = MagicMock() + store.retire_workflow_run.side_effect = RuntimeError("database unavailable") + log_exception = MagicMock() + enqueue = MagicMock() + monkeypatch.setattr(layer_module, "WorkflowAgentWorkspaceStore", MagicMock(return_value=store)) + monkeypatch.setattr(layer_module.logger, "exception", log_exception) + monkeypatch.setattr(layer_module, "enqueue_agent_resource_collection", enqueue) + layer = WorkflowAgentWorkspaceRetirementLayer(dify_run_context=_run_context()) + layer.initialize(MagicMock(), MagicMock()) + monkeypatch.setattr(layer_module, "get_system_text", lambda *_: "workflow-run-1") + + layer.on_event(MagicMock(spec=GraphRunSucceededEvent)) + + log_exception.assert_called_once_with( + "Failed to retire Workflow Agent Workspaces", + extra={ + "tenant_id": "tenant-1", + "app_id": "app-1", + "workflow_run_id": "workflow-run-1", + }, + ) + enqueue.assert_not_called() diff --git a/api/tests/unit_tests/core/workflow/nodes/human_input/test_human_input_form_filled_event.py b/api/tests/unit_tests/core/workflow/nodes/human_input/test_human_input_form_filled_event.py index a1cb0af85fe..33d1a41fad8 100644 --- a/api/tests/unit_tests/core/workflow/nodes/human_input/test_human_input_form_filled_event.py +++ b/api/tests/unit_tests/core/workflow/nodes/human_input/test_human_input_form_filled_event.py @@ -73,13 +73,15 @@ def _create_human_input_node( node_data=node_data, file_reference_factory=_TestFileReferenceFactory(), ) - return HumanInputNode( + node = HumanInputNode( node_id=config["id"], data=node_data, graph_init_params=graph_init_params, graph_runtime_state=graph_runtime_state, hitl_callback=callback, ) + node.bind_execution_id("00000000-0000-4000-8000-000000000001") + return node def _build_node( diff --git a/api/tests/unit_tests/core/workflow/nodes/iteration/test_iteration_child_engine_errors.py b/api/tests/unit_tests/core/workflow/nodes/iteration/test_iteration_child_engine_errors.py deleted file mode 100644 index 18ed7a0b1d6..00000000000 --- a/api/tests/unit_tests/core/workflow/nodes/iteration/test_iteration_child_engine_errors.py +++ /dev/null @@ -1,95 +0,0 @@ -from collections.abc import Mapping -from typing import Any - -import pytest - -from core.workflow.system_variables import default_system_variables -from graphon.entities import GraphInitParams -from graphon.nodes.iteration.entities import IterationNodeData -from graphon.nodes.iteration.exc import IterationGraphNotFoundError -from graphon.nodes.iteration.iteration_node import IterationNode -from graphon.runtime import ( - ChildEngineBuilderNotConfiguredError, - ChildGraphNotFoundError, - GraphRuntimeState, - VariablePool, -) -from tests.workflow_test_utils import build_test_graph_init_params - - -class _MissingGraphBuilder: - def build_child_engine( - self, - *, - workflow_id: str, - graph_init_params: GraphInitParams, - parent_graph_runtime_state: GraphRuntimeState, - root_node_id: str, - variable_pool: VariablePool | None = None, - ) -> object: - raise ChildGraphNotFoundError(f"child graph root node '{root_node_id}' not found") - - -def _build_runtime_state() -> GraphRuntimeState: - return GraphRuntimeState( - variable_pool=VariablePool.from_bootstrap(system_variables=default_system_variables(), user_inputs={}), - start_at=0.0, - ) - - -def _build_iteration_node( - *, - graph_config: Mapping[str, Any], - runtime_state: GraphRuntimeState, - start_node_id: str, -) -> IterationNode: - init_params = build_test_graph_init_params(graph_config=graph_config) - return IterationNode( - node_id="iteration-node", - data=IterationNodeData( - type="iteration", - title="Iteration", - iterator_selector=["start", "items"], - output_selector=["iteration-node", "output"], - start_node_id=start_node_id, - ), - graph_init_params=init_params, - graph_runtime_state=runtime_state, - ) - - -def test_graph_runtime_state_raises_specific_error_when_child_builder_is_missing(): - runtime_state = _build_runtime_state() - graph_init_params = build_test_graph_init_params() - - with pytest.raises(ChildEngineBuilderNotConfiguredError): - runtime_state.create_child_engine( - workflow_id="workflow", - graph_init_params=graph_init_params, - root_node_id="root", - ) - - -def test_iteration_node_only_translates_child_graph_not_found_error(): - runtime_state = _build_runtime_state() - runtime_state.bind_child_engine_builder(_MissingGraphBuilder()) - node = _build_iteration_node( - graph_config={"nodes": [{"id": "present-node"}], "edges": []}, - runtime_state=runtime_state, - start_node_id="missing-node", - ) - - with pytest.raises(IterationGraphNotFoundError): - node._create_graph_engine(index=0, item="item") - - -def test_iteration_node_propagates_non_graph_not_found_errors(): - runtime_state = _build_runtime_state() - node = _build_iteration_node( - graph_config={"nodes": [{"id": "start-node"}], "edges": []}, - runtime_state=runtime_state, - start_node_id="start-node", - ) - - with pytest.raises(ChildEngineBuilderNotConfiguredError): - node._create_graph_engine(index=0, item="item") diff --git a/api/tests/unit_tests/core/workflow/nodes/knowledge_index/test_knowledge_index_node.py b/api/tests/unit_tests/core/workflow/nodes/knowledge_index/test_knowledge_index_node.py index cefcf04ae02..f131987c189 100644 --- a/api/tests/unit_tests/core/workflow/nodes/knowledge_index/test_knowledge_index_node.py +++ b/api/tests/unit_tests/core/workflow/nodes/knowledge_index/test_knowledge_index_node.py @@ -1,9 +1,11 @@ import time import uuid -from unittest.mock import MagicMock, Mock +from unittest.mock import Mock import pytest from pytest_mock import MockerFixture +from sqlalchemy import event +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom from core.rag.index_processor.constant.index_type import IndexTechniqueType @@ -248,7 +250,6 @@ class TestKnowledgeIndexNode: def test_run_preview_mode_success( self, - mocker: MockerFixture, mock_graph_init_params, mock_graph_runtime_state, mock_index_processor, @@ -283,14 +284,6 @@ class TestKnowledgeIndexNode: total_segments=2, ) mock_index_processor.get_preview_output.return_value = mock_preview - session = MagicMock() - session_context = MagicMock() - session_context.__enter__.return_value = session - mocker.patch( - "core.workflow.nodes.knowledge_index.knowledge_index_node.session_factory.create_session", - return_value=session_context, - ) - node_id = str(uuid.uuid4()) config = { "id": node_id, @@ -310,7 +303,7 @@ class TestKnowledgeIndexNode: # Assert assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED assert result.outputs is not None - assert mock_index_processor.get_preview_output.call_args.kwargs["session"] is session + assert isinstance(mock_index_processor.get_preview_output.call_args.kwargs["session"], Session) def test_run_production_mode_success( self, @@ -548,6 +541,7 @@ class TestKnowledgeIndexNode: mock_index_processor, mock_summary_index_service, sample_node_data, + sqlite_session: Session, ): # Arrange dataset_id = str(uuid.uuid4()) @@ -572,7 +566,9 @@ class TestKnowledgeIndexNode: ) # Act - session = MagicMock() + session = sqlite_session + commits: list[str] = [] + event.listen(session, "after_commit", lambda _session: commits.append("commit")) result = node._invoke_knowledge_index( session=session, dataset_id=dataset_id, @@ -587,7 +583,7 @@ class TestKnowledgeIndexNode: # Assert assert mock_summary_index_service.generate_and_vectorize_summary.called assert mock_index_processor.index_and_clean.called - session.commit.assert_called_once() + assert commits == ["commit"] assert result == {"status": "indexed"} def test_version_method(self): @@ -637,6 +633,7 @@ class TestInvokeKnowledgeIndex: mock_index_processor, mock_summary_index_service, sample_node_data, + sqlite_session: Session, ): # Arrange dataset_id = str(uuid.uuid4()) @@ -662,7 +659,9 @@ class TestInvokeKnowledgeIndex: ) # Act - session = MagicMock() + session = sqlite_session + commits: list[str] = [] + event.listen(session, "after_commit", lambda _session: commits.append("commit")) result = node._invoke_knowledge_index( session=session, dataset_id=dataset_id, @@ -679,7 +678,13 @@ class TestInvokeKnowledgeIndex: dataset_id, document_id, False, summary_setting ) mock_index_processor.index_and_clean.assert_called_once_with( - dataset_id, document_id, original_document_id, chunks, batch, summary_setting, session=session + dataset_id, + document_id, + original_document_id, + chunks, + batch, + summary_setting, + session=session, ) - session.commit.assert_called_once() + assert commits == ["commit"] assert result == {"status": "indexed"} diff --git a/api/tests/unit_tests/core/workflow/nodes/list_operator/node_spec.py b/api/tests/unit_tests/core/workflow/nodes/list_operator/node_spec.py index 20b94d5d509..2c47ba93926 100644 --- a/api/tests/unit_tests/core/workflow/nodes/list_operator/node_spec.py +++ b/api/tests/unit_tests/core/workflow/nodes/list_operator/node_spec.py @@ -1,4 +1,3 @@ -from types import SimpleNamespace from unittest.mock import MagicMock import pytest @@ -36,7 +35,6 @@ class TestListOperatorNode: """Create mock GraphRuntimeState.""" mock_state = MagicMock(spec=GraphRuntimeState) mock_variable_pool = MagicMock() - mock_variable_pool.convert_template.side_effect = lambda value: SimpleNamespace(text=value) mock_state.variable_pool = mock_variable_pool return mock_state diff --git a/api/tests/unit_tests/core/workflow/nodes/llm/test_node.py b/api/tests/unit_tests/core/workflow/nodes/llm/test_node.py index d437c565949..d6f6771a5c2 100644 --- a/api/tests/unit_tests/core/workflow/nodes/llm/test_node.py +++ b/api/tests/unit_tests/core/workflow/nodes/llm/test_node.py @@ -351,6 +351,33 @@ def test_fetch_model_config_hydrates_model_instance_runtime_settings(model_confi provider_model.raise_for_status.assert_called_once() +@pytest.mark.parametrize( + ("provider", "model_name"), + [ + ("", "gpt-3.5-turbo"), + ("openai", ""), + ], +) +def test_fetch_model_config_rejects_unconfigured_model(provider: str, model_name: str): + credentials_provider = mock.MagicMock(spec=CredentialsProvider) + model_factory = mock.MagicMock(spec=DifyModelFactory) + + with pytest.raises(ValueError, match="LLM provider and model are required"): + fetch_model_config( + node_data_model=ModelConfig( + provider=provider, + name=model_name, + mode="chat", + completion_params={}, + ), + credentials_provider=credentials_provider, + model_factory=model_factory, + ) + + credentials_provider.fetch.assert_not_called() + model_factory.init_model_instance.assert_not_called() + + def test_fetch_model_config_reuses_validated_provider_model_from_dify_credentials_provider( model_config: ModelConfigWithCredentialsEntity, ): diff --git a/api/tests/unit_tests/core/workflow/nodes/test_if_else.py b/api/tests/unit_tests/core/workflow/nodes/test_if_else.py index 5965645c4f4..38d1e94bbc9 100644 --- a/api/tests/unit_tests/core/workflow/nodes/test_if_else.py +++ b/api/tests/unit_tests/core/workflow/nodes/test_if_else.py @@ -1,13 +1,12 @@ import time import uuid -from unittest.mock import MagicMock, Mock +from unittest.mock import Mock import pytest from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, InvokeFrom, UserFrom from core.workflow.node_factory import DifyNodeFactory from core.workflow.system_variables import build_system_variables -from extensions.ext_database import db from graphon.enums import WorkflowNodeExecutionStatus from graphon.file import File, FileTransferMethod, FileType from graphon.graph import Graph @@ -124,9 +123,6 @@ def test_execute_if_else_result_true(): graph_runtime_state=graph_runtime_state, ) - # Mock db.session.close() - db.session.close = MagicMock() - # execute node result = node._run() @@ -188,9 +184,6 @@ def test_execute_if_else_result_false(): graph_runtime_state=graph_runtime_state, ) - # Mock db.session.close() - db.session.close = MagicMock() - # execute node result = node._run() @@ -335,9 +328,6 @@ def test_execute_if_else_boolean_conditions(condition: Condition): graph_runtime_state=graph_runtime_state, ) - # Mock db.session.close() - db.session.close = MagicMock() - # execute node result = node._run() @@ -400,9 +390,6 @@ def test_execute_if_else_boolean_false_conditions(): graph_runtime_state=graph_runtime_state, ) - # Mock db.session.close() - db.session.close = MagicMock() - # execute node result = node._run() @@ -468,9 +455,6 @@ def test_execute_if_else_boolean_cases_structure(): graph_runtime_state=graph_runtime_state, ) - # Mock db.session.close() - db.session.close = MagicMock() - # execute node result = node._run() diff --git a/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node.py b/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node.py index 0ee70256d7f..f898a8a8f4b 100644 --- a/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node.py +++ b/api/tests/unit_tests/core/workflow/nodes/tool/test_tool_node.py @@ -239,7 +239,6 @@ def test_image_link_messages_use_tool_file_id_metadata(tool_node: ToolNode): def test_tool_node_passes_node_execution_id_when_runtime_accepts_it(tool_node: ToolNode): runtime_handle = ToolRuntimeHandle(raw=object()) tool_node._runtime.get_runtime = MagicMock(return_value=runtime_handle) - tool_node.ensure_execution_id = MagicMock(return_value="node-execution-id") result = tool_node._get_tool_runtime( variable_pool=tool_node.graph_runtime_state.variable_pool, diff --git a/api/tests/unit_tests/core/workflow/test_human_input_adapter.py b/api/tests/unit_tests/core/workflow/test_human_input_adapter.py index 7a6328ffb4b..30352738d0b 100644 --- a/api/tests/unit_tests/core/workflow/test_human_input_adapter.py +++ b/api/tests/unit_tests/core/workflow/test_human_input_adapter.py @@ -18,12 +18,12 @@ from core.workflow.human_input_adapter import ( ) from graphon.enums import BuiltinNodeTypes from graphon.nodes.base.variable_template_parser import VariableTemplateParser +from graphon.runtime import VariablePool def test_email_delivery_config_helpers_render_and_sanitize_text() -> None: - variable_pool = SimpleNamespace( - convert_template=lambda body: SimpleNamespace(text=body.replace("{{#node.value#}}", "42")) - ) + variable_pool = VariablePool() + variable_pool.add(["node", "value"], "42") rendered = EmailDeliveryConfig.render_body_template( body="Open {{#url#}} and use {{#node.value#}}", diff --git a/api/tests/unit_tests/core/workflow/test_human_input_callback.py b/api/tests/unit_tests/core/workflow/test_human_input_callback.py index 4e4479b1594..88477049d84 100644 --- a/api/tests/unit_tests/core/workflow/test_human_input_callback.py +++ b/api/tests/unit_tests/core/workflow/test_human_input_callback.py @@ -59,6 +59,27 @@ def test_dify_hitl_callback_creates_pause_requested_for_new_form() -> None: assert params.node_id == "node-1" +def test_dify_hitl_callback_scopes_form_to_node_execution() -> None: + repository = MagicMock(spec=HumanInputFormRepository) + repository.get_form.return_value = None + repository.create_form.return_value = SimpleNamespace(id="execution-1") + callback = DifyHITLCallback( + form_repository=repository, + node_data=HumanInputNodeData( + title="Approval", + form_content="Please approve", + user_actions=[UserActionConfig(id="approve", title="Approve")], + ), + execution_id_getter=lambda: "execution-1", + ) + + callback(_ctx("run-1", "node-1")) + + repository.get_form.assert_called_once_with("node-1", form_id="execution-1") + params: FormCreateParams = repository.create_form.call_args.args[0] + assert params.form_id == "execution-1" + + def test_dify_hitl_callback_returns_completed_for_submitted_form() -> None: repository = MagicMock(spec=HumanInputFormRepository) repository.get_form.return_value = SimpleNamespace( diff --git a/api/tests/unit_tests/core/workflow/test_llm_node.py b/api/tests/unit_tests/core/workflow/test_llm_node.py new file mode 100644 index 00000000000..85a77dfd59e --- /dev/null +++ b/api/tests/unit_tests/core/workflow/test_llm_node.py @@ -0,0 +1,56 @@ +from collections.abc import Generator +from types import SimpleNamespace +from typing import cast +from unittest.mock import Mock, sentinel + +import pytest + +from core.workflow.llm_node import DifyLLMNode +from graphon.nodes.llm.node import LLMNode +from graphon.nodes.llm.runtime_protocols import LLMPollingCapableProtocol + + +def test_dify_llm_node_finalizes_polling_when_generator_is_closed(monkeypatch: pytest.MonkeyPatch) -> None: + def invoke(*args: object, **kwargs: object) -> Generator[object, None, None]: + _ = args, kwargs + yield sentinel.event + yield sentinel.unconsumed + + monkeypatch.setattr(LLMNode, "_invoke_llm_with_polling", invoke) + finalizer = Mock() + node = object.__new__(DifyLLMNode) + node._polling_finalizer = finalizer + + events = node._invoke_llm_with_polling( + polling_model=cast(LLMPollingCapableProtocol, SimpleNamespace()), + prompt_messages=[], + stop=None, + ) + + assert next(events) is sentinel.event + events.close() + + finalizer.assert_called_once_with() + + +def test_dify_llm_node_finalizes_polling_when_polling_fails(monkeypatch: pytest.MonkeyPatch) -> None: + def invoke(*args: object, **kwargs: object) -> Generator[object, None, None]: + _ = args, kwargs + yield sentinel.event + raise RuntimeError("polling failed") + + monkeypatch.setattr(LLMNode, "_invoke_llm_with_polling", invoke) + finalizer = Mock() + node = object.__new__(DifyLLMNode) + node._polling_finalizer = finalizer + events = node._invoke_llm_with_polling( + polling_model=cast(LLMPollingCapableProtocol, SimpleNamespace()), + prompt_messages=[], + stop=None, + ) + + assert next(events) is sentinel.event + with pytest.raises(RuntimeError, match="polling failed"): + next(events) + + finalizer.assert_called_once_with() diff --git a/api/tests/unit_tests/core/workflow/test_node_factory.py b/api/tests/unit_tests/core/workflow/test_node_factory.py index 3d305d3a8f4..f5dd034a649 100644 --- a/api/tests/unit_tests/core/workflow/test_node_factory.py +++ b/api/tests/unit_tests/core/workflow/test_node_factory.py @@ -12,6 +12,7 @@ from core.plugin.impl.model_runtime import PluginModelRuntime from core.plugin.plugin_service import PluginService from core.workflow import node_factory from core.workflow import template_rendering as workflow_template_rendering +from core.workflow.llm_node import DifyLLMNode from core.workflow.node_runtime import DifyPreparedLLM from core.workflow.nodes.knowledge_index import KNOWLEDGE_INDEX_NODE_TYPE from graphon.entities.base_node_data import BaseNodeData @@ -24,7 +25,7 @@ from graphon.nodes.llm.entities import LLMNodeData from graphon.nodes.llm.node import LLMNode from graphon.nodes.llm.runtime_protocols import LLMPollingCapableProtocol from graphon.nodes.parameter_extractor.entities import ParameterExtractorNodeData -from graphon.variables.segments import ArrayObjectSegment, StringSegment +from graphon.variables.segments import ArrayObjectSegment, ObjectSegment, StringSegment from models.base import TypeBase from models.model import AppMode, Conversation, ConversationFromSource @@ -324,6 +325,19 @@ class TestDifyNodeFactoryInit: graph_runtime_state=sentinel.graph_runtime_state, ) + def test_with_runtime_state_rebinds_factory(self): + factory = object.__new__(node_factory.DifyNodeFactory) + factory.graph_init_params = sentinel.graph_init_params + + with patch.object(node_factory, "DifyNodeFactory", return_value=sentinel.factory) as factory_cls: + rebound = factory.with_runtime_state(sentinel.graph_runtime_state) + + assert rebound is sentinel.factory + factory_cls.assert_called_once_with( + graph_init_params=sentinel.graph_init_params, + graph_runtime_state=sentinel.graph_runtime_state, + ) + def test_init_builds_default_dependencies(self): graph_init_params = SimpleNamespace(run_context={"context": "value"}) graph_runtime_state = sentinel.graph_runtime_state @@ -675,7 +689,7 @@ class TestDifyNodeFactoryCreateNode: }, } ) - wrapped_model_instance = sentinel.wrapped_model_instance + wrapped_model_instance = MagicMock(spec=DifyPreparedLLM) memory = sentinel.memory factory._build_model_instance_for_llm_node = MagicMock(return_value=sentinel.model_instance) factory._build_memory_for_llm_node = MagicMock(return_value=memory) @@ -704,6 +718,120 @@ class TestDifyNodeFactoryCreateNode: request_metadata={"app_id": "app-id"}, ) assert kwargs["model_instance"] is wrapped_model_instance + assert kwargs["polling_finalizer"] is wrapped_model_instance.finalize_llm_polling + + def test_resolve_llm_model_reference_uses_shared_model_and_parameters(self, factory): + node_data = LLMNodeData.model_validate( + { + "type": BuiltinNodeTypes.LLM, + "title": "LLM", + "model": { + "provider": "old-provider", + "name": "old-model", + "mode": "chat", + "completion_params": {"temperature": 0.2}, + }, + "model_selector": ["env", "for_summarize"], + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + } + ) + factory.graph_runtime_state.variable_pool.get.return_value = ObjectSegment( + value={ + "provider": "new-provider", + "name": "new-model", + "mode": "chat", + "completion_params": {"temperature": 0.8}, + } + ) + + result = factory._resolve_llm_model_reference(node_data) + + assert result.model.provider == "new-provider" + assert result.model.name == "new-model" + assert result.model.mode == node_data.model.mode + assert result.model.completion_params == {"temperature": 0.8} + factory.graph_runtime_state.variable_pool.get.assert_called_once_with(("env", "for_summarize")) + + def test_resolve_llm_model_reference_keeps_node_parameters_for_legacy_variable(self, factory): + node_data = LLMNodeData.model_validate( + { + "type": BuiltinNodeTypes.LLM, + "title": "LLM", + "model": { + "provider": "old-provider", + "name": "old-model", + "mode": "chat", + "completion_params": {"temperature": 0.2}, + }, + "model_selector": ["env", "for_summarize"], + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + } + ) + factory.graph_runtime_state.variable_pool.get.return_value = ObjectSegment( + value={"provider": "new-provider", "name": "new-model", "mode": "chat"} + ) + + result = factory._resolve_llm_model_reference(node_data) + + assert result.model.completion_params == {"temperature": 0.2} + + def test_resolve_llm_model_reference_rejects_mode_mismatch(self, factory): + node_data = LLMNodeData.model_validate( + { + "type": BuiltinNodeTypes.LLM, + "title": "LLM", + "model": {"provider": "provider", "name": "model", "mode": "chat"}, + "model_selector": ["env", "shared_model"], + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + } + ) + factory.graph_runtime_state.variable_pool.get.return_value = ObjectSegment( + value={"provider": "provider", "name": "model", "mode": "completion"} + ) + + with pytest.raises(ValueError, match="uses mode 'completion'.*uses mode 'chat'"): + factory._resolve_llm_model_reference(node_data) + + def test_resolve_llm_model_reference_rejects_missing_variable(self, factory): + node_data = LLMNodeData.model_validate( + { + "type": BuiltinNodeTypes.LLM, + "title": "LLM", + "model": {"provider": "provider", "name": "model", "mode": "chat"}, + "model_selector": ["env", "shared_model"], + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + } + ) + factory.graph_runtime_state.variable_pool.get.return_value = None + + with pytest.raises(ValueError, match="shared_model.*not found"): + factory._resolve_llm_model_reference(node_data) + + def test_resolve_llm_model_reference_keeps_static_model_for_legacy_non_environment_selector(self, factory): + node_data = LLMNodeData.model_validate( + { + "type": BuiltinNodeTypes.LLM, + "title": "LLM", + "model": {"provider": "provider", "name": "model", "mode": "chat"}, + "model_selector": ["conversation", "shared_model"], + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + } + ) + + result = factory._resolve_llm_model_reference(node_data) + + assert result is node_data + factory.graph_runtime_state.variable_pool.get.assert_not_called() def test_build_llm_compatible_node_init_kwargs_uses_polling_wrapper_for_polling_llm_node(self, factory): node_data = LLMNodeData.model_validate( @@ -845,6 +973,44 @@ class TestDifyNodeFactoryCreateNode: assert node.node_data.structured_output_switch_on is True assert node.node_data.structured_output_enabled is True + def test_create_node_uses_dify_llm_node_for_persisted_version_one(self, monkeypatch, factory): + factory.graph_init_params = SimpleNamespace( + workflow_id="workflow-id", + graph_config={}, + run_context={}, + call_depth=0, + ) + monkeypatch.setattr( + factory, + "_build_llm_compatible_node_init_kwargs", + MagicMock( + return_value={ + "model_instance": sentinel.model_instance, + "llm_file_saver": sentinel.llm_file_saver, + "prompt_message_serializer": sentinel.prompt_message_serializer, + "polling_finalizer": MagicMock(), + } + ), + ) + + node = factory.create_node( + { + "id": "llm-node-id", + "data": { + "type": BuiltinNodeTypes.LLM, + "version": "1", + "title": "LLM", + "model": {"provider": "provider", "name": "model", "mode": "chat"}, + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + }, + } + ) + + assert isinstance(node, DifyLLMNode) + assert node.version() == "1" + @pytest.mark.parametrize( ("node_type", "constructor_name", "expected_extra_kwargs"), [ diff --git a/api/tests/unit_tests/core/workflow/test_node_runtime.py b/api/tests/unit_tests/core/workflow/test_node_runtime.py index adfd7ed2c5f..39cfa96c243 100644 --- a/api/tests/unit_tests/core/workflow/test_node_runtime.py +++ b/api/tests/unit_tests/core/workflow/test_node_runtime.py @@ -11,6 +11,7 @@ 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 from core.llm_generator.output_parser.errors import OutputParserError +from core.model_manager import QuotaManagedModelInstance from core.plugin.impl.exc import PluginLLMPollingUnsupportedError from core.plugin.impl.model import PluginModelClient from core.plugin.impl.model_runtime import PluginModelRuntime @@ -148,6 +149,12 @@ class _ModelInstanceStub: self.invoke_llm = Mock(return_value=invoke_llm_result) +class _QuotaManagedModelInstanceStub(_ModelInstanceStub, QuotaManagedModelInstance): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self.reserve_quota = Mock() + + def _build_run_context(*, invoke_from: InvokeFrom | str = InvokeFrom.DEBUGGER) -> dict[str, object]: return build_test_run_context( tenant_id="tenant-id", @@ -357,6 +364,146 @@ def test_dify_prepared_polling_llm_delegates_to_plugin_runtime() -> None: ) +def test_dify_prepared_polling_llm_commits_successful_reservation() -> None: + running_result = LLMPollingResult( + status=LLMPollingStatus.RUNNING, + plugin_state={"task_id": "poll-1"}, + ) + usage = node_runtime.LLMUsage.empty_usage().model_copy(update={"total_tokens": 5}) + succeeded_result = LLMPollingResult( + status=LLMPollingStatus.SUCCEEDED, + result=node_runtime.LLMResult( + model="gpt-4o-mini", + prompt_messages=[], + message=AssistantPromptMessage(content="done"), + usage=usage, + ), + ) + plugin_runtime = PluginModelRuntime( + tenant_id="tenant-id", + user_id="user-id", + client=Mock(spec=PluginModelClient), + plugin_service=PluginService, + ) + plugin_runtime.start_llm_polling = Mock(return_value=running_result) # type: ignore[method-assign] + plugin_runtime.check_llm_polling = Mock(return_value=succeeded_result) # type: ignore[method-assign] + reservation = MagicMock() + model_instance = _QuotaManagedModelInstanceStub( + model_schema=_build_model_schema(features=[ModelFeature.POLLING]), + model_runtime=plugin_runtime, + ) + model_instance.reserve_quota.return_value = reservation + prepared = DifyPreparedPollingLLM(model_instance) + + prepared.start_llm_polling( + prompt_messages=[], + model_parameters={}, + tools=None, + stop=None, + json_schema=None, + ) + prepared.check_llm_polling(plugin_state={"task_id": "poll-1"}) + + reservation.commit.assert_called_once_with(usage) + reservation.release.assert_not_called() + + +def test_dify_prepared_polling_llm_releases_previous_reservation_on_restart() -> None: + running_result = LLMPollingResult( + status=LLMPollingStatus.RUNNING, + plugin_state={"task_id": "poll-1"}, + ) + plugin_runtime = PluginModelRuntime( + tenant_id="tenant-id", + user_id="user-id", + client=Mock(spec=PluginModelClient), + plugin_service=PluginService, + ) + plugin_runtime.start_llm_polling = Mock(return_value=running_result) # type: ignore[method-assign] + first_reservation = MagicMock() + second_reservation = MagicMock() + model_instance = _QuotaManagedModelInstanceStub( + model_schema=_build_model_schema(features=[ModelFeature.POLLING]), + model_runtime=plugin_runtime, + ) + model_instance.reserve_quota.side_effect = [first_reservation, second_reservation] + prepared = DifyPreparedPollingLLM(model_instance) + + for _ in range(2): + prepared.start_llm_polling( + prompt_messages=[], + model_parameters={}, + tools=None, + stop=None, + json_schema=None, + ) + + first_reservation.release.assert_called_once_with() + second_reservation.release.assert_not_called() + assert model_instance.reserve_quota.call_count == 2 + + +def test_dify_prepared_polling_llm_releases_reservation_when_finalized() -> None: + running_result = LLMPollingResult( + status=LLMPollingStatus.RUNNING, + plugin_state={"task_id": "poll-1"}, + ) + polling_runtime = SimpleNamespace( + start_llm_polling=Mock(return_value=running_result), + check_llm_polling=Mock(), + ) + reservation = MagicMock() + model_instance = _QuotaManagedModelInstanceStub( + model_schema=_build_model_schema(features=[ModelFeature.POLLING]), + model_runtime=polling_runtime, + ) + model_instance.reserve_quota.return_value = reservation + prepared = DifyPreparedPollingLLM(model_instance) + + prepared.start_llm_polling( + prompt_messages=[], + model_parameters={}, + tools=None, + stop=None, + json_schema=None, + ) + prepared.finalize_llm_polling() + prepared.finalize_llm_polling() + + reservation.release.assert_called_once_with() + + +def test_dify_prepared_polling_llm_releases_reservation_when_check_fails() -> None: + running_result = LLMPollingResult( + status=LLMPollingStatus.RUNNING, + plugin_state={"task_id": "poll-1"}, + ) + polling_runtime = SimpleNamespace( + start_llm_polling=Mock(return_value=running_result), + check_llm_polling=Mock(side_effect=RuntimeError("polling failed")), + ) + reservation = MagicMock() + model_instance = _QuotaManagedModelInstanceStub( + model_schema=_build_model_schema(features=[ModelFeature.POLLING]), + model_runtime=polling_runtime, + ) + model_instance.reserve_quota.return_value = reservation + prepared = DifyPreparedPollingLLM(model_instance) + + prepared.start_llm_polling( + prompt_messages=[], + model_parameters={}, + tools=None, + stop=None, + json_schema=None, + ) + + with pytest.raises(RuntimeError, match="polling failed"): + prepared.check_llm_polling(plugin_state={"task_id": "poll-1"}) + + reservation.release.assert_called_once_with() + + def test_dify_prepared_polling_llm_raise_exception_when_polling_is_unsupported() -> None: llm_result = node_runtime.LLMResult( model="gpt-4o-mini", diff --git a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py index 41037233b8c..847acee37e5 100644 --- a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py +++ b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py @@ -1,5 +1,5 @@ from collections import UserString -from contextlib import nullcontext +from datetime import datetime from types import SimpleNamespace from unittest.mock import MagicMock, patch, sentinel @@ -14,234 +14,17 @@ from graphon.enums import NodeType, WorkflowNodeExecutionStatus from graphon.errors import WorkflowNodeRunFailedError from graphon.file import File, FileTransferMethod, FileType from graphon.filters import ResponseStreamFilter -from graphon.graph import Graph -from graphon.graph_events import GraphRunFailedEvent -from graphon.model_runtime.entities.llm_entities import LLMMode, LLMUsage +from graphon.graph_events import GraphRunFailedEvent, NodeRunSucceededEvent from graphon.node_events import NodeRunResult from graphon.nodes import BuiltinNodeTypes -from graphon.nodes.base.node import Node -from graphon.nodes.llm.entities import ContextConfig, LLMNodeData, ModelConfig -from graphon.nodes.question_classifier.entities import QuestionClassifierNodeData -from graphon.runtime import ChildGraphNotFoundError, VariablePool +from graphon.runtime import VariablePool from graphon.variables.variables import StringVariable -from tests.workflow_test_utils import build_test_graph_init_params, build_test_variable_pool def _build_typed_node_config(node_type: NodeType): return {"id": "node-id", "data": BaseNodeData(type=node_type)} -def _build_model_config(*, provider: str = "openai", model_name: str = "gpt-4o") -> ModelConfig: - return ModelConfig(provider=provider, name=model_name, mode=LLMMode.CHAT) - - -def _build_llm_node_data(*, provider: str = "openai", model_name: str = "gpt-4o") -> LLMNodeData: - return LLMNodeData( - type=BuiltinNodeTypes.LLM, - title="Child Model", - model=_build_model_config(provider=provider, model_name=model_name), - prompt_template=[], - context=ContextConfig(enabled=False), - ) - - -def _build_question_classifier_node_data( - *, provider: str = "openai", model_name: str = "gpt-4o" -) -> QuestionClassifierNodeData: - return QuestionClassifierNodeData( - type=BuiltinNodeTypes.QUESTION_CLASSIFIER, - title="Child Model", - query_variable_selector=["sys", "query"], - model=_build_model_config(provider=provider, model_name=model_name), - classes=[], - ) - - -class _FakeModelNodeMixin: - @classmethod - def version(cls) -> str: - return "1" - - def post_init(self) -> None: - self.model_instance = SimpleNamespace(provider="stale-provider", model_name="stale-model") - self.usage_snapshot = LLMUsage.empty_usage() - self.usage_snapshot.total_tokens = 1 - - def _run(self) -> NodeRunResult: - return NodeRunResult( - status=WorkflowNodeExecutionStatus.SUCCEEDED, - inputs={ - "model_provider": self.node_data.model.provider, - "model_name": self.node_data.model.name, - }, - llm_usage=self.usage_snapshot, - ) - - -class _FakeLLMNode(_FakeModelNodeMixin, Node[LLMNodeData]): - node_type = BuiltinNodeTypes.LLM - - -class _FakeQuestionClassifierNode(_FakeModelNodeMixin, Node[QuestionClassifierNodeData]): - node_type = BuiltinNodeTypes.QUESTION_CLASSIFIER - - -class TestWorkflowChildEngineBuilder: - @pytest.mark.parametrize( - ("graph_config", "node_id", "expected"), - [ - ({"nodes": [{"id": "root"}]}, "root", True), - ({"nodes": [{"id": "root"}]}, "other", False), - ({"nodes": "invalid"}, "root", None), - ({"nodes": ["invalid"]}, "root", None), - ], - ) - def test_has_node_id(self, graph_config, node_id, expected): - result = workflow_entry._WorkflowChildEngineBuilder._has_node_id(graph_config, node_id) - - assert result is expected - - def test_build_child_engine_raises_when_root_node_is_missing(self): - builder = workflow_entry._WorkflowChildEngineBuilder(tenant_id="tenant-id") - graph_init_params = SimpleNamespace(graph_config={"nodes": []}) - parent_graph_runtime_state = SimpleNamespace( - execution_context=sentinel.execution_context, - variable_pool=sentinel.variable_pool, - ) - - with patch.object(workflow_entry, "DifyNodeFactory", return_value=sentinel.factory): - with pytest.raises(ChildGraphNotFoundError, match="child graph root node 'missing' not found"): - builder.build_child_engine( - workflow_id="workflow-id", - graph_init_params=graph_init_params, - parent_graph_runtime_state=parent_graph_runtime_state, - root_node_id="missing", - ) - - def test_build_child_engine_constructs_graph_engine_with_quota_layer_only(self): - builder = workflow_entry._WorkflowChildEngineBuilder(tenant_id="tenant-id") - graph_init_params = SimpleNamespace(graph_config={"nodes": [{"id": "root"}]}) - parent_graph_runtime_state = SimpleNamespace( - execution_context=sentinel.execution_context, - variable_pool=sentinel.parent_variable_pool, - ) - child_graph = sentinel.child_graph - child_graph_runtime_state = sentinel.child_graph_runtime_state - child_engine = MagicMock() - - with ( - patch.object(workflow_entry.time, "perf_counter", return_value=123.0), - patch.object( - workflow_entry, - "GraphRuntimeState", - return_value=child_graph_runtime_state, - ) as graph_runtime_state_cls, - patch.object(workflow_entry, "DifyNodeFactory", return_value=sentinel.factory) as dify_node_factory, - patch.object(workflow_entry.Graph, "init", return_value=child_graph) as graph_init, - patch.object(workflow_entry, "GraphEngine", return_value=child_engine) as graph_engine_cls, - patch.object(workflow_entry, "GraphEngineConfig", return_value=sentinel.graph_engine_config), - patch.object(workflow_entry, "InMemoryChannel", return_value=sentinel.command_channel), - patch.object(workflow_entry, "LLMQuotaLayer", return_value=sentinel.llm_quota_layer) as llm_quota_layer_cls, - ): - result = builder.build_child_engine( - workflow_id="workflow-id", - graph_init_params=graph_init_params, - parent_graph_runtime_state=parent_graph_runtime_state, - root_node_id="root", - variable_pool=sentinel.child_variable_pool, - ) - - assert result is child_engine - graph_runtime_state_cls.assert_called_once_with( - variable_pool=sentinel.child_variable_pool, - start_at=123.0, - execution_context=sentinel.execution_context, - ) - dify_node_factory.assert_called_once_with( - graph_init_params=graph_init_params, - graph_runtime_state=child_graph_runtime_state, - ) - graph_init.assert_called_once_with( - graph_config={"nodes": [{"id": "root"}]}, - node_factory=sentinel.factory, - root_node_id="root", - ) - graph_engine_cls.assert_called_once_with( - workflow_id="workflow-id", - graph=child_graph, - graph_runtime_state=child_graph_runtime_state, - command_channel=sentinel.command_channel, - config=sentinel.graph_engine_config, - child_engine_builder=builder, - ) - llm_quota_layer_cls.assert_called_once_with(tenant_id="tenant-id") - assert child_engine.layer.call_args_list == [((sentinel.llm_quota_layer,), {})] - - @pytest.mark.parametrize("node_cls", [_FakeLLMNode, _FakeQuestionClassifierNode]) - def test_build_child_engine_runs_llm_quota_layer_for_child_model_nodes(self, node_cls): - builder = workflow_entry._WorkflowChildEngineBuilder(tenant_id="tenant-id") - graph_init_params = build_test_graph_init_params( - graph_config={"nodes": [{"id": "root"}], "edges": []}, - ) - parent_graph_runtime_state = SimpleNamespace( - execution_context=nullcontext(None), - variable_pool=build_test_variable_pool(), - ) - created_node: dict[str, _FakeLLMNode | _FakeQuestionClassifierNode] = {} - - def build_graph(*, graph_config, node_factory, root_node_id): - _ = graph_config - node_data = _build_llm_node_data() if node_cls is _FakeLLMNode else _build_question_classifier_node_data() - node = node_cls( - node_id=root_node_id, - data=node_data, - graph_init_params=node_factory.graph_init_params, - graph_runtime_state=node_factory.graph_runtime_state, - ) - created_node["node"] = node - return Graph( - nodes={root_node_id: node}, - edges={}, - in_edges={}, - out_edges={}, - root_node=node, - ) - - with ( - patch.object( - workflow_entry, - "DifyNodeFactory", - side_effect=lambda graph_init_params, graph_runtime_state: SimpleNamespace( - graph_init_params=graph_init_params, - graph_runtime_state=graph_runtime_state, - ), - ), - patch.object(workflow_entry.Graph, "init", side_effect=build_graph), - patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model") as ensure_quota, - patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model") as deduct_quota, - ): - child_engine = builder.build_child_engine( - workflow_id="workflow-id", - graph_init_params=graph_init_params, - parent_graph_runtime_state=parent_graph_runtime_state, - root_node_id="root", - ) - list(child_engine.run()) - - node = created_node["node"] - ensure_quota.assert_called_once_with( - tenant_id="tenant-id", - provider=node.node_data.model.provider, - model=node.node_data.model.name, - ) - deduct_quota.assert_called_once_with( - tenant_id="tenant-id", - provider=node.node_data.model.provider, - model=node.node_data.model.name, - usage=node.usage_snapshot, - ) - - def _build_minimal_workflow_entry( monkeypatch: pytest.MonkeyPatch, *, @@ -249,13 +32,12 @@ def _build_minimal_workflow_entry( ) -> workflow_entry.WorkflowEntry: """Construct a minimal WorkflowEntry with GraphEngine construction mocked out.""" graph_engine = MagicMock() - graph_runtime_state = SimpleNamespace(execution_context=None) + graph_runtime_state = SimpleNamespace(_execution_context=None) monkeypatch.setattr(workflow_entry, "capture_current_context", lambda: sentinel.execution_context) monkeypatch.setattr(workflow_entry, "GraphEngine", MagicMock(return_value=graph_engine)) monkeypatch.setattr(workflow_entry, "GraphEngineConfig", MagicMock(return_value=sentinel.graph_engine_config)) monkeypatch.setattr(workflow_entry, "InMemoryChannel", MagicMock(return_value=sentinel.command_channel)) - monkeypatch.setattr(workflow_entry, "LLMQuotaLayer", MagicMock(return_value=sentinel.llm_quota_layer)) return workflow_entry.WorkflowEntry( tenant_id="tenant-id", @@ -294,10 +76,9 @@ class TestWorkflowEntryInit: def test_applies_debug_and_observability_layers(self): graph_engine = MagicMock() - graph_runtime_state = SimpleNamespace(execution_context=None) + graph_runtime_state = SimpleNamespace(_execution_context=None) debug_layer = sentinel.debug_layer execution_limits_layer = sentinel.execution_limits_layer - llm_quota_layer = sentinel.llm_quota_layer observability_layer = sentinel.observability_layer with ( @@ -314,7 +95,6 @@ class TestWorkflowEntryInit: "ExecutionLimitsLayer", return_value=execution_limits_layer, ) as execution_limits_layer_cls, - patch.object(workflow_entry, "LLMQuotaLayer", return_value=llm_quota_layer) as llm_quota_layer_cls, patch.object(workflow_entry, "ObservabilityLayer", return_value=observability_layer), ): entry = workflow_entry.WorkflowEntry( @@ -339,9 +119,8 @@ class TestWorkflowEntryInit: graph_runtime_state=graph_runtime_state, command_channel=sentinel.command_channel, config=sentinel.graph_engine_config, - child_engine_builder=entry._child_engine_builder, ) - assert graph_runtime_state.execution_context is sentinel.execution_context + assert graph_runtime_state._execution_context is sentinel.execution_context debug_logging_layer.assert_called_once_with( level="DEBUG", include_inputs=True, @@ -353,11 +132,9 @@ class TestWorkflowEntryInit: max_steps=workflow_entry.dify_config.WORKFLOW_MAX_EXECUTION_STEPS, max_time=workflow_entry.dify_config.WORKFLOW_MAX_EXECUTION_TIME, ) - llm_quota_layer_cls.assert_called_once_with(tenant_id="tenant-id") assert graph_engine.layer.call_args_list == [ ((debug_layer,), {}), ((execution_limits_layer,), {}), - ((llm_quota_layer,), {}), ((observability_layer,), {}), ] @@ -443,6 +220,21 @@ class TestWorkflowEntryRun: class TestWorkflowEntrySingleStepRun: + @pytest.mark.parametrize("node_type", [BuiltinNodeTypes.LOOP, BuiltinNodeTypes.ITERATION]) + def test_rejects_container_nodes(self, node_type): + workflow = SimpleNamespace( + get_node_config_by_id=lambda _node_id: _build_typed_node_config(node_type), + ) + + with pytest.raises(ValueError, match="engine-backed debug endpoints"): + workflow_entry.WorkflowEntry.single_step_run( + workflow=workflow, + node_id="node-id", + user_id="user-id", + user_inputs={}, + variable_pool=sentinel.variable_pool, + ) + def test_preloads_constructor_variables_before_creating_memory_node(self): class FakeLLMNode: id = "node-id" @@ -482,7 +274,7 @@ class TestWorkflowEntrySingleStepRun: patch.object(workflow_entry.WorkflowEntry, "mapping_user_inputs_to_variable_pool"), patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", return_value=iter(["event"]), ), ): @@ -545,7 +337,7 @@ class TestWorkflowEntrySingleStepRun: ) as mapping_user_inputs_to_variable_pool, patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", return_value=iter(["event"]), ), ): @@ -614,7 +406,7 @@ class TestWorkflowEntrySingleStepRun: ) as mapping_user_inputs_to_variable_pool, patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", return_value=iter(["event"]), ), ): @@ -645,7 +437,7 @@ class TestWorkflowEntrySingleStepRun: ) mapping_user_inputs_to_variable_pool.assert_not_called() - def test_wraps_traced_node_run_failures(self): + def test_wraps_layered_node_run_failures(self): class FakeNode: id = "node-id" title = "Node Title" @@ -671,7 +463,7 @@ class TestWorkflowEntrySingleStepRun: patch.object(workflow_entry.WorkflowEntry, "mapping_user_inputs_to_variable_pool"), patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", side_effect=RuntimeError("boom"), ), ): @@ -790,7 +582,7 @@ class TestWorkflowEntryHelpers: ) as mapping_user_inputs_to_variable_pool, patch.object( workflow_entry.WorkflowEntry, - "_traced_node_run", + "_run_node_with_layers", return_value=iter(["event"]), ), ): @@ -953,41 +745,74 @@ class TestMappingUserInputsBranches: ) -class TestWorkflowEntryTracing: - def test_traced_node_run_reports_success(self): - layer = MagicMock() +class TestWorkflowEntryNodeLayers: + def test_run_node_with_layers_reports_success(self): + observability_layer = MagicMock() + result_event = NodeRunSucceededEvent( + id="execution-id", + node_id="node-id", + node_type=BuiltinNodeTypes.START, + start_at=datetime.now(), + node_run_result=NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED), + ) class FakeNode: - def ensure_execution_id(self): + graph_runtime_state = sentinel.graph_runtime_state + + def bind_execution_id(self, _execution_id): return None def run(self): - yield "event" + yield result_event - with patch.object(workflow_entry, "ObservabilityLayer", return_value=layer): - events = list(workflow_entry.WorkflowEntry._traced_node_run(FakeNode())) + node = FakeNode() + with ( + patch.object(workflow_entry, "ObservabilityLayer", return_value=observability_layer), + patch.object(workflow_entry, "InMemoryChannel", return_value=sentinel.command_channel), + patch.object( + workflow_entry, + "ReadOnlyGraphRuntimeStateWrapper", + return_value=sentinel.read_only_runtime_state, + ) as runtime_state_wrapper, + ): + events = list(workflow_entry.WorkflowEntry._run_node_with_layers(node, tenant_id="tenant-id")) - assert events == ["event"] - layer.on_graph_start.assert_called_once_with() - layer.on_node_run_start.assert_called_once() - layer.on_node_run_end.assert_called_once_with( - layer.on_node_run_start.call_args.args[0], - None, - ) + assert events == [result_event] + runtime_state_wrapper.assert_called_once_with(sentinel.graph_runtime_state) + for layer in (observability_layer,): + layer.initialize.assert_called_once_with(sentinel.read_only_runtime_state, sentinel.command_channel) + layer.on_graph_start.assert_called_once_with() + layer.on_node_run_start.assert_called_once_with(node) + layer.on_node_run_end.assert_called_once_with(node, None, result_event) + layer.on_graph_end.assert_called_once_with(None) - def test_traced_node_run_reports_errors(self): - layer = MagicMock() + def test_run_node_with_layers_reports_errors(self): + observability_layer = MagicMock() class FakeNode: - def ensure_execution_id(self): + graph_runtime_state = sentinel.graph_runtime_state + + def bind_execution_id(self, _execution_id): return None def run(self): raise RuntimeError("boom") yield - with patch.object(workflow_entry, "ObservabilityLayer", return_value=layer): + node = FakeNode() + with ( + patch.object(workflow_entry, "ObservabilityLayer", return_value=observability_layer), + patch.object( + workflow_entry, + "ReadOnlyGraphRuntimeStateWrapper", + return_value=sentinel.read_only_runtime_state, + ), + ): with pytest.raises(RuntimeError, match="boom"): - list(workflow_entry.WorkflowEntry._traced_node_run(FakeNode())) + list(workflow_entry.WorkflowEntry._run_node_with_layers(node, tenant_id="tenant-id")) - assert isinstance(layer.on_node_run_end.call_args.args[1], RuntimeError) + for layer in (observability_layer,): + assert layer.on_node_run_end.call_args.args[0] is node + assert isinstance(layer.on_node_run_end.call_args.args[1], RuntimeError) + assert layer.on_node_run_end.call_args.args[2] is None + assert isinstance(layer.on_graph_end.call_args.args[0], RuntimeError) diff --git a/api/tests/unit_tests/enums/test_quota_type.py b/api/tests/unit_tests/enums/test_quota_type.py index f256ff3b4e1..7a7d569d59b 100644 --- a/api/tests/unit_tests/enums/test_quota_type.py +++ b/api/tests/unit_tests/enums/test_quota_type.py @@ -4,7 +4,7 @@ from unittest.mock import patch import pytest -from enums.quota_type import QuotaType +from enums import DeploymentEdition, QuotaType from services.quota_service import QuotaCharge, QuotaService, unlimited @@ -21,19 +21,19 @@ class TestQuotaType: class TestQuotaService: - def test_reserve_billing_disabled(self): + def test_reserve_outside_cloud_edition(self): with ( patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService"), ): - mock_cfg.BILLING_ENABLED = False + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY charge = QuotaService.reserve(QuotaType.TRIGGER, "t1") assert charge.success is True assert charge.charge_id is None def test_reserve_zero_amount_raises(self): with patch("services.quota_service.dify_config") as mock_cfg: - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD with pytest.raises(ValueError, match="greater than 0"): QuotaService.reserve(QuotaType.TRIGGER, "t1", amount=0) @@ -42,7 +42,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_bs.quota_reserve.return_value = {"reservation_id": "rid-1", "available": 99} charge = QuotaService.reserve(QuotaType.TRIGGER, "t1", amount=1) @@ -61,7 +61,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_bs.quota_reserve.return_value = {} with pytest.raises(QuotaExceededError): @@ -74,7 +74,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_bs.quota_reserve.side_effect = QuotaExceededError(feature="trigger", tenant_id="t1", required=1) with pytest.raises(QuotaExceededError): @@ -85,7 +85,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_bs.quota_reserve.side_effect = RuntimeError("network") charge = QuotaService.reserve(QuotaType.TRIGGER, "t1") @@ -97,7 +97,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_bs.quota_reserve.return_value = {"reservation_id": "rid-c"} mock_bs.quota_commit.return_value = {} @@ -105,14 +105,14 @@ class TestQuotaService: assert charge.success is True mock_bs.quota_commit.assert_called_once() - def test_check_billing_disabled(self): + def test_check_outside_cloud_edition(self): with patch("services.quota_service.dify_config") as mock_cfg: - mock_cfg.BILLING_ENABLED = False + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY assert QuotaService.check(QuotaType.TRIGGER, "t1") is True def test_check_zero_amount_raises(self): with patch("services.quota_service.dify_config") as mock_cfg: - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD with pytest.raises(ValueError, match="greater than 0"): QuotaService.check(QuotaType.TRIGGER, "t1", amount=0) @@ -121,7 +121,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch.object(QuotaService, "get_remaining", return_value=100), ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD assert QuotaService.check(QuotaType.TRIGGER, "t1", amount=50) is True def test_check_insufficient_quota(self): @@ -129,7 +129,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch.object(QuotaService, "get_remaining", return_value=5), ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD assert QuotaService.check(QuotaType.TRIGGER, "t1", amount=10) is False def test_check_unlimited_quota(self): @@ -137,7 +137,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch.object(QuotaService, "get_remaining", return_value=-1), ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD assert QuotaService.check(QuotaType.TRIGGER, "t1", amount=999) is True def test_check_exception_returns_true(self): @@ -145,15 +145,15 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch.object(QuotaService, "get_remaining", side_effect=RuntimeError), ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD assert QuotaService.check(QuotaType.TRIGGER, "t1") is True - def test_release_billing_disabled(self): + def test_release_outside_cloud_edition(self): with ( patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = False + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY QuotaService.release(QuotaType.TRIGGER, "rid-1", "t1", "trigger_event") mock_bs.quota_release.assert_not_called() @@ -162,7 +162,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD QuotaService.release(QuotaType.TRIGGER, "", "t1", "trigger_event") mock_bs.quota_release.assert_not_called() @@ -171,7 +171,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_bs.quota_release.return_value = {} QuotaService.release(QuotaType.TRIGGER, "rid-1", "t1", "trigger_event") mock_bs.quota_release.assert_called_once_with( @@ -183,7 +183,7 @@ class TestQuotaService: patch("services.quota_service.dify_config") as mock_cfg, patch("services.billing_service.BillingService") as mock_bs, ): - mock_cfg.BILLING_ENABLED = True + mock_cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_bs.quota_release.side_effect = RuntimeError("fail") QuotaService.release(QuotaType.TRIGGER, "rid-1", "t1", "trigger_event") diff --git a/api/tests/unit_tests/events/test_app_event_signals.py b/api/tests/unit_tests/events/test_app_event_signals.py index 78e441f7eee..6feedf01273 100644 --- a/api/tests/unit_tests/events/test_app_event_signals.py +++ b/api/tests/unit_tests/events/test_app_event_signals.py @@ -10,6 +10,7 @@ from sqlalchemy.orm import Session from events.app_event import app_was_deleted, app_was_updated from models.account import Account +from models.agent import Agent, AgentWorkspace from models.dataset import AppDatasetJoin from models.model import App, AppMode, AppModelConfig, IconType, InstalledApp from services.app_service import AppService @@ -61,7 +62,7 @@ def _make_collector(target: list[App]): return handler -@pytest.mark.parametrize("sqlite_session", [(App, Account)], indirect=True) +@pytest.mark.parametrize("sqlite_session", [(App, Account, Agent, AgentWorkspace)], indirect=True) @pytest.mark.usefixtures("_mock_deps") class TestAppWasDeletedSignal: def test_sends_signal(self, app_model: App, sqlite_session: Session) -> None: diff --git a/api/tests/unit_tests/events/test_update_provider_when_message_created.py b/api/tests/unit_tests/events/test_update_provider_when_message_created.py index a31f0ecdb97..a700783b31c 100644 --- a/api/tests/unit_tests/events/test_update_provider_when_message_created.py +++ b/api/tests/unit_tests/events/test_update_provider_when_message_created.py @@ -11,14 +11,14 @@ from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import ChatAppGenerateEntity from core.entities.provider_entities import ProviderQuotaType, QuotaUnit from events.event_handlers import update_provider_when_message_created -from models import TenantCreditPool +from models import Message, TenantCreditPool +from models.enums import ProviderQuotaType as ModelProviderQuotaType from models.provider import ProviderType @pytest.fixture def credit_pool_session_factory(sqlite_engine: Engine) -> Iterator[sessionmaker[Session]]: """Bind message-created accounting to fixture-owned SQLite sessions.""" - TenantCreditPool.__table__.create(sqlite_engine) session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) with patch("events.event_handlers.update_provider_when_message_created.db.session", session_factory): yield session_factory @@ -31,7 +31,7 @@ def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_ pool_id = str(uuid4()) pool = TenantCreditPool( tenant_id=tenant_id, - pool_type=ProviderQuotaType.TRIAL, + pool_type=ModelProviderQuotaType.TRIAL, quota_limit=10, quota_used=9, ) @@ -62,7 +62,7 @@ def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_ ), ), ) - message = SimpleNamespace(message_tokens=2, answer_tokens=1) + message = Message(message_tokens=2, answer_tokens=1) with ( patch.object(update_provider_when_message_created, "_execute_provider_updates"), @@ -103,7 +103,7 @@ def test_message_created_paid_credit_accounting_uses_paid_pool() -> None: ), ), ) - message = SimpleNamespace(message_tokens=2, answer_tokens=1) + message = Message(message_tokens=2, answer_tokens=1) with ( patch.object(update_provider_when_message_created, "_deduct_credit_pool_quota_capped") as mock_deduct, diff --git a/api/tests/unit_tests/extensions/logstore/repositories/test_logstore_api_workflow_node_execution_repository.py b/api/tests/unit_tests/extensions/logstore/repositories/test_logstore_api_workflow_node_execution_repository.py index 913e8ea6722..4d760748839 100644 --- a/api/tests/unit_tests/extensions/logstore/repositories/test_logstore_api_workflow_node_execution_repository.py +++ b/api/tests/unit_tests/extensions/logstore/repositories/test_logstore_api_workflow_node_execution_repository.py @@ -1,4 +1,5 @@ -from unittest.mock import patch +import json +from unittest.mock import MagicMock, patch from extensions.logstore.repositories.logstore_api_workflow_node_execution_repository import ( LogstoreAPIWorkflowNodeExecutionRepository, @@ -13,3 +14,28 @@ def test_load_full_process_data_returns_logstore_mapping() -> None: execution.process_data = '{"__dify_retry_history": [{"retry_index": 1}]}' assert repository.load_full_process_data(execution) == {"__dify_retry_history": [{"retry_index": 1}]} + + +def test_get_execution_by_id_keeps_process_data_from_highest_failed_log_version() -> None: + with patch("extensions.logstore.repositories.logstore_api_workflow_node_execution_repository.AliyunLogStore"): + repository = LogstoreAPIWorkflowNodeExecutionRepository(session_maker=None) + repository.logstore_client = MagicMock(supports_pg_protocol=False) + repository.logstore_client.get_logs.return_value = [ + { + "id": "execution-1", + "log_version": "1", + "process_data": "{}", + }, + { + "id": "execution-1", + "log_version": "2", + "status": "failed", + "process_data": json.dumps({"workflow_agent_binding_id": "binding-1"}), + }, + ] + + execution = repository.get_execution_by_id("execution-1") + + assert execution is not None + assert execution.status.value == "failed" + assert execution.process_data_dict == {"workflow_agent_binding_id": "binding-1"} diff --git a/api/tests/unit_tests/extensions/logstore/repositories/test_logstore_workflow_node_execution_repository.py b/api/tests/unit_tests/extensions/logstore/repositories/test_logstore_workflow_node_execution_repository.py new file mode 100644 index 00000000000..b3ba655eaa4 --- /dev/null +++ b/api/tests/unit_tests/extensions/logstore/repositories/test_logstore_workflow_node_execution_repository.py @@ -0,0 +1,38 @@ +from types import SimpleNamespace +from typing import cast +from unittest.mock import MagicMock, patch + +import pytest +from sqlalchemy.orm import Session, sessionmaker + +from extensions.logstore.repositories.logstore_workflow_node_execution_repository import ( + LogstoreWorkflowNodeExecutionRepository, +) +from models.account import Account +from models.workflow import WorkflowNodeExecutionTriggeredFrom + + +def test_save_synchronously_writes_sql_when_dual_write_is_disabled( + monkeypatch: pytest.MonkeyPatch, sqlite_session_factory: sessionmaker[Session] +) -> None: + monkeypatch.delenv("LOGSTORE_DUAL_WRITE_ENABLED", raising=False) + with ( + patch("extensions.logstore.repositories.logstore_workflow_node_execution_repository.AliyunLogStore"), + patch( + "extensions.logstore.repositories.logstore_workflow_node_execution_repository." + "SQLAlchemyWorkflowNodeExecutionRepository" + ) as sql_repository_type, + ): + repository = LogstoreWorkflowNodeExecutionRepository( + session_factory=sqlite_session_factory, + tenant_id="tenant-1", + user=cast(Account, SimpleNamespace(id="account-1")), + app_id="app-1", + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + ) + + execution = MagicMock() + repository.save_synchronously(execution) + + assert repository._enable_dual_write is False + sql_repository_type.return_value.save_synchronously.assert_called_once_with(execution) diff --git a/api/tests/unit_tests/extensions/otel/test_flask_instrumentation.py b/api/tests/unit_tests/extensions/otel/test_flask_instrumentation.py new file mode 100644 index 00000000000..a9ea47f06be --- /dev/null +++ b/api/tests/unit_tests/extensions/otel/test_flask_instrumentation.py @@ -0,0 +1,57 @@ +""" +Guards the OpenTelemetry Flask instrumentation contract Dify relies on. + +Flask can run teardown handlers more than once for a request. Instrumentation +versions before 0.63b0 left the span activation and context token in +``request.environ`` after the first teardown, so the second one detached an +already-used token and the OTel API logged "Failed to detach context" for every +affected request. +""" + +import logging +from collections.abc import Iterator + +import flask +import pytest +from opentelemetry.instrumentation.flask import FlaskInstrumentor + + +@pytest.fixture +def instrumented_app() -> Iterator[flask.Flask]: + app = flask.Flask(__name__) + + @app.route("/ping") + def ping() -> str: + return "pong" + + instrumentor = FlaskInstrumentor() + instrumentor.instrument_app(app) + yield app + instrumentor.uninstrument_app(app) + + +@pytest.mark.usefixtures("tracer_provider_with_memory_exporter") +def test_duplicate_teardown_does_not_log_detach_failure( + instrumented_app: flask.Flask, caplog: pytest.LogCaptureFixture +) -> None: + handlers = instrumented_app.teardown_request_funcs[None] + # instrument_app registers its teardown handler last. + index = len(handlers) - 1 + instrumentation_teardown = handlers[index] + calls = 0 + + def teardown_twice(exc: BaseException | None) -> None: + nonlocal calls + calls += 1 + instrumentation_teardown(exc) + instrumentation_teardown(exc) + + handlers[index] = teardown_twice + try: + with caplog.at_level(logging.ERROR, logger="opentelemetry.context"): + assert instrumented_app.test_client().get("/ping").status_code == 200 + finally: + handlers[index] = instrumentation_teardown + + assert calls == 1, "the instrumentation teardown never ran, so nothing was exercised" + assert "Failed to detach context" not in caplog.text diff --git a/api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py b/api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py index 586d39f8f60..824028316fe 100644 --- a/api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py +++ b/api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py @@ -51,11 +51,12 @@ def test_multiple_retrieve_preserves_otel_context_in_dataset_thread( ) -> None: """Per-dataset retrieval spans must remain in the workflow node trace.""" retrieval = DatasetRetrieval() - dataset = MagicMock(spec=Dataset) - dataset.id = str(uuid4()) - dataset.indexing_technique = "high_quality" - dataset.embedding_model = "text-embedding-3-small" - dataset.embedding_model_provider = "openai" + dataset = Dataset( + id=str(uuid4()), + indexing_technique="high_quality", + embedding_model="text-embedding-3-small", + embedding_model_provider="openai", + ) observed_trace_ids: list[int] = [] def record_active_trace(**_kwargs: object) -> None: @@ -120,3 +121,47 @@ def test_retriever_thread_exception_sets_error_span_and_is_collected( assert retrieval_span.status.status_code == StatusCode.ERROR assert cancel_event.is_set() assert thread_exceptions == [expected_error] + + +def test_retriever_thread_exception_emits_skip_event_when_requested( + app, + memory_span_exporter, + tracer_provider_with_memory_exporter, +) -> None: + retrieval = DatasetRetrieval() + cancel_event = threading.Event() + thread_exceptions: list[Exception] = [] + expected_error = RuntimeError("retrieval failed") + dataset_id = str(uuid4()) + + with ( + patch("extensions.otel.decorators.base.dify_config.ENABLE_OTEL", True), + patch("core.rag.retrieval.dataset_retrieval.session_factory.create_session"), + patch.object(retrieval, "_retriever", side_effect=expected_error), + get_tracer(__name__).start_as_current_span("dataset-retrieval-parent") as parent_span, + ): + retrieval._run_retriever_thread_safely( + flask_app=app, + dataset_id=dataset_id, + query="test query", + top_k=4, + all_documents=[], + document_ids_filter=None, + metadata_condition=None, + attachment_ids=None, + cancel_event=cancel_event, + thread_exceptions=thread_exceptions, + skip_on_error=True, + ) + + retrieval_span = next( + span + for span in memory_span_exporter.get_finished_spans() + if span.name.endswith("DatasetRetrieval._run_retriever_thread") + ) + skip_event = next(event for event in parent_span.events if event.name == "dataset_retrieval.skipped") + assert retrieval_span.status.status_code == StatusCode.ERROR + assert skip_event.attributes["dataset_id"] == dataset_id + assert skip_event.attributes["error.message"] == "retrieval failed" + assert not cancel_event.is_set() + assert thread_exceptions == [] diff --git a/api/tests/unit_tests/extensions/test_celery_ssl.py b/api/tests/unit_tests/extensions/test_celery_ssl.py index 366e45d86d8..ad8192ecea4 100644 --- a/api/tests/unit_tests/extensions/test_celery_ssl.py +++ b/api/tests/unit_tests/extensions/test_celery_ssl.py @@ -3,6 +3,8 @@ import ssl from unittest.mock import MagicMock, patch +from enums import DeploymentEdition + class TestCelerySSLConfiguration: """Test suite for Celery SSL configuration.""" @@ -193,7 +195,7 @@ class TestCelerySSLConfiguration: assert "redis_backend_use_ssl" in celery_app.conf assert celery_app.conf["redis_backend_use_ssl"] is not None - def test_celery_init_applies_global_keyprefix_to_broker_and_backend_transport(self): + def test_celery_init_applies_global_keyprefix_and_registers_agent_resource_collector(self): mock_config = MagicMock() mock_config.BROKER_USE_SSL = False mock_config.REDIS_KEY_PREFIX = "enterprise-a" @@ -226,7 +228,7 @@ class TestCelerySSLConfiguration: mock_config.TRIGGER_PROVIDER_REFRESH_INTERVAL = 15 mock_config.ENABLE_API_TOKEN_LAST_USED_UPDATE_TASK = False mock_config.API_TOKEN_LAST_USED_UPDATE_INTERVAL = 30 - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY mock_config.ENTERPRISE_TELEMETRY_ENABLED = False with patch("extensions.ext_celery.dify_config", mock_config): @@ -238,3 +240,4 @@ class TestCelerySSLConfiguration: assert celery_app.conf["broker_transport_options"]["global_keyprefix"] == "enterprise-a:" assert celery_app.conf["result_backend_transport_options"]["global_keyprefix"] == "enterprise-a:" + assert "tasks.collect_agent_resources_task" in celery_app.conf["imports"] diff --git a/api/tests/unit_tests/extensions/test_community_telemetry_celery.py b/api/tests/unit_tests/extensions/test_community_telemetry_celery.py new file mode 100644 index 00000000000..58ec7641457 --- /dev/null +++ b/api/tests/unit_tests/extensions/test_community_telemetry_celery.py @@ -0,0 +1,43 @@ +from types import SimpleNamespace +from unittest.mock import Mock + +from extensions.ext_celery import _enqueue_initial_community_telemetry_heartbeat + + +def test_beat_start_enqueues_community_telemetry_heartbeat() -> None: + task = Mock() + sender = SimpleNamespace( + app=SimpleNamespace( + conf=SimpleNamespace(beat_schedule={"community_telemetry_heartbeat": {}}), + tasks={"community_telemetry.send_heartbeat": task}, + ) + ) + + _enqueue_initial_community_telemetry_heartbeat(sender) + + task.apply_async.assert_called_once_with() + + +def test_beat_start_skips_community_telemetry_when_not_scheduled() -> None: + task = Mock() + sender = SimpleNamespace( + app=SimpleNamespace( + conf=SimpleNamespace(beat_schedule={}), + tasks={"community_telemetry.send_heartbeat": task}, + ) + ) + + _enqueue_initial_community_telemetry_heartbeat(sender) + + task.apply_async.assert_not_called() + + +def test_beat_start_skips_community_telemetry_when_task_is_unavailable() -> None: + sender = SimpleNamespace( + app=SimpleNamespace( + conf=SimpleNamespace(beat_schedule={"community_telemetry_heartbeat": {}}), + tasks={}, + ) + ) + + _enqueue_initial_community_telemetry_heartbeat(sender) diff --git a/api/tests/unit_tests/extensions/test_ext_application_services.py b/api/tests/unit_tests/extensions/test_ext_application_services.py new file mode 100644 index 00000000000..8122b953c10 --- /dev/null +++ b/api/tests/unit_tests/extensions/test_ext_application_services.py @@ -0,0 +1,136 @@ +"""Tests for application-service dependency wiring.""" + +from unittest.mock import MagicMock, patch + +import pytest +from flask import Flask +from sqlalchemy.orm import Session, sessionmaker + +from enums import DeploymentEdition +from extensions import ext_application_services +from extensions.ext_redis import RedisClientWrapper +from models.model import DifySetup +from services.init_validation_service import InvalidInitializationPasswordError + + +@pytest.mark.parametrize( + ("deployment_edition", "initialization_password", "session_validated", "setup_exists", "expected"), + [ + pytest.param(DeploymentEdition.CLOUD, "expected", False, False, True, id="cloud"), + pytest.param(DeploymentEdition.COMMUNITY, "", False, False, True, id="no-password"), + pytest.param(DeploymentEdition.COMMUNITY, "expected", False, False, False, id="not-validated"), + pytest.param(DeploymentEdition.ENTERPRISE, "expected", False, False, False, id="enterprise"), + pytest.param(DeploymentEdition.COMMUNITY, "expected", True, False, True, id="browser-session"), + pytest.param(DeploymentEdition.COMMUNITY, "expected", False, True, True, id="setup-record"), + ], +) +def test_build_application_services_configures_init_validation( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + deployment_edition: DeploymentEdition, + initialization_password: str, + session_validated: bool, + setup_exists: bool, + expected: bool, +) -> None: + if setup_exists: + sqlite_session.add(DifySetup(version="test-version")) + sqlite_session.commit() + + services = ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=deployment_edition, + initialization_password=initialization_password, + redis=MagicMock(spec=RedisClientWrapper), + ) + + assert services.init_validation.is_validated(session_validated=session_validated) is expected + + +def test_build_application_services_passes_the_expected_password( + sqlite_session_factory: sessionmaker[Session], +) -> None: + services = ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=DeploymentEdition.COMMUNITY, + initialization_password="expected", + redis=MagicMock(spec=RedisClientWrapper), + ) + + services.init_validation.validate_password("expected") + with pytest.raises(InvalidInitializationPasswordError): + services.init_validation.validate_password("wrong") + + +def test_init_app_registers_services_for_the_current_app( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], +) -> None: + app = Flask(__name__) + monkeypatch.setattr(ext_application_services, "get_session_maker", lambda: sqlite_session_factory) + monkeypatch.setattr( + ext_application_services.dify_config, + "DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ) + monkeypatch.setattr(ext_application_services.dify_config, "INIT_PASSWORD", "expected") + + ext_application_services.init_app(app) + + with app.app_context(): + services = ext_application_services.application_services() + assert services is app.extensions["application_services"] + assert services.init_validation.is_validated(session_validated=False) is False + + +@pytest.mark.parametrize( + ("deployment_edition", "setup_completed"), + [ + pytest.param(DeploymentEdition.CLOUD, True, id="cloud"), + pytest.param(DeploymentEdition.COMMUNITY, False, id="community"), + pytest.param(DeploymentEdition.ENTERPRISE, False, id="enterprise"), + ], +) +def test_build_application_services_configures_setup_policy( + sqlite_session_factory: sessionmaker[Session], + deployment_edition: DeploymentEdition, + setup_completed: bool, +) -> None: + services = ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=deployment_edition, + initialization_password="", + redis=MagicMock(spec=RedisClientWrapper), + ) + + assert services.setup.get_status().completed is setup_completed + + +def test_build_application_services_wires_builtin_schema_definitions( + sqlite_session_factory: sessionmaker[Session], +) -> None: + services = ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=DeploymentEdition.COMMUNITY, + initialization_password="", + redis=MagicMock(spec=RedisClientWrapper), + ) + + definitions = services.schema_definitions.list() + + assert definitions + assert all({"name", "label", "schema"} <= definition.keys() for definition in definitions) + + +def test_build_application_services_does_not_construct_schema_manager( + sqlite_session_factory: sessionmaker[Session], +) -> None: + with patch("extensions.ext_application_services.SchemaManager") as schema_manager: + ext_application_services.build_application_services( + database_client=sqlite_session_factory, + deployment_edition=DeploymentEdition.COMMUNITY, + initialization_password="", + redis=MagicMock(spec=RedisClientWrapper), + ) + + schema_manager.assert_not_called() diff --git a/api/tests/unit_tests/extensions/test_ext_login.py b/api/tests/unit_tests/extensions/test_ext_login.py index 2207dac1e6a..bcf182c180e 100644 --- a/api/tests/unit_tests/extensions/test_ext_login.py +++ b/api/tests/unit_tests/extensions/test_ext_login.py @@ -3,11 +3,13 @@ from typing import cast from unittest import mock import pytest -from flask import Response +from flask import Flask, Response, request +from constants import COOKIE_NAME_ACCESS_TOKEN from core.logging.context import clear_request_context, get_identity_context from extensions import ext_login from extensions.ext_login import unauthorized_handler +from models.account import TenantAccountRole @pytest.fixture(autouse=True) @@ -30,9 +32,10 @@ def test_unauthorized_handler_returns_json_response() -> None: def test_on_user_logged_in_sets_account_logging_identity() -> None: - account = mock.Mock(spec=ext_login.Account) + account = ext_login.Account(name="Test Account", email="test@example.com") account.id = "account-id" - account.current_tenant_id = "tenant-id" + account._current_tenant = ext_login.Tenant(name="Test Tenant") + account._current_tenant.id = "tenant-id" clear_request_context() ext_login.on_user_logged_in(None, account) @@ -41,10 +44,11 @@ def test_on_user_logged_in_sets_account_logging_identity() -> None: def test_on_user_logged_in_sets_end_user_logging_identity() -> None: - end_user = mock.Mock(spec=ext_login.EndUser) - end_user.id = "end-user-id" - end_user.tenant_id = "tenant-id" - end_user.type = "browser" + end_user = ext_login.EndUser( + id="end-user-id", + tenant_id="tenant-id", + type="browser", + ) clear_request_context() ext_login.on_user_logged_in(None, end_user) @@ -53,12 +57,19 @@ def test_on_user_logged_in_sets_end_user_logging_identity() -> None: def test_on_user_logged_in_does_not_break_auth_when_identity_is_unavailable(caplog: pytest.LogCaptureFixture) -> None: - account = mock.Mock(spec=ext_login.Account) - type(account).current_tenant_id = mock.PropertyMock(side_effect=RuntimeError("unavailable")) + account = ext_login.Account(name="Test Account", email="test@example.com") account.id = "account-id" clear_request_context() - with caplog.at_level("ERROR", logger=ext_login.logger.name): + with ( + mock.patch.object( + ext_login.Account, + "current_tenant_id", + new_callable=mock.PropertyMock, + side_effect=RuntimeError("unavailable"), + ), + caplog.at_level("ERROR", logger=ext_login.logger.name), + ): ext_login.on_user_logged_in(None, account) assert get_identity_context() == ("", "", "") @@ -74,3 +85,36 @@ def test_on_user_logged_in_logs_unsupported_user_type(caplog: pytest.LogCaptureF assert get_identity_context() == ("", "", "") assert "Failed to set logging identity context" in caplog.text + + +def test_admin_api_key_header_takes_precedence_over_console_cookie(monkeypatch: pytest.MonkeyPatch) -> None: + app = Flask(__name__) + session = mock.Mock(spec=ext_login.Session) + tenant = ext_login.Tenant(name="Test Tenant") + tenant_account_join = ext_login.TenantAccountJoin( + tenant_id="tenant-id", + account_id="account-id", + role=TenantAccountRole.NORMAL, + ) + account = ext_login.Account(name="Test Account", email="test@example.com") + session.execute.return_value.one_or_none.return_value = (tenant, tenant_account_join) + session.scalar.side_effect = [account, tenant_account_join] + session.scalars.return_value.one.return_value = tenant + monkeypatch.setattr(ext_login.dify_config, "ADMIN_API_KEY_ENABLE", True) + monkeypatch.setattr(ext_login.dify_config, "ADMIN_API_KEY", "admin-key") + monkeypatch.setattr(ext_login.dify_config, "CONSOLE_WEB_URL", "http://console.example.com") + monkeypatch.setattr(ext_login.dify_config, "CONSOLE_API_URL", "http://api.example.com") + monkeypatch.setattr(ext_login.dify_config, "COOKIE_DOMAIN", "") + + with app.test_request_context( + "/console/api/test", + headers={ + "Authorization": "Bearer admin-key", + "Cookie": f"{COOKIE_NAME_ACCESS_TOKEN}=console-session", + "X-WORKSPACE-ID": "workspace-id", + }, + ): + result = ext_login._load_user_from_request(request, session) + + assert result is account + assert account.current_tenant is tenant diff --git a/api/tests/unit_tests/extensions/test_ext_socketio.py b/api/tests/unit_tests/extensions/test_ext_socketio.py index f9f85106960..27be8c72971 100644 --- a/api/tests/unit_tests/extensions/test_ext_socketio.py +++ b/api/tests/unit_tests/extensions/test_ext_socketio.py @@ -1,5 +1,6 @@ import ssl +import pytest import socketio from extensions import ext_socketio @@ -9,7 +10,7 @@ def test_socketio_server_uses_redis_manager() -> None: assert isinstance(ext_socketio.sio.manager, socketio.RedisManager) -def test_create_socketio_client_manager_uses_pubsub_url_and_prefixed_channel(monkeypatch) -> None: +def test_create_socketio_client_manager_uses_pubsub_url_and_prefixed_channel(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(ext_socketio.dify_config, "PUBSUB_REDIS_URL", "redis://redis.example.com:6380/3") monkeypatch.setattr(ext_socketio.dify_config, "REDIS_KEY_PREFIX", "tenant-a") @@ -19,7 +20,7 @@ def test_create_socketio_client_manager_uses_pubsub_url_and_prefixed_channel(mon assert manager.channel == "tenant-a:socketio" -def test_build_redis_options_includes_tls_options_for_rediss(monkeypatch) -> None: +def test_build_redis_options_includes_tls_options_for_rediss(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(ext_socketio.dify_config, "REDIS_SSL_CERT_REQS", "CERT_REQUIRED") monkeypatch.setattr(ext_socketio.dify_config, "REDIS_SSL_CA_CERTS", "/ca.pem") monkeypatch.setattr(ext_socketio.dify_config, "REDIS_SSL_CERTFILE", "/cert.pem") @@ -31,3 +32,15 @@ def test_build_redis_options_includes_tls_options_for_rediss(monkeypatch) -> Non assert options["ssl_ca_certs"] == "/ca.pem" assert options["ssl_certfile"] == "/cert.pem" assert options["ssl_keyfile"] == "/key.pem" + + +def test_build_redis_options_omits_socket_timeout(monkeypatch: pytest.MonkeyPatch) -> None: + # socket_timeout must not be passed to RedisManager because the pub/sub + # listen loop blocks indefinitely between messages; a read timeout there + # triggers an infinite reconnect storm (issue #39423). + monkeypatch.setattr(ext_socketio.dify_config, "REDIS_SOCKET_TIMEOUT", 5.0) + + options = ext_socketio._build_redis_options("redis://redis.example.com:6380/3") + + assert "socket_timeout" not in options + assert "socket_connect_timeout" in options diff --git a/api/tests/unit_tests/factories/test_variable_factory.py b/api/tests/unit_tests/factories/test_variable_factory.py index fdaf98c86a0..6219336a88e 100644 --- a/api/tests/unit_tests/factories/test_variable_factory.py +++ b/api/tests/unit_tests/factories/test_variable_factory.py @@ -7,6 +7,7 @@ import pytest from hypothesis import HealthCheck, given, settings from hypothesis import strategies as st +from core.workflow.llm_environment_variable import LLMEnvironmentVariable, dump_environment_variable from factories import variable_factory from factories.variable_factory import TypeMismatchError, build_segment, build_segment_with_type from graphon.file import File, FileTransferMethod, FileType @@ -62,6 +63,57 @@ def test_secret_variable(): assert isinstance(result, SecretVariable) +def test_llm_environment_variable(): + result = variable_factory.build_environment_variable_from_mapping( + { + "value_type": "llm", + "name": "for_summarize", + "value": { + "provider": "langgenius/openai/openai", + "name": "gpt-4o-mini", + "mode": "chat", + "completion_params": {"temperature": 0.8}, + }, + } + ) + + assert isinstance(result, LLMEnvironmentVariable) + assert result.value_type == SegmentType.OBJECT + assert result.selector == ["env", "for_summarize"] + dumped = dump_environment_variable(result, mode="json") + assert dumped["value_type"] == "llm" + assert dumped["value"]["completion_params"] == {"temperature": 0.8} + + +def test_llm_environment_variable_normalizes_selector_to_name(): + result = variable_factory.build_environment_variable_from_mapping( + { + "value_type": "llm", + "name": "for_summarize", + "selector": ["env", "different_name"], + "value": {"provider": "langgenius/openai/openai", "name": "gpt-4o-mini", "mode": "chat"}, + } + ) + + assert result.selector == ["env", "for_summarize"] + + +@pytest.mark.parametrize( + "value", + [ + {"provider": "provider", "name": "model"}, + {"provider": "provider", "name": "model", "mode": "embedding"}, + {"provider": "", "name": "model", "mode": "chat"}, + {"provider": "provider", "name": "model", "mode": "chat", "completion_params": []}, + ], +) +def test_llm_environment_variable_rejects_invalid_model_selection(value): + with pytest.raises(VariableError, match="invalid LLM environment variable"): + variable_factory.build_environment_variable_from_mapping( + {"value_type": "llm", "name": "shared_model", "value": value} + ) + + def test_invalid_value_type(): test_data = {"value_type": "unknown", "name": "test_invalid", "value": "value"} with pytest.raises(VariableError): diff --git a/api/tests/unit_tests/fields/test_file_fields.py b/api/tests/unit_tests/fields/test_file_fields.py index 2e33172aedd..c35ca5f7b1d 100644 --- a/api/tests/unit_tests/fields/test_file_fields.py +++ b/api/tests/unit_tests/fields/test_file_fields.py @@ -67,20 +67,24 @@ def test_remote_file_info_and_upload_config() -> None: config = UploadConfig( file_size_limit=1, + knowledge_file_size_limit=11, batch_count_limit=2, file_upload_limit=3, image_file_size_limit=4, video_file_size_limit=5, audio_file_size_limit=6, - workflow_file_upload_limit=7, - image_file_batch_limit=8, - single_chunk_attachment_limit=9, - attachment_image_file_size_limit=10, + skill_file_size_limit=7, + workflow_file_upload_limit=8, + image_file_batch_limit=9, + single_chunk_attachment_limit=10, + attachment_image_file_size_limit=11, ) dumped = config.model_dump(mode="json") assert dumped["file_upload_limit"] == 3 - assert dumped["attachment_image_file_size_limit"] == 10 + assert dumped["knowledge_file_size_limit"] == 11 + assert dumped["skill_file_size_limit"] == 7 + assert dumped["attachment_image_file_size_limit"] == 11 @pytest.mark.parametrize( diff --git a/api/tests/unit_tests/libs/broadcast_channel/redis/test_channel_unit_tests.py b/api/tests/unit_tests/libs/broadcast_channel/redis/test_channel_unit_tests.py index b74d494134b..b323e4c950f 100644 --- a/api/tests/unit_tests/libs/broadcast_channel/redis/test_channel_unit_tests.py +++ b/api/tests/unit_tests/libs/broadcast_channel/redis/test_channel_unit_tests.py @@ -27,6 +27,7 @@ from libs.broadcast_channel.redis.sharded_channel import ( ShardedTopic, _RedisShardedSubscription, ) +from libs.broadcast_channel.signals import SIG_CLOSE class TestBroadcastChannel: @@ -1239,6 +1240,30 @@ class TestRedisSubscriptionCommon: subscription_type, _ = subscription_params assert subscription._get_subscription_type() == subscription_type + def test_listener_ignores_close_signal_from_another_subscription(self, subscription, subscription_params): + subscription_type, _ = subscription_params + topic = f"test-{subscription_type}-topic" + message_type = "message" if subscription_type == "regular" else "smessage" + messages = iter( + [ + {"type": message_type, "channel": topic, "data": SIG_CLOSE}, + {"type": message_type, "channel": topic, "data": b"next-event"}, + ] + ) + + def get_message(): + try: + return next(messages) + except StopIteration: + subscription._closed.set() + return None + + subscription._get_message = get_message + subscription._listen() + + assert subscription._queue.get_nowait() == b"next-event" + assert subscription._queue.empty() + # ==================== Lifecycle Tests ==================== def test_start_if_needed_first_call(self, subscription, subscription_params, mock_pubsub: MagicMock): diff --git a/api/tests/unit_tests/libs/broadcast_channel/redis/test_streams_channel_unit_tests.py b/api/tests/unit_tests/libs/broadcast_channel/redis/test_streams_channel_unit_tests.py index 9a8cb861abf..5022f572153 100644 --- a/api/tests/unit_tests/libs/broadcast_channel/redis/test_streams_channel_unit_tests.py +++ b/api/tests/unit_tests/libs/broadcast_channel/redis/test_streams_channel_unit_tests.py @@ -12,6 +12,7 @@ from libs.broadcast_channel.redis.streams_channel import ( StreamsTopic, _StreamsSubscription, ) +from libs.broadcast_channel.signals import SIG_CLOSE class FakeStreamsRedis: @@ -282,6 +283,34 @@ class TestStreamsSubscription: assert received == case.expected_messages + def test_listener_ignores_close_signal_from_another_subscription(self): + class OneShotRedis: + def __init__(self) -> None: + self._calls = 0 + + def xread(self, streams: dict[str, Any], block: int | None = None, count: int | None = None): + self._calls += 1 + if self._calls == 1: + key = next(iter(streams)) + return [ + ( + key, + [ + ("1-0", {b"data": SIG_CLOSE}), + ("2-0", {b"data": b"next-event"}), + ], + ) + ] + subscription._closed = True + return [] + + subscription = _StreamsSubscription(OneShotRedis(), "stream:close-signal") + subscription._listen() + + assert subscription._queue.get_nowait() == b"next-event" + assert subscription._queue.get_nowait() is subscription._SENTINEL + assert subscription._queue.empty() + def test_iterator_yields_messages_until_subscription_is_closed(self, streams_channel: StreamsBroadcastChannel): topic = streams_channel.topic("iter") subscription = topic.subscribe() diff --git a/api/tests/unit_tests/libs/key_providers/test_azure_keyvault_key_provider.py b/api/tests/unit_tests/libs/key_providers/test_azure_keyvault_key_provider.py new file mode 100644 index 00000000000..748785371f7 --- /dev/null +++ b/api/tests/unit_tests/libs/key_providers/test_azure_keyvault_key_provider.py @@ -0,0 +1,265 @@ +""" +Unit tests for AzureKeyVaultKeyProvider, with the Azure SDK mocked out (no network calls). + +The main thing under test is that ciphertext is pinned to the key *version* that wrapped it, +so that Key Vault's native key rotation (a new "current" version appearing) does not break +decryption of tokens encrypted before the rotation. +""" + +import datetime +import json +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from azure.core.exceptions import HttpResponseError +from azure.keyvault.keys import KeyRotationPolicy + +from configs import dify_config +from libs.key_providers.azure_keyvault_key_provider import AzureKeyVaultKeyProvider + + +class FakeCryptographyClient: + """ + Stands in for azure.keyvault.keys.crypto.CryptographyClient. wrap_key/unwrap_key here just + tag the payload with the bound key version, so unwrapping with the "wrong" version's client + can be detected -- mirroring how a real RSA key from a different version can't unwrap data + wrapped under another version's key. + + Constructed from a key *id string* (e.g. ".../keys//"), mirroring how the + real provider now avoids ever handing CryptographyClient an already-materialized key -- + see AzureKeyVaultKeyProvider._get_crypto_client(). + """ + + # Class-level so a test can mark a version "disabled" (mirroring Key Vault's real behavior + # for a disabled/deleted key version) without threading state through every fixture. + disabled_versions: set[str] = set() + + def __init__(self, key: str, credential: object = None) -> None: + self.version = key.rsplit("/", 1)[-1] + self.credential = credential + self.closed = False + + def close(self) -> None: + self.closed = True + + def wrap_key(self, algorithm: object, key_bytes: bytes) -> SimpleNamespace: + self.last_wrap_algorithm = algorithm + return SimpleNamespace(encrypted_key=f"{self.version}:".encode() + key_bytes) + + def unwrap_key(self, algorithm: object, wrapped: bytes) -> SimpleNamespace: + self.last_unwrap_algorithm = algorithm + if self.version in self.disabled_versions: + raise HttpResponseError(message="Operation unwrapKey is not allowed on a disabled key.") + prefix = f"{self.version}:".encode() + if not wrapped.startswith(prefix): + raise ValueError(f"key version {self.version} cannot unwrap data wrapped by another version") + return SimpleNamespace(key=wrapped[len(prefix) :]) + + +class FakeKeyClient: + """Stands in for azure.keyvault.keys.KeyClient.""" + + def __init__(self, vault_url: str | None = None, credential: object = None) -> None: + self.vault_url = vault_url + self.credential = credential + self.current_version = "v1" + self.rotation_policies: dict[str, KeyRotationPolicy] = {} + self.created_keys: dict[str, int] = {} + # version -> created_on, in creation order; mirrors what list_properties_of_key_versions + # reports in real Key Vault, which the provider now relies on (instead of GetKey) to + # resolve "the current version" without ever fetching key material. + self._version_created_on: dict[str, datetime.datetime] = {} + + def _register_current_version(self) -> None: + if self.current_version not in self._version_created_on: + self._version_created_on[self.current_version] = datetime.datetime(2024, 1, 1) + datetime.timedelta( + seconds=len(self._version_created_on) + ) + + def create_rsa_key(self, name: str, size: int = 2048) -> SimpleNamespace: + self.created_keys[name] = size + self._register_current_version() + return SimpleNamespace(properties=SimpleNamespace(version=self.current_version)) + + def list_properties_of_key_versions(self, name: str) -> list[SimpleNamespace]: + assert name in self.created_keys, f"key {name} was never created" + self._register_current_version() + return [ + SimpleNamespace(version=version, created_on=created_on, id=f"{self.vault_url}/keys/{name}/{version}") + for version, created_on in self._version_created_on.items() + ] + + def update_key_rotation_policy(self, name: str, policy: KeyRotationPolicy) -> KeyRotationPolicy: + self.rotation_policies[name] = policy + return policy + + def get_cryptography_client(self, name: str, *, key_version: str | None = None) -> "FakeCryptographyClient": + assert name in self.created_keys, f"key {name} was never created" + key_id = f"{self.vault_url}/keys/{name}/{key_version}" + return FakeCryptographyClient(key_id, credential=self.credential) + + +@pytest.fixture +def fake_key_client(monkeypatch: pytest.MonkeyPatch) -> FakeKeyClient: + client = FakeKeyClient() + monkeypatch.setattr( + "libs.key_providers.azure_keyvault_key_provider.KeyClient", + MagicMock(return_value=client), + ) + monkeypatch.setattr( + "libs.key_providers.azure_keyvault_key_provider.CryptographyClient", + FakeCryptographyClient, + ) + monkeypatch.setattr( + "libs.key_providers.azure_keyvault_key_provider.DefaultAzureCredential", + MagicMock(), + ) + FakeCryptographyClient.disabled_versions = set() + return client + + +@pytest.fixture(autouse=True) +def azure_keyvault_config(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(dify_config, "AZURE_KEYVAULT_VAULT_URL", "https://fake-vault.vault.azure.net") + monkeypatch.setattr(dify_config, "AZURE_KEYVAULT_KEY_SIZE", 2048) + monkeypatch.setattr(dify_config, "AZURE_KEYVAULT_ROTATION_INTERVAL_DAYS", None) + + +@pytest.mark.usefixtures("fake_key_client") +def test_missing_vault_url_raises(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(dify_config, "AZURE_KEYVAULT_VAULT_URL", None) + with pytest.raises(ValueError, match="AZURE_KEYVAULT_VAULT_URL"): + AzureKeyVaultKeyProvider() + + +@pytest.mark.usefixtures("fake_key_client") +def test_encrypt_decrypt_roundtrip() -> None: + provider = AzureKeyVaultKeyProvider() + provider.generate_key_pair("tenant-1") + + encrypted = provider.encrypt("tenant-1", "super-secret") + decoding = provider.get_decrypt_decoding("tenant-1") + assert provider.decrypt_with_decoding(encrypted, decoding) == "super-secret" + + +def _embedded_key_version(envelope: bytes) -> str: + """Peel out the {"key_version": ...} metadata Dify embeds in each ciphertext, for assertions.""" + body = envelope[len(b"HYBRID:") :] + metadata_len = int.from_bytes(body[1:3], "big") + metadata = json.loads(body[3 : 3 + metadata_len]) + return metadata["key_version"] + + +def test_rotation_does_not_break_decryption_of_old_ciphertext(fake_key_client: FakeKeyClient) -> None: + """ + The core guarantee: a token encrypted before rotation must still decrypt correctly after + Key Vault promotes a new "current" version, and new tokens must promptly pick up the new + version rather than keep using a stale cached "current" client. + """ + provider = AzureKeyVaultKeyProvider() + provider.generate_key_pair("tenant-1") + + encrypted_v1 = provider.encrypt("tenant-1", "secret-before-rotation") + assert _embedded_key_version(encrypted_v1) == "v1" + + # Simulate Key Vault's rotation policy promoting a new version in the background. + fake_key_client.current_version = "v2" + + encrypted_v2 = provider.encrypt("tenant-1", "secret-after-rotation") + # The new version must be picked up immediately, not after some cache TTL elapses. + assert _embedded_key_version(encrypted_v2) == "v2" + + decoding = provider.get_decrypt_decoding("tenant-1") + assert provider.decrypt_with_decoding(encrypted_v1, decoding) == "secret-before-rotation" + assert provider.decrypt_with_decoding(encrypted_v2, decoding) == "secret-after-rotation" + + +def test_generate_key_pair_without_rotation_interval_does_not_set_policy(fake_key_client: FakeKeyClient) -> None: + provider = AzureKeyVaultKeyProvider() + provider.generate_key_pair("tenant-1") + assert fake_key_client.rotation_policies == {} + + +def test_generate_key_pair_with_rotation_interval_sets_time_after_create_only( + monkeypatch: pytest.MonkeyPatch, fake_key_client: FakeKeyClient +) -> None: + monkeypatch.setattr(dify_config, "AZURE_KEYVAULT_ROTATION_INTERVAL_DAYS", 30) + + provider = AzureKeyVaultKeyProvider() + provider.generate_key_pair("tenant-1") + + policy = fake_key_client.rotation_policies["dify-tenant-tenant-1"] + assert policy.expires_in is None + (action,) = policy.lifetime_actions + assert action.time_after_create == "P30D" + assert action.time_before_expiry is None + + +@pytest.mark.usefixtures("fake_key_client") +def test_decrypt_rejects_unknown_envelope_version() -> None: + provider = AzureKeyVaultKeyProvider() + provider.generate_key_pair("tenant-1") + + encrypted = bytearray(provider.encrypt("tenant-1", "secret")) + # Byte right after the "HYBRID:" prefix is the envelope version. + encrypted[len(b"HYBRID:")] = 99 + + with pytest.raises(ValueError, match="envelope version"): + provider.decrypt_with_decoding(bytes(encrypted), "tenant-1") + + +@pytest.mark.usefixtures("fake_key_client") +def test_decrypt_rejects_unrecognized_prefix() -> None: + provider = AzureKeyVaultKeyProvider() + with pytest.raises(ValueError, match="Unsupported ciphertext format"): + provider.decrypt_with_decoding(b"not-a-valid-envelope", "tenant-1") + + +@pytest.mark.usefixtures("fake_key_client") +def test_decrypt_rejects_truncated_envelope_missing_version_byte() -> None: + provider = AzureKeyVaultKeyProvider() + with pytest.raises(ValueError, match="Malformed Azure Key Vault envelope"): + provider.decrypt_with_decoding(b"HYBRID:", "tenant-1") + + +@pytest.mark.usefixtures("fake_key_client") +@pytest.mark.parametrize( + "truncate_at", + [ + len(b"HYBRID:") + 2, # cut inside the metadata-length prefix + len(b"HYBRID:") + 3, # cut inside the metadata JSON blob + ], +) +def test_decrypt_rejects_truncated_envelope_raises_value_error(truncate_at: int) -> None: + provider = AzureKeyVaultKeyProvider() + provider.generate_key_pair("tenant-1") + + encrypted = provider.encrypt("tenant-1", "secret") + truncated = encrypted[:truncate_at] + + with pytest.raises(ValueError): + provider.decrypt_with_decoding(truncated, "tenant-1") + + +@pytest.mark.usefixtures("fake_key_client") +def test_decrypt_translates_disabled_key_version_into_value_error() -> None: + """ + A specific key_version can become unusable after a credential was encrypted (disabled, + deleted, or expired in Key Vault). Callers of decrypt_token_with_decoding + (core/provider_manager.py, services/model_load_balancing_service.py) only catch ValueError + for "this credential can't be decrypted right now" -- if the underlying Azure SDK error + leaked through untranslated, it would crash the whole call chain (e.g. building a tenant's + full provider configuration just to create an unrelated new credential). + """ + provider = AzureKeyVaultKeyProvider() + provider.generate_key_pair("tenant-1") + + encrypted = provider.encrypt("tenant-1", "secret") + assert _embedded_key_version(encrypted) == "v1" + + # Simulate the vault operator disabling/deleting the version that encrypted this credential. + FakeCryptographyClient.disabled_versions.add("v1") + + with pytest.raises(ValueError, match="Failed to unwrap credential"): + provider.decrypt_with_decoding(encrypted, "tenant-1") diff --git a/api/tests/unit_tests/libs/test_datetime_utils.py b/api/tests/unit_tests/libs/test_datetime_utils.py index 57314d29d4b..e2dce54d1e0 100644 --- a/api/tests/unit_tests/libs/test_datetime_utils.py +++ b/api/tests/unit_tests/libs/test_datetime_utils.py @@ -4,7 +4,7 @@ from unittest.mock import patch import pytest import pytz -from libs.datetime_utils import naive_utc_now, parse_time_range +from libs.datetime_utils import naive_utc_now, parse_time_range, to_utc_timestamp def test_naive_utc_now(monkeypatch: pytest.MonkeyPatch): @@ -24,6 +24,18 @@ def test_naive_utc_now(monkeypatch: pytest.MonkeyPatch): assert naive_time == utc_time +@pytest.mark.parametrize( + "value", + [ + datetime.datetime(2024, 1, 1), + datetime.datetime(2024, 1, 1, tzinfo=datetime.UTC), + datetime.datetime(2024, 1, 1, 9, tzinfo=datetime.timezone(datetime.timedelta(hours=9))), + ], +) +def test_to_utc_timestamp(value: datetime.datetime): + assert to_utc_timestamp(value) == 1704067200 + + class TestParseTimeRange: """Test cases for parse_time_range function.""" diff --git a/api/tests/unit_tests/libs/test_email_i18n.py b/api/tests/unit_tests/libs/test_email_i18n.py index b4c0eaf7ee2..0fc37bb016e 100644 --- a/api/tests/unit_tests/libs/test_email_i18n.py +++ b/api/tests/unit_tests/libs/test_email_i18n.py @@ -21,7 +21,7 @@ from libs.email_i18n import ( create_default_email_config, get_email_i18n_service, ) -from services.feature_service import BrandingModel +from services.entities.feature_entities import BrandingModel class MockEmailRenderer: diff --git a/api/tests/unit_tests/libs/test_helper.py b/api/tests/unit_tests/libs/test_helper.py index dbc2cd6cba8..d5f78e67abc 100644 --- a/api/tests/unit_tests/libs/test_helper.py +++ b/api/tests/unit_tests/libs/test_helper.py @@ -2,7 +2,7 @@ from datetime import datetime import pytest -from libs.helper import OptionalTimestampField, email, escape_like_pattern, extract_tenant_id +from libs.helper import OptionalTimestampField, alphanumeric, email, escape_like_pattern, extract_tenant_id from models.account import Account from models.model import EndUser @@ -153,3 +153,47 @@ class TestEmailValidator: def test_invalid_email_rejected(self): with pytest.raises(ValueError, match="not a valid email"): email("not-an-email") + + +class TestAlphanumericValidator: + """Tests for the alphanumeric() validator — regression for #39666.""" + + def test_valid_alphanumeric_accepted(self): + assert alphanumeric("tool_name") == "tool_name" + assert alphanumeric("Tool123") == "Tool123" + assert alphanumeric("_underscore_start") == "_underscore_start" + assert alphanumeric("a") == "a" + + def test_trailing_newline_rejected(self): + # re.match with $ accepts a trailing \n in Python; re.fullmatch does not. + # This was the pre-fix behaviour: alphanumeric("tool\n") returned "tool\n". + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("tool_name\n") + + def test_trailing_carriage_return_rejected(self): + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("tool_name\r") + + def test_trailing_crlf_rejected(self): + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("tool_name\r\n") + + def test_leading_newline_rejected(self): + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("\ntool_name") + + def test_embedded_whitespace_rejected(self): + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("tool name") + + def test_empty_string_rejected(self): + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("") + + def test_special_characters_rejected(self): + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("tool-name") + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("tool.name") + with pytest.raises(ValueError, match="not a valid alphanumeric value"): + alphanumeric("tool/name") diff --git a/api/tests/unit_tests/libs/test_login.py b/api/tests/unit_tests/libs/test_login.py index 8b32e448d64..8155dbd4c9d 100644 --- a/api/tests/unit_tests/libs/test_login.py +++ b/api/tests/unit_tests/libs/test_login.py @@ -266,10 +266,10 @@ class TestCurrentAccountWithTenant: current_user_proxy._get_current_object.return_value = account mocker.patch.object(login_module, "current_user", new=current_user_proxy) - user, tenant_id = login_module.current_account_with_tenant() + account_with_tenant = login_module.current_account_with_tenant() - assert user is account - assert tenant_id == "tenant-123" + assert account_with_tenant.account is account + assert account_with_tenant.tenant_id == "tenant-123" current_user_proxy._get_current_object.assert_called_once_with() def test_raises_when_current_user_is_not_account(self, mocker: MockerFixture): @@ -334,7 +334,11 @@ class TestResolveTenantIdFallback: tenant = Tenant(name="Test Tenant") tenant.id = "tenant-123" account._current_tenant = tenant - mocker.patch.object(login_module, "current_account_with_tenant", return_value=(account, tenant.id)) + mocker.patch.object( + login_module, + "current_account_with_tenant", + return_value=login_module.AccountWithTenant(account=account, tenant_id=tenant.id), + ) tenant_id = login_module.resolve_tenant_id_fallback() diff --git a/api/tests/unit_tests/libs/test_token.py b/api/tests/unit_tests/libs/test_token.py index 97d156478d6..f129a1d86ce 100644 --- a/api/tests/unit_tests/libs/test_token.py +++ b/api/tests/unit_tests/libs/test_token.py @@ -92,3 +92,21 @@ def test_non_whitelisted_path_requires_csrf(): with pytest.raises(Unauthorized): token.check_csrf_token(request, "account-1") + + +def test_admin_api_key_header_bypasses_csrf_when_console_cookie_is_present(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(token.dify_config, "ADMIN_API_KEY_ENABLE", True) + monkeypatch.setattr(token.dify_config, "ADMIN_API_KEY", "admin-key") + monkeypatch.setattr(token.dify_config, "CONSOLE_WEB_URL", "http://console.example.com") + monkeypatch.setattr(token.dify_config, "CONSOLE_API_URL", "http://api.example.com") + monkeypatch.setattr(token.dify_config, "COOKIE_DOMAIN", "") + request = cast( + Request, + MockRequest( + headers={"Authorization": "Bearer admin-key"}, + cookies={COOKIE_NAME_ACCESS_TOKEN: "console-session"}, + args={}, + ), + ) + + token.check_csrf_token(request, "account-1") diff --git a/api/tests/unit_tests/libs/test_workspace_member_helper.py b/api/tests/unit_tests/libs/test_workspace_member_helper.py index f4933e7f594..d35a83e6430 100644 --- a/api/tests/unit_tests/libs/test_workspace_member_helper.py +++ b/api/tests/unit_tests/libs/test_workspace_member_helper.py @@ -1,22 +1,66 @@ -"""Unit tests for require_workspace_member.""" +"""SQLite-backed unit tests for workspace membership enforcement.""" from __future__ import annotations import uuid -from unittest.mock import MagicMock, patch +from collections.abc import Iterator +from dataclasses import dataclass +from unittest.mock import Mock import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden +from enums import DeploymentEdition +from libs import oauth_bearer from libs.oauth_bearer import AuthContext, Scope, SubjectType, TokenType, require_workspace_member +from models.account import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole + +pytestmark = pytest.mark.usefixtures("community_edition") -def _ctx(verified: dict[str, bool] | None = None, *, account: bool = True) -> AuthContext: +@dataclass(frozen=True) +class Database: + """Real ORM binding and executed-statement log for one isolated test.""" + + session: Session + statements: list[str] + + +@pytest.fixture +def database(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Iterator[Database]: + statements: list[str] = [] + + def record_statement(_connection, _cursor, statement, _parameters, _context, _executemany) -> None: + statements.append(statement) + + engine = sqlite_session.get_bind() + event.listen(engine, "before_cursor_execute", record_statement) + binding = Database(session=sqlite_session, statements=statements) + monkeypatch.setattr(oauth_bearer, "db", binding) + try: + yield binding + finally: + event.remove(engine, "before_cursor_execute", record_statement) + + +@pytest.fixture +def community_edition(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(oauth_bearer.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + + +def _ctx( + verified: dict[str, bool] | None = None, + *, + account_id: uuid.UUID | None = None, + account: bool = True, +) -> AuthContext: return AuthContext( subject_type=SubjectType.ACCOUNT if account else SubjectType.EXTERNAL_SSO, subject_email="e@example.com", subject_issuer=None, - account_id=uuid.uuid4() if account else None, + account_id=account_id or (uuid.uuid4() if account else None), client_id="difyctl", scopes=frozenset({Scope.FULL}), token_id=uuid.uuid4(), @@ -27,68 +71,130 @@ def _ctx(verified: dict[str, bool] | None = None, *, account: bool = True) -> Au ) -@patch("libs.oauth_bearer.dify_config") -def test_skips_when_enterprise_enabled(mock_cfg): - mock_cfg.ENTERPRISE_ENABLED = True - require_workspace_member(_ctx(), "t1") +def _persist_membership( + session: Session, + *, + account_id: uuid.UUID, + tenant_id: str, + status: AccountStatus = AccountStatus.ACTIVE, +) -> None: + account = _account(account_id, status=status) + tenant = Tenant(name=f"Tenant {tenant_id}") + tenant.id = tenant_id + membership = TenantAccountJoin( + tenant_id=tenant_id, + account_id=account_id.hex, + role=TenantAccountRole.NORMAL, + ) + session.add_all([account, tenant, membership]) + session.commit() -@patch("libs.oauth_bearer.dify_config") -def test_skips_for_external_sso(mock_cfg): - mock_cfg.ENTERPRISE_ENABLED = False - require_workspace_member(_ctx(account=False), "t1") +def _account(account_id: uuid.UUID, *, status: AccountStatus = AccountStatus.ACTIVE) -> Account: + account = Account(name="Workspace member", email=f"{account_id}@example.com", status=status) + # SQLite's StringUUID adapter binds UUID objects as compact hex, while + # PostgreSQL binds their dashed string form. Persist the SQLite-bound form + # so the production query can keep accepting the AuthContext UUID object. + account.id = account_id.hex + return account -@patch("libs.oauth_bearer.db") -@patch("libs.oauth_bearer.dify_config") -def test_uses_cached_ok_no_db_access(mock_cfg, mock_db): - mock_cfg.ENTERPRISE_ENABLED = False - require_workspace_member(_ctx({"t1": True}), "t1") - mock_db.session.execute.assert_not_called() +def test_skips_for_enterprise_edition(database: Database, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(oauth_bearer.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) + before = len(database.statements) + + require_workspace_member(_ctx(), "tenant-1") + + assert len(database.statements) == before -@patch("libs.oauth_bearer.db") -@patch("libs.oauth_bearer.dify_config") -def test_uses_cached_denied(mock_cfg, mock_db): - mock_cfg.ENTERPRISE_ENABLED = False +def test_skips_for_external_sso(database: Database) -> None: + before = len(database.statements) + + require_workspace_member(_ctx(account=False), "tenant-1") + + assert len(database.statements) == before + + +def test_uses_cached_allow_without_database_access(database: Database) -> None: + before = len(database.statements) + + require_workspace_member(_ctx({"tenant-1": True}), "tenant-1") + + assert len(database.statements) == before + + +def test_uses_cached_denial_without_database_access(database: Database) -> None: + before = len(database.statements) + with pytest.raises(Forbidden, match="workspace_membership_revoked"): - require_workspace_member(_ctx({"t1": False}), "t1") - mock_db.session.execute.assert_not_called() + require_workspace_member(_ctx({"tenant-1": False}), "tenant-1") + + assert len(database.statements) == before -@patch("libs.oauth_bearer.record_layer0_verdict") -@patch("libs.oauth_bearer.db") -@patch("libs.oauth_bearer.dify_config") -def test_denies_when_no_membership(mock_cfg, mock_db, mock_record): - mock_cfg.ENTERPRISE_ENABLED = False - mock_db.session.execute.return_value.scalar_one_or_none.return_value = None +@pytest.mark.usefixtures("database") +def test_denies_when_membership_is_absent( + monkeypatch: pytest.MonkeyPatch, +) -> None: + record_verdict = Mock() + monkeypatch.setattr(oauth_bearer, "record_layer0_verdict", record_verdict) + with pytest.raises(Forbidden, match="workspace_membership_revoked"): - require_workspace_member(_ctx({}), "t1") - mock_record.assert_called_once_with("h1", "t1", False) + require_workspace_member(_ctx(), "tenant-1") + + record_verdict.assert_called_once_with("h1", "tenant-1", False) -@patch("libs.oauth_bearer.record_layer0_verdict") -@patch("libs.oauth_bearer.db") -@patch("libs.oauth_bearer.dify_config") -def test_denies_when_account_inactive(mock_cfg, mock_db, mock_record): - mock_cfg.ENTERPRISE_ENABLED = False - mock_db.session.execute.side_effect = [ - MagicMock(scalar_one_or_none=MagicMock(return_value="join-id")), - MagicMock(scalar_one_or_none=MagicMock(return_value="banned")), - ] +def test_denies_membership_from_another_tenant( + database: Database, + monkeypatch: pytest.MonkeyPatch, +) -> None: + account_id = uuid.uuid4() + requested_tenant_member_id = uuid.uuid4() + status_decoy_id = uuid.uuid4() + _persist_membership(database.session, account_id=account_id, tenant_id="tenant-2") + _persist_membership(database.session, account_id=requested_tenant_member_id, tenant_id="tenant-1") + database.session.add(_account(status_decoy_id, status=AccountStatus.BANNED)) + database.session.commit() + record_verdict = Mock() + monkeypatch.setattr(oauth_bearer, "record_layer0_verdict", record_verdict) + with pytest.raises(Forbidden, match="workspace_membership_revoked"): - require_workspace_member(_ctx({}), "t1") - mock_record.assert_called_once_with("h1", "t1", False) + require_workspace_member(_ctx(account_id=account_id), "tenant-1") + + record_verdict.assert_called_once_with("h1", "tenant-1", False) -@patch("libs.oauth_bearer.record_layer0_verdict") -@patch("libs.oauth_bearer.db") -@patch("libs.oauth_bearer.dify_config") -def test_allows_active_member(mock_cfg, mock_db, mock_record): - mock_cfg.ENTERPRISE_ENABLED = False - mock_db.session.execute.side_effect = [ - MagicMock(scalar_one_or_none=MagicMock(return_value="join-id")), - MagicMock(scalar_one_or_none=MagicMock(return_value="active")), - ] - require_workspace_member(_ctx({}), "t1") - mock_record.assert_called_once_with("h1", "t1", True) +def test_denies_when_account_is_inactive( + database: Database, + monkeypatch: pytest.MonkeyPatch, +) -> None: + account_id = uuid.uuid4() + _persist_membership( + database.session, + account_id=account_id, + tenant_id="tenant-1", + status=AccountStatus.BANNED, + ) + record_verdict = Mock() + monkeypatch.setattr(oauth_bearer, "record_layer0_verdict", record_verdict) + + with pytest.raises(Forbidden, match="workspace_membership_revoked"): + require_workspace_member(_ctx(account_id=account_id), "tenant-1") + + record_verdict.assert_called_once_with("h1", "tenant-1", False) + + +def test_allows_active_member_and_records_verdict( + database: Database, + monkeypatch: pytest.MonkeyPatch, +) -> None: + account_id = uuid.uuid4() + _persist_membership(database.session, account_id=account_id, tenant_id="tenant-1") + record_verdict = Mock() + monkeypatch.setattr(oauth_bearer, "record_layer0_verdict", record_verdict) + + require_workspace_member(_ctx(account_id=account_id), "tenant-1") + + record_verdict.assert_called_once_with("h1", "tenant-1", True) diff --git a/api/tests/unit_tests/libs/test_workspace_permission.py b/api/tests/unit_tests/libs/test_workspace_permission.py index e0c425e7c1e..2d3523e1fad 100644 --- a/api/tests/unit_tests/libs/test_workspace_permission.py +++ b/api/tests/unit_tests/libs/test_workspace_permission.py @@ -4,6 +4,7 @@ from unittest.mock import Mock, patch import pytest from werkzeug.exceptions import Forbidden +from enums import DeploymentEdition from libs.workspace_permission import ( check_workspace_member_invite_permission, check_workspace_owner_transfer_permission, @@ -17,7 +18,7 @@ class TestWorkspacePermissionHelper: @patch("libs.workspace_permission.EnterpriseService") def test_community_edition_allows_invite(self, mock_enterprise_service, mock_config): """Community edition should always allow invitations without calling any service.""" - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY # Should not raise check_workspace_member_invite_permission("test-workspace-id") @@ -29,7 +30,7 @@ class TestWorkspacePermissionHelper: @patch("libs.workspace_permission.FeatureService") def test_community_edition_allows_transfer(self, mock_feature_service, mock_config): """Community edition should check billing plan but not call enterprise service.""" - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY mock_features = Mock() mock_features.is_allow_transfer_workspace = True mock_feature_service.get_features.return_value = mock_features @@ -43,7 +44,7 @@ class TestWorkspacePermissionHelper: @patch("libs.workspace_permission.dify_config") def test_enterprise_blocks_invite_when_disabled(self, mock_config, mock_enterprise_service): """Enterprise edition should block invitations when workspace policy is False.""" - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_permission = Mock() mock_permission.allow_member_invite = False @@ -58,7 +59,7 @@ class TestWorkspacePermissionHelper: @patch("libs.workspace_permission.dify_config") def test_enterprise_allows_invite_when_enabled(self, mock_config, mock_enterprise_service): """Enterprise edition should allow invitations when workspace policy is True.""" - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_permission = Mock() mock_permission.allow_member_invite = True @@ -74,7 +75,7 @@ class TestWorkspacePermissionHelper: @patch("libs.workspace_permission.FeatureService") def test_billing_plan_blocks_transfer(self, mock_feature_service, mock_config, mock_enterprise_service): """SANDBOX billing plan should block owner transfer before checking enterprise policy.""" - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_features = Mock() mock_features.is_allow_transfer_workspace = False # SANDBOX plan mock_feature_service.get_features.return_value = mock_features @@ -90,7 +91,7 @@ class TestWorkspacePermissionHelper: @patch("libs.workspace_permission.FeatureService") def test_enterprise_blocks_transfer_when_disabled(self, mock_feature_service, mock_config, mock_enterprise_service): """Enterprise edition should block transfer when workspace policy is False.""" - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_features = Mock() mock_features.is_allow_transfer_workspace = True # Billing plan allows mock_feature_service.get_features.return_value = mock_features @@ -111,7 +112,7 @@ class TestWorkspacePermissionHelper: self, mock_feature_service, mock_config, mock_enterprise_service ): """Enterprise edition should allow transfer when both billing and workspace policy allow.""" - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_features = Mock() mock_features.is_allow_transfer_workspace = True # Billing plan allows mock_feature_service.get_features.return_value = mock_features @@ -131,7 +132,7 @@ class TestWorkspacePermissionHelper: self, mock_config, mock_enterprise_service, caplog: pytest.LogCaptureFixture ): """On enterprise service error, should fail-open (allow) and log error.""" - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE # Simulate enterprise service error mock_enterprise_service.WorkspacePermissionService.get_permission.side_effect = Exception("Service unavailable") diff --git a/api/tests/unit_tests/migrations/test_agent_workspace_migration.py b/api/tests/unit_tests/migrations/test_agent_workspace_migration.py new file mode 100644 index 00000000000..b75921d8d81 --- /dev/null +++ b/api/tests/unit_tests/migrations/test_agent_workspace_migration.py @@ -0,0 +1,144 @@ +from __future__ import annotations + +import importlib.util +from collections.abc import Callable +from pathlib import Path +from typing import Protocol, cast + +import sqlalchemy as sa +from alembic.migration import MigrationContext +from alembic.operations import Operations + +_MIGRATION_PATH = ( + Path(__file__).resolve().parents[3] + / "migrations/versions/2026_07_23_0203-f6e4c5686857_replace_agent_runtime_sessions_with_.py" +) + + +class _MigrationModule(Protocol): + op: Operations + upgrade: Callable[[], None] + + +def _load_migration_module() -> _MigrationModule: + spec = importlib.util.spec_from_file_location("agent_workspace_migration", _MIGRATION_PATH) + if spec is None or spec.loader is None: + raise RuntimeError("failed to load Agent Workspace migration") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return cast(_MigrationModule, module) + + +def _create_pre_upgrade_schema(engine: sa.Engine) -> None: + metadata = sa.MetaData() + runtime_sessions = sa.Table( + "agent_runtime_sessions", + metadata, + sa.Column("id", sa.String(36), primary_key=True), + sa.Column("tenant_id", sa.String(36), nullable=False), + sa.Column("conversation_id", sa.String(36)), + sa.Column("workflow_run_id", sa.String(36)), + sa.Column("node_id", sa.String(255)), + sa.Column("binding_id", sa.String(36)), + sa.Column("agent_id", sa.String(36), nullable=False), + sa.Column("agent_config_snapshot_id", sa.String(36)), + sa.Column("home_snapshot_id", sa.String(36), nullable=True), + sa.Column("backend_run_id", sa.String(255)), + sa.Column("status", sa.String(32), nullable=False), + ) + sa.Index("agent_runtime_session_backend_run_idx", runtime_sessions.c.backend_run_id) + sa.Index( + "agent_runtime_session_conversation_lookup_idx", + runtime_sessions.c.tenant_id, + runtime_sessions.c.conversation_id, + runtime_sessions.c.status, + ) + sa.Index( + "agent_runtime_session_conversation_scope_unique", + runtime_sessions.c.tenant_id, + runtime_sessions.c.conversation_id, + runtime_sessions.c.agent_id, + runtime_sessions.c.agent_config_snapshot_id, + runtime_sessions.c.home_snapshot_id, + unique=True, + ) + sa.Index( + "agent_runtime_session_workflow_lookup_idx", + runtime_sessions.c.tenant_id, + runtime_sessions.c.workflow_run_id, + runtime_sessions.c.node_id, + runtime_sessions.c.status, + ) + sa.Index( + "agent_runtime_session_workflow_scope_unique", + runtime_sessions.c.tenant_id, + runtime_sessions.c.workflow_run_id, + runtime_sessions.c.node_id, + runtime_sessions.c.binding_id, + runtime_sessions.c.agent_id, + unique=True, + ) + sa.Table( + "agent_home_snapshots", + metadata, + sa.Column("id", sa.String(36), primary_key=True), + sa.Column("tenant_id", sa.String(36), nullable=False), + sa.Column("agent_id", sa.String(36), nullable=False), + sa.Column("snapshot_ref", sa.String(255), nullable=False), + ) + for table_name in ("conversations", "agent_config_drafts", "workflow_node_executions"): + sa.Table(table_name, metadata, sa.Column("id", sa.String(36), primary_key=True)) + metadata.create_all(engine) + + +def _run_upgrade(module: _MigrationModule, engine: sa.Engine) -> None: + with engine.begin() as connection: + operations = Operations(MigrationContext.configure(connection)) + original_op = module.op + module.op = operations + try: + module.upgrade() + finally: + module.op = original_op + + +def test_upgrade_replaces_runtime_sessions_with_workspace_schema() -> None: + engine = sa.create_engine("sqlite:///:memory:") + _create_pre_upgrade_schema(engine) + + _run_upgrade(_load_migration_module(), engine) + + inspector = sa.inspect(engine) + assert "agent_runtime_sessions" not in inspector.get_table_names() + binding_column_definitions = { + column["name"]: column for column in inspector.get_columns("agent_workspace_bindings") + } + binding_columns = set(binding_column_definitions) + assert { + "workspace_id", + "agent_id", + "base_home_snapshot_id", + "agent_config_version_id", + "agent_config_version_kind", + "backend_binding_ref", + "session_snapshot", + "status", + "retired_at", + "pending_form_id", + "pending_tool_call_id", + }.issubset(binding_columns) + assert "active_guard" not in binding_columns + assert binding_column_definitions["base_home_snapshot_id"]["nullable"] is True + binding_indexes = {index["name"] for index in inspector.get_indexes("agent_workspace_bindings")} + assert "agent_workspace_binding_agent_active_unique" not in binding_indexes + workspace_columns = {column["name"] for column in inspector.get_columns("agent_workspaces")} + assert {"owner_type", "owner_id", "owner_scope_key", "backend_workspace_ref", "status"}.issubset(workspace_columns) + workspace_indexes = {index["name"] for index in inspector.get_indexes("agent_workspaces")} + assert "agent_workspace_owner_active_unique" in workspace_indexes + home_columns = {column["name"] for column in inspector.get_columns("agent_home_snapshots")} + assert {"status", "retired_at"}.issubset(home_columns) + home_indexes = {index["name"] for index in inspector.get_indexes("agent_home_snapshots")} + assert "agent_home_snapshot_status_retired_idx" in home_indexes + for table_name in ("conversations", "agent_config_drafts", "workflow_node_executions"): + caller_columns = {column["name"] for column in inspector.get_columns(table_name)} + assert "agent_workspace_binding_id" in caller_columns diff --git a/api/tests/unit_tests/migrations/test_nullable_home_snapshot_migration.py b/api/tests/unit_tests/migrations/test_nullable_home_snapshot_migration.py new file mode 100644 index 00000000000..adebda34291 --- /dev/null +++ b/api/tests/unit_tests/migrations/test_nullable_home_snapshot_migration.py @@ -0,0 +1,181 @@ +from __future__ import annotations + +import importlib.util +from collections.abc import Callable +from pathlib import Path +from typing import Protocol, cast + +import sqlalchemy as sa +from alembic.migration import MigrationContext +from alembic.operations import Operations + +_VERSIONS_DIR = Path(__file__).resolve().parents[3] / "migrations/versions" +_MIGRATION_PATHS = ( + _VERSIONS_DIR / "2026_07_21_2251-2f39536b3feb_add_agent_home_snapshot_ledger.py", + _VERSIONS_DIR / "2026_07_23_0203-f6e4c5686857_replace_agent_runtime_sessions_with_.py", + _VERSIONS_DIR / "2026_07_28_2331-e4708db55c1d_make_home_snapshot_references_nullable.py", +) + + +class _MigrationModule(Protocol): + op: Operations + upgrade: Callable[[], None] + + +def _load_migration(path: Path) -> _MigrationModule: + spec = importlib.util.spec_from_file_location(path.stem, path) + if spec is None or spec.loader is None: + raise RuntimeError(f"failed to load migration {path.name}") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return cast(_MigrationModule, module) + + +def _run_upgrade(module: _MigrationModule, engine: sa.Engine) -> None: + with engine.begin() as connection: + original_op = module.op + module.op = Operations(MigrationContext.configure(connection)) + try: + module.upgrade() + finally: + module.op = original_op + + +def _create_pre_home_snapshot_schema(engine: sa.Engine) -> None: + metadata = sa.MetaData() + drafts = sa.Table("agent_config_drafts", metadata, sa.Column("id", sa.String(36), primary_key=True)) + snapshots = sa.Table("agent_config_snapshots", metadata, sa.Column("id", sa.String(36), primary_key=True)) + conversations = sa.Table("conversations", metadata, sa.Column("id", sa.String(36), primary_key=True)) + executions = sa.Table("workflow_node_executions", metadata, sa.Column("id", sa.String(36), primary_key=True)) + runtime_sessions = sa.Table( + "agent_runtime_sessions", + metadata, + sa.Column("id", sa.String(36), primary_key=True), + sa.Column("tenant_id", sa.String(36), nullable=False), + sa.Column("conversation_id", sa.String(36)), + sa.Column("workflow_run_id", sa.String(36)), + sa.Column("node_id", sa.String(255)), + sa.Column("binding_id", sa.String(36)), + sa.Column("agent_id", sa.String(36), nullable=False), + sa.Column("agent_config_snapshot_id", sa.String(36)), + sa.Column("backend_run_id", sa.String(255)), + sa.Column("status", sa.String(32), nullable=False), + ) + sa.Index("agent_runtime_session_backend_run_idx", runtime_sessions.c.backend_run_id) + sa.Index( + "agent_runtime_session_conversation_lookup_idx", + runtime_sessions.c.tenant_id, + runtime_sessions.c.conversation_id, + runtime_sessions.c.status, + ) + sa.Index( + "agent_runtime_session_conversation_scope_unique", + runtime_sessions.c.tenant_id, + runtime_sessions.c.conversation_id, + runtime_sessions.c.agent_id, + runtime_sessions.c.agent_config_snapshot_id, + unique=True, + ) + sa.Index( + "agent_runtime_session_workflow_lookup_idx", + runtime_sessions.c.tenant_id, + runtime_sessions.c.workflow_run_id, + runtime_sessions.c.node_id, + runtime_sessions.c.status, + ) + sa.Index( + "agent_runtime_session_workflow_scope_unique", + runtime_sessions.c.tenant_id, + runtime_sessions.c.workflow_run_id, + runtime_sessions.c.node_id, + runtime_sessions.c.binding_id, + runtime_sessions.c.agent_id, + unique=True, + ) + metadata.create_all(engine) + with engine.begin() as connection: + connection.execute(drafts.insert().values(id="draft-1")) + connection.execute(snapshots.insert().values(id="snapshot-1")) + connection.execute(conversations.insert().values(id="conversation-1")) + connection.execute(executions.insert().values(id="execution-1")) + connection.execute( + runtime_sessions.insert().values( + id="runtime-1", + tenant_id="tenant-1", + conversation_id="conversation-1", + agent_id="agent-1", + agent_config_snapshot_id="snapshot-1", + status="active", + ) + ) + + +def _create_old_f6_schema(engine: sa.Engine) -> None: + metadata = sa.MetaData() + drafts = sa.Table( + "agent_config_drafts", + metadata, + sa.Column("id", sa.String(36), primary_key=True), + sa.Column("home_snapshot_id", sa.String(36), nullable=False), + ) + snapshots = sa.Table( + "agent_config_snapshots", + metadata, + sa.Column("id", sa.String(36), primary_key=True), + sa.Column("home_snapshot_id", sa.String(36), nullable=False), + ) + bindings = sa.Table( + "agent_workspace_bindings", + metadata, + sa.Column("id", sa.String(36), primary_key=True), + sa.Column("base_home_snapshot_id", sa.String(36), nullable=False), + ) + metadata.create_all(engine) + with engine.begin() as connection: + connection.execute(drafts.insert().values(id="draft-1", home_snapshot_id="home-draft")) + connection.execute(snapshots.insert().values(id="snapshot-1", home_snapshot_id="home-snapshot")) + connection.execute(bindings.insert().values(id="binding-1", base_home_snapshot_id="home-binding")) + + +def test_historical_rows_upgrade_through_nullable_home_snapshot_chain() -> None: + engine = sa.create_engine("sqlite:///:memory:") + _create_pre_home_snapshot_schema(engine) + + for path in _MIGRATION_PATHS: + _run_upgrade(_load_migration(path), engine) + + inspector = sa.inspect(engine) + for table_name, column_name in ( + ("agent_config_drafts", "home_snapshot_id"), + ("agent_config_snapshots", "home_snapshot_id"), + ("agent_workspace_bindings", "base_home_snapshot_id"), + ): + columns = {column["name"]: column for column in inspector.get_columns(table_name)} + assert columns[column_name]["nullable"] is True + + with engine.connect() as connection: + assert connection.scalar(sa.text("SELECT home_snapshot_id FROM agent_config_drafts")) is None + assert connection.scalar(sa.text("SELECT home_snapshot_id FROM agent_config_snapshots")) is None + + +def test_old_f6_schema_converges_to_nullable_without_changing_data() -> None: + engine = sa.create_engine("sqlite:///:memory:") + _create_old_f6_schema(engine) + + _run_upgrade(_load_migration(_MIGRATION_PATHS[-1]), engine) + + inspector = sa.inspect(engine) + for table_name, column_name in ( + ("agent_config_drafts", "home_snapshot_id"), + ("agent_config_snapshots", "home_snapshot_id"), + ("agent_workspace_bindings", "base_home_snapshot_id"), + ): + columns = {column["name"]: column for column in inspector.get_columns(table_name)} + assert columns[column_name]["nullable"] is True + + with engine.connect() as connection: + assert connection.scalar(sa.text("SELECT home_snapshot_id FROM agent_config_drafts")) == "home-draft" + assert connection.scalar(sa.text("SELECT home_snapshot_id FROM agent_config_snapshots")) == "home-snapshot" + assert ( + connection.scalar(sa.text("SELECT base_home_snapshot_id FROM agent_workspace_bindings")) == "home-binding" + ) diff --git a/api/tests/unit_tests/models/test_agent_config_entities.py b/api/tests/unit_tests/models/test_agent_config_entities.py index d95b7130f5c..70b98ea2f36 100644 --- a/api/tests/unit_tests/models/test_agent_config_entities.py +++ b/api/tests/unit_tests/models/test_agent_config_entities.py @@ -2,6 +2,7 @@ import pytest from core.workflow.file_reference import build_file_reference from models.agent_config_entities import ( + AgentSoulModelSettings, DeclaredArrayItem, DeclaredOutputChildConfig, DeclaredOutputConfig, @@ -9,6 +10,22 @@ from models.agent_config_entities import ( ) +def test_agent_soul_model_settings_preserves_plugin_declared_parameters() -> None: + settings = AgentSoulModelSettings.model_validate( + { + "temperature": 0.7, + "enable_thinking": True, + "thinking_budget": 4096, + } + ) + + dumped = settings.model_dump(mode="json", exclude_none=True) + + assert dumped["temperature"] == 0.7 + assert dumped["enable_thinking"] is True + assert dumped["thinking_budget"] == 4096 + + def test_file_default_value_accepts_canonical_reference_mapping() -> None: reference = build_file_reference(record_id="tool-file-1") diff --git a/api/tests/unit_tests/models/test_app_models.py b/api/tests/unit_tests/models/test_app_models.py index 5ef0f52d3d8..c90c8d72fc8 100644 --- a/api/tests/unit_tests/models/test_app_models.py +++ b/api/tests/unit_tests/models/test_app_models.py @@ -191,11 +191,12 @@ class TestAppModelValidation: enable_site=True, enable_api=False, created_by=str(uuid4()), + app_model_config_id="config-1", + ) + app_model_config = AppModelConfig( + app_id="app-id", + agent_mode=json.dumps({"enabled": True, "strategy": "react"}), ) - app.app_model_config_id = "config-1" - app_model_config = MagicMock(spec=AppModelConfig) - app_model_config.agent_mode = "agent" - app_model_config.agent_mode_dict = {"enabled": True, "strategy": "react"} session = MagicMock() session.get.return_value = app_model_config @@ -582,18 +583,19 @@ class TestConversationModel: model_id="model-1", model_provider="provider-1", ) - app_model_config = MagicMock(spec=AppModelConfig) - app_model_config.app_id = "app-1" - app_model_config.to_dict.return_value = {} + app_model_config = AppModelConfig(app_id="app-1") session = MagicMock() session.scalar.return_value = app_model_config annotation_reply = {"enabled": False} - with patch("models.model.load_annotation_reply_config", return_value=annotation_reply) as load_config: + with ( + patch.object(AppModelConfig, "to_dict", return_value={}) as to_dict, + patch("models.model.load_annotation_reply_config", return_value=annotation_reply) as load_config, + ): result = conversation.model_config_with_session(session=session) load_config.assert_called_once_with(session, "app-1") - app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + to_dict.assert_called_once_with(annotation_reply=annotation_reply) assert result["model_id"] == "model-1" assert result["provider"] == "provider-1" @@ -607,20 +609,19 @@ class TestConversationModel: from_end_user_id=str(uuid4()), override_model_configs=json.dumps({"model": {}}), ) - app_model_config = MagicMock(spec=AppModelConfig) - app_model_config.app_id = "app-1" - app_model_config.to_dict.return_value = {} + app_model_config = AppModelConfig(app_id="app-1") session = MagicMock() annotation_reply = {"enabled": False} with ( patch.object(AppModelConfig, "from_model_config_dict", return_value=app_model_config), + patch.object(AppModelConfig, "to_dict", return_value={}) as to_dict, patch("models.model.load_annotation_reply_config", return_value=annotation_reply) as load_config, ): conversation.model_config_with_session(session=session) load_config.assert_called_once_with(session, "app-1") - app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + to_dict.assert_called_once_with(annotation_reply=annotation_reply) session.scalar.assert_not_called() def test_conversation_in_debug_mode(self): @@ -718,8 +719,8 @@ class TestMessageModel: answer_unit_price=Decimal("0.0002"), currency="USD", from_source=ConversationFromSource.API, + _inputs=inputs, ) - message._inputs = inputs # Act result = message.inputs @@ -835,11 +836,11 @@ class TestMessageModel: currency="USD", from_source=ConversationFromSource.API, status="normal", + id=str(uuid4()), + _inputs={"query": "test"}, + created_at=now, + updated_at=now, ) - message.id = str(uuid4()) - message._inputs = {"query": "test"} - message.created_at = now - message.updated_at = now # Act result = message.to_dict() @@ -1178,8 +1179,8 @@ class TestModelIntegration: enable_site=True, enable_api=True, created_by=created_by, + id=app_id, ) - app.id = app_id # Create conversation conversation = Conversation( @@ -1203,8 +1204,8 @@ class TestModelIntegration: answer_unit_price=Decimal("0.0002"), currency="USD", from_source=ConversationFromSource.API, + id=message_id, ) - message.id = message_id # Assert assert app.id == app_id @@ -1229,8 +1230,8 @@ class TestModelIntegration: enable_site=True, enable_api=True, created_by=created_user_id, + id=app_id, ) - app.id = app_id # Create annotation setting setting = AppAnnotationSetting( @@ -1264,8 +1265,8 @@ class TestModelIntegration: answer_unit_price=Decimal("0.0002"), currency="USD", from_source=ConversationFromSource.API, + id=message_id, ) - message.id = message_id # Create annotation annotation = MessageAnnotation( @@ -1332,8 +1333,8 @@ class TestModelIntegration: enable_site=True, enable_api=True, created_by=str(uuid4()), + id=app_id, ) - app.id = app_id # Create site site = Site( diff --git a/api/tests/unit_tests/models/test_dataset_models.py b/api/tests/unit_tests/models/test_dataset_models.py index 0a7a0722933..e7a3a6f83ae 100644 --- a/api/tests/unit_tests/models/test_dataset_models.py +++ b/api/tests/unit_tests/models/test_dataset_models.py @@ -18,6 +18,7 @@ from urllib.parse import parse_qs, urlparse from uuid import uuid4 import pytest +from sqlalchemy.orm import Session from core.rag.entities import ParentMode from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType @@ -97,8 +98,8 @@ class TestDatasetModelValidation: name="Test Dataset", data_source_type=DataSourceType.UPLOAD_FILE, created_by=str(uuid4()), + id=str(uuid4()), ) - dataset.id = str(uuid4()) account = Mock() process_rule = Mock() session = Mock() @@ -112,14 +113,40 @@ class TestDatasetModelValidation: session.get.assert_called_once_with(Account, dataset.created_by) assert session.scalar.call_count == 2 + def test_get_doc_form_ignores_foreign_tenant_document(self, sqlite_session: Session) -> None: + dataset_id = str(uuid4()) + tenant_id = str(uuid4()) + created_by = str(uuid4()) + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + ) + foreign_document = Document( + tenant_id=str(uuid4()), + dataset_id=dataset_id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="foreign", + name="Foreign", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, + doc_form=IndexStructureType.PARENT_CHILD_INDEX, + ) + sqlite_session.add_all([dataset, foreign_document]) + + assert dataset.get_doc_form(session=sqlite_session) is None + def test_get_dataset_keyword_table_uses_caller_session(self): dataset = Dataset( tenant_id=str(uuid4()), name="Test Dataset", data_source_type=DataSourceType.UPLOAD_FILE, created_by=str(uuid4()), + id=str(uuid4()), ) - dataset.id = str(uuid4()) keyword_table = Mock() session = Mock() session.scalar.return_value = keyword_table @@ -137,8 +164,8 @@ class TestDatasetModelValidation: created_by=str(uuid4()), provider="vendor", built_in_field_enabled=False, + id=str(uuid4()), ) - dataset.id = str(uuid4()) account = Mock(name="account", name_value="Ada") account.name = "Ada" session = Mock() @@ -274,9 +301,13 @@ class TestDatasetModelValidation: created_by=str(uuid4()), provider="external", ) - binding = Mock(spec=ExternalKnowledgeBindings) - binding.external_knowledge_id = "knowledge-1" - binding.external_knowledge_api_id = str(uuid4()) + binding = ExternalKnowledgeBindings( + tenant_id="tenant-id", + external_knowledge_api_id=str(uuid4()), + dataset_id="dataset-id", + external_knowledge_id="knowledge-1", + created_by="account-id", + ) session = Mock() session.scalar.side_effect = [binding, None] @@ -686,7 +717,7 @@ class TestDocumentModelRelationships: created_from=DocumentCreatedFrom.WEB, created_by=str(uuid4()), ) - dataset = Mock(spec=Dataset) + dataset = Dataset() session = Mock() session.get.return_value = dataset @@ -752,20 +783,31 @@ class TestDocumentSegmentIndexing: tokens=5, created_by=str(uuid4()), ) - document = Mock(spec=Document) + document = Document() process_rule = Mock(mode="hierarchical", rules_dict={"parent_mode": "paragraph"}) - document.get_dataset_process_rule.return_value = process_rule - child_chunk = Mock(spec=ChildChunk) + child_chunk = ChildChunk( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + segment_id="segment-id", + position=1, + content="", + word_count=0, + created_by="account-id", + ) session = Mock() session.get.return_value = document session.scalars.return_value.all.return_value = [child_chunk] - with patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode="paragraph")): + with ( + patch.object(Document, "get_dataset_process_rule", return_value=process_rule) as get_process_rule, + patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode="paragraph")), + ): result = segment.get_child_chunks(session=session) assert result == [child_chunk] session.get.assert_called_once_with(Document, segment.document_id) - document.get_dataset_process_rule.assert_called_once_with(session=session) + get_process_rule.assert_called_once_with(session=session) session.scalars.assert_called_once() def test_get_child_chunks_includes_full_doc_unless_explicitly_hidden(self): @@ -779,17 +821,29 @@ class TestDocumentSegmentIndexing: tokens=5, created_by=str(uuid4()), ) - document = Mock(spec=Document) - document.get_dataset_process_rule.return_value = Mock( + document = Document() + process_rule = Mock( mode="hierarchical", rules_dict={"parent_mode": ParentMode.FULL_DOC}, ) session = Mock() session.get.return_value = document - child_chunk = Mock(spec=ChildChunk) + child_chunk = ChildChunk( + tenant_id="tenant-id", + dataset_id="dataset-id", + document_id="document-id", + segment_id="segment-id", + position=1, + content="", + word_count=0, + created_by="account-id", + ) session.scalars.return_value.all.return_value = [child_chunk] - with patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode=ParentMode.FULL_DOC)): + with ( + patch.object(Document, "get_dataset_process_rule", return_value=process_rule), + patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode=ParentMode.FULL_DOC)), + ): result = segment.get_child_chunks(session=session) response_result = segment.get_child_chunks(session=session, include_full_doc=False) @@ -808,8 +862,8 @@ class TestDocumentSegmentIndexing: tokens=5, created_by=str(uuid4()), ) - dataset = Mock(spec=Dataset) - document = Mock(spec=Document) + dataset = Dataset() + document = Document() session = Mock() session.get.side_effect = [dataset, document] @@ -1369,8 +1423,8 @@ class TestModelIntegration: data_source_type=DataSourceType.UPLOAD_FILE, created_by=created_by, indexing_technique=IndexTechniqueType.HIGH_QUALITY, + id=dataset_id, ) - dataset.id = dataset_id # Create document document = Document( @@ -1383,8 +1437,8 @@ class TestModelIntegration: created_from=DocumentCreatedFrom.WEB, created_by=created_by, word_count=100, + id=document_id, ) - document.id = document_id # Create segment segment = DocumentSegment( diff --git a/api/tests/unit_tests/models/test_snippet.py b/api/tests/unit_tests/models/test_snippet.py index 17f7cb3c9d4..12500bb2d0c 100644 --- a/api/tests/unit_tests/models/test_snippet.py +++ b/api/tests/unit_tests/models/test_snippet.py @@ -1,10 +1,31 @@ +"""Snippet model properties backed by the shared SQLite test session.""" + import json -from types import SimpleNamespace -from unittest.mock import Mock import pytest +from sqlalchemy.orm import Session +from models import snippet as snippet_module +from models.account import Account +from models.enums import TagType +from models.model import Tag, TagBinding from models.snippet import CustomizedSnippet +from models.workflow import Workflow, WorkflowType + +TENANT_ID = "11111111-1111-1111-1111-111111111111" +WORKFLOW_ID = "22222222-2222-2222-2222-222222222222" +APP_ID = "33333333-3333-3333-3333-333333333333" +SNIPPET_ID = "44444444-4444-4444-4444-444444444444" +ACCOUNT_1_ID = "55555555-5555-5555-5555-555555555555" +ACCOUNT_2_ID = "55555555-5555-5555-5555-555555555556" +SQLITE_MODELS = (Workflow, Tag, TagBinding, Account) + + +@pytest.fixture +def snippet_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Expose the shared SQLite session to model properties that use the global Flask session.""" + monkeypatch.setattr(snippet_module.db, "session", sqlite_session) + return sqlite_session def test_graph_dict_returns_empty_without_workflow_id() -> None: @@ -13,20 +34,28 @@ def test_graph_dict_returns_empty_without_workflow_id() -> None: assert snippet.graph_dict == {} -def test_graph_dict_loads_published_workflow_graph(monkeypatch: pytest.MonkeyPatch) -> None: - workflow = SimpleNamespace(graph=json.dumps({"nodes": [{"id": "llm-1"}], "edges": []})) - session = SimpleNamespace(get=Mock(return_value=workflow)) - monkeypatch.setattr("models.snippet.db.session", session) - snippet = CustomizedSnippet(workflow_id="workflow-1") +@pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True) +def test_graph_dict_loads_published_workflow_graph(snippet_session: Session) -> None: + workflow = Workflow( + tenant_id=TENANT_ID, + app_id=APP_ID, + type=WorkflowType.WORKFLOW, + version="1", + graph=json.dumps({"nodes": [{"id": "llm-1"}], "edges": []}), + _features="{}", + created_by=ACCOUNT_1_ID, + ) + workflow.id = WORKFLOW_ID + snippet_session.add(workflow) + snippet_session.commit() + snippet = CustomizedSnippet(workflow_id=WORKFLOW_ID) assert snippet.graph_dict == {"nodes": [{"id": "llm-1"}], "edges": []} - session.get.assert_called_once() -def test_graph_dict_returns_empty_when_workflow_missing(monkeypatch: pytest.MonkeyPatch) -> None: - session = SimpleNamespace(get=Mock(return_value=None)) - monkeypatch.setattr("models.snippet.db.session", session) - snippet = CustomizedSnippet(workflow_id="missing-workflow") +@pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True) +def test_graph_dict_returns_empty_when_workflow_missing(snippet_session: Session) -> None: + snippet = CustomizedSnippet(workflow_id=WORKFLOW_ID) assert snippet.graph_dict == {} @@ -38,26 +67,30 @@ def test_input_fields_list_parses_json_or_returns_empty() -> None: ] -def test_tags_returns_query_results_or_empty(monkeypatch: pytest.MonkeyPatch) -> None: - tags = [SimpleNamespace(id="tag-1")] - session = SimpleNamespace(scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=tags)))) - monkeypatch.setattr("models.snippet.db.session", session) - snippet = CustomizedSnippet(id="snippet-1", tenant_id="tenant-1") +@pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True) +def test_tags_returns_query_results_or_empty(snippet_session: Session) -> None: + tag = Tag(tenant_id=TENANT_ID, type=TagType.SNIPPET, name="Reusable", created_by=ACCOUNT_1_ID) + binding = TagBinding(tenant_id=TENANT_ID, tag_id=tag.id, target_id=SNIPPET_ID, created_by=ACCOUNT_1_ID) + snippet_session.add_all((tag, binding)) + snippet_session.commit() + snippet = CustomizedSnippet(id=SNIPPET_ID, tenant_id=TENANT_ID) - assert snippet.tags == tags + assert snippet.tags == [tag] - session.scalars.return_value.all.return_value = None + snippet_session.delete(binding) + snippet_session.commit() assert snippet.tags == [] -def test_account_properties_and_author_name(monkeypatch: pytest.MonkeyPatch) -> None: - account = SimpleNamespace(id="account-1", name="Ada") - updated_account = SimpleNamespace(id="account-2", name="Grace") - session = SimpleNamespace( - get=Mock(side_effect=lambda _model, account_id: account if account_id == "account-1" else updated_account) - ) - monkeypatch.setattr("models.snippet.db.session", session) - snippet = CustomizedSnippet(created_by="account-1", updated_by="account-2") +@pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True) +def test_account_properties_and_author_name(snippet_session: Session) -> None: + account = Account(name="Ada", email="ada@example.com") + account.id = ACCOUNT_1_ID + updated_account = Account(name="Grace", email="grace@example.com") + updated_account.id = ACCOUNT_2_ID + snippet_session.add_all((account, updated_account)) + snippet_session.commit() + snippet = CustomizedSnippet(created_by=ACCOUNT_1_ID, updated_by=ACCOUNT_2_ID) assert snippet.created_by_account is account assert snippet.author_name == "Ada" diff --git a/api/tests/unit_tests/models/test_workflow.py b/api/tests/unit_tests/models/test_workflow.py index 12f45e21beb..9ec7383e1dd 100644 --- a/api/tests/unit_tests/models/test_workflow.py +++ b/api/tests/unit_tests/models/test_workflow.py @@ -6,6 +6,7 @@ from uuid import uuid4 from constants import HIDDEN_VALUE from core.helper import encrypter from core.workflow.file_reference import build_file_reference +from core.workflow.llm_environment_variable import LLMEnvironmentVariable from factories.variable_factory import build_segment from graphon.file import File, FileTransferMethod, FileType from graphon.variables import FloatVariable, IntegerVariable, SecretVariable, StringVariable @@ -53,6 +54,32 @@ def test_environment_variables(): assert workflow.environment_variables == variables +def test_llm_environment_variable_round_trip(): + workflow = Workflow( + tenant_id="tenant_id", + app_id="app_id", + type="workflow", + version="draft", + graph="{}", + features="{}", + created_by="account_id", + environment_variables=[], + conversation_variables=[], + ) + variable = LLMEnvironmentVariable( + name="for_research", + value={"provider": "langgenius/anthropic/anthropic", "name": "claude-sonnet", "mode": "chat"}, + id=str(uuid4()), + selector=["env", "for_research"], + ) + + workflow.environment_variables = [variable] + + assert workflow.environment_variables == [variable] + assert json.loads(workflow._environment_variables)["for_research"]["value_type"] == "llm" + assert workflow.to_dict()["environment_variables"][0]["value_type"] == "llm" + + def test_update_environment_variables(): # tenant_id context variable removed - using current_user.current_tenant_id directly @@ -146,10 +173,10 @@ def test_workflow_account_getters_use_caller_session(): created_by="created-account-id", environment_variables=[], conversation_variables=[], + updated_by="updated-account-id", ) - workflow.updated_by = "updated-account-id" - created_account = mock.Mock(spec=Account) - updated_account = mock.Mock(spec=Account) + created_account = Account(name="Test Account", email="test@example.com") + updated_account = Account(name="Test Account", email="test@example.com") session = mock.Mock() session.get.side_effect = [created_account, updated_account] @@ -218,8 +245,9 @@ def test_normalize_environment_variable_mappings_keeps_hidden_value(): class TestWorkflowNodeExecution: def test_execution_metadata_dict(self): - node_exec = WorkflowNodeExecutionModel() - node_exec.execution_metadata = None + node_exec = WorkflowNodeExecutionModel( + execution_metadata=None, + ) assert node_exec.execution_metadata_dict == {} original = {"a": 1, "b": ["2"]} @@ -382,8 +410,9 @@ class TestWorkflowDraftVariableGetValue: size=12, storage_key="canonical-storage-key", ) - draft_var = WorkflowDraftVariable() - draft_var.app_id = "app-1" + draft_var = WorkflowDraftVariable( + app_id="app-1", + ) draft_var.set_value(build_segment(persisted_file)) draft_var._WorkflowDraftVariable__value = None diff --git a/api/tests/unit_tests/pyrefly.toml b/api/tests/unit_tests/pyrefly.toml new file mode 100644 index 00000000000..d5db248382f --- /dev/null +++ b/api/tests/unit_tests/pyrefly.toml @@ -0,0 +1,956 @@ +preset = "strict" +project-includes = ["."] +search-path = ["../.."] +python-platform = "linux" +python-version = "3.12.0" +infer-with-first-use = true +min-severity = "warn" + +# Existing strict-mode debt. Remove a file when bringing it under strict checking. +project-excludes = [ + "clients/agent_backend/test_cleanup_composition_compositor_integration.py", + "clients/agent_backend/test_client.py", + "clients/agent_backend/test_event_adapter.py", + "clients/agent_backend/test_fake_client.py", + "clients/agent_backend/test_request_builder.py", + "clients/agent_backend/test_session_cleanup.py", + "commands/test_archive_workflow_runs.py", + "commands/test_clean_expired_messages.py", + "commands/test_data_migration_commands.py", + "commands/test_data_migration_wizard.py", + "commands/test_fix_app_site_missing.py", + "commands/test_generate_swagger_markdown_docs.py", + "commands/test_generate_swagger_specs.py", + "commands/test_legacy_model_type_migration.py", + "commands/test_lint_response_contracts.py", + "commands/test_reset_encrypt_key_pair.py", + "commands/test_upgrade_db.py", + "configs/test_dify_config.py", + "configs/test_env_consistency.py", + "configs/test_nacos_http_client.py", + "controllers/common/test_agent_app_parameters.py", + "controllers/common/test_app_access.py", + "controllers/common/test_errors.py", + "controllers/common/test_fields.py", + "controllers/common/test_file_response.py", + "controllers/common/test_helpers.py", + "controllers/common/test_schema.py", + "controllers/common/test_session.py", + "controllers/console/agent/test_agent_controllers.py", + "controllers/console/app/test_agent_app_sandbox.py", + "controllers/console/app/test_agent_config_inspector.py", + "controllers/console/app/test_agent_drive_inspector.py", + "controllers/console/app/test_agent_manage_guard.py", + "controllers/console/app/test_agent_skills.py", + "controllers/console/app/test_annotation_api.py", + "controllers/console/app/test_annotation_security.py", + "controllers/console/app/test_app_apis.py", + "controllers/console/app/test_app_import_api.py", + "controllers/console/app/test_app_response_models.py", + "controllers/console/app/test_audio.py", + "controllers/console/app/test_conversation_api.py", + "controllers/console/app/test_conversation_variables_api.py", + "controllers/console/app/test_create_app_payload.py", + "controllers/console/app/test_description_validation.py", + "controllers/console/app/test_generator_api.py", + "controllers/console/app/test_generator_api_missing.py", + "controllers/console/app/test_mcp_server_response.py", + "controllers/console/app/test_message_api.py", + "controllers/console/app/test_model_config_api.py", + "controllers/console/app/test_ops_trace_api.py", + "controllers/console/app/test_statistic_api.py", + "controllers/console/app/test_workflow.py", + "controllers/console/app/test_workflow_app_log_api.py", + "controllers/console/app/test_workflow_comment_api.py", + "controllers/console/app/test_workflow_convert_api.py", + "controllers/console/app/test_workflow_node_output_inspector.py", + "controllers/console/app/test_workflow_pause_details_api.py", + "controllers/console/app/test_workflow_run_api.py", + "controllers/console/app/test_workflow_trigger_api.py", + "controllers/console/app/test_wraps.py", + "controllers/console/app/workflow_draft_variables_test.py", + "controllers/console/auth/test_account_activation.py", + "controllers/console/auth/test_authentication_security.py", + "controllers/console/auth/test_data_source_bearer_auth.py", + "controllers/console/auth/test_email_register.py", + "controllers/console/auth/test_email_register_language.py", + "controllers/console/auth/test_email_verification.py", + "controllers/console/auth/test_forgot_password.py", + "controllers/console/auth/test_login_logout.py", + "controllers/console/auth/test_oauth.py", + "controllers/console/auth/test_oauth_timezone.py", + "controllers/console/auth/test_password_reset.py", + "controllers/console/auth/test_token_refresh.py", + "controllers/console/billing/test_billing.py", + "controllers/console/datasets/rag_pipeline/test_datasource_auth.py", + "controllers/console/datasets/rag_pipeline/test_datasource_content_preview.py", + "controllers/console/datasets/rag_pipeline/test_rag_pipeline.py", + "controllers/console/datasets/rag_pipeline/test_rag_pipeline_draft_variable.py", + "controllers/console/datasets/rag_pipeline/test_rag_pipeline_workflow.py", + "controllers/console/datasets/test_data_source.py", + "controllers/console/datasets/test_datasets.py", + "controllers/console/datasets/test_datasets_document.py", + "controllers/console/datasets/test_datasets_document_download.py", + "controllers/console/datasets/test_datasets_segments.py", + "controllers/console/datasets/test_external.py", + "controllers/console/datasets/test_hit_testing.py", + "controllers/console/datasets/test_hit_testing_base.py", + "controllers/console/datasets/test_metadata.py", + "controllers/console/datasets/test_website.py", + "controllers/console/datasets/test_wraps.py", + "controllers/console/explore/test_audio.py", + "controllers/console/explore/test_banner.py", + "controllers/console/explore/test_completion.py", + "controllers/console/explore/test_installed_app.py", + "controllers/console/explore/test_message.py", + "controllers/console/explore/test_parameter.py", + "controllers/console/explore/test_recommended_app.py", + "controllers/console/explore/test_saved_message.py", + "controllers/console/explore/test_trial.py", + "controllers/console/explore/test_workflow.py", + "controllers/console/explore/test_wraps.py", + "controllers/console/snippets/test_snippet_workflow.py", + "controllers/console/snippets/test_snippet_workflow_draft_variable.py", + "controllers/console/tag/test_tags.py", + "controllers/console/test_document_detail_api_data_source_info.py", + "controllers/console/test_extension.py", + "controllers/console/test_fastopenapi_ping.py", + "controllers/console/test_fastopenapi_setup.py", + "controllers/console/test_fastopenapi_version.py", + "controllers/console/test_feature.py", + "controllers/console/test_files.py", + "controllers/console/test_files_security.py", + "controllers/console/test_human_input_form.py", + "controllers/console/test_knowledge_fs_proxy.py", + "controllers/console/test_remote_files.py", + "controllers/console/test_spec.py", + "controllers/console/test_version.py", + "controllers/console/test_workflow_run_archive.py", + "controllers/console/test_workspace_account.py", + "controllers/console/test_workspace_members.py", + "controllers/console/test_wraps.py", + "controllers/console/workspace/test_accounts.py", + "controllers/console/workspace/test_agent_providers.py", + "controllers/console/workspace/test_endpoint.py", + "controllers/console/workspace/test_load_balancing_config.py", + "controllers/console/workspace/test_members.py", + "controllers/console/workspace/test_model_providers.py", + "controllers/console/workspace/test_models.py", + "controllers/console/workspace/test_plugin.py", + "controllers/console/workspace/test_rbac.py", + "controllers/console/workspace/test_snippets.py", + "controllers/console/workspace/test_tool_providers.py", + "controllers/console/workspace/test_workspace.py", + "controllers/files/test_image_preview.py", + "controllers/files/test_tool_files.py", + "controllers/files/test_upload.py", + "controllers/inner_api/app/test_dsl.py", + "controllers/inner_api/plugin/test_agent_config.py", + "controllers/inner_api/plugin/test_agent_drive.py", + "controllers/inner_api/plugin/test_plugin.py", + "controllers/inner_api/plugin/test_plugin_wraps.py", + "controllers/inner_api/test_auth_wraps.py", + "controllers/inner_api/test_knowledge_retrieval.py", + "controllers/inner_api/test_mail.py", + "controllers/inner_api/test_runtime_credentials.py", + "controllers/inner_api/workspace/test_workspace.py", + "controllers/mcp/test_mcp.py", + "controllers/openapi/auth/test_composition.py", + "controllers/openapi/auth/test_conditions.py", + "controllers/openapi/auth/test_data.py", + "controllers/openapi/auth/test_flow.py", + "controllers/openapi/auth/test_pipeline.py", + "controllers/openapi/auth/test_prepare.py", + "controllers/openapi/auth/test_verify.py", + "controllers/openapi/conftest.py", + "controllers/openapi/test_account.py", + "controllers/openapi/test_app_describe_builder.py", + "controllers/openapi/test_app_list_query.py", + "controllers/openapi/test_app_payloads.py", + "controllers/openapi/test_app_run_dispatch.py", + "controllers/openapi/test_app_run_rate_limit.py", + "controllers/openapi/test_app_run_streaming.py", + "controllers/openapi/test_apps_permitted_external_query.py", + "controllers/openapi/test_audit_app_run.py", + "controllers/openapi/test_contract.py", + "controllers/openapi/test_cors.py", + "controllers/openapi/test_device_approve_deny.py", + "controllers/openapi/test_device_code.py", + "controllers/openapi/test_device_lookup.py", + "controllers/openapi/test_device_sso.py", + "controllers/openapi/test_device_token.py", + "controllers/openapi/test_error_contract.py", + "controllers/openapi/test_health.py", + "controllers/openapi/test_human_input_form.py", + "controllers/openapi/test_input_schema.py", + "controllers/openapi/test_meta_version.py", + "controllers/openapi/test_models.py", + "controllers/openapi/test_oauth_sso_claims.py", + "controllers/openapi/test_oauth_sso_csrf.py", + "controllers/openapi/test_oauth_sso_host_header.py", + "controllers/openapi/test_pagination_envelope.py", + "controllers/openapi/test_supported_app_type.py", + "controllers/openapi/test_version_gate.py", + "controllers/openapi/test_workflow_events_openapi.py", + "controllers/openapi/test_workspaces.py", + "controllers/openapi/test_workspaces_members.py", + "controllers/service_api/app/test_annotation.py", + "controllers/service_api/app/test_app.py", + "controllers/service_api/app/test_audio.py", + "controllers/service_api/app/test_chat_request_payload.py", + "controllers/service_api/app/test_completion.py", + "controllers/service_api/app/test_conversation.py", + "controllers/service_api/app/test_file.py", + "controllers/service_api/app/test_file_preview.py", + "controllers/service_api/app/test_hitl_service_api.py", + "controllers/service_api/app/test_human_input_form.py", + "controllers/service_api/app/test_message.py", + "controllers/service_api/app/test_workflow.py", + "controllers/service_api/app/test_workflow_events.py", + "controllers/service_api/conftest.py", + "controllers/service_api/dataset/rag_pipeline/test_rag_pipeline_workflow.py", + "controllers/service_api/dataset/test_dataset_segment.py", + "controllers/service_api/dataset/test_document.py", + "controllers/service_api/dataset/test_hit_testing.py", + "controllers/service_api/dataset/test_metadata.py", + "controllers/service_api/dataset/test_rag_pipeline_file_upload_serialization.py", + "controllers/service_api/dataset/test_rag_pipeline_route_registration.py", + "controllers/service_api/test_index.py", + "controllers/service_api/test_trace_session_id_parsing.py", + "controllers/service_api/test_wraps.py", + "controllers/test_compare_versions.py", + "controllers/test_conversation_rename_payload.py", + "controllers/test_swagger.py", + "controllers/trigger/test_trigger.py", + "controllers/trigger/test_webhook.py", + "controllers/web/conftest.py", + "controllers/web/test_app.py", + "controllers/web/test_completion.py", + "controllers/web/test_human_input_file_upload.py", + "controllers/web/test_human_input_form.py", + "controllers/web/test_message_endpoints.py", + "controllers/web/test_message_list.py", + "controllers/web/test_pydantic_models.py", + "controllers/web/test_web_forgot_password.py", + "controllers/web/test_web_login.py", + "controllers/web/test_web_passport.py", + "controllers/web/test_workflow.py", + "controllers/web/test_wraps.py", + "core/agent/conftest.py", + "core/agent/output_parser/test_cot_output_parser.py", + "core/agent/strategy/test_base.py", + "core/agent/strategy/test_plugin.py", + "core/agent/test_base_agent_runner.py", + "core/agent/test_cot_agent_runner.py", + "core/agent/test_cot_chat_agent_runner.py", + "core/agent/test_cot_completion_agent_runner.py", + "core/agent/test_fc_agent_runner.py", + "core/agent/test_plugin_entities.py", + "core/app/app_config/common/test_parameters_mapping.py", + "core/app/app_config/common/test_sensitive_word_avoidance_manager.py", + "core/app/app_config/easy_ui_based_app/test_agent_manager.py", + "core/app/app_config/easy_ui_based_app/test_dataset_manager.py", + "core/app/app_config/easy_ui_based_app/test_model_config_converter.py", + "core/app/app_config/easy_ui_based_app/test_model_config_manager.py", + "core/app/app_config/easy_ui_based_app/test_prompt_template_manager.py", + "core/app/app_config/easy_ui_based_app/test_variables_manager.py", + "core/app/app_config/features/file_upload/test_manager.py", + "core/app/app_config/features/test_additional_feature_managers.py", + "core/app/app_config/test_base_app_config_manager.py", + "core/app/app_config/test_entities.py", + "core/app/app_config/workflow_ui_based_app/test_workflow_ui_based_app_manager.py", + "core/app/apps/advanced_chat/test_app_config_manager.py", + "core/app/apps/advanced_chat/test_app_generator.py", + "core/app/apps/advanced_chat/test_app_runner_conversation_variables.py", + "core/app/apps/advanced_chat/test_app_runner_input_moderation.py", + "core/app/apps/advanced_chat/test_generate_response_converter.py", + "core/app/apps/advanced_chat/test_generate_task_pipeline.py", + "core/app/apps/advanced_chat/test_generate_task_pipeline_core.py", + "core/app/apps/agent_app/test_app_config_manager.py", + "core/app/apps/agent_app/test_app_generator.py", + "core/app/apps/agent_app/test_app_runner.py", + "core/app/apps/agent_app/test_input_guards.py", + "core/app/apps/agent_app/test_resolve_agent.py", + "core/app/apps/agent_app/test_runtime_request_builder.py", + "core/app/apps/agent_app/test_session_store.py", + "core/app/apps/agent_chat/test_agent_chat_app_config_manager.py", + "core/app/apps/agent_chat/test_agent_chat_app_generator.py", + "core/app/apps/agent_chat/test_agent_chat_app_runner.py", + "core/app/apps/agent_chat/test_agent_chat_generate_response_converter.py", + "core/app/apps/chat/test_app_config_manager.py", + "core/app/apps/chat/test_app_generator_and_runner.py", + "core/app/apps/chat/test_base_app_runner_multimodal.py", + "core/app/apps/chat/test_generate_response_converter.py", + "core/app/apps/common/test_graph_runtime_state_support.py", + "core/app/apps/common/test_workflow_response_converter.py", + "core/app/apps/common/test_workflow_response_converter_human_input.py", + "core/app/apps/common/test_workflow_response_converter_resumption.py", + "core/app/apps/common/test_workflow_response_converter_truncation.py", + "core/app/apps/completion/test_app_runner.py", + "core/app/apps/completion/test_completion_app_config_manager.py", + "core/app/apps/completion/test_completion_completion_app_generator.py", + "core/app/apps/completion/test_completion_generate_response_converter.py", + "core/app/apps/pipeline/test_pipeline_config_manager.py", + "core/app/apps/pipeline/test_pipeline_generate_response_converter.py", + "core/app/apps/pipeline/test_pipeline_generator.py", + "core/app/apps/pipeline/test_pipeline_queue_manager.py", + "core/app/apps/pipeline/test_pipeline_runner.py", + "core/app/apps/test_advanced_chat_app_generator.py", + "core/app/apps/test_base_app_generate_response_converter.py", + "core/app/apps/test_base_app_generator.py", + "core/app/apps/test_base_app_queue_manager.py", + "core/app/apps/test_base_app_runner.py", + "core/app/apps/test_exc.py", + "core/app/apps/test_message_based_app_generator.py", + "core/app/apps/test_message_based_app_queue_manager.py", + "core/app/apps/test_message_generator.py", + "core/app/apps/test_pause_resume.py", + "core/app/apps/test_streaming_utils.py", + "core/app/apps/test_trace_session_id_generate_extras.py", + "core/app/apps/test_workflow_app_generator.py", + "core/app/apps/test_workflow_app_runner_core.py", + "core/app/apps/test_workflow_app_runner_notifications.py", + "core/app/apps/test_workflow_app_runner_single_node.py", + "core/app/apps/test_workflow_pause_events.py", + "core/app/apps/workflow/test_active_workflow_tasks.py", + "core/app/apps/workflow/test_app_config_manager.py", + "core/app/apps/workflow/test_app_generator_extra.py", + "core/app/apps/workflow/test_app_queue_manager.py", + "core/app/apps/workflow/test_command_channels.py", + "core/app/apps/workflow/test_errors.py", + "core/app/apps/workflow/test_generate_response_converter.py", + "core/app/apps/workflow/test_generate_task_pipeline.py", + "core/app/apps/workflow/test_generate_task_pipeline_core.py", + "core/app/entities/test_app_invoke_entities.py", + "core/app/entities/test_queue_entities.py", + "core/app/entities/test_rag_pipeline_invoke_entities.py", + "core/app/entities/test_task_entities.py", + "core/app/features/rate_limiting/conftest.py", + "core/app/features/rate_limiting/test_rate_limit.py", + "core/app/features/test_annotation_reply.py", + "core/app/features/test_hosting_moderation.py", + "core/app/layers/test_conversation_variable_persist_layer.py", + "core/app/layers/test_pause_state_persist_layer.py", + "core/app/layers/test_suspend_layer.py", + "core/app/layers/test_timeslice_layer.py", + "core/app/layers/test_trigger_post_layer.py", + "core/app/task_pipeline/test_based_generate_task_pipeline.py", + "core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline.py", + "core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py", + "core/app/task_pipeline/test_exc.py", + "core/app/task_pipeline/test_message_cycle_manager_optimization.py", + "core/app/test_easy_ui_model_config_manager.py", + "core/app/test_invoke_from.py", + "core/app/test_llm_quota.py", + "core/app/workflow/layers/test_persistence.py", + "core/app/workflow/layers/test_persistence_inspector_publish.py", + "core/app/workflow/test_file_runtime.py", + "core/app/workflow/test_node_factory.py", + "core/app/workflow/test_observability_layer_extra.py", + "core/app/workflow/test_persistence_layer.py", + "core/base/test_app_generator_tts_publisher.py", + "core/callback_handler/test_agent_tool_callback_handler.py", + "core/callback_handler/test_index_tool_callback_handler.py", + "core/callback_handler/test_workflow_tool_callback_handler.py", + "core/datasource/__base/test_datasource_plugin.py", + "core/datasource/__base/test_datasource_provider.py", + "core/datasource/__base/test_datasource_runtime.py", + "core/datasource/entities/test_api_entities.py", + "core/datasource/entities/test_common_entities.py", + "core/datasource/entities/test_datasource_entities.py", + "core/datasource/local_file/test_local_file_plugin.py", + "core/datasource/local_file/test_local_file_provider.py", + "core/datasource/online_document/test_online_document_plugin.py", + "core/datasource/online_document/test_online_document_provider.py", + "core/datasource/online_drive/test_online_drive_plugin.py", + "core/datasource/online_drive/test_online_drive_provider.py", + "core/datasource/test_datasource_file_manager.py", + "core/datasource/test_datasource_manager.py", + "core/datasource/test_errors.py", + "core/datasource/test_file_upload.py", + "core/datasource/test_notion_provider.py", + "core/datasource/test_website_crawl.py", + "core/datasource/utils/test_message_transformer.py", + "core/datasource/website_crawl/test_website_crawl_plugin.py", + "core/datasource/website_crawl/test_website_crawl_provider.py", + "core/entities/test_entities_mcp_provider.py", + "core/entities/test_entities_provider_configuration.py", + "core/extension/test_api_based_extension_requestor.py", + "core/extension/test_extensible.py", + "core/extension/test_extension.py", + "core/external_data_tool/api/test_api.py", + "core/external_data_tool/test_base.py", + "core/external_data_tool/test_external_data_fetch.py", + "core/external_data_tool/test_factory.py", + "core/file/test_models.py", + "core/file/test_remote_fetcher.py", + "core/helper/code_executor/javascript/test_javascript_transformer.py", + "core/helper/code_executor/jinja2/test_jinja2_sandbox.py", + "core/helper/code_executor/python3/test_python3_transformer.py", + "core/helper/code_executor/test_code_executor.py", + "core/helper/code_executor/test_code_node_provider.py", + "core/helper/code_executor/test_template_transformer.py", + "core/helper/test_creators.py", + "core/helper/test_credential_utils.py", + "core/helper/test_csv_sanitizer.py", + "core/helper/test_encrypter.py", + "core/helper/test_ssrf_proxy.py", + "core/helper/test_trace_id_helper.py", + "core/llm_generator/output_parser/test_rule_config_generator.py", + "core/llm_generator/output_parser/test_structured_output.py", + "core/llm_generator/test_llm_generator.py", + "core/llm_generator/test_llm_generator_missing.py", + "core/logging/test_context.py", + "core/logging/test_filters.py", + "core/logging/test_structured_formatter.py", + "core/logging/test_trace_helpers.py", + "core/mcp/auth/test_auth_flow.py", + "core/mcp/client/test_session.py", + "core/mcp/client/test_sse.py", + "core/mcp/client/test_streamable_http.py", + "core/mcp/server/test_streamable_http.py", + "core/mcp/session/test_base_session.py", + "core/mcp/session/test_client_session.py", + "core/mcp/test_auth_client_inheritance.py", + "core/mcp/test_entities.py", + "core/mcp/test_error.py", + "core/mcp/test_mcp_client.py", + "core/mcp/test_types.py", + "core/mcp/test_utils.py", + "core/memory/test_token_buffer_memory.py", + "core/model_runtime/test_model_provider_factory.py", + "core/moderation/api/test_api.py", + "core/moderation/test_content_moderation.py", + "core/moderation/test_input_moderation.py", + "core/moderation/test_output_moderation.py", + "core/moderation/test_sensitive_word_filter.py", + "core/ops/test_base_trace_instance.py", + "core/ops/test_config_entity.py", + "core/ops/test_lookup_helpers.py", + "core/ops/test_ops_trace_manager.py", + "core/ops/test_trace_queue_manager.py", + "core/ops/test_trace_session_metadata.py", + "core/ops/test_utils.py", + "core/plugin/impl/test_agent_client.py", + "core/plugin/impl/test_asset_manager.py", + "core/plugin/impl/test_base_client_impl.py", + "core/plugin/impl/test_datasource_manager.py", + "core/plugin/impl/test_debugging_client.py", + "core/plugin/impl/test_endpoint_client_impl.py", + "core/plugin/impl/test_exc_impl.py", + "core/plugin/impl/test_model_client.py", + "core/plugin/impl/test_model_runtime_factory.py", + "core/plugin/impl/test_oauth_handler.py", + "core/plugin/impl/test_tool_manager.py", + "core/plugin/impl/test_trigger_client.py", + "core/plugin/test_backwards_invocation_app.py", + "core/plugin/test_backwards_invocation_model.py", + "core/plugin/test_endpoint_client.py", + "core/plugin/test_model_runtime_adapter.py", + "core/plugin/test_plugin_entities.py", + "core/plugin/test_plugin_manager.py", + "core/plugin/test_plugin_runtime.py", + "core/plugin/utils/test_chunk_merger.py", + "core/plugin/utils/test_http_parser.py", + "core/prompt/test_advanced_prompt_transform.py", + "core/prompt/test_agent_history_prompt_transform.py", + "core/prompt/test_extract_thread_messages.py", + "core/prompt/test_prompt_message.py", + "core/prompt/test_prompt_transform.py", + "core/prompt/test_simple_prompt_transform.py", + "core/rag/cleaner/test_clean_processor.py", + "core/rag/data_post_processor/test_data_post_processor.py", + "core/rag/datasource/keyword/jieba/test_jieba.py", + "core/rag/datasource/keyword/jieba/test_jieba_keyword_table_handler.py", + "core/rag/datasource/keyword/jieba/test_stopwords.py", + "core/rag/datasource/keyword/test_keyword_base.py", + "core/rag/datasource/keyword/test_keyword_factory.py", + "core/rag/datasource/test_datasource_retrieval.py", + "core/rag/datasource/vdb/test_field.py", + "core/rag/datasource/vdb/test_vector_base.py", + "core/rag/datasource/vdb/test_vector_factory.py", + "core/rag/docstore/test_dataset_docstore.py", + "core/rag/embedding/test_cached_embedding.py", + "core/rag/embedding/test_embedding_base.py", + "core/rag/embedding/test_embedding_service.py", + "core/rag/extractor/blob/test_blob.py", + "core/rag/extractor/firecrawl/test_firecrawl.py", + "core/rag/extractor/test_csv_extractor.py", + "core/rag/extractor/test_excel_extractor.py", + "core/rag/extractor/test_extract_processor.py", + "core/rag/extractor/test_extractor_base.py", + "core/rag/extractor/test_helpers.py", + "core/rag/extractor/test_html_extractor.py", + "core/rag/extractor/test_jina_reader_extractor.py", + "core/rag/extractor/test_markdown_extractor.py", + "core/rag/extractor/test_notion_extractor.py", + "core/rag/extractor/test_pdf_extractor.py", + "core/rag/extractor/test_text_extractor.py", + "core/rag/extractor/test_word_extractor.py", + "core/rag/extractor/unstructured/test_unstructured_extractors.py", + "core/rag/extractor/watercrawl/test_watercrawl.py", + "core/rag/indexing/processor/test_paragraph_index_processor.py", + "core/rag/indexing/processor/test_qa_index_processor.py", + "core/rag/indexing/test_index_processor.py", + "core/rag/indexing/test_index_processor_base.py", + "core/rag/indexing/test_indexing_runner.py", + "core/rag/pipeline/test_queue.py", + "core/rag/rerank/test_reranker.py", + "core/rag/retrieval/test_dataset_retrieval.py", + "core/rag/retrieval/test_dataset_retrieval_methods.py", + "core/rag/retrieval/test_multi_dataset_function_call_router.py", + "core/rag/retrieval/test_multi_dataset_react_route.py", + "core/rag/splitter/test_text_splitter.py", + "core/repositories/test_celery_workflow_execution_repository.py", + "core/repositories/test_celery_workflow_node_execution_repository.py", + "core/repositories/test_factory.py", + "core/repositories/test_human_input_form_repository_impl.py", + "core/repositories/test_human_input_repository.py", + "core/repositories/test_sqlalchemy_workflow_execution_repository.py", + "core/repositories/test_sqlalchemy_workflow_node_execution_repository.py", + "core/repositories/test_workflow_node_execution_conflict_handling.py", + "core/repositories/test_workflow_node_execution_truncation.py", + "core/schemas/test_registry.py", + "core/schemas/test_resolver.py", + "core/schemas/test_schema_manager.py", + "core/telemetry/test_facade.py", + "core/telemetry/test_gateway_integration.py", + "core/test_file.py", + "core/test_model_manager.py", + "core/test_provider_configuration.py", + "core/test_provider_manager.py", + "core/test_trigger_debug_event_selectors.py", + "core/tools/entities/test_api_entities.py", + "core/tools/test_base_tool.py", + "core/tools/test_builtin_tool_base.py", + "core/tools/test_builtin_tool_provider.py", + "core/tools/test_builtin_tools_extra.py", + "core/tools/test_custom_tool.py", + "core/tools/test_custom_tool_provider.py", + "core/tools/test_dataset_retriever_tool.py", + "core/tools/test_mcp_tool.py", + "core/tools/test_mcp_tool_provider.py", + "core/tools/test_plugin_tool.py", + "core/tools/test_plugin_tool_provider.py", + "core/tools/test_tool_engine.py", + "core/tools/test_tool_entities.py", + "core/tools/test_tool_file_manager.py", + "core/tools/test_tool_label_manager.py", + "core/tools/test_tool_manager.py", + "core/tools/test_tool_parameter_type.py", + "core/tools/test_tool_provider_controller.py", + "core/tools/utils/test_configuration.py", + "core/tools/utils/test_encryption.py", + "core/tools/utils/test_message_transformer.py", + "core/tools/utils/test_misc_utils_extra.py", + "core/tools/utils/test_model_invocation_utils.py", + "core/tools/utils/test_parser.py", + "core/tools/utils/test_system_oauth_encryption.py", + "core/tools/utils/test_tool_engine_serialization.py", + "core/tools/utils/test_web_reader_tool.py", + "core/tools/utils/test_workflow_configuration_sync.py", + "core/tools/workflow_as_tool/test_provider.py", + "core/tools/workflow_as_tool/test_tool.py", + "core/trigger/conftest.py", + "core/trigger/debug/test_debug_event_bus.py", + "core/trigger/debug/test_debug_event_selectors.py", + "core/trigger/test_provider.py", + "core/trigger/test_trigger_manager.py", + "core/trigger/utils/test_utils_encryption.py", + "core/trigger/utils/test_utils_endpoint.py", + "core/trigger/utils/test_utils_locks.py", + "core/variables/test_segment.py", + "core/variables/test_segment_type.py", + "core/variables/test_segment_type_validation.py", + "core/variables/test_variables.py", + "core/workflow/context/test_execution_context.py", + "core/workflow/context/test_flask_app_context.py", + "core/workflow/entities/test_private_workflow_pause.py", + "core/workflow/generator/test_prompts.py", + "core/workflow/generator/test_runner.py", + "core/workflow/generator/test_runner_missing.py", + "core/workflow/generator/test_tool_catalogue.py", + "core/workflow/graph_engine/layers/conftest.py", + "core/workflow/graph_engine/layers/test_llm_quota.py", + "core/workflow/graph_engine/layers/test_observability.py", + "core/workflow/graph_engine/test_mock_config.py", + "core/workflow/graph_engine/test_mock_factory.py", + "core/workflow/graph_engine/test_mock_nodes.py", + "core/workflow/graph_engine/test_parallel_human_input_join_resume.py", + "core/workflow/graph_engine/test_table_runner.py", + "core/workflow/graph_engine/test_tool_in_chatflow.py", + "core/workflow/nodes/agent/test_message_transformer.py", + "core/workflow/nodes/agent/test_runtime_support.py", + "core/workflow/nodes/agent_v2/test_agent_node.py", + "core/workflow/nodes/agent_v2/test_binding_resolver.py", + "core/workflow/nodes/agent_v2/test_dify_tools_builder.py", + "core/workflow/nodes/agent_v2/test_file_tenant_validator.py", + "core/workflow/nodes/agent_v2/test_output_adapter.py", + "core/workflow/nodes/agent_v2/test_output_failure_orchestrator.py", + "core/workflow/nodes/agent_v2/test_output_file_rebacker.py", + "core/workflow/nodes/agent_v2/test_output_type_checker.py", + "core/workflow/nodes/agent_v2/test_runtime_request_builder.py", + "core/workflow/nodes/agent_v2/test_session_cleanup_layer.py", + "core/workflow/nodes/agent_v2/test_session_store.py", + "core/workflow/nodes/agent_v2/test_validators.py", + "core/workflow/nodes/answer/test_answer.py", + "core/workflow/nodes/base/test_base_node.py", + "core/workflow/nodes/base/test_get_node_type_classes_mapping.py", + "core/workflow/nodes/code/code_node_spec.py", + "core/workflow/nodes/datasource/test_datasource_node.py", + "core/workflow/nodes/http_request/test_http_request_executor.py", + "core/workflow/nodes/http_request/test_http_request_node.py", + "core/workflow/nodes/human_input/test_dify_owned_contracts.py", + "core/workflow/nodes/human_input/test_email_delivery_config.py", + "core/workflow/nodes/human_input/test_entities.py", + "core/workflow/nodes/human_input/test_human_input_form_filled_event.py", + "core/workflow/nodes/iteration/test_iteration_child_engine_errors.py", + "core/workflow/nodes/knowledge_index/test_knowledge_index_node.py", + "core/workflow/nodes/knowledge_retrieval/test_knowledge_retrieval_node.py", + "core/workflow/nodes/list_operator/node_spec.py", + "core/workflow/nodes/llm/test_llm_utils.py", + "core/workflow/nodes/llm/test_node.py", + "core/workflow/nodes/parameter_extractor/test_parameter_extractor_node.py", + "core/workflow/nodes/template_transform/template_transform_node_spec.py", + "core/workflow/nodes/template_transform/test_template_transform_node.py", + "core/workflow/nodes/test_base_node.py", + "core/workflow/nodes/test_document_extractor_node.py", + "core/workflow/nodes/test_if_else.py", + "core/workflow/nodes/test_list_operator.py", + "core/workflow/nodes/test_start_node_json_object.py", + "core/workflow/nodes/tool/test_tool_node.py", + "core/workflow/nodes/tool/test_tool_node_runtime.py", + "core/workflow/nodes/webhook/test_entities.py", + "core/workflow/nodes/webhook/test_exceptions.py", + "core/workflow/nodes/webhook/test_webhook_file_conversion.py", + "core/workflow/nodes/webhook/test_webhook_node.py", + "core/workflow/test_enrich_pause_reasons.py", + "core/workflow/test_form_input_serialization_compat.py", + "core/workflow/test_graph_topology.py", + "core/workflow/test_human_input_adapter.py", + "core/workflow/test_human_input_callback.py", + "core/workflow/test_human_input_forms.py", + "core/workflow/test_node_factory.py", + "core/workflow/test_node_mapping_bootstrap.py", + "core/workflow/test_node_runtime.py", + "core/workflow/test_system_variable.py", + "core/workflow/test_variable_pool.py", + "core/workflow/test_workflow_entry.py", + "core/workflow/test_workflow_entry_helpers.py", + "core/workflow/test_workflow_entry_redis_channel.py", + "dev/test_generate_knowledge_fs_contract.py", + "enterprise/telemetry/test_contracts.py", + "enterprise/telemetry/test_draft_trace.py", + "enterprise/telemetry/test_enterprise_trace.py", + "enterprise/telemetry/test_event_handlers.py", + "enterprise/telemetry/test_exporter.py", + "enterprise/telemetry/test_gateway.py", + "enterprise/telemetry/test_metric_handler.py", + "enums/test_quota_type.py", + "events/event_handlers/test_clean_when_document_deleted.py", + "events/event_handlers/test_delete_tool_parameters_cache_when_sync_draft_workflow.py", + "events/test_app_event_signals.py", + "events/test_events_package_compat.py", + "extensions/logstore/repositories/test_logstore_api_workflow_node_execution_repository.py", + "extensions/logstore/test_sql_escape.py", + "extensions/otel/conftest.py", + "extensions/otel/decorators/handlers/test_generate_handler.py", + "extensions/otel/decorators/handlers/test_workflow_app_runner_handler.py", + "extensions/otel/decorators/test_base.py", + "extensions/otel/decorators/test_handler.py", + "extensions/otel/test_celery_sqlcommenter.py", + "extensions/otel/test_context.py", + "extensions/otel/test_retrieval_tracing.py", + "extensions/otel/test_runtime.py", + "extensions/storage/test_supabase_storage.py", + "extensions/test_celery_ssl.py", + "extensions/test_ext_blueprints_openapi.py", + "extensions/test_ext_login.py", + "extensions/test_ext_request_logging.py", + "extensions/test_ext_socketio.py", + "extensions/test_pubsub_channel.py", + "extensions/test_redis.py", + "extensions/test_set_secretkey.py", + "extensions/test_workflow_warm_shutdown.py", + "factories/test_build_from_mapping.py", + "factories/test_file_factory.py", + "factories/test_file_validation.py", + "factories/test_variable_factory.py", + "fields/test_dataset_fields.py", + "fields/test_file_fields.py", + "fields/test_message_fields.py", + "libs/_human_input/support.py", + "libs/_human_input/test_form_service.py", + "libs/_human_input/test_models.py", + "libs/broadcast_channel/redis/test_channel_unit_tests.py", + "libs/broadcast_channel/redis/test_streams_channel_unit_tests.py", + "libs/test_api_token_cache.py", + "libs/test_archive_storage.py", + "libs/test_cron_compatibility.py", + "libs/test_custom_inputs.py", + "libs/test_datetime_utils.py", + "libs/test_email.py", + "libs/test_email_i18n.py", + "libs/test_encryption.py", + "libs/test_external_api.py", + "libs/test_file_utils.py", + "libs/test_flask_utils.py", + "libs/test_helper.py", + "libs/test_json_in_md_parser.py", + "libs/test_jwt_imports.py", + "libs/test_login.py", + "libs/test_oauth_base.py", + "libs/test_oauth_bearer.py", + "libs/test_oauth_bearer_layer0_cache.py", + "libs/test_oauth_bearer_rate_limit_ordering.py", + "libs/test_oauth_bearer_require_scope.py", + "libs/test_oauth_clients.py", + "libs/test_orjson.py", + "libs/test_pagination.py", + "libs/test_pandas.py", + "libs/test_passport.py", + "libs/test_password.py", + "libs/test_rate_limit_bearer.py", + "libs/test_rate_limiter.py", + "libs/test_rsa.py", + "libs/test_schedule_utils_enhanced.py", + "libs/test_sendgrid_client.py", + "libs/test_smtp_client.py", + "libs/test_time_parser.py", + "libs/test_token.py", + "libs/test_token_manager.py", + "libs/test_uuid_utils.py", + "libs/test_workspace_member_helper.py", + "libs/test_workspace_permission.py", + "libs/test_yarl.py", + "migrations/test_agent_drive_skill_metadata_refactor.py", + "migrations/test_uuidv7_pg18_migration.py", + "models/test_account_models.py", + "models/test_agent.py", + "models/test_app_models.py", + "models/test_base.py", + "models/test_conversation_variable.py", + "models/test_dataset_models.py", + "models/test_end_user_type.py", + "models/test_enums_creator_user_role.py", + "models/test_model.py", + "models/test_plugin_entities.py", + "models/test_provider_models.py", + "models/test_tool_models.py", + "models/test_types.py", + "models/test_workflow.py", + "models/test_workflow_models.py", + "models/test_workflow_node_execution_offload.py", + "oss/__mock/aliyun_oss.py", + "oss/__mock/baidu_obs.py", + "oss/__mock/base.py", + "oss/__mock/local.py", + "oss/__mock/tencent_cos.py", + "oss/__mock/volcengine_tos.py", + "oss/aliyun_oss/aliyun_oss/test_aliyun_oss.py", + "oss/baidu_obs/test_baidu_obs.py", + "oss/opendal/test_opendal.py", + "oss/tencent_cos/test_tencent_cos.py", + "oss/volcengine_tos/test_volcengine_tos.py", + "repositories/test_sqlalchemy_api_workflow_run_repository.py", + "services/agent/test_agent_composer_entities.py", + "services/agent/test_agent_dsl_service.py", + "services/agent/test_agent_observability_service.py", + "services/agent/test_agent_services.py", + "services/agent/test_composer_candidates.py", + "services/agent/test_composer_mention_validation.py", + "services/agent/test_prompt_mentions.py", + "services/agent/test_skill_package_service.py", + "services/agent/test_skill_standardize_service.py", + "services/agent/test_skill_tool_inference_service.py", + "services/agent/test_workflow_publish_service.py", + "services/auth/test_api_key_auth_base.py", + "services/auth/test_api_key_auth_factory.py", + "services/auth/test_api_key_auth_service.py", + "services/auth/test_auth_type.py", + "services/auth/test_firecrawl_auth.py", + "services/auth/test_jina_auth.py", + "services/auth/test_jina_auth_standalone_module.py", + "services/auth/test_watercrawl_auth.py", + "services/controller_api.py", + "services/data_migration/test_dependency_discovery_service.py", + "services/data_migration/test_entities.py", + "services/data_migration/test_export_service.py", + "services/data_migration/test_import_service.py", + "services/data_migration/test_package_service.py", + "services/data_migration/test_report_service.py", + "services/dataset_service_test_helpers.py", + "services/document_service_validation.py", + "services/enterprise/test_account_deletion_sync.py", + "services/enterprise/test_app_permitted_service.py", + "services/enterprise/test_enterprise_service.py", + "services/enterprise/test_plugin_manager_service.py", + "services/enterprise/test_rbac_service.py", + "services/enterprise/test_traceparent_propagation.py", + "services/hit_service.py", + "services/openapi/test_mint_policy.py", + "services/plugin/conftest.py", + "services/plugin/test_dependencies_analysis.py", + "services/plugin/test_endpoint_service.py", + "services/plugin/test_oauth_service.py", + "services/plugin/test_plugin_migration.py", + "services/plugin/test_plugin_parameter_service.py", + "services/plugin/test_plugin_service.py", + "services/plugin/test_plugin_service_installation.py", + "services/rag_pipeline/pipeline_template/test_built_in_retrieval.py", + "services/rag_pipeline/pipeline_template/test_pipeline_template_base.py", + "services/rag_pipeline/test_pipeline_generate_service.py", + "services/rag_pipeline/test_rag_pipeline_dsl_service.py", + "services/rag_pipeline/test_rag_pipeline_service.py", + "services/rag_pipeline/test_rag_pipeline_task_proxy.py", + "services/rag_pipeline/test_rag_pipeline_transform_service.py", + "services/recommend_app/test_buildin_retrieval.py", + "services/recommend_app/test_category_order.py", + "services/recommend_app/test_recommend_app_factory.py", + "services/recommend_app/test_recommend_app_type.py", + "services/recommend_app/test_remote_retrieval.py", + "services/retention/test_messages_clean_policy.py", + "services/retention/workflow_run/test_archive_download_preparation.py", + "services/retention/workflow_run/test_archive_download_task_cache.py", + "services/retention/workflow_run/test_archive_log_service.py", + "services/retention/workflow_run/test_bundle_archive_maintenance.py", + "services/retention/workflow_run/test_clear_free_plan_expired_workflow_run_logs.py", + "services/retention/workflow_run/test_restore_archived_workflow_run.py", + "services/test_agent_app_feature_service.py", + "services/test_agent_app_sandbox_service.py", + "services/test_agent_config_service.py", + "services/test_agent_drive_service.py", + "services/test_annotation_service.py", + "services/test_api_token_service.py", + "services/test_app_generate_service.py", + "services/test_app_generate_service_streaming_integration.py", + "services/test_app_model_config_service.py", + "services/test_app_service.py", + "services/test_app_task_service.py", + "services/test_archive_workflow_run_logs.py", + "services/test_async_workflow_service.py", + "services/test_audio_service.py", + "services/test_billing_service.py", + "services/test_clear_free_plan_expired_workflow_run_logs.py", + "services/test_clear_free_plan_tenant_expired_logs.py", + "services/test_code_based_extension_service.py", + "services/test_conversation_service.py", + "services/test_credit_pool_service.py", + "services/test_dataset_service_dataset.py", + "services/test_dataset_service_document.py", + "services/test_dataset_service_lock_not_owned.py", + "services/test_dataset_service_segment.py", + "services/test_datasource_provider_service.py", + "services/test_document_indexing_task_proxy.py", + "services/test_duplicate_document_indexing_task_proxy.py", + "services/test_export_app_messages.py", + "services/test_external_dataset_service.py", + "services/test_feature_service_app_dsl_version.py", + "services/test_feature_service_enable_app_deploy.py", + "services/test_feature_service_human_input_email_delivery.py", + "services/test_feature_service_learn_app.py", + "services/test_feature_service_licensed_seats.py", + "services/test_feature_service_trial_models.py", + "services/test_feature_service_vector_space.py", + "services/test_feature_service_webapp_public_access.py", + "services/test_feedback_service.py", + "services/test_file_service.py", + "services/test_human_input_delivery_test_service.py", + "services/test_human_input_file_upload_service.py", + "services/test_knowledge_retrieval_inner_service.py", + "services/test_knowledge_service.py", + "services/test_message_service.py", + "services/test_messages_clean_service.py", + "services/test_metadata_nullable_bug.py", + "services/test_metadata_service_session_boundary.py", + "services/test_model_load_balancing_service.py", + "services/test_model_provider_service.py", + "services/test_model_provider_service_sanitization.py", + "services/test_oauth_device_flow.py", + "services/test_oauth_server_service.py", + "services/test_operation_service.py", + "services/test_rag_pipeline_task_proxy.py", + "services/test_schedule_service.py", + "services/test_snippet_dsl_service.py", + "services/test_snippet_generate_service.py", + "services/test_snippet_service.py", + "services/test_step_by_step_tour_service.py", + "services/test_summary_index_service.py", + "services/test_telemetry_service.py", + "services/test_trigger_provider_service.py", + "services/test_variable_truncator.py", + "services/test_vector_service.py", + "services/test_webhook_service.py", + "services/test_webhook_service_additional.py", + "services/test_website_service.py", + "services/test_workflow_app_service_metadata.py", + "services/test_workflow_collaboration_service.py", + "services/test_workflow_comment_service.py", + "services/test_workflow_generator_service.py", + "services/test_workflow_node_execution_trace_service.py", + "services/test_workflow_run_service.py", + "services/test_workflow_run_service_pause.py", + "services/test_workflow_service.py", + "services/tools/test_api_tools_manage_service.py", + "services/tools/test_builtin_tools_manage_service.py", + "services/tools/test_mcp_tools_transform.py", + "services/tools/test_tool_labels_service.py", + "services/tools/test_tools_manage_service.py", + "services/tools/test_tools_transform_service.py", + "services/workflow/test_draft_var_loader_simple.py", + "services/workflow/test_inspector_events.py", + "services/workflow/test_node_output_inspector_service.py", + "services/workflow/test_queue_dispatcher.py", + "services/workflow/test_scheduler.py", + "services/workflow/test_workflow_converter_additional.py", + "services/workflow/test_workflow_draft_variable_service.py", + "services/workflow/test_workflow_event_snapshot_service.py", + "services/workflow/test_workflow_event_snapshot_service_additional.py", + "services/workflow/test_workflow_human_input_delivery.py", + "services/workflow/test_workflow_restore.py", + "tasks/test_agent_backend_session_cleanup_task.py", + "tasks/test_async_workflow_tasks.py", + "tasks/test_batch_clean_document_task.py", + "tasks/test_clean_dataset_task.py", + "tasks/test_clean_document_task.py", + "tasks/test_community_telemetry_task.py", + "tasks/test_dataset_indexing_task.py", + "tasks/test_document_indexing_sync_task.py", + "tasks/test_document_indexing_update_task.py", + "tasks/test_duplicate_document_indexing_task.py", + "tasks/test_enable_segment_index_tasks.py", + "tasks/test_enterprise_telemetry_task.py", + "tasks/test_human_input_timeout_tasks.py", + "tasks/test_initialize_created_app_rbac_access_task.py", + "tasks/test_install_default_plugins_task.py", + "tasks/test_mail_human_input_delivery_task.py", + "tasks/test_mail_send_task.py", + "tasks/test_ops_trace_task.py", + "tasks/test_process_tenant_plugin_autoupgrade_check_task.py", + "tasks/test_refresh_billing_vector_space_task.py", + "tasks/test_remove_app_and_related_data_task.py", + "tasks/test_resume_agent_app_task.py", + "tasks/test_summary_queue_isolation.py", + "tasks/test_trigger_processing_tasks.py", + "tasks/test_workflow_execute_task.py", + "test_app_factory.py", + "test_makefile_backend_tests.py", + "test_pytest_dify.py", + "tools/test_api_tool.py", + "tools/test_mcp_tool.py", + "utils/encryption/test_system_encryption.py", + "utils/http_parser/test_oauth_convert_request_to_raw_data.py", + "utils/position_helper/test_position_helper.py", + "utils/structured_output_parser/test_structured_output_parser.py", + "utils/test_text_processing.py", + "utils/yaml/test_yaml_utils.py", +] + +[errors] +missing-override-decorator = "error" +redundant-cast = true +unannotated-return = true +unnecessary-type-conversion = true +unused-ignore = true +implicit-any-lambda = "info" +missing-attribute = "warn" diff --git a/api/tests/unit_tests/repositories/test_installation_state_repository.py b/api/tests/unit_tests/repositories/test_installation_state_repository.py new file mode 100644 index 00000000000..9f82ef8cd25 --- /dev/null +++ b/api/tests/unit_tests/repositories/test_installation_state_repository.py @@ -0,0 +1,40 @@ +"""Persistence tests for installation setup and tenant-existence state.""" + +from sqlalchemy.orm import Session, sessionmaker + +from models.account import Tenant +from models.model import DifySetup +from repositories.installation_state_repository import InstallationStateRepository + + +def test_empty_database_has_no_installation_state(sqlite_session_factory: sessionmaker[Session]) -> None: + repository = InstallationStateRepository(client=sqlite_session_factory) + + assert repository.get_setup_at() is None + assert repository.is_setup() is False + assert repository.has_tenants() is False + + +def test_get_setup_at_returns_persisted_timestamp( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + setup = DifySetup(version="test-version") + sqlite_session.add(setup) + sqlite_session.commit() + sqlite_session.refresh(setup) + repository = InstallationStateRepository(client=sqlite_session_factory) + + assert repository.get_setup_at() == setup.setup_at + assert repository.is_setup() is True + + +def test_has_tenants_detects_existing_tenant( + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + sqlite_session.add(Tenant(name="Existing workspace")) + sqlite_session.commit() + repository = InstallationStateRepository(client=sqlite_session_factory) + + assert repository.has_tenants() is True diff --git a/api/tests/unit_tests/repositories/test_workspace_member_query_repository.py b/api/tests/unit_tests/repositories/test_workspace_member_query_repository.py new file mode 100644 index 00000000000..852dc543e67 --- /dev/null +++ b/api/tests/unit_tests/repositories/test_workspace_member_query_repository.py @@ -0,0 +1,126 @@ +from datetime import datetime + +from sqlalchemy.orm import Session, sessionmaker + +from models.account import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole +from repositories.workspace_member_query_repository import WorkspaceMemberQueryRepository +from services.workspace_member_query_service import WorkspaceMemberRecord + + +def make_account( + account_id: str, + *, + status: AccountStatus, + created_at: datetime, +) -> Account: + account = Account( + name=f"Member {account_id}", + email=f"{account_id}@example.com", + avatar=f"{account_id}.png", + status=status, + ) + account.id = account_id + account.last_login_at = created_at + account.last_active_at = created_at + account.created_at = created_at + return account + + +def make_tenant(tenant_id: str) -> Tenant: + tenant = Tenant(name=f"Workspace {tenant_id}") + tenant.id = tenant_id + return tenant + + +def test_list_for_workspace_uses_join_membership_and_preserves_account_lifecycle( + sqlite_session_factory: sessionmaker[Session], +) -> None: + created_at = datetime(2026, 1, 1) + active = make_account("active", status=AccountStatus.ACTIVE, created_at=created_at) + uninitialized = make_account("uninitialized", status=AccountStatus.UNINITIALIZED, created_at=created_at) + pending = make_account("pending", status=AccountStatus.PENDING, created_at=created_at) + banned = make_account("banned", status=AccountStatus.BANNED, created_at=created_at) + closed = make_account("closed", status=AccountStatus.CLOSED, created_at=created_at) + other_workspace_member = make_account("other", status=AccountStatus.ACTIVE, created_at=created_at) + unjoined = make_account("unjoined", status=AccountStatus.ACTIVE, created_at=created_at) + workspace = make_tenant("workspace-1") + other_workspace = make_tenant("workspace-2") + + with sqlite_session_factory() as session: + session.add_all( + [ + workspace, + other_workspace, + active, + uninitialized, + pending, + banned, + closed, + other_workspace_member, + unjoined, + TenantAccountJoin( + tenant_id=workspace.id, + account_id=active.id, + role=TenantAccountRole.OWNER, + ), + TenantAccountJoin( + tenant_id=workspace.id, + account_id=uninitialized.id, + role=TenantAccountRole.NORMAL, + ), + TenantAccountJoin( + tenant_id=workspace.id, + account_id=pending.id, + role=TenantAccountRole.NORMAL, + ), + TenantAccountJoin( + tenant_id=workspace.id, + account_id=banned.id, + role=TenantAccountRole.ADMIN, + ), + TenantAccountJoin( + tenant_id=workspace.id, + account_id=closed.id, + role=TenantAccountRole.EDITOR, + ), + TenantAccountJoin( + tenant_id=other_workspace.id, + account_id=other_workspace_member.id, + role=TenantAccountRole.ADMIN, + ), + ] + ) + session.commit() + + result = WorkspaceMemberQueryRepository(sqlite_session_factory).list_for_workspace(workspace.id) + + by_id = {member.id: member for member in result} + assert set(by_id) == {"active", "uninitialized", "pending", "banned", "closed"} + assert by_id["active"] == WorkspaceMemberRecord( + id=active.id, + name=active.name, + email=active.email, + avatar=active.avatar, + last_login_at=created_at, + last_active_at=created_at, + created_at=created_at, + status=AccountStatus.ACTIVE.value, + legacy_role=TenantAccountRole.OWNER.value, + ) + assert by_id["uninitialized"].status == AccountStatus.UNINITIALIZED.value + assert by_id["pending"].status == AccountStatus.PENDING.value + assert by_id["pending"].legacy_role == TenantAccountRole.NORMAL.value + assert by_id["banned"].status == AccountStatus.BANNED.value + assert by_id["closed"].status == AccountStatus.CLOSED.value + + +def test_list_for_workspace_returns_empty_tuple_without_membership( + sqlite_session_factory: sessionmaker[Session], +) -> None: + with sqlite_session_factory() as session: + session.add(make_tenant("workspace-1")) + session.commit() + + result = WorkspaceMemberQueryRepository(sqlite_session_factory).list_for_workspace("workspace-1") + + assert result == () diff --git a/api/tests/unit_tests/services/agent/test_agent_composer_entities.py b/api/tests/unit_tests/services/agent/test_agent_composer_entities.py index 4aaae11b7dc..efc29ffb602 100644 --- a/api/tests/unit_tests/services/agent/test_agent_composer_entities.py +++ b/api/tests/unit_tests/services/agent/test_agent_composer_entities.py @@ -365,6 +365,51 @@ def test_knowledge_runtime_requirements_block_publish_but_not_draft_save(knowled ComposerConfigValidator.validate_publish_payload(publish_payload) +def test_manual_metadata_filtering_condition_accepts_composer_ui_identifiers(): + """The composer's condition editor always sends ``id`` (list row key) and + ``metadata_id`` (selected metadata field reference) alongside every + condition. Rejecting them as unknown fields broke every save of a manual + metadata filter (see GH issue #40169).""" + payload = ComposerSavePayload.model_validate( + { + "variant": ComposerVariant.AGENT_APP, + "save_strategy": ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION, + "agent_soul": { + "knowledge": { + "sets": [ + { + "id": "support", + "name": "Support KB", + "datasets": [{"id": "dataset-1"}], + "query": {"mode": "generated_query"}, + "retrieval": {"mode": "multiple", "top_k": 4}, + "metadata_filtering": { + "mode": "manual", + "conditions": { + "logical_operator": "and", + "conditions": [ + { + "id": "b149eceb-191a-40a2-9f11-61cf21ebd147", + "metadata_id": "ad6cf326-eadf-46e8-a2d5-9cb892d2cc84", + "name": "category", + "comparison_operator": "is", + "value": "auth", + } + ], + }, + }, + }, + ] + } + }, + } + ) + + condition = payload.agent_soul.knowledge.sets[0].metadata_filtering.conditions.conditions[0] + assert condition.id == "b149eceb-191a-40a2-9f11-61cf21ebd147" + assert condition.metadata_id == "ad6cf326-eadf-46e8-a2d5-9cb892d2cc84" + + def test_agent_soul_model_config_is_first_class_without_credentials(): config = AgentSoulConfig( model=AgentSoulModelConfig( diff --git a/api/tests/unit_tests/services/agent/test_agent_dsl_service.py b/api/tests/unit_tests/services/agent/test_agent_dsl_service.py index 4ceecc12a65..4ebf30c25ad 100644 --- a/api/tests/unit_tests/services/agent/test_agent_dsl_service.py +++ b/api/tests/unit_tests/services/agent/test_agent_dsl_service.py @@ -51,6 +51,7 @@ def _snapshot(*, snapshot_id: str = "snapshot-1", soul: AgentSoulConfig | None = tenant_id="tenant-1", agent_id="agent-1", version=1, + home_snapshot_id="home-1", config_snapshot=soul or AgentSoulConfig(), created_by="account-1", ) @@ -190,7 +191,7 @@ def test_agent_package_rejects_null_file_id_for_available_assets(asset: dict) -> AgentPackage.model_validate(package) -def test_import_warnings_cover_runtime_setup_removed_from_package(monkeypatch) -> None: +def test_import_warnings_cover_runtime_setup_removed_from_package(monkeypatch: pytest.MonkeyPatch) -> None: soul = AgentSoulConfig.model_validate( { "tools": { @@ -325,7 +326,7 @@ def test_graph_without_package_bindings_removes_portable_fields() -> None: assert AGENT_NODE_JOB_DSL_KEY in graph["nodes"][0]["data"] -def test_import_agent_app_package_creates_config_and_unpublished_draft(monkeypatch) -> None: +def test_import_agent_app_package_creates_config_and_unpublished_draft(monkeypatch: pytest.MonkeyPatch) -> None: session = Mock() service = AgentDslService(session) soul = AgentSoulConfig(config_note="portable") @@ -385,7 +386,11 @@ def test_import_workflow_packages_materializes_every_package_binding_as_inline() } for node in graph["nodes"][:3]: node["data"][AGENT_NODE_JOB_DSL_KEY] = {"workflow_prompt": node["id"]} - old_binding = SimpleNamespace(id="old-binding") + old_binding = SimpleNamespace( + id="old-binding", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id="old-inline-agent", + ) session = Mock() session.scalars.return_value.all.return_value = [old_binding] service = AgentDslService(session) @@ -406,7 +411,7 @@ def test_import_workflow_packages_materializes_every_package_binding_as_inline() graph="{}", ) - result, warnings = service.import_workflow_packages( + result, warnings, retirement_candidates = service.import_workflow_packages( workflow=workflow, portable_graph=graph, raw_packages={"agent_1": package.model_dump(mode="json")}, @@ -414,6 +419,7 @@ def test_import_workflow_packages_materializes_every_package_binding_as_inline() ) session.delete.assert_called_once_with(old_binding) + assert retirement_candidates == {"old-inline-agent"} assert service._create_imported_inline_agent.call_count == 3 assert [call.kwargs["node_id"] for call in service._create_imported_inline_agent.call_args_list] == [ "roster-1", @@ -459,7 +465,7 @@ def test_import_workflow_packages_rejects_invalid_package_binding(binding: dict, ) -def test_clone_inline_binding_copies_soul_and_drive_rows(monkeypatch) -> None: +def test_clone_inline_binding_copies_soul_and_drive_rows(monkeypatch: pytest.MonkeyPatch) -> None: session = Mock() service = AgentDslService(session) target_agent = SimpleNamespace(id="target-agent") @@ -499,7 +505,7 @@ def test_clone_inline_binding_copies_soul_and_drive_rows(monkeypatch) -> None: ) -def test_extract_package_dependencies_covers_model_tools_and_knowledge(monkeypatch) -> None: +def test_extract_package_dependencies_covers_model_tools_and_knowledge(monkeypatch: pytest.MonkeyPatch) -> None: model_dependency = Mock(side_effect=lambda provider: f"model:{provider}") tool_dependency = Mock(side_effect=lambda provider: f"tool:{provider}") monkeypatch.setattr( @@ -582,7 +588,7 @@ def test_create_imported_inline_agent_uses_import_provenance() -> None: ) -def test_create_workflow_only_agent_sets_backing_app_and_snapshot(monkeypatch) -> None: +def test_create_workflow_only_agent_sets_backing_app_and_snapshot(monkeypatch: pytest.MonkeyPatch) -> None: session = Mock() service = AgentDslService(session) roster_service = Mock() @@ -612,7 +618,7 @@ def test_create_workflow_only_agent_sets_backing_app_and_snapshot(monkeypatch) - assert session.flush.call_count == 2 -def test_resolve_package_soul_preserves_existing_and_marks_missing_knowledge(monkeypatch) -> None: +def test_resolve_package_soul_preserves_existing_and_marks_missing_knowledge(monkeypatch: pytest.MonkeyPatch) -> None: soul = AgentSoulConfig.model_validate( { "config_skills": [{"name": "skill", "file_kind": "tool_file", "file_id": "skill-file"}], @@ -692,6 +698,7 @@ def test_create_snapshot_increments_version_and_records_revision() -> None: ) assert snapshot.version == 3 + assert snapshot.home_snapshot_id is None assert isinstance(session.add.call_args_list[0].args[0], AgentConfigSnapshot) revision = session.add.call_args_list[1].args[0] assert isinstance(revision, AgentConfigRevision) diff --git a/api/tests/unit_tests/services/agent/test_agent_observability_service.py b/api/tests/unit_tests/services/agent/test_agent_observability_service.py index 4fcb0466a09..4ad06f887a2 100644 --- a/api/tests/unit_tests/services/agent/test_agent_observability_service.py +++ b/api/tests/unit_tests/services/agent/test_agent_observability_service.py @@ -1,12 +1,19 @@ from datetime import UTC, datetime from decimal import Decimal from types import SimpleNamespace +from typing import Protocol import pytest from core.app.entities.app_invoke_entities import InvokeFrom from graphon.enums import WorkflowNodeExecutionStatus -from models.enums import ConversationFromSource, CreatorUserRole, MessageStatus +from models.enums import ( + ConversationFromSource, + CreatorUserRole, + FeedbackFromSource, + FeedbackRating, + MessageStatus, +) from services.agent import observability_service as observability_service_module from services.agent.observability_service import AgentLogQueryParams, AgentObservabilityService @@ -278,7 +285,8 @@ def test_list_log_messages_merges_deduplicates_and_sorts_sources(monkeypatch: py {"id": "workflow-only", "created_at": 20, "updated_at": 10}, ] monkeypatch.setattr(service, "_list_webapp_messages", lambda **kwargs: [webapp_message]) - monkeypatch.setattr(service, "serialize_log_message", lambda message: webapp_row) + monkeypatch.setattr(service, "_list_message_feedbacks", lambda **kwargs: {}) + monkeypatch.setattr(service, "serialize_log_message", lambda message, feedbacks=(): webapp_row) monkeypatch.setattr(service, "_list_workflow_messages", lambda **kwargs: workflow_rows) payload = service.list_log_messages( @@ -301,6 +309,60 @@ def test_list_log_messages_merges_deduplicates_and_sorts_sources(monkeypatch: py } +def test_list_webapp_conversation_logs_includes_feedback_rates(monkeypatch: pytest.MonkeyPatch) -> None: + timestamp = datetime(2026, 7, 23, 7, 0, 19, tzinfo=UTC) + conversation = SimpleNamespace( + id="conversation-1", + name="Feedback conversation", + from_end_user_id="end-user-1", + read_at=None, + ) + + class FakeRow: + message_count = 2 + paused_count = 0 + failed_count = 0 + created_at = timestamp + updated_at = timestamp + + def __getitem__(self, index: int) -> SimpleNamespace: + if index != 0: + raise IndexError(index) + return conversation + + class FakeResult: + def all(self) -> list[FakeRow]: + return [FakeRow()] + + class FakeSession: + def execute(self, stmt: object) -> FakeResult: + str(stmt) + return FakeResult() + + app = SimpleNamespace( + id="app-1", + name="Agent WebApp", + icon_type=None, + icon=None, + icon_background=None, + ) + service = AgentObservabilityService(FakeSession()) + monkeypatch.setattr( + service, + "_list_conversation_feedback_rates", + lambda **kwargs: {"conversation-1": {"user_rate": 0.5, "operation_rate": 1.0}}, + ) + + rows = service._list_webapp_conversation_logs( + app=app, # type: ignore[arg-type] + params=AgentLogQueryParams(), + source_filter=AgentObservabilityService.resolve_source_filter("webapp"), + ) + + assert rows[0]["user_rate"] == 0.5 + assert rows[0]["operation_rate"] == 1.0 + + def test_list_workflow_logs_uses_node_executions_without_messages() -> None: created_at = datetime(2026, 7, 23, 7, 0, 19, tzinfo=UTC) node_execution = SimpleNamespace( @@ -551,8 +613,24 @@ def test_serialize_log_message_returns_frontend_log_shape() -> None: updated_at=updated_at, ) conversation = SimpleNamespace(name="Debug conversation") + feedbacks = [ + SimpleNamespace( + rating=FeedbackRating.LIKE, + content="Useful", + from_source=FeedbackFromSource.USER, + ), + SimpleNamespace( + rating=FeedbackRating.DISLIKE, + content="Needs more detail", + from_source=FeedbackFromSource.ADMIN, + ), + ] - payload = AgentObservabilityService.serialize_log_message(message, conversation) # type: ignore[arg-type] + payload = AgentObservabilityService.serialize_log_message( # type: ignore[arg-type] + message, + conversation, + feedbacks, + ) assert payload == { "id": "message-1", @@ -567,6 +645,11 @@ def test_serialize_log_message_returns_frontend_log_shape() -> None: "from_source": "console", "from_end_user_id": None, "from_account_id": "account-1", + "feedback_enabled": True, + "feedbacks": [ + {"rating": "like", "content": "Useful", "from_source": "user"}, + {"rating": "dislike", "content": "Needs more detail", "from_source": "admin"}, + ], "message_tokens": 3, "answer_tokens": 4, "total_tokens": 7, @@ -615,6 +698,8 @@ def test_serialize_workflow_node_message_returns_frontend_log_shape() -> None: "error": None, "from_end_user_id": "end-user-1", "from_account_id": None, + "feedback_enabled": False, + "feedbacks": [], "message_tokens": 10, "answer_tokens": 5, "total_tokens": 15, @@ -676,6 +761,97 @@ def test_serialize_workflow_node_message_handles_sparse_runtime_data() -> None: assert payload["updated_at"] == int(created_at.timestamp()) +def test_positive_feedback_rate_uses_rated_messages_as_denominator() -> None: + assert AgentObservabilityService._positive_feedback_rate(like_count=2, total_count=4) == 0.5 + assert AgentObservabilityService._positive_feedback_rate(like_count=0, total_count=1) == 0 + assert AgentObservabilityService._positive_feedback_rate(like_count=None, total_count=0) is None + + +def test_list_message_feedbacks_groups_feedbacks_by_message() -> None: + feedbacks = [ + SimpleNamespace(message_id="message-1"), + SimpleNamespace(message_id="message-1"), + SimpleNamespace(message_id="message-2"), + ] + + class Compilable(Protocol): + def compile(self) -> object: ... + + class FakeScalarResult: + def all(self) -> list[SimpleNamespace]: + return feedbacks + + class FakeSession: + def __init__(self) -> None: + self.scalar_calls = 0 + + def scalars(self, stmt: Compilable) -> FakeScalarResult: + stmt.compile() + self.scalar_calls += 1 + return FakeScalarResult() + + session = FakeSession() + service = AgentObservabilityService(session) # type: ignore[arg-type] + + grouped_feedbacks = service._list_message_feedbacks( + app=SimpleNamespace(id="app-1"), # type: ignore[arg-type] + messages=[SimpleNamespace(id="message-1"), SimpleNamespace(id="message-2")], # type: ignore[list-item] + ) + + assert grouped_feedbacks == { + "message-1": feedbacks[:2], + "message-2": feedbacks[2:], + } + assert service._list_message_feedbacks(app=SimpleNamespace(id="app-1"), messages=[]) == {} # type: ignore[arg-type] + assert session.scalar_calls == 1 + + +def test_list_conversation_feedback_rates_maps_user_and_admin_sources() -> None: + class FakeResult: + def all(self): + return [ + SimpleNamespace( + conversation_id="conversation-1", + from_source=FeedbackFromSource.USER, + like_count=2, + total_count=4, + ), + SimpleNamespace( + conversation_id="conversation-1", + from_source=FeedbackFromSource.ADMIN, + like_count=1, + total_count=1, + ), + SimpleNamespace( + conversation_id="conversation-without-ratings", + from_source=FeedbackFromSource.USER, + like_count=0, + total_count=0, + ), + ] + + class FakeSession: + def execute(self, stmt): + stmt.compile() + return FakeResult() + + service = AgentObservabilityService(FakeSession()) + + rates = service._list_conversation_feedback_rates( + app=SimpleNamespace(id="app-1"), # type: ignore[arg-type] + conversation_ids=["conversation-1"], + ) + + assert rates == {"conversation-1": {"user_rate": 0.5, "operation_rate": 1.0}} + assert ( + service._list_conversation_feedback_rates( + app=SimpleNamespace(id="app-1"), # type: ignore[arg-type] + conversation_ids=[], + ) + == {} + ) + + def test_workflow_node_serialization_helpers_handle_invalid_values() -> None: assert AgentObservabilityService._json_mapping(None) == {} assert AgentObservabilityService._json_mapping("not-json") == {} diff --git a/api/tests/unit_tests/services/agent/test_agent_services.py b/api/tests/unit_tests/services/agent/test_agent_services.py index c1946968e67..479324c3167 100644 --- a/api/tests/unit_tests/services/agent/test_agent_services.py +++ b/api/tests/unit_tests/services/agent/test_agent_services.py @@ -1,26 +1,34 @@ import json +from contextlib import nullcontext from datetime import UTC, datetime from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import MagicMock, call import pytest -from agenton.compositor import CompositorSessionSnapshot -from dify_agent.protocol import RuntimeLayerSpec +from sqlalchemy import event, select from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session from core.workflow.nodes.agent_v2.validators import WorkflowAgentNodeValidationError +from models.account import Account from models.agent import ( Agent, AgentConfigDraft, AgentConfigDraftType, + AgentConfigRevision, AgentConfigRevisionOperation, AgentConfigSnapshot, + AgentConfigVersionKind, AgentDebugConversation, AgentDriveFile, + AgentDriveFileKind, + AgentHomeSnapshot, AgentKind, AgentScope, AgentSource, AgentStatus, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, WorkflowAgentBindingType, WorkflowAgentNodeBinding, ) @@ -32,13 +40,14 @@ from models.agent_config_entities import ( WorkflowNodeJobConfig, ) from models.enums import AppStatus, ConversationFromSource, ConversationStatus -from models.model import App, AppMode, Conversation, IconType -from models.workflow import Workflow +from models.model import App, AppMode, Conversation, IconType, Message +from models.workflow import Workflow, WorkflowType from services.agent import composer_service, roster_service from services.agent.agent_soul_state import agent_soul_has_model from services.agent.composer_service import AgentComposerService from services.agent.composer_validator import ComposerConfigValidator from services.agent.errors import ( + AgentBuildSandboxNotFoundError, AgentModelNotConfiguredError, AgentNameConflictError, AgentNotFoundError, @@ -46,59 +55,14 @@ from services.agent.errors import ( AgentVersionNotFoundError, InvalidComposerConfigError, ) +from services.agent.home_snapshot_service import AgentHomeSnapshotService from services.agent.roster_service import AgentRosterService from services.agent.workflow_publish_service import WorkflowAgentPublishService +from services.agent.workspace_service import AgentWorkspaceService from services.app_service import AppListParams, AppService from services.entities.agent_entities import AgentSoulConfig, ComposerSavePayload, ComposerSaveStrategy, ComposerVariant -class FakeScalarResult: - def __init__(self, values): - self.values = values - - def all(self): - return self.values - - -class FakeSession: - def __init__(self, *, scalars=None, scalar=None): - self._scalars = list(scalars or []) - self._scalar = list(scalar or []) - self.added = [] - self.deleted = [] - self.commits = 0 - self.flushes = 0 - self.rollbacks = 0 - - def scalar(self, _stmt): - if self._scalar: - return self._scalar.pop(0) - return None - - def scalars(self, _stmt): - if self._scalars: - return FakeScalarResult(self._scalars.pop(0)) - return FakeScalarResult([]) - - def add(self, value): - self.added.append(value) - - def delete(self, value): - self.deleted.append(value) - - def flush(self): - self.flushes += 1 - for index, value in enumerate(self.added, start=1): - if getattr(value, "id", None) is None: - value.id = f"generated-{index}" - - def commit(self): - self.commits += 1 - - def rollback(self): - self.rollbacks += 1 - - def _agent_soul_with_model() -> AgentSoulConfig: return AgentSoulConfig.model_validate( { @@ -111,64 +75,178 @@ def _agent_soul_with_model() -> AgentSoulConfig: ) +def _agent( + *, + agent_id: str = "agent-1", + tenant_id: str = "tenant-1", + name: str = "Researcher", + scope: AgentScope = AgentScope.ROSTER, + source: AgentSource = AgentSource.ROSTER, + app_id: str | None = None, +) -> Agent: + return Agent( + id=agent_id, + tenant_id=tenant_id, + name=name, + description="desc", + role="assistant", + agent_kind=AgentKind.DIFY_AGENT, + scope=scope, + source=source, + app_id=app_id, + status=AgentStatus.ACTIVE, + created_by="account-1", + updated_by="account-1", + ) + + +def _snapshot( + *, + snapshot_id: str = "snapshot-1", + tenant_id: str = "tenant-1", + agent_id: str = "agent-1", + version: int = 1, + agent_soul: AgentSoulConfig | None = None, +) -> AgentConfigSnapshot: + return AgentConfigSnapshot( + id=snapshot_id, + tenant_id=tenant_id, + agent_id=agent_id, + version=version, + config_snapshot=agent_soul or _agent_soul_with_model(), + created_by="account-1", + ) + + +def _conversation(*, conversation_id: str = "conversation-1", account_id: str = "account-1") -> Conversation: + return Conversation( + id=conversation_id, + app_id="app-1", + override_model_configs="{}", + mode=AppMode.AGENT_CHAT, + name="Debug", + summary="", + _inputs={}, + introduction="", + system_instruction="", + status=ConversationStatus.NORMAL, + from_source=ConversationFromSource.CONSOLE, + from_account_id=account_id, + dialogue_count=0, + ) + + +def _workflow(*, workflow_id: str = "workflow-1", tenant_id: str = "tenant-1", app_id: str = "app-1") -> Workflow: + return Workflow( + id=workflow_id, + tenant_id=tenant_id, + app_id=app_id, + type=WorkflowType.WORKFLOW, + version=Workflow.VERSION_DRAFT, + graph='{"nodes": [], "edges": []}', + _features="{}", + created_by="account-1", + _environment_variables="{}", + _conversation_variables="{}", + _rag_pipeline_variables="{}", + ) + + +def _app( + *, + app_id: str = "app-1", + tenant_id: str = "tenant-1", + name: str = "Agent App", + mode: AppMode = AppMode.AGENT_CHAT, +) -> App: + return App( + id=app_id, + tenant_id=tenant_id, + name=name, + description="", + mode=mode, + icon_type=IconType.EMOJI, + icon="🤖", + icon_background="#fff", + status=AppStatus.NORMAL, + enable_site=False, + enable_api=True, + max_active_requests=None, + created_by="account-1", + ) + + def test_agent_soul_has_model(): assert agent_soul_has_model(_agent_soul_with_model()) is True assert agent_soul_has_model(AgentSoulConfig()) is False -def test_get_published_agent_soul_for_app_uses_active_snapshot(): +def test_get_published_agent_soul_for_app_uses_active_snapshot(sqlite_session: Session): agent_soul = AgentSoulConfig.model_validate({"app_features": {"speech_to_text": {"enabled": True}}}) - agent = SimpleNamespace(id="agent-1", active_config_snapshot_id="version-1") - version = SimpleNamespace(config_snapshot_dict=agent_soul.model_dump(mode="json")) - service = AgentRosterService(FakeSession(scalar=[agent, version])) + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") + version = _snapshot(snapshot_id="version-1", agent_soul=agent_soul) + agent.active_config_snapshot_id = version.id + session.add_all([agent, version]) + session.commit() - result = service.get_published_agent_soul_for_app(tenant_id="tenant-1", app_id="app-1") + result = AgentRosterService(session).get_published_agent_soul_for_app(tenant_id="tenant-1", app_id="app-1") assert result == agent_soul -def test_get_published_agent_soul_for_app_returns_none_without_backing_agent(): - service = AgentRosterService(FakeSession(scalar=[None])) +def test_get_published_agent_soul_for_app_returns_none_without_backing_agent(sqlite_session: Session): + service = AgentRosterService(sqlite_session) result = service.get_published_agent_soul_for_app(tenant_id="tenant-1", app_id="legacy-app-1") assert result is None -def test_peek_authz_app_id_uses_the_parent_app_not_the_hidden_backing_app(): +def test_peek_authz_app_id_uses_the_parent_app_not_the_hidden_backing_app(sqlite_session: Session): """A workflow-only Agent is authorized against its parent workflow App.""" - agent = SimpleNamespace(id="agent-1", backing_app_id="backing-app-1", app_id="parent-app-1") - service = AgentRosterService(FakeSession(scalar=[agent])) + session = sqlite_session + agent = _agent( + scope=AgentScope.WORKFLOW_ONLY, + source=AgentSource.WORKFLOW, + app_id="parent-app-1", + ) + agent.backing_app_id = "backing-app-1" + agent.workflow_id = "workflow-1" + agent.workflow_node_id = "node-1" + session.add(agent) + session.commit() - result = service.peek_authz_app_id(tenant_id="tenant-1", agent_id="agent-1") + result = AgentRosterService(session).peek_authz_app_id(tenant_id="tenant-1", agent_id="agent-1") assert result == "parent-app-1" -def test_peek_authz_app_id_uses_the_roster_agent_app(): - agent = SimpleNamespace(id="agent-1", backing_app_id=None, app_id="roster-app-1") - service = AgentRosterService(FakeSession(scalar=[agent])) +def test_peek_authz_app_id_uses_the_roster_agent_app(sqlite_session: Session): + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="roster-app-1") + session.add(agent) + session.commit() - result = service.peek_authz_app_id(tenant_id="tenant-1", agent_id="agent-1") + result = AgentRosterService(session).peek_authz_app_id(tenant_id="tenant-1", agent_id="agent-1") assert result == "roster-app-1" -def test_peek_authz_app_id_returns_none_without_creating_a_backing_app(): +def test_peek_authz_app_id_returns_none_without_creating_a_backing_app(sqlite_session: Session): """Authorization checks must not materialize the hidden backing App.""" - session = FakeSession(scalar=[None]) + session = sqlite_session service = AgentRosterService(session) result = service.peek_authz_app_id(tenant_id="tenant-1", agent_id="agent-1") assert result is None - assert session.added == [] - assert session.commits == 0 - assert session.flushes == 0 + assert not session.new + assert not session.dirty -def test_load_workflow_composer_returns_empty_state(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_load_workflow_composer_returns_empty_state(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: None) @@ -188,8 +266,8 @@ def test_load_workflow_composer_returns_empty_state(monkeypatch: pytest.MonkeyPa assert files_output["array_item"] == {"type": "file", "description": None, "children": []} -def test_load_workflow_composer_serializes_existing_binding(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_load_workflow_composer_serializes_existing_binding(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session binding = SimpleNamespace( agent_id="agent-1", binding_type=WorkflowAgentBindingType.ROSTER_AGENT, @@ -220,8 +298,8 @@ def test_load_workflow_composer_serializes_existing_binding(monkeypatch: pytest. assert result == {"agent": "agent-1", "version": "version-1"} -def test_load_workflow_composer_uses_roster_preview_snapshot(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_load_workflow_composer_uses_roster_preview_snapshot(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session binding = SimpleNamespace( agent_id="agent-1", binding_type=WorkflowAgentBindingType.ROSTER_AGENT, @@ -257,8 +335,8 @@ def test_load_workflow_composer_uses_roster_preview_snapshot(monkeypatch: pytest assert result == {"binding_snapshot_id": "binding-version", "version": "preview-version"} -def test_load_workflow_composer_uses_inline_preview_snapshot(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_load_workflow_composer_uses_inline_preview_snapshot(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session binding = SimpleNamespace( agent_id="inline-agent-1", binding_type=WorkflowAgentBindingType.INLINE_AGENT, @@ -301,15 +379,15 @@ def test_load_workflow_composer_uses_inline_preview_snapshot(monkeypatch: pytest assert result == {"agent": "inline-agent-1", "version": "inline-preview-version"} -def test_workflow_inline_debug_conversation_seed(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_workflow_inline_debug_conversation_seed(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session captured: dict[str, object] = {} class FakeRosterService: def __init__(self, session): captured["session"] = session - def get_or_create_agent_app_debug_conversation_id(self, **kwargs): + def get_or_create_build_conversation(self, **kwargs): captured.update(kwargs) return "debug-conversation-1" @@ -333,8 +411,10 @@ def test_workflow_inline_debug_conversation_seed(monkeypatch: pytest.MonkeyPatch assert captured["commit"] is False -def test_workflow_inline_debug_conversation_seed_skips_non_inline(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_workflow_inline_debug_conversation_seed_skips_non_inline( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session class UnexpectedRosterService: def __init__(self, session): @@ -364,8 +444,10 @@ def test_workflow_inline_debug_conversation_seed_skips_non_inline(monkeypatch: p ) -def test_load_workflow_composer_rejects_preview_without_binding(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_load_workflow_composer_rejects_preview_without_binding( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: None) @@ -389,8 +471,8 @@ def test_load_workflow_composer_rejects_preview_without_binding(monkeypatch: pyt (ComposerSaveStrategy.SAVE_TO_ROSTER, "_save_to_roster"), ], ) -def test_save_workflow_composer_dispatches_save_strategy(monkeypatch, strategy, helper_name): - fake_session = FakeSession() +def test_save_workflow_composer_dispatches_save_strategy(monkeypatch, strategy, helper_name, sqlite_session: Session): + session = sqlite_session binding = SimpleNamespace( agent_id="agent-1", binding_type=WorkflowAgentBindingType.ROSTER_AGENT, @@ -399,7 +481,6 @@ def test_save_workflow_composer_dispatches_save_strategy(monkeypatch, strategy, calls = [] serialize_calls = [] - session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_draft_save_payload", lambda payload: None) monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: None) @@ -446,11 +527,75 @@ def test_save_workflow_composer_dispatches_save_strategy(monkeypatch, strategy, assert result == {"state": "ok"} assert calls assert serialize_calls[0]["account_id"] == "account-1" - assert fake_session.flushes >= 1 -def test_save_workflow_composer_rejects_agent_app_variant(): - session = FakeSession() +def test_save_workflow_composer_commits_before_retiring_replaced_inline_agent( + monkeypatch, sqlite_session: Session +) -> None: + session = sqlite_session + events: list[str] = [] + old_binding = SimpleNamespace( + agent_id="old-inline-agent", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + ) + new_binding = SimpleNamespace( + agent_id="new-agent", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + ) + monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **_kwargs: SimpleNamespace(id="workflow-1")) + monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **_kwargs: old_binding) + monkeypatch.setattr(AgentComposerService, "_save_as_new_agent", lambda **_kwargs: new_binding) + monkeypatch.setattr( + AgentComposerService, + "_get_agent_if_present", + lambda **_kwargs: SimpleNamespace(id="new-agent", active_config_snapshot_id="version-1"), + ) + monkeypatch.setattr( + AgentComposerService, + "_get_version_if_present", + lambda **_kwargs: SimpleNamespace(id="version-1"), + ) + monkeypatch.setattr(AgentComposerService, "_serialize_workflow_state", lambda **_kwargs: {"state": "ok"}) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **_kwargs: None) + monkeypatch.setattr(AgentComposerService, "collect_validation_findings", lambda **_kwargs: {}) + event.listen(session, "after_commit", lambda _session: events.append("commit")) + + def retire_unowned(**kwargs): + assert kwargs["agent_ids"] == {"old-inline-agent"} + events.append("retire") + return ["binding-1"], ["home-1"] + + monkeypatch.setattr(composer_service.WorkflowAgentRetirementService, "retire_unowned", retire_unowned) + monkeypatch.setattr( + composer_service, + "enqueue_agent_resource_collection", + MagicMock(side_effect=lambda **_kwargs: events.append("enqueue")), + ) + payload = ComposerSavePayload.model_validate( + { + "variant": ComposerVariant.WORKFLOW, + "save_strategy": ComposerSaveStrategy.SAVE_AS_NEW_AGENT, + "agent_soul": _agent_soul_with_model().model_dump(mode="json"), + "new_agent_name": "New Agent", + "soul_lock": {"locked": False}, + } + ) + + AgentComposerService.save_workflow_composer( + session=session, + tenant_id="tenant-1", + app_id="app-1", + node_id="node-1", + account_id="account-1", + payload=payload, + ) + + assert events == ["commit", "retire", "enqueue"] + + +def test_save_workflow_composer_rejects_agent_app_variant(sqlite_session: Session): + session = sqlite_session payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.AGENT_APP.value, @@ -510,11 +655,14 @@ def test_publish_save_strategies_run_publish_validation(strategy: ComposerSaveSt composer_service._validate_composer_payload_for_strategy(_duplicate_env_secret_payload(strategy)) -def test_save_agent_app_composer_creates_agent_when_missing(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession(scalar=[None]) - saved_draft = SimpleNamespace(id="draft-1", config_snapshot_dict={"prompt": {"system_prompt": "x"}}) +def test_save_agent_app_composer_creates_agent_when_missing(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session + saved_draft = SimpleNamespace( + id="draft-1", + home_snapshot_id="home-initial", + config_snapshot_dict={"prompt": {"system_prompt": "x"}}, + ) - session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_draft_save_payload", lambda payload: None) monkeypatch.setattr(AgentComposerService, "_save_agent_draft", lambda **kwargs: saved_draft) monkeypatch.setattr(AgentComposerService, "load_agent_composer", lambda **kwargs: {"loaded": True}) @@ -537,17 +685,22 @@ def test_save_agent_app_composer_creates_agent_when_missing(monkeypatch: pytest. assert result.pop("validation") == {"warnings": [], "knowledge_retrieval_placeholder": []} assert result == {"loaded": True} - assert fake_session.added[0].name == "Analyst" - assert fake_session.added[0].active_config_snapshot_id is None - assert fake_session.added[0].active_config_is_published is False - assert fake_session.flushes >= 1 + agent = session.scalar(select(Agent).where(Agent.app_id == "app-1")) + assert agent is not None + snapshot = session.scalar(select(AgentConfigSnapshot).where(AgentConfigSnapshot.agent_id == agent.id)) + assert snapshot is not None + assert agent.name == "Analyst" + assert agent.active_config_snapshot_id == snapshot.id + assert snapshot.home_snapshot_id is None + assert agent.active_config_is_published is False -def test_load_agent_app_composer_exposes_draft_save_only(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_load_agent_app_composer_exposes_draft_save_only(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session agent = SimpleNamespace( id="agent-1", active_config_snapshot_id="version-1", + active_config_is_published=True, updated_by="account-1", created_by="account-1", app_id="app-1", @@ -567,10 +720,11 @@ def test_load_agent_app_composer_exposes_draft_save_only(monkeypatch: pytest.Mon result = AgentComposerService.load_agent_app_composer(session=session, tenant_id="tenant-1", app_id="app-1") assert result["save_options"] == [ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION.value] + assert result["active_config_is_published"] is True -def test_save_agent_app_composer_rejects_version_save_strategy(): - session = FakeSession() +def test_save_agent_app_composer_rejects_version_save_strategy(sqlite_session: Session): + session = sqlite_session payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.AGENT_APP.value, @@ -589,28 +743,31 @@ def test_save_agent_app_composer_rejects_version_save_strategy(): ) -def test_save_agent_app_composer_updates_normal_draft(monkeypatch: pytest.MonkeyPatch): - agent = SimpleNamespace( - id="agent-1", - tenant_id="tenant-1", - source=AgentSource.AGENT_APP, - active_config_snapshot_id="version-1", - active_config_is_published=True, - updated_by=None, +def test_save_agent_app_composer_updates_normal_draft(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") + agent.active_config_snapshot_id = "version-1" + agent.active_config_is_published = True + agent.updated_by = None + session.add(agent) + session.commit() + active_version = SimpleNamespace( + home_snapshot_id="home-initial", config_snapshot_dict=AgentSoulConfig().model_dump(mode="json") ) - active_version = SimpleNamespace(config_snapshot_dict=AgentSoulConfig().model_dump(mode="json")) - fake_session = FakeSession(scalar=[agent]) saved = {} - session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_draft_save_payload", lambda payload: None) monkeypatch.setattr( AgentComposerService, "_save_agent_draft", - lambda **kwargs: saved.update(kwargs) or SimpleNamespace(id="draft-1"), + lambda **kwargs: saved.update(kwargs) or SimpleNamespace(id="draft-1", home_snapshot_id="home-initial"), ) monkeypatch.setattr(AgentComposerService, "_get_version_if_present", lambda **_kwargs: active_version) - monkeypatch.setattr(AgentComposerService, "load_agent_composer", lambda **kwargs: {"loaded": True}) + monkeypatch.setattr( + AgentComposerService, + "load_agent_composer", + lambda **kwargs: {"loaded": True, "active_config_is_published": agent.active_config_is_published}, + ) payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.AGENT_APP.value, @@ -628,36 +785,46 @@ def test_save_agent_app_composer_updates_normal_draft(monkeypatch: pytest.Monkey ) assert result.pop("validation") == {"warnings": [], "knowledge_retrieval_placeholder": []} - assert result == {"loaded": True} + assert result == {"loaded": True, "active_config_is_published": False} assert saved["draft_type"] == AgentConfigDraftType.DRAFT assert saved["agent_soul"].model_dump(mode="json") == _agent_soul_with_model().model_dump(mode="json") assert agent.active_config_is_published is False - assert fake_session._scalar == [] - assert fake_session.flushes >= 1 -def test_save_agent_app_composer_keeps_published_when_draft_matches_active_snapshot(monkeypatch: pytest.MonkeyPatch): +def test_save_agent_app_composer_keeps_published_when_draft_matches_active_snapshot( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): agent_soul = _agent_soul_with_model() - agent = SimpleNamespace( - id="agent-1", + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") + agent.active_config_snapshot_id = "version-1" + agent.active_config_is_published = False + agent.updated_by = None + revision = AgentConfigRevision( tenant_id="tenant-1", - source=AgentSource.AGENT_APP, - active_config_snapshot_id="version-1", - active_config_is_published=False, - updated_by=None, + agent_id=agent.id, + current_snapshot_id="version-1", + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + ) + session.add_all([agent, revision]) + session.commit() + active_version = SimpleNamespace( + home_snapshot_id="home-initial", config_snapshot_dict=agent_soul.model_dump(mode="json") ) - active_version = SimpleNamespace(config_snapshot_dict=agent_soul.model_dump(mode="json")) - fake_session = FakeSession(scalar=[agent, "publish-revision-1"]) - session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_draft_save_payload", lambda payload: None) monkeypatch.setattr( AgentComposerService, "_save_agent_draft", - lambda **_kwargs: SimpleNamespace(id="draft-1"), + lambda **_kwargs: SimpleNamespace(id="draft-1", home_snapshot_id="home-initial"), ) monkeypatch.setattr(AgentComposerService, "_get_version_if_present", lambda **_kwargs: active_version) - monkeypatch.setattr(AgentComposerService, "load_agent_composer", lambda **_kwargs: {"loaded": True}) + monkeypatch.setattr( + AgentComposerService, + "load_agent_composer", + lambda **_kwargs: {"loaded": True, "active_config_is_published": agent.active_config_is_published}, + ) payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.AGENT_APP.value, @@ -666,7 +833,7 @@ def test_save_agent_app_composer_keeps_published_when_draft_matches_active_snaps } ) - AgentComposerService.save_agent_app_composer( + result = AgentComposerService.save_agent_app_composer( session=session, tenant_id="tenant-1", app_id="app-1", @@ -675,10 +842,11 @@ def test_save_agent_app_composer_keeps_published_when_draft_matches_active_snaps ) assert agent.active_config_is_published is True - assert fake_session.flushes >= 1 + assert result["active_config_is_published"] is True -def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.MonkeyPatch): +def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -688,6 +856,7 @@ def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.Monke scope=AgentScope.ROSTER, source=AgentSource.AGENT_APP, status=AgentStatus.ACTIVE, + app_id="app-1", active_config_snapshot_id="version-1", active_config_is_published=False, ) @@ -699,7 +868,8 @@ def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.Monke base_snapshot_id="version-1", config_snapshot=AgentSoulConfig(), ) - fake_session = FakeSession(scalar=[agent, draft, None]) + session.add_all([agent, draft]) + session.commit() def fail_create_config_version(**_kwargs): raise AssertionError("config version must not be created when Agent Soul has no model") @@ -708,6 +878,7 @@ def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.Monke raise AssertionError("knowledge datasets must not be validated when Agent Soul has no model") monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda payload: None) + monkeypatch.setattr(composer_service, "agent_has_workflow_callable_active_snapshot", lambda **_kwargs: False) monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", fail_validate_knowledge_datasets) monkeypatch.setattr(AgentComposerService, "_create_config_version", fail_create_config_version) @@ -717,18 +888,19 @@ def test_publish_agent_app_draft_rejects_missing_model(monkeypatch: pytest.Monke agent_id="agent-1", account_id="account-1", version_note="ship it", - session=fake_session, + session=session, ) assert exc_info.value.error_code == "agent_model_not_configured" assert agent.active_config_snapshot_id == "version-1" assert agent.active_config_is_published is False assert draft.base_snapshot_id == "version-1" - assert fake_session.flushes == 0 - assert fake_session.commits == 0 + assert not session.new + assert not session.dirty -def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest.MonkeyPatch): +def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -738,6 +910,7 @@ def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest. scope=AgentScope.ROSTER, source=AgentSource.AGENT_APP, status=AgentStatus.ACTIVE, + app_id="app-1", active_config_snapshot_id="version-1", ) draft = AgentConfigDraft( @@ -746,19 +919,30 @@ def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest. draft_type=AgentConfigDraftType.DRAFT, draft_owner_key="", base_snapshot_id="version-1", + home_snapshot_id=None, config_snapshot=_agent_soul_with_model(), ) version = SimpleNamespace(id="version-2") - fake_session = FakeSession(scalar=[agent, draft]) + app = _app(mode=AppMode.AGENT) + app.enable_site = False + app.enable_api = False + session.add_all([agent, draft, app]) + session.commit() created: dict[str, object] = {} + calls: list[str] = [] - session = fake_session monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda payload: None) + monkeypatch.setattr(composer_service, "agent_has_workflow_callable_active_snapshot", lambda **_kwargs: False) monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **kwargs: None) + monkeypatch.setattr( + composer_service, + "validate_home_snapshot_binding", + lambda **kwargs: calls.append("validate_home"), + ) monkeypatch.setattr( AgentComposerService, "_create_config_version", - lambda **kwargs: created.update(kwargs) or version, + lambda **kwargs: calls.append("create_version") or created.update(kwargs) or version, ) monkeypatch.setattr(AgentComposerService, "_serialize_version", lambda _version: {"id": _version.id}) @@ -775,13 +959,96 @@ def test_publish_agent_app_draft_creates_published_snapshot(monkeypatch: pytest. assert result["draft"]["base_snapshot_id"] == "version-2" assert created["operation"] == AgentConfigRevisionOperation.PUBLISH_DRAFT assert created["previous_snapshot_id"] == "version-1" + assert created["home_snapshot_id"] is None + assert calls == ["validate_home", "create_version"] assert agent.active_config_snapshot_id == "version-2" assert agent.active_config_has_model is True assert agent.active_config_is_published is True - assert fake_session.flushes >= 1 + assert app.enable_site is True + assert app.enable_api is True + assert app.updated_by == "account-1" -def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkeypatch: pytest.MonkeyPatch): +def test_repeated_publish_reuses_normal_draft_home_without_creating_resources( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + app_id="app-1", + active_config_snapshot_id="version-1", + ) + draft = AgentConfigDraft( + id="draft-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id="home-1", + config_snapshot=_agent_soul_with_model(), + ) + app = _app(mode=AppMode.AGENT) + app.enable_site = False + app.enable_api = False + session.add_all([agent, draft, app]) + session.commit() + published_homes: list[str] = [] + versions = iter([SimpleNamespace(id="version-2"), SimpleNamespace(id="version-3")]) + create_from_build = MagicMock() + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda _payload: None) + publish_visibility = iter([False, True]) + monkeypatch.setattr( + composer_service, + "agent_has_workflow_callable_active_snapshot", + lambda **_kwargs: next(publish_visibility), + ) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **_kwargs: None) + monkeypatch.setattr(composer_service, "validate_home_snapshot_binding", lambda **_kwargs: None) + monkeypatch.setattr( + AgentComposerService, + "_create_config_version", + lambda **kwargs: published_homes.append(kwargs["home_snapshot_id"]) or next(versions), + ) + monkeypatch.setattr(AgentComposerService, "_serialize_version", lambda version: {"id": version.id}) + monkeypatch.setattr(AgentComposerService, "_serialize_draft", lambda value: {"id": value.id}) + monkeypatch.setattr(AgentHomeSnapshotService, "create_for_build_apply", create_from_build) + + first = AgentComposerService.publish_agent_app_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + app.enable_site = False + app.enable_api = False + second = AgentComposerService.publish_agent_app_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + assert first["active_config_snapshot_id"] == "version-2" + assert second["active_config_snapshot_id"] == "version-3" + assert published_homes == ["home-1", "home-1"] + assert draft.home_snapshot_id == "home-1" + assert app.enable_site is False + assert app.enable_api is False + create_from_build.assert_not_called() + + +def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -801,11 +1068,26 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey account_id=None, draft_owner_key="", base_snapshot_id="version-1", + home_snapshot_id="home-initial", config_snapshot=_agent_soul_with_model(), ) - fake_session = FakeSession(scalar=[agent, normal_draft, None]) - - session = fake_session + active_version = AgentConfigSnapshot( + id="version-1", + tenant_id="tenant-1", + agent_id=agent.id, + version=1, + home_snapshot_id=normal_draft.home_snapshot_id, + config_snapshot=normal_draft.config_snapshot, + ) + publish_revision = AgentConfigRevision( + tenant_id="tenant-1", + agent_id=agent.id, + current_snapshot_id=active_version.id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + ) + session.add_all([agent, normal_draft, active_version, publish_revision]) + session.commit() checked_out = AgentComposerService.checkout_agent_app_build_draft( session=session, @@ -814,19 +1096,30 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey account_id="account-1", ) - build_draft = fake_session.added[0] + build_draft = session.scalar( + select(AgentConfigDraft).where( + AgentConfigDraft.agent_id == agent.id, + AgentConfigDraft.draft_type == AgentConfigDraftType.DEBUG_BUILD, + AgentConfigDraft.account_id == "account-1", + ) + ) + assert build_draft is not None assert checked_out["draft"]["id"] == build_draft.id assert checked_out["draft"]["draft_type"] == AgentConfigDraftType.DEBUG_BUILD.value assert checked_out["draft"]["account_id"] == "account-1" assert checked_out["draft"]["base_snapshot_id"] == "version-1" + assert build_draft.home_snapshot_id == "home-initial" assert checked_out["agent_soul"] == normal_draft.config_snapshot_dict - assert fake_session.flushes >= 1 - active_version = SimpleNamespace(config_snapshot_dict=build_draft.config_snapshot_dict) - fake_session = FakeSession( - scalar=[agent, build_draft, normal_draft, active_version, "publish-revision-1"], - ) - session = fake_session + source_binding_id = "binding-1" + build_draft.agent_workspace_binding_id = source_binding_id + session.commit() + create_home = MagicMock(return_value=SimpleNamespace(id="home-build", snapshot_ref="backend-home-build")) + monkeypatch.setattr(AgentHomeSnapshotService, "create_for_build_apply", create_home) + retire_binding = MagicMock() + enqueue_collection = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) applied = AgentComposerService.apply_agent_app_build_draft( session=session, @@ -838,9 +1131,909 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey assert applied["result"] == "success" assert applied["draft"]["id"] == normal_draft.id assert normal_draft.config_snapshot_dict == build_draft.config_snapshot_dict - assert agent.active_config_is_published is True - assert fake_session.deleted == [build_draft] - assert fake_session.flushes >= 1 + assert normal_draft.home_snapshot_id == "home-build" + assert agent.active_config_is_published is False + assert session.get(AgentConfigDraft, build_draft.id) is None + create_home.assert_called_once_with( + session=session, + build_draft=build_draft, + ) + retire_binding.assert_called_once_with(session=session, tenant_id="tenant-1", binding_id=source_binding_id) + enqueue_collection.assert_called_once_with(tenant_id="tenant-1", binding_ids=[source_binding_id]) + + +@pytest.mark.parametrize( + ("scope", "source", "app_id", "backing_app_id", "expected_runtime_app_id"), + [ + (AgentScope.ROSTER, AgentSource.AGENT_APP, "app-1", None, "app-1"), + ( + AgentScope.WORKFLOW_ONLY, + AgentSource.WORKFLOW, + "workflow-app-1", + "runtime-app-1", + "runtime-app-1", + ), + ], +) +def test_force_build_draft_checkout_collects_retired_binding_after_commit( + monkeypatch: pytest.MonkeyPatch, + scope: AgentScope, + source: AgentSource, + app_id: str, + backing_app_id: str | None, + expected_runtime_app_id: str, + sqlite_session: Session, +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=scope, + source=source, + status=AgentStatus.ACTIVE, + app_id=app_id, + backing_app_id=backing_app_id, + ) + normal_draft = AgentConfigDraft( + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id="home-1", + config_snapshot=_agent_soul_with_model(), + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-1", + agent_workspace_binding_id="binding-1", + config_snapshot=_agent_soul_with_model(), + ) + binding = SimpleNamespace( + agent_id=agent.id, + base_home_snapshot_id="home-1", + agent_config_version_id=build_draft.id, + agent_config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) + get_active_binding = MagicMock(return_value=binding) + retire_binding = MagicMock(return_value="binding-1") + enqueue_collection = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active_binding) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) + session.add_all([agent, normal_draft, build_draft]) + session.commit() + + AgentComposerService.checkout_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id=agent.id, + account_id="account-1", + ) + + assert build_draft.agent_workspace_binding_id == "binding-1" + retire_binding.assert_not_called() + enqueue_collection.assert_not_called() + + AgentComposerService.checkout_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id=agent.id, + account_id="account-1", + force=True, + ) + + assert build_draft.agent_workspace_binding_id is None + owner_scope = get_active_binding.call_args.kwargs["expected_owner_scope"] + assert owner_scope.app_id == expected_runtime_app_id + assert owner_scope.owner_type is AgentWorkspaceOwnerType.BUILD_DRAFT + assert owner_scope.owner_id == build_draft.id + retire_binding.assert_called_once_with( + session=session, + tenant_id="tenant-1", + binding_id="binding-1", + ) + enqueue_collection.assert_called_once_with( + tenant_id="tenant-1", + binding_ids=("binding-1",), + ) + + +def test_force_build_draft_checkout_rejects_unavailable_pointed_binding( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + app_id="app-1", + ) + normal_draft = AgentConfigDraft( + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id="home-1", + config_snapshot=_agent_soul_with_model(), + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-1", + agent_workspace_binding_id="binding-missing", + config_snapshot=_agent_soul_with_model(), + ) + session.add_all([agent, normal_draft, build_draft]) + session.commit() + retire_binding = MagicMock() + enqueue_collection = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=None)) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) + + with pytest.raises(AgentBuildSandboxNotFoundError): + AgentComposerService.checkout_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id=agent.id, + account_id="account-1", + force=True, + ) + + assert build_draft.agent_workspace_binding_id == "binding-missing" + assert session.get(AgentConfigDraft, build_draft.id).agent_workspace_binding_id == "binding-missing" + retire_binding.assert_not_called() + enqueue_collection.assert_not_called() + + +def test_build_apply_checkpoints_binding_updates_normal_draft_then_collects( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + app_id="app-1", + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id=None, + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + normal_draft = AgentConfigDraft( + id="draft-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id=None, + config_snapshot=AgentSoulConfig(), + ) + source_binding = SimpleNamespace( + id="binding-1", + backend_binding_ref="backend-binding-1", + agent_id="agent-1", + base_home_snapshot_id=None, + agent_config_version_id="build-1", + agent_config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) + session.add_all([agent, build_draft]) + session.commit() + client = MagicMock() + client.create_home_snapshot_from_binding_sync.return_value = SimpleNamespace(snapshot_ref="snapshot-ref-new") + lifecycle: list[str] = [] + event.listen(session, "after_commit", lambda _session: lifecycle.append("commit")) + retire = MagicMock(return_value=source_binding.id) + enqueue_collection = MagicMock(side_effect=lambda **_kwargs: lifecycle.append("enqueue")) + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda _payload: None) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **_kwargs: None) + monkeypatch.setattr(AgentHomeSnapshotService, "_client", lambda: nullcontext(client)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=source_binding)) + monkeypatch.setattr(AgentComposerService, "_save_agent_draft", lambda **_kwargs: normal_draft) + monkeypatch.setattr(AgentComposerService, "_agent_soul_matches_active_config", lambda **_kwargs: False) + monkeypatch.setattr(AgentComposerService, "_serialize_draft", lambda draft: {"id": draft.id}) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire) + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) + + result = AgentComposerService.apply_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + request = client.create_home_snapshot_from_binding_sync.call_args.args[0] + assert request.backend_binding_ref == "backend-binding-1" + created_home = session.scalar(select(AgentHomeSnapshot).where(AgentHomeSnapshot.snapshot_ref == "snapshot-ref-new")) + assert created_home is not None + assert created_home.snapshot_ref == "snapshot-ref-new" + assert normal_draft.home_snapshot_id == created_home.id + retire.assert_called_once_with(session=session, tenant_id="tenant-1", binding_id="binding-1") + enqueue_collection.assert_called_once_with(tenant_id="tenant-1", binding_ids=["binding-1"]) + assert lifecycle == ["commit", "enqueue"] + assert result == {"result": "success", "draft": {"id": "draft-1"}} + + +def test_build_apply_retires_normal_preview_binding_before_replacing_draft_home( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + app_id="app-1", + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-build", + agent_workspace_binding_id="binding-build", + config_snapshot=AgentSoulConfig(), + ) + normal_draft = AgentConfigDraft( + id="draft-1", + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id="home-preview-old", + config_snapshot=AgentSoulConfig(), + ) + preview_mapping = AgentDebugConversation( + tenant_id=agent.tenant_id, + agent_id=agent.id, + app_id="app-1", + account_id="account-2", + draft_type=AgentConfigDraftType.DRAFT, + conversation_id="conversation-preview", + ) + empty_preview_mapping = AgentDebugConversation( + tenant_id=agent.tenant_id, + agent_id=agent.id, + app_id="app-1", + account_id="account-3", + draft_type=AgentConfigDraftType.DRAFT, + conversation_id="conversation-preview-empty", + ) + preview_conversation = _conversation(conversation_id="conversation-preview", account_id="account-2") + preview_conversation.agent_workspace_binding_id = "binding-preview" + empty_preview_conversation = _conversation(conversation_id="conversation-preview-empty", account_id="account-3") + preview_binding = SimpleNamespace( + id="binding-preview", + agent_id=agent.id, + base_home_snapshot_id="home-preview-old", + agent_config_version_id=normal_draft.id, + agent_config_version_kind=AgentConfigVersionKind.DRAFT, + ) + session.add_all( + [ + agent, + build_draft, + preview_mapping, + empty_preview_mapping, + preview_conversation, + empty_preview_conversation, + ] + ) + session.commit() + lifecycle: list[str] = [] + event.listen(session, "after_commit", lambda _session: lifecycle.append("commit")) + monkeypatch.setattr( + AgentHomeSnapshotService, + "create_for_build_apply", + MagicMock(return_value=SimpleNamespace(id="home-applied")), + ) + monkeypatch.setattr(AgentComposerService, "_save_agent_draft", MagicMock(return_value=normal_draft)) + monkeypatch.setattr(AgentComposerService, "_agent_soul_matches_active_config", MagicMock(return_value=False)) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", MagicMock()) + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", MagicMock()) + get_active_binding = MagicMock(return_value=preview_binding) + retire_binding = MagicMock(side_effect=["binding-preview", "binding-build"]) + validate_generation = MagicMock() + enqueue_collection = MagicMock(side_effect=lambda **_kwargs: lifecycle.append("enqueue")) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active_binding) + monkeypatch.setattr(AgentWorkspaceService, "validate_binding_generation", validate_generation) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) + + AgentComposerService.apply_agent_app_build_draft( + session=session, + tenant_id=agent.tenant_id, + agent_id=agent.id, + account_id="account-1", + ) + + assert preview_conversation.agent_workspace_binding_id is None + validate_generation.assert_called_once_with( + preview_binding, + base_home_snapshot_id="home-preview-old", + agent_config_version_id=normal_draft.id, + agent_config_version_kind=AgentConfigVersionKind.DRAFT, + ) + enqueue_collection.assert_called_once_with( + tenant_id=agent.tenant_id, + binding_ids=["binding-preview", "binding-build"], + ) + assert get_active_binding.call_count == 1 + assert get_active_binding.call_args.kwargs["binding_id"] == "binding-preview" + assert retire_binding.call_args_list == [ + call(session=session, tenant_id=agent.tenant_id, binding_id="binding-preview"), + call(session=session, tenant_id=agent.tenant_id, binding_id="binding-build"), + ] + assert lifecycle == ["commit", "enqueue"] + + +def test_build_apply_validates_before_resolving_or_snapshotting_sandbox( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-1", + config_snapshot=_agent_soul_with_model(), + ) + session.add_all([agent, build_draft]) + session.commit() + validation = MagicMock(side_effect=InvalidComposerConfigError("invalid Build Draft")) + create_home = MagicMock() + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", validation) + monkeypatch.setattr(AgentHomeSnapshotService, "create_for_build_apply", create_home) + + with pytest.raises(InvalidComposerConfigError, match="invalid Build Draft"): + AgentComposerService.apply_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + validation.assert_called_once() + create_home.assert_not_called() + + +def test_build_apply_without_model_snapshots_source_sandbox_and_updates_normal_draft( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-old", + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + normal_draft = AgentConfigDraft( + id="draft-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id="home-old", + config_snapshot=AgentSoulConfig(), + ) + session.add_all([agent, build_draft]) + session.commit() + create_home = MagicMock(return_value=SimpleNamespace(id="home-new", snapshot_ref="backend-home-new")) + validate_knowledge = MagicMock() + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda _payload: None) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", validate_knowledge) + monkeypatch.setattr(AgentHomeSnapshotService, "create_for_build_apply", create_home) + monkeypatch.setattr(AgentComposerService, "_save_agent_draft", lambda **_kwargs: normal_draft) + monkeypatch.setattr(AgentComposerService, "_agent_soul_matches_active_config", lambda **_kwargs: False) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", MagicMock()) + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", MagicMock()) + monkeypatch.setattr(AgentComposerService, "_serialize_draft", lambda draft: {"id": draft.id}) + + result = AgentComposerService.apply_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + create_home.assert_called_once_with( + session=session, + build_draft=build_draft, + ) + validate_knowledge.assert_called_once() + assert normal_draft.home_snapshot_id == "home-new" + assert session.get(AgentConfigDraft, build_draft.id) is None + assert result == {"result": "success", "draft": {"id": "draft-1"}} + + +def test_build_apply_requires_retained_sandbox_before_creating_home( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-old", + config_snapshot=AgentSoulConfig(), + ) + create_home = MagicMock() + save_draft = MagicMock() + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda _payload: None) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **_kwargs: None) + monkeypatch.setattr(AgentHomeSnapshotService, "create_for_build_apply", create_home) + monkeypatch.setattr(AgentComposerService, "_save_agent_draft", save_draft) + session.add_all([agent, build_draft]) + session.commit() + + with pytest.raises(AgentBuildSandboxNotFoundError): + AgentComposerService.apply_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + create_home.assert_not_called() + save_draft.assert_not_called() + assert build_draft.home_snapshot_id == "home-old" + assert session.get(AgentConfigDraft, build_draft.id) is not None + + +def test_build_apply_fails_when_locked_source_cannot_be_retired( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-old", + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + normal_draft = AgentConfigDraft( + id="draft-1", + tenant_id="tenant-1", + agent_id=agent.id, + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id="home-old", + config_snapshot=AgentSoulConfig(), + ) + session.add_all([agent, build_draft]) + session.commit() + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda _payload: None) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **_kwargs: None) + monkeypatch.setattr( + AgentHomeSnapshotService, + "create_for_build_apply", + MagicMock(return_value=SimpleNamespace(id="home-new", snapshot_ref="backend-home-new")), + ) + monkeypatch.setattr(AgentComposerService, "_save_agent_draft", MagicMock(return_value=normal_draft)) + monkeypatch.setattr(AgentComposerService, "_agent_soul_matches_active_config", lambda **_kwargs: False) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", MagicMock(return_value=None)) + enqueue_collection = MagicMock() + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) + + with pytest.raises(AgentBuildSandboxNotFoundError): + AgentComposerService.apply_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id=agent.id, + account_id="account-1", + ) + + assert session.get(AgentConfigDraft, build_draft.id) is not None + enqueue_collection.assert_not_called() + + +def test_build_apply_home_create_failure_leaves_drafts_untouched( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-old", + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + normal_draft = AgentConfigDraft( + id="draft-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id="home-old", + config_snapshot=AgentSoulConfig(), + ) + save_draft = MagicMock(return_value=normal_draft) + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda _payload: None) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **_kwargs: None) + monkeypatch.setattr( + AgentHomeSnapshotService, + "create_for_build_apply", + MagicMock(side_effect=RuntimeError("snapshot failed")), + ) + monkeypatch.setattr(AgentComposerService, "_save_agent_draft", save_draft) + session.add_all([agent, build_draft]) + session.commit() + + with pytest.raises(RuntimeError, match="snapshot failed"): + AgentComposerService.apply_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + save_draft.assert_not_called() + assert build_draft.home_snapshot_id == "home-old" + assert normal_draft.home_snapshot_id == "home-old" + assert session.get(AgentConfigDraft, build_draft.id) is not None + + +def test_build_apply_commit_failure_rolls_back_and_preserves_physical_home( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-old", + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + normal_draft = AgentConfigDraft( + id="draft-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DRAFT, + draft_owner_key="", + home_snapshot_id="home-old", + config_snapshot=AgentSoulConfig(), + ) + session.add_all([agent, build_draft]) + session.commit() + + def fail_commit(_session: Session) -> None: + raise RuntimeError("commit failed") + + event.listen(session, "before_commit", fail_commit) + delete_home = MagicMock() + monkeypatch.setattr(composer_service.ComposerConfigValidator, "validate_publish_payload", lambda _payload: None) + monkeypatch.setattr(AgentComposerService, "validate_knowledge_datasets", lambda **_kwargs: None) + monkeypatch.setattr( + AgentHomeSnapshotService, + "create_for_build_apply", + lambda **_kwargs: SimpleNamespace(id="home-new", snapshot_ref="backend-home-new"), + ) + monkeypatch.setattr(AgentHomeSnapshotService, "delete", delete_home) + monkeypatch.setattr(AgentComposerService, "_save_agent_draft", lambda **_kwargs: normal_draft) + monkeypatch.setattr(AgentComposerService, "_agent_soul_matches_active_config", lambda **_kwargs: False) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", MagicMock(return_value="binding-1")) + enqueue_collection = MagicMock() + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) + + with pytest.raises(RuntimeError, match="commit failed"): + AgentComposerService.apply_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + assert session.get(AgentConfigDraft, build_draft.id) is not None + delete_home.assert_not_called() + enqueue_collection.assert_not_called() + + +def test_build_draft_save_and_discard_do_not_manage_home_resources( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + ) + build_draft = AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-existing", + config_snapshot=AgentSoulConfig(), + ) + create_from_build = MagicMock() + delete_home = MagicMock() + monkeypatch.setattr(AgentHomeSnapshotService, "create_for_build_apply", create_from_build) + monkeypatch.setattr(AgentHomeSnapshotService, "delete", delete_home) + monkeypatch.setattr(AgentComposerService, "_require_agent", lambda **_kwargs: agent) + monkeypatch.setattr(AgentComposerService, "_save_agent_draft", lambda **_kwargs: build_draft) + monkeypatch.setattr(AgentComposerService, "_serialize_build_draft_state", lambda draft: {"id": draft.id}) + save_session = sqlite_session + + AgentComposerService.save_agent_app_build_draft( + session=save_session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + payload=ComposerSavePayload( + variant=ComposerVariant.AGENT_APP, + save_strategy=ComposerSaveStrategy.SAVE_TO_CURRENT_VERSION, + agent_soul=AgentSoulConfig(), + ), + ) + discard_session = sqlite_session + discard_session.add(build_draft) + discard_session.commit() + AgentComposerService.discard_agent_app_build_draft( + session=discard_session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + create_from_build.assert_not_called() + delete_home.assert_not_called() + + +def test_build_discard_retires_then_commits_before_enqueue( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") + build_draft = AgentConfigDraft( + id="build-1", + tenant_id=agent.tenant_id, + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-1", + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + binding = SimpleNamespace( + agent_id=agent.id, + base_home_snapshot_id="home-1", + agent_config_version_id=build_draft.id, + agent_config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) + session.add_all([agent, build_draft]) + session.commit() + events: list[str] = [] + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=binding)) + monkeypatch.setattr( + AgentWorkspaceService, + "retire_binding", + MagicMock(side_effect=lambda **_kwargs: events.append("retire") or "binding-1"), + ) + event.listen(session, "after_commit", lambda _session: events.append("commit")) + monkeypatch.setattr( + composer_service, + "enqueue_agent_resource_collection", + MagicMock(side_effect=lambda **_kwargs: events.append("enqueue")), + ) + + AgentComposerService.discard_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + assert events == ["retire", "commit", "enqueue"] + + +def test_build_discard_commit_failure_does_not_enqueue( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") + build_draft = AgentConfigDraft( + id="build-1", + tenant_id=agent.tenant_id, + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-1", + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + binding = SimpleNamespace( + agent_id=agent.id, + base_home_snapshot_id="home-1", + agent_config_version_id=build_draft.id, + agent_config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) + session.add_all([agent, build_draft]) + session.commit() + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=binding)) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", MagicMock(return_value="binding-1")) + event.listen( + session, + "before_commit", + lambda _session: (_ for _ in ()).throw(RuntimeError("commit failed")), + ) + enqueue_collection = MagicMock() + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) + + with pytest.raises(RuntimeError, match="commit failed"): + AgentComposerService.discard_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", + ) + + enqueue_collection.assert_not_called() + + +def test_build_discard_rejects_unavailable_pointed_binding( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") + build_draft = AgentConfigDraft( + id="build-1", + tenant_id=agent.tenant_id, + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id="home-1", + agent_workspace_binding_id="binding-missing", + config_snapshot=AgentSoulConfig(), + ) + session.add_all([agent, build_draft]) + session.commit() + retire_binding = MagicMock() + enqueue_collection = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=None)) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", enqueue_collection) + + with pytest.raises(AgentBuildSandboxNotFoundError): + AgentComposerService.discard_agent_app_build_draft( + session=session, + tenant_id="tenant-1", + agent_id=agent.id, + account_id="account-1", + ) + + assert build_draft.agent_workspace_binding_id == "binding-missing" + assert session.get(AgentConfigDraft, build_draft.id) is not None + retire_binding.assert_not_called() + enqueue_collection.assert_not_called() @pytest.mark.parametrize( @@ -853,7 +2046,9 @@ def test_agent_app_build_draft_checkout_and_apply_use_user_isolated_draft(monkey def test_load_agent_soul_for_debug_selects_requested_draft( draft_type: AgentConfigDraftType, account_id: str | None, + sqlite_session: Session, ): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -873,33 +2068,22 @@ def test_load_agent_soul_for_debug_selects_requested_draft( draft_owner_key=account_id or "", config_snapshot=agent_soul, ) - fake_session = FakeSession( - scalar=[draft] if draft_type == AgentConfigDraftType.DEBUG_BUILD else [agent, draft, None] - ) + session.add_all([agent, draft]) + session.commit() result = AgentComposerService.load_agent_soul_for_debug( tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", draft_type=draft_type, - session=fake_session, + session=session, ) assert result == agent_soul -def test_load_agent_soul_for_debug_requires_existing_build_draft(): - agent = Agent( - id="agent-1", - tenant_id="tenant-1", - name="Iris", - description="", - agent_kind=AgentKind.DIFY_AGENT, - scope=AgentScope.ROSTER, - source=AgentSource.AGENT_APP, - status=AgentStatus.ACTIVE, - ) - fake_session = FakeSession(scalar=[None]) +def test_load_agent_soul_for_debug_requires_existing_build_draft(sqlite_session: Session): + session = sqlite_session with pytest.raises(AgentVersionNotFoundError): AgentComposerService.load_agent_soul_for_debug( @@ -907,11 +2091,14 @@ def test_load_agent_soul_for_debug_requires_existing_build_draft(): agent_id="agent-1", account_id="account-1", draft_type=AgentConfigDraftType.DEBUG_BUILD, - session=fake_session, + session=session, ) -def test_agent_app_build_draft_apply_marks_unpublished_when_build_draft_differs(monkeypatch: pytest.MonkeyPatch): +def test_agent_app_build_draft_apply_marks_unpublished_when_build_draft_differs( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -940,6 +2127,7 @@ def test_agent_app_build_draft_apply_marks_unpublished_when_build_draft_differs( account_id="account-1", draft_owner_key="account-1", base_snapshot_id="version-1", + agent_workspace_binding_id="binding-1", config_snapshot=build_agent_soul, ) normal_draft = AgentConfigDraft( @@ -951,9 +2139,20 @@ def test_agent_app_build_draft_apply_marks_unpublished_when_build_draft_differs( base_snapshot_id="version-1", config_snapshot=active_agent_soul, ) - active_version = SimpleNamespace(config_snapshot_dict=active_agent_soul.model_dump(mode="json")) - fake_session = FakeSession(scalar=[agent, build_draft, normal_draft, active_version]) - session = fake_session + active_version = AgentConfigSnapshot( + id="version-1", + tenant_id=agent.tenant_id, + agent_id=agent.id, + version=1, + home_snapshot_id="home-initial", + config_snapshot=active_agent_soul, + ) + session.add_all([agent, build_draft, normal_draft, active_version]) + session.commit() + create_home = MagicMock(return_value=SimpleNamespace(id="home-build", snapshot_ref="backend-home-build")) + monkeypatch.setattr(AgentHomeSnapshotService, "create_for_build_apply", create_home) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", MagicMock()) + monkeypatch.setattr(composer_service, "enqueue_agent_resource_collection", MagicMock()) AgentComposerService.apply_agent_app_build_draft( session=session, @@ -964,16 +2163,31 @@ def test_agent_app_build_draft_apply_marks_unpublished_when_build_draft_differs( assert normal_draft.config_snapshot_dict == build_draft.config_snapshot_dict assert agent.active_config_is_published is False - assert fake_session.deleted == [build_draft] - assert fake_session.flushes >= 1 + assert session.get(AgentConfigDraft, build_draft.id) is None + create_home.assert_called_once_with( + session=session, + build_draft=build_draft, + ) -def test_agent_app_composer_candidates_and_impact(monkeypatch: pytest.MonkeyPatch): +def test_agent_app_composer_candidates_and_impact(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session bindings = [ - SimpleNamespace(app_id="app-1", workflow_id="workflow-1", node_id="node-1"), - SimpleNamespace(app_id="app-1", workflow_id="workflow-1", node_id="node-2"), + WorkflowAgentNodeBinding( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_version=Workflow.VERSION_DRAFT, + node_id=node_id, + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + agent_id="agent-1", + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), + ) + for node_id in ("node-1", "node-2") ] - session = FakeSession(scalars=[bindings]) + session.add_all(bindings) + session.commit() # Candidates assembly is covered in test_composer_candidates.py; here we stub # the IO loaders and assert the response envelope per variant (ENG-615). @@ -1007,8 +2221,10 @@ def test_agent_app_composer_candidates_and_impact(monkeypatch: pytest.MonkeyPatc assert impact["bindings"][1]["node_id"] == "node-2" -def test_serialize_workflow_state_changes_lock_and_save_options(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_serialize_workflow_state_changes_lock_and_save_options( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session binding = WorkflowAgentNodeBinding( id="binding-1", tenant_id="tenant-1", @@ -1051,8 +2267,10 @@ def test_serialize_workflow_state_changes_lock_and_save_options(monkeypatch: pyt assert effective_names == ["text", "files", "json"] -def test_serialize_workflow_state_passes_user_declared_outputs_through_effective(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_serialize_workflow_state_passes_user_declared_outputs_through_effective( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session binding = WorkflowAgentNodeBinding( id="binding-1", tenant_id="tenant-1", @@ -1090,8 +2308,9 @@ def test_serialize_workflow_state_passes_user_declared_outputs_through_effective def test_serialize_workflow_state_includes_inline_debug_conversation_message_state( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ): - session = FakeSession() + session = sqlite_session binding = WorkflowAgentNodeBinding( id="binding-1", tenant_id="tenant-1", @@ -1137,9 +2356,8 @@ def test_serialize_workflow_state_includes_inline_debug_conversation_message_sta assert state["debug_conversation_message_count"] == 2 -def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession() - session = fake_session +def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session workflow_agent = SimpleNamespace(id="inline-agent-1", active_config_snapshot_id="inline-version-1") roster_agent = SimpleNamespace( id="roster-agent-1", @@ -1150,6 +2368,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk icon_type="emoji", icon="source", icon_background="#FFFFFF", + scope=AgentScope.WORKFLOW_ONLY, ) create_roster_calls = [] copy_drive_calls = [] @@ -1178,6 +2397,7 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk tenant_id="tenant-1", agent_id="roster-agent-1", version=1, + home_snapshot_id="home-source", config_snapshot='{"prompt":{"system_prompt":"old"}}', ), ) @@ -1186,7 +2406,6 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk "_create_config_version", lambda **kwargs: AgentConfigSnapshot(id="new-version-1", version=2), ) - payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.WORKFLOW.value, @@ -1282,9 +2501,8 @@ def test_composer_save_helpers_create_and_rebind_agents(monkeypatch: pytest.Monk ] -def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession() - session = fake_session +def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session inline_agent = SimpleNamespace( id="inline-agent-1", scope=AgentScope.WORKFLOW_ONLY, @@ -1297,6 +2515,7 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch tenant_id="tenant-1", agent_id="inline-agent-1", version=1, + home_snapshot_id="home-inline-1", config_snapshot='{"prompt":{"system_prompt":"old"}}', ) next_snapshot = AgentConfigSnapshot( @@ -1304,6 +2523,7 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch tenant_id="tenant-1", agent_id="inline-agent-1", version=2, + home_snapshot_id="home-inline-2", config_snapshot=AgentSoulConfig.model_validate( { "model": { @@ -1323,6 +2543,7 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch account_id=None, draft_owner_key="", base_snapshot_id="inline-version-1", + home_snapshot_id="home-inline-1", config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "old"}}), ) @@ -1376,11 +2597,13 @@ def test_node_job_only_updates_inline_agent_soul(monkeypatch: pytest.MonkeyPatch 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.home_snapshot_id == "home-inline-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(): +def test_get_or_create_normal_agent_draft_rebases_stale_workflow_only_draft(sqlite_session: Session): + session = sqlite_session agent = Agent( id="inline-agent-1", tenant_id="tenant-1", @@ -1402,6 +2625,7 @@ def test_get_or_create_normal_agent_draft_rebases_stale_workflow_only_draft(): account_id=None, draft_owner_key="", base_snapshot_id="inline-version-1", + home_snapshot_id="home-inline-1", config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "old"}}), ) active_snapshot = AgentConfigSnapshot( @@ -1409,9 +2633,11 @@ def test_get_or_create_normal_agent_draft_rebases_stale_workflow_only_draft(): tenant_id="tenant-1", agent_id=agent.id, version=2, + home_snapshot_id="home-inline-2", config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "new"}}), ) - session = FakeSession(scalar=[draft, active_snapshot]) + session.add_all([agent, draft, active_snapshot]) + session.commit() resolved = AgentComposerService.get_or_create_normal_agent_draft( session=session, @@ -1423,12 +2649,13 @@ def test_get_or_create_normal_agent_draft_rebases_stale_workflow_only_draft(): assert resolved is draft assert resolved.id == "draft-1" assert resolved.base_snapshot_id == "inline-version-2" + assert resolved.home_snapshot_id == "home-inline-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(): +def test_get_or_create_normal_agent_draft_keeps_roster_draft_edits(sqlite_session: Session): + session = sqlite_session agent = Agent( id="roster-agent-1", tenant_id="tenant-1", @@ -1450,7 +2677,8 @@ def test_get_or_create_normal_agent_draft_keeps_roster_draft_edits(): base_snapshot_id="version-1", config_snapshot=AgentSoulConfig.model_validate({"prompt": {"system_prompt": "local edit"}}), ) - session = FakeSession(scalar=[draft]) + session.add_all([agent, draft]) + session.commit() resolved = AgentComposerService.get_or_create_normal_agent_draft( session=session, @@ -1462,12 +2690,12 @@ def test_get_or_create_normal_agent_draft_keeps_roster_draft_edits(): 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): - fake_session = FakeSession() - session = fake_session +def test_node_job_only_switches_roster_binding_to_inline_agent( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session created_agent = SimpleNamespace(id="inline-agent-1", active_config_snapshot_id="inline-version-1") captured: dict[str, object] = {} @@ -1522,11 +2750,10 @@ def test_node_job_only_switches_roster_binding_to_inline_agent(monkeypatch: pyte assert captured["node_id"] == "node-1" assert captured["account_id"] == "account-1" assert captured["agent_soul"].prompt.system_prompt == "start from scratch" - assert fake_session.flushes == 1 -def test_node_job_only_rejects_start_from_scratch_with_existing_inline_binding_id(): - session = FakeSession() +def test_node_job_only_rejects_start_from_scratch_with_existing_inline_binding_id(sqlite_session: Session): + session = sqlite_session binding = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", @@ -1562,9 +2789,10 @@ def test_node_job_only_rejects_start_from_scratch_with_existing_inline_binding_i ) -def test_node_job_only_rejects_inline_binding_pointing_to_roster_agent(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession() - session = fake_session +def test_node_job_only_rejects_inline_binding_pointing_to_roster_agent( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session current_snapshot = AgentConfigSnapshot( id="inline-version-1", tenant_id="tenant-1", @@ -1612,9 +2840,9 @@ def test_node_job_only_rejects_inline_binding_pointing_to_roster_agent(monkeypat def test_copy_workflow_composer_from_roster_creates_inline_agent_and_preserves_node_job( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ): - fake_session = FakeSession(scalar=["publish-revision-1"]) - session = fake_session + session = sqlite_session workflow = SimpleNamespace(id="workflow-1") node_job = WorkflowNodeJobConfig(workflow_prompt="keep this node task") binding = WorkflowAgentNodeBinding( @@ -1646,6 +2874,16 @@ def test_copy_workflow_composer_from_roster_creates_inline_agent_and_preserves_n version=2, config_snapshot='{"prompt":{"system_prompt":"copy me"}}', ) + session.add( + AgentConfigRevision( + tenant_id="tenant-1", + agent_id=roster_agent.id, + current_snapshot_id=source_version.id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + ) + ) + session.commit() inline_agent = Agent( id="inline-agent-1", tenant_id="tenant-1", @@ -1707,11 +2945,12 @@ def test_copy_workflow_composer_from_roster_creates_inline_agent_and_preserves_n drive_kwargs = captured["drive"] assert drive_kwargs["source_agent_id"] == "roster-agent-1" assert drive_kwargs["target_agent_id"] == "inline-agent-1" - assert fake_session.flushes >= 1 -def test_copy_workflow_composer_from_roster_rejects_stale_source_snapshot(monkeypatch: pytest.MonkeyPatch): - session = FakeSession(scalar=["publish-revision-1"]) +def test_copy_workflow_composer_from_roster_rejects_stale_source_snapshot( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr( AgentComposerService, @@ -1744,6 +2983,16 @@ def test_copy_workflow_composer_from_roster_rejects_stale_source_snapshot(monkey version=2, config_snapshot='{"prompt":{"system_prompt":"copy me"}}', ) + session.add( + AgentConfigRevision( + tenant_id="tenant-1", + agent_id=roster_agent.id, + current_snapshot_id=source_version.id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + ) + ) + session.commit() monkeypatch.setattr(AgentComposerService, "_require_agent", lambda **kwargs: roster_agent) monkeypatch.setattr(AgentComposerService, "_require_version", lambda **kwargs: source_version) @@ -1759,8 +3008,10 @@ def test_copy_workflow_composer_from_roster_rejects_stale_source_snapshot(monkey ) -def test_copy_workflow_composer_from_roster_rejects_unpublished_source(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_copy_workflow_composer_from_roster_rejects_unpublished_source( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session binding = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", @@ -1799,10 +3050,13 @@ def test_copy_workflow_composer_from_roster_rejects_unpublished_source(monkeypat ) require_version.assert_not_called() - assert session.flushes == 0 + assert not session.new + assert not session.dirty -def test_copy_workflow_composer_from_roster_is_idempotent_when_already_inline(monkeypatch: pytest.MonkeyPatch): +def test_copy_workflow_composer_from_roster_is_idempotent_when_already_inline( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): inline_binding = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", @@ -1830,7 +3084,7 @@ def test_copy_workflow_composer_from_roster_is_idempotent_when_already_inline(mo config_snapshot='{"prompt":{"system_prompt":"inline"}}', ) serialize_calls = [] - session = FakeSession() + session = sqlite_session monkeypatch.setattr(AgentComposerService, "_get_draft_workflow", lambda **kwargs: SimpleNamespace(id="workflow-1")) monkeypatch.setattr(AgentComposerService, "_get_workflow_binding", lambda **kwargs: inline_binding) monkeypatch.setattr(AgentComposerService, "_get_agent_if_present", lambda **kwargs: inline_agent) @@ -1896,8 +3150,9 @@ def test_copy_workflow_composer_from_roster_rejects_invalid_source_binding( source_scope: AgentScope, source_status: AgentStatus, expected_message: str, + sqlite_session: Session, ): - session = FakeSession() + session = sqlite_session binding = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", @@ -1933,7 +3188,8 @@ def test_copy_workflow_composer_from_roster_rejects_invalid_source_binding( ) -def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest.MonkeyPatch): +def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session skill_row = AgentDriveFile( tenant_id="tenant-1", agent_id="roster-agent-1", @@ -1966,8 +3222,8 @@ def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest size=30, mime_type="application/pdf", ) - fake_session = FakeSession(scalars=[[skill_row, script_row, file_row], []]) - session = fake_session + session.add_all([skill_row, script_row, file_row]) + session.commit() agent_soul = AgentSoulConfig.model_validate( { "prompt": { @@ -1989,21 +3245,31 @@ def test_copy_agent_drive_rows_copies_skill_prefix_and_files(monkeypatch: pytest node_job=node_job, ) - copied = [row for row in fake_session.added if isinstance(row, AgentDriveFile)] - assert [row.key for row in copied] == [ + session.flush() + copied = list( + session.scalars( + select(AgentDriveFile).where( + AgentDriveFile.tenant_id == "tenant-1", + AgentDriveFile.agent_id == "inline-agent-1", + ) + ) + ) + assert {row.key for row in copied} == { "tender-analyzer/SKILL.md", "tender-analyzer/scripts/run.sh", "files/qna.pdf", - ] + } assert {row.agent_id for row in copied} == {"inline-agent-1"} - assert copied[0].file_id == "tool-file-1" - assert copied[0].is_skill is True - assert copied[2].value_owned_by_drive is False + copied_by_key = {row.key: row for row in copied} + assert copied_by_key["tender-analyzer/SKILL.md"].file_id == "tool-file-1" + assert copied_by_key["tender-analyzer/SKILL.md"].is_skill is True + assert copied_by_key["files/qna.pdf"].value_owned_by_drive is False -def test_copy_agent_drive_rows_skips_when_no_referenced_drive_keys(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession() - session = fake_session +def test_copy_agent_drive_rows_skips_when_no_referenced_drive_keys( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session agent_soul = AgentSoulConfig.model_validate({"prompt": {"system_prompt": "No drive mentions."}}) AgentComposerService._copy_agent_drive_rows( @@ -2015,10 +3281,11 @@ def test_copy_agent_drive_rows_skips_when_no_referenced_drive_keys(monkeypatch: agent_soul=agent_soul, ) - assert fake_session.added == [] + assert not session.new -def test_copy_agent_drive_rows_skips_existing_target_keys(monkeypatch: pytest.MonkeyPatch): +def test_copy_agent_drive_rows_skips_existing_target_keys(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session source_row = AgentDriveFile( tenant_id="tenant-1", agent_id="roster-agent-1", @@ -2029,8 +3296,18 @@ def test_copy_agent_drive_rows_skips_existing_target_keys(monkeypatch: pytest.Mo size=30, mime_type="application/pdf", ) - fake_session = FakeSession(scalars=[[source_row], ["files/qna.pdf"]]) - session = fake_session + target_row = AgentDriveFile( + tenant_id="tenant-1", + agent_id="inline-agent-1", + key=source_row.key, + file_kind=source_row.file_kind, + file_id=source_row.file_id, + value_owned_by_drive=source_row.value_owned_by_drive, + size=source_row.size, + mime_type=source_row.mime_type, + ) + session.add_all([source_row, target_row]) + session.commit() agent_soul = AgentSoulConfig.model_validate({"prompt": {"system_prompt": "[§file:files/qna.pdf:qna.pdf§]"}}) AgentComposerService._copy_agent_drive_rows( @@ -2042,7 +3319,16 @@ def test_copy_agent_drive_rows_skips_existing_target_keys(monkeypatch: pytest.Mo agent_soul=agent_soul, ) - assert [row for row in fake_session.added if isinstance(row, AgentDriveFile)] == [] + session.flush() + target_rows = list( + session.scalars( + select(AgentDriveFile).where( + AgentDriveFile.tenant_id == "tenant-1", + AgentDriveFile.agent_id == "inline-agent-1", + ) + ) + ) + assert [row.key for row in target_rows] == ["files/qna.pdf"] def test_drive_copy_scopes_include_declared_output_benchmark_files(): @@ -2087,9 +3373,11 @@ def test_drive_copy_scopes_include_declared_output_benchmark_files(): assert prefixes == {"tender-analyzer/"} -def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession() - session = fake_session +def test_composer_create_agents_syncs_active_config_has_model( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + session = sqlite_session created_apps = [] hidden_backing_apps = [] backing_agent = Agent( @@ -2126,13 +3414,15 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes monkeypatch.setattr( AgentComposerService, "_require_version", - lambda **kwargs: SimpleNamespace(id="empty-version-1", tenant_id="tenant-1", agent_id="roster-agent-1"), - ) - monkeypatch.setattr( - AgentComposerService, - "_create_config_version", - lambda **kwargs: SimpleNamespace(id="version-with-model"), + lambda **kwargs: SimpleNamespace( + id="empty-version-1", + tenant_id="tenant-1", + agent_id="roster-agent-1", + home_snapshot_id=None, + ), ) + create_config_version = MagicMock(return_value=SimpleNamespace(id="version-with-model")) + monkeypatch.setattr(AgentComposerService, "_create_config_version", create_config_version) workflow_agent = AgentComposerService._create_workflow_only_agent( session=session, @@ -2156,6 +3446,8 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes assert workflow_agent.active_config_snapshot_id == "version-with-model" assert workflow_agent.active_config_has_model is True assert workflow_agent.backing_app_id == "hidden-app-1" + assert create_config_version.call_count == 2 + assert all(call.kwargs["home_snapshot_id"] is None for call in create_config_version.call_args_list) assert hidden_backing_apps[0]["name"] == "Workflow Agent node-1" assert roster_agent.active_config_snapshot_id == "version-with-model" assert roster_agent.active_config_has_model is True @@ -2168,23 +3460,26 @@ def test_composer_create_agents_syncs_active_config_has_model(monkeypatch: pytes assert created_account.id == "account-1" -def test_composer_require_account(monkeypatch: pytest.MonkeyPatch): - account = SimpleNamespace(id="account-1") - session = SimpleNamespace(get=lambda model, account_id: account) +def test_composer_require_account(sqlite_session: Session): + session = sqlite_session + account = Account(name="Tester", email="tester@example.com") + account.id = "account-1" + session.add(account) + session.commit() assert AgentComposerService._require_account(session=session, account_id="account-1") is account -def test_composer_require_account_raises_when_missing(monkeypatch: pytest.MonkeyPatch): - session = SimpleNamespace(get=lambda model, account_id: None) - +def test_composer_require_account_raises_when_missing(sqlite_session: Session): with pytest.raises(ValueError, match="Account not found"): - AgentComposerService._require_account(session=session, account_id="missing-account") + AgentComposerService._require_account(session=sqlite_session, account_id="missing-account") -def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession() - session = fake_session +def test_composer_create_roster_agent_maps_name_conflict_without_owning_rollback( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session + transaction = session.begin() class FakeAppService: def create_app(self, tenant_id, params, account, *, session): @@ -2204,12 +3499,13 @@ def test_composer_create_roster_agent_rolls_back_name_conflict(monkeypatch: pyte version_note=None, ) - assert fake_session.rollbacks == 1 + assert transaction.is_active -def test_composer_create_roster_agent_raises_when_backing_agent_missing(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession() - session = fake_session +def test_composer_create_roster_agent_raises_when_backing_agent_missing( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session class FakeAppService: def create_app(self, tenant_id, params, account, *, session): @@ -2238,7 +3534,9 @@ def test_composer_create_roster_agent_raises_when_backing_agent_missing(monkeypa ) -def test_agent_app_draft_match_does_not_mark_create_version_as_published(monkeypatch: pytest.MonkeyPatch): +def test_agent_app_draft_match_does_not_mark_create_version_as_published( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): agent_soul = AgentSoulConfig() agent = Agent( id="agent-1", @@ -2246,9 +3544,8 @@ def test_agent_app_draft_match_does_not_mark_create_version_as_published(monkeyp source=AgentSource.AGENT_APP, active_config_snapshot_id="snapshot-1", ) - snapshot = SimpleNamespace(config_snapshot_dict=agent_soul) - fake_session = FakeSession() - session = fake_session + snapshot = SimpleNamespace(config_snapshot_dict=agent_soul, home_snapshot_id="home-1") + session = sqlite_session monkeypatch.setattr(AgentComposerService, "_get_version_if_present", lambda **kwargs: snapshot) assert ( @@ -2257,12 +3554,15 @@ def test_agent_app_draft_match_does_not_mark_create_version_as_published(monkeyp tenant_id="tenant-1", agent=agent, agent_soul=agent_soul, + home_snapshot_id="home-1", ) is False ) -def test_agent_app_draft_match_marks_publish_visible_revision_as_published(monkeypatch: pytest.MonkeyPatch): +def test_agent_app_draft_match_marks_publish_visible_revision_as_published( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): agent_soul = AgentSoulConfig() agent = Agent( id="agent-1", @@ -2270,9 +3570,18 @@ def test_agent_app_draft_match_marks_publish_visible_revision_as_published(monke source=AgentSource.AGENT_APP, active_config_snapshot_id="snapshot-1", ) - snapshot = SimpleNamespace(config_snapshot_dict=agent_soul) - fake_session = FakeSession(scalar=["publish-revision-1"]) - session = fake_session + snapshot = SimpleNamespace(config_snapshot_dict=agent_soul, home_snapshot_id="home-1") + session = sqlite_session + session.add( + AgentConfigRevision( + tenant_id=agent.tenant_id, + agent_id=agent.id, + current_snapshot_id=agent.active_config_snapshot_id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + ) + ) + session.commit() monkeypatch.setattr(AgentComposerService, "_get_version_if_present", lambda **kwargs: snapshot) assert ( @@ -2281,29 +3590,27 @@ def test_agent_app_draft_match_marks_publish_visible_revision_as_published(monke tenant_id="tenant-1", agent=agent, agent_soul=agent_soul, + home_snapshot_id="home-1", ) is True ) -def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession( - scalar=[ - 1, - 3, - 2, - 4, - SimpleNamespace(id="workflow-1"), - None, - SimpleNamespace(id="agent-1"), - None, - SimpleNamespace(id="version-1"), - None, - ] +def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session + agent = _agent() + initial_snapshot = _snapshot(snapshot_id="version-1") + initial_revision = AgentConfigRevision( + tenant_id="tenant-1", + agent_id=agent.id, + current_snapshot_id=initial_snapshot.id, + revision=1, + operation=AgentConfigRevisionOperation.CREATE_VERSION, ) - session = fake_session + workflow_row = _workflow() + session.add_all([agent, initial_snapshot, initial_revision, workflow_row]) + session.commit() agent_soul = AgentSoulConfig.model_validate({"prompt": {"system_prompt": "new"}}) - version = AgentComposerService._create_config_version( session=session, tenant_id="tenant-1", @@ -2312,6 +3619,7 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa agent_soul=agent_soul, operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, version_note="note", + home_snapshot_id="home-1", ) updated_snapshot = AgentComposerService._update_current_version( session=session, @@ -2320,6 +3628,7 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa tenant_id="tenant-1", agent_id="agent-1", version=1, + home_snapshot_id="home-1", config_snapshot='{"prompt":{"system_prompt":"old"}}', ), account_id="account-1", @@ -2336,7 +3645,7 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa ) with pytest.raises(composer_service.AgentNotFoundError): AgentComposerService._require_agent(session=session, tenant_id="tenant-1", agent_id=None) - assert AgentComposerService._get_agent_if_present(session=session, tenant_id="tenant-1", agent_id="agent-1") is None + assert AgentComposerService._get_agent_if_present(session=session, tenant_id="tenant-1", agent_id="missing") is None assert ( AgentComposerService._require_version( session=session, @@ -2356,12 +3665,13 @@ def test_composer_version_helpers_and_lookup_errors(monkeypatch: pytest.MonkeyPa assert version.version == 2 assert updated_snapshot.version == 3 + assert version.home_snapshot_id == "home-1" + assert updated_snapshot.home_snapshot_id == "home-1" assert workflow.id == "workflow-1" -def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession(scalar=[2]) - session = fake_session +def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session payload = ComposerSavePayload.model_validate( { "variant": ComposerVariant.WORKFLOW.value, @@ -2376,6 +3686,7 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc tenant_id="tenant-1", agent_id="agent-1", version=1, + home_snapshot_id="home-1", config_snapshot='{"prompt":{"system_prompt":"old"}}', ) monkeypatch.setattr(AgentComposerService, "_require_version", lambda **kwargs: version) @@ -2384,7 +3695,6 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc "_require_agent", lambda **kwargs: SimpleNamespace(updated_by=None, active_config_is_published=False), ) - result = AgentComposerService._save_to_current_version( session=session, tenant_id="tenant-1", @@ -2395,6 +3705,9 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc assert result.updated_by == "account-1" assert result.current_snapshot_id != "version-1" + created_version = session.get(AgentConfigSnapshot, result.current_snapshot_id) + assert created_version is not None + assert created_version.home_snapshot_id == "home-1" with pytest.raises(ValueError): AgentComposerService._require_binding(None) with pytest.raises(ValueError): @@ -2415,7 +3728,8 @@ def test_composer_current_version_and_error_paths(monkeypatch: pytest.MonkeyPatc ) -def test_roster_list_and_invite_options(monkeypatch: pytest.MonkeyPatch): +def test_roster_list_and_invite_options(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session created_at = datetime(2026, 1, 2, 3, 4, 5, tzinfo=UTC) updated_at = datetime(2026, 1, 3, 3, 4, 5, tzinfo=UTC) version_created_at = datetime(2026, 1, 4, 3, 4, 5, tzinfo=UTC) @@ -2433,7 +3747,11 @@ def test_roster_list_and_invite_options(monkeypatch: pytest.MonkeyPatch): agent.created_at = created_at agent.updated_at = updated_at version = AgentConfigSnapshot( - id="version-1", agent_id="agent-1", version=1, config_snapshot=_agent_soul_with_model() + id="version-1", + tenant_id="tenant-1", + agent_id="agent-1", + version=1, + config_snapshot=_agent_soul_with_model(), ) version.created_at = version_created_at agent.active_config_snapshot_id = "version-1" @@ -2452,23 +3770,36 @@ def test_roster_list_and_invite_options(monkeypatch: pytest.MonkeyPatch): ) unconfigured_agent.active_config_snapshot_id = "version-2" unconfigured_agent.active_config_has_model = False + unconfigured_agent.updated_at = datetime(2026, 1, 1, 3, 4, 5, tzinfo=UTC) unconfigured_version = AgentConfigSnapshot( - id="version-2", agent_id="agent-2", version=1, config_snapshot=AgentSoulConfig() + id="version-2", + tenant_id="tenant-1", + agent_id="agent-2", + version=1, + config_snapshot=AgentSoulConfig(), ) - fake_session = FakeSession( - scalar=[2, 1, SimpleNamespace(id="workflow-1")], - scalars=[ - [agent, unconfigured_agent], - [agent], - [SimpleNamespace(agent_id="agent-1", node_id="node-1")], - ], + workflow_row = _workflow() + binding = WorkflowAgentNodeBinding( + tenant_id="tenant-1", + app_id="app-1", + workflow_id=workflow_row.id, + workflow_version=Workflow.VERSION_DRAFT, + node_id="node-1", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + agent_id=agent.id, + current_snapshot_id=version.id, + node_job_config=WorkflowNodeJobConfig(), ) - service = AgentRosterService(fake_session) - monkeypatch.setattr( - service, - "_load_versions_by_id", - lambda version_ids: {"version-1": version, "version-2": unconfigured_version}, + publish_revision = AgentConfigRevision( + tenant_id="tenant-1", + agent_id=agent.id, + current_snapshot_id=version.id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, ) + session.add_all([agent, unconfigured_agent, version, unconfigured_version, workflow_row, binding, publish_revision]) + session.commit() + service = AgentRosterService(session) monkeypatch.setattr(service, "_load_published_references_by_agent_id", lambda **kwargs: {}) monkeypatch.setattr(service, "_load_reference_counts_by_agent_id", lambda **kwargs: {"agent-1": 1}) @@ -2490,7 +3821,8 @@ def test_roster_list_and_invite_options(monkeypatch: pytest.MonkeyPatch): assert invited["data"][0]["existing_node_ids"] == ["node-1"] -def test_invite_options_uses_db_filtered_pagination(monkeypatch: pytest.MonkeyPatch): +def test_invite_options_uses_db_filtered_pagination(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session configured_agent = Agent( id="agent-2", tenant_id="tenant-1", @@ -2502,18 +3834,25 @@ def test_invite_options_uses_db_filtered_pagination(monkeypatch: pytest.MonkeyPa status=AgentStatus.ACTIVE, active_config_snapshot_id="version-2", active_config_has_model=True, + active_config_is_published=True, ) - fake_session = FakeSession(scalar=[1], scalars=[[configured_agent]]) - service = AgentRosterService(fake_session) - monkeypatch.setattr( - service, - "_load_versions_by_id", - lambda version_ids: { - "version-2": AgentConfigSnapshot( - id="version-2", agent_id="agent-2", version=1, config_snapshot=_agent_soul_with_model() - ) - }, + version = AgentConfigSnapshot( + id="version-2", + tenant_id="tenant-1", + agent_id="agent-2", + version=1, + config_snapshot=_agent_soul_with_model(), ) + publish_revision = AgentConfigRevision( + tenant_id="tenant-1", + agent_id=configured_agent.id, + current_snapshot_id=version.id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + ) + session.add_all([configured_agent, version, publish_revision]) + session.commit() + service = AgentRosterService(session) monkeypatch.setattr(service, "_load_published_references_by_agent_id", lambda **kwargs: {}) monkeypatch.setattr(service, "_load_reference_counts_by_agent_id", lambda **kwargs: {}) @@ -2524,7 +3863,7 @@ def test_invite_options_uses_db_filtered_pagination(monkeypatch: pytest.MonkeyPa assert [item["id"] for item in result["data"]] == ["agent-2"] -def test_active_config_is_published_flags_use_stored_agent_state(): +def test_active_config_is_published_flags_use_stored_agent_state(sqlite_session: Session): agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -2561,7 +3900,7 @@ def test_active_config_is_published_flags_use_stored_agent_state(): active_config_snapshot_id="version-3", active_config_is_published=False, ) - service = AgentRosterService(FakeSession()) + service = AgentRosterService(sqlite_session) flags = service.load_active_config_is_published_by_agent_id( tenant_id="tenant-1", agents=[agent, draft_agent, dirty_agent] @@ -2569,13 +3908,13 @@ def test_active_config_is_published_flags_use_stored_agent_state(): assert flags == {"agent-1": True, "agent-2": False, "agent-3": False} assert service.active_config_is_published(tenant_id="tenant-1", agent=agent) is True - assert AgentRosterService(FakeSession()).load_active_config_is_published_by_agent_id( + assert AgentRosterService(sqlite_session).load_active_config_is_published_by_agent_id( tenant_id="tenant-1", agents=[draft_agent], ) == {"agent-2": False} -def test_active_config_is_published_skips_empty_agent_ids(): +def test_active_config_is_published_skips_empty_agent_ids(sqlite_session: Session): empty_id_agent = Agent( id="", tenant_id="tenant-1", @@ -2587,19 +3926,19 @@ def test_active_config_is_published_skips_empty_agent_ids(): status=AgentStatus.ACTIVE, active_config_snapshot_id=None, ) - fake_session = FakeSession(scalars=[["should-not-be-read"]]) + session = sqlite_session assert ( - AgentRosterService(fake_session).load_active_config_is_published_by_agent_id( + AgentRosterService(session).load_active_config_is_published_by_agent_id( tenant_id="tenant-1", agents=[empty_id_agent], ) == {} ) - assert fake_session._scalars == [["should-not-be-read"]] -def test_load_app_backing_agents_skips_empty_agent_ids(): +def test_load_app_backing_agents_skips_empty_agent_ids(sqlite_session: Session): + session = sqlite_session valid_agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -2623,7 +3962,9 @@ def test_load_app_backing_agents_skips_empty_agent_ids(): status=AgentStatus.ACTIVE, ) - result = AgentRosterService(FakeSession(scalars=[[valid_agent, empty_id_agent]])).load_app_backing_agents_by_app_id( + session.add_all([valid_agent, empty_id_agent]) + session.commit() + result = AgentRosterService(session).load_app_backing_agents_by_app_id( tenant_id="tenant-1", app_ids=["app-1", "app-2"], ) @@ -2631,50 +3972,48 @@ def test_load_app_backing_agents_skips_empty_agent_ids(): assert result == {"app-1": valid_agent} -def test_published_references_include_app_display_fields_and_sort_by_updated_at(): +def test_published_references_include_app_display_fields_and_sort_by_updated_at(sqlite_session: Session): + session = sqlite_session recent_updated_at = datetime(2026, 1, 7, 3, 4, 5, tzinfo=UTC) stale_updated_at = datetime(2026, 1, 6, 3, 4, 5, tzinfo=UTC) bindings = [ - SimpleNamespace( + WorkflowAgentNodeBinding( tenant_id="tenant-1", agent_id="agent-1", app_id="app-stale", workflow_id="workflow-stale", workflow_version="published-stale", node_id="node-b", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), - SimpleNamespace( + WorkflowAgentNodeBinding( tenant_id="tenant-1", agent_id="agent-1", app_id="app-recent", workflow_id="workflow-recent", workflow_version="published-recent", node_id="node-a", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), ] - apps = [ - SimpleNamespace( - id="app-stale", - name="Stale Workflow", - mode="advanced-chat", - workflow_id="workflow-stale", - icon_type=SimpleNamespace(value="emoji"), - icon="old", - icon_background="#F3F4F6", - updated_at=stale_updated_at, - ), - SimpleNamespace( - id="app-recent", - name="Recent Workflow", - mode="advanced-chat", - workflow_id="workflow-recent", - icon_type=SimpleNamespace(value="image"), - icon="upload-file-id", - icon_background="#E0F2FE", - updated_at=recent_updated_at, - ), - ] - service = AgentRosterService(FakeSession(scalars=[bindings, apps])) + stale_app = _app(app_id="app-stale", name="Stale Workflow", mode=AppMode.ADVANCED_CHAT) + stale_app.workflow_id = "workflow-stale" + stale_app.icon = "old" + stale_app.icon_background = "#F3F4F6" + stale_app.updated_at = stale_updated_at + recent_app = _app(app_id="app-recent", name="Recent Workflow", mode=AppMode.ADVANCED_CHAT) + recent_app.workflow_id = "workflow-recent" + recent_app.icon_type = IconType.IMAGE + recent_app.icon = "upload-file-id" + recent_app.icon_background = "#E0F2FE" + recent_app.updated_at = recent_updated_at + session.add_all([*bindings, stale_app, recent_app]) + session.commit() + service = AgentRosterService(session) result = service._load_published_references_by_agent_id(tenant_id="tenant-1", agent_ids=["agent-1"]) @@ -2687,64 +4026,102 @@ def test_published_references_include_app_display_fields_and_sort_by_updated_at( assert references[0]["workflow_version"] == "published-recent" -def test_reference_counts_include_draft_and_published_bindings_once_per_app(): +def test_reference_counts_include_draft_and_published_bindings_once_per_app(sqlite_session: Session): + session = sqlite_session bindings = [ - SimpleNamespace( + WorkflowAgentNodeBinding( + tenant_id="tenant-1", agent_id="agent-1", app_id="app-1", workflow_id="workflow-draft", workflow_version=Workflow.VERSION_DRAFT, + node_id="node-draft", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), - SimpleNamespace( + WorkflowAgentNodeBinding( + tenant_id="tenant-1", agent_id="agent-1", app_id="app-1", workflow_id="workflow-published", workflow_version="v1", + node_id="node-published", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), - SimpleNamespace( + WorkflowAgentNodeBinding( + tenant_id="tenant-1", agent_id="agent-1", app_id="app-2", workflow_id="workflow-stale", workflow_version="old-version", + node_id="node-stale", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), ] - apps = [ - SimpleNamespace(id="app-1", workflow_id="workflow-published"), - SimpleNamespace(id="app-2", workflow_id="workflow-stale"), - ] - workflows = [ - SimpleNamespace(id="workflow-draft", app_id="app-1", version=Workflow.VERSION_DRAFT), - SimpleNamespace(id="workflow-published", app_id="app-1", version="v1"), - SimpleNamespace(id="workflow-stale", app_id="app-2", version="current-version"), - ] - service = AgentRosterService(FakeSession(scalars=[bindings, apps, workflows])) + app_one = _app(app_id="app-1", mode=AppMode.ADVANCED_CHAT) + app_one.workflow_id = "workflow-published" + app_two = _app(app_id="app-2", mode=AppMode.ADVANCED_CHAT) + app_two.workflow_id = "workflow-stale" + draft_workflow = _workflow(workflow_id="workflow-draft", app_id="app-1") + published_workflow = _workflow(workflow_id="workflow-published", app_id="app-1") + published_workflow.version = "v1" + stale_workflow = _workflow(workflow_id="workflow-stale", app_id="app-2") + stale_workflow.version = "current-version" + session.add_all([*bindings, app_one, app_two, draft_workflow, published_workflow, stale_workflow]) + session.commit() + service = AgentRosterService(session) result = service._load_reference_counts_by_agent_id(tenant_id="tenant-1", agent_ids=["agent-1"]) assert result == {"agent-1": 1} -def test_roster_update_archive_versions_and_detail(monkeypatch: pytest.MonkeyPatch): - listed_version = AgentConfigSnapshot(id="version-4", agent_id="agent-1", version=4) +def test_roster_update_archive_versions_and_detail(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session + listed_version = AgentConfigSnapshot( + id="version-4", + tenant_id="tenant-1", + agent_id="agent-1", + version=4, + config_snapshot=AgentSoulConfig(), + ) listed_version_created_at = datetime(2026, 1, 5, 3, 4, 5, tzinfo=UTC) listed_version.created_at = listed_version_created_at - older_listed_version = AgentConfigSnapshot(id="version-2", agent_id="agent-1", version=2) + older_listed_version = AgentConfigSnapshot( + id="version-2", + tenant_id="tenant-1", + agent_id="agent-1", + version=2, + config_snapshot='{"prompt":{}}', + ) older_listed_version.created_at = datetime(2026, 1, 4, 3, 4, 5, tzinfo=UTC) revision_created_at = datetime(2026, 1, 6, 3, 4, 5, tzinfo=UTC) - revision = SimpleNamespace( + revision = AgentConfigRevision( id="revision-1", + tenant_id="tenant-1", + agent_id="agent-1", previous_snapshot_id=None, current_snapshot_id="version-2", revision=1, - operation=AgentConfigRevisionOperation.CREATE_VERSION, + operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, summary=None, version_note=None, created_by="account-1", - created_at=revision_created_at, ) - fake_session = FakeSession( - scalar=["visible-revision"], - scalars=[[listed_version, older_listed_version], [older_listed_version, listed_version], [revision]], + revision.created_at = revision_created_at + listed_revision = AgentConfigRevision( + tenant_id="tenant-1", + agent_id="agent-1", + previous_snapshot_id="version-2", + current_snapshot_id="version-4", + revision=2, + operation=AgentConfigRevisionOperation.SAVE_NEW_VERSION, + created_by="account-1", ) agent = Agent( id="agent-1", @@ -2756,12 +4133,11 @@ def test_roster_update_archive_versions_and_detail(monkeypatch: pytest.MonkeyPat source=AgentSource.AGENT_APP, status=AgentStatus.ACTIVE, ) - version = AgentConfigSnapshot(id="version-2", agent_id="agent-1", version=2, config_snapshot='{"prompt":{}}') - version.created_at = datetime(2026, 1, 4, 3, 4, 5, tzinfo=UTC) - - service = AgentRosterService(fake_session) - monkeypatch.setattr(service, "_get_agent", lambda **kwargs: agent) - monkeypatch.setattr(service, "_get_version", lambda **kwargs: version) + session.add_all([agent, listed_version, older_listed_version, revision, listed_revision]) + session.commit() + service = AgentRosterService(session) + retire_snapshots = MagicMock(return_value=[]) + monkeypatch.setattr(AgentHomeSnapshotService, "retire_all_for_agent", retire_snapshots) monkeypatch.setattr( service, "get_roster_agent_detail", @@ -2780,6 +4156,7 @@ def test_roster_update_archive_versions_and_detail(monkeypatch: pytest.MonkeyPat assert updated["description"] == "new" assert agent.status == AgentStatus.ARCHIVED + retire_snapshots.assert_called_once_with(session=session, tenant_id="tenant-1", agent_id="agent-1") assert versions[0]["id"] == "version-4" assert versions[0]["version"] == 2 assert versions[0]["display_version"] == 2 @@ -2792,21 +4169,76 @@ def test_roster_update_archive_versions_and_detail(monkeypatch: pytest.MonkeyPat assert detail["display_version"] == 1 assert detail["snapshot_version"] == 2 assert detail["config_snapshot"] == {"prompt": {}} - assert detail["created_at"] == int(version.created_at.timestamp()) + assert detail["created_at"] == int(older_listed_version.created_at.timestamp()) assert detail["revisions"][0]["created_at"] == int(revision_created_at.timestamp()) -def test_roster_create_detail_and_lookup_helpers(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession( - scalar=[ - SimpleNamespace(id="agent-1"), - None, - SimpleNamespace(id="version-1"), - None, - ], - scalars=[[AgentConfigSnapshot(id="version-1", agent_id="agent-1", version=1)]], +def test_roster_archive_retires_then_commits_before_enqueue( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + service = AgentRosterService(session) + agent = _agent() + binding = AgentWorkspaceBinding( + id="binding-1", + tenant_id=agent.tenant_id, + app_id="app-1", + workspace_id="workspace-1", + agent_id=agent.id, + agent_config_version_id="version-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + backend_binding_ref="backend-binding-1", ) - service = AgentRosterService(fake_session) + session.add_all([agent, binding]) + session.commit() + events: list[str] = [] + monkeypatch.setattr( + AgentWorkspaceService, + "retire_binding", + MagicMock(side_effect=lambda **_kwargs: events.append("retire-binding") or "binding-1"), + ) + monkeypatch.setattr( + AgentHomeSnapshotService, + "retire_all_for_agent", + MagicMock(side_effect=lambda **_kwargs: events.append("retire-home") or ["home-1"]), + ) + event.listen(session, "after_commit", lambda _session: events.append("commit")) + monkeypatch.setattr( + roster_service, + "enqueue_agent_resource_collection", + MagicMock(side_effect=lambda **_kwargs: events.append("enqueue")), + ) + + service.archive_roster_agent(tenant_id="tenant-1", agent_id="agent-1", account_id="account-1") + + assert events == ["retire-binding", "retire-home", "commit", "enqueue"] + + +def test_roster_archive_commit_failure_does_not_enqueue( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + session = sqlite_session + service = AgentRosterService(session) + session.add(_agent()) + session.commit() + monkeypatch.setattr(AgentHomeSnapshotService, "retire_all_for_agent", MagicMock(return_value=["home-1"])) + event.listen( + session, + "before_commit", + lambda _session: (_ for _ in ()).throw(RuntimeError("commit failed")), + ) + enqueue_collection = MagicMock() + monkeypatch.setattr(roster_service, "enqueue_agent_resource_collection", enqueue_collection) + + with pytest.raises(RuntimeError, match="commit failed"): + service.archive_roster_agent(tenant_id="tenant-1", agent_id="agent-1", account_id="account-1") + + enqueue_collection.assert_not_called() + + +def test_roster_create_detail_and_lookup_helpers(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): + session = sqlite_session + service = AgentRosterService(session) monkeypatch.setattr( AgentRosterService, "_get_or_create_agent_app_debug_conversation", @@ -2831,29 +4263,42 @@ def test_roster_create_detail_and_lookup_helpers(monkeypatch: pytest.MonkeyPatch name="Backing Agent", role="Support agent", ) - found_agent = service._get_agent(tenant_id="tenant-1", agent_id="agent-1") + found_agent = service._get_agent(tenant_id="tenant-1", agent_id=created.id) with pytest.raises(roster_service.AgentNotFoundError): service._get_agent(tenant_id="tenant-1", agent_id="missing") - found_version = service._get_version(tenant_id="tenant-1", agent_id="agent-1", version_id="version-1") + found_version = service._get_version( + tenant_id="tenant-1", + agent_id=created.id, + version_id=created.active_config_snapshot_id, + ) with pytest.raises(roster_service.AgentVersionNotFoundError): - service._get_version(tenant_id="tenant-1", agent_id="agent-1", version_id=None) - loaded_versions = service._load_versions_by_id(["version-1"]) + service._get_version(tenant_id="tenant-1", agent_id=created.id, version_id=None) + loaded_versions = service._load_versions_by_id([created.active_config_snapshot_id]) assert service._load_versions_by_id([]) == {} assert created.name == "Analyst" assert created.role == "Research assistant" assert created.source == AgentSource.ROSTER assert created.active_config_snapshot_id is not None + snapshots = list( + session.scalars( + select(AgentConfigSnapshot) + .where(AgentConfigSnapshot.agent_id.in_([created.id, backing_agent.id])) + .order_by(AgentConfigSnapshot.agent_id) + ) + ) + assert [snapshot.home_snapshot_id for snapshot in snapshots] == [None, None] assert created.active_config_has_model is False assert backing_agent.role == "Support agent" assert backing_agent.active_config_snapshot_id is not None assert backing_agent.active_config_has_model is False - assert found_agent.id == "agent-1" - assert found_version.id == "version-1" - assert loaded_versions["version-1"].agent_id == "agent-1" + assert found_agent.id == created.id + assert found_version.id == created.active_config_snapshot_id + assert loaded_versions[created.active_config_snapshot_id].agent_id == created.id -def test_get_agent_runtime_app_model_creates_hidden_backing_app_for_existing_inline_agent(): +def test_get_agent_runtime_app_model_creates_hidden_backing_app_for_existing_inline_agent(sqlite_session: Session): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -2869,27 +4314,20 @@ def test_get_agent_runtime_app_model_creates_hidden_backing_app_for_existing_inl created_by="account-1", updated_by="account-1", ) - backing_app = App( - id="generated-1", - tenant_id="tenant-1", - name="Inline Agent", - mode=AppMode.AGENT, - status=AppStatus.NORMAL, - ) - session = FakeSession(scalar=[agent, backing_app]) + session.add(agent) + session.commit() service = AgentRosterService(session) resolved_app = service.get_agent_runtime_app_model(tenant_id="tenant-1", agent_id="agent-1") - assert resolved_app is backing_app - assert agent.backing_app_id == "generated-1" - assert session.commits == 1 - created_app = next(value for value in session.added if isinstance(value, App)) - assert created_app.enable_site is False - assert created_app.enable_api is False + assert resolved_app.id == agent.backing_app_id + assert resolved_app.enable_site is False + assert resolved_app.enable_api is False + assert session.get(App, resolved_app.id) is resolved_app -def test_agent_app_debug_conversation_create_reuse_and_recreate(): +def test_agent_app_build_conversation_create_reuse_and_recreate(sqlite_session: Session): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -2901,15 +4339,25 @@ def test_agent_app_debug_conversation_create_reuse_and_recreate(): source=AgentSource.AGENT_APP, status=AgentStatus.ACTIVE, ) + session.add(agent) + session.commit() + service = AgentRosterService(session) - create_session = FakeSession(scalar=[agent, None]) - created_id = AgentRosterService(create_session).get_or_create_agent_app_debug_conversation_id( + created_id = service.get_or_create_build_conversation( tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", ) - created_conversation = next(value for value in create_session.added if isinstance(value, Conversation)) - created_mapping = next(value for value in create_session.added if isinstance(value, AgentDebugConversation)) + created_conversation = session.get(Conversation, created_id) + created_mapping = session.scalar( + select(AgentDebugConversation).where( + AgentDebugConversation.agent_id == agent.id, + AgentDebugConversation.account_id == "account-1", + AgentDebugConversation.draft_type == AgentConfigDraftType.DEBUG_BUILD, + ) + ) + assert created_conversation is not None + assert created_mapping is not None assert created_id == created_mapping.conversation_id assert created_conversation.app_id == "app-1" assert created_conversation.from_account_id == "account-1" @@ -2917,47 +4365,28 @@ def test_agent_app_debug_conversation_create_reuse_and_recreate(): 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( - tenant_id="tenant-1", - 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"]) - reused_id = AgentRosterService(reuse_session).get_or_create_agent_app_debug_conversation_id( + reused_id = service.get_or_create_build_conversation( tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", ) - assert reused_id == "existing-conversation" - assert reuse_session.added == [] - assert reuse_session.commits == 1 + assert reused_id == created_id - stale_mapping = AgentDebugConversation( - tenant_id="tenant-1", - 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]) - recreated_id = AgentRosterService(recreate_session).get_or_create_agent_app_debug_conversation_id( + created_conversation.is_deleted = True + session.commit() + recreated_id = service.get_or_create_build_conversation( tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", ) - assert recreated_id == stale_mapping.conversation_id - assert recreated_id != "deleted-conversation" - assert any(isinstance(value, Conversation) for value in recreate_session.added) - assert recreate_session.commits == 1 + assert recreated_id == created_mapping.conversation_id + assert recreated_id != created_id + assert session.get(Conversation, recreated_id) is not None -def test_agent_app_debug_conversations_are_isolated_by_draft_type(): +def test_agent_app_debug_conversations_are_isolated_by_draft_type(sqlite_session: Session): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -2969,23 +4398,22 @@ def test_agent_app_debug_conversations_are_isolated_by_draft_type(): source=AgentSource.AGENT_APP, status=AgentStatus.ACTIVE, ) - session = FakeSession(scalar=[agent, None, agent, None]) + session.add(agent) + session.commit() service = AgentRosterService(session) - build_conversation_id = service.get_or_create_agent_app_debug_conversation_id( + build_conversation_id = service.get_or_create_build_conversation( 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( + preview_conversation_id = service.rotate_preview_conversation( 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)] + mappings = list(session.scalars(select(AgentDebugConversation).where(AgentDebugConversation.agent_id == agent.id))) assert build_conversation_id != preview_conversation_id assert {mapping.draft_type for mapping in mappings} == { AgentConfigDraftType.DRAFT, @@ -2993,8 +4421,28 @@ def test_agent_app_debug_conversations_are_isolated_by_draft_type(): } -def test_agent_app_debug_conversation_message_count(): - session = FakeSession(scalar=[3]) +def test_agent_app_debug_conversation_message_count(sqlite_session: Session): + session = sqlite_session + for index in range(3): + session.add( + Message( + id=f"message-{index}", + app_id="app-1", + conversation_id="debug-conversation-1", + _inputs={}, + query="q", + message={}, + message_unit_price=0, + answer="a", + answer_unit_price=0, + total_price=0, + currency="USD", + from_source=ConversationFromSource.CONSOLE, + from_account_id="account-1", + app_mode=AppMode.AGENT_CHAT, + ) + ) + session.commit() count = AgentRosterService(session).count_agent_app_debug_conversation_messages( conversation_id="debug-conversation-1", @@ -3003,7 +4451,7 @@ def test_agent_app_debug_conversation_message_count(): assert count == 3 -def test_agent_app_debug_conversation_requires_app_binding(): +def test_agent_app_debug_conversation_requires_app_binding(sqlite_session: Session): agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3016,14 +4464,15 @@ def test_agent_app_debug_conversation_requires_app_binding(): ) with pytest.raises(roster_service.AgentNotFoundError): - AgentRosterService(FakeSession())._get_or_create_agent_app_debug_conversation( + AgentRosterService(sqlite_session)._get_or_create_agent_app_debug_conversation( agent=agent, account_id="account-1", draft_type=AgentConfigDraftType.DEBUG_BUILD, ) -def test_load_or_create_agent_app_debug_conversations_supports_runtime_backed_agents(): +def test_load_or_create_build_conversations_supports_runtime_backed_agents(sqlite_session: Session): + session = sqlite_session valid_agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3053,10 +4502,13 @@ def test_load_or_create_agent_app_debug_conversations_supports_runtime_backed_ag scope=AgentScope.WORKFLOW_ONLY, source=AgentSource.WORKFLOW, status=AgentStatus.ACTIVE, + workflow_id="workflow-1", + workflow_node_id="node-1", ) - fake_session = FakeSession(scalar=[None]) - result = AgentRosterService(fake_session).load_or_create_agent_app_debug_conversation_ids_by_agent_id( + session.add_all([valid_agent, wrong_tenant_agent, workflow_agent]) + session.commit() + result = AgentRosterService(session).load_or_create_build_conversation_ids_by_agent_id( tenant_id="tenant-1", agents=[valid_agent, wrong_tenant_agent, workflow_agent], account_id="account-1", @@ -3065,8 +4517,9 @@ def test_load_or_create_agent_app_debug_conversations_supports_runtime_backed_ag assert list(result) == ["agent-1", "agent-3"] assert result["agent-1"] assert result["agent-3"] - assert fake_session.commits == 1 - mappings = [value for value in fake_session.added if isinstance(value, AgentDebugConversation)] + mappings = list( + session.scalars(select(AgentDebugConversation).where(AgentDebugConversation.tenant_id == "tenant-1")) + ) assert len(mappings) == 2 assert all(mapping.draft_type == AgentConfigDraftType.DEBUG_BUILD for mapping in mappings) @@ -3090,9 +4543,11 @@ def test_agent_app_visible_versions_exclude_draft_saves(): assert AgentConfigRevisionOperation.SAVE_CURRENT_VERSION not in roster_operations -def test_restore_roster_agent_version_switches_active_snapshot(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession(scalar=["version-2", None]) - service = AgentRosterService(fake_session) +def test_restore_roster_agent_version_switches_active_snapshot( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session + service = AgentRosterService(session) agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3110,11 +4565,18 @@ def test_restore_roster_agent_version_switches_active_snapshot(monkeypatch: pyte tenant_id="tenant-1", agent_id="agent-1", version=2, + home_snapshot_id="home-version-2", config_snapshot=_agent_soul_with_model(), ) - - monkeypatch.setattr(service, "_get_agent", lambda **kwargs: agent) - monkeypatch.setattr(service, "_get_version", lambda **kwargs: version) + revision = AgentConfigRevision( + tenant_id="tenant-1", + agent_id=agent.id, + current_snapshot_id=version.id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + ) + session.add_all([agent, version, revision]) + session.commit() restored = service.restore_agent_version( tenant_id="tenant-1", @@ -3126,25 +4588,28 @@ def test_restore_roster_agent_version_switches_active_snapshot(monkeypatch: pyte assert restored == { "result": "success", "active_config_snapshot_id": "version-4", - "draft_config_id": fake_session.added[0].id, + "draft_config_id": restored["draft_config_id"], "restored_version_id": "version-2", } assert agent.active_config_snapshot_id == "version-4" assert agent.active_config_is_published is False assert agent.updated_by == "account-1" - assert fake_session.commits == 1 - draft = fake_session.added[0] + draft = session.get(AgentConfigDraft, restored["draft_config_id"]) + assert draft is not None assert draft.tenant_id == "tenant-1" assert draft.agent_id == "agent-1" assert draft.draft_type == AgentConfigDraftType.DRAFT assert draft.base_snapshot_id == "version-2" + assert draft.home_snapshot_id == "home-version-2" assert draft.config_snapshot_dict == _agent_soul_with_model().model_dump(mode="json") assert draft.updated_by == "account-1" -def test_restore_roster_agent_version_rejects_invisible_versions(monkeypatch: pytest.MonkeyPatch): - fake_session = FakeSession(scalar=[None]) - service = AgentRosterService(fake_session) +def test_restore_roster_agent_version_rejects_invisible_versions( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session + service = AgentRosterService(session) agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3156,8 +4621,15 @@ def test_restore_roster_agent_version_rejects_invisible_versions(monkeypatch: py status=AgentStatus.ACTIVE, active_config_snapshot_id="version-4", ) - - monkeypatch.setattr(service, "_get_agent", lambda **kwargs: agent) + version = AgentConfigSnapshot( + id="version-2", + tenant_id="tenant-1", + agent_id=agent.id, + version=2, + config_snapshot=_agent_soul_with_model(), + ) + session.add_all([agent, version]) + session.commit() with pytest.raises(roster_service.AgentVersionNotFoundError): service.restore_agent_version( @@ -3168,23 +4640,18 @@ def test_restore_roster_agent_version_rejects_invisible_versions(monkeypatch: py ) assert agent.active_config_snapshot_id == "version-4" - assert fake_session.added == [] - assert fake_session.commits == 0 + assert session.scalar(select(AgentConfigDraft).where(AgentConfigDraft.agent_id == agent.id)) is None -def test_app_list_all_excludes_agent_apps_by_default(): - filters = AppService._build_app_list_filters( - "account-1", "tenant-1", AppListParams(mode="all"), FakeSession(scalar=None, scalars=None) - ) +def test_app_list_all_excludes_agent_apps_by_default(sqlite_session: Session): + filters = AppService._build_app_list_filters("account-1", "tenant-1", AppListParams(mode="all"), sqlite_session) sql = " ".join(str(filter_) for filter_ in filters) assert "apps.mode != :mode_1" in sql -def test_app_list_agent_mode_requires_visible_roster_backing_agent(): - filters = AppService._build_app_list_filters( - "account-1", "tenant-1", AppListParams(mode="agent"), FakeSession(scalar=None, scalars=None) - ) +def test_app_list_agent_mode_requires_visible_roster_backing_agent(sqlite_session: Session): + filters = AppService._build_app_list_filters("account-1", "tenant-1", AppListParams(mode="agent"), sqlite_session) sql = " ".join(str(filter_) for filter_ in filters) assert "EXISTS" in sql @@ -3384,10 +4851,12 @@ class TestAgentAppBackingAgent: ``Agent.app_id``. ``AppService.create_app`` builds the backing agent inside its own transaction, so the helper must add+flush without committing.""" - def test_create_backing_agent_for_app_links_app_and_seeds_default_soul(self): - session = FakeSession() + def test_create_backing_agent_for_app_links_app_and_seeds_default_soul( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + ): + session = sqlite_session + transaction = session.begin() service = AgentRosterService(session) - agent = service.create_backing_agent_for_app( tenant_id="tenant-1", account_id="account-1", @@ -3406,22 +4875,30 @@ class TestAgentAppBackingAgent: assert agent.name == "Iris" assert agent.role == "research assistant" # A v1 snapshot + revision are seeded and wired as the active version. - snapshots = [a for a in session.added if isinstance(a, AgentConfigSnapshot)] + snapshots = list(session.scalars(select(AgentConfigSnapshot).where(AgentConfigSnapshot.agent_id == agent.id))) assert len(snapshots) == 1 assert snapshots[0].version == 1 + assert snapshots[0].home_snapshot_id is None assert agent.active_config_snapshot_id == snapshots[0].id - revisions = [ - a for a in session.added if getattr(a, "operation", None) == AgentConfigRevisionOperation.CREATE_VERSION - ] + revisions = list( + session.scalars( + select(AgentConfigRevision).where( + AgentConfigRevision.agent_id == agent.id, + AgentConfigRevision.operation == AgentConfigRevisionOperation.CREATE_VERSION, + ) + ) + ) assert len(revisions) == 1 - conversations = [a for a in session.added if isinstance(a, Conversation)] + conversations = list(session.scalars(select(Conversation).where(Conversation.app_id == "app-1"))) assert len(conversations) == 1 assert conversations[0].app_id == "app-1" assert conversations[0].mode == "agent" assert conversations[0].status == ConversationStatus.NORMAL assert conversations[0].from_source == ConversationFromSource.CONSOLE assert conversations[0].from_account_id == "account-1" - debug_mappings = [a for a in session.added if isinstance(a, AgentDebugConversation)] + debug_mappings = list( + session.scalars(select(AgentDebugConversation).where(AgentDebugConversation.agent_id == agent.id)) + ) assert len(debug_mappings) == 1 assert debug_mappings[0].tenant_id == "tenant-1" assert debug_mappings[0].agent_id == agent.id @@ -3429,24 +4906,27 @@ class TestAgentAppBackingAgent: assert debug_mappings[0].account_id == "account-1" assert debug_mappings[0].conversation_id == conversations[0].id # Caller (AppService.create_app) owns the commit — helper must not commit. - assert session.commits == 0 + assert transaction.is_active - def test_get_app_backing_agent_queries_active_agent_app_agent(self): - sentinel = SimpleNamespace(id="agent-1", app_id="app-1") - session = FakeSession(scalar=[sentinel]) + def test_get_app_backing_agent_queries_active_agent_app_agent(self, sqlite_session: Session): + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") + session.add(agent) + session.commit() service = AgentRosterService(session) result = service.get_app_backing_agent(tenant_id="tenant-1", app_id="app-1") - assert result is sentinel + assert result is agent - def test_get_app_backing_agent_returns_none_when_unbound(self): - session = FakeSession() + def test_get_app_backing_agent_returns_none_when_unbound(self, sqlite_session: Session): + session = sqlite_session service = AgentRosterService(session) assert service.get_app_backing_agent(tenant_id="tenant-1", app_id="app-x") is None - def test_get_agent_app_model_resolves_app_backing_agent(self): + def test_get_agent_app_model_resolves_app_backing_agent(self, sqlite_session: Session): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3458,20 +4938,22 @@ class TestAgentAppBackingAgent: status=AgentStatus.ACTIVE, app_id="app-1", ) - app = SimpleNamespace(id="app-1", mode="agent", status="normal") - session = FakeSession(scalar=[agent, app]) + app = _app(mode=AppMode.AGENT) + session.add_all([agent, app]) + session.commit() service = AgentRosterService(session) assert service.get_agent_app_model(tenant_id="tenant-1", agent_id="agent-1") is app - def test_get_agent_app_model_rejects_unbound_agent(self): - session = FakeSession() + def test_get_agent_app_model_rejects_unbound_agent(self, sqlite_session: Session): + session = sqlite_session service = AgentRosterService(session) with pytest.raises(roster_service.AgentNotFoundError): service.get_agent_app_model(tenant_id="tenant-1", agent_id="agent-x") - def test_refresh_agent_app_debug_conversation_creates_mapping(self): + def test_reset_build_conversation_creates_mapping(self, sqlite_session: Session): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3483,22 +4965,25 @@ class TestAgentAppBackingAgent: status=AgentStatus.ACTIVE, app_id="app-1", ) - session = FakeSession(scalar=[agent, None]) + session.add(agent) + session.commit() service = AgentRosterService(session) - conversation_id = service.refresh_agent_app_debug_conversation_id( + conversation_id = service.reset_build_conversation( tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", ) - conversations = [a for a in session.added if isinstance(a, Conversation)] + conversations = list(session.scalars(select(Conversation).where(Conversation.id == conversation_id))) assert len(conversations) == 1 assert conversations[0].id == conversation_id assert conversations[0].app_id == "app-1" assert conversations[0].from_source == ConversationFromSource.CONSOLE assert conversations[0].from_account_id == "account-1" - mappings = [a for a in session.added if isinstance(a, AgentDebugConversation)] + mappings = list( + session.scalars(select(AgentDebugConversation).where(AgentDebugConversation.agent_id == agent.id)) + ) assert len(mappings) == 1 assert mappings[0].tenant_id == "tenant-1" assert mappings[0].agent_id == "agent-1" @@ -3506,43 +4991,11 @@ class TestAgentAppBackingAgent: 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 - def test_refresh_agent_app_debug_conversation_replaces_existing_mapping(self): - agent = Agent( - id="agent-1", - tenant_id="tenant-1", - name="Iris", - description="", - agent_kind=AgentKind.DIFY_AGENT, - scope=AgentScope.ROSTER, - source=AgentSource.AGENT_APP, - status=AgentStatus.ACTIVE, - app_id="app-1", - ) - mapping = SimpleNamespace(app_id="old-app", conversation_id="old-conversation") - session = FakeSession(scalar=[agent, mapping]) - service = AgentRosterService(session) - - conversation_id = service.refresh_agent_app_debug_conversation_id( - tenant_id="tenant-1", - agent_id="agent-1", - account_id="account-1", - ) - - assert mapping.app_id == "app-1" - assert mapping.conversation_id == conversation_id - assert [a for a in session.added if isinstance(a, AgentDebugConversation)] == [] - conversations = [a for a in session.added if isinstance(a, Conversation)] - assert len(conversations) == 1 - assert conversations[0].id == conversation_id - assert session.deleted == [] - assert session.commits == 1 - - def test_refresh_agent_app_debug_conversation_rotates_preview_mapping_each_time( - self, monkeypatch: pytest.MonkeyPatch + def test_rotate_preview_conversation_retires_exact_binding( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3554,125 +5007,58 @@ class TestAgentAppBackingAgent: status=AgentStatus.ACTIVE, app_id="app-1", ) - session = FakeSession(scalar=[agent, None]) - service = AgentRosterService(session) - cleanup = MagicMock() - monkeypatch.setattr(service, "_cleanup_debug_conversation_runtime_sessions", cleanup) - - first_conversation_id = service.refresh_agent_app_debug_conversation_id( - tenant_id="tenant-1", - agent_id="agent-1", - account_id="account-1", - draft_type=AgentConfigDraftType.DRAFT, - ) - mapping = next(value for value in session.added if isinstance(value, AgentDebugConversation)) - session._scalar.extend([agent, mapping]) - second_conversation_id = service.refresh_agent_app_debug_conversation_id( - tenant_id="tenant-1", - agent_id="agent-1", - account_id="account-1", - draft_type=AgentConfigDraftType.DRAFT, - ) - - assert first_conversation_id != second_conversation_id - assert mapping.draft_type == AgentConfigDraftType.DRAFT - assert mapping.conversation_id == second_conversation_id - assert len([value for value in session.added if isinstance(value, Conversation)]) == 2 - cleanup.assert_called_once_with( - tenant_id="tenant-1", - agent_id="agent-1", - account_id="account-1", - draft_type=AgentConfigDraftType.DRAFT, - app_id="app-1", - conversation_id=first_conversation_id, - ) - - @pytest.mark.parametrize( - "draft_type", - [AgentConfigDraftType.DRAFT, AgentConfigDraftType.DEBUG_BUILD], - ) - def test_refresh_agent_app_debug_conversation_enqueues_cleanup_for_old_runtime_sessions( - self, - monkeypatch: pytest.MonkeyPatch, - draft_type: AgentConfigDraftType, - ): - agent = Agent( - id="agent-1", - tenant_id="tenant-1", - name="Iris", - description="", - agent_kind=AgentKind.DIFY_AGENT, - scope=AgentScope.ROSTER, - source=AgentSource.AGENT_APP, - status=AgentStatus.ACTIVE, - app_id="app-1", - ) - mapping = SimpleNamespace(app_id="old-app", conversation_id="old-conversation") - stored_session = SimpleNamespace( - scope=SimpleNamespace( - tenant_id="tenant-1", - app_id="old-app", - conversation_id="old-conversation", - agent_id="agent-9", - agent_config_snapshot_id="snap-9", - ), - session_snapshot=CompositorSessionSnapshot(layers=[]), - backend_run_id="run-old", - runtime_layer_specs=[RuntimeLayerSpec(name="history", type="pydantic_ai.history")], - ) - session = FakeSession(scalar=[agent, mapping]) - service = AgentRosterService(session) - events: list[str] = [] - original_commit = session.commit - - def record_commit() -> None: - events.append("commit") - original_commit() - - def list_active_sessions(**kwargs: object) -> list[object]: - events.append("cleanup") - return [stored_session] - - cleanup_delay = MagicMock() - cleanup_store = MagicMock() - cleanup_store.list_active_sessions_for_conversation.side_effect = list_active_sessions - monkeypatch.setattr(session, "commit", record_commit) - monkeypatch.setattr(roster_service, "AgentAppRuntimeSessionStore", lambda: cleanup_store) - monkeypatch.setattr(roster_service.cleanup_conversation_agent_runtime_session, "delay", cleanup_delay) - - conversation_id = service.refresh_agent_app_debug_conversation_id( - tenant_id="tenant-1", - agent_id="agent-1", - account_id="account-1", - draft_type=draft_type, - ) - - cleanup_store.list_active_sessions_for_conversation.assert_called_once_with( - tenant_id="tenant-1", + mapping = AgentDebugConversation( + tenant_id=agent.tenant_id, + agent_id=agent.id, app_id="old-app", + account_id="account-1", + draft_type=AgentConfigDraftType.DRAFT, conversation_id="old-conversation", ) - cleanup_delay.assert_called_once() - 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"] == draft_type.value - assert ( - payload["idempotency_key"] - == f"tenant-1:agent-1:account-1:{draft_type.value}:old-conversation:debug-session-cleanup:" - "agent-9:snap-9:run-old" + previous_conversation = _conversation(conversation_id="old-conversation") + previous_conversation.app_id = "old-app" + previous_conversation.agent_workspace_binding_id = "binding-1" + binding = SimpleNamespace(agent_id=agent.id) + session.add_all([agent, mapping, previous_conversation]) + session.commit() + service = AgentRosterService(session) + events: list[str] = [] + event.listen(session, "after_commit", lambda _session: events.append("commit")) + get_active_binding = MagicMock(return_value=binding) + retire_binding = MagicMock(side_effect=lambda **_kwargs: events.append("retire") or "binding-1") + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active_binding) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + monkeypatch.setattr( + roster_service, + "enqueue_agent_resource_collection", + MagicMock(side_effect=lambda **_kwargs: events.append("enqueue")), ) - cleanup_store.mark_cleaned.assert_called_once_with( - scope=stored_session.scope, - backend_run_id="run-old", + + conversation_id = service.rotate_preview_conversation( + tenant_id="tenant-1", + agent_id="agent-1", + account_id="account-1", ) + assert mapping.app_id == "app-1" assert mapping.conversation_id == conversation_id - assert events == ["commit", "cleanup"] + conversations = list(session.scalars(select(Conversation).where(Conversation.id == conversation_id))) + assert len(conversations) == 1 + assert conversations[0].id == conversation_id + assert events == ["retire", "commit", "enqueue"] + owner_scope = get_active_binding.call_args.kwargs["expected_owner_scope"] + assert owner_scope.owner_type is AgentWorkspaceOwnerType.CONVERSATION + assert owner_scope.owner_id == previous_conversation.id + retire_binding.assert_called_once_with( + session=session, + tenant_id="tenant-1", + binding_id="binding-1", + ) - def test_refresh_agent_app_debug_conversation_does_not_cleanup_when_commit_fails( - self, monkeypatch: pytest.MonkeyPatch + def test_reset_build_conversation_retires_build_draft_binding( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3684,29 +5070,109 @@ class TestAgentAppBackingAgent: status=AgentStatus.ACTIVE, app_id="app-1", ) - mapping = SimpleNamespace(app_id="old-app", conversation_id="old-conversation") - session = FakeSession(scalar=[agent, mapping]) - service = AgentRosterService(session) - runtime_store_factory = MagicMock() - cleanup_delay = MagicMock() - monkeypatch.setattr(session, "commit", MagicMock(side_effect=RuntimeError("database unavailable"))) - monkeypatch.setattr(roster_service, "AgentAppRuntimeSessionStore", runtime_store_factory) - monkeypatch.setattr(roster_service.cleanup_conversation_agent_runtime_session, "delay", cleanup_delay) + mapping = AgentDebugConversation( + tenant_id=agent.tenant_id, + agent_id=agent.id, + app_id="app-1", + account_id="account-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + conversation_id="old-build-conversation", + ) + build_draft = AgentConfigDraft( + id="build-draft-1", + tenant_id=agent.tenant_id, + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + session.add_all([agent, mapping, build_draft]) + session.commit() + get_active_binding = MagicMock(return_value=SimpleNamespace(agent_id=agent.id)) + retire_binding = MagicMock(return_value="binding-1") + enqueue_collection = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active_binding) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + monkeypatch.setattr(roster_service, "enqueue_agent_resource_collection", enqueue_collection) - with pytest.raises(RuntimeError, match="database unavailable"): - service.refresh_agent_app_debug_conversation_id( + conversation_id = AgentRosterService(session).reset_build_conversation( + tenant_id="tenant-1", + agent_id=agent.id, + account_id="account-1", + ) + + assert mapping.conversation_id == conversation_id + assert build_draft.agent_workspace_binding_id is None + owner_scope = get_active_binding.call_args.kwargs["expected_owner_scope"] + assert owner_scope.owner_type is AgentWorkspaceOwnerType.BUILD_DRAFT + assert owner_scope.owner_id == build_draft.id + retire_binding.assert_called_once_with( + session=session, + tenant_id="tenant-1", + binding_id="binding-1", + ) + enqueue_collection.assert_called_once_with( + tenant_id="tenant-1", + binding_ids=("binding-1",), + ) + + def test_preview_rotation_commit_failure_rolls_back_before_enqueue( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + ): + session = sqlite_session + agent = Agent( + id="agent-1", + tenant_id="tenant-1", + name="Iris", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + app_id="app-1", + ) + mapping = AgentDebugConversation( + tenant_id=agent.tenant_id, + agent_id=agent.id, + app_id="app-1", + account_id="account-1", + draft_type=AgentConfigDraftType.DRAFT, + conversation_id="old-conversation", + ) + previous_conversation = _conversation(conversation_id="old-conversation") + previous_conversation.agent_workspace_binding_id = "binding-1" + session.add_all([agent, mapping, previous_conversation]) + session.commit() + event.listen( + session, + "before_commit", + lambda _session: (_ for _ in ()).throw(RuntimeError("commit failed")), + ) + monkeypatch.setattr( + AgentWorkspaceService, + "get_active_binding", + MagicMock(return_value=SimpleNamespace(agent_id=agent.id)), + ) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", MagicMock(return_value="binding-1")) + enqueue_collection = MagicMock() + monkeypatch.setattr(roster_service, "enqueue_agent_resource_collection", enqueue_collection) + + with pytest.raises(RuntimeError, match="commit failed"): + AgentRosterService(session).rotate_preview_conversation( tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", - draft_type=AgentConfigDraftType.DRAFT, ) - runtime_store_factory.assert_not_called() - cleanup_delay.assert_not_called() + assert session.get(AgentDebugConversation, mapping.id).conversation_id == "old-conversation" + enqueue_collection.assert_not_called() - def test_refresh_agent_app_debug_conversation_marks_old_runtime_sessions_clean_when_enqueue_fails( - self, monkeypatch: pytest.MonkeyPatch + def test_build_reset_commit_failure_rolls_back_before_enqueue( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ): + session = sqlite_session agent = Agent( id="agent-1", tenant_id="tenant-1", @@ -3718,85 +5184,66 @@ class TestAgentAppBackingAgent: status=AgentStatus.ACTIVE, app_id="app-1", ) - mapping = SimpleNamespace(app_id="old-app", conversation_id="old-conversation") - stored_session = SimpleNamespace( - scope=SimpleNamespace( - tenant_id="tenant-1", - app_id="old-app", - conversation_id="old-conversation", - agent_id="agent-9", - agent_config_snapshot_id="snap-9", - ), - session_snapshot=CompositorSessionSnapshot(layers=[]), - backend_run_id="run-old", - runtime_layer_specs=[RuntimeLayerSpec(name="history", type="pydantic_ai.history")], - ) - session = FakeSession(scalar=[agent, mapping]) - service = AgentRosterService(session) - cleanup_store = MagicMock() - cleanup_store.list_active_sessions_for_conversation.return_value = [stored_session] - monkeypatch.setattr(roster_service, "AgentAppRuntimeSessionStore", lambda: cleanup_store) - monkeypatch.setattr( - roster_service.cleanup_conversation_agent_runtime_session, - "delay", - MagicMock(side_effect=RuntimeError("queue down")), - ) - - service.refresh_agent_app_debug_conversation_id( - tenant_id="tenant-1", - agent_id="agent-1", - account_id="account-1", - ) - - cleanup_store.mark_cleaned.assert_called_once_with( - scope=stored_session.scope, - backend_run_id="run-old", - ) - - def test_refresh_agent_app_debug_conversation_ignores_mark_cleaned_failure(self, monkeypatch: pytest.MonkeyPatch): - agent = Agent( - id="agent-1", - tenant_id="tenant-1", - name="Iris", - description="", - agent_kind=AgentKind.DIFY_AGENT, - scope=AgentScope.ROSTER, - source=AgentSource.AGENT_APP, - status=AgentStatus.ACTIVE, + mapping = AgentDebugConversation( + tenant_id=agent.tenant_id, + agent_id=agent.id, app_id="app-1", - ) - mapping = SimpleNamespace(app_id="old-app", conversation_id="old-conversation") - stored_session = SimpleNamespace( - scope=SimpleNamespace( - tenant_id="tenant-1", - app_id="old-app", - conversation_id="old-conversation", - agent_id="agent-9", - agent_config_snapshot_id="snap-9", - ), - session_snapshot=CompositorSessionSnapshot(layers=[]), - backend_run_id="run-old", - runtime_layer_specs=[RuntimeLayerSpec(name="history", type="pydantic_ai.history")], - ) - session = FakeSession(scalar=[agent, mapping]) - service = AgentRosterService(session) - cleanup_store = MagicMock() - cleanup_store.list_active_sessions_for_conversation.return_value = [stored_session] - cleanup_store.mark_cleaned.side_effect = RuntimeError("cleanup bookkeeping failed") - monkeypatch.setattr(roster_service, "AgentAppRuntimeSessionStore", lambda: cleanup_store) - monkeypatch.setattr(roster_service.cleanup_conversation_agent_runtime_session, "delay", MagicMock()) - - conversation_id = service.refresh_agent_app_debug_conversation_id( - tenant_id="tenant-1", - agent_id="agent-1", account_id="account-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + conversation_id="old-build-conversation", ) + build_draft = AgentConfigDraft( + id="build-draft-1", + tenant_id=agent.tenant_id, + agent_id=agent.id, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + session.add_all([agent, mapping, build_draft]) + session.commit() + events: list[str] = [] - assert mapping.app_id == "app-1" - assert mapping.conversation_id == conversation_id - assert session.commits == 1 + def fail_commit(_session: Session) -> None: + assert build_draft.agent_workspace_binding_id is None + events.append("commit") + raise RuntimeError("commit failed") - def test_duplicate_agent_app_copies_app_config_and_active_soul(self, monkeypatch: pytest.MonkeyPatch): + def rollback(_session: Session) -> None: + events.append("rollback") + + event.listen(session, "before_commit", fail_commit) + event.listen(session, "after_rollback", rollback) + get_active_binding = MagicMock(return_value=SimpleNamespace(agent_id=agent.id)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_active_binding) + retire_binding = MagicMock(side_effect=lambda **_kwargs: events.append("retire") or "binding-1") + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + enqueue_collection = MagicMock() + monkeypatch.setattr(roster_service, "enqueue_agent_resource_collection", enqueue_collection) + + with pytest.raises(RuntimeError, match="commit failed"): + AgentRosterService(session).reset_build_conversation( + tenant_id="tenant-1", + agent_id=agent.id, + account_id="account-1", + ) + + assert events == ["retire", "commit", "rollback"] + owner_scope = get_active_binding.call_args.kwargs["expected_owner_scope"] + assert owner_scope.owner_type is AgentWorkspaceOwnerType.BUILD_DRAFT + assert owner_scope.owner_id == build_draft.id + retire_binding.assert_called_once_with( + session=session, + tenant_id="tenant-1", + binding_id="binding-1", + ) + enqueue_collection.assert_not_called() + + def test_duplicate_agent_app_copies_app_config_and_active_soul( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + ): source_config = SimpleNamespace( opening_statement="hello", suggested_questions='["q1"]', @@ -3841,8 +5288,8 @@ class TestAgentAppBackingAgent: id="target-app", app_model_config=target_config, app_model_config_with_session=lambda *, session: target_config, - enable_site=True, - enable_api=True, + enable_site=False, + enable_api=False, use_icon_as_answer_icon=False, tracing=None, ) @@ -3859,42 +5306,21 @@ class TestAgentAppBackingAgent: app_id="source-app", active_config_snapshot_id="source-version", active_config_has_model=True, - ) - target_agent = Agent( - id="target-agent", - tenant_id="tenant-1", - name="Iris copy", - description="source desc", - role="Analyst", - agent_kind=AgentKind.DIFY_AGENT, - scope=AgentScope.ROSTER, - source=AgentSource.AGENT_APP, - status=AgentStatus.ACTIVE, - app_id="target-app", - active_config_snapshot_id="target-version", + active_config_is_published=True, ) source_version = AgentConfigSnapshot( id="source-version", tenant_id="tenant-1", agent_id="source-agent", version=1, + home_snapshot_id="home-source", config_snapshot=_agent_soul_with_model(), summary="configured", version_note="v1", created_by="account-1", ) - target_version = AgentConfigSnapshot( - id="target-version", - tenant_id="tenant-1", - agent_id="target-agent", - version=1, - config_snapshot=AgentSoulConfig(), - created_by="account-1", - ) - session = FakeSession( - scalar=[source_agent, source_app, source_agent, target_agent, source_version, target_version], - scalars=[[]], - ) + session = sqlite_session + service = AgentRosterService(session) captured: dict[str, object] = {} class FakeAppService: @@ -3902,9 +5328,46 @@ class TestAgentAppBackingAgent: captured["tenant_id"] = tenant_id captured["params"] = params captured["account"] = account + target_agent = AgentRosterService(session).create_backing_agent_for_app( + tenant_id=tenant_id, + account_id="account-1", + app_id=target_app.id, + name=params.name, + description=params.description, + role=params.agent_role, + ) + target_version = session.scalar( + select(AgentConfigSnapshot).where(AgentConfigSnapshot.agent_id == target_agent.id) + ) + assert target_version is not None + captured["target_agent"] = target_agent + captured["target_version"] = target_version return target_app monkeypatch.setattr(roster_service, "AppService", FakeAppService) + monkeypatch.setattr( + AgentRosterService, + "_get_or_create_agent_app_debug_conversation", + lambda _self, **_kwargs: None, + ) + monkeypatch.setattr(service, "get_agent_app_model", lambda **_kwargs: source_app) + monkeypatch.setattr( + service, + "get_app_backing_agent", + lambda *, tenant_id, app_id: ( + source_agent + if app_id == source_app.id + else session.scalar(select(Agent).where(Agent.tenant_id == tenant_id, Agent.app_id == app_id)) + ), + ) + monkeypatch.setattr( + service, + "_get_version", + lambda *, version_id, **_kwargs: ( + source_version if version_id == source_version.id else session.get(AgentConfigSnapshot, version_id) + ), + ) + monkeypatch.setattr(service, "_next_duplicate_agent_name", lambda **_kwargs: "Iris copy") monkeypatch.setattr( roster_service.FeatureService, "get_system_features", @@ -3912,7 +5375,7 @@ class TestAgentAppBackingAgent: ) account = SimpleNamespace(id="account-1") - duplicated = AgentRosterService(session).duplicate_agent_app( + duplicated = service.duplicate_agent_app( tenant_id="tenant-1", agent_id="source-agent", account=account, @@ -3924,20 +5387,27 @@ class TestAgentAppBackingAgent: assert params.mode == "agent" assert params.agent_role == "Analyst" assert target_app.enable_site is False - assert target_app.enable_api is True + assert target_app.enable_api is False assert target_app.use_icon_as_answer_icon is True assert target_app.tracing == "{}" assert target_config.opening_statement == "hello" assert target_config.file_upload == '{"image": {"enabled": true}}' assert target_config.updated_by == "account-1" + target_agent = captured["target_agent"] + target_version = captured["target_version"] assert target_version.config_snapshot.model.model == "gpt-4o" + assert source_version.home_snapshot_id == "home-source" + assert target_version.home_snapshot_id is None assert target_version.summary == "configured" assert target_version.version_note == "v1" assert target_agent.active_config_has_model is True + assert target_agent.active_config_is_published is False assert target_agent.updated_by == "account-1" - assert session.commits == 1 + assert session.get(Agent, target_agent.id) is target_agent - def test_duplicate_agent_app_inherits_webapp_access_mode(self, monkeypatch: pytest.MonkeyPatch): + def test_duplicate_agent_app_inherits_webapp_access_mode( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + ): source_app = SimpleNamespace( id="source-app", tenant_id="tenant-1", @@ -3956,7 +5426,7 @@ class TestAgentAppBackingAgent: ) source_agent = SimpleNamespace(id="source-agent", role="Analyst") target_app = SimpleNamespace(id="target-app") - session = FakeSession() + session = sqlite_session service = AgentRosterService(session) monkeypatch.setattr(service, "get_agent_app_model", lambda **_: source_app) monkeypatch.setattr(service, "get_app_backing_agent", lambda **_: source_agent) @@ -4001,7 +5471,9 @@ class TestAgentAppBackingAgent: assert captured["params"].agent_role == "Custom Analyst" assert access_mode_updates == [("target-app", "private")] - def test_duplicate_agent_app_falls_back_to_public_access_mode(self, monkeypatch: pytest.MonkeyPatch): + def test_duplicate_agent_app_falls_back_to_public_access_mode( + self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + ): source_app = SimpleNamespace( id="source-app", tenant_id="tenant-1", @@ -4020,7 +5492,7 @@ class TestAgentAppBackingAgent: ) source_agent = SimpleNamespace(id="source-agent", role="Analyst") target_app = SimpleNamespace(id="target-app") - session = FakeSession() + session = sqlite_session service = AgentRosterService(session) monkeypatch.setattr(service, "get_agent_app_model", lambda **_: source_app) monkeypatch.setattr(service, "get_app_backing_agent", lambda **_: source_agent) @@ -4066,25 +5538,50 @@ class TestAgentAppBackingAgent: class TestListWorkflowsReferencingAppAgent: - def test_groups_bindings_by_workflow_app_and_sorts_by_name(self): - agent = SimpleNamespace(id="agent-1") + def test_groups_bindings_by_workflow_app_and_sorts_by_name(self, sqlite_session: Session): + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") bindings = [ - SimpleNamespace( - agent_id="agent-1", app_id="wf-app-1", workflow_id="wf-1", workflow_version="v1", node_id="node-b" + WorkflowAgentNodeBinding( + tenant_id="tenant-1", + agent_id="agent-1", + app_id="wf-app-1", + workflow_id="wf-1", + workflow_version="v1", + node_id="node-b", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), - SimpleNamespace( - agent_id="agent-1", app_id="wf-app-1", workflow_id="wf-1", workflow_version="v1", node_id="node-a" + WorkflowAgentNodeBinding( + tenant_id="tenant-1", + agent_id="agent-1", + app_id="wf-app-1", + workflow_id="wf-1", + workflow_version="v1", + node_id="node-a", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), - SimpleNamespace( - agent_id="agent-1", app_id="wf-app-2", workflow_id="wf-2", workflow_version="v2", node_id="node-a" + WorkflowAgentNodeBinding( + tenant_id="tenant-1", + agent_id="agent-1", + app_id="wf-app-2", + workflow_id="wf-2", + workflow_version="v2", + node_id="node-a", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), ] - apps = [ - SimpleNamespace(id="wf-app-1", name="Beta Flow", mode="workflow", workflow_id="wf-1"), - SimpleNamespace(id="wf-app-2", name="Alpha Flow", mode="advanced-chat", workflow_id="wf-2"), - ] - # scalar -> backing agent; scalars -> bindings, then resolved apps. - session = FakeSession(scalar=[agent], scalars=[bindings, apps]) + beta_app = _app(app_id="wf-app-1", name="Beta Flow", mode=AppMode.WORKFLOW) + beta_app.workflow_id = "wf-1" + alpha_app = _app(app_id="wf-app-2", name="Alpha Flow", mode=AppMode.ADVANCED_CHAT) + alpha_app.workflow_id = "wf-2" + session.add_all([agent, *bindings, beta_app, alpha_app]) + session.commit() service = AgentRosterService(session) result = service.list_workflows_referencing_app_agent(tenant_id="tenant-1", app_id="app-1") @@ -4095,43 +5592,71 @@ class TestListWorkflowsReferencingAppAgent: assert beta["workflow_id"] == "wf-1" assert beta["workflow_version"] == "v1" - def test_returns_empty_when_no_backing_agent(self): - session = FakeSession() # scalar() -> None + def test_returns_empty_when_no_backing_agent(self, sqlite_session: Session): + session = sqlite_session # scalar() -> None service = AgentRosterService(session) assert service.list_workflows_referencing_app_agent(tenant_id="tenant-1", app_id="app-x") == [] - def test_returns_empty_when_no_bindings(self): - agent = SimpleNamespace(id="agent-1") - session = FakeSession(scalar=[agent], scalars=[[]]) + def test_returns_empty_when_no_bindings(self, sqlite_session: Session): + session = sqlite_session + session.add(_agent(source=AgentSource.AGENT_APP, app_id="app-1")) + session.commit() service = AgentRosterService(session) assert service.list_workflows_referencing_app_agent(tenant_id="tenant-1", app_id="app-1") == [] - def test_skips_orphaned_binding_whose_app_is_gone(self): - agent = SimpleNamespace(id="agent-1") - bindings = [ - SimpleNamespace( - agent_id="agent-1", app_id="wf-app-gone", workflow_id="wf-9", workflow_version="v9", node_id="node-a" - ) - ] - session = FakeSession(scalar=[agent], scalars=[bindings, []]) # no apps resolved + def test_skips_orphaned_binding_whose_app_is_gone(self, sqlite_session: Session): + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") + binding = WorkflowAgentNodeBinding( + tenant_id="tenant-1", + agent_id=agent.id, + app_id="wf-app-gone", + workflow_id="wf-9", + workflow_version="v9", + node_id="node-a", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), + ) + session.add_all([agent, binding]) + session.commit() service = AgentRosterService(session) assert service.list_workflows_referencing_app_agent(tenant_id="tenant-1", app_id="app-1") == [] - def test_skips_historical_published_workflow_versions(self): - agent = SimpleNamespace(id="agent-1") + def test_skips_historical_published_workflow_versions(self, sqlite_session: Session): + session = sqlite_session + agent = _agent(source=AgentSource.AGENT_APP, app_id="app-1") bindings = [ - SimpleNamespace( - agent_id="agent-1", app_id="wf-app-1", workflow_id="old-wf", workflow_version="old", node_id="old" + WorkflowAgentNodeBinding( + tenant_id="tenant-1", + agent_id=agent.id, + app_id="wf-app-1", + workflow_id="old-wf", + workflow_version="old", + node_id="old", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), - SimpleNamespace( - agent_id="agent-1", app_id="wf-app-1", workflow_id="current-wf", workflow_version="v2", node_id="new" + WorkflowAgentNodeBinding( + tenant_id="tenant-1", + agent_id=agent.id, + app_id="wf-app-1", + workflow_id="current-wf", + workflow_version="v2", + node_id="new", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + current_snapshot_id="version-1", + node_job_config=WorkflowNodeJobConfig(), ), ] - apps = [SimpleNamespace(id="wf-app-1", name="Flow", mode="workflow", workflow_id="current-wf")] - session = FakeSession(scalar=[agent], scalars=[bindings, apps]) + app = _app(app_id="wf-app-1", name="Flow", mode=AppMode.WORKFLOW) + app.workflow_id = "current-wf" + session.add_all([agent, *bindings, app]) + session.commit() service = AgentRosterService(session) result = service.list_workflows_referencing_app_agent(tenant_id="tenant-1", app_id="app-1") @@ -4143,18 +5668,14 @@ class TestListWorkflowsReferencingAppAgent: class TestWorkflowAgentDraftBindingSync: def _agent_workflow(self) -> Workflow: - return Workflow( - id="workflow-1", - tenant_id="tenant-1", - app_id="app-1", - version=Workflow.VERSION_DRAFT, - graph=json.dumps( - { - "nodes": [{"id": "agent-node", "data": {"type": "agent", "version": "2"}}], - "edges": [], - } - ), + workflow = _workflow() + workflow.graph = json.dumps( + { + "nodes": [{"id": "agent-node", "data": {"type": "agent", "version": "2"}}], + "edges": [], + } ) + return workflow def _agent_binding(self) -> WorkflowAgentNodeBinding: return WorkflowAgentNodeBinding( @@ -4175,6 +5696,9 @@ class TestWorkflowAgentDraftBindingSync: id="agent-1", tenant_id="tenant-1", name="Iris", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.ROSTER, status=AgentStatus.ACTIVE, active_config_snapshot_id="snapshot-1", ) @@ -4188,35 +5712,42 @@ class TestWorkflowAgentDraftBindingSync: config_snapshot=agent_soul, ) + def _publish_revision(self, snapshot_id: str = "snapshot-2") -> AgentConfigRevision: + return AgentConfigRevision( + id="revision-1", + tenant_id="tenant-1", + agent_id="agent-1", + current_snapshot_id=snapshot_id, + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + created_by="account-1", + ) + def _sync_roster_agent_task_refs( self, *, agent_task: str, existing_ref_selectors: list[list[str]] | None = None, + sqlite_session: Session, ) -> WorkflowNodeJobConfig: - workflow = Workflow( - id="workflow-1", - tenant_id="tenant-1", - app_id="app-1", - version=Workflow.VERSION_DRAFT, - graph=json.dumps( - { - "nodes": [ - { - "id": "agent-node", - "data": { - "type": "agent", - "version": "2", - "agent_task": agent_task, - "agent_binding": { - "binding_type": "roster_agent", - "agent_id": "agent-1", - }, + workflow = _workflow() + workflow.graph = json.dumps( + { + "nodes": [ + { + "id": "agent-node", + "data": { + "type": "agent", + "version": "2", + "agent_task": agent_task, + "agent_binding": { + "binding_type": "roster_agent", + "agent_id": "agent-1", }, - } - ] - } - ), + }, + } + ] + } ) agent = Agent( id="agent-1", @@ -4228,10 +5759,9 @@ class TestWorkflowAgentDraftBindingSync: status=AgentStatus.ACTIVE, active_config_snapshot_id="snapshot-2", ) + session = sqlite_session existing_binding = None - if existing_ref_selectors is None: - session = FakeSession(scalar=[agent], scalars=[[]]) - else: + if existing_ref_selectors is not None: existing_binding = WorkflowAgentNodeBinding( id="binding-1", tenant_id="tenant-1", @@ -4249,7 +5779,10 @@ class TestWorkflowAgentDraftBindingSync: } ), ) - session = FakeSession(scalar=[agent], scalars=[[existing_binding]]) + session.add_all([agent, self._publish_revision()]) + if existing_binding is not None: + session.add(existing_binding) + session.commit() WorkflowAgentPublishService.sync_roster_agent_bindings_for_draft( session=session, @@ -4257,10 +5790,14 @@ class TestWorkflowAgentDraftBindingSync: account_id="account-1", ) - binding = existing_binding or next(item for item in session.added if isinstance(item, WorkflowAgentNodeBinding)) + binding = existing_binding or session.scalar( + select(WorkflowAgentNodeBinding).where(WorkflowAgentNodeBinding.node_id == "agent-node") + ) + assert binding is not None return WorkflowNodeJobConfig.model_validate(binding.node_job_config_dict) - def test_publish_validation_rejects_agent_soul_publish_only_errors(self): + def test_publish_validation_rejects_agent_soul_publish_only_errors(self, sqlite_session: Session): + session = sqlite_session binding = self._agent_binding() agent_soul = AgentSoulConfig.model_validate( { @@ -4275,7 +5812,8 @@ class TestWorkflowAgentDraftBindingSync: ) agent = self._publish_agent() snapshot = self._snapshot(agent_soul) - session = FakeSession(scalar=[binding, agent, snapshot, agent, snapshot], scalars=[[binding]]) + session.add_all([binding, agent, snapshot]) + session.commit() with pytest.raises(InvalidComposerConfigError, match="human_involvement_not_referenced"): WorkflowAgentPublishService.validate_agent_nodes_for_publish( @@ -4283,7 +5821,8 @@ class TestWorkflowAgentDraftBindingSync: draft_workflow=self._agent_workflow(), ) - def test_publish_validation_rejects_dangling_agent_soul_drive_refs(self): + def test_publish_validation_rejects_dangling_agent_soul_drive_refs(self, sqlite_session: Session): + session = sqlite_session binding = self._agent_binding() agent_soul = AgentSoulConfig.model_validate( { @@ -4297,7 +5836,8 @@ class TestWorkflowAgentDraftBindingSync: ) agent = self._publish_agent() snapshot = self._snapshot(agent_soul) - session = FakeSession(scalar=[binding, agent, snapshot, agent, snapshot], scalars=[[binding], []]) + session.add_all([binding, agent, snapshot]) + session.commit() with pytest.raises(WorkflowAgentNodeValidationError, match="skill_ref_dangling"): WorkflowAgentPublishService.validate_agent_nodes_for_publish( @@ -4327,7 +5867,8 @@ class TestWorkflowAgentDraftBindingSync: with pytest.raises(InvalidComposerConfigError, match="config_asset_missing.*skill:research.*file:guide.txt"): ComposerConfigValidator.validate_publish_payload(payload) - def test_projects_binding_declared_outputs_to_draft_graph_response(self): + def test_projects_binding_declared_outputs_to_draft_graph_response(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4384,7 +5925,8 @@ class TestWorkflowAgentDraftBindingSync: ], ), ) - session = FakeSession(scalars=[[binding]]) + session.add(binding) + session.commit() graph = WorkflowAgentPublishService.project_draft_bindings_to_graph( session=session, @@ -4406,7 +5948,8 @@ class TestWorkflowAgentDraftBindingSync: assert profile_output["children"][1]["array_item"]["children"][0]["name"] == "city" assert "agent_declared_outputs" not in workflow.graph_dict["nodes"][0]["data"] - def test_projects_inline_binding_over_pending_inline_graph_response(self): + def test_projects_inline_binding_over_pending_inline_graph_response(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4440,8 +5983,10 @@ class TestWorkflowAgentDraftBindingSync: binding_type=WorkflowAgentBindingType.INLINE_AGENT, agent_id="inline-agent-1", current_snapshot_id="inline-snapshot-1", + node_job_config=WorkflowNodeJobConfig(), ) - session = FakeSession(scalars=[[binding]]) + session.add(binding) + session.commit() graph = WorkflowAgentPublishService.project_draft_bindings_to_graph( session=session, @@ -4457,7 +6002,8 @@ class TestWorkflowAgentDraftBindingSync: "binding_type": "inline_agent", } - def test_keeps_pending_inline_graph_response_over_existing_roster_binding(self): + def test_keeps_pending_inline_graph_response_over_existing_roster_binding(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4491,8 +6037,10 @@ class TestWorkflowAgentDraftBindingSync: binding_type=WorkflowAgentBindingType.ROSTER_AGENT, agent_id="agent-1", current_snapshot_id="snapshot-1", + node_job_config=WorkflowNodeJobConfig(), ) - session = FakeSession(scalars=[[binding]]) + session.add(binding) + session.commit() graph = WorkflowAgentPublishService.project_draft_bindings_to_graph( session=session, @@ -4503,7 +6051,8 @@ class TestWorkflowAgentDraftBindingSync: "binding_type": "inline_agent", } - def test_creates_roster_binding_from_agent_node_graph(self): + def test_creates_roster_binding_from_agent_node_graph(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4545,7 +6094,8 @@ class TestWorkflowAgentDraftBindingSync: status=AgentStatus.ACTIVE, active_config_snapshot_id="snapshot-2", ) - session = FakeSession(scalar=[agent], scalars=[[]]) + session.add_all([agent, self._publish_revision()]) + session.commit() WorkflowAgentPublishService.sync_roster_agent_bindings_for_draft( session=session, @@ -4553,7 +6103,10 @@ class TestWorkflowAgentDraftBindingSync: account_id="account-1", ) - binding = next(item for item in session.added if isinstance(item, WorkflowAgentNodeBinding)) + binding = session.scalar( + select(WorkflowAgentNodeBinding).where(WorkflowAgentNodeBinding.node_id == "agent-node") + ) + assert binding is not None assert binding.binding_type == WorkflowAgentBindingType.ROSTER_AGENT assert binding.agent_id == "agent-1" assert binding.current_snapshot_id == "snapshot-2" @@ -4568,33 +6121,37 @@ class TestWorkflowAgentDraftBindingSync: ], ).model_dump(mode="json") - def test_creates_roster_binding_deriving_previous_node_refs_from_agent_task(self): + def test_creates_roster_binding_deriving_previous_node_refs_from_agent_task(self, sqlite_session: Session): node_job = self._sync_roster_agent_task_refs( agent_task="Review {{#previous-node.report#}} for {{#sys.query#}}.", + sqlite_session=sqlite_session, ) assert node_job.workflow_prompt == "Review {{#previous-node.report#}} for {{#sys.query#}}." assert [ref.selector for ref in node_job.previous_node_output_refs] == [["previous-node", "report"]] - def test_updates_existing_roster_binding_clearing_legacy_only_previous_node_refs(self): + def test_updates_existing_roster_binding_clearing_legacy_only_previous_node_refs(self, sqlite_session: Session): node_job = self._sync_roster_agent_task_refs( agent_task="Review [§node_output:previous-node.report:PREV/report§].", existing_ref_selectors=[["previous-node", "report"]], + sqlite_session=sqlite_session, ) assert node_job.workflow_prompt == "Review [§node_output:previous-node.report:PREV/report§]." assert node_job.previous_node_output_refs == [] - def test_updates_existing_roster_binding_clearing_stale_previous_node_refs(self): + def test_updates_existing_roster_binding_clearing_stale_previous_node_refs(self, sqlite_session: Session): node_job = self._sync_roster_agent_task_refs( agent_task="Review the current request without upstream context.", existing_ref_selectors=[["previous-node", "report"]], + sqlite_session=sqlite_session, ) assert node_job.workflow_prompt == "Review the current request without upstream context." assert node_job.previous_node_output_refs == [] - def test_creates_inline_binding_from_agent_node_graph(self): + def test_creates_inline_binding_from_agent_node_graph(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4640,7 +6197,8 @@ class TestWorkflowAgentDraftBindingSync: version=1, config_snapshot=AgentSoulConfig(), ) - session = FakeSession(scalar=[agent, snapshot], scalars=[[]]) + session.add_all([agent, snapshot]) + session.commit() WorkflowAgentPublishService.sync_agent_bindings_for_draft( session=session, @@ -4648,7 +6206,10 @@ class TestWorkflowAgentDraftBindingSync: account_id="account-1", ) - binding = next(item for item in session.added if isinstance(item, WorkflowAgentNodeBinding)) + binding = session.scalar( + select(WorkflowAgentNodeBinding).where(WorkflowAgentNodeBinding.node_id == "agent-node") + ) + assert binding is not None assert binding.binding_type == WorkflowAgentBindingType.INLINE_AGENT assert binding.agent_id == "inline-agent-1" assert binding.current_snapshot_id == "inline-snapshot-1" @@ -4656,7 +6217,8 @@ class TestWorkflowAgentDraftBindingSync: workflow_prompt="Use the current node context.", ).model_dump(mode="json") - def test_keeps_pending_inline_binding_in_draft_graph_without_db_binding(self): + def test_keeps_pending_inline_binding_in_draft_graph_without_db_binding(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4689,8 +6251,10 @@ class TestWorkflowAgentDraftBindingSync: binding_type=WorkflowAgentBindingType.ROSTER_AGENT, agent_id="agent-1", current_snapshot_id="snapshot-1", + node_job_config=WorkflowNodeJobConfig(), ) - session = FakeSession(scalars=[[existing_binding]]) + session.add(existing_binding) + session.commit() WorkflowAgentPublishService.sync_agent_bindings_for_draft( session=session, @@ -4698,11 +6262,11 @@ class TestWorkflowAgentDraftBindingSync: account_id="account-1", ) - assert session.deleted == [] - assert session.added == [] - assert session.flushes == 1 + assert session.get(WorkflowAgentNodeBinding, existing_binding.id) is existing_binding + assert not session.new - def test_clones_inline_binding_for_agent_owned_by_another_node(self, monkeypatch): + def test_clones_inline_binding_for_agent_owned_by_another_node(self, monkeypatch, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4740,7 +6304,8 @@ class TestWorkflowAgentDraftBindingSync: status=AgentStatus.ACTIVE, active_config_snapshot_id="inline-snapshot-1", ) - session = FakeSession(scalar=[agent], scalars=[[]]) + session.add(agent) + session.commit() clone = MagicMock(return_value=(SimpleNamespace(id="cloned-agent"), "cloned-snapshot")) monkeypatch.setattr(WorkflowAgentPublishService, "_clone_inline_graph_binding_for_node", clone) @@ -4751,10 +6316,15 @@ class TestWorkflowAgentDraftBindingSync: ) clone.assert_called_once() - assert session.added[0].agent_id == "cloned-agent" - assert session.added[0].current_snapshot_id == "cloned-snapshot" + binding = session.scalar( + select(WorkflowAgentNodeBinding).where(WorkflowAgentNodeBinding.node_id == "agent-node") + ) + assert binding is not None + assert binding.agent_id == "cloned-agent" + assert binding.current_snapshot_id == "cloned-snapshot" - def test_rejects_agent_node_graph_binding_with_unsupported_type(self): + def test_rejects_agent_node_graph_binding_with_unsupported_type(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4781,12 +6351,13 @@ class TestWorkflowAgentDraftBindingSync: with pytest.raises(ValueError, match="unsupported agent_binding type"): WorkflowAgentPublishService.sync_agent_bindings_for_draft( - session=FakeSession(scalars=[[]]), + session=session, draft_workflow=workflow, account_id="account-1", ) - def test_treats_partial_inline_binding_as_pending_draft_state(self): + def test_treats_partial_inline_binding_as_pending_draft_state(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4811,19 +6382,16 @@ class TestWorkflowAgentDraftBindingSync: ), ) - session = FakeSession(scalars=[[]]) - WorkflowAgentPublishService.sync_agent_bindings_for_draft( session=session, draft_workflow=workflow, account_id="account-1", ) - assert session.added == [] - assert session.deleted == [] - assert session.flushes == 1 + assert session.scalar(select(WorkflowAgentNodeBinding)) is None - def test_rejects_inline_binding_with_missing_snapshot(self): + def test_rejects_inline_binding_with_missing_snapshot(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4861,15 +6429,18 @@ class TestWorkflowAgentDraftBindingSync: status=AgentStatus.ACTIVE, active_config_snapshot_id="inline-snapshot-1", ) + session.add(agent) + session.commit() with pytest.raises(ValueError, match="missing inline agent config snapshot"): WorkflowAgentPublishService.sync_agent_bindings_for_draft( - session=FakeSession(scalar=[agent, None], scalars=[[]]), + session=session, draft_workflow=workflow, account_id="account-1", ) - def test_updates_existing_roster_binding_prompt_from_agent_node_graph(self): + def test_updates_existing_roster_binding_prompt_from_agent_node_graph(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4921,7 +6492,8 @@ class TestWorkflowAgentDraftBindingSync: ], ), ) - session = FakeSession(scalar=[agent], scalars=[[existing_binding]]) + session.add_all([agent, self._publish_revision(), existing_binding]) + session.commit() WorkflowAgentPublishService.sync_roster_agent_bindings_for_draft( session=session, @@ -4934,7 +6506,8 @@ class TestWorkflowAgentDraftBindingSync: assert [output.name for output in node_job.declared_outputs] == ["summary"] assert existing_binding.current_snapshot_id == "snapshot-2" - def test_updates_existing_roster_binding_declared_outputs_from_agent_node_graph(self): + def test_updates_existing_roster_binding_declared_outputs_from_agent_node_graph(self, sqlite_session: Session): + session = sqlite_session workflow = Workflow( id="workflow-1", tenant_id="tenant-1", @@ -4991,7 +6564,8 @@ class TestWorkflowAgentDraftBindingSync: ], ), ) - session = FakeSession(scalar=[agent], scalars=[[existing_binding]]) + session.add_all([agent, self._publish_revision(), existing_binding]) + session.commit() WorkflowAgentPublishService.sync_roster_agent_bindings_for_draft( session=session, @@ -5004,38 +6578,213 @@ class TestWorkflowAgentDraftBindingSync: assert node_job.declared_outputs == [] assert existing_binding.current_snapshot_id == "snapshot-2" - def test_deletes_draft_binding_when_agent_node_removed(self): + @pytest.mark.parametrize( + "sqlite_session", + [(Agent, AgentConfigRevision, AgentConfigSnapshot, WorkflowAgentNodeBinding)], + indirect=True, + ) + def test_deletes_draft_binding_and_returns_only_replaced_inline_agent( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + ) -> None: workflow = Workflow( id="workflow-1", tenant_id="tenant-1", app_id="app-1", version=Workflow.VERSION_DRAFT, - graph='{"nodes":[]}', + graph=json.dumps( + { + "nodes": [ + { + "id": "kept-node", + "data": { + "type": "agent", + "version": "2", + "agent_binding": { + "binding_type": "inline_agent", + "agent_id": "inline-kept", + "current_snapshot_id": "snapshot-kept", + }, + }, + }, + { + "id": "roster-node", + "data": { + "type": "agent", + "version": "2", + "agent_binding": { + "binding_type": "roster_agent", + "agent_id": "roster-new", + }, + }, + }, + { + "id": "inline-replaced-node", + "data": { + "type": "agent", + "version": "2", + "agent_binding": { + "binding_type": "inline_agent", + "agent_id": "inline-new", + "current_snapshot_id": "snapshot-inline-new", + }, + }, + }, + ] + } + ), ) - stale_binding = WorkflowAgentNodeBinding( - id="binding-1", + removed_inline = WorkflowAgentNodeBinding( + id="binding-inline-removed", tenant_id="tenant-1", app_id="app-1", workflow_id="workflow-1", workflow_version=Workflow.VERSION_DRAFT, node_id="removed-node", - binding_type=WorkflowAgentBindingType.ROSTER_AGENT, - agent_id="agent-1", - current_snapshot_id="snapshot-1", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id="inline-removed", + current_snapshot_id="snapshot-removed", node_job_config=WorkflowNodeJobConfig(), ) - session = FakeSession(scalars=[[stale_binding]]) - - WorkflowAgentPublishService.sync_roster_agent_bindings_for_draft( - session=session, + kept_inline = WorkflowAgentNodeBinding( + id="binding-inline-kept", + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_version=Workflow.VERSION_DRAFT, + node_id="kept-node", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id="inline-kept", + current_snapshot_id="snapshot-kept", + node_job_config=WorkflowNodeJobConfig(), + ) + removed_roster = WorkflowAgentNodeBinding( + id="binding-roster-removed", + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_version=Workflow.VERSION_DRAFT, + node_id="removed-roster-node", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + agent_id="roster-removed", + current_snapshot_id="snapshot-roster", + node_job_config=WorkflowNodeJobConfig(), + ) + old_inline_replaced_by_roster = WorkflowAgentNodeBinding( + id="binding-inline-to-roster", + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_version=Workflow.VERSION_DRAFT, + node_id="roster-node", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id="inline-old-roster", + current_snapshot_id="snapshot-inline-old-roster", + node_job_config=WorkflowNodeJobConfig(), + ) + old_inline_replaced_by_inline = WorkflowAgentNodeBinding( + id="binding-inline-to-inline", + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_version=Workflow.VERSION_DRAFT, + node_id="inline-replaced-node", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id="inline-old-inline", + current_snapshot_id="snapshot-inline-old-inline", + node_job_config=WorkflowNodeJobConfig(), + ) + kept_agent = Agent( + id="inline-kept", + tenant_id="tenant-1", + name="Kept inline", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.WORKFLOW_ONLY, + source=AgentSource.WORKFLOW, + status=AgentStatus.ACTIVE, + app_id="app-1", + workflow_id="workflow-1", + workflow_node_id="kept-node", + active_config_snapshot_id="snapshot-kept", + ) + replacement_agent = Agent( + id="inline-new", + tenant_id="tenant-1", + name="Replacement inline", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.WORKFLOW_ONLY, + source=AgentSource.WORKFLOW, + status=AgentStatus.ACTIVE, + app_id="app-1", + workflow_id="workflow-1", + workflow_node_id="inline-replaced-node", + active_config_snapshot_id="snapshot-inline-new", + ) + roster_agent = Agent( + id="roster-new", + tenant_id="tenant-1", + name="Roster replacement", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.ROSTER, + status=AgentStatus.ACTIVE, + active_config_snapshot_id="snapshot-roster-new", + ) + kept_snapshot = AgentConfigSnapshot( + id="snapshot-kept", + tenant_id="tenant-1", + agent_id="inline-kept", + version=1, + config_snapshot=AgentSoulConfig(), + home_snapshot_id="home-kept", + ) + replacement_snapshot = AgentConfigSnapshot( + id="snapshot-inline-new", + tenant_id="tenant-1", + agent_id="inline-new", + version=1, + config_snapshot=AgentSoulConfig(), + home_snapshot_id="home-inline-new", + ) + sqlite_session.add_all( + [ + kept_agent, + replacement_agent, + roster_agent, + kept_snapshot, + replacement_snapshot, + removed_inline, + kept_inline, + removed_roster, + old_inline_replaced_by_roster, + old_inline_replaced_by_inline, + ] + ) + sqlite_session.commit() + retirement_candidates = WorkflowAgentPublishService.sync_agent_bindings_for_draft( + session=sqlite_session, draft_workflow=workflow, account_id="account-1", ) - assert session.deleted == [stale_binding] + assert sqlite_session.get(WorkflowAgentNodeBinding, removed_inline.id) is None + assert sqlite_session.get(WorkflowAgentNodeBinding, removed_roster.id) is None + assert kept_inline.binding_type == WorkflowAgentBindingType.INLINE_AGENT + assert kept_inline.agent_id == "inline-kept" + assert old_inline_replaced_by_roster.binding_type == WorkflowAgentBindingType.ROSTER_AGENT + assert old_inline_replaced_by_roster.agent_id == "roster-new" + assert old_inline_replaced_by_roster.current_snapshot_id == "snapshot-roster-new" + assert old_inline_replaced_by_inline.binding_type == WorkflowAgentBindingType.INLINE_AGENT + assert old_inline_replaced_by_inline.agent_id == "inline-new" + assert old_inline_replaced_by_inline.current_snapshot_id == "snapshot-inline-new" + assert retirement_candidates == {"inline-removed", "inline-old-roster", "inline-old-inline"} -def test_dataset_rows_filters_malformed_ids(monkeypatch: pytest.MonkeyPatch): +def test_dataset_rows_filters_malformed_ids(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): """Mention ids are user-editable text: a non-UUID id must read as missing (placeholder semantics), never reach the UUID-typed dataset query (E2E 500).""" captured = {} @@ -5051,18 +6800,20 @@ def test_dataset_rows_filters_malformed_ids(monkeypatch: pytest.MonkeyPatch): valid = "550e8400-e29b-41d4-a716-446655440000" rows = get_tenant_knowledge_dataset_rows( - session=FakeSession(), tenant_id="tenant-1", dataset_ids=["9999dead-beef", valid] + session=sqlite_session, tenant_id="tenant-1", dataset_ids=["9999dead-beef", valid] ) assert rows == {} assert captured["ids"] == [valid] # all-malformed input never touches the DB captured.clear() - assert get_tenant_knowledge_dataset_rows(session=FakeSession(), tenant_id="tenant-1", dataset_ids=["nope"]) == {} + assert get_tenant_knowledge_dataset_rows(session=sqlite_session, tenant_id="tenant-1", dataset_ids=["nope"]) == {} assert captured == {} -def test_composer_save_rejects_malformed_knowledge_dataset_ids(monkeypatch: pytest.MonkeyPatch): +def test_composer_save_rejects_malformed_knowledge_dataset_ids( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): captured = {"calls": 0} def fake_get_datasets_by_ids(ids, tenant_id, *, session): @@ -5093,13 +6844,15 @@ def test_composer_save_rejects_malformed_knowledge_dataset_ids(monkeypatch: pyte with pytest.raises(InvalidComposerConfigError, match="not-a-uuid"): AgentComposerService.validate_knowledge_datasets( - session=FakeSession(), tenant_id="tenant-1", agent_soul=agent_soul + session=sqlite_session, tenant_id="tenant-1", agent_soul=agent_soul ) assert captured == {"calls": 0} -def test_composer_save_rejects_missing_or_out_of_scope_knowledge_datasets(monkeypatch: pytest.MonkeyPatch): +def test_composer_save_rejects_missing_or_out_of_scope_knowledge_datasets( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): captured = {} missing_dataset_id = "550e8400-e29b-41d4-a716-446655440000" @@ -5130,24 +6883,44 @@ def test_composer_save_rejects_missing_or_out_of_scope_knowledge_datasets(monkey with pytest.raises(InvalidComposerConfigError, match=missing_dataset_id): AgentComposerService.validate_knowledge_datasets( - session=FakeSession(), tenant_id="tenant-1", agent_soul=agent_soul + session=sqlite_session, tenant_id="tenant-1", agent_soul=agent_soul ) assert captured == {"ids": [missing_dataset_id], "tenant_id": "tenant-1"} -def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pytest.MonkeyPatch): - agent = SimpleNamespace( - id="agent-1", - tenant_id="tenant-1", +def test_save_agent_composer_allows_incomplete_knowledge_draft( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session + agent = _agent( source=AgentSource.AGENT_APP, - active_config_snapshot_id="version-1", - active_config_is_published=True, - updated_by=None, + app_id="app-1", + ) + agent.active_config_snapshot_id = "version-1" + agent.active_config_is_published = True + revision = AgentConfigRevision( + id="revision-1", + tenant_id="tenant-1", + agent_id=agent.id, + current_snapshot_id="version-1", + revision=1, + operation=AgentConfigRevisionOperation.PUBLISH_DRAFT, + created_by="account-1", + ) + session.add_all([agent, revision]) + session.commit() + active_version = SimpleNamespace( + home_snapshot_id="home-initial", config_snapshot_dict=AgentSoulConfig().model_dump(mode="json") ) - active_version = SimpleNamespace(config_snapshot_dict=AgentSoulConfig().model_dump(mode="json")) - fake_session = FakeSession(scalar=[agent]) saved = {} + flushes = 0 + + def count_flush(_session: Session, _flush_context: object) -> None: + nonlocal flushes + flushes += 1 + + event.listen(session, "after_flush", count_flush) import services.dataset_service as dataset_service_module @@ -5159,7 +6932,7 @@ def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pyte monkeypatch.setattr( AgentComposerService, "_save_agent_draft", - lambda **kwargs: saved.update(kwargs) or SimpleNamespace(id="draft-1"), + lambda **kwargs: saved.update(kwargs) or SimpleNamespace(id="draft-1", home_snapshot_id="home-initial"), ) monkeypatch.setattr(AgentComposerService, "_get_version_if_present", lambda **_kwargs: active_version) monkeypatch.setattr(AgentComposerService, "load_agent_composer", lambda **_kwargs: {"loaded": True}) @@ -5186,7 +6959,7 @@ def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pyte ) result = AgentComposerService.save_agent_composer( - session=fake_session, + session=session, tenant_id="tenant-1", agent_id="agent-1", account_id="account-1", @@ -5199,7 +6972,8 @@ def test_save_agent_composer_allows_incomplete_knowledge_draft(monkeypatch: pyte assert saved["agent_soul"].knowledge.sets[0].retrieval.model is None assert saved["agent_soul"].knowledge.sets[0].metadata_filtering.mode == "automatic" assert saved["agent_soul"].knowledge.sets[0].metadata_filtering.metadata_model_config is None - assert fake_session.flushes == 1 + assert flushes >= 1 + assert not session.dirty def test_workspace_dify_tools_returns_provider_and_tool_granularities(monkeypatch: pytest.MonkeyPatch): @@ -5259,20 +7033,27 @@ def _drive_soul(**overrides): return AgentSoulConfig.model_validate(base) -def _patch_drive_keys(monkeypatch, existing_keys): - - captured: dict[str, object] = {} - - def fake_scalars(stmt): - captured["stmt"] = stmt - return list(existing_keys) - - captured["session"] = type("S", (), {"scalars": staticmethod(fake_scalars)})() - return captured +def _session_with_drive_keys(sqlite_session: Session, existing_keys: list[str]) -> Session: + session = sqlite_session + session.add_all( + [ + AgentDriveFile( + id=f"drive-file-{index}", + tenant_id="tenant-1", + agent_id="agent-1", + key=key, + file_kind=AgentDriveFileKind.UPLOAD_FILE, + file_id=f"upload-{index}", + ) + for index, key in enumerate(existing_keys, start=1) + ] + ) + session.commit() + return session -def test_drive_mention_findings_reports_missing_keys(monkeypatch: pytest.MonkeyPatch): - session = _patch_drive_keys(monkeypatch, existing_keys=["tender-analyzer/SKILL.md"])["session"] +def test_drive_mention_findings_reports_missing_keys(sqlite_session: Session): + session = _session_with_drive_keys(sqlite_session, ["tender-analyzer/SKILL.md"]) findings = AgentComposerService._drive_mention_findings( session=session, @@ -5286,8 +7067,11 @@ def test_drive_mention_findings_reports_missing_keys(monkeypatch: pytest.MonkeyP assert str(findings[0]["message"]).startswith("file 'sample.pdf' has no drive entry") -def test_drive_mention_findings_clean_when_all_keys_exist(monkeypatch: pytest.MonkeyPatch): - session = _patch_drive_keys(monkeypatch, existing_keys=["tender-analyzer/SKILL.md", "files/sample.pdf"])["session"] +def test_drive_mention_findings_clean_when_all_keys_exist(sqlite_session: Session): + session = _session_with_drive_keys( + sqlite_session, + ["tender-analyzer/SKILL.md", "files/sample.pdf"], + ) assert ( AgentComposerService._drive_mention_findings( @@ -5300,8 +7084,8 @@ def test_drive_mention_findings_clean_when_all_keys_exist(monkeypatch: pytest.Mo ) -def test_drive_mention_findings_skips_prompt_without_drive_mentions(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_drive_mention_findings_skips_prompt_without_drive_mentions(sqlite_session: Session): + session = sqlite_session # No drive-backed mention at all -> no DB roundtrip, no findings. soul = _drive_soul(prompt={"system_prompt": "Use [§knowledge:kb-1:Docs§]."}) findings = AgentComposerService._drive_mention_findings( @@ -5314,11 +7098,11 @@ def test_drive_mention_findings_skips_prompt_without_drive_mentions(monkeypatch: def test_collect_validation_findings_appends_drive_mention_findings_with_agent_context( - monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ): from services.entities.agent_entities import ComposerSavePayload - session = _patch_drive_keys(monkeypatch, existing_keys=[])["session"] + session = _session_with_drive_keys(sqlite_session, []) payload = ComposerSavePayload.model_validate( { "variant": "agent_app", @@ -5347,15 +7131,24 @@ def test_collect_validation_findings_appends_drive_mention_findings_with_agent_c # ── ENG-623/625: resolver helpers + save-path drive guard ──────────────────── -def test_resolve_bound_agent_id_queries_active_roster_agent(monkeypatch: pytest.MonkeyPatch): - from types import SimpleNamespace - - session = SimpleNamespace(scalar=lambda stmt: "agent-9") +def test_resolve_bound_agent_id_queries_active_roster_agent(sqlite_session: Session): + session = sqlite_session + session.add( + _agent( + agent_id="agent-9", + tenant_id="t-1", + source=AgentSource.ROSTER, + app_id="app-1", + ) + ) + session.commit() assert AgentComposerService.resolve_bound_agent_id(session=session, tenant_id="t-1", app_id="app-1") == "agent-9" -def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(monkeypatch: pytest.MonkeyPatch): - session = FakeSession() +def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): + session = sqlite_session from types import SimpleNamespace def boom(cls, **kwargs): @@ -5387,7 +7180,9 @@ def test_resolve_workflow_node_agent_id_degrades_without_workflow_or_binding(mon ) -def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only(monkeypatch: pytest.MonkeyPatch): +def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): payload = ComposerSavePayload.model_validate( { "variant": "workflow", @@ -5406,7 +7201,7 @@ def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only( agent_id="agent-1", current_snapshot_id="version-1", ) - session = FakeSession() + session = sqlite_session monkeypatch.setattr( AgentComposerService, "_get_draft_workflow", classmethod(lambda cls, **kwargs: SimpleNamespace(id="wf-1")) ) @@ -5450,7 +7245,9 @@ def test_save_workflow_composer_reports_drive_mentions_for_inline_node_job_only( assert guarded == {"tenant_id": "t-1", "agent_id": "agent-1"} -def test_save_workflow_composer_reports_drive_mentions_for_roster_node_job_only(monkeypatch: pytest.MonkeyPatch): +def test_save_workflow_composer_reports_drive_mentions_for_roster_node_job_only( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +): payload = ComposerSavePayload.model_validate( { "variant": "workflow", @@ -5469,7 +7266,7 @@ def test_save_workflow_composer_reports_drive_mentions_for_roster_node_job_only( agent_id="agent-1", current_snapshot_id="version-1", ) - session = FakeSession() + session = sqlite_session monkeypatch.setattr( AgentComposerService, "_get_draft_workflow", classmethod(lambda cls, **kwargs: SimpleNamespace(id="wf-1")) ) diff --git a/api/tests/unit_tests/services/agent/test_home_snapshot_service.py b/api/tests/unit_tests/services/agent/test_home_snapshot_service.py new file mode 100644 index 00000000000..2f7e7aceccc --- /dev/null +++ b/api/tests/unit_tests/services/agent/test_home_snapshot_service.py @@ -0,0 +1,186 @@ +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from sqlalchemy.orm import Session + +from models.agent import ( + Agent, + AgentConfigDraft, + AgentConfigDraftType, + AgentConfigSnapshot, + AgentHomeSnapshot, + AgentWorkingResourceStatus, +) +from models.agent_config_entities import AgentSoulConfig +from services.agent.errors import AgentBuildSandboxNotFoundError +from services.agent.home_snapshot_service import AgentHomeSnapshotService, validate_home_snapshot_binding +from services.agent.workspace_service import AgentWorkspaceService + + +def _build_draft(*, home_snapshot_id: str | None = "home-old") -> AgentConfigDraft: + return AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + home_snapshot_id=home_snapshot_id, + agent_workspace_binding_id="binding-1", + config_snapshot=AgentSoulConfig(), + ) + + +def _client(*, snapshot_ref: str = "snapshot-ref-1") -> MagicMock: + client = MagicMock() + client.create_home_snapshot_from_binding_sync.return_value = SimpleNamespace(snapshot_ref=snapshot_ref) + return client + + +def test_validate_home_snapshot_binding_accepts_default_home_without_ledger_lookup() -> None: + session = MagicMock() + validate_home_snapshot_binding( + session=session, + agent=Agent(id="agent-1"), + home_snapshot_id=None, + ) + + session.scalar.assert_not_called() + + +@pytest.mark.parametrize( + ("app_id", "backing_app_id", "expected_runtime_app_id"), + [ + ("app-1", None, "app-1"), + ("workflow-app-1", "runtime-app-1", "runtime-app-1"), + ], +) +def test_build_apply_checkpoints_exact_active_binding( + monkeypatch: pytest.MonkeyPatch, + app_id: str, + backing_app_id: str | None, + expected_runtime_app_id: str, +) -> None: + session = MagicMock() + session.scalar.return_value = SimpleNamespace(app_id=app_id, backing_app_id=backing_app_id) + binding = SimpleNamespace( + backend_binding_ref="binding-ref-1", + agent_id="agent-1", + base_home_snapshot_id="home-old", + agent_config_version_id="build-1", + agent_config_version_kind="build_draft", + ) + get_binding = MagicMock(return_value=binding) + client = _client(snapshot_ref="snapshot-ref-2") + monkeypatch.setattr(AgentHomeSnapshotService, "_client", lambda: nullcontext(client)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_binding) + validate_generation = MagicMock() + monkeypatch.setattr(AgentWorkspaceService, "validate_binding_generation", validate_generation) + + snapshot = AgentHomeSnapshotService.create_for_build_apply( + session=session, + build_draft=_build_draft(), + ) + + assert get_binding.call_args.kwargs["binding_id"] == "binding-1" + assert get_binding.call_args.kwargs["expected_owner_scope"].app_id == expected_runtime_app_id + request = client.create_home_snapshot_from_binding_sync.call_args.args[0] + assert request.backend_binding_ref == "binding-ref-1" + assert snapshot.snapshot_ref == "snapshot-ref-2" + assert validate_generation.call_args.kwargs["base_home_snapshot_id"] == "home-old" + + +def test_build_apply_forwards_default_home_generation(monkeypatch: pytest.MonkeyPatch) -> None: + session = MagicMock() + session.scalar.return_value = SimpleNamespace(app_id="app-1", backing_app_id=None) + binding = SimpleNamespace( + backend_binding_ref="binding-ref-1", + agent_id="agent-1", + base_home_snapshot_id=None, + agent_config_version_id="build-1", + agent_config_version_kind="build_draft", + ) + client = _client(snapshot_ref="snapshot-ref-2") + validate_generation = MagicMock() + monkeypatch.setattr(AgentHomeSnapshotService, "_client", lambda: nullcontext(client)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", MagicMock(return_value=binding)) + monkeypatch.setattr(AgentWorkspaceService, "validate_binding_generation", validate_generation) + + snapshot = AgentHomeSnapshotService.create_for_build_apply( + session=session, + build_draft=_build_draft(home_snapshot_id=None), + ) + + assert snapshot.snapshot_ref == "snapshot-ref-2" + assert validate_generation.call_args.kwargs["base_home_snapshot_id"] is None + + +def test_build_apply_fails_fast_without_source_binding() -> None: + session = MagicMock() + build_draft = _build_draft() + build_draft.agent_workspace_binding_id = None + + with pytest.raises(AgentBuildSandboxNotFoundError): + AgentHomeSnapshotService.create_for_build_apply( + session=session, + build_draft=build_draft, + ) + + +def test_home_snapshot_collection_database_failure_is_best_effort(monkeypatch: pytest.MonkeyPatch) -> None: + context = MagicMock() + session = context.__enter__.return_value + session.scalar.side_effect = RuntimeError("database unavailable") + log_exception = MagicMock() + monkeypatch.setattr( + "services.agent.home_snapshot_service.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr("services.agent.home_snapshot_service.logger.exception", log_exception) + + AgentHomeSnapshotService.collect_retired_home_snapshot( + tenant_id="tenant-1", + home_snapshot_id="home-1", + ) + + session.scalar.assert_called_once() + log_exception.assert_called_once_with( + "Failed to collect retired Agent Home Snapshot", + extra={"tenant_id": "tenant-1", "home_snapshot_id": "home-1"}, + ) + + +@pytest.mark.parametrize( + "sqlite_session", + [(AgentHomeSnapshot, AgentConfigDraft, AgentConfigSnapshot)], + indirect=True, +) +def test_home_snapshot_collection_final_delete_failure_is_best_effort( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + snapshot = AgentHomeSnapshot( + id="home-1", + tenant_id="tenant-1", + agent_id="agent-1", + snapshot_ref="snapshot-ref-1", + status=AgentWorkingResourceStatus.RETIRED, + ) + sqlite_session.add(snapshot) + sqlite_session.commit() + commit = MagicMock(side_effect=RuntimeError("database unavailable")) + delete = MagicMock() + monkeypatch.setattr( + "services.agent.home_snapshot_service.session_factory.create_session", + lambda: nullcontext(sqlite_session), + ) + monkeypatch.setattr(sqlite_session, "commit", commit) + monkeypatch.setattr(AgentHomeSnapshotService, "delete", delete) + + AgentHomeSnapshotService.collect_retired_home_snapshot( + tenant_id="tenant-1", + home_snapshot_id="home-1", + ) + + delete.assert_called_once_with(snapshot_ref="snapshot-ref-1") diff --git a/api/tests/unit_tests/services/agent/test_prompt_mentions.py b/api/tests/unit_tests/services/agent/test_prompt_mentions.py index 82c42a5bd86..48e4978a3bc 100644 --- a/api/tests/unit_tests/services/agent/test_prompt_mentions.py +++ b/api/tests/unit_tests/services/agent/test_prompt_mentions.py @@ -246,7 +246,7 @@ def test_node_job_resolver_resolves_each_kind(node_job: WorkflowNodeJobConfig): "Read START/tenders and produce qna_report (file output; create the file locally, run " "`dify-agent file upload `, then set final_output.qna_report to a `tool_file` mapping " "using the returned `reference`; if replying to the user in natural language, use the returned " - "`download_url`; do not call final_output before upload succeeds, and do not use the local path, " + "`public_download_url`; do not call final_output before upload succeeds, and do not use the local path, " "filename, URL, or a synthesized dify-file-ref as the reference); " "if unsure contact EMAIL · David Hayes." ) diff --git a/api/tests/unit_tests/services/agent/test_retirement_service.py b/api/tests/unit_tests/services/agent/test_retirement_service.py new file mode 100644 index 00000000000..71362eecfe3 --- /dev/null +++ b/api/tests/unit_tests/services/agent/test_retirement_service.py @@ -0,0 +1,205 @@ +from contextlib import nullcontext +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from sqlalchemy.orm import Session + +from models.agent import ( + Agent, + AgentConfigVersionKind, + AgentHomeSnapshot, + AgentKind, + AgentScope, + AgentSource, + AgentStatus, + AgentWorkingResourceStatus, + AgentWorkspace, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, + WorkflowAgentBindingType, + WorkflowAgentNodeBinding, +) +from models.enums import AppStatus +from models.model import App, AppMode +from models.workflow import Workflow, WorkflowType +from services.agent.home_snapshot_service import AgentHomeSnapshotService +from services.agent.retirement_service import WorkflowAgentRetirementService +from services.agent.workspace_service import AgentWorkspaceService + + +def test_retire_unowned_commits_resource_retirement(monkeypatch: pytest.MonkeyPatch) -> None: + context = MagicMock() + session = context.__enter__.return_value + session.scalars.return_value.all.return_value = [SimpleNamespace(id="binding-1")] + monkeypatch.setattr( + "services.agent.retirement_service.session_factory.create_session", + lambda: context, + ) + monkeypatch.setattr( + WorkflowAgentRetirementService, + "archive_unowned", + MagicMock(return_value=["agent-1"]), + ) + monkeypatch.setattr( + AgentWorkspaceService, + "retire_binding", + MagicMock(return_value="binding-1"), + ) + monkeypatch.setattr( + AgentHomeSnapshotService, + "retire_all_for_agent", + MagicMock(return_value=["home-1"]), + ) + + result = WorkflowAgentRetirementService.retire_unowned( + tenant_id="tenant-1", + agent_ids=["agent-1"], + account_id="account-1", + ) + + assert result == (["binding-1"], ["home-1"]) + session.commit.assert_called_once_with() + + +def _workflow_only_agent() -> Agent: + return Agent( + id="agent-1", + tenant_id="tenant-1", + name="Inline Agent", + description="", + role="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.WORKFLOW_ONLY, + source=AgentSource.WORKFLOW, + status=AgentStatus.ACTIVE, + ) + + +@pytest.mark.parametrize( + "sqlite_session", + [(Agent, App, Workflow, WorkflowAgentNodeBinding)], + indirect=True, +) +def test_retire_unowned_keeps_effectively_owned_agent_active( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + agent = _workflow_only_agent() + app = App( + id="app-1", + tenant_id="tenant-1", + name="Workflow", + mode=AppMode.WORKFLOW, + status=AppStatus.NORMAL, + enable_site=True, + enable_api=True, + ) + workflow = Workflow.new( + tenant_id="tenant-1", + app_id=app.id, + type=WorkflowType.WORKFLOW.value, + version=Workflow.VERSION_DRAFT, + graph="{}", + features="{}", + created_by="account-1", + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + binding = WorkflowAgentNodeBinding( + tenant_id="tenant-1", + app_id=app.id, + workflow_id=workflow.id, + workflow_version=workflow.version, + node_id="agent-node", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id=agent.id, + current_snapshot_id="config-1", + node_job_config={}, + ) + sqlite_session.add_all([agent, app, workflow, binding]) + sqlite_session.commit() + monkeypatch.setattr( + "services.agent.retirement_service.session_factory.create_session", + lambda: nullcontext(sqlite_session), + ) + + result = WorkflowAgentRetirementService.retire_unowned( + tenant_id="tenant-1", + agent_ids=[agent.id], + account_id="account-1", + ) + + assert result == ([], []) + stored_agent = sqlite_session.get(Agent, agent.id) + assert stored_agent is not None + assert stored_agent.status is AgentStatus.ACTIVE + + +@pytest.mark.parametrize( + "sqlite_session", + [(Agent, App, Workflow, WorkflowAgentNodeBinding, AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], + indirect=True, +) +def test_retire_unowned_archives_orphan_and_retires_resources( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + agent = _workflow_only_agent() + home = AgentHomeSnapshot( + id="home-1", + tenant_id="tenant-1", + agent_id=agent.id, + snapshot_ref="home-ref", + status=AgentWorkingResourceStatus.ACTIVE, + ) + workspace = AgentWorkspace( + id="workspace-1", + tenant_id="tenant-1", + app_id="app-1", + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id="conversation-1", + owner_scope_key="root", + backend_workspace_ref="workspace-ref", + status=AgentWorkingResourceStatus.ACTIVE, + active_guard=1, + ) + binding = AgentWorkspaceBinding( + id="binding-1", + tenant_id="tenant-1", + app_id="app-1", + workspace_id=workspace.id, + agent_id=agent.id, + base_home_snapshot_id=home.id, + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + backend_binding_ref="binding-ref", + status=AgentWorkingResourceStatus.ACTIVE, + ) + sqlite_session.add_all([agent, home, workspace, binding]) + sqlite_session.commit() + monkeypatch.setattr( + "services.agent.retirement_service.session_factory.create_session", + lambda: nullcontext(sqlite_session), + ) + + result = WorkflowAgentRetirementService.retire_unowned( + tenant_id="tenant-1", + agent_ids=[agent.id], + account_id="account-1", + ) + + assert result == ([binding.id], [home.id]) + stored_agent = sqlite_session.get(Agent, agent.id) + stored_binding = sqlite_session.get(AgentWorkspaceBinding, binding.id) + stored_workspace = sqlite_session.get(AgentWorkspace, workspace.id) + stored_home = sqlite_session.get(AgentHomeSnapshot, home.id) + assert stored_agent is not None + assert stored_binding is not None + assert stored_workspace is not None + assert stored_home is not None + assert stored_agent.status is AgentStatus.ARCHIVED + assert stored_binding.status is AgentWorkingResourceStatus.RETIRED + assert stored_workspace.status is AgentWorkingResourceStatus.RETIRED + assert stored_home.status is AgentWorkingResourceStatus.RETIRED diff --git a/api/tests/unit_tests/services/agent/test_skill_package_service.py b/api/tests/unit_tests/services/agent/test_skill_package_service.py index 5c4e479e187..451069a3f9a 100644 --- a/api/tests/unit_tests/services/agent/test_skill_package_service.py +++ b/api/tests/unit_tests/services/agent/test_skill_package_service.py @@ -210,10 +210,10 @@ def test_validate_and_normalize_rejects_archive_too_large_uncompressed(monkeypat def test_validate_and_normalize_rejects_archive_too_large_uploaded_bytes(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(skill_package_service_module, "_MAX_ARCHIVE_BYTES", 8) + monkeypatch.setattr(skill_package_service_module.dify_config, "UPLOAD_SKILL_FILE_SIZE_LIMIT", 1) with pytest.raises(SkillPackageError) as exc_info: - SkillPackageService().validate_and_normalize(content=b"x" * 9, filename="skill.zip") + SkillPackageService().validate_and_normalize(content=b"x" * (1024 * 1024 + 1), filename="skill.zip") assert exc_info.value.code == "archive_too_large" diff --git a/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py b/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py index dd61ba79c19..68c07abb377 100644 --- a/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py +++ b/api/tests/unit_tests/services/agent/test_skill_tool_inference_service.py @@ -139,7 +139,7 @@ def test_binary_skill_md_maps_to_404(sqlite_session: Session): # ── real-path coverage: _invoke / passthrough ──────────────────────────────── -def test_invoke_maps_missing_default_model_to_400(monkeypatch): +def test_invoke_maps_missing_default_model_to_400(monkeypatch: pytest.MonkeyPatch): import services.agent.skill_tool_inference_service as module from core.errors.error import ProviderTokenNotInitError @@ -153,7 +153,7 @@ def test_invoke_maps_missing_default_model_to_400(monkeypatch): assert exc_info.value.status_code == 400 -def test_invoke_maps_model_failure_to_422_and_success_returns_text(monkeypatch): +def test_invoke_maps_model_failure_to_422_and_success_returns_text(monkeypatch: pytest.MonkeyPatch): import services.agent.skill_tool_inference_service as module fake_manager = MagicMock() diff --git a/api/tests/unit_tests/services/agent/test_workflow_publish_service.py b/api/tests/unit_tests/services/agent/test_workflow_publish_service.py index ef687714717..cb577eed287 100644 --- a/api/tests/unit_tests/services/agent/test_workflow_publish_service.py +++ b/api/tests/unit_tests/services/agent/test_workflow_publish_service.py @@ -2,9 +2,13 @@ from types import SimpleNamespace from unittest.mock import ANY, Mock import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session -from models.agent import WorkflowAgentBindingType, WorkflowAgentNodeBinding +from models.agent import Agent, WorkflowAgentBindingType, WorkflowAgentNodeBinding from models.agent_config_entities import WorkflowNodeJobConfig +from models.enums import AppStatus +from models.model import App, AppMode from models.workflow import Workflow, WorkflowType from services.agent.dsl_service import AgentDslService from services.agent.workflow_publish_service import WorkflowAgentPublishService, _InlineAgentOwnershipError @@ -25,7 +29,7 @@ def _workflow(*, workflow_id: str = "workflow-1", version: str = Workflow.VERSIO ) -def test_inline_binding_from_another_node_is_cloned(monkeypatch) -> None: +def test_inline_binding_from_another_node_is_cloned(monkeypatch: pytest.MonkeyPatch) -> None: session = Mock() draft_workflow = _workflow() monkeypatch.setattr( @@ -63,16 +67,28 @@ def test_inline_binding_from_another_node_is_cloned(monkeypatch) -> None: assert binding.node_job_config.workflow_prompt == "Summarize the input" -def test_restore_replaces_draft_bindings_with_published_bindings() -> None: - existing = WorkflowAgentNodeBinding( +def test_restore_replaces_bindings_and_returns_only_replaced_inline_agent() -> None: + existing_inline = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", workflow_id="draft-workflow", workflow_version=Workflow.VERSION_DRAFT, - node_id="old-node", + node_id="old-inline-node", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id="old-inline-agent", + current_snapshot_id="old-inline-snapshot", + node_job_config={}, + created_by="account-1", + ) + existing_roster = WorkflowAgentNodeBinding( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="draft-workflow", + workflow_version=Workflow.VERSION_DRAFT, + node_id="old-roster-node", binding_type=WorkflowAgentBindingType.ROSTER_AGENT, - agent_id="old-agent", - current_snapshot_id="old-snapshot", + agent_id="old-roster-agent", + current_snapshot_id="old-roster-snapshot", node_job_config={}, created_by="account-1", ) @@ -89,16 +105,21 @@ def test_restore_replaces_draft_bindings_with_published_bindings() -> None: created_by="account-1", ) session = Mock() - session.scalars.side_effect = [SimpleNamespace(all=lambda: [existing]), SimpleNamespace(all=lambda: [source])] - - WorkflowAgentPublishService.restore_agent_node_bindings_to_draft( + session.scalars.side_effect = [ + SimpleNamespace(all=lambda: [existing_inline, existing_roster]), + SimpleNamespace(all=lambda: [source]), + ] + retirement_candidates = WorkflowAgentPublishService.restore_agent_node_bindings_to_draft( session=session, source_workflow=_workflow(workflow_id="published-workflow", version="2026-07-13 00:00:00"), draft_workflow=_workflow(workflow_id="draft-workflow"), account_id="account-2", ) - session.delete.assert_called_once_with(existing) + assert {item.args[0].agent_id for item in session.delete.call_args_list} == { + "old-inline-agent", + "old-roster-agent", + } restored = session.add.call_args.args[0] assert isinstance(restored, WorkflowAgentNodeBinding) assert restored.workflow_id == "draft-workflow" @@ -107,9 +128,93 @@ def test_restore_replaces_draft_bindings_with_published_bindings() -> None: assert restored.current_snapshot_id == "published-snapshot" assert restored.node_job_config.workflow_prompt == "Use the roster agent" session.flush.assert_called_once() + assert retirement_candidates == {"old-inline-agent"} -def test_inline_binding_reuses_existing_node_owned_agent(monkeypatch) -> None: +@pytest.mark.parametrize( + "sqlite_session", + [(App, Agent, WorkflowAgentNodeBinding)], + indirect=True, +) +def test_publish_binding_replacement_returns_only_previous_inline_agent( + sqlite_session: Session, +) -> None: + draft_workflow = _workflow() + draft_workflow.graph = '{"nodes":[{"id":"agent-node","data":{"type":"agent","version":"2"}}],"edges":[]}' + published_workflow = _workflow(workflow_id="published-new", version="published-new") + app = App( + id="app-1", + tenant_id="tenant-1", + name="Workflow", + description="", + mode=AppMode.WORKFLOW, + workflow_id="published-previous", + status=AppStatus.NORMAL, + enable_site=False, + enable_api=False, + api_rpm=0, + api_rph=0, + ) + previous_inline_binding = WorkflowAgentNodeBinding( + id="previous-inline-binding", + tenant_id="tenant-1", + app_id="app-1", + workflow_id="published-previous", + workflow_version="published-previous", + node_id="previous-inline-node", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id="previous-inline-agent", + current_snapshot_id="previous-inline-snapshot", + node_job_config={}, + created_by="account-1", + ) + previous_roster_binding = WorkflowAgentNodeBinding( + id="previous-roster-binding", + tenant_id="tenant-1", + app_id="app-1", + workflow_id="published-previous", + workflow_version="published-previous", + node_id="previous-roster-node", + binding_type=WorkflowAgentBindingType.ROSTER_AGENT, + agent_id="previous-roster-agent", + current_snapshot_id="previous-roster-snapshot", + node_job_config={}, + created_by="account-1", + ) + draft_binding = WorkflowAgentNodeBinding( + id="draft-binding", + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_version=Workflow.VERSION_DRAFT, + node_id="agent-node", + binding_type=WorkflowAgentBindingType.INLINE_AGENT, + agent_id="draft-inline-agent", + current_snapshot_id="draft-inline-snapshot", + node_job_config={"workflow_prompt": "work"}, + created_by="account-1", + ) + sqlite_session.add_all([app, previous_inline_binding, previous_roster_binding, draft_binding]) + sqlite_session.commit() + retirement_candidates = WorkflowAgentPublishService.copy_agent_node_bindings_to_published( + session=sqlite_session, + draft_workflow=draft_workflow, + published_workflow=published_workflow, + ) + + assert retirement_candidates == {"previous-inline-agent"} + copied = sqlite_session.scalar( + select(WorkflowAgentNodeBinding).where( + WorkflowAgentNodeBinding.workflow_id == "published-new", + WorkflowAgentNodeBinding.workflow_version == "published-new", + ) + ) + assert copied is not None + assert copied.agent_id == "draft-inline-agent" + assert copied.current_snapshot_id == "draft-inline-snapshot" + + +def test_inline_binding_reuses_existing_node_owned_agent(monkeypatch: pytest.MonkeyPatch) -> None: session = Mock() draft_workflow = _workflow() existing_binding = WorkflowAgentNodeBinding( @@ -157,7 +262,7 @@ def test_inline_binding_reuses_existing_node_owned_agent(monkeypatch) -> None: clone.assert_not_called() -def test_resolve_existing_inline_binding_agent_returns_valid_agent_or_none(monkeypatch) -> None: +def test_resolve_existing_inline_binding_agent_returns_valid_agent_or_none(monkeypatch: pytest.MonkeyPatch) -> None: binding = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", @@ -209,7 +314,7 @@ def test_resolve_roster_binding_rejects_unpublished_agent() -> None: ) -def test_clone_inline_graph_binding_for_node_clones_source(monkeypatch) -> None: +def test_clone_inline_graph_binding_for_node_clones_source(monkeypatch: pytest.MonkeyPatch) -> None: session = Mock() source_agent = SimpleNamespace(id="source-agent") source_snapshot = SimpleNamespace(id="source-snapshot") @@ -258,7 +363,7 @@ def test_clone_inline_graph_binding_for_node_rejects_missing_source(scalar_resul ) -def test_restore_clones_inline_binding_owned_by_published_workflow(monkeypatch) -> None: +def test_restore_clones_inline_binding_owned_by_published_workflow(monkeypatch: pytest.MonkeyPatch) -> None: source = WorkflowAgentNodeBinding( tenant_id="tenant-1", app_id="app-1", diff --git a/api/tests/unit_tests/services/agent/test_workspace_service.py b/api/tests/unit_tests/services/agent/test_workspace_service.py new file mode 100644 index 00000000000..ade845e33f7 --- /dev/null +++ b/api/tests/unit_tests/services/agent/test_workspace_service.py @@ -0,0 +1,475 @@ +from contextlib import nullcontext +from datetime import datetime, timedelta +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session + +from models.agent import ( + AgentConfigVersionKind, + AgentHomeSnapshot, + AgentWorkingResourceStatus, + AgentWorkspace, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, +) +from services.agent.workspace_service import AgentWorkspaceNotFoundError, AgentWorkspaceService, WorkspaceOwnerScope + + +def _scope() -> WorkspaceOwnerScope: + return WorkspaceOwnerScope( + tenant_id="tenant-1", + app_id="app-1", + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id="conversation-1", + ) + + +def _backend_client() -> MagicMock: + client = MagicMock() + client.create_execution_binding_sync.return_value = SimpleNamespace( + binding_ref="binding-ref", + workspace_ref="workspace-ref", + ) + return client + + +def _workspace( + *, + workspace_id: str = "workspace-1", + tenant_id: str = "tenant-1", + app_id: str = "app-1", + owner_type: AgentWorkspaceOwnerType = AgentWorkspaceOwnerType.CONVERSATION, + owner_id: str = "conversation-1", + status: AgentWorkingResourceStatus = AgentWorkingResourceStatus.ACTIVE, + updated_at: datetime | None = None, + backend_workspace_ref: str = "workspace-ref", +) -> AgentWorkspace: + return AgentWorkspace( + id=workspace_id, + tenant_id=tenant_id, + app_id=app_id, + owner_type=owner_type, + owner_id=owner_id, + owner_scope_key="root", + backend_workspace_ref=backend_workspace_ref, + status=status, + active_guard=1 if status is AgentWorkingResourceStatus.ACTIVE else None, + updated_at=updated_at, + ) + + +def _binding( + *, + binding_id: str = "binding-1", + tenant_id: str = "tenant-1", + app_id: str = "app-1", + workspace_id: str = "workspace-1", + agent_id: str = "agent-1", + status: AgentWorkingResourceStatus = AgentWorkingResourceStatus.ACTIVE, + config_kind: AgentConfigVersionKind = AgentConfigVersionKind.SNAPSHOT, + updated_at: datetime | None = None, +) -> AgentWorkspaceBinding: + return AgentWorkspaceBinding( + id=binding_id, + tenant_id=tenant_id, + app_id=app_id, + workspace_id=workspace_id, + agent_id=agent_id, + base_home_snapshot_id="home-1", + agent_config_version_id="config-1", + agent_config_version_kind=config_kind, + backend_binding_ref=f"{binding_id}-ref", + status=status, + updated_at=updated_at, + ) + + +@pytest.mark.parametrize( + "sqlite_session", + [(AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], + indirect=True, +) +def test_create_binding_success_persists_new_workspace_and_binding( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + sqlite_session.add( + AgentHomeSnapshot( + id="home-1", + tenant_id="tenant-1", + agent_id="agent-1", + snapshot_ref="home-ref", + status=AgentWorkingResourceStatus.ACTIVE, + ) + ) + sqlite_session.commit() + client = _backend_client() + monkeypatch.setattr(AgentWorkspaceService, "_client", lambda: nullcontext(client)) + + binding = AgentWorkspaceService.create_binding( + session=sqlite_session, + scope=_scope(), + agent_id="agent-1", + base_home_snapshot_id="home-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + sqlite_session.commit() + + workspace = sqlite_session.scalar(select(AgentWorkspace)) + stored_binding = sqlite_session.get(AgentWorkspaceBinding, binding.id) + assert workspace is not None + assert stored_binding is not None + assert binding.workspace_id == workspace.id + assert stored_binding.backend_binding_ref == "binding-ref" + request = client.create_execution_binding_sync.call_args.args[0] + assert request.existing_workspace_ref is None + assert request.workspace_id == workspace.id + assert request.home_snapshot_ref == "home-ref" + + +@pytest.mark.parametrize( + "sqlite_session", + [(AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], + indirect=True, +) +def test_create_binding_without_home_snapshot_uses_backend_default( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + client = _backend_client() + monkeypatch.setattr(AgentWorkspaceService, "_client", lambda: nullcontext(client)) + + binding = AgentWorkspaceService.create_binding( + session=sqlite_session, + scope=_scope(), + agent_id="agent-1", + base_home_snapshot_id=None, + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + sqlite_session.commit() + + stored_binding = sqlite_session.get(AgentWorkspaceBinding, binding.id) + assert stored_binding is not None + assert stored_binding.base_home_snapshot_id is None + request = client.create_execution_binding_sync.call_args.args[0] + assert request.home_snapshot_ref is None + + +@pytest.mark.parametrize( + "sqlite_session", + [(AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], + indirect=True, +) +def test_create_binding_rejects_missing_explicit_home_snapshot_before_backend_call( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + client = _backend_client() + monkeypatch.setattr(AgentWorkspaceService, "_client", lambda: nullcontext(client)) + + with pytest.raises(AgentWorkspaceNotFoundError, match="base Home Snapshot is unavailable"): + AgentWorkspaceService.create_binding( + session=sqlite_session, + scope=_scope(), + agent_id="agent-1", + base_home_snapshot_id="missing-home", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + + client.create_execution_binding_sync.assert_not_called() + + +@pytest.mark.parametrize( + "sqlite_session", + [(AgentHomeSnapshot, AgentWorkspace, AgentWorkspaceBinding)], + indirect=True, +) +def test_create_second_binding_reuses_existing_workspace( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + workspace = _workspace() + first_binding = _binding() + home = AgentHomeSnapshot( + id="home-1", + tenant_id="tenant-1", + agent_id="agent-1", + snapshot_ref="home-ref", + status=AgentWorkingResourceStatus.ACTIVE, + ) + sqlite_session.add_all([workspace, first_binding, home]) + sqlite_session.commit() + client = _backend_client() + monkeypatch.setattr(AgentWorkspaceService, "_client", lambda: nullcontext(client)) + + binding = AgentWorkspaceService.create_binding( + session=sqlite_session, + scope=_scope(), + agent_id="agent-1", + base_home_snapshot_id="home-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + ) + sqlite_session.commit() + + assert binding.workspace_id == workspace.id + assert sqlite_session.scalars(select(AgentWorkspace)).all() == [workspace] + assert {row.id for row in sqlite_session.scalars(select(AgentWorkspaceBinding)).all()} == { + first_binding.id, + binding.id, + } + request = client.create_execution_binding_sync.call_args.args[0] + assert request.existing_workspace_ref == "workspace-ref" + assert request.workspace_id == workspace.id + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_get_active_binding_resolves_exact_participant(sqlite_session: Session) -> None: + conversation_workspace = _workspace(workspace_id="workspace-conversation") + build_workspace = _workspace( + workspace_id="workspace-build", + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id="draft-1", + ) + conversation_binding = _binding( + binding_id="binding-conversation", + workspace_id=conversation_workspace.id, + ) + build_binding = _binding( + binding_id="binding-build", + workspace_id=build_workspace.id, + config_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) + sqlite_session.add_all([conversation_workspace, build_workspace, conversation_binding, build_binding]) + sqlite_session.commit() + + resolved = AgentWorkspaceService.get_active_binding( + session=sqlite_session, + tenant_id="tenant-1", + binding_id=conversation_binding.id, + expected_owner_scope=_scope(), + ) + + assert resolved is not None + assert resolved.id == conversation_binding.id + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_get_active_binding_rejects_wrong_owner(sqlite_session: Session) -> None: + build_workspace = _workspace( + workspace_id="workspace-build", + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id="draft-1", + ) + build_binding = _binding( + binding_id="binding-build", + workspace_id=build_workspace.id, + config_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) + sqlite_session.add_all([build_workspace, build_binding]) + sqlite_session.commit() + + resolved = AgentWorkspaceService.get_active_binding( + session=sqlite_session, + tenant_id="tenant-1", + binding_id=build_binding.id, + expected_owner_scope=_scope(), + ) + + assert resolved is None + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_retire_non_final_binding_keeps_workspace_active(sqlite_session: Session) -> None: + binding = _binding() + other_binding = _binding(binding_id="binding-2", agent_id="agent-2") + workspace = _workspace() + sqlite_session.add_all([workspace, binding, other_binding]) + sqlite_session.commit() + + retired_id = AgentWorkspaceService.retire_binding( + session=sqlite_session, + tenant_id="tenant-1", + binding_id=binding.id, + ) + + assert retired_id == binding.id + assert binding.status is AgentWorkingResourceStatus.RETIRED + assert binding.retired_at is not None + assert workspace.status is AgentWorkingResourceStatus.ACTIVE + assert other_binding.status is AgentWorkingResourceStatus.ACTIVE + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_retire_final_binding_retires_workspace(sqlite_session: Session) -> None: + binding = _binding() + workspace = _workspace() + sqlite_session.add_all([workspace, binding]) + sqlite_session.commit() + + AgentWorkspaceService.retire_binding(session=sqlite_session, tenant_id="tenant-1", binding_id=binding.id) + + assert binding.status is AgentWorkingResourceStatus.RETIRED + assert workspace.status is AgentWorkingResourceStatus.RETIRED + assert workspace.active_guard is None + assert workspace.retired_at == binding.retired_at + + +def test_retire_workspace_retires_all_active_bindings() -> None: + workspace = _workspace() + bindings = [_binding(), _binding(binding_id="binding-2", agent_id="agent-2")] + session = MagicMock() + session.scalar.return_value = workspace + session.scalars.return_value.all.return_value = bindings + + retired_id = AgentWorkspaceService.retire_workspace( + session=session, + tenant_id="tenant-1", + workspace_id=workspace.id, + ) + + assert retired_id == workspace.id + assert workspace.status is AgentWorkingResourceStatus.RETIRED + assert all(binding.status is AgentWorkingResourceStatus.RETIRED for binding in bindings) + assert all(binding.retired_at == workspace.retired_at for binding in bindings) + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_retire_all_for_app_retires_only_active_workspaces_for_that_app(sqlite_session: Session) -> None: + active = _workspace(workspace_id="workspace-active", owner_id="conversation-active") + already_retired = _workspace( + workspace_id="workspace-retired", + owner_id="conversation-retired", + status=AgentWorkingResourceStatus.RETIRED, + ) + other_app = _workspace( + workspace_id="workspace-other", + app_id="app-2", + owner_id="conversation-other", + ) + active_binding = _binding(binding_id="binding-active", workspace_id=active.id) + other_binding = _binding( + binding_id="binding-other", + app_id="app-2", + workspace_id=other_app.id, + ) + sqlite_session.add_all([active, already_retired, other_app, active_binding, other_binding]) + sqlite_session.commit() + + retired_ids = AgentWorkspaceService.retire_all_for_app( + session=sqlite_session, + tenant_id="tenant-1", + app_id="app-1", + ) + + assert retired_ids == [active.id] + assert active.status is AgentWorkingResourceStatus.RETIRED + assert active_binding.status is AgentWorkingResourceStatus.RETIRED + assert already_retired.status is AgentWorkingResourceStatus.RETIRED + assert other_app.status is AgentWorkingResourceStatus.ACTIVE + assert other_binding.status is AgentWorkingResourceStatus.ACTIVE + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_collect_binding_without_retired_workspace_destroys_binding_only( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + binding = _binding(status=AgentWorkingResourceStatus.RETIRED) + workspace = _workspace() + sqlite_session.add_all([workspace, binding]) + sqlite_session.commit() + client = MagicMock() + monkeypatch.setattr( + "services.agent.workspace_service.session_factory.create_session", lambda: nullcontext(sqlite_session) + ) + monkeypatch.setattr(AgentWorkspaceService, "_client", lambda: nullcontext(client)) + + AgentWorkspaceService.collect_retired_binding(tenant_id="tenant-1", binding_id=binding.id) + + request = client.destroy_execution_binding_sync.call_args.args[0] + assert request.binding_ref == binding.backend_binding_ref + assert request.destroy_workspace is False + assert request.workspace_ref is None + assert sqlite_session.get(AgentWorkspaceBinding, binding.id) is None + assert sqlite_session.get(AgentWorkspace, workspace.id) is not None + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_collect_workspace_destroys_workspace_then_remaining_bindings( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + workspace = _workspace(status=AgentWorkingResourceStatus.RETIRED) + anchor = _binding(status=AgentWorkingResourceStatus.RETIRED) + remaining = _binding( + binding_id="binding-2", + agent_id="agent-2", + status=AgentWorkingResourceStatus.RETIRED, + ) + anchor.created_at = datetime(2026, 7, 23, 10) + remaining.created_at = anchor.created_at + timedelta(minutes=1) + sqlite_session.add_all([workspace, anchor, remaining]) + sqlite_session.commit() + client = MagicMock() + monkeypatch.setattr( + "services.agent.workspace_service.session_factory.create_session", lambda: nullcontext(sqlite_session) + ) + monkeypatch.setattr(AgentWorkspaceService, "_client", lambda: nullcontext(client)) + + AgentWorkspaceService.collect_retired_workspace(tenant_id="tenant-1", workspace_id=workspace.id) + + requests = [call.args[0] for call in client.destroy_execution_binding_sync.call_args_list] + assert len(requests) == 2 + assert requests[0].binding_ref == anchor.backend_binding_ref + assert requests[0].workspace_ref == workspace.backend_workspace_ref + assert requests[0].destroy_workspace is True + assert requests[1].binding_ref == remaining.backend_binding_ref + assert requests[1].destroy_workspace is False + assert sqlite_session.get(AgentWorkspace, workspace.id) is None + assert sqlite_session.get(AgentWorkspaceBinding, anchor.id) is None + assert sqlite_session.get(AgentWorkspaceBinding, remaining.id) is None + + +def test_binding_collection_database_failure_is_best_effort(monkeypatch: pytest.MonkeyPatch) -> None: + context = MagicMock() + session = context.__enter__.return_value + session.scalar.side_effect = RuntimeError("database unavailable") + log_exception = MagicMock() + monkeypatch.setattr("services.agent.workspace_service.session_factory.create_session", lambda: context) + monkeypatch.setattr("services.agent.workspace_service.logger.exception", log_exception) + + AgentWorkspaceService.collect_retired_binding(tenant_id="tenant-1", binding_id="binding-1") + + session.scalar.assert_called_once() + log_exception.assert_called_once_with( + "Failed to collect retired Agent Workspace Binding", + extra={"tenant_id": "tenant-1", "binding_id": "binding-1"}, + ) + + +@pytest.mark.parametrize("sqlite_session", [(AgentWorkspace, AgentWorkspaceBinding)], indirect=True) +def test_workspace_collection_final_delete_failure_is_best_effort( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + workspace = _workspace(status=AgentWorkingResourceStatus.RETIRED) + anchor = _binding(status=AgentWorkingResourceStatus.RETIRED) + sqlite_session.add_all([workspace, anchor]) + sqlite_session.commit() + commit = MagicMock(side_effect=RuntimeError("database unavailable")) + client = MagicMock() + log_exception = MagicMock() + monkeypatch.setattr( + "services.agent.workspace_service.session_factory.create_session", lambda: nullcontext(sqlite_session) + ) + monkeypatch.setattr(sqlite_session, "commit", commit) + monkeypatch.setattr(AgentWorkspaceService, "_client", lambda: nullcontext(client)) + monkeypatch.setattr("services.agent.workspace_service.logger.exception", log_exception) + + AgentWorkspaceService.collect_retired_workspace(tenant_id="tenant-1", workspace_id="workspace-1") + + client.destroy_execution_binding_sync.assert_called_once() + log_exception.assert_called_once_with( + "Failed to collect retired Agent Workspace", + extra={"tenant_id": "tenant-1", "workspace_id": "workspace-1"}, + ) diff --git a/api/tests/unit_tests/services/auth/test_watercrawl_auth.py b/api/tests/unit_tests/services/auth/test_watercrawl_auth.py index 6d76d046874..c8bc8414cd1 100644 --- a/api/tests/unit_tests/services/auth/test_watercrawl_auth.py +++ b/api/tests/unit_tests/services/auth/test_watercrawl_auth.py @@ -77,6 +77,7 @@ class TestWatercrawlAuth: mock_get.assert_called_once_with( "https://app.watercrawl.dev/api/v1/core/crawl-requests/", headers={"Content-Type": "application/json", "X-API-KEY": "test_api_key_123"}, + timeout=httpx.Timeout(10.0), ) @pytest.mark.parametrize( diff --git a/api/tests/unit_tests/services/controller_api.py b/api/tests/unit_tests/services/controller_api.py index 48805da884e..4dd0019cfc0 100644 --- a/api/tests/unit_tests/services/controller_api.py +++ b/api/tests/unit_tests/services/controller_api.py @@ -97,6 +97,7 @@ from controllers.console.datasets.external import ( ExternalApiTemplateListApi, ) from controllers.console.datasets.hit_testing import HitTestingApi +from enums import DeploymentEdition from models.account import Account, AccountStatus, TenantAccountRole from models.dataset import Dataset, DatasetPermissionEnum @@ -192,34 +193,27 @@ class ControllerApiTestDataFactory: tenant_id: str = "tenant-123", permission: DatasetPermissionEnum = DatasetPermissionEnum.ONLY_ME, **kwargs, - ) -> Mock: + ) -> Dataset: """ - Create a mock Dataset instance. + Create a Dataset instance. Args: dataset_id: Unique identifier for the dataset name: Name of the dataset tenant_id: Tenant identifier permission: Dataset permission level - **kwargs: Additional attributes to set on the mock + **kwargs: Additional mapped attributes for the dataset Returns: - Mock object configured as a Dataset instance + Configured Dataset instance """ - dataset = Mock(spec=Dataset) - dataset.id = dataset_id - dataset.name = name - dataset.tenant_id = tenant_id - dataset.permission = permission - dataset.to_dict.return_value = { - "id": dataset_id, - "name": name, - "tenant_id": tenant_id, - "permission": permission.value, - } - for key, value in kwargs.items(): - setattr(dataset, key, value) - return dataset + return Dataset( + id=dataset_id, + name=name, + tenant_id=tenant_id, + permission=permission, + **kwargs, + ) @staticmethod def create_user_mock( @@ -240,15 +234,14 @@ class ControllerApiTestDataFactory: Returns: Mock object configured as a user/account instance """ - user = Mock() - user.id = user_id - user.current_tenant_id = tenant_id - user.is_dataset_editor = is_dataset_editor - user.has_edit_permission = True - user.is_dataset_operator = False - for key, value in kwargs.items(): - setattr(user, key, value) - return user + return Mock( + id=user_id, + current_tenant_id=tenant_id, + is_dataset_editor=is_dataset_editor, + has_edit_permission=True, + is_dataset_operator=False, + **kwargs, + ) @staticmethod def create_paginated_response(items, total, page=1, per_page=20): @@ -822,23 +815,22 @@ class TestExternalDatasetApi: @pytest.fixture def mock_current_account_context(self, app: Flask) -> Iterator[Mock]: """Provide the wrapper auth context required by HTTP-client controller tests.""" - mock_user = Account(name="Test User", email="user-123@example.com") + mock_user = Account( + name="Test User", + email="user-123@example.com", + status=AccountStatus.ACTIVE, + ) mock_user.id = "user-123" - mock_user.status = AccountStatus.ACTIVE mock_user.role = TenantAccountRole.EDITOR def load_user_from_request_context() -> None: g._login_user = mock_user - setattr( # noqa: B010 - app, - "login_manager", - SimpleNamespace(load_user_from_request_context=load_user_from_request_context), - ) + app.login_manager = SimpleNamespace(load_user_from_request_context=load_user_from_request_context) with ( patch("controllers.console.wraps.current_account_with_tenant") as mock_get_user, - patch("controllers.console.wraps.dify_config.EDITION", "CLOUD"), + patch("controllers.console.wraps.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("libs.login.check_csrf_token", return_value=None), ): mock_tenant_id = "tenant-123" diff --git a/api/tests/unit_tests/services/data_migration/test_export_service.py b/api/tests/unit_tests/services/data_migration/test_export_service.py index dbb05386cb3..faad5d20521 100644 --- a/api/tests/unit_tests/services/data_migration/test_export_service.py +++ b/api/tests/unit_tests/services/data_migration/test_export_service.py @@ -98,7 +98,7 @@ def test_export_config_parser_rejects_unsupported_app_modes(): ) -def test_secret_free_api_tool_export_uses_masking_and_omits_credentials(monkeypatch): +def test_secret_free_api_tool_export_uses_masking_and_omits_credentials(monkeypatch: pytest.MonkeyPatch): calls = [] def fake_get_api_provider(provider: str, tenant_id: str, mask: bool = True): diff --git a/api/tests/unit_tests/services/data_migration/test_import_service.py b/api/tests/unit_tests/services/data_migration/test_import_service.py index 10460aba470..02a897a8f20 100644 --- a/api/tests/unit_tests/services/data_migration/test_import_service.py +++ b/api/tests/unit_tests/services/data_migration/test_import_service.py @@ -1,7 +1,15 @@ +from dataclasses import dataclass +from typing import cast + import pytest import yaml +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session -from models.tools import MCPToolProvider, WorkflowToolProvider +from core.tools.entities.tool_entities import ApiProviderSchemaType +from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole +from models.model import App, AppMode +from models.tools import ApiToolProvider, MCPToolProvider, WorkflowToolProvider from services.app_dsl_service import Import from services.data_migration import import_service from services.data_migration.entities import ( @@ -19,6 +27,142 @@ from services.data_migration.import_service import ImportRequest, ImportTargetRe from services.entities.dsl_entities import ImportStatus +@dataclass(frozen=True) +class Database: + """Typed database binding used by import code that still reads ``db.engine``.""" + + engine: Engine + session: Session + + +@pytest.fixture +def database(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Database: + database = Database(engine=cast(Engine, sqlite_session.get_bind()), session=sqlite_session) + monkeypatch.setattr(import_service, "db", database) + return database + + +def _persist_tenant_account( + session: Session, + *, + tenant_id: str = "tenant-1", + tenant_name: str = "target", + account_id: str = "account-1", +) -> tuple[Tenant, Account]: + tenant = Tenant(name=tenant_name) + tenant.id = tenant_id + account = Account(name="Owner", email="owner@example.com") + account.id = account_id + join = TenantAccountJoin( + tenant_id=tenant_id, + account_id=account_id, + current=True, + role=TenantAccountRole.OWNER, + ) + session.add_all([tenant, account, join]) + session.commit() + return tenant, account + + +def _persist_app( + session: Session, + *, + app_id: str, + tenant_id: str = "tenant-1", +) -> App: + app = App( + id=app_id, + tenant_id=tenant_id, + name=f"App {app_id}", + description="Migration fixture", + mode=AppMode.WORKFLOW, + icon_type=None, + icon="", + icon_background=None, + workflow_id=None, + enable_site=False, + enable_api=False, + max_active_requests=None, + created_by="account-1", + maintainer="account-1", + ) + session.add(app) + session.commit() + return app + + +def _persist_workflow_provider( + session: Session, + *, + provider_id: str, + app_id: str, + name: str = "embedded_workflow_as_tool", + tenant_id: str = "tenant-1", +) -> WorkflowToolProvider: + provider = WorkflowToolProvider( + name=name, + label=name, + icon="{}", + app_id=app_id, + version="", + user_id="account-1", + tenant_id=tenant_id, + description="", + parameter_configuration="[]", + ) + provider.id = provider_id + session.add(provider) + session.commit() + return provider + + +def _persist_api_provider( + session: Session, + *, + provider_id: str, + name: str = "weather", + tenant_id: str = "tenant-1", +) -> ApiToolProvider: + provider = ApiToolProvider( + name=name, + icon="{}", + schema="openapi: 3.0.0", + schema_type_str=ApiProviderSchemaType.OPENAPI, + user_id="account-1", + tenant_id=tenant_id, + description="", + tools_str="[]", + credentials_str="{}", + ) + provider.id = provider_id + session.add(provider) + session.commit() + return provider + + +def _persist_mcp_provider( + session: Session, + *, + provider_id: str, + server_identifier: str = "my-test-mcp", + name: str = "my-test-mcp", + tenant_id: str = "tenant-1", +) -> MCPToolProvider: + provider = MCPToolProvider( + name=name, + server_identifier=server_identifier, + server_url="http://localhost:3000/mcp", + server_url_hash=f"hash-{provider_id}", + icon=None, + tenant_id=tenant_id, + user_id="account-1", + ) + provider.id = provider_id + session.add(provider) + session.commit() + return provider + + def test_target_tenant_precedence_cli_then_config_then_package(): package = MigrationPackage.from_mapping( { @@ -77,31 +221,18 @@ def test_target_tenant_name_is_not_treated_as_uuid(): assert resolver._is_uuid("49a99e46-bc2c-4885-91fa-47615f6192b5") is True -def test_package_target_tenant_id_ignores_invalid_uuid(monkeypatch): +def test_package_target_tenant_id_ignores_invalid_uuid(database: Database): package = MigrationPackage.from_mapping( {"metadata": {"version": "1", "source_scope": "single", "target_tenant": {"id": "not-a-uuid"}}} ) - class StubSession: - def get(self, model, identifier): - raise AssertionError("invalid UUID should not be passed to session.get") - - def scalars(self, statement): - class EmptyResult: - def all(self): - return [] - - return EmptyResult() - - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - with pytest.raises(MigrationDataError, match="Target tenant not found"): - ImportTargetResolver().resolve(ImportRequest(package=package), session=import_service.db.session) + ImportTargetResolver().resolve(ImportRequest(package=package), session=database.session) + + assert database.session.query(Tenant).count() == 0 -def test_options_override_replaces_package_defaults(): +def test_options_override_replaces_package_defaults(database: Database): package = MigrationPackage.from_mapping( { "metadata": { @@ -142,7 +273,7 @@ def test_options_override_replaces_package_defaults(): CapturingImportService(target_resolver=StubResolver()).import_package( ImportRequest(package=package, options_override=override), - session=import_service.db.session, + session=database.session, ) assert captured_options == [override] @@ -155,74 +286,55 @@ def test_only_preserve_id_strategy_reuses_source_app_id(): assert service._should_preserve_source_app_id(ImportOptions(id_strategy=IdStrategy.GENERATE_NEW_ID)) is False -def test_find_existing_app_ignores_invalid_uuid(monkeypatch): - class StubSession: - def scalar(self, statement): - raise AssertionError("invalid UUID should not be queried against App.id") - - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - - assert ( - MigrationImportService()._find_existing_app("not-a-uuid", "tenant-1", session=import_service.db.session) is None - ) +def test_find_existing_app_ignores_invalid_uuid(database: Database): + _persist_app(database.session, app_id="app-other", tenant_id="tenant-2") + assert MigrationImportService()._find_existing_app("not-a-uuid", "tenant-1", session=database.session) is None -def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(monkeypatch): - captured = [] +def test_find_existing_workflow_tool_does_not_compare_invalid_uuid(database: Database): + provider = _persist_workflow_provider(database.session, provider_id="provider-1", app_id="app-id", name="tool-name") + statements = [] - class StubSession: - def scalar(self, statement): - captured.append(statement) + def capture_statement(orm_execute_state) -> None: + statements.append(orm_execute_state.statement) - from services.data_migration import import_service + event.listen(database.session, "do_orm_execute", capture_statement) + try: + result = MigrationImportService()._find_existing_workflow_tool( + "tenant-1", "not-a-uuid", "tool-name", "app-id", session=database.session + ) + finally: + event.remove(database.session, "do_orm_execute", capture_statement) - monkeypatch.setattr(import_service.db, "session", StubSession()) - - MigrationImportService()._find_existing_workflow_tool( - "tenant-1", "not-a-uuid", "tool-name", "app-id", session=import_service.db.session - ) - - where_clause = str(captured[0].whereclause) + assert result is provider + where_clause = str(statements[0].whereclause) assert f"{WorkflowToolProvider.__tablename__}.id" not in where_clause -def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(monkeypatch): - captured = [] +def test_find_existing_mcp_tool_does_not_compare_invalid_uuid(database: Database): + provider = _persist_mcp_provider(database.session, provider_id="provider-1") + statements = [] - class StubSession: - def scalar(self, statement): - captured.append(statement) + def capture_statement(orm_execute_state) -> None: + statements.append(orm_execute_state.statement) - from services.data_migration import import_service + event.listen(database.session, "do_orm_execute", capture_statement) + try: + result = MigrationImportService()._find_existing_mcp_tool( + "tenant-1", "my-test-mcp", "my-test-mcp", session=database.session + ) + finally: + event.remove(database.session, "do_orm_execute", capture_statement) - monkeypatch.setattr(import_service.db, "session", StubSession()) - - MigrationImportService()._find_existing_mcp_tool( - "tenant-1", "my-test-mcp", "my-test-mcp", session=import_service.db.session - ) - - where_clause = str(captured[0].whereclause) + assert result is provider + where_clause = str(statements[0].whereclause) assert f"{MCPToolProvider.__tablename__}.id" not in where_clause assert f"{MCPToolProvider.__tablename__}.name" not in where_clause -def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction(monkeypatch): - class FailingNestedTransaction: - def __enter__(self): - raise AssertionError("nested transaction should not be opened") - - def __exit__(self, exc_type, exc_value, traceback): - return False - - class StubSession: - def begin_nested(self): - return FailingNestedTransaction() - - def commit(self): - return None - +def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction( + monkeypatch: pytest.MonkeyPatch, database: Database +): class StubAppDslService: def __init__(self, session): self.session = session @@ -230,22 +342,30 @@ def test_workflow_app_import_does_not_wrap_app_dsl_import_in_nested_transaction( def import_app(self, **kwargs): return Import(id="import-id", status=ImportStatus.COMPLETED, app_id="imported-app-id") - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "AppDslService", StubAppDslService) + nested_transactions = [] - imported_app_id = MigrationImportService()._import_workflow_app( - account=object(), - workflow_data={"name": "main_chatflow"}, - dsl_content="app:\n mode: workflow\n", - app_id="source-app-id", - existing_app=None, - options=ImportOptions(id_strategy=IdStrategy.PRESERVE_ID), - session=import_service.db.session, - ) + def capture_transaction(_session, transaction) -> None: + if transaction.nested: + nested_transactions.append(transaction) + + event.listen(database.session, "after_transaction_create", capture_transaction) + + try: + imported_app_id = MigrationImportService()._import_workflow_app( + account=object(), + workflow_data={"name": "main_chatflow"}, + dsl_content="app:\n mode: workflow\n", + app_id="source-app-id", + existing_app=None, + options=ImportOptions(id_strategy=IdStrategy.PRESERVE_ID), + session=database.session, + ) + finally: + event.remove(database.session, "after_transaction_create", capture_transaction) assert imported_app_id == "imported-app-id" + assert nested_transactions == [] def test_rewrite_workflow_dsl_replaces_tool_provider_ids(): @@ -330,211 +450,149 @@ def test_source_api_provider_ids_are_discovered_from_workflow_dsl(): assert MigrationImportService()._source_api_provider_ids_by_name(package) == {"weather": {"source-api-provider-id"}} -def test_workflow_tool_import_publishes_referenced_app_before_create(monkeypatch): +def test_workflow_tool_import_publishes_referenced_app_before_create( + monkeypatch: pytest.MonkeyPatch, database: Database +): events = [] - account = type("Account", (), {"id": "account-1"})() - - class StubSession: - def get(self, model, identifier): - return account + _persist_tenant_account(database.session) + app_id = "00000000-0000-0000-0000-000000000001" + _persist_app(database.session, app_id=app_id) class PublishingImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id, session): - return object() - - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): - if ("created", app_id) in events: - return type("WorkflowToolProvider", (), {"id": workflow_tool_id or "created-workflow-tool-id"})() - return None - def _ensure_workflow_app_is_published(self, target, account, app_id, session): events.append(("published", app_id)) - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - monkeypatch.setattr( - import_service.WorkflowToolManageService, - "create_workflow_tool", - lambda **kwargs: events.append(("created", kwargs["workflow_app_id"])), - ) + def create_workflow_tool(**kwargs) -> None: + events.append(("created", kwargs["workflow_app_id"])) + _persist_workflow_provider( + database.session, + provider_id="00000000-0000-0000-0000-000000000010", + app_id=kwargs["workflow_app_id"], + ) + monkeypatch.setattr(import_service.WorkflowToolManageService, "create_workflow_tool", create_workflow_tool) PublishingImportService()._import_workflow_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, - "workflow_tools": [ - { - "id": "workflow-tool-1", - "name": "embedded_workflow_as_tool", - "app_id": "workflow-app-1", - } - ], + "workflow_tools": [{"id": "workflow-tool-1", "name": "embedded_workflow_as_tool", "app_id": app_id}], } ), - ImportTarget( - tenant_id="tenant-1", - tenant_name="target", - operator_id="account-1", - operator_email="owner@example.com", - ), + ImportTarget("tenant-1", "target", "account-1", "owner@example.com"), ImportOptions(), {}, [], [], - session=import_service.db.session, + session=database.session, ) - assert events == [("published", "workflow-app-1"), ("created", "workflow-app-1")] + assert events == [("published", app_id), ("created", app_id)] -@pytest.mark.parametrize( - ("id_strategy", "expected_import_id"), - [ - (IdStrategy.PRESERVE_ID, "source-workflow-tool-id"), - (IdStrategy.GENERATE_NEW_ID, ""), - ], -) -def test_workflow_tool_import_id_follows_id_strategy(monkeypatch: pytest.MonkeyPatch, id_strategy, expected_import_id): +@pytest.mark.parametrize("id_strategy", [IdStrategy.PRESERVE_ID, IdStrategy.GENERATE_NEW_ID]) +def test_workflow_tool_import_id_follows_id_strategy( + monkeypatch: pytest.MonkeyPatch, database: Database, id_strategy: IdStrategy +): created_kwargs = [] - target_provider = type("WorkflowToolProvider", (), {"id": "target-workflow-tool-id"})() - account = type("Account", (), {"id": "account-1"})() - id_mapping = {"source-app-id": "target-app-id"} + _persist_tenant_account(database.session) + source_app_id = "00000000-0000-0000-0000-000000000021" + target_app_id = "00000000-0000-0000-0000-000000000022" + source_provider_id = "00000000-0000-0000-0000-000000000020" + generated_provider_id = "00000000-0000-0000-0000-000000000023" + _persist_app(database.session, app_id=target_app_id) + id_mapping = {source_app_id: target_app_id} id_mapping_details = [] - class StubSession: - def get(self, model, identifier): - return account - class StrategyImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id, session): - return object() - - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): - return target_provider if created_kwargs else None - def _ensure_workflow_app_is_published(self, target, account, app_id, session): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - monkeypatch.setattr( - import_service.WorkflowToolManageService, - "create_workflow_tool", - lambda **kwargs: created_kwargs.append(kwargs), - ) + def create_workflow_tool(**kwargs) -> None: + created_kwargs.append(kwargs) + _persist_workflow_provider( + database.session, + provider_id=kwargs["import_id"] or generated_provider_id, + app_id=kwargs["workflow_app_id"], + ) + monkeypatch.setattr(import_service.WorkflowToolManageService, "create_workflow_tool", create_workflow_tool) StrategyImportService()._import_workflow_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, "workflow_tools": [ - { - "id": "source-workflow-tool-id", - "name": "embedded_workflow_as_tool", - "app_id": "source-app-id", - } + {"id": source_provider_id, "name": "embedded_workflow_as_tool", "app_id": source_app_id} ], } ), - ImportTarget( - tenant_id="tenant-1", - tenant_name="target", - operator_id="account-1", - operator_email="owner@example.com", - ), + ImportTarget("tenant-1", "target", "account-1", "owner@example.com"), ImportOptions(id_strategy=id_strategy), id_mapping, id_mapping_details, [], - session=import_service.db.session, + session=database.session, ) + expected_import_id = source_provider_id if id_strategy == IdStrategy.PRESERVE_ID else "" + target_id = expected_import_id or generated_provider_id assert created_kwargs[0]["import_id"] == expected_import_id - assert id_mapping["source-workflow-tool-id"] == "target-workflow-tool-id" + assert id_mapping[source_provider_id] == target_id assert id_mapping_details == [ - ResourceIdMapping( - ResourceType.WORKFLOW_TOOL, - "embedded_workflow_as_tool", - "source-workflow-tool-id", - "target-workflow-tool-id", - ) + ResourceIdMapping(ResourceType.WORKFLOW_TOOL, "embedded_workflow_as_tool", source_provider_id, target_id) ] -def test_workflow_tool_skip_records_id_mapping(monkeypatch): - account = type("Account", (), {"id": "account-1"})() - existing_provider = type("WorkflowToolProvider", (), {"id": "existing-workflow-tool-id"})() - id_mapping = {"source-app-id": "target-app-id"} - - class StubSession: - def get(self, model, identifier): - return account +def test_workflow_tool_skip_records_id_mapping(database: Database): + _persist_tenant_account(database.session) + source_app_id = "00000000-0000-0000-0000-000000000031" + target_app_id = "00000000-0000-0000-0000-000000000032" + source_provider_id = "00000000-0000-0000-0000-000000000034" + _persist_app(database.session, app_id=target_app_id) + existing = _persist_workflow_provider( + database.session, provider_id="00000000-0000-0000-0000-000000000033", app_id=target_app_id + ) + id_mapping = {source_app_id: target_app_id} class SkipImportService(MigrationImportService): - def _find_existing_app(self, app_id, tenant_id, session): - return object() - - def _find_existing_workflow_tool(self, tenant_id, workflow_tool_id, tool_name, app_id, session): - return existing_provider - def _ensure_workflow_app_is_published(self, target, account, app_id, session): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) - SkipImportService()._import_workflow_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, "workflow_tools": [ - { - "id": "source-workflow-tool-id", - "name": "embedded_workflow_as_tool", - "app_id": "source-app-id", - } + {"id": source_provider_id, "name": "embedded_workflow_as_tool", "app_id": source_app_id} ], } ), - ImportTarget( - tenant_id="tenant-1", - tenant_name="target", - operator_id="account-1", - operator_email="owner@example.com", - ), + ImportTarget("tenant-1", "target", "account-1", "owner@example.com"), ImportOptions(conflict_strategy=ConflictStrategy.SKIP, id_strategy=IdStrategy.GENERATE_NEW_ID), id_mapping, [], [], - session=import_service.db.session, + session=database.session, ) - assert id_mapping["source-workflow-tool-id"] == "existing-workflow-tool-id" + assert id_mapping[source_provider_id] == existing.id @pytest.mark.parametrize("conflict_strategy", [ConflictStrategy.SKIP, ConflictStrategy.UPDATE]) -def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_strategy): - target_provider = type("ApiToolProvider", (), {"id": "target-api-provider-id", "name": "weather"})() +def test_api_tool_existing_provider_records_id_mapping( + monkeypatch: pytest.MonkeyPatch, database: Database, conflict_strategy: ConflictStrategy +): + target_provider = _persist_api_provider(database.session, provider_id="target-api-provider-id") + _persist_api_provider(database.session, provider_id="other-tenant-provider-id", tenant_id="tenant-2") id_mapping = {} id_mapping_details = [] report_items = [] - class ExistingApiImportService(MigrationImportService): - def _find_api_tool_provider(self, tenant_id, provider_name, session): - return target_provider - - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: target_provider) monkeypatch.setattr( import_service.ApiToolManageService, "parser_api_schema", lambda schema: {"schema_type": "openapi"} ) monkeypatch.setattr(import_service.ApiToolManageService, "update_api_tool_provider", lambda **kwargs: None) - ExistingApiImportService()._import_api_tools( + MigrationImportService()._import_api_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -552,7 +610,7 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str id_mapping, id_mapping_details, {"weather": {"source-api-provider-id-from-dsl"}}, - session=import_service.db.session, + session=database.session, ) assert id_mapping == { @@ -565,27 +623,19 @@ def test_api_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str ) -def test_api_tool_create_records_id_mapping(monkeypatch): - target_provider = type("ApiToolProvider", (), {"id": "target-api-provider-id", "name": "weather"})() +def test_api_tool_create_records_id_mapping(monkeypatch: pytest.MonkeyPatch, database: Database): id_mapping = {} - class StubSession: - def scalar(self, statement): - return None - - class CreatedApiImportService(MigrationImportService): - def _find_api_tool_provider(self, tenant_id, provider_name, session): - return target_provider - - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr( import_service.ApiToolManageService, "parser_api_schema", lambda schema: {"schema_type": "openapi"} ) - monkeypatch.setattr(import_service.ApiToolManageService, "create_api_tool_provider", lambda **kwargs: None) - CreatedApiImportService()._import_api_tools( + def create_api_tool_provider(**_kwargs) -> None: + _persist_api_provider(database.session, provider_id="target-api-provider-id") + + monkeypatch.setattr(import_service.ApiToolManageService, "create_api_tool_provider", create_api_tool_provider) + + MigrationImportService()._import_api_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -603,25 +653,16 @@ def test_api_tool_create_records_id_mapping(monkeypatch): id_mapping, [], {}, - session=import_service.db.session, + session=database.session, ) assert id_mapping["source-api-provider-id"] == "target-api-provider-id" -def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): - provider = type( - "Provider", (), {"id": "target-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"} - )() +def test_mcp_tool_import_restores_exported_tool_list(monkeypatch: pytest.MonkeyPatch, database: Database): + provider = _persist_mcp_provider(database.session, provider_id="target-provider-id") report_items = [] - class StubSession: - def scalar(self, statement): - return provider - - def commit(self): - return None - class StubMCPToolManageService: def __init__(self, session): self.session = session @@ -629,9 +670,6 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): def update_provider(self, **kwargs): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) MigrationImportService()._import_mcp_tools( @@ -660,29 +698,29 @@ def test_mcp_tool_import_restores_exported_tool_list(monkeypatch): report_items, {}, [], - session=import_service.db.session, + session=database.session, ) + database.session.refresh(provider) assert provider.tools == '[{"name": "echo"}]' assert provider.authed is True @pytest.mark.parametrize("conflict_strategy", [ConflictStrategy.SKIP, ConflictStrategy.UPDATE]) -def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_strategy): - provider = type( - "Provider", (), {"id": "target-mcp-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"} - )() +def test_mcp_tool_existing_provider_records_id_mapping( + monkeypatch: pytest.MonkeyPatch, database: Database, conflict_strategy: ConflictStrategy +): + provider = _persist_mcp_provider(database.session, provider_id="target-mcp-provider-id") + _persist_mcp_provider( + database.session, + provider_id="other-mcp-provider-id", + server_identifier="other-mcp", + name="other-mcp", + tenant_id="tenant-2", + ) id_mapping = {} id_mapping_details = [] - class StubSession: - def commit(self): - return None - - class ExistingMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier, session): - return provider - class StubMCPToolManageService: def __init__(self, session): self.session = session @@ -690,12 +728,9 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str def update_provider(self, **kwargs): return None - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) - ExistingMCPImportService()._import_mcp_tools( + MigrationImportService()._import_mcp_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -721,7 +756,7 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str [], id_mapping, id_mapping_details, - session=import_service.db.session, + session=database.session, ) assert id_mapping["source-mcp-provider-id"] == "target-mcp-provider-id" @@ -731,35 +766,25 @@ def test_mcp_tool_existing_provider_records_id_mapping(monkeypatch, conflict_str ] -def test_mcp_tool_create_records_id_mapping(monkeypatch): - provider = type( - "Provider", (), {"id": "target-mcp-provider-id", "tools": "[]", "authed": False, "identity_mode": "off"} - )() +def test_mcp_tool_create_records_id_mapping(monkeypatch: pytest.MonkeyPatch, database: Database): id_mapping = {} - provider_created = False - - class StubSession: - def commit(self): - return None - - class CreatedMCPImportService(MigrationImportService): - def _find_existing_mcp_tool(self, tenant_id, provider_id, server_identifier, session): - return provider if provider_created else None class StubMCPToolManageService: def __init__(self, session): self.session = session def create_provider(self, **kwargs): - nonlocal provider_created - provider_created = True + _persist_mcp_provider( + database.session, + provider_id="target-mcp-provider-id", + server_identifier=kwargs["server_identifier"], + name=kwargs["name"], + tenant_id=kwargs["tenant_id"], + ) - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db, "session", StubSession()) monkeypatch.setattr(import_service, "MCPToolManageService", StubMCPToolManageService) - CreatedMCPImportService()._import_mcp_tools( + MigrationImportService()._import_mcp_tools( MigrationPackage.from_mapping( { "metadata": {"version": "1", "source_scope": "single"}, @@ -784,13 +809,13 @@ def test_mcp_tool_create_records_id_mapping(monkeypatch): [], id_mapping, [], - session=import_service.db.session, + session=database.session, ) assert id_mapping["source-mcp-provider-id"] == "target-mcp-provider-id" -def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_workflow_context(monkeypatch): +def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_workflow_context(database: Database): report_items = [] package = MigrationPackage.from_mapping( { @@ -829,10 +854,6 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work } ) - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: None) - MigrationImportService()._preflight_dependency_only_mcp( package, ImportTarget( @@ -842,7 +863,7 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work operator_email="owner@example.com", ), report_items, - session=import_service.db.session, + session=database.session, ) assert report_items == [ @@ -857,29 +878,30 @@ def test_dependency_only_mcp_preflight_reports_missing_target_provider_with_work ] -def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(monkeypatch): - captured = [] +def test_dependency_only_mcp_lookup_does_not_compare_non_uuid_identifier_to_uuid_id(database: Database): + provider = _persist_mcp_provider(database.session, provider_id="provider-1", server_identifier="my-test-mcp-server") + statements = [] - class StubSession: - def scalar(self, statement): - captured.append(statement) + def capture_statement(orm_execute_state) -> None: + statements.append(orm_execute_state.statement) - from services.data_migration import import_service + event.listen(database.session, "do_orm_execute", capture_statement) + try: + result = MigrationImportService()._find_dependency_only_mcp_provider( + "tenant-1", + "my-test-mcp-server", + "my-test-mcp", + session=database.session, + ) + finally: + event.remove(database.session, "do_orm_execute", capture_statement) - monkeypatch.setattr(import_service.db, "session", StubSession()) - - MigrationImportService()._find_dependency_only_mcp_provider( - "tenant-1", - "my-test-mcp-server", - "my-test-mcp", - session=import_service.db.session, - ) - - where_clause = str(captured[0].whereclause) + assert result is provider + where_clause = str(statements[0].whereclause) assert f"{MCPToolProvider.__tablename__}.id" not in where_clause -def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeypatch): +def test_dependency_only_mcp_preflight_reports_available_target_provider(database: Database): report_items = [] package = MigrationPackage.from_mapping( { @@ -887,15 +909,18 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp "dependencies": [{"kind": "mcp_tool", "provider_id": "my-test-mcp-server"}], } ) - provider = type( - "Provider", - (), - {"id": "target-provider-id", "name": "my-test-mcp", "server_identifier": "my-test-mcp-server"}, - )() - - from services.data_migration import import_service - - monkeypatch.setattr(import_service.db.session, "scalar", lambda statement: provider) + _persist_mcp_provider( + database.session, + provider_id="target-provider-id", + server_identifier="my-test-mcp-server", + ) + _persist_mcp_provider( + database.session, + provider_id="other-provider-id", + server_identifier="other-server", + name="other-provider", + tenant_id="tenant-2", + ) MigrationImportService()._preflight_dependency_only_mcp( package, @@ -906,7 +931,7 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp operator_email="owner@example.com", ), report_items, - session=import_service.db.session, + session=database.session, ) assert report_items == [ @@ -920,7 +945,7 @@ def test_dependency_only_mcp_preflight_reports_available_target_provider(monkeyp ] -def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): +def test_import_package_imports_workflow_tool_provider_apps_before_consumers(database: Database): events = [] class StubResolver(ImportTargetResolver): @@ -996,7 +1021,7 @@ def test_import_package_imports_workflow_tool_provider_apps_before_consumers(): ) OrderedImportService(target_resolver=StubResolver()).import_package( - ImportRequest(package=package), session=import_service.db.session + ImportRequest(package=package), session=database.session ) assert events == [ diff --git a/api/tests/unit_tests/services/dataset_service_test_helpers.py b/api/tests/unit_tests/services/dataset_service_test_helpers.py index d3dfe2fce79..99157c63cab 100644 --- a/api/tests/unit_tests/services/dataset_service_test_helpers.py +++ b/api/tests/unit_tests/services/dataset_service_test_helpers.py @@ -6,6 +6,7 @@ document, and segment service test modules that exercise """ import json +from datetime import datetime from types import SimpleNamespace from typing import Any from unittest.mock import MagicMock, Mock, create_autospec, patch @@ -18,7 +19,7 @@ from core.rag.entities import PreProcessingRule, Rule, Segmentation from core.rag.index_processor.constant.built_in_field import BuiltInField from core.rag.index_processor.constant.index_type import IndexStructureType from core.rag.retrieval.retrieval_methods import RetrievalMethod -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from graphon.model_runtime.entities.model_entities import ModelFeature, ModelType from models import Account, TenantAccountRole from models.dataset import ( @@ -169,25 +170,23 @@ class DatasetServiceUnitDataFactory: enable_api: bool = False, summary_index_setting: dict[str, Any] | None = None, **kwargs, - ) -> Mock: - dataset = Mock(spec=Dataset) - dataset.id = dataset_id - dataset.tenant_id = tenant_id - dataset.permission = permission - dataset.created_by = created_by - dataset.indexing_technique = indexing_technique - dataset.embedding_model_provider = embedding_model_provider - dataset.embedding_model = embedding_model - dataset.built_in_field_enabled = built_in_field_enabled - dataset.doc_form = doc_form - dataset.get_doc_form.return_value = doc_form - dataset.enable_api = enable_api - dataset.updated_by = None - dataset.updated_at = None - dataset.summary_index_setting = summary_index_setting - for key, value in kwargs.items(): - setattr(dataset, key, value) - return dataset + ) -> Dataset: + return Dataset( + id=dataset_id, + tenant_id=tenant_id, + permission=permission, + created_by=created_by, + indexing_technique=indexing_technique, + embedding_model_provider=embedding_model_provider, + embedding_model=embedding_model, + built_in_field_enabled=built_in_field_enabled, + chunk_structure=doc_form, + enable_api=enable_api, + updated_by=None, + updated_at=None, + summary_index_setting=summary_index_setting, + **kwargs, + ) @staticmethod def create_user_mock( @@ -196,14 +195,12 @@ class DatasetServiceUnitDataFactory: role: str = TenantAccountRole.OWNER, **kwargs, ) -> SimpleNamespace: - user = SimpleNamespace( + return SimpleNamespace( id=user_id, current_tenant_id=tenant_id, current_role=role, + **kwargs, ) - for key, value in kwargs.items(): - setattr(user, key, value) - return user @staticmethod def create_document_mock( @@ -224,34 +221,45 @@ class DatasetServiceUnitDataFactory: doc_metadata: dict[str, Any] | None = None, name: str = "Document", **kwargs, - ) -> Mock: - document = Mock(spec=Document) - document.id = document_id - document.dataset_id = dataset_id - document.tenant_id = tenant_id - document.indexing_status = indexing_status - document.is_paused = is_paused - document.paused_by = None - document.paused_at = None - document.archived = archived - document.enabled = enabled - document.data_source_type = data_source_type - document.data_source_info_dict = data_source_info_dict or {} - document.data_source_info = data_source_info - document.doc_form = doc_form - document.need_summary = need_summary - document.position = position - document.doc_metadata = doc_metadata - document.name = name - for key, value in kwargs.items(): - setattr(document, key, value) - return document + ) -> Document: + return Document( + id=document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + indexing_status=indexing_status, + is_paused=is_paused, + paused_by=None, + paused_at=None, + archived=archived, + enabled=enabled, + data_source_type=data_source_type, + data_source_info=data_source_info + if data_source_info is not None + else json.dumps(data_source_info_dict or {}), + doc_form=doc_form, + need_summary=need_summary, + position=position, + doc_metadata=doc_metadata, + name=name, + **kwargs, + ) @staticmethod - def create_upload_file_mock(file_id: str = "file-123", name: str = "upload.txt") -> Mock: - upload_file = Mock(spec=UploadFile) + def create_upload_file_mock(file_id: str = "file-123", name: str = "upload.txt") -> UploadFile: + upload_file = UploadFile( + tenant_id="tenant-id", + storage_type="opendal", + key="test-key", + name=name, + size=0, + extension="txt", + mime_type="text/plain", + created_by_role="account", + created_by="account-id", + created_at=datetime.now(), + used=False, + ) upload_file.id = file_id - upload_file.name = name return upload_file @@ -281,21 +289,20 @@ def _make_dataset( tenant_id: str = "tenant-1", data_source_type: str | None = None, indexing_technique: str | None = "economy", - latest_process_rule=None, -) -> Mock: - dataset = Mock(spec=Dataset) - dataset.id = dataset_id - dataset.tenant_id = tenant_id - dataset.data_source_type = data_source_type - dataset.indexing_technique = indexing_technique - dataset.latest_process_rule = latest_process_rule - dataset.get_latest_process_rule.return_value = latest_process_rule - dataset.embedding_model_provider = "provider" - dataset.embedding_model = "embedding-model" - dataset.summary_index_setting = None - dataset.retrieval_model = None - dataset.collection_binding_id = None - return dataset + doc_form: str = IndexStructureType.PARAGRAPH_INDEX, +) -> Dataset: + return Dataset( + id=dataset_id, + tenant_id=tenant_id, + data_source_type=data_source_type, + indexing_technique=indexing_technique, + chunk_structure=doc_form, + embedding_model_provider="provider", + embedding_model="embedding-model", + summary_index_setting=None, + retrieval_model=None, + collection_binding_id=None, + ) def _make_document( @@ -310,31 +317,29 @@ def _make_document( enabled: bool = True, archived: bool = False, indexing_status: str = "completed", - display_status: str = "available", -) -> Mock: - document = Mock(spec=Document) - document.id = document_id - document.dataset_id = dataset_id - document.tenant_id = tenant_id - document.batch = batch - document.doc_form = doc_form - document.word_count = word_count - document.name = name - document.enabled = enabled - document.archived = archived - document.indexing_status = indexing_status - document.display_status = display_status - document.data_source_type = "upload_file" - document.data_source_info = "{}" - document.completed_at = SimpleNamespace() - document.processing_started_at = "started" - document.parsing_completed_at = "parsed" - document.cleaning_completed_at = "cleaned" - document.splitting_completed_at = "split" - document.updated_at = None - document.created_from = None - document.dataset_process_rule_id = "process-rule-1" - return document +) -> Document: + return Document( + id=document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + batch=batch, + doc_form=doc_form, + word_count=word_count, + name=name, + enabled=enabled, + archived=archived, + indexing_status=indexing_status, + data_source_type="upload_file", + data_source_info="{}", + completed_at=SimpleNamespace(), + processing_started_at="started", + parsing_completed_at="parsed", + cleaning_completed_at="cleaned", + splitting_completed_at="split", + updated_at=None, + created_from=None, + dataset_process_rule_id="process-rule-1", + ) def _make_segment( @@ -347,21 +352,26 @@ def _make_segment( index_node_id: str = "node-1", dataset_id: str = "dataset-1", document_id: str = "doc-1", -) -> Mock: - segment = Mock(spec=DocumentSegment) +) -> DocumentSegment: + segment = DocumentSegment( + tenant_id="tenant-id", + dataset_id=dataset_id, + document_id=document_id, + position=1, + content=content, + word_count=word_count, + tokens=0, + created_by="account-id", + enabled=enabled, + keywords=keywords or [], + answer=None, + index_node_id=index_node_id, + disabled_at=None, + disabled_by=None, + status="completed", + error=None, + ) segment.id = segment_id - segment.dataset_id = dataset_id - segment.document_id = document_id - segment.content = content - segment.word_count = word_count - segment.enabled = enabled - segment.keywords = keywords or [] - segment.answer = None - segment.index_node_id = index_node_id - segment.disabled_at = None - segment.disabled_by = None - segment.status = "completed" - segment.error = None return segment diff --git a/api/tests/unit_tests/services/document_service_validation.py b/api/tests/unit_tests/services/document_service_validation.py index 8fad09e4d4c..26aca66bd31 100644 --- a/api/tests/unit_tests/services/document_service_validation.py +++ b/api/tests/unit_tests/services/document_service_validation.py @@ -109,6 +109,7 @@ This test suite follows a comprehensive testing strategy that covers: from unittest.mock import Mock, patch import pytest +from sqlalchemy.orm import Session from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.rag.entities import PreProcessingRule, Rule, Segmentation @@ -156,9 +157,9 @@ class DocumentValidationTestDataFactory: embedding_model_provider: str = "openai", embedding_model: str = "text-embedding-ada-002", **kwargs, - ) -> Mock: + ) -> Dataset: """ - Create a mock Dataset with specified attributes. + Create a Dataset with specified attributes. Args: dataset_id: Unique identifier for the dataset @@ -167,22 +168,20 @@ class DocumentValidationTestDataFactory: indexing_technique: Indexing technique embedding_model_provider: Embedding model provider embedding_model: Embedding model name - **kwargs: Additional attributes to set on the mock + **kwargs: Additional mapped attributes for the dataset Returns: - Mock object configured as a Dataset instance + Configured Dataset instance """ - dataset = Mock(spec=Dataset) - dataset.id = dataset_id - dataset.tenant_id = tenant_id - dataset.doc_form = doc_form - dataset.get_doc_form.return_value = doc_form - dataset.indexing_technique = indexing_technique - dataset.embedding_model_provider = embedding_model_provider - dataset.embedding_model = embedding_model - for key, value in kwargs.items(): - setattr(dataset, key, value) - return dataset + return Dataset( + id=dataset_id, + tenant_id=tenant_id, + chunk_structure=doc_form, + indexing_technique=indexing_technique, + embedding_model_provider=embedding_model_provider, + embedding_model=embedding_model, + **kwargs, + ) @staticmethod def create_knowledge_config_mock( @@ -205,14 +204,14 @@ class DocumentValidationTestDataFactory: Returns: Mock object configured as a KnowledgeConfig instance """ - config = Mock(spec=KnowledgeConfig) - config.data_source = data_source - config.process_rule = process_rule - config.doc_form = doc_form - config.indexing_technique = indexing_technique - for key, value in kwargs.items(): - setattr(config, key, value) - return config + return Mock( + spec=KnowledgeConfig, + data_source=data_source, + process_rule=process_rule, + doc_form=doc_form, + indexing_technique=indexing_technique, + **kwargs, + ) @staticmethod def create_data_source_mock( @@ -313,6 +312,10 @@ class TestDatasetServiceCheckDocForm: - Various form type combinations """ + @pytest.fixture(autouse=True) + def _bind_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + def test_check_doc_form_matching_forms_success(self): """ Test successful validation when form types match. @@ -328,10 +331,8 @@ class TestDatasetServiceCheckDocForm: # Arrange dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=IndexStructureType.PARAGRAPH_INDEX) doc_form = IndexStructureType.PARAGRAPH_INDEX - session = Mock() - # Act (should not raise) - DatasetService.check_doc_form(dataset, doc_form, session=session) + DatasetService.check_doc_form(dataset, doc_form, session=self.session) # Assert # No exception should be raised @@ -351,10 +352,8 @@ class TestDatasetServiceCheckDocForm: # Arrange dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=None) doc_form = IndexStructureType.PARAGRAPH_INDEX - session = Mock() - # Act (should not raise) - DatasetService.check_doc_form(dataset, doc_form, session=session) + DatasetService.check_doc_form(dataset, doc_form, session=self.session) # Assert # No exception should be raised @@ -374,11 +373,9 @@ class TestDatasetServiceCheckDocForm: # Arrange dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=IndexStructureType.PARAGRAPH_INDEX) doc_form = IndexStructureType.PARENT_CHILD_INDEX # Different form - session = Mock() - # Act & Assert with pytest.raises(ValueError, match="doc_form is different from the dataset doc_form"): - DatasetService.check_doc_form(dataset, doc_form, session=session) + DatasetService.check_doc_form(dataset, doc_form, session=self.session) def test_check_doc_form_different_form_types_error(self): """ @@ -394,11 +391,9 @@ class TestDatasetServiceCheckDocForm: # Arrange dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form="knowledge_card") doc_form = IndexStructureType.PARAGRAPH_INDEX # Different form - session = Mock() - # Act & Assert with pytest.raises(ValueError, match="doc_form is different from the dataset doc_form"): - DatasetService.check_doc_form(dataset, doc_form, session=session) + DatasetService.check_doc_form(dataset, doc_form, session=self.session) # ============================================================================ @@ -543,7 +538,7 @@ class TestDatasetServiceCheckDatasetModelSetting: error_description = "Provider token not initialized" mock_instance = Mock() - mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(description=error_description) + mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(error_description) mock_model_manager.return_value = mock_instance # Act & Assert @@ -665,7 +660,7 @@ class TestDatasetServiceCheckEmbeddingModelSetting: error_description = "Provider token not initialized" mock_instance = Mock() - mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(description=error_description) + mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(error_description) mock_model_manager.return_value = mock_instance # Act & Assert @@ -787,7 +782,7 @@ class TestDatasetServiceCheckRerankingModelSetting: error_description = "Provider token not initialized" mock_instance = Mock() - mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(description=error_description) + mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(error_description) mock_model_manager.return_value = mock_instance # Act & Assert @@ -1094,14 +1089,13 @@ class TestDocumentServiceDataSourceArgsValidate: def test_data_source_args_validate_missing_info_list_error(self): """ - Test error when info_list is missing. + Test current failure behavior when info_list is missing. - Verifies that when info_list is None, a ValueError is raised. + The validator currently dereferences info_list before its explicit + missing-value guard, so this path raises AttributeError. This test ensures: - - Missing info_list is rejected - - Error message is clear - - Error type is correct + - Missing info_list is rejected before any database access """ # Arrange data_source = Mock(spec=DataSource) @@ -1109,7 +1103,7 @@ class TestDocumentServiceDataSourceArgsValidate: knowledge_config = DocumentValidationTestDataFactory.create_knowledge_config_mock(data_source=data_source) # Act & Assert - with pytest.raises(ValueError, match="Data source info is required"): + with pytest.raises(AttributeError, match="data_source_type"): DocumentService.data_source_args_validate(knowledge_config) def test_data_source_args_validate_missing_file_info_error(self): diff --git a/api/tests/unit_tests/services/enterprise/test_account_deletion_sync.py b/api/tests/unit_tests/services/enterprise/test_account_deletion_sync.py index 0624b5ac778..fd631dc91ec 100644 --- a/api/tests/unit_tests/services/enterprise/test_account_deletion_sync.py +++ b/api/tests/unit_tests/services/enterprise/test_account_deletion_sync.py @@ -13,6 +13,7 @@ import pytest from redis import RedisError from sqlalchemy.orm import Session +from enums import DeploymentEdition from models.account import TenantAccountJoin from services.enterprise.account_deletion_sync import ( _queue_task, @@ -48,21 +49,21 @@ class TestSyncWorkspaceMemberRemoval: mock_queue.return_value = True yield mock_queue - def test_sync_workspace_member_removal_enterprise_enabled(self, mock_queue_task): + def test_sync_workspace_member_removal_enterprise_edition(self, mock_queue_task): workspace_id = str(uuid4()) member_id = str(uuid4()) with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE result = sync_workspace_member_removal(workspace_id=workspace_id, member_id=member_id, source="removed") assert result is True mock_queue_task.assert_called_once_with(workspace_id=workspace_id, member_id=member_id, source="removed") - def test_sync_workspace_member_removal_enterprise_disabled(self, mock_queue_task): + def test_sync_workspace_member_removal_non_enterprise_edition(self, mock_queue_task): with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY result = sync_workspace_member_removal( workspace_id=str(uuid4()), member_id=str(uuid4()), source="test_source" @@ -75,7 +76,7 @@ class TestSyncWorkspaceMemberRemoval: mock_queue_task.return_value = False with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE result = sync_workspace_member_removal( workspace_id=str(uuid4()), member_id=str(uuid4()), source="test_source" @@ -92,9 +93,9 @@ class TestSyncAccountDeletion: mock_queue.return_value = True yield mock_queue - def test_sync_account_deletion_enterprise_disabled(self, mock_queue_task, sqlite_session: Session) -> None: + def test_sync_account_deletion_non_enterprise_edition(self, mock_queue_task, sqlite_session: Session) -> None: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY result = sync_account_deletion(account_id=str(uuid4()), source="account_deleted", session=sqlite_session) @@ -111,7 +112,7 @@ class TestSyncAccountDeletion: sqlite_session.commit() with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE result = sync_account_deletion(account_id=account_id, source="account_deleted", session=sqlite_session) @@ -123,7 +124,7 @@ class TestSyncAccountDeletion: def test_sync_account_deletion_no_workspaces(self, sqlite_session: Session, mock_queue_task) -> None: with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE result = sync_account_deletion(account_id=str(uuid4()), source="account_deleted", session=sqlite_session) @@ -146,7 +147,7 @@ class TestSyncAccountDeletion: mock_queue_task.side_effect = queue_side_effect with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE result = sync_account_deletion(account_id=account_id, source="account_deleted", session=sqlite_session) @@ -164,7 +165,7 @@ class TestSyncAccountDeletion: mock_queue_task.return_value = False with patch("services.enterprise.account_deletion_sync.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE result = sync_account_deletion(account_id=account_id, source="account_deleted", session=sqlite_session) diff --git a/api/tests/unit_tests/services/enterprise/test_enterprise_service.py b/api/tests/unit_tests/services/enterprise/test_enterprise_service.py index f64c7233b9d..7ec5d1c01e6 100644 --- a/api/tests/unit_tests/services/enterprise/test_enterprise_service.py +++ b/api/tests/unit_tests/services/enterprise/test_enterprise_service.py @@ -10,6 +10,8 @@ from unittest.mock import patch import pytest +from enums import DeploymentEdition +from services.enterprise.base import MCPIdentityRefreshError, MCPNoRefreshTokenError, MCPTokenError from services.enterprise.enterprise_service import ( INVALID_LICENSE_CACHE_TTL, LICENSE_STATUS_CACHE_KEY, @@ -20,6 +22,8 @@ from services.enterprise.enterprise_service import ( WorkspacePermission, try_join_default_workspace, ) +from services.entities.feature_entities import LicenseStatus +from services.errors.enterprise import EnterpriseAPIError, EnterpriseAPIForbiddenError, EnterpriseAPIUnauthorizedError MODULE = "services.enterprise.enterprise_service" @@ -265,12 +269,12 @@ class TestJoinDefaultWorkspace: class TestTryJoinDefaultWorkspace: - def test_try_join_default_workspace_enterprise_disabled_noop(self): + def test_try_join_default_workspace_non_enterprise_edition_noop(self): with ( patch("services.enterprise.enterprise_service.dify_config") as mock_config, patch("services.enterprise.enterprise_service.EnterpriseService.join_default_workspace") as mock_join, ): - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY try_join_default_workspace("11111111-1111-1111-1111-111111111111") @@ -283,7 +287,7 @@ class TestTryJoinDefaultWorkspace: patch("services.enterprise.enterprise_service.dify_config") as mock_config, patch("services.enterprise.enterprise_service.EnterpriseService.join_default_workspace") as mock_join, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_join.return_value = DefaultWorkspaceJoinResult( workspace_id="22222222-2222-2222-2222-222222222222", joined=True, @@ -302,7 +306,7 @@ class TestTryJoinDefaultWorkspace: patch("services.enterprise.enterprise_service.dify_config") as mock_config, patch("services.enterprise.enterprise_service.EnterpriseService.join_default_workspace") as mock_join, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_join.return_value = DefaultWorkspaceJoinResult( workspace_id="", joined=False, @@ -321,7 +325,7 @@ class TestTryJoinDefaultWorkspace: patch("services.enterprise.enterprise_service.dify_config") as mock_config, patch("services.enterprise.enterprise_service.EnterpriseService.join_default_workspace") as mock_join, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_join.side_effect = Exception("network failure") # Should not raise @@ -331,7 +335,7 @@ class TestTryJoinDefaultWorkspace: def test_try_join_default_workspace_invalid_account_id_soft_fails(self): with patch("services.enterprise.enterprise_service.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE # Should not raise even though UUID parsing fails inside join_default_workspace try_join_default_workspace("not-a-uuid") @@ -347,21 +351,19 @@ _EE_SVC = "services.enterprise.enterprise_service" class TestGetCachedLicenseStatus: """Tests for EnterpriseService.get_cached_license_status.""" - def test_returns_none_when_enterprise_disabled(self): + def test_returns_none_outside_enterprise_edition(self): with patch(f"{_EE_SVC}.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY assert EnterpriseService.get_cached_license_status() is None def test_cache_hit_returns_license_status_enum(self): - from services.feature_service import LicenseStatus - with ( patch(f"{_EE_SVC}.dify_config") as mock_config, patch(f"{_EE_SVC}.redis_client") as mock_redis, patch.object(EnterpriseService, "get_info") as mock_get_info, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_redis.get.return_value = b"active" result = EnterpriseService.get_cached_license_status() @@ -371,14 +373,12 @@ class TestGetCachedLicenseStatus: mock_get_info.assert_not_called() def test_cache_miss_fetches_api_and_caches_valid_status(self): - from services.feature_service import LicenseStatus - with ( patch(f"{_EE_SVC}.dify_config") as mock_config, patch(f"{_EE_SVC}.redis_client") as mock_redis, patch.object(EnterpriseService, "get_info") as mock_get_info, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_redis.get.return_value = None mock_get_info.return_value = {"License": {"status": "active"}} @@ -390,14 +390,12 @@ class TestGetCachedLicenseStatus: ) def test_cache_miss_fetches_api_and_caches_invalid_status_with_short_ttl(self): - from services.feature_service import LicenseStatus - with ( patch(f"{_EE_SVC}.dify_config") as mock_config, patch(f"{_EE_SVC}.redis_client") as mock_redis, patch.object(EnterpriseService, "get_info") as mock_get_info, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_redis.get.return_value = None mock_get_info.return_value = {"License": {"status": "expired"}} @@ -409,14 +407,12 @@ class TestGetCachedLicenseStatus: ) def test_redis_read_failure_falls_through_to_api(self): - from services.feature_service import LicenseStatus - with ( patch(f"{_EE_SVC}.dify_config") as mock_config, patch(f"{_EE_SVC}.redis_client") as mock_redis, patch.object(EnterpriseService, "get_info") as mock_get_info, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_redis.get.side_effect = ConnectionError("redis down") mock_get_info.return_value = {"License": {"status": "active"}} @@ -426,14 +422,12 @@ class TestGetCachedLicenseStatus: mock_get_info.assert_called_once() def test_redis_write_failure_still_returns_status(self): - from services.feature_service import LicenseStatus - with ( patch(f"{_EE_SVC}.dify_config") as mock_config, patch(f"{_EE_SVC}.redis_client") as mock_redis, patch.object(EnterpriseService, "get_info") as mock_get_info, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_redis.get.return_value = None mock_redis.setex.side_effect = ConnectionError("redis down") mock_get_info.return_value = {"License": {"status": "expiring"}} @@ -448,7 +442,7 @@ class TestGetCachedLicenseStatus: patch(f"{_EE_SVC}.redis_client") as mock_redis, patch.object(EnterpriseService, "get_info") as mock_get_info, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_redis.get.return_value = None mock_get_info.side_effect = Exception("network failure") @@ -460,7 +454,7 @@ class TestGetCachedLicenseStatus: patch(f"{_EE_SVC}.redis_client") as mock_redis, patch.object(EnterpriseService, "get_info") as mock_get_info, ): - mock_config.ENTERPRISE_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE mock_redis.get.return_value = None mock_get_info.return_value = {} # no "License" key @@ -519,18 +513,12 @@ class TestIssueMCPToken: assert body["app_id"] == "app-uuid" def test_401_maps_to_identity_refresh_error(self): - from services.enterprise.base import MCPIdentityRefreshError - from services.errors.enterprise import EnterpriseAPIUnauthorizedError - with patch(f"{MODULE}.EnterpriseRequest") as req: req.send_request.side_effect = EnterpriseAPIUnauthorizedError("refresh rejected by IdP") with pytest.raises(MCPIdentityRefreshError, match="refresh rejected"): self._call() def test_428_maps_to_no_refresh_token_error(self): - from services.enterprise.base import MCPNoRefreshTokenError - from services.errors.enterprise import EnterpriseAPIError - with patch(f"{MODULE}.EnterpriseRequest") as req: # 428 PreconditionRequired is what EE returns when there's no # stored SSO refresh token for the user. @@ -539,34 +527,24 @@ class TestIssueMCPToken: self._call() def test_403_maps_to_identity_refresh_error_for_license(self): - from services.enterprise.base import MCPIdentityRefreshError - from services.errors.enterprise import EnterpriseAPIForbiddenError - with patch(f"{MODULE}.EnterpriseRequest") as req: req.send_request.side_effect = EnterpriseAPIForbiddenError("not licensed for MCP forwarding") with pytest.raises(MCPIdentityRefreshError, match="not licensed"): self._call() def test_other_status_maps_to_generic_token_error(self): - from services.enterprise.base import MCPTokenError - from services.errors.enterprise import EnterpriseAPIError - with patch(f"{MODULE}.EnterpriseRequest") as req: req.send_request.side_effect = EnterpriseAPIError("upstream 502", status_code=502) with pytest.raises(MCPTokenError, match="status=502"): self._call() def test_malformed_response_shape_raises_token_error(self): - from services.enterprise.base import MCPTokenError - with patch(f"{MODULE}.EnterpriseRequest") as req: req.send_request.return_value = "not-a-dict" with pytest.raises(MCPTokenError, match="invalid response shape"): self._call() def test_missing_token_field_raises_token_error(self): - from services.enterprise.base import MCPTokenError - with patch(f"{MODULE}.EnterpriseRequest") as req: req.send_request.return_value = {"expires_at": 1700000000} # no token with pytest.raises(MCPTokenError, match="missing or non-string token"): @@ -583,8 +561,6 @@ class TestIssueMCPToken: def test_bool_expires_at_is_rejected(self): """bool is a subclass of int — must NOT be accepted as expires_at.""" - from services.enterprise.base import MCPTokenError - with patch(f"{MODULE}.EnterpriseRequest") as req: req.send_request.return_value = {"token": "t", "expires_at": True} with pytest.raises(MCPTokenError, match="non-numeric expires_at"): diff --git a/api/tests/unit_tests/services/enterprise/test_rbac_service.py b/api/tests/unit_tests/services/enterprise/test_rbac_service.py index 96bd81aec6b..ca9f7d1764d 100644 --- a/api/tests/unit_tests/services/enterprise/test_rbac_service.py +++ b/api/tests/unit_tests/services/enterprise/test_rbac_service.py @@ -133,7 +133,7 @@ class TestRoles: call = _call_args(mock_send) assert call.method == "GET" assert call.endpoint == "/rbac/roles/item" - assert call.params == {"id": "role-1"} + assert call.params == {"billing_enabled": True, "id": "role-1"} def test_members_forwards_role_id_and_pagination(self, mock_send: MagicMock): mock_send.return_value = { @@ -636,6 +636,7 @@ class TestMyPermissions: assert not any(key.startswith("billing.") for key in out.workspace.permission_keys) if role == "editor": assert "app.acl.log_and_annotation" in out.app.default_permission_keys + assert "app.acl.deploy" not in out.app.default_permission_keys @pytest.mark.parametrize( ("role", "expected_snippet_keys"), @@ -743,6 +744,7 @@ class TestMemberRoles: assert "snippets.create_and_modify" in out.roles[0].permission_keys assert "app.acl.preview" in out.roles[0].permission_keys assert "dataset.acl.preview" in out.roles[0].permission_keys + assert "app.acl.deploy" not in out.roles[0].permission_keys def test_replace(self, mock_send: MagicMock, sqlite_session: Session): mock_send.return_value = {"account_id": "acct-2", "roles": []} @@ -838,7 +840,10 @@ class TestResourcePermissions: 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"]}, + { + "resource_id": "app-1", + "permission_keys": ["app.acl.view_layout", "app.acl.edit", "app.acl.deploy"], + }, {"resource_id": "app-2", "permission_keys": []}, ] } @@ -853,7 +858,7 @@ class TestResourcePermissions: assert call.endpoint == "/rbac/apps/permission-keys/batch" assert call.json == {"app_ids": ["app-1", "app-2"]} assert out == { - "app-1": ["app.acl.view_layout", "app.acl.edit"], + "app-1": ["app.acl.view_layout", "app.acl.edit", "app.acl.deploy"], "app-2": [], } @@ -874,6 +879,7 @@ class TestResourcePermissions: "app-1": svc._LEGACY_APP_EDITOR_KEYS, "app-2": svc._LEGACY_APP_EDITOR_KEYS, } + assert all("app.acl.deploy" not in permission_keys for permission_keys in out.values()) def test_dataset_permissions_batch_get(self, mock_send: MagicMock, sqlite_session: Session): mock_send.return_value = { diff --git a/api/tests/unit_tests/services/hit_service.py b/api/tests/unit_tests/services/hit_service.py index 2a456dc4b9d..6ef10ad0726 100644 --- a/api/tests/unit_tests/services/hit_service.py +++ b/api/tests/unit_tests/services/hit_service.py @@ -16,7 +16,7 @@ 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 import Account, Tenant from models.dataset import Dataset, DatasetQuery from services.hit_testing_service import HitTestingService @@ -41,27 +41,27 @@ class HitTestingTestDataFactory: provider: str = "vendor", retrieval_model: dict[str, Any] | None = None, **kwargs, - ) -> Mock: + ) -> Dataset: """ - Create a mock dataset with specified attributes. + Create a dataset with specified attributes. Args: dataset_id: Unique identifier for the dataset tenant_id: Tenant identifier provider: Dataset provider (vendor, external, etc.) retrieval_model: Optional retrieval model configuration - **kwargs: Additional attributes to set on the mock + **kwargs: Additional mapped attributes for the dataset Returns: - Mock object configured as a Dataset instance + Configured Dataset instance """ - dataset = Mock(spec=Dataset) - dataset.id = dataset_id - dataset.tenant_id = tenant_id - dataset.provider = provider - dataset.retrieval_model = retrieval_model - for key, value in kwargs.items(): - setattr(dataset, key, value) + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + provider=provider, + retrieval_model=retrieval_model, + **kwargs, + ) return dataset @staticmethod @@ -69,7 +69,7 @@ class HitTestingTestDataFactory: user_id: str = "user-789", tenant_id: str = "tenant-123", **kwargs, - ) -> Mock: + ) -> Account: """ Create a mock user (Account) with specified attributes. @@ -81,12 +81,10 @@ class HitTestingTestDataFactory: Returns: Mock object configured as an Account instance """ - user = Mock(spec=Account) + user = Account(name="Test User", email="test@example.com", **kwargs) user.id = user_id - user.current_tenant_id = tenant_id - user.name = "Test User" - for key, value in kwargs.items(): - setattr(user, key, value) + user._current_tenant = Tenant(name="Test Tenant") + user._current_tenant.id = tenant_id return user @staticmethod @@ -106,12 +104,7 @@ class HitTestingTestDataFactory: Returns: Mock object configured as a Document instance """ - document = Mock(spec=Document) - document.page_content = content - document.metadata = metadata or {} - for key, value in kwargs.items(): - setattr(document, key, value) - return document + return Mock(spec=Document, page_content=content, metadata=metadata or {}, **kwargs) @staticmethod def create_retrieval_record_mock( diff --git a/api/tests/unit_tests/services/openapi/test_mint_policy.py b/api/tests/unit_tests/services/openapi/test_mint_policy.py index 7409a064a97..80425725f39 100644 --- a/api/tests/unit_tests/services/openapi/test_mint_policy.py +++ b/api/tests/unit_tests/services/openapi/test_mint_policy.py @@ -10,6 +10,7 @@ from __future__ import annotations import pytest +from enums import DeploymentEdition from libs.oauth_bearer import MINTABLE_PROFILES, Scope, SubjectType from services.openapi.mint_policy import MintPolicyViolation, validate_mint_policy @@ -84,7 +85,7 @@ def test_license_required_decorator_skips_on_ce(): return "ok" with patch("services.openapi.license_gate.dify_config") as cfg: - cfg.ENTERPRISE_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY assert view() == "ok" @@ -103,7 +104,7 @@ def test_license_required_decorator_403_on_invalid_ee_license(): patch("services.openapi.license_gate.dify_config") as cfg, patch("services.openapi.license_gate._is_license_valid", return_value=False), ): - cfg.ENTERPRISE_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE with pytest.raises(Forbidden) as exc: view() assert "license_required" in exc.value.description @@ -122,5 +123,5 @@ def test_license_required_decorator_passes_on_valid_ee_license(): patch("services.openapi.license_gate.dify_config") as cfg, patch("services.openapi.license_gate._is_license_valid", return_value=True), ): - cfg.ENTERPRISE_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.ENTERPRISE assert view() == "ok" diff --git a/api/tests/unit_tests/services/plugin/conftest.py b/api/tests/unit_tests/services/plugin/conftest.py index cb300584941..5345db65dc9 100644 --- a/api/tests/unit_tests/services/plugin/conftest.py +++ b/api/tests/unit_tests/services/plugin/conftest.py @@ -6,7 +6,7 @@ from unittest.mock import MagicMock import pytest -from services.feature_service import PluginInstallationScope +from services.entities.feature_entities import PluginInstallationScope def make_features( diff --git a/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py b/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py index e66bb3fff04..0296f16dd2c 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_auto_upgrade_service.py @@ -1,89 +1,105 @@ import logging from types import SimpleNamespace from unittest.mock import MagicMock, patch +from uuid import uuid4 import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session from models.account import ( TenantPluginAutoUpgradeCategory, TenantPluginAutoUpgradeMode, + TenantPluginAutoUpgradeStrategy, TenantPluginAutoUpgradeStrategySetting, ) +from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService MODULE = "services.plugin.plugin_auto_upgrade_service" PLUGIN_CATEGORY = TenantPluginAutoUpgradeCategory.TOOL +STRATEGY_MODELS = (TenantPluginAutoUpgradeStrategy,) -def _patched_session(): - """Return a mock SQLAlchemy session for service calls.""" - session = MagicMock() - return session +def _strategy( + tenant_id: str, + *, + category: TenantPluginAutoUpgradeCategory = PLUGIN_CATEGORY, + setting: TenantPluginAutoUpgradeStrategySetting = TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, + mode: TenantPluginAutoUpgradeMode = TenantPluginAutoUpgradeMode.EXCLUDE, + exclude: list[str] | None = None, + include: list[str] | None = None, + upgrade_time: int = 0, +) -> TenantPluginAutoUpgradeStrategy: + return TenantPluginAutoUpgradeStrategy( + tenant_id=tenant_id, + category=category, + strategy_setting=setting, + upgrade_time_of_day=upgrade_time, + upgrade_mode=mode, + exclude_plugins=exclude or [], + include_plugins=include or [], + ) class TestGetStrategy: - def test_returns_strategy_when_found(self): - session = _patched_session() - strategy = MagicMock() - session.scalar.return_value = strategy + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_returns_strategy_when_found(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) + strategy = _strategy(tenant_id) + sqlite_session.add(strategy) + sqlite_session.commit() - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - - result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY, session=session) + result = PluginAutoUpgradeService.get_strategy(tenant_id, PLUGIN_CATEGORY, session=sqlite_session) assert result is strategy - def test_returns_none_when_not_found(self): - session = _patched_session() - session.scalar.return_value = None - - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - - result = PluginAutoUpgradeService.get_strategy("t1", PLUGIN_CATEGORY, session=session) - - assert result is None + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_returns_none_when_not_found(self, sqlite_session: Session) -> None: + assert PluginAutoUpgradeService.get_strategy(str(uuid4()), PLUGIN_CATEGORY, session=sqlite_session) is None class TestChangeStrategy: - def test_creates_new_strategy(self): - session = _patched_session() - session.scalar.return_value = None - - with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: - strat_cls.return_value = MagicMock() - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - - result = PluginAutoUpgradeService.change_strategy( - "t1", - TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, - 3, - TenantPluginAutoUpgradeMode.ALL, - [], - [], - category=PLUGIN_CATEGORY, - session=session, - ) - - assert result is True - session.add.assert_called_once() - - def test_updates_existing_strategy(self): - session = _patched_session() - existing = MagicMock() - session.scalar.return_value = existing - - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_creates_new_strategy(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) result = PluginAutoUpgradeService.change_strategy( - "t1", + tenant_id, + TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, + 3, + TenantPluginAutoUpgradeMode.ALL, + [], + [], + category=PLUGIN_CATEGORY, + session=sqlite_session, + ) + + strategy = sqlite_session.scalar(select(TenantPluginAutoUpgradeStrategy)) + assert result is True + assert strategy is not None + assert strategy.tenant_id == tenant_id + assert strategy.upgrade_time_of_day == 3 + assert strategy.upgrade_mode == TenantPluginAutoUpgradeMode.ALL + + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_updates_existing_strategy(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) + existing = _strategy(tenant_id) + sqlite_session.add(existing) + sqlite_session.commit() + + result = PluginAutoUpgradeService.change_strategy( + tenant_id, TenantPluginAutoUpgradeStrategySetting.LATEST, 5, TenantPluginAutoUpgradeMode.PARTIAL, ["p1"], ["p2"], category=PLUGIN_CATEGORY, - session=session, + session=sqlite_session, ) + sqlite_session.refresh(existing) assert result is True assert existing.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST assert existing.upgrade_time_of_day == 5 @@ -93,157 +109,115 @@ class TestChangeStrategy: class TestExcludePlugin: - def test_creates_default_strategy_when_none_exists(self): - session = _patched_session() - session.scalar.return_value = None + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_creates_default_strategy_when_none_exists(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) - with ( - patch(f"{MODULE}.select"), - patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy"), - ): - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - - result = PluginAutoUpgradeService.exclude_plugin( - "t1", - "plugin-1", - PLUGIN_CATEGORY, - session=session, - ) + result = PluginAutoUpgradeService.exclude_plugin(tenant_id, "plugin-1", PLUGIN_CATEGORY, session=sqlite_session) + strategy = sqlite_session.scalar(select(TenantPluginAutoUpgradeStrategy)) assert result is True - session.add.assert_called_once() + assert strategy is not None + assert strategy.exclude_plugins == ["plugin-1"] - def test_appends_to_exclude_list_in_exclude_mode(self): - session = _patched_session() - existing = MagicMock() - existing.upgrade_mode = TenantPluginAutoUpgradeMode.EXCLUDE - existing.exclude_plugins = ["p-existing"] - session.scalar.return_value = existing + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_appends_to_exclude_list_in_exclude_mode(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) + existing = _strategy(tenant_id, exclude=["p-existing"]) + sqlite_session.add(existing) + sqlite_session.commit() - with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: - strat_cls.UpgradeMode.EXCLUDE = "exclude" - strat_cls.UpgradeMode.PARTIAL = "partial" - strat_cls.UpgradeMode.ALL = "all" - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + PluginAutoUpgradeService.exclude_plugin(tenant_id, "p-new", PLUGIN_CATEGORY, session=sqlite_session) - result = PluginAutoUpgradeService.exclude_plugin("t1", "p-new", PLUGIN_CATEGORY, session=session) - - assert result is True + sqlite_session.refresh(existing) assert existing.exclude_plugins == ["p-existing", "p-new"] - def test_removes_from_include_list_in_partial_mode(self): - session = _patched_session() - existing = MagicMock() - existing.upgrade_mode = TenantPluginAutoUpgradeMode.PARTIAL - existing.include_plugins = ["p1", "p2"] - session.scalar.return_value = existing + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_removes_from_include_list_in_partial_mode(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) + existing = _strategy(tenant_id, mode=TenantPluginAutoUpgradeMode.PARTIAL, include=["p1", "p2"]) + sqlite_session.add(existing) + sqlite_session.commit() - with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: - strat_cls.UpgradeMode.EXCLUDE = "exclude" - strat_cls.UpgradeMode.PARTIAL = "partial" - strat_cls.UpgradeMode.ALL = "all" - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + PluginAutoUpgradeService.exclude_plugin(tenant_id, "p1", PLUGIN_CATEGORY, session=sqlite_session) - result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) - - assert result is True + sqlite_session.refresh(existing) assert existing.include_plugins == ["p2"] - def test_switches_to_exclude_mode_from_all(self): - session = _patched_session() - existing = MagicMock() - existing.upgrade_mode = TenantPluginAutoUpgradeMode.ALL - session.scalar.return_value = existing + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_switches_to_exclude_mode_from_all(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) + existing = _strategy(tenant_id, mode=TenantPluginAutoUpgradeMode.ALL) + sqlite_session.add(existing) + sqlite_session.commit() - with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: - strat_cls.UpgradeMode.EXCLUDE = "exclude" - strat_cls.UpgradeMode.PARTIAL = "partial" - strat_cls.UpgradeMode.ALL = "all" - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService + PluginAutoUpgradeService.exclude_plugin(tenant_id, "p1", PLUGIN_CATEGORY, session=sqlite_session) - result = PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) - - assert result is True + sqlite_session.refresh(existing) assert existing.upgrade_mode == TenantPluginAutoUpgradeMode.EXCLUDE assert existing.exclude_plugins == ["p1"] - def test_no_duplicate_in_exclude_list(self): - session = _patched_session() - existing = MagicMock() - existing.upgrade_mode = TenantPluginAutoUpgradeMode.EXCLUDE - existing.exclude_plugins = ["p1"] - session.scalar.return_value = existing + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_no_duplicate_in_exclude_list(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) + existing = _strategy(tenant_id, exclude=["p1"]) + sqlite_session.add(existing) + sqlite_session.commit() - with patch(f"{MODULE}.select"), patch(f"{MODULE}.TenantPluginAutoUpgradeStrategy") as strat_cls: - strat_cls.UpgradeMode.EXCLUDE = "exclude" - strat_cls.UpgradeMode.PARTIAL = "partial" - strat_cls.UpgradeMode.ALL = "all" - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - - PluginAutoUpgradeService.exclude_plugin("t1", "p1", PLUGIN_CATEGORY, session=session) + PluginAutoUpgradeService.exclude_plugin(tenant_id, "p1", PLUGIN_CATEGORY, session=sqlite_session) + sqlite_session.refresh(existing) assert existing.exclude_plugins == ["p1"] class TestBackfillStrategyCategories: - def test_creates_default_missing_categories_without_fetching_daemon(self): - session = _patched_session() - tool_strategy = SimpleNamespace( - category=TenantPluginAutoUpgradeCategory.TOOL, - strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, - upgrade_time_of_day=0, - upgrade_mode=TenantPluginAutoUpgradeMode.EXCLUDE, - exclude_plugins=[], - include_plugins=[], - ) - session.scalars.return_value.all.return_value = [tool_strategy] + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_creates_default_missing_categories_without_fetching_daemon(self, sqlite_session: Session) -> None: + tenant_id = str(uuid4()) + tool_strategy = _strategy(tenant_id) + sqlite_session.add(tool_strategy) + sqlite_session.commit() installer = MagicMock() with patch(f"{MODULE}.PluginInstaller", return_value=installer): - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - - result = PluginAutoUpgradeService.backfill_strategy_categories("t1", session=session) - expected_time = PluginAutoUpgradeService.default_upgrade_time_of_day("t1") + result = PluginAutoUpgradeService.backfill_strategy_categories(tenant_id, session=sqlite_session) + expected_time = PluginAutoUpgradeService.default_upgrade_time_of_day(tenant_id) + strategies = list(sqlite_session.scalars(select(TenantPluginAutoUpgradeStrategy)).all()) assert result.created_count == len(TenantPluginAutoUpgradeCategory) - 1 assert result.normalized is False installer.list_plugins.assert_not_called() + assert len(strategies) == len(TenantPluginAutoUpgradeCategory) assert tool_strategy.upgrade_time_of_day == expected_time - created_strategies = [call.args[0] for call in session.add.call_args_list] model_strategy = next( - strategy for strategy in created_strategies if strategy.category == TenantPluginAutoUpgradeCategory.MODEL + strategy for strategy in strategies if strategy.category == TenantPluginAutoUpgradeCategory.MODEL ) assert model_strategy.strategy_setting == TenantPluginAutoUpgradeStrategySetting.LATEST assert model_strategy.upgrade_time_of_day == expected_time - def test_default_upgrade_time_is_aligned_to_fifteen_minutes(self): - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - - default_time = PluginAutoUpgradeService.default_upgrade_time_of_day("t1") - + def test_default_upgrade_time_is_aligned_to_fifteen_minutes(self) -> None: + default_time = PluginAutoUpgradeService.default_upgrade_time_of_day(str(uuid4())) assert default_time % (15 * 60) == 0 assert 0 <= default_time < 24 * 60 * 60 - def test_creates_missing_categories_and_splits_known_plugins(self, caplog: pytest.LogCaptureFixture): - session = _patched_session() - tool_strategy = SimpleNamespace( - category=TenantPluginAutoUpgradeCategory.TOOL, - strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, - upgrade_time_of_day=0, - upgrade_mode=TenantPluginAutoUpgradeMode.EXCLUDE, - exclude_plugins=["tool-plugin", "model-plugin", "unknown-plugin"], - include_plugins=["model-plugin", "tool-plugin"], + @pytest.mark.parametrize("sqlite_session", [STRATEGY_MODELS], indirect=True) + def test_creates_missing_categories_and_splits_known_plugins( + self, sqlite_session: Session, caplog: pytest.LogCaptureFixture + ) -> None: + tenant_id = str(uuid4()) + tool_strategy = _strategy( + tenant_id, + exclude=["tool-plugin", "model-plugin", "unknown-plugin"], + include=["model-plugin", "tool-plugin"], ) - model_strategy = SimpleNamespace( + model_strategy = _strategy( + tenant_id, category=TenantPluginAutoUpgradeCategory.MODEL, - strategy_setting=TenantPluginAutoUpgradeStrategySetting.FIX_ONLY, - upgrade_time_of_day=0, - upgrade_mode=TenantPluginAutoUpgradeMode.EXCLUDE, - exclude_plugins=["tool-plugin", "model-plugin", "unknown-plugin"], - include_plugins=["model-plugin", "tool-plugin"], + exclude=["tool-plugin", "model-plugin", "unknown-plugin"], + include=["model-plugin", "tool-plugin"], ) - session.scalars.return_value.all.return_value = [tool_strategy, model_strategy] - + sqlite_session.add_all([tool_strategy, model_strategy]) + sqlite_session.commit() installed_plugins = [ SimpleNamespace( plugin_id="tool-plugin", @@ -261,18 +235,17 @@ class TestBackfillStrategyCategories: patch(f"{MODULE}.PluginInstaller", return_value=installer), caplog.at_level(logging.WARNING, logger=MODULE), ): - from services.plugin.plugin_auto_upgrade_service import PluginAutoUpgradeService - - result = PluginAutoUpgradeService.backfill_strategy_categories("t1", session=session) + result = PluginAutoUpgradeService.backfill_strategy_categories(tenant_id, session=sqlite_session) + strategies = list(sqlite_session.scalars(select(TenantPluginAutoUpgradeStrategy)).all()) assert result.created_count == len(TenantPluginAutoUpgradeCategory) - 2 assert result.normalized is True - assert session.add.call_count == len(TenantPluginAutoUpgradeCategory) - 2 + assert len(strategies) == len(TenantPluginAutoUpgradeCategory) assert tool_strategy.exclude_plugins == ["tool-plugin"] assert tool_strategy.include_plugins == ["tool-plugin"] assert model_strategy.exclude_plugins == ["model-plugin"] assert model_strategy.include_plugins == ["model-plugin"] assert ( "Skipped unknown plugin IDs while backfilling plugin auto-upgrade strategies: " - "tenant_id=t1, field=exclude_plugins, plugin_ids=['unknown-plugin']" in caplog.messages + f"tenant_id={tenant_id}, field=exclude_plugins, plugin_ids=['unknown-plugin']" in caplog.messages ) diff --git a/api/tests/unit_tests/services/plugin/test_plugin_migration.py b/api/tests/unit_tests/services/plugin/test_plugin_migration.py index 27b9749bf11..58f7d8cef16 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_migration.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_migration.py @@ -1,9 +1,12 @@ import json +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest from pytest_mock import MockerFixture +from sqlalchemy.orm import Session +from models.model import App, AppMode from services.plugin.plugin_migration import PluginMigration MIGRATION_MODULE = "services.plugin.plugin_migration" @@ -33,24 +36,28 @@ def test_fetch_latest_package_identifier_calls_marketplace_when_enabled(mocker: assert result == "langgenius/openai:1.0.0@abc" -def test_extract_app_tables_checks_agent_mode_with_its_session(mocker: MockerFixture) -> None: - app = mocker.MagicMock(app_model_config_id=None, mode="chat") - app.is_agent_with_session.return_value = False - apps_result = mocker.MagicMock() - apps_result.all.return_value = [app] - configs_result = mocker.MagicMock() - configs_result.all.return_value = [] - session = mocker.MagicMock() - session.scalars.side_effect = [apps_result, configs_result] - session_context = mocker.MagicMock() - session_context.__enter__.return_value = session - mocker.patch(f"{MIGRATION_MODULE}.Session", return_value=session_context) - mocker.patch(f"{MIGRATION_MODULE}.db") +def test_extract_app_tables_checks_agent_mode_with_its_session(mocker: MockerFixture, sqlite_session: Session) -> None: + app = App( + id="app-1", + tenant_id="tenant-1", + name="Chat app", + description="", + mode=AppMode.CHAT, + icon_type=None, + icon="", + icon_background=None, + enable_site=False, + enable_api=False, + created_by="account-1", + max_active_requests=0, + ) + sqlite_session.add(app) + sqlite_session.commit() + mocker.patch(f"{MIGRATION_MODULE}.db", SimpleNamespace(engine=sqlite_session.get_bind())) result = PluginMigration.extract_app_tables("tenant-1") assert result == [] - app.is_agent_with_session.assert_called_once_with(session=session) class TestHandlePluginInstanceInstall: diff --git a/api/tests/unit_tests/services/plugin/test_plugin_parameter_service.py b/api/tests/unit_tests/services/plugin/test_plugin_parameter_service.py index b51e13013a4..916217ce0b1 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_parameter_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_parameter_service.py @@ -140,7 +140,7 @@ class TestGetDynamicSelectOptionsTrigger: result = PluginParameterService.get_dynamic_select_options( tenant_id="t1", user_id="u1", - plugin_id="p1", + plugin_id="org/plugin", provider="github", action="on_push", parameter="branch", @@ -149,6 +149,11 @@ class TestGetDynamicSelectOptionsTrigger: ) assert result == ["opt"] + builder_call = mock_builder_svc.get_subscription_builder.call_args.kwargs + assert builder_call["tenant_id"] == "t1" + assert builder_call["user_id"] == "u1" + assert str(builder_call["provider_id"]) == "org/plugin/github" + assert builder_call["subscription_builder_id"] == "builder-1" @patch("services.plugin.plugin_parameter_service.DynamicSelectClient") @patch("services.plugin.plugin_parameter_service.TriggerProviderService") @@ -166,7 +171,7 @@ class TestGetDynamicSelectOptionsTrigger: result = PluginParameterService.get_dynamic_select_options( tenant_id="t1", user_id="u1", - plugin_id="p1", + plugin_id="org/plugin", provider="github", action="on_push", parameter="branch", @@ -186,7 +191,7 @@ class TestGetDynamicSelectOptionsTrigger: PluginParameterService.get_dynamic_select_options( tenant_id="t1", user_id="u1", - plugin_id="p1", + plugin_id="org/plugin", provider="github", action="on_push", parameter="branch", diff --git a/api/tests/unit_tests/services/plugin/test_plugin_service.py b/api/tests/unit_tests/services/plugin/test_plugin_service.py index 278898926b9..b0b727c3466 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_service.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_service.py @@ -7,48 +7,58 @@ import pytest import zstandard from pydantic import TypeAdapter from redis import RedisError +from sqlalchemy.orm import Session +from core.helper.model_provider_cache import ProviderCredentialsCacheType from core.plugin.entities.plugin import PluginCategory, PluginInstallationSource -from core.plugin.entities.plugin_daemon import PluginInstallTask, PluginInstallTaskStatus, PluginModelProviderEntity +from core.plugin.entities.plugin_daemon import ( + PluginInstallTask, + PluginInstallTaskStatus, + PluginModelProviderDeclaration, + PluginModelProviderEntity, +) +from core.provider_manager import ProviderConfigurationCacheSource, ProviderManager +from enums import DeploymentEdition from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.provider_entities import ConfigurateMethod, ProviderEntity +from models.provider import Provider, ProviderCredential, ProviderType, TenantPreferredModelProvider +from services.entities.feature_entities import PluginInstallationPermissionModel, PluginInstallationScope MODULE = "core.plugin.plugin_service" +TENANT_ID = "11111111-1111-1111-1111-111111111111" +OTHER_TENANT_ID = "22222222-2222-2222-2222-222222222222" +USER_ID = "33333333-3333-3333-3333-333333333333" -class _FakeSession: - def __init__(self) -> None: - self.execute = Mock() - self.scalars = Mock(return_value=SimpleNamespace(all=Mock(return_value=[]))) - - def __enter__(self) -> "_FakeSession": - return self - - def __exit__(self, exc_type, exc, traceback) -> None: - return None - - def begin(self) -> "_FakeSession": - return self - - -def _build_provider_entity(provider: str = "openai") -> ProviderEntity: - return ProviderEntity( +def _build_provider_entity( + provider: str = "openai", + installation_source: PluginInstallationSource | None = PluginInstallationSource.Marketplace, +) -> PluginModelProviderDeclaration: + return PluginModelProviderDeclaration( provider=f"langgenius/{provider}/{provider}", + plugin_unique_identifier=f"langgenius/{provider}:1.0.0@checksum", + installation_source=installation_source, label=I18nObject(en_US=provider.title()), supported_model_types=[], configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL], ) -def _build_plugin_model_provider(*, tenant_id: str = "tenant-1", provider: str = "openai") -> PluginModelProviderEntity: +def _build_plugin_model_provider( + *, + tenant_id: str = "tenant-1", + provider: str = "openai", + installation_source: PluginInstallationSource | None = PluginInstallationSource.Marketplace, +) -> PluginModelProviderEntity: return PluginModelProviderEntity( id=uuid.uuid4().hex, created_at=datetime.datetime.now(), updated_at=datetime.datetime.now(), provider=provider, tenant_id=tenant_id, - plugin_unique_identifier=f"langgenius/{provider}/{provider}", + plugin_unique_identifier=f"langgenius/{provider}:1.0.0@checksum", plugin_id=f"langgenius/{provider}", + installation_source=installation_source, declaration=ProviderEntity( provider=provider, label=I18nObject(en_US=provider.title()), @@ -144,7 +154,7 @@ class TestPluginModelProviderCache: """Large provider metadata payloads are compressed before being stored in Redis.""" large_provider = _build_provider_entity() large_provider.label = I18nObject(en_US="OpenAI " * 10000) - raw_payload = TypeAdapter(list[ProviderEntity]).dump_json([large_provider]) + raw_payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json([large_provider]) cache_key = _provider_cache_key("tenant-1", 0) with ( @@ -171,7 +181,7 @@ class TestPluginModelProviderCache: """Compressed tenant cache entries are decoded before provider schema validation.""" cached_provider = _build_provider_entity() cached_provider.label = I18nObject(en_US="OpenAI " * 10000) - cached_payload = TypeAdapter(list[ProviderEntity]).dump_json([cached_provider]) + cached_payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json([cached_provider]) generation_key = _provider_generation_key("tenant-1") cache_key = _provider_cache_key("tenant-1", 0) @@ -189,6 +199,7 @@ class TestPluginModelProviderCache: result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + assert result[0].plugin_unique_identifier == "langgenius/openai:1.0.0@checksum" assert result[0].label.en_us == "OpenAI " * 10000 client.fetch_model_providers.assert_not_called() redis_client.setex.assert_not_called() @@ -196,9 +207,9 @@ class TestPluginModelProviderCache: redis_client.mget.assert_called_once_with([cache_key]) def test_fetch_plugin_model_providers_returns_cached_provider_without_calling_daemon(self) -> None: - """A valid tenant cache entry is reused across runtime calls without plugin daemon access.""" - cached_provider = _build_provider_entity() - cached_payload = TypeAdapter(list[ProviderEntity]).dump_json([cached_provider]) + """A cached package source remains available to the system configuration guard.""" + cached_provider = _build_provider_entity(installation_source=PluginInstallationSource.Package) + cached_payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json([cached_provider]) generation_key = _provider_generation_key("tenant-1") cache_key = _provider_cache_key("tenant-1", 0) @@ -212,13 +223,33 @@ class TestPluginModelProviderCache: result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) assert [provider.provider for provider in result] == ["langgenius/openai/openai"] + assert result[0].installation_source == PluginInstallationSource.Package + provider_manager = ProviderManager(model_runtime=Mock()) + with ( + patch( + "core.provider_manager.ext_hosting_provider.hosting_configuration.provider_map", + {result[0].provider: SimpleNamespace(enabled=True)}, + ), + patch("core.plugin.plugin_service.PluginService.is_plugin_verified") as is_plugin_verified, + ): + system_configuration = provider_manager._to_system_configuration("tenant-1", result[0], []) + + assert system_configuration.enabled is False + is_plugin_verified.assert_not_called() client.fetch_model_providers.assert_not_called() redis_client.setex.assert_not_called() redis_client.get.assert_called_once_with(generation_key) redis_client.mget.assert_called_once_with([cache_key]) - def test_fetch_plugin_model_providers_deletes_invalid_cache_and_refetches(self) -> None: - """Invalid generation-scoped cache payloads are removed before falling back to the daemon.""" + def test_fetch_plugin_model_providers_invalidates_legacy_cache_without_plugin_identity(self) -> None: + """Legacy provider cache entries are refreshed before they can reach system configuration.""" + legacy_provider = ProviderEntity( + provider="langgenius/openai/openai", + label=I18nObject(en_US="OpenAI"), + supported_model_types=[], + configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL], + ) + legacy_payload = TypeAdapter(list[ProviderEntity]).dump_json([legacy_provider]) generation_key = _provider_generation_key("tenant-1") cache_key = _provider_cache_key("tenant-1", 0) with ( @@ -226,7 +257,7 @@ class TestPluginModelProviderCache: patch(f"{MODULE}.dify_config") as mock_config, ): redis_client.get.side_effect = [None, None, None] - redis_client.mget.side_effect = [["not-json"], [None]] + redis_client.mget.side_effect = [[legacy_payload], [None]] mock_config.PLUGIN_MODEL_PROVIDERS_CACHE_TTL = 86400 client = Mock() client.fetch_model_providers.return_value = [_build_plugin_model_provider()] @@ -246,6 +277,51 @@ class TestPluginModelProviderCache: call([cache_key]), ] + def test_fetch_plugin_model_providers_bypasses_redis_when_cache_disabled(self) -> None: + """With the cache disabled the daemon is the only source, and Redis is never touched.""" + with patch(f"{MODULE}.redis_client") as redis_client, patch(f"{MODULE}.dify_config") as config: + config.PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED = False + client = Mock() + client.fetch_model_providers.return_value = [_build_plugin_model_provider()] + + from core.plugin.plugin_service import PluginService + + first = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + second = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + assert [provider.provider for provider in first] == ["langgenius/openai/openai"] + assert [provider.provider for provider in second] == ["langgenius/openai/openai"] + assert first[0].plugin_unique_identifier == "langgenius/openai:1.0.0@checksum" + assert client.fetch_model_providers.call_count == 2 + redis_client.get.assert_not_called() + redis_client.mget.assert_not_called() + redis_client.setex.assert_not_called() + redis_client.lock.assert_not_called() + + def test_fetch_plugin_model_providers_resolves_missing_installation_source(self) -> None: + provider = _build_plugin_model_provider(installation_source=None) + installation = SimpleNamespace( + plugin_unique_identifier=provider.plugin_unique_identifier, + source=PluginInstallationSource.Package, + ) + + from core.plugin.plugin_service import PluginService + + with ( + patch(f"{MODULE}.dify_config") as config, + patch.object( + PluginService, "list_installations_from_ids", return_value=[installation] + ) as list_installations, + ): + config.PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED = False + client = Mock() + client.fetch_model_providers.return_value = [provider] + + result = PluginService.fetch_plugin_model_providers(tenant_id="tenant-1", client=client) + + list_installations.assert_called_once_with("tenant-1", [provider.plugin_id]) + assert result[0].installation_source == PluginInstallationSource.Package + def test_fetch_plugin_model_providers_refetches_when_cache_read_fails(self) -> None: """Redis read failures do not block provider discovery for the tenant.""" with patch(f"{MODULE}.redis_client") as redis_client: @@ -296,7 +372,7 @@ class TestPluginModelProviderCache: def test_fetch_plugin_model_providers_waits_for_concurrent_refresh_cache_fill(self) -> None: """A cache miss waits for the active tenant refresh instead of stampeding the daemon.""" cached_provider = _build_provider_entity() - cached_payload = TypeAdapter(list[ProviderEntity]).dump_json([cached_provider]) + cached_payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json([cached_provider]) cache_key = _provider_cache_key("tenant-1", 0) with ( @@ -623,7 +699,7 @@ class TestPluginModelProviderCache: def test_fetch_plugin_model_providers_reuses_cached_empty_provider_list(self) -> None: """A cached empty list should prevent repeated daemon fetches for tenants without plugin models.""" - empty_payload = TypeAdapter(list[ProviderEntity]).dump_json([]) + empty_payload = TypeAdapter(list[PluginModelProviderDeclaration]).dump_json([]) cache_key = _provider_cache_key("tenant-1", 0) with patch(f"{MODULE}.redis_client") as redis_client: @@ -806,7 +882,105 @@ class TestPluginListEndpointCounts: assert tool_plugin.endpoints_active == 0 +class TestPluginCategoryList: + def test_list_by_category_forwards_search_and_tag_filters(self) -> None: + plugins = SimpleNamespace(list=[], has_more=False) + + with patch(f"{MODULE}.PluginInstaller") as installer_cls: + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category( + "tenant-1", + PluginCategory.Tool, + 2, + 25, + query="weather", + tags=["search", "rag"], + language="zh_Hans", + ) + + assert result is plugins + installer_cls.return_value.list_plugins_by_category.assert_called_once_with( + "tenant-1", + PluginCategory.Tool, + 2, + 25, + query="weather", + tags=["search", "rag"], + language="zh_Hans", + ) + + def test_filtered_model_category_does_not_reconcile_from_a_partial_result(self) -> None: + plugins = SimpleNamespace(list=[], has_more=False) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category( + "tenant-1", + PluginCategory.Model, + 1, + 100, + query="openai", + tags=[], + language="en_US", + ) + + assert result is plugins + invalidate_cache.assert_not_called() + store_marker.assert_not_called() + + +class TestInstalledPluginIds: + def test_list_installed_plugin_ids_uses_lightweight_daemon_endpoint(self) -> None: + with patch(f"{MODULE}.PluginInstaller") as installer_cls: + installer_cls.return_value.list_installed_plugin_ids.return_value = [ + "langgenius/openai", + "langgenius/anthropic", + ] + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_installed_plugin_ids("tenant-1", PluginCategory.Tool) + + assert result == ["langgenius/openai", "langgenius/anthropic"] + installer_cls.return_value.list_installed_plugin_ids.assert_called_once_with("tenant-1", PluginCategory.Tool) + + class TestPluginModelProviderCacheInvalidation: + def test_list_model_provider_bindings_reconciles_remote_provider_cache(self) -> None: + """The summary binding read owns the remote marker once the full category list leaves the first-load path.""" + remote_binding = _build_remote_model_plugin() + client = MagicMock() + client.fetch_model_provider_bindings.return_value = [remote_binding] + remote_plugin_marker = "langgenius/debug-model:langgenius/debug-model:1.0.0" + + with ( + patch( + f"{MODULE}.PluginService._should_invalidate_model_provider_cache_for_remote_model_plugins", + return_value=True, + ) as should_invalidate, + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + from core.plugin.plugin_service import PluginService + + result = PluginService.list_model_provider_bindings("tenant-1", client=client) + + assert result == [remote_binding] + client.fetch_model_provider_bindings.assert_called_once_with("tenant-1") + should_invalidate.assert_called_once_with("tenant-1", [remote_binding]) + invalidate_cache.assert_called_once_with("tenant-1") + store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) + def test_get_debugging_key_does_not_invalidate_model_provider_cache(self) -> None: """Reading a debug key does not mean a debug runtime has registered a model provider.""" with ( @@ -850,7 +1024,13 @@ class TestPluginModelProviderCacheInvalidation: assert result is plugins installer_cls.return_value.list_plugins_by_category.assert_called_once_with( - "tenant-1", PluginCategory.Model, 1, 100 + "tenant-1", + PluginCategory.Model, + 1, + 100, + query="", + tags=(), + language="en_US", ) invalidate_cache.assert_called_once_with("tenant-1") store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) @@ -917,14 +1097,42 @@ class TestPluginModelProviderCacheInvalidation: invalidate_cache.assert_not_called() store_marker.assert_called_once_with("tenant-1", remote_plugin_marker) - def test_list_model_category_invalidates_when_remote_model_plugin_disconnects(self) -> None: - """The current model category result clears provider cache when the previous debug model disappears.""" + @pytest.mark.parametrize(("page", "has_more"), [(1, True), (2, False)]) + def test_list_model_category_does_not_reconcile_partial_page(self, page: int, has_more: bool) -> None: + """Only an unfiltered, complete first page may write the remote model marker.""" installed_plugin = SimpleNamespace( plugin_id="langgenius/openai", plugin_unique_identifier="langgenius/openai:1.0.0", source=PluginInstallationSource.Marketplace, ) - plugins = SimpleNamespace(list=[installed_plugin], has_more=True) + plugins = SimpleNamespace(list=[installed_plugin], has_more=has_more) + + with ( + patch(f"{MODULE}.PluginInstaller") as installer_cls, + patch( + f"{MODULE}.PluginService._load_cached_remote_model_plugin_marker", + return_value="langgenius/debug-model:langgenius/debug-model:1.0.0", + ), + patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch(f"{MODULE}.PluginService._store_cached_remote_model_plugin_marker") as store_marker, + ): + installer_cls.return_value.list_plugins_by_category.return_value = plugins + + from core.plugin.plugin_service import PluginService + + result = PluginService.list_by_category("tenant-1", PluginCategory.Model, page, 100) + + assert result is plugins + invalidate_cache.assert_not_called() + store_marker.assert_not_called() + + def test_list_model_category_complete_first_page_reconciles_remote_plugin_disconnect(self) -> None: + installed_plugin = SimpleNamespace( + plugin_id="langgenius/openai", + plugin_unique_identifier="langgenius/openai:1.0.0", + source=PluginInstallationSource.Marketplace, + ) + plugins = SimpleNamespace(list=[installed_plugin], has_more=False) with ( patch(f"{MODULE}.PluginInstaller") as installer_cls, @@ -1006,11 +1214,15 @@ class TestPluginModelProviderCacheInvalidation: patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, ): mock_config.MARKETPLACE_ENABLED = True - feature_service.get_system_features.return_value = SimpleNamespace( - plugin_installation_permission=SimpleNamespace(restrict_to_marketplace_only=False) + feature_service.get_plugin_installation_permission.return_value = PluginInstallationPermissionModel( + restrict_to_marketplace_only=False, + plugin_installation_scope=PluginInstallationScope.ALL, ) installer = installer_cls.return_value installer.fetch_plugin_manifest.return_value = MagicMock() + decode_response = MagicMock() + decode_response.verification = None + installer.decode_plugin_from_identifier.return_value = decode_response installer.upgrade_plugin.return_value = "task-id" from core.plugin.plugin_service import PluginService @@ -1146,30 +1358,109 @@ class TestPluginModelProviderCacheInvalidation: assert result is True invalidate_cache.assert_called_once_with("tenant-1") - def test_uninstall_existing_plugin_invalidates_cache_after_credential_cleanup(self) -> None: + @pytest.mark.parametrize( + "sqlite_session", [(Provider, ProviderCredential, TenantPreferredModelProvider)], indirect=True + ) + def test_uninstall_existing_plugin_invalidates_cache_after_credential_cleanup( + self, sqlite_session: Session + ) -> None: """Successful uninstall with plugin metadata also invalidates the mutated tenant provider cache.""" + plugin_id = "langgenius/openai" + provider_name = f"{plugin_id}/openai" plugin = SimpleNamespace( installation_id="installation-1", - plugin_id="langgenius/openai", + plugin_id=plugin_id, plugin_unique_identifier="langgenius/openai:1.0.0", ) - session = _FakeSession() + credential = ProviderCredential( + tenant_id=TENANT_ID, + provider_name=provider_name, + credential_name="Target credential", + encrypted_config="{}", + user_id=USER_ID, + ) + other_credential = ProviderCredential( + tenant_id=OTHER_TENANT_ID, + provider_name=provider_name, + credential_name="Other credential", + encrypted_config="{}", + user_id=USER_ID, + ) + sqlite_session.add_all([credential, other_credential]) + sqlite_session.flush() + provider = Provider( + tenant_id=TENANT_ID, + provider_name=provider_name, + provider_type=ProviderType.CUSTOM, + credential_id=credential.id, + ) + other_provider = Provider( + tenant_id=OTHER_TENANT_ID, + provider_name=provider_name, + provider_type=ProviderType.CUSTOM, + credential_id=other_credential.id, + ) + preferred_provider = TenantPreferredModelProvider( + tenant_id=TENANT_ID, + provider_name=provider_name, + preferred_provider_type=ProviderType.CUSTOM, + ) + other_preferred_provider = TenantPreferredModelProvider( + tenant_id=OTHER_TENANT_ID, + provider_name=provider_name, + preferred_provider_type=ProviderType.CUSTOM, + ) + sqlite_session.add_all([provider, other_provider, preferred_provider, other_preferred_provider]) + sqlite_session.commit() + credential_id = credential.id + other_credential_id = other_credential.id + provider_id = provider.id + other_provider_id = other_provider.id + preferred_provider_id = preferred_provider.id + other_preferred_provider_id = other_preferred_provider.id + with ( - patch(f"{MODULE}.db", SimpleNamespace(engine=object())), + patch(f"{MODULE}.db", SimpleNamespace(engine=sqlite_session.get_bind())), patch(f"{MODULE}.dify_config") as mock_config, patch(f"{MODULE}.PluginInstaller") as installer_cls, - patch(f"{MODULE}.Session", return_value=session), + patch(f"{MODULE}.ProviderCredentialsCache") as credentials_cache, patch(f"{MODULE}.PluginService.invalidate_plugin_model_providers_cache") as invalidate_cache, + patch("core.provider_manager.ProviderManager.invalidate_configurations_cache") as invalidate_configurations, ): - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY installer = installer_cls.return_value installer.list_plugins.return_value = [plugin] installer.uninstall.return_value = True from core.plugin.plugin_service import PluginService - result = PluginService.uninstall("tenant-1", "installation-1") + result = PluginService.uninstall(TENANT_ID, "installation-1") assert result is True - installer.uninstall.assert_called_once_with("tenant-1", "installation-1") - invalidate_cache.assert_called_once_with("tenant-1") + installer.uninstall.assert_called_once_with(TENANT_ID, "installation-1") + invalidate_cache.assert_called_once_with(TENANT_ID) + invalidate_configurations.assert_called_once_with( + TENANT_ID, + sources=( + ProviderConfigurationCacheSource.PREFERRED_MODEL_PROVIDERS, + ProviderConfigurationCacheSource.PROVIDER_CREDENTIALS, + ), + ) + credentials_cache.assert_called_once_with( + tenant_id=TENANT_ID, + identity_id=provider_id, + cache_type=ProviderCredentialsCacheType.PROVIDER, + ) + credentials_cache.return_value.delete.assert_called_once_with() + + sqlite_session.expunge_all() + assert sqlite_session.get(ProviderCredential, credential_id) is None + persisted_provider = sqlite_session.get(Provider, provider_id) + assert persisted_provider is not None + assert persisted_provider.credential_id is None + assert sqlite_session.get(TenantPreferredModelProvider, preferred_provider_id) is None + assert sqlite_session.get(ProviderCredential, other_credential_id) is not None + persisted_other_provider = sqlite_session.get(Provider, other_provider_id) + assert persisted_other_provider is not None + assert persisted_other_provider.credential_id == other_credential_id + assert sqlite_session.get(TenantPreferredModelProvider, other_preferred_provider_id) is not None diff --git a/api/tests/unit_tests/services/plugin/test_plugin_service_installation.py b/api/tests/unit_tests/services/plugin/test_plugin_service_installation.py index ca1a227b010..a98a249457e 100644 --- a/api/tests/unit_tests/services/plugin/test_plugin_service_installation.py +++ b/api/tests/unit_tests/services/plugin/test_plugin_service_installation.py @@ -8,6 +8,7 @@ verification, marketplace upgrade flows, and uninstall with credential cleanup. from __future__ import annotations from collections.abc import Iterator +from typing import cast from unittest.mock import MagicMock, patch from uuid import uuid4 @@ -19,28 +20,24 @@ from sqlalchemy.orm import Session from core.plugin.entities.plugin import PluginInstallationSource from core.plugin.entities.plugin_daemon import PluginVerification from core.plugin.plugin_service import PluginService -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from models import ProviderType from models.engine import db from models.provider import Provider, ProviderCredential, TenantPreferredModelProvider -from services.errors.plugin import PluginInstallationForbiddenError -from services.feature_service import ( +from services.entities.feature_entities import ( PluginInstallationPermissionModel, PluginInstallationScope, - SystemFeatureModel, ) +from services.errors.plugin import PluginInstallationForbiddenError -def _make_features( +def _make_permission( restrict_to_marketplace: bool = False, scope: PluginInstallationScope = PluginInstallationScope.ALL, -) -> SystemFeatureModel: - return SystemFeatureModel( - deployment_edition=DeploymentEdition.COMMUNITY, - plugin_installation_permission=PluginInstallationPermissionModel( - restrict_to_marketplace_only=restrict_to_marketplace, - plugin_installation_scope=scope, - ), +) -> PluginInstallationPermissionModel: + return PluginInstallationPermissionModel( + restrict_to_marketplace_only=restrict_to_marketplace, + plugin_installation_scope=scope, ) @@ -119,22 +116,31 @@ class TestFetchLatestPluginVersion: class TestCheckMarketplaceOnlyPermission: @patch("core.plugin.plugin_service.FeatureService") def test_raises_when_restricted(self, mock_fs): - mock_fs.get_system_features.return_value = _make_features(restrict_to_marketplace=True) + mock_fs.get_plugin_installation_permission.return_value = _make_permission(restrict_to_marketplace=True) with pytest.raises(PluginInstallationForbiddenError): PluginService._check_marketplace_only_permission() @patch("core.plugin.plugin_service.FeatureService") def test_passes_when_not_restricted(self, mock_fs): - mock_fs.get_system_features.return_value = _make_features(restrict_to_marketplace=False) + mock_fs.get_plugin_installation_permission.return_value = _make_permission(restrict_to_marketplace=False) PluginService._check_marketplace_only_permission() # should not raise + @patch("core.plugin.plugin_service.FeatureService") + def test_raises_when_scope_denies_all(self, mock_fs): + mock_fs.get_plugin_installation_permission.return_value = _make_permission(scope=PluginInstallationScope.NONE) + + with pytest.raises(PluginInstallationForbiddenError, match="not allowed"): + PluginService._check_marketplace_only_permission() + class TestCheckPluginInstallationScope: @patch("core.plugin.plugin_service.FeatureService") def test_official_only_allows_langgenius(self, mock_fs): - mock_fs.get_system_features.return_value = _make_features(scope=PluginInstallationScope.OFFICIAL_ONLY) + mock_fs.get_plugin_installation_permission.return_value = _make_permission( + scope=PluginInstallationScope.OFFICIAL_ONLY + ) verification = MagicMock() verification.authorized_category = PluginVerification.AuthorizedCategory.Langgenius @@ -142,14 +148,16 @@ class TestCheckPluginInstallationScope: @patch("core.plugin.plugin_service.FeatureService") def test_official_only_rejects_third_party(self, mock_fs): - mock_fs.get_system_features.return_value = _make_features(scope=PluginInstallationScope.OFFICIAL_ONLY) + mock_fs.get_plugin_installation_permission.return_value = _make_permission( + scope=PluginInstallationScope.OFFICIAL_ONLY + ) with pytest.raises(PluginInstallationForbiddenError): PluginService._check_plugin_installation_scope(None) @patch("core.plugin.plugin_service.FeatureService") def test_official_and_partners_allows_partner(self, mock_fs): - mock_fs.get_system_features.return_value = _make_features( + mock_fs.get_plugin_installation_permission.return_value = _make_permission( scope=PluginInstallationScope.OFFICIAL_AND_SPECIFIC_PARTNERS ) verification = MagicMock() @@ -159,7 +167,7 @@ class TestCheckPluginInstallationScope: @patch("core.plugin.plugin_service.FeatureService") def test_official_and_partners_rejects_none(self, mock_fs): - mock_fs.get_system_features.return_value = _make_features( + mock_fs.get_plugin_installation_permission.return_value = _make_permission( scope=PluginInstallationScope.OFFICIAL_AND_SPECIFIC_PARTNERS ) @@ -168,7 +176,7 @@ class TestCheckPluginInstallationScope: @patch("core.plugin.plugin_service.FeatureService") def test_none_scope_always_raises(self, mock_fs): - mock_fs.get_system_features.return_value = _make_features(scope=PluginInstallationScope.NONE) + mock_fs.get_plugin_installation_permission.return_value = _make_permission(scope=PluginInstallationScope.NONE) verification = MagicMock() verification.authorized_category = PluginVerification.AuthorizedCategory.Langgenius @@ -177,10 +185,19 @@ class TestCheckPluginInstallationScope: @patch("core.plugin.plugin_service.FeatureService") def test_all_scope_passes_any(self, mock_fs): - mock_fs.get_system_features.return_value = _make_features(scope=PluginInstallationScope.ALL) + mock_fs.get_plugin_installation_permission.return_value = _make_permission(scope=PluginInstallationScope.ALL) PluginService._check_plugin_installation_scope(None) # should not raise + @patch("core.plugin.plugin_service.FeatureService") + def test_unknown_scope_always_raises(self, mock_fs): + permission = _make_permission() + permission.plugin_installation_scope = cast(PluginInstallationScope, "unknown-scope") + mock_fs.get_plugin_installation_permission.return_value = permission + + with pytest.raises(PluginInstallationForbiddenError, match="policy is invalid"): + PluginService._check_plugin_installation_scope(None) + class TestGetPluginIconUrl: @patch("core.plugin.plugin_service.dify_config") @@ -248,7 +265,7 @@ class TestUpgradePluginWithMarketplace: @patch("core.plugin.plugin_service.dify_config") def test_skips_download_when_already_installed(self, mock_config, mock_installer_cls, mock_fs, mock_marketplace): mock_config.MARKETPLACE_ENABLED = True - mock_fs.get_system_features.return_value = _make_features() + mock_fs.get_plugin_installation_permission.return_value = _make_permission() installer = mock_installer_cls.return_value installer.fetch_plugin_manifest.return_value = MagicMock() installer.upgrade_plugin.return_value = MagicMock() @@ -264,7 +281,7 @@ class TestUpgradePluginWithMarketplace: @patch("core.plugin.plugin_service.dify_config") def test_downloads_when_not_installed(self, mock_config, mock_installer_cls, mock_fs, mock_download): mock_config.MARKETPLACE_ENABLED = True - mock_fs.get_system_features.return_value = _make_features() + mock_fs.get_plugin_installation_permission.return_value = _make_permission() installer = mock_installer_cls.return_value installer.fetch_plugin_manifest.side_effect = RuntimeError("not found") mock_download.return_value = b"pkg-bytes" @@ -278,12 +295,77 @@ class TestUpgradePluginWithMarketplace: mock_download.assert_called_once_with("new-uid") installer.upload_pkg.assert_called_once() + @pytest.mark.parametrize( + "scope", + [PluginInstallationScope.OFFICIAL_ONLY, PluginInstallationScope.OFFICIAL_AND_SPECIFIC_PARTNERS], + ) + @patch("core.plugin.plugin_service.download_plugin_pkg") + @patch("core.plugin.plugin_service.marketplace") + @patch("core.plugin.plugin_service.FeatureService") + @patch("core.plugin.plugin_service.PluginInstaller") + @patch("core.plugin.plugin_service.dify_config") + def test_rejects_cached_pkg_outside_scope( + self, mock_config, mock_installer_cls, mock_fs, mock_marketplace, mock_download, scope + ): + mock_config.MARKETPLACE_ENABLED = True + mock_fs.get_plugin_installation_permission.return_value = _make_permission(scope=scope) + installer = mock_installer_cls.return_value + installer.fetch_plugin_manifest.return_value = MagicMock() + decode_resp = MagicMock() + decode_resp.verification.authorized_category = PluginVerification.AuthorizedCategory.Community + installer.decode_plugin_from_identifier.return_value = decode_resp + + with pytest.raises(PluginInstallationForbiddenError): + PluginService.upgrade_plugin_with_marketplace("t1", "old-uid", "new-uid") + + installer.upgrade_plugin.assert_not_called() + mock_marketplace.record_install_plugin_event.assert_not_called() + # the rejection must not fall through to the download branch and cache the pkg again + mock_download.assert_not_called() + installer.upload_pkg.assert_not_called() + + @patch("core.plugin.plugin_service.FeatureService") + @patch("core.plugin.plugin_service.PluginInstaller") + @patch("core.plugin.plugin_service.dify_config") + def test_rejects_before_touching_daemon_when_scope_is_none(self, mock_config, mock_installer_cls, mock_fs): + mock_config.MARKETPLACE_ENABLED = True + mock_fs.get_plugin_installation_permission.return_value = _make_permission(scope=PluginInstallationScope.NONE) + installer = mock_installer_cls.return_value + + with pytest.raises(PluginInstallationForbiddenError): + PluginService.upgrade_plugin_with_marketplace("t1", "old-uid", "new-uid") + + installer.fetch_plugin_manifest.assert_not_called() + installer.upgrade_plugin.assert_not_called() + + @patch("core.plugin.plugin_service.marketplace") + @patch("core.plugin.plugin_service.FeatureService") + @patch("core.plugin.plugin_service.PluginInstaller") + @patch("core.plugin.plugin_service.dify_config") + def test_allows_cached_official_pkg_under_official_only( + self, mock_config, mock_installer_cls, mock_fs, mock_marketplace + ): + mock_config.MARKETPLACE_ENABLED = True + mock_fs.get_plugin_installation_permission.return_value = _make_permission( + scope=PluginInstallationScope.OFFICIAL_ONLY + ) + installer = mock_installer_cls.return_value + installer.fetch_plugin_manifest.return_value = MagicMock() + decode_resp = MagicMock() + decode_resp.verification.authorized_category = PluginVerification.AuthorizedCategory.Langgenius + installer.decode_plugin_from_identifier.return_value = decode_resp + + PluginService.upgrade_plugin_with_marketplace("t1", "old-uid", "new-uid") + + mock_marketplace.record_install_plugin_event.assert_called_once_with("new-uid") + installer.upgrade_plugin.assert_called_once() + class TestUpgradePluginWithGithub: @patch("core.plugin.plugin_service.FeatureService") @patch("core.plugin.plugin_service.PluginInstaller") def test_checks_marketplace_permission_and_delegates(self, mock_installer_cls: MagicMock, mock_fs: MagicMock): - mock_fs.get_system_features.return_value = _make_features() + mock_fs.get_plugin_installation_permission.return_value = _make_permission() installer = mock_installer_cls.return_value installer.upgrade_plugin.return_value = MagicMock() @@ -298,7 +380,7 @@ class TestUploadPkg: @patch("core.plugin.plugin_service.FeatureService") @patch("core.plugin.plugin_service.PluginInstaller") def test_runs_permission_and_scope_checks(self, mock_installer_cls: MagicMock, mock_fs: MagicMock): - mock_fs.get_system_features.return_value = _make_features() + mock_fs.get_plugin_installation_permission.return_value = _make_permission() upload_resp = MagicMock() upload_resp.verification = None mock_installer_cls.return_value.upload_pkg.return_value = upload_resp @@ -322,7 +404,7 @@ class TestInstallFromMarketplacePkg: @patch("core.plugin.plugin_service.dify_config") def test_downloads_when_not_cached(self, mock_config, mock_installer_cls, mock_fs, mock_download): mock_config.MARKETPLACE_ENABLED = True - mock_fs.get_system_features.return_value = _make_features() + mock_fs.get_plugin_installation_permission.return_value = _make_permission() installer = mock_installer_cls.return_value installer.fetch_plugin_manifest.side_effect = RuntimeError("not found") mock_download.return_value = b"pkg" @@ -344,7 +426,7 @@ class TestInstallFromMarketplacePkg: @patch("core.plugin.plugin_service.dify_config") def test_uses_cached_when_already_downloaded(self, mock_config, mock_installer_cls: MagicMock, mock_fs: MagicMock): mock_config.MARKETPLACE_ENABLED = True - mock_fs.get_system_features.return_value = _make_features() + mock_fs.get_plugin_installation_permission.return_value = _make_permission() installer = mock_installer_cls.return_value installer.fetch_plugin_manifest.return_value = MagicMock() decode_resp = MagicMock() @@ -358,6 +440,29 @@ class TestInstallFromMarketplacePkg: call_args = installer.install_from_identifiers.call_args[0] assert call_args[1] == ["uid-1"] + @patch("core.plugin.plugin_service.download_plugin_pkg") + @patch("core.plugin.plugin_service.FeatureService") + @patch("core.plugin.plugin_service.PluginInstaller") + @patch("core.plugin.plugin_service.dify_config") + def test_rejects_cached_pkg_outside_scope(self, mock_config, mock_installer_cls, mock_fs, mock_download): + mock_config.MARKETPLACE_ENABLED = True + mock_fs.get_plugin_installation_permission.return_value = _make_permission( + scope=PluginInstallationScope.OFFICIAL_ONLY + ) + installer = mock_installer_cls.return_value + installer.fetch_plugin_manifest.return_value = MagicMock() + decode_resp = MagicMock() + decode_resp.verification.authorized_category = PluginVerification.AuthorizedCategory.Community + installer.decode_plugin_from_identifier.return_value = decode_resp + + with pytest.raises(PluginInstallationForbiddenError): + PluginService.install_from_marketplace_pkg("t1", ["uid-1"]) + + installer.install_from_identifiers.assert_not_called() + # the rejection must not fall through to the download branch and cache the pkg again + mock_download.assert_not_called() + installer.upload_pkg.assert_not_called() + class TestUninstall: @patch("core.plugin.plugin_service.PluginInstaller") @@ -412,7 +517,7 @@ class TestUninstall: installer.uninstall.return_value = True with patch("core.plugin.plugin_service.dify_config") as mock_config: - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY result = PluginService.uninstall(tenant_id, "install-1") assert result is True @@ -436,3 +541,78 @@ class TestUninstall: ) ).all() assert len(remaining_prefs) == 0 + + @patch("core.plugin.plugin_service.PluginInstaller") + def test_preserves_credentials_when_replacing_plugin( + self, mock_installer_cls: MagicMock, plugin_db: Session + ) -> None: + tenant_id = str(uuid4()) + plugin_id = "org/myplugin" + provider_name = f"{plugin_id}/model-provider" + + credential = ProviderCredential( + tenant_id=tenant_id, + provider_name=provider_name, + credential_name="default", + encrypted_config="{}", + ) + plugin_db.add(credential) + plugin_db.flush() + + provider = Provider( + tenant_id=tenant_id, + provider_name=provider_name, + credential_id=credential.id, + ) + plugin_db.add(provider) + + preferred_provider = TenantPreferredModelProvider( + tenant_id=tenant_id, + provider_name=provider_name, + preferred_provider_type=ProviderType.CUSTOM, + ) + plugin_db.add(preferred_provider) + plugin_db.commit() + + plugin = MagicMock(installation_id="install-1", plugin_id=plugin_id) + installer = mock_installer_cls.return_value + installer.list_plugins.return_value = [plugin] + installer.uninstall.return_value = True + + with patch("core.plugin.plugin_service.dify_config") as mock_config: + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY + result = PluginService.uninstall(tenant_id, "install-1", preserve_credentials=True) + + assert result is True + plugin_db.expire_all() + assert plugin_db.get(ProviderCredential, credential.id) is not None + assert plugin_db.get(Provider, provider.id).credential_id == credential.id + assert plugin_db.get(TenantPreferredModelProvider, preferred_provider.id) is not None + + @patch("core.plugin.plugin_service.PluginInstaller") + def test_preserves_credentials_when_daemon_uninstall_fails( + self, mock_installer_cls: MagicMock, plugin_db: Session + ) -> None: + tenant_id = str(uuid4()) + plugin_id = "org/myplugin" + credential = ProviderCredential( + tenant_id=tenant_id, + provider_name=f"{plugin_id}/model-provider", + credential_name="default", + encrypted_config="{}", + ) + plugin_db.add(credential) + plugin_db.commit() + + plugin = MagicMock(installation_id="install-1", plugin_id=plugin_id) + installer = mock_installer_cls.return_value + installer.list_plugins.return_value = [plugin] + installer.uninstall.return_value = False + + with patch("core.plugin.plugin_service.dify_config") as mock_config: + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY + result = PluginService.uninstall(tenant_id, "install-1") + + assert result is False + plugin_db.expire_all() + assert plugin_db.get(ProviderCredential, credential.id) is not None diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/conftest.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/conftest.py deleted file mode 100644 index 2e65be48086..00000000000 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/conftest.py +++ /dev/null @@ -1,18 +0,0 @@ -from collections.abc import Callable -from unittest.mock import Mock - -import pytest -from pytest_mock import MockerFixture - - -@pytest.fixture -def patch_session_factory(mocker: MockerFixture) -> Callable[[str, Mock], Mock]: - def _patch(module_path: str, session_mock: Mock) -> Mock: - session_context = mocker.MagicMock() - session_context.__enter__.return_value = session_mock - session_context.__exit__.return_value = None - session_maker = mocker.Mock(return_value=session_context) - mocker.patch(f"{module_path}.session_factory.get_session_maker", return_value=session_maker) - return session_maker - - return _patch diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py index 0576f44babd..9ea47c0a141 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_built_in_retrieval.py @@ -45,7 +45,7 @@ def test_get_pipeline_template_detail(mocker: MockerFixture, sqlite_session: Ses ) retrieval = BuiltInPipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail("tpl-1", session=sqlite_session) + detail = retrieval.get_pipeline_template_detail("tpl-1", "tenant-1", session=sqlite_session) assert detail == {"id": "tpl-1", "name": "Template 1"} assert not sqlite_session.in_transaction() @@ -79,7 +79,7 @@ def test_get_pipeline_template_detail_returns_none_for_unknown_id( ) retrieval = BuiltInPipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail("nonexistent-id", session=sqlite_session) + result = retrieval.get_pipeline_template_detail("nonexistent-id", "tenant-1", session=sqlite_session) assert result is None assert not sqlite_session.in_transaction() diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py index 0245c2b2aa9..d883a0ac07d 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_customized_retrieval.py @@ -81,7 +81,7 @@ def test_get_pipeline_template_detail_returns_detail(monkeypatch: pytest.MonkeyP monkeypatch.setattr("models.dataset.db", SimpleNamespace(session=sqlite_session)) retrieval = CustomizedPipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail(TEMPLATE_ID, session=sqlite_session) + detail = retrieval.get_pipeline_template_detail(TEMPLATE_ID, TENANT_ID, session=sqlite_session) assert detail == { "id": TEMPLATE_ID, @@ -97,10 +97,12 @@ def test_get_pipeline_template_detail_returns_detail(monkeypatch: pytest.MonkeyP @pytest.mark.parametrize("sqlite_session", [(PipelineCustomizedTemplate,)], indirect=True) -def test_get_pipeline_template_detail_returns_none_when_not_found(sqlite_session: Session) -> None: +def test_get_pipeline_template_detail_rejects_other_tenant(sqlite_session: Session) -> None: + sqlite_session.add(_template(tenant_id=OTHER_TENANT_ID)) + sqlite_session.commit() retrieval = CustomizedPipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail(TEMPLATE_ID, session=sqlite_session) + result = retrieval.get_pipeline_template_detail(TEMPLATE_ID, TENANT_ID, session=sqlite_session) assert result is None assert sqlite_session.in_transaction() diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py index 338efe7558e..243121a6d71 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_database_retrieval.py @@ -68,7 +68,7 @@ def test_get_pipeline_template_detail_returns_detail(sqlite_session: Session) -> sqlite_session.commit() retrieval = DatabasePipelineTemplateRetrieval() - detail = retrieval.get_pipeline_template_detail(TEMPLATE_ID, session=sqlite_session) + detail = retrieval.get_pipeline_template_detail(TEMPLATE_ID, "tenant-1", session=sqlite_session) assert detail == { "id": TEMPLATE_ID, @@ -86,7 +86,7 @@ def test_get_pipeline_template_detail_returns_detail(sqlite_session: Session) -> def test_get_pipeline_template_detail_returns_none_when_not_found(sqlite_session: Session) -> None: retrieval = DatabasePipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail(TEMPLATE_ID, session=sqlite_session) + result = retrieval.get_pipeline_template_detail(TEMPLATE_ID, "tenant-1", session=sqlite_session) assert result is None assert sqlite_session.in_transaction() diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py index 25472552e8e..ad19b411a43 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_pipeline_template_base.py @@ -9,8 +9,8 @@ class DummyRetrieval(PipelineTemplateRetrievalBase): del session return {"language": language} - def get_pipeline_template_detail(self, template_id: str, *, session) -> dict | None: - del session + def get_pipeline_template_detail(self, template_id: str, current_tenant_id: str, *, session) -> dict | None: + del current_tenant_id, session return {"id": template_id} def get_type(self) -> str: @@ -22,6 +22,6 @@ def test_pipeline_template_retrieval_base_concrete_implementation(sqlite_session retrieval = DummyRetrieval() assert retrieval.get_pipeline_templates("en-US", session=sqlite_session) == {"language": "en-US"} - assert retrieval.get_pipeline_template_detail("tpl-1", session=sqlite_session) == {"id": "tpl-1"} + assert retrieval.get_pipeline_template_detail("tpl-1", "tenant-1", session=sqlite_session) == {"id": "tpl-1"} assert retrieval.get_type() == "dummy" assert not sqlite_session.in_transaction() diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py index 5db648a3014..273f89b8db3 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py @@ -46,7 +46,7 @@ def test_get_pipeline_template_detail_fallbacks_to_database_on_error( ) retrieval = RemotePipelineTemplateRetrieval() - result = retrieval.get_pipeline_template_detail("tpl-1", session=sqlite_session) + result = retrieval.get_pipeline_template_detail("tpl-1", "tenant-1", session=sqlite_session) assert result == {"id": "db-1"} fetch_mock.assert_called_once_with("tpl-1") diff --git a/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py b/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py index 1ad4e070edd..ac9303b2bc4 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_pipeline_generate_service.py @@ -3,14 +3,51 @@ from typing import cast import pytest from pytest_mock import MockerFixture +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom -from models.dataset import Dataset, Pipeline +from core.rag.index_processor.constant.index_type import IndexStructureType +from models.dataset import Dataset, Document, Pipeline +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus from models.model import Account, App, EndUser from services.dataset_ref_service import DatasetRefService from services.rag_pipeline.pipeline_generate_service import PipelineGenerateService +def _make_pipeline(*, tenant_id: str = "tenant-1") -> Pipeline: + pipeline = Pipeline(tenant_id=tenant_id, name="Pipeline", description="") + pipeline.id = "pipeline-1" + return pipeline + + +def _make_dataset(*, dataset_id: str = "dataset-1", tenant_id: str = "tenant-1") -> Dataset: + return Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Dataset", + created_by="user-1", + pipeline_id="pipeline-1", + ) + + +def _make_document( + *, document_id: str = "doc-1", dataset_id: str = "dataset-1", tenant_id: str = "tenant-1" +) -> Document: + return Document( + id=document_id, + tenant_id=tenant_id, + dataset_id=dataset_id, + position=1, + data_source_type=DataSourceType.LOCAL_FILE, + batch="batch", + name="Document", + created_from=DocumentCreatedFrom.API, + created_by="user-1", + indexing_status=IndexingStatus.COMPLETED, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + + def test_get_max_active_requests_uses_smallest_non_zero_limit(mocker: MockerFixture) -> None: mocker.patch("services.rag_pipeline.pipeline_generate_service.dify_config.APP_DEFAULT_ACTIVE_REQUESTS", 5) mocker.patch("services.rag_pipeline.pipeline_generate_service.dify_config.APP_MAX_ACTIVE_REQUESTS", 3) @@ -41,14 +78,14 @@ def test_get_max_active_requests_returns_zero_when_all_unlimited(mocker: MockerF (InvokeFrom.DEBUGGER, SimpleNamespace(id="wf-1"), None), ], ) -def test_get_workflow(mocker: MockerFixture, invoke_from, workflow, expected_error) -> None: +def test_get_workflow(mocker: MockerFixture, invoke_from, workflow, expected_error, sqlite_session: Session) -> None: rag_pipeline_service_cls = mocker.patch("services.rag_pipeline.pipeline_generate_service.RagPipelineService") rag_pipeline_service = rag_pipeline_service_cls.return_value rag_pipeline_service.get_draft_workflow.return_value = workflow rag_pipeline_service.get_published_workflow.return_value = workflow pipeline = cast(Pipeline, SimpleNamespace(id="pipeline-1")) - session = mocker.Mock() + session = sqlite_session if expected_error: with pytest.raises(ValueError, match=expected_error): @@ -58,7 +95,9 @@ def test_get_workflow(mocker: MockerFixture, invoke_from, workflow, expected_err assert result == workflow -def test_generate_updates_document_status_and_returns_event_stream(mocker: MockerFixture) -> None: +def test_generate_updates_document_status_and_returns_event_stream( + mocker: MockerFixture, sqlite_session: Session +) -> None: dataset = cast(Dataset, SimpleNamespace(id="dataset-1", tenant_id="tenant-1")) pipeline = cast( Pipeline, @@ -70,7 +109,6 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke ) user = cast(Account | EndUser, SimpleNamespace(id="user-1")) args = {"original_document_id": "doc-1", "query": "hello"} - session_mock = mocker.Mock() mocker.patch.object(PipelineGenerateService, "_get_workflow", return_value=SimpleNamespace(id="wf-1")) update_status_mock = mocker.patch.object(PipelineGenerateService, "update_document_status") @@ -86,7 +124,7 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke args=args, invoke_from=InvokeFrom.WEB_APP, streaming=True, - session=session_mock, + session=sqlite_session, ) assert result == "stream-events" @@ -94,11 +132,11 @@ def test_generate_updates_document_status_and_returns_event_stream(mocker: Mocke assert document_ref.dataset.tenant_id == "tenant-1" assert document_ref.dataset.dataset_id == "dataset-1" assert document_ref.document_id == "doc-1" - update_status_mock.assert_called_once_with(document_ref, session=session_mock) - assert generator_instance.generate.call_args.kwargs["session"] is session_mock + update_status_mock.assert_called_once_with(document_ref, session=sqlite_session) + assert generator_instance.generate.call_args.kwargs["session"] is sqlite_session -def test_generate_rejects_pipeline_dataset_from_another_tenant(mocker: MockerFixture) -> None: +def test_generate_rejects_pipeline_dataset_from_another_tenant(mocker: MockerFixture, sqlite_session: Session) -> None: dataset = cast(Dataset, SimpleNamespace(id="dataset-1", tenant_id="tenant-2")) pipeline = cast( Pipeline, @@ -117,7 +155,7 @@ def test_generate_rejects_pipeline_dataset_from_another_tenant(mocker: MockerFix user=cast(Account, SimpleNamespace(id="user-1")), args={"original_document_id": "doc-1"}, invoke_from=InvokeFrom.WEB_APP, - session=mocker.Mock(), + session=sqlite_session, ) update_status_mock.assert_not_called() @@ -125,18 +163,13 @@ def test_generate_rejects_pipeline_dataset_from_another_tenant(mocker: MockerFix def test_generate_rejects_original_document_outside_pipeline_dataset_before_dispatch( mocker: MockerFixture, + sqlite_session: Session, ) -> None: - dataset = cast(Dataset, SimpleNamespace(id="dataset-1", tenant_id="tenant-1")) - pipeline = cast( - Pipeline, - SimpleNamespace( - id="pipeline-1", - tenant_id="tenant-1", - retrieve_dataset=mocker.Mock(return_value=dataset), - ), - ) - session = mocker.Mock() - session.scalar.return_value = None + dataset = _make_dataset() + pipeline = _make_pipeline() + outside_document = _make_document(document_id="foreign-doc", dataset_id="other-dataset", tenant_id="tenant-2") + sqlite_session.add_all([dataset, pipeline, outside_document]) + sqlite_session.commit() mocker.patch.object(PipelineGenerateService, "_get_workflow", return_value=SimpleNamespace(id="wf-1")) generator_cls = mocker.patch("services.rag_pipeline.pipeline_generate_service.PipelineGenerator") @@ -146,29 +179,25 @@ def test_generate_rejects_original_document_outside_pipeline_dataset_before_disp user=cast(Account, SimpleNamespace(id="user-1")), args={"original_document_id": "foreign-doc"}, invoke_from=InvokeFrom.PUBLISHED_PIPELINE, - session=session, + session=sqlite_session, ) - statement = session.scalar.call_args.args[0] - assert {"foreign-doc", "dataset-1", "tenant-1"} <= set(statement.compile().params.values()) + sqlite_session.refresh(outside_document) + assert outside_document.indexing_status == IndexingStatus.COMPLETED generator_cls.assert_not_called() -def test_update_document_status_updates_existing_document(mocker: MockerFixture) -> None: - document = SimpleNamespace(indexing_status="completed") - dataset = cast(Dataset, SimpleNamespace(id="dataset-1", tenant_id="tenant-1")) +def test_update_document_status_updates_existing_document(sqlite_session: Session) -> None: + document = _make_document() + dataset = _make_dataset() + sqlite_session.add_all([dataset, document]) + sqlite_session.commit() dataset_ref = DatasetRefService.create_dataset_ref(dataset) document_ref = DatasetRefService.create_document_ref_from_id(dataset_ref, "doc-1") - session_mock = mocker.Mock() - get_document_mock = mocker.patch.object(DatasetRefService, "get_document_by_ref", return_value=document) - add_mock = session_mock.add + PipelineGenerateService.update_document_status(document_ref, session=sqlite_session) - PipelineGenerateService.update_document_status(document_ref, session=session_mock) - - assert document.indexing_status == "waiting" - get_document_mock.assert_called_once_with(document_ref, session=session_mock) - add_mock.assert_called_once_with(document) + assert document.indexing_status == IndexingStatus.WAITING @pytest.mark.parametrize( @@ -179,42 +208,28 @@ def test_update_document_status_updates_existing_document(mocker: MockerFixture) ], ) def test_update_document_status_rejects_document_outside_owner( - mocker: MockerFixture, document_tenant_id: str, document_dataset_id: str, + sqlite_session: Session, ) -> None: dataset = cast(Dataset, SimpleNamespace(id="dataset-1", tenant_id="tenant-1")) dataset_ref = DatasetRefService.create_dataset_ref(dataset) document_ref = DatasetRefService.create_document_ref_from_id(dataset_ref, "doc-1") - outside_document = SimpleNamespace( - id="doc-1", - tenant_id=document_tenant_id, - dataset_id=document_dataset_id, - indexing_status="completed", - ) - session_mock = mocker.Mock() - - def resolve_document(statement): - params = set(statement.compile().params.values()) - outside_owner = {outside_document.id, outside_document.dataset_id, outside_document.tenant_id} - return outside_document if outside_owner <= params else None - - session_mock.scalar.side_effect = resolve_document - add_mock = session_mock.add + outside_document = _make_document(tenant_id=document_tenant_id, dataset_id=document_dataset_id) + sqlite_session.add(outside_document) + sqlite_session.commit() with pytest.raises(ValueError, match="Pipeline document not found"): - PipelineGenerateService.update_document_status(document_ref, session=session_mock) + PipelineGenerateService.update_document_status(document_ref, session=sqlite_session) - statement = session_mock.scalar.call_args.args[0] - assert {"doc-1", "dataset-1", "tenant-1"} <= set(statement.compile().params.values()) - assert outside_document.indexing_status == "completed" - add_mock.assert_not_called() + sqlite_session.refresh(outside_document) + assert outside_document.indexing_status == IndexingStatus.COMPLETED # --- generate_single_iteration --- -def test_generate_single_iteration_delegates(mocker: MockerFixture) -> None: +def test_generate_single_iteration_delegates(mocker: MockerFixture, sqlite_session: Session) -> None: mocker.patch.object(PipelineGenerateService, "_get_workflow", return_value=SimpleNamespace(id="wf-1")) generator_cls = mocker.patch("services.rag_pipeline.pipeline_generate_service.PipelineGenerator") @@ -224,7 +239,7 @@ def test_generate_single_iteration_delegates(mocker: MockerFixture) -> None: pipeline = cast(Pipeline, SimpleNamespace(id="p1")) user = cast(Account, SimpleNamespace(id="u1")) - session = mocker.Mock() + session = sqlite_session result = PipelineGenerateService.generate_single_iteration(pipeline, user, "node-1", {"key": "val"}, session) @@ -236,7 +251,7 @@ def test_generate_single_iteration_delegates(mocker: MockerFixture) -> None: # --- generate_single_loop --- -def test_generate_single_loop_delegates(mocker: MockerFixture) -> None: +def test_generate_single_loop_delegates(mocker: MockerFixture, sqlite_session: Session) -> None: mocker.patch.object(PipelineGenerateService, "_get_workflow", return_value=SimpleNamespace(id="wf-1")) generator_cls = mocker.patch("services.rag_pipeline.pipeline_generate_service.PipelineGenerator") @@ -246,7 +261,7 @@ def test_generate_single_loop_delegates(mocker: MockerFixture) -> None: pipeline = cast(Pipeline, SimpleNamespace(id="p1")) user = cast(Account, SimpleNamespace(id="u1")) - session = mocker.Mock() + session = sqlite_session result = PipelineGenerateService.generate_single_loop(pipeline, user, "node-1", {"key": "val"}, session) diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py index 5cdb2afd093..a2cc34741ef 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_dsl_service.py @@ -1,25 +1,176 @@ +"""SQLite-backed tests for the RAG pipeline DSL service. + +Database behavior and transaction ownership use the shared SQLite fixtures. Plugin, +Redis, validation, and HTTP boundaries remain isolated so each test can exercise the +service branch that owns the behavior under test. +""" + +from __future__ import annotations + +import json +from collections.abc import Generator +from contextlib import contextmanager from types import SimpleNamespace from typing import Any, cast -from unittest.mock import MagicMock, Mock +from unittest.mock import MagicMock, Mock, call import pytest import yaml -from pytest_mock import MockerFixture -from sqlalchemy.orm import Session +from sqlalchemy import Engine, event, select +from sqlalchemy.orm import Session, sessionmaker +from core.workflow.llm_environment_variable import LLMEnvironmentVariable from core.workflow.nodes.knowledge_index import KNOWLEDGE_INDEX_NODE_TYPE from graphon.enums import BuiltinNodeTypes +from models import Account, Tenant +from models.dataset import ( + Dataset, + Pipeline, + PipelineCustomizedTemplate, +) +from models.enums import DataSourceType +from models.workflow import Workflow, WorkflowKind, WorkflowType from services.dsl_version import check_version_compatibility from services.entities.knowledge_entities.rag_pipeline_entities import IconInfo, RagPipelineDatasetCreateEntity -from services.rag_pipeline import rag_pipeline_dsl_service +from services.rag_pipeline import rag_pipeline_dsl_service as module from services.rag_pipeline.rag_pipeline_dsl_service import ( + ImportMode, ImportStatus, RagPipelineDslService, + RagPipelinePendingData, ) +@pytest.fixture +def service(sqlite_session: Session) -> RagPipelineDslService: + return RagPipelineDslService(session=sqlite_session) + + +def _account(*, tenant_id: str = "tenant-1", account_id: str = "account-1") -> Account: + tenant = Tenant(name="Tenant") + tenant.id = tenant_id + account = Account(name="Account", email="account@example.com") + account.id = account_id + account._current_tenant = tenant + return account + + +def _pipeline( + session: Session, + *, + tenant_id: str = "tenant-1", + name: str = "Pipeline", + published: bool = False, +) -> Pipeline: + pipeline = Pipeline( + tenant_id=tenant_id, + name=name, + description="description", + is_published=published, + created_by="account-1", + updated_by="account-1", + ) + session.add(pipeline) + session.commit() + return pipeline + + +def _workflow(session: Session, pipeline: Pipeline, *, graph: dict[str, Any] | None = None) -> Workflow: + workflow = Workflow( + id=f"workflow-{pipeline.id}", + tenant_id=pipeline.tenant_id, + app_id=pipeline.id, + type=WorkflowType.RAG_PIPELINE, + kind=WorkflowKind.STANDARD, + version=Workflow.VERSION_DRAFT, + graph=json.dumps(graph or {"nodes": [], "edges": []}), + features="{}", + created_by="account-1", + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + pipeline.workflow_id = workflow.id + session.add(workflow) + session.commit() + return workflow + + +def _dataset( + session: Session, + pipeline: Pipeline, + *, + name: str = "Dataset", + chunk_structure: str = "text_model", +) -> Dataset: + dataset = Dataset( + tenant_id=pipeline.tenant_id, + name=name, + description="description", + data_source_type=DataSourceType.UPLOAD_FILE, + indexing_technique="high_quality", + created_by="account-1", + maintainer="account-1", + chunk_structure=chunk_structure, + pipeline_id=pipeline.id, + icon_info={"icon": "📙", "icon_type": "emoji"}, + ) + session.add(dataset) + session.commit() + return dataset + + +def _knowledge_configuration() -> SimpleNamespace: + return SimpleNamespace( + indexing_technique="high_quality", + embedding_model="text-embedding", + embedding_model_provider="openai", + chunk_structure="text_model", + retrieval_model=SimpleNamespace(model_dump=lambda: {}), + summary_index_setting=None, + keyword_number=10, + ) + + +def _valid_dsl(*, version: str = "0.1.0", name: str = "Imported") -> str: + return f""" +version: {version} +kind: rag_pipeline +rag_pipeline: + name: {name} + description: description +workflow: + graph: + nodes: + - id: knowledge-index + data: + type: {KNOWLEDGE_INDEX_NODE_TYPE} + edges: [] +""" + + +@contextmanager +def _raise_on_workflow_insert(engine: Engine) -> Generator[None]: + def raise_error( + _conn: object, + _cursor: object, + statement: str, + _parameters: object, + _context: object, + _executemany: object, + ) -> None: + if statement.lstrip().upper().startswith("INSERT") and "workflows" in statement: + raise RuntimeError("forced workflow INSERT") + + event.listen(engine, "before_cursor_execute", raise_error) + try: + yield + finally: + event.remove(engine, "before_cursor_execute", raise_error) + + @pytest.mark.parametrize( - ("imported_version", "expected_status"), + ("version", "expected"), [ ("invalid", ImportStatus.FAILED), ("1.0.0", ImportStatus.PENDING), @@ -27,1131 +178,43 @@ from services.rag_pipeline.rag_pipeline_dsl_service import ( ("0.1.0", ImportStatus.COMPLETED), ], ) -def test_check_version_compatibility(imported_version: str, expected_status: ImportStatus) -> None: - assert ( - check_version_compatibility(imported_version, rag_pipeline_dsl_service.CURRENT_DSL_VERSION) == expected_status - ) +def test_version_compatibility(version: str, expected: ImportStatus) -> None: + assert check_version_compatibility(version, module.CURRENT_DSL_VERSION) == expected -def test_encrypt_decrypt_dataset_id_roundtrip() -> None: - service = RagPipelineDslService(session=Mock()) - +def test_dataset_id_encryption_roundtrip_and_invalid(service: RagPipelineDslService) -> None: encrypted = service.encrypt_dataset_id("dataset-1", "tenant-1") - decrypted = service.decrypt_dataset_id(encrypted, "tenant-1") + assert service.decrypt_dataset_id(encrypted, "tenant-1") == "dataset-1" + assert service.decrypt_dataset_id("not-base64", "tenant-1") is None - assert decrypted == "dataset-1" +def test_dependency_helpers_keep_plugin_analysis_external( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService +) -> None: + assert service.get_leaked_dependencies("tenant-1", []) == [] + dependency = MagicMock() + leaked = [MagicMock()] + analyze = Mock(return_value=leaked) + monkeypatch.setattr(module.DependenciesAnalysisService, "get_leaked_dependencies", analyze) + assert service.get_leaked_dependencies("tenant-1", [dependency]) == leaked + assert service._extract_dependencies_from_model_config({}) == [] + assert service._extract_dependencies_from_workflow_graph({}) == [] -def test_decrypt_dataset_id_returns_none_for_invalid_payload() -> None: - service = RagPipelineDslService(session=Mock()) - result = service.decrypt_dataset_id("not-base64", "tenant-1") - - assert result is None - - -def test_get_leaked_dependencies_returns_empty_list_for_empty_input() -> None: - result = RagPipelineDslService.get_leaked_dependencies("tenant-1", []) - - assert result == [] - - -def test_get_leaked_dependencies_delegates_to_analysis_service(mocker: MockerFixture) -> None: - expected = [Mock()] - get_leaked_mock = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.get_leaked_dependencies", - return_value=expected, - ) - - dependency = Mock() - result = RagPipelineDslService.get_leaked_dependencies("tenant-1", [dependency]) - - assert result == expected - get_leaked_mock.assert_called_once_with(tenant_id="tenant-1", dependencies=[dependency]) - - -# --- check_dependencies --- - - -def test_check_dependencies_returns_empty_when_no_redis_data(mocker: MockerFixture) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", - return_value=None, - ) - service = RagPipelineDslService(session=Mock()) - pipeline = Mock(id="p1", tenant_id="t1") - - result = service.check_dependencies(pipeline=pipeline) - - assert result.leaked_dependencies == [] - - -def test_check_dependencies_returns_leaked_deps_from_redis(mocker: MockerFixture) -> None: - from core.plugin.entities.plugin import PluginDependency, PluginDependencyType - from services.rag_pipeline.rag_pipeline_dsl_service import CheckDependenciesPendingData - - dep = PluginDependency( - type=PluginDependencyType.Marketplace, - value=PluginDependency.Marketplace(marketplace_plugin_unique_identifier="test/plugin:0.1.0"), - ) - pending_data = CheckDependenciesPendingData( - dependencies=[dep], - pipeline_id="p1", - ) - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", - return_value=pending_data.model_dump_json(), - ) - leaked = [dep] - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.get_leaked_dependencies", - return_value=leaked, - ) - service = RagPipelineDslService(session=Mock()) - pipeline = Mock(id="p1", tenant_id="t1") - - result = service.check_dependencies(pipeline=pipeline) - - assert result.leaked_dependencies == leaked - - -# --- _extract_dependencies_from_model_config --- - - -def test_extract_dependencies_from_model_config_extracts_model(mocker: MockerFixture) -> None: - analyze_mock = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", - return_value="langgenius/openai", - ) - config = {"model": {"provider": "openai"}} - - result = RagPipelineDslService._extract_dependencies_from_model_config(config) - - assert "langgenius/openai" in result - analyze_mock.assert_called_with("openai") - - -def test_extract_dependencies_from_model_config_extracts_tools(mocker: MockerFixture) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", - return_value="x", - ) - analyze_tool_mock = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_tool_dependency", - return_value="langgenius/google", - ) - config = { - "model": {"provider": "openai"}, - "agent_mode": {"tools": [{"provider_id": "google"}]}, - } - - result = RagPipelineDslService._extract_dependencies_from_model_config(config) - - assert "langgenius/google" in result - analyze_tool_mock.assert_called_with("google") - - -def test_extract_dependencies_from_model_config_empty_config() -> None: - result = RagPipelineDslService._extract_dependencies_from_model_config({}) - - assert result == [] - - -# --- _extract_dependencies_from_workflow_graph --- - - -def test_extract_dependencies_from_workflow_graph_ignores_unknown_types(mocker: MockerFixture) -> None: - service = RagPipelineDslService(session=Mock()) - graph = {"nodes": [{"data": {"type": "some-unknown-type"}}]} - - result = service._extract_dependencies_from_workflow_graph(graph) - - assert result == [] - - -def test_extract_dependencies_from_workflow_graph_handles_empty_graph() -> None: - service = RagPipelineDslService(session=Mock()) - - result = service._extract_dependencies_from_workflow_graph({}) - - assert result == [] - - -def test_extract_dependencies_from_workflow_graph_handles_malformed_node(mocker: MockerFixture) -> None: - service = RagPipelineDslService(session=Mock()) - # Node with TOOL type but invalid data should be caught by exception handler - from graphon.enums import BuiltinNodeTypes - - graph = {"nodes": [{"data": {"type": BuiltinNodeTypes.TOOL}}]} - - result = service._extract_dependencies_from_workflow_graph(graph) - - # Should not raise, error is caught internally - assert isinstance(result, list) - - -# --- export_rag_pipeline_dsl --- - - -def test_export_rag_pipeline_dsl_raises_when_dataset_missing() -> None: - pipeline = Mock() - pipeline.retrieve_dataset.return_value = None - - service = RagPipelineDslService(session=Mock()) - - with pytest.raises(ValueError, match="Missing dataset"): - service.export_rag_pipeline_dsl(pipeline=pipeline) - - -# --- import_rag_pipeline --- - - -def test_import_rag_pipeline_url_fetch_error(mocker: MockerFixture) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.remote_fetcher.make_request", - side_effect=Exception("fetch failed"), - ) - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline( - account=account, import_mode="yaml-url", yaml_url="https://example.com/dsl.yml" - ) - - assert result.status == ImportStatus.FAILED - assert "fetch failed" in result.error - - -def test_import_rag_pipeline_yaml_content_success(mocker: MockerFixture) -> None: - yaml_content = """ -version: 0.1.0 -kind: rag_pipeline -rag_pipeline: - name: Test Pipeline -workflow: - graph: - nodes: - - data: - type: knowledge-index -""" - pipeline = Mock() - pipeline.name = "Test Pipeline" - pipeline.description = "desc" - pipeline.id = "p1" - pipeline.is_published = False - mocker.patch.object(RagPipelineDslService, "_create_or_update_pipeline", return_value=pipeline) - - config_mock = Mock() - config_mock.indexing_technique = "high_quality" - config_mock.embedding_model = "m" - config_mock.embedding_model_provider = "p" - config_mock.summary_index_setting = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", - return_value=config_mock, - ) - - dataset_mock = Mock() - dataset_mock.id = "d1" - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Dataset", return_value=dataset_mock) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock()) - - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - session.scalars.return_value.all.return_value = [] - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content=yaml_content) - - if result.status == ImportStatus.FAILED: - print(f"DEBUG: {result.error}") - assert result.status == ImportStatus.COMPLETED - session.commit.assert_not_called() - session.flush.assert_called() - - -def test_import_rag_pipeline_flushes_new_collection_binding_without_commit(mocker: MockerFixture) -> None: - yaml_content = """ -version: 0.1.0 -kind: rag_pipeline -rag_pipeline: - name: Test Pipeline -workflow: - graph: - nodes: - - data: - type: knowledge-index -""" - pipeline = Mock(id="p1", description="desc", is_published=False) - pipeline.name = "Test Pipeline" - mocker.patch.object(RagPipelineDslService, "_create_or_update_pipeline", return_value=pipeline) - - config_mock = Mock() - config_mock.indexing_technique = "high_quality" - config_mock.embedding_model = "m" - config_mock.embedding_model_provider = "p" - config_mock.chunk_structure = "text_model" - config_mock.retrieval_model.model_dump.return_value = {} - config_mock.summary_index_setting = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", - return_value=config_mock, - ) - - dataset_mock = Mock(id="d1") - binding_mock = Mock(id="b1") - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Dataset", return_value=dataset_mock) - binding_cls = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DatasetCollectionBinding", - return_value=binding_mock, - ) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock()) - - session = cast(MagicMock, Mock()) - session.scalar.return_value = None - session.scalars.return_value.all.return_value = [] - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content=yaml_content) - - assert result.status == ImportStatus.COMPLETED - binding_cls.assert_called_once() - assert dataset_mock.collection_binding_id == "b1" - session.commit.assert_not_called() - assert session.flush.call_count >= 2 - - -def test_import_rag_pipeline_pending_version(mocker: MockerFixture) -> None: - yaml_content = "version: 1.0.0\nkind: rag_pipeline\nrag_pipeline: {name: x}" - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.setex") - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1", id="u1") - - result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content=yaml_content) - - assert result.status == ImportStatus.PENDING - assert result.imported_dsl_version == "1.0.0" - - -# --- confirm_import --- - - -def test_confirm_import_success(mocker: MockerFixture) -> None: - from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelinePendingData - - yaml_content = """ -version: 0.1.0 -kind: rag_pipeline -rag_pipeline: - name: Test Pipeline -workflow: - graph: - nodes: - - data: - type: knowledge-index -""" - pending = RagPipelinePendingData(import_mode="yaml-content", yaml_content=yaml_content, pipeline_id="p1") - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", - return_value=pending.model_dump_json(), - ) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.delete") - - pipeline = Mock() - pipeline.id = "p1" - pipeline.name = "Test Pipeline" - pipeline.description = "desc" - pipeline.retrieve_dataset.return_value = None - - mocker.patch.object(RagPipelineDslService, "_create_or_update_pipeline", return_value=pipeline) - - config_mock = Mock() - config_mock.indexing_technique = "high_quality" - config_mock.embedding_model = "m" - config_mock.embedding_model_provider = "p" - config_mock.chunk_structure = "text_model" - config_mock.retrieval_model.model_dump.return_value = {} - config_mock.summary_index_setting = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", - return_value=config_mock, - ) - - dataset_mock = Mock() - dataset_mock.id = "d1" - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Dataset", return_value=dataset_mock) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.DatasetCollectionBinding", return_value=Mock(id="b1")) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock()) - - service = RagPipelineDslService(session=Mock()) - # Mocking self._session.scalar for the pipeline lookup - service._session.scalar.return_value = pipeline - - account = Mock() - account.id = "u1" - account.current_tenant_id = "t1" - - result = service.confirm_import(account=account, import_id="imp-1") - - assert result.status == ImportStatus.COMPLETED - assert result.pipeline_id == "p1" - assert result.dataset_id == "d1" - - -def test_confirm_import_flushes_new_collection_binding_without_commit(mocker: MockerFixture) -> None: - from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelinePendingData - - yaml_content = """ -version: 0.1.0 -kind: rag_pipeline -rag_pipeline: - name: Test Pipeline -workflow: - graph: - nodes: - - data: - type: knowledge-index -""" - pending = RagPipelinePendingData(import_mode="yaml-content", yaml_content=yaml_content, pipeline_id="p1") - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", - return_value=pending.model_dump_json(), - ) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.delete") - - pipeline = Mock(id="p1", description="desc") - pipeline.name = "Test Pipeline" - pipeline.retrieve_dataset.return_value = None - mocker.patch.object(RagPipelineDslService, "_create_or_update_pipeline", return_value=pipeline) - - config_mock = Mock() - config_mock.indexing_technique = "high_quality" - config_mock.embedding_model = "m" - config_mock.embedding_model_provider = "p" - config_mock.chunk_structure = "text_model" - config_mock.retrieval_model.model_dump.return_value = {} - config_mock.summary_index_setting = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", - return_value=config_mock, - ) - - dataset_mock = Mock(id="d1") - binding_mock = Mock(id="b1") - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Dataset", return_value=dataset_mock) - binding_cls = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DatasetCollectionBinding", - return_value=binding_mock, - ) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock()) - - session = cast(MagicMock, Mock()) - session.scalar.side_effect = [pipeline, None] - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(id="u1", current_tenant_id="t1") - - result = service.confirm_import(account=account, import_id="imp-1") - - assert result.status == ImportStatus.COMPLETED - binding_cls.assert_called_once() - assert dataset_mock.collection_binding_id == "b1" - session.commit.assert_not_called() - assert session.flush.call_count >= 2 - - -# --- _extract_dependencies_from_workflow_graph all types --- - - -@pytest.mark.parametrize( - "node_type", - [ - BuiltinNodeTypes.TOOL, - BuiltinNodeTypes.LLM, - BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL, - BuiltinNodeTypes.PARAMETER_EXTRACTOR, - BuiltinNodeTypes.QUESTION_CLASSIFIER, - ], -) -def test_extract_dependencies_from_workflow_graph_types(mocker, node_type) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_tool_dependency", - return_value="t1", - ) - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", - return_value="m1", - ) - - # Mock all potential node data classes - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.ToolNodeData.model_validate", - return_value=Mock(provider_id="p1"), - ) - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.LLMNodeData.model_validate", - return_value=Mock(model=Mock(provider="p1")), - ) - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeRetrievalNodeData.model_validate", - return_value=Mock( - retrieval_mode="single", - single_retrieval_config=Mock(model=Mock(provider="p1")), - ), - ) - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.ParameterExtractorNodeData.model_validate", - return_value=Mock(model=Mock(provider="p1")), - ) - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.QuestionClassifierNodeData.model_validate", - return_value=Mock(model=Mock(provider="p1")), - ) - - service = RagPipelineDslService(session=Mock()) - graph = {"nodes": [{"data": {"type": node_type}}]} - - result = service._extract_dependencies_from_workflow_graph(graph) - - assert len(result) > 0 - - -# --- _create_or_update_pipeline --- - - -def test_create_or_update_pipeline_create_new(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(current_tenant_id="t1", id="u1") - data = { - "rag_pipeline": {"name": "New", "description": "desc"}, - "workflow": {"graph": {"nodes": []}}, - } - - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.current_user", SimpleNamespace(id="u1")) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Workflow", return_value=Mock()) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock()) - pipeline_cls = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Pipeline") - pipeline_instance = pipeline_cls.return_value - pipeline_instance.tenant_id = "t1" - pipeline_instance.id = "p1" - pipeline_instance.name = "P" - pipeline_instance.is_published = False - session.scalar.return_value = None - - result = service._create_or_update_pipeline(pipeline=None, data=data, account=account, dependencies=[]) - - assert result == pipeline_instance - session.add.assert_called() - session.commit.assert_not_called() - session.flush.assert_called() - - -# --- export_rag_pipeline_dsl comprehensive --- - - -def test_export_rag_pipeline_dsl_with_workflow(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - pipeline = Mock() - pipeline.id = "p1" - pipeline.tenant_id = "t1" - pipeline.name = "P" - pipeline.description = "d" - - dataset = Mock() - dataset.id = "d1" - dataset.name = "D" - dataset.chunk_structure = "text_model" - dataset.doc_form = "text_model" - dataset.icon_info = {"icon": "i"} - pipeline.retrieve_dataset.return_value = dataset - - workflow = Mock() - workflow.app_id = "p1" - workflow.graph_dict = {"nodes": []} - workflow.environment_variables = [] - workflow.conversation_variables = [] - workflow.rag_pipeline_variables = [] - workflow.to_dict.return_value = {"graph": {"nodes": []}} - - session.scalar.return_value = workflow - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.generate_dependencies", - return_value=[], - ) - - result_yaml = service.export_rag_pipeline_dsl(pipeline=pipeline) - data = yaml.safe_load(result_yaml) - - assert data["kind"] == "rag_pipeline" - assert data["rag_pipeline"]["name"] == "D" - assert "workflow" in data - - -# --- _extract_dependencies_from_workflow_graph more types --- - - -def test_extract_dependencies_from_workflow_graph_datasource(mocker: MockerFixture) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DatasourceNodeData.model_validate", - return_value=Mock(provider_type="online", plugin_id="ds1"), - ) - service = RagPipelineDslService(session=Mock()) - graph = {"nodes": [{"data": {"type": BuiltinNodeTypes.DATASOURCE}}]} - - result = service._extract_dependencies_from_workflow_graph(graph) - - assert "ds1" in result - - -def test_import_rag_pipeline_raises_for_invalid_mode() -> None: - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - with pytest.raises(ValueError, match="Invalid import_mode"): - service.import_rag_pipeline(account=account, import_mode="invalid-mode") - - -def test_import_rag_pipeline_yaml_url_requires_url() -> None: - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline(account=account, import_mode="yaml-url", yaml_url=None) - - assert result.status == ImportStatus.FAILED - assert "yaml_url is required" in result.error - - -def test_import_rag_pipeline_yaml_content_requires_content() -> None: - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content=None) - - assert result.status == ImportStatus.FAILED - assert "yaml_content is required" in result.error - - -def test_import_rag_pipeline_rejects_oversized_yaml_content_before_parsing( +def test_extract_dependencies_from_model_config_covers_models_rerankers_and_tools( monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr("services.rag_pipeline.rag_pipeline_dsl_service.DSL_MAX_SIZE", 3) - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") + def model_dependency(provider: str) -> str: + return f"model:{provider}" - result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content="你你") + def tool_dependency(provider: str) -> str: + return f"tool:{provider}" - assert result.status == ImportStatus.FAILED - assert result.error == "File size exceeds the limit of 10MB" - - -def test_import_rag_pipeline_yaml_content_requires_mapping() -> None: - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content="- one\n- two") - - assert result.status == ImportStatus.FAILED - assert "content must be a mapping" in result.error - - -def test_import_rag_pipeline_rejects_oversized_yaml_content_by_bytes( - monkeypatch: pytest.MonkeyPatch, -) -> None: - monkeypatch.setattr("services.rag_pipeline.rag_pipeline_dsl_service.DSL_MAX_SIZE", 1) - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content="é") - - assert result.status == ImportStatus.FAILED - assert "10MB" in result.error - - -def test_confirm_import_returns_failed_when_pending_data_is_invalid_type(mocker: MockerFixture) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", return_value=object()) - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.confirm_import(import_id="imp-1", account=account) - - assert result.status == ImportStatus.FAILED - assert "Invalid import information" in result.error - - -def test_append_workflow_export_data_filters_credentials(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - workflow = Mock() - workflow.graph_dict = {"nodes": []} - workflow.to_dict.return_value = { - "graph": { - "nodes": [ - { - "data": { - "type": BuiltinNodeTypes.TOOL, - "credential_id": "secret", - } - }, - { - "data": { - "type": BuiltinNodeTypes.AGENT, - "agent_parameters": {"tools": {"value": [{"credential_id": "secret-agent"}]}}, - } - }, - ] - } - } - session.scalar.return_value = workflow - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.generate_dependencies", - return_value=[], - ) - export_data: dict[str, Any] = {} - pipeline = Mock(id="p1", tenant_id="t1") - - service._append_workflow_export_data(export_data=export_data, pipeline=pipeline, include_secret=False) - - nodes = export_data["workflow"]["graph"]["nodes"] - assert "credential_id" not in nodes[0]["data"] - assert "credential_id" not in nodes[1]["data"]["agent_parameters"]["tools"]["value"][0] - - -def test_create_rag_pipeline_dataset_raises_when_name_conflicts(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - session.scalar.return_value = Mock() - create_entity = RagPipelineDatasetCreateEntity( - name="Existing Name", - description="", - icon_info=IconInfo(icon="book"), - permission="only_me", - yaml_content="x", - ) - - with pytest.raises(ValueError, match="already exists"): - service.create_rag_pipeline_dataset("tenant-1", create_entity) - - -def test_create_rag_pipeline_dataset_generates_name_when_missing(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - session.scalar.return_value = None - session.scalars.return_value.all.return_value = [Mock(name="Untitled")] - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.generate_incremental_name", return_value="Untitled 2") - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.current_user", Mock(id="u1", current_tenant_id="t1")) - mocker.patch.object( - service, - "import_rag_pipeline", - return_value=SimpleNamespace( - id="imp-1", - dataset_id="d1", - pipeline_id="p1", - status=ImportStatus.COMPLETED, - imported_dsl_version="0.1.0", - current_dsl_version="0.1.0", - error="", - ), - ) - create_entity = RagPipelineDatasetCreateEntity( - name="", - description="", - icon_info=IconInfo(icon="book"), - permission="only_me", - yaml_content="x", - ) - - result = service.create_rag_pipeline_dataset("tenant-1", create_entity) - - assert create_entity.name == "Untitled 2" - assert result["status"] == ImportStatus.COMPLETED - - -def test_append_workflow_export_data_encrypts_knowledge_retrieval_dataset_ids(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - workflow = Mock() - workflow.graph_dict = {"nodes": []} - workflow.to_dict.return_value = { - "graph": { - "nodes": [ - { - "data": { - "type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL, - "dataset_ids": ["d1", "d2"], - } - } - ] - } - } - session.scalar.return_value = workflow - mocker.patch.object(service, "encrypt_dataset_id", side_effect=lambda dataset_id, tenant_id: f"enc-{dataset_id}") - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.generate_dependencies", - return_value=[], - ) - export_data: dict[str, Any] = {} - pipeline = Mock(id="p1", tenant_id="t1") - - service._append_workflow_export_data(export_data=export_data, pipeline=pipeline, include_secret=False) - - ids = export_data["workflow"]["graph"]["nodes"][0]["data"]["dataset_ids"] - assert ids == ["enc-d1", "enc-d2"] - - -def test_confirm_import_updates_existing_dataset(mocker: MockerFixture) -> None: - from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelinePendingData - - yaml_content = ( - "version: 0.1.0\n" - "kind: rag_pipeline\n" - "rag_pipeline: {name: x}\n" - "workflow: {graph: {nodes: [{data: {type: knowledge-index}}]}}" - ) - pending = RagPipelinePendingData(import_mode="yaml-content", yaml_content=yaml_content, pipeline_id="p1") - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", - return_value=pending.model_dump_json(), - ) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.delete") - pipeline = Mock(id="p1", name="P", description="D") - dataset = Mock(id="d1") - pipeline.retrieve_dataset.return_value = dataset - mocker.patch.object(RagPipelineDslService, "_create_or_update_pipeline", return_value=pipeline) - config_mock = Mock() - config_mock.indexing_technique = "economy" - config_mock.keyword_number = 3 - config_mock.retrieval_model.model_dump.return_value = {"top_k": 3} - config_mock.chunk_structure = "text_model" - config_mock.summary_index_setting = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", - return_value=config_mock, - ) - service = RagPipelineDslService(session=Mock()) - service._session.scalar.return_value = pipeline - account = Mock(id="u1", current_tenant_id="t1") - - result = service.confirm_import(import_id="imp-1", account=account) - - assert result.status == ImportStatus.COMPLETED - assert dataset.indexing_technique == "economy" - - -def test_import_rag_pipeline_yaml_url_handles_empty_content_after_github_rewrite(mocker: MockerFixture) -> None: - response = Mock() - response.raise_for_status.return_value = None - response.content = b"" - get_mock = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.remote_fetcher.make_request", - return_value=response, - ) - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline( - account=account, - import_mode="yaml-url", - yaml_url="https://github.com/langgenius/dify/blob/main/pipeline.yml", - ) - - assert result.status == ImportStatus.FAILED - assert "Empty content from url" in result.error - called_url = get_mock.call_args.args[1] - assert "raw.githubusercontent.com" in called_url - - -def test_create_or_update_pipeline_decrypts_knowledge_retrieval_dataset_ids(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(id="u1", current_tenant_id="t1") - pipeline = Mock(id="p1", tenant_id="t1", name="N", description="D") - data = { - "rag_pipeline": {"name": "N2", "description": "D2"}, - "workflow": { - "graph": { - "nodes": [ - { - "data": { - "type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL, - "dataset_ids": ["enc-1", "enc-2"], - } - } - ] - } - }, - } - draft_workflow = Mock(id="wf1") - session.scalar.return_value = draft_workflow - mocker.patch.object(service, "decrypt_dataset_id", side_effect=["d1", None]) - - result = service._create_or_update_pipeline(pipeline=pipeline, data=data, account=account) - - assert result is pipeline - assert data["workflow"]["graph"]["nodes"][0]["data"]["dataset_ids"] == ["d1"] - assert draft_workflow.graph is not None - - -def test_create_or_update_pipeline_creates_draft_when_missing(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(id="u1", current_tenant_id="t1") - pipeline = Mock(id="p1", tenant_id="t1", name="N", description="D") - data = {"rag_pipeline": {"name": "N2", "description": "D2"}, "workflow": {"graph": {"nodes": []}}} - session.scalar.return_value = None - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock()) - workflow_cls = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Workflow") - workflow_cls.return_value.id = "wf-new" - - service._create_or_update_pipeline(pipeline=pipeline, data=data, account=account) - - assert pipeline.workflow_id == "wf-new" - - -def test_import_rag_pipeline_url_size_exceeds_limit(mocker: MockerFixture) -> None: - response = Mock() - response.raise_for_status.return_value = None - response.content = b"x" * (10 * 1024 * 1024 + 1) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.remote_fetcher.make_request", return_value=response) - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline( - account=account, - import_mode="yaml-url", - yaml_url="https://example.com/pipeline.yaml", - ) - - assert result.status == ImportStatus.FAILED - assert "10MB" in result.error - - -def test_import_rag_pipeline_rejects_oversized_yaml_url_bytes_before_decode( - mocker: MockerFixture, - monkeypatch: pytest.MonkeyPatch, -) -> None: - monkeypatch.setattr("services.rag_pipeline.rag_pipeline_dsl_service.DSL_MAX_SIZE", 1) - response = Mock() - response.raise_for_status.return_value = None - response.content = b"\xff\xff" - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.remote_fetcher.make_request", return_value=response) - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline( - account=account, - import_mode="yaml-url", - yaml_url="https://example.com/pipeline.yaml", - ) - - assert result.status == ImportStatus.FAILED - assert "10MB" in result.error - - -def test_import_rag_pipeline_returns_decode_error_for_invalid_yaml_url_bytes(mocker: MockerFixture) -> None: - response = Mock() - response.raise_for_status.return_value = None - response.content = b"\xff" - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.remote_fetcher.make_request", return_value=response) - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline( - account=account, - import_mode="yaml-url", - yaml_url="https://example.com/pipeline.yaml", - ) - - assert result.status == ImportStatus.FAILED - assert "utf-8" in result.error - - -def test_import_rag_pipeline_fails_when_rag_pipeline_data_missing() -> None: - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - result = service.import_rag_pipeline( - account=account, - import_mode="yaml-content", - yaml_content="version: 0.1.0\nkind: rag_pipeline\nworkflow: {}", - ) - - assert result.status == ImportStatus.FAILED - assert "Missing rag_pipeline data" in result.error - - -def test_import_rag_pipeline_fails_when_pipeline_id_not_found() -> None: - session = cast(MagicMock, Mock()) - session.scalar.return_value = None - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline( - account=account, - import_mode="yaml-content", - yaml_content="version: 0.1.0\nkind: rag_pipeline\nrag_pipeline: {name: x}\nworkflow: {}", - pipeline_id="missing-pipeline", - ) - - assert result.status == ImportStatus.FAILED - assert "Pipeline not found" in result.error - - -def test_import_rag_pipeline_fails_for_non_string_version_type() -> None: - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1") - - result = service.import_rag_pipeline( - account=account, - import_mode="yaml-content", - yaml_content="version: 1\nkind: rag_pipeline\nrag_pipeline: {name: x}\nworkflow: {}", - ) - - assert result.status == ImportStatus.FAILED - assert "Invalid version type" in result.error - - -def test_append_workflow_export_data_raises_when_draft_workflow_missing() -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - session.scalar.return_value = None - - with pytest.raises(ValueError, match="Missing draft workflow configuration"): - service._append_workflow_export_data(export_data={}, pipeline=Mock(tenant_id="t1"), include_secret=False) - - -def test_append_workflow_export_data_keeps_secret_fields_when_include_secret_true(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - workflow = Mock() - workflow.graph_dict = {"nodes": []} - workflow.to_dict.return_value = { - "graph": { - "nodes": [ - {"data": {"type": BuiltinNodeTypes.TOOL, "credential_id": "tool-secret"}}, - { - "data": { - "type": BuiltinNodeTypes.AGENT, - "agent_parameters": {"tools": {"value": [{"credential_id": "agent-secret"}]}}, - } - }, - ] - } - } - session.scalar.return_value = workflow - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.generate_dependencies", - return_value=[], - ) - - export_data: dict[str, object] = {} - service._append_workflow_export_data(export_data=export_data, pipeline=Mock(tenant_id="t1"), include_secret=True) - - workflow_data = cast(dict[str, object], export_data["workflow"]) - graph = cast(dict[str, object], workflow_data["graph"]) - nodes = cast(list[dict[str, object]], graph["nodes"]) - node0_data = cast(dict[str, object], nodes[0]["data"]) - node1_data = cast(dict[str, object], nodes[1]["data"]) - agent_parameters = cast(dict[str, object], node1_data["agent_parameters"]) - tools = cast(dict[str, object], agent_parameters["tools"]) - tool_values = cast(list[dict[str, object]], tools["value"]) - assert node0_data["credential_id"] == "tool-secret" - assert tool_values[0]["credential_id"] == "agent-secret" - - -def test_extract_dependencies_from_workflow_graph_skips_local_file_datasource(mocker: MockerFixture) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DatasourceNodeData.model_validate", - return_value=Mock(provider_type="local_file", plugin_id="plugin-x"), - ) - service = RagPipelineDslService(session=Mock()) - - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": BuiltinNodeTypes.DATASOURCE}}]} - ) - - assert result == [] - - -def test_extract_dependencies_from_workflow_graph_knowledge_index_reranking(mocker: MockerFixture) -> None: - analyze = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", - side_effect=lambda provider: f"dep:{provider}", - ) - knowledge = Mock() - knowledge.indexing_technique = "high_quality" - knowledge.embedding_model_provider = "embed-provider" - knowledge.retrieval_model.reranking_mode = "reranking_model" - knowledge.retrieval_model.reranking_enable = True - knowledge.retrieval_model.reranking_model.reranking_provider_name = "rerank-provider" - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", - return_value=knowledge, - ) - service = RagPipelineDslService(session=Mock()) - - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": KNOWLEDGE_INDEX_NODE_TYPE}}]} - ) - - assert result == ["dep:embed-provider", "dep:rerank-provider"] - assert analyze.call_count == 2 - - -def test_extract_dependencies_from_workflow_graph_multiple_retrieval_weighted_score(mocker: MockerFixture) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", - return_value="dep:weighted", - ) - retrieval = Mock() - retrieval.retrieval_mode = "multiple" - retrieval.multiple_retrieval_config.reranking_mode = "weighted_score" - retrieval.multiple_retrieval_config.weights.vector_setting.embedding_provider_name = "emb-provider" - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeRetrievalNodeData.model_validate", - return_value=retrieval, - ) - service = RagPipelineDslService(session=Mock()) - - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL}}]} - ) - - assert result == ["dep:weighted"] - - -def test_extract_dependencies_from_workflow_graph_multiple_retrieval_reranking_model(mocker: MockerFixture) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", - return_value="dep:rerank", - ) - retrieval = Mock() - retrieval.retrieval_mode = "multiple" - retrieval.multiple_retrieval_config.reranking_mode = "reranking_model" - retrieval.multiple_retrieval_config.reranking_model.provider = "rerank-provider" - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeRetrievalNodeData.model_validate", - return_value=retrieval, - ) - service = RagPipelineDslService(session=Mock()) - - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL}}]} - ) - - assert result == ["dep:rerank"] - - -def test_extract_dependencies_from_model_config_includes_dataset_reranking_and_tools(mocker: MockerFixture) -> None: - model_analyze = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", - side_effect=["dep:model", "dep:rerank"], - ) - tool_analyze = mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_tool_dependency", - return_value="dep:tool", - ) - config = { + analyze_model = Mock(side_effect=model_dependency) + analyze_tool = Mock(side_effect=tool_dependency) + monkeypatch.setattr(module.DependenciesAnalysisService, "analyze_model_provider_dependency", analyze_model) + monkeypatch.setattr(module.DependenciesAnalysisService, "analyze_tool_dependency", analyze_tool) + model_config: dict[str, Any] = { "model": {"provider": "openai"}, "dataset_configs": { "datasets": { @@ -1167,360 +230,742 @@ def test_extract_dependencies_from_model_config_includes_dataset_reranking_and_t "agent_mode": {"tools": [{"provider_id": "google"}]}, } - deps = RagPipelineDslService._extract_dependencies_from_model_config(config) + dependencies = RagPipelineDslService._extract_dependencies_from_model_config(model_config) - assert deps == ["dep:model", "dep:rerank", "dep:tool"] - assert model_analyze.call_count == 2 - tool_analyze.assert_called_once_with("google") + assert dependencies == ["model:openai", "model:cohere", "tool:google"] + assert analyze_model.call_args_list == [call("openai"), call("cohere")] + analyze_tool.assert_called_once_with("google") -def test_check_version_compatibility_hits_major_older_branch(mocker: MockerFixture) -> None: - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.CURRENT_DSL_VERSION", "1.0.0") - - status = check_version_compatibility("0.9.0", rag_pipeline_dsl_service.CURRENT_DSL_VERSION) - - assert status == ImportStatus.PENDING - - -def test_import_rag_pipeline_sets_default_version_and_kind(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(current_tenant_id="t1") - pipeline = Mock(id="p1", name="P", description="D", is_published=False) - mocker.patch.object(service, "_create_or_update_pipeline", return_value=pipeline) - config = Mock() - config.indexing_technique = "economy" - config.keyword_number = 2 - config.retrieval_model.model_dump.return_value = {} - config.summary_index_setting = None - config.chunk_structure = "text_model" - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", - return_value=config, +def test_extract_workflow_dependencies_uses_llm_environment_variable_provider( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService +) -> None: + workflow = SimpleNamespace( + graph_dict={ + "nodes": [ + { + "id": "llm-node", + "data": { + "type": "llm", + "title": "LLM", + "model": {"provider": "old-provider", "name": "old-model", "mode": "chat"}, + "model_selector": ["env", "shared_model"], + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + }, + } + ] + }, + environment_variables=[ + LLMEnvironmentVariable( + name="shared_model", + value={"provider": "new-provider", "name": "new-model", "mode": "chat"}, + ) + ], ) - dataset = Mock(id="d1") - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Dataset", return_value=dataset) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock()) - session.scalars.return_value.all.return_value = [] - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.generate_incremental_name", return_value="P") - - result = service.import_rag_pipeline( - account=account, - import_mode="yaml-content", - yaml_content="rag_pipeline: {name: x}\nworkflow: {graph: {nodes: [{data: {type: knowledge-index}}]}}", + analyze_dependency = Mock(side_effect=lambda provider: provider) + monkeypatch.setattr( + module.DependenciesAnalysisService, + "analyze_model_provider_dependency", + analyze_dependency, ) - assert result.status == ImportStatus.COMPLETED - assert result.imported_dsl_version == "0.1.0" + result = service._extract_dependencies_from_workflow(cast(Workflow, workflow)) + + assert result == ["new-provider"] + analyze_dependency.assert_called_once_with("new-provider") -def test_import_rag_pipeline_creates_pending_for_dependencies(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(current_tenant_id="t1") - setex = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.setex") - yaml_content = """ -version: 1.0.0 -kind: rag_pipeline -rag_pipeline: {name: x} -dependencies: - - type: marketplace - value: - marketplace_plugin_unique_identifier: langgenius/example:0.1.0 -workflow: {graph: {nodes: []}} -""" - - result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content=yaml_content) - - assert result.status == ImportStatus.PENDING - setex.assert_called_once() - - -def test_confirm_import_returns_failed_when_pending_pipeline_missing(mocker: MockerFixture) -> None: - from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelinePendingData - - pending = RagPipelinePendingData(import_mode="yaml-content", yaml_content="version: 0.1.0", pipeline_id="p1") - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", return_value=pending.model_dump_json() +@pytest.mark.parametrize("model_selector", [[], ["env", "missing_model"]]) +def test_extract_workflow_dependencies_tolerates_unresolved_llm_environment_reference( + monkeypatch: pytest.MonkeyPatch, + service: RagPipelineDslService, + model_selector: list[str], +) -> None: + workflow = SimpleNamespace( + graph_dict={ + "nodes": [ + { + "id": "llm-node", + "data": { + "type": "llm", + "title": "LLM", + "model": {"provider": "old-provider", "name": "old-model", "mode": "chat"}, + "model_selector": model_selector, + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + }, + } + ] + }, + environment_variables=[], ) - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - session.scalar.return_value = None - mocker.patch.object(RagPipelineDslService, "_create_or_update_pipeline", side_effect=ValueError("pipeline missing")) - - result = service.confirm_import(import_id="imp-1", account=Mock(current_tenant_id="t1")) - - assert result.status == ImportStatus.FAILED - - -def test_append_workflow_export_data_skips_empty_node_data(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - workflow = Mock() - workflow.graph_dict = {"nodes": []} - workflow.to_dict.return_value = {"graph": {"nodes": [{"data": {}}, {}]}} - session.scalar.return_value = workflow - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.generate_dependencies", - return_value=[], - ) - export_data = {} - - service._append_workflow_export_data(export_data=export_data, pipeline=Mock(tenant_id="t1"), include_secret=False) - - assert "workflow" in export_data - - -def test_extract_dependencies_from_workflow_graph_multiple_config_none(mocker: MockerFixture) -> None: - retrieval = Mock() - retrieval.retrieval_mode = "multiple" - retrieval.multiple_retrieval_config = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeRetrievalNodeData.model_validate", - return_value=retrieval, - ) - service = RagPipelineDslService(session=Mock()) - - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL}}]} + analyze_dependency = Mock(side_effect=lambda provider: provider) + monkeypatch.setattr( + module.DependenciesAnalysisService, + "analyze_model_provider_dependency", + analyze_dependency, ) - assert result == [] + result = service._extract_dependencies_from_workflow(cast(Workflow, workflow)) + + assert result == ["old-provider"] + analyze_dependency.assert_called_once_with("old-provider") -def test_extract_dependencies_from_workflow_graph_single_config_none(mocker: MockerFixture) -> None: - retrieval = Mock() - retrieval.retrieval_mode = "single" - retrieval.single_retrieval_config = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeRetrievalNodeData.model_validate", - return_value=retrieval, +def test_extract_dependencies_from_workflow_graph_covers_plugin_and_model_nodes( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService +) -> None: + def model_dependency(provider: str) -> str: + return f"model:{provider}" + + def tool_dependency(provider: str) -> str: + return f"tool:{provider}" + + monkeypatch.setattr( + module.DependenciesAnalysisService, + "analyze_model_provider_dependency", + Mock(side_effect=model_dependency), ) - service = RagPipelineDslService(session=Mock()) + monkeypatch.setattr( + module.DependenciesAnalysisService, + "analyze_tool_dependency", + Mock(side_effect=tool_dependency), + ) + monkeypatch.setattr( + module.ToolNodeData, + "model_validate", + Mock(return_value=SimpleNamespace(provider_id="tool-provider")), + ) + monkeypatch.setattr( + module.DatasourceNodeData, + "model_validate", + Mock( + side_effect=[ + SimpleNamespace(provider_type="online_document", plugin_id="datasource-plugin"), + SimpleNamespace(provider_type="local_file", plugin_id="ignored-local-file"), + ] + ), + ) + monkeypatch.setattr( + module.LLMNodeData, + "model_validate", + Mock(return_value=SimpleNamespace(model=SimpleNamespace(provider="llm-provider"))), + ) + monkeypatch.setattr( + module.QuestionClassifierNodeData, + "model_validate", + Mock(return_value=SimpleNamespace(model=SimpleNamespace(provider="classifier-provider"))), + ) + monkeypatch.setattr( + module.ParameterExtractorNodeData, + "model_validate", + Mock(return_value=SimpleNamespace(model=SimpleNamespace(provider="extractor-provider"))), + ) + graph: dict[str, Any] = { + "nodes": [ + {"data": {"type": BuiltinNodeTypes.TOOL}}, + {"data": {"type": BuiltinNodeTypes.DATASOURCE}}, + {"data": {"type": BuiltinNodeTypes.DATASOURCE}}, + {"data": {"type": BuiltinNodeTypes.LLM}}, + {"data": {"type": BuiltinNodeTypes.QUESTION_CLASSIFIER}}, + {"data": {"type": BuiltinNodeTypes.PARAMETER_EXTRACTOR}}, + {"data": {"type": "unknown"}}, + ] + } - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL}}]} + dependencies = service._extract_dependencies_from_workflow_graph(graph) + + assert dependencies == [ + "tool:tool-provider", + "datasource-plugin", + "model:llm-provider", + "model:classifier-provider", + "model:extractor-provider", + ] + + +def test_extract_dependencies_from_workflow_graph_covers_knowledge_variants( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService +) -> None: + def model_dependency(provider: str) -> str: + return f"model:{provider}" + + monkeypatch.setattr( + module.DependenciesAnalysisService, + "analyze_model_provider_dependency", + Mock(side_effect=model_dependency), + ) + knowledge_config = SimpleNamespace( + indexing_technique="high_quality", + embedding_model_provider="embedding-provider", + retrieval_model=SimpleNamespace( + reranking_mode="reranking_model", + reranking_enable=True, + reranking_model=SimpleNamespace(reranking_provider_name="knowledge-reranker"), + ), + ) + monkeypatch.setattr(module.KnowledgeConfiguration, "model_validate", Mock(return_value=knowledge_config)) + retrieval_configs = [ + SimpleNamespace( + retrieval_mode="multiple", + multiple_retrieval_config=SimpleNamespace( + reranking_mode="weighted_score", + weights=SimpleNamespace(vector_setting=SimpleNamespace(embedding_provider_name="weighted-embedding")), + ), + ), + SimpleNamespace( + retrieval_mode="multiple", + multiple_retrieval_config=SimpleNamespace( + reranking_mode="reranking_model", + reranking_model=SimpleNamespace(provider="retrieval-reranker"), + ), + ), + SimpleNamespace( + retrieval_mode="single", + single_retrieval_config=SimpleNamespace(model=SimpleNamespace(provider="single-provider")), + ), + SimpleNamespace(retrieval_mode="multiple", multiple_retrieval_config=None), + SimpleNamespace(retrieval_mode="single", single_retrieval_config=None), + ] + monkeypatch.setattr( + module.KnowledgeRetrievalNodeData, + "model_validate", + Mock(side_effect=retrieval_configs), + ) + graph: dict[str, Any] = { + "nodes": [ + {"data": {"type": KNOWLEDGE_INDEX_NODE_TYPE}}, + *[{"data": {"type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL}} for _ in retrieval_configs], + ] + } + + dependencies = service._extract_dependencies_from_workflow_graph(graph) + + assert dependencies == [ + "model:embedding-provider", + "model:knowledge-reranker", + "model:weighted-embedding", + "model:retrieval-reranker", + "model:single-provider", + ] + + +def test_extract_dependencies_from_workflow_graph_ignores_malformed_nodes( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService +) -> None: + monkeypatch.setattr(module.ToolNodeData, "model_validate", Mock(side_effect=ValueError("invalid tool"))) + + dependencies = service._extract_dependencies_from_workflow_graph( + {"nodes": [{"data": {"type": BuiltinNodeTypes.TOOL}}]} ) - assert result == [] + assert dependencies == [] -def test_create_or_update_pipeline_raises_when_workflow_missing() -> None: - service = RagPipelineDslService(session=Mock()) - account = Mock(current_tenant_id="t1", id="u1") - - with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"): - service._create_or_update_pipeline(pipeline=None, data={"rag_pipeline": {"name": "x"}}, account=account) - - -def test_import_rag_pipeline_with_pipeline_id_uses_existing_dataset(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - existing_dataset = Mock(id="d1", chunk_structure="text_model") - existing_pipeline = Mock(id="p1", name="P", description="D", is_published=False) - existing_pipeline.retrieve_dataset.return_value = existing_dataset - session.scalar.return_value = existing_pipeline - mocker.patch.object(service, "_create_or_update_pipeline", return_value=existing_pipeline) - config = Mock() - config.indexing_technique = "economy" - config.keyword_number = 3 - config.chunk_structure = "text_model" - config.summary_index_setting = {"enabled": True} - config.retrieval_model.model_dump.return_value = {"top_k": 3} - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", return_value=config - ) - - yaml_content = ( - "version: 0.1.0\n" - "kind: rag_pipeline\n" - "rag_pipeline: {name: x}\n" - "workflow: {graph: {nodes: [{data: {type: knowledge-index}}]}}" - ) - - result = service.import_rag_pipeline( - account=Mock(id="u1", current_tenant_id="t1"), - import_mode="yaml-content", - yaml_content=yaml_content, - pipeline_id="p1", - ) - - assert result.status == ImportStatus.COMPLETED - assert result.dataset_id == "d1" - - -def test_import_rag_pipeline_raises_for_chunk_structure_mismatch_on_published(mocker: MockerFixture) -> None: - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - existing_dataset = Mock(id="d1", chunk_structure="hierarchical_model") - existing_pipeline = Mock(id="p1", name="P", description="D", is_published=True) - existing_pipeline.retrieve_dataset.return_value = existing_dataset - session.scalar.return_value = existing_pipeline - mocker.patch.object(service, "_create_or_update_pipeline", return_value=existing_pipeline) - config = Mock() - config.chunk_structure = "text_model" - config.indexing_technique = "economy" - config.keyword_number = 3 - config.summary_index_setting = None - config.retrieval_model.model_dump.return_value = {} - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", return_value=config - ) - - yaml_content = ( - "version: 0.1.0\n" - "kind: rag_pipeline\n" - "rag_pipeline: {name: x}\n" - "workflow: {graph: {nodes: [{data: {type: knowledge-index}}]}}" - ) - - result = service.import_rag_pipeline( - account=Mock(id="u1", current_tenant_id="t1"), - import_mode="yaml-content", - yaml_content=yaml_content, - pipeline_id="p1", - ) - - assert result.status == ImportStatus.FAILED - assert "Chunk structure is not compatible" in result.error - - -def test_import_rag_pipeline_fails_when_no_knowledge_index_node(mocker: MockerFixture) -> None: - service = RagPipelineDslService(session=Mock()) - pipeline = Mock(id="p1", name="P", description="D", is_published=False) - mocker.patch.object(service, "_create_or_update_pipeline", return_value=pipeline) - - yaml_content = ( - "version: 0.1.0\n" - "kind: rag_pipeline\n" - "rag_pipeline: {name: x}\n" - "workflow: {graph: {nodes: [{data: {type: start}}]}}" - ) - - result = service.import_rag_pipeline( - account=Mock(id="u1", current_tenant_id="t1"), - import_mode="yaml-content", - yaml_content=yaml_content, - ) - - assert result.status == ImportStatus.FAILED - assert "Knowledge Index node" in result.error - - -def test_confirm_import_fails_when_no_knowledge_index_node(mocker: MockerFixture) -> None: - from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelinePendingData - - yaml_content = ( - "version: 0.1.0\n" - "kind: rag_pipeline\n" - "rag_pipeline: {name: x}\n" - "workflow: {graph: {nodes: [{data: {type: start}}]}}" - ) - - pending = RagPipelinePendingData( - import_mode="yaml-content", - yaml_content=yaml_content, - pipeline_id=None, - ) - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get", return_value=pending.model_dump_json() - ) - service = RagPipelineDslService(session=Mock()) - pipeline = Mock(id="p1", name="P", description="D") - pipeline.retrieve_dataset.return_value = None - mocker.patch.object(service, "_create_or_update_pipeline", return_value=pipeline) - - result = service.confirm_import(import_id="imp-1", account=Mock(id="u1", current_tenant_id="t1")) - - assert result.status == ImportStatus.FAILED - assert "Knowledge Index node" in result.error - - -def test_create_or_update_pipeline_saves_dependencies_to_redis(mocker: MockerFixture) -> None: +def test_check_dependencies_reads_redis_for_persisted_pipeline( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService, sqlite_session: Session +) -> None: from core.plugin.entities.plugin import PluginDependency, PluginDependencyType + from services.rag_pipeline.rag_pipeline_dsl_service import CheckDependenciesPendingData - session = cast(MagicMock, Mock()) - service = RagPipelineDslService(session=cast(Session, session)) - account = Mock(id="u1", current_tenant_id="t1") - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.current_user", SimpleNamespace(id="u1")) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Workflow", return_value=Mock(id="wf-1")) - mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock()) - pipeline_cls = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Pipeline") - pipeline = pipeline_cls.return_value - pipeline.tenant_id = "t1" - pipeline.id = "p1" - session.scalar.return_value = None - setex = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.setex") + pipeline = _pipeline(sqlite_session) + monkeypatch.setattr(module.redis_client, "get", Mock(return_value=None)) + assert service.check_dependencies(pipeline=pipeline).leaked_dependencies == [] dependency = PluginDependency( type=PluginDependencyType.Marketplace, - value=PluginDependency.Marketplace(marketplace_plugin_unique_identifier="langgenius/example:0.1.0"), + value=PluginDependency.Marketplace(marketplace_plugin_unique_identifier="test/plugin:0.1.0"), + ) + pending = CheckDependenciesPendingData(dependencies=[dependency], pipeline_id=pipeline.id) + monkeypatch.setattr(module.redis_client, "get", Mock(return_value=pending.model_dump_json())) + monkeypatch.setattr(module.DependenciesAnalysisService, "get_leaked_dependencies", Mock(return_value=[dependency])) + assert service.check_dependencies(pipeline=pipeline).leaked_dependencies == [dependency] + + +@pytest.mark.parametrize( + ("import_mode", "yaml_content", "yaml_url", "error"), + [ + (ImportMode.YAML_URL.value, None, None, "yaml_url is required"), + (ImportMode.YAML_CONTENT.value, None, None, "yaml_content is required"), + ( + ImportMode.YAML_CONTENT.value, + "- item", + None, + "content must be a mapping", + ), + ( + ImportMode.YAML_CONTENT.value, + "version: 1\nkind: rag_pipeline", + None, + "Invalid version type", + ), + ( + ImportMode.YAML_CONTENT.value, + "version: 0.1.0\nkind: rag_pipeline", + None, + "Missing rag_pipeline data", + ), + ], +) +def test_import_validation_errors( + service: RagPipelineDslService, + import_mode: str, + yaml_content: str | None, + yaml_url: str | None, + error: str, +) -> None: + result = service.import_rag_pipeline( + account=_account(), + import_mode=import_mode, + yaml_content=yaml_content, + yaml_url=yaml_url, + ) + assert result.status == ImportStatus.FAILED + assert error in result.error + + +def test_import_rejects_invalid_mode(service: RagPipelineDslService) -> None: + with pytest.raises(ValueError, match="Invalid import_mode"): + service.import_rag_pipeline(account=_account(), import_mode="invalid") + + +def test_import_url_boundary_failure(monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService) -> None: + monkeypatch.setattr(module.remote_fetcher, "make_request", Mock(side_effect=RuntimeError("network down"))) + result = service.import_rag_pipeline( + account=_account(), import_mode=ImportMode.YAML_URL.value, yaml_url="https://example.com/pipeline.yml" + ) + assert result.status == ImportStatus.FAILED + assert "network down" in result.error + + +def test_import_rejects_oversized_unicode_content_by_encoded_size( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService +) -> None: + monkeypatch.setattr(module, "DSL_MAX_SIZE", 3) + + result = service.import_rag_pipeline( + account=_account(), + import_mode=ImportMode.YAML_CONTENT.value, + yaml_content="你你", ) - service._create_or_update_pipeline( - pipeline=None, - data={"rag_pipeline": {"name": "x"}, "workflow": {"graph": {"nodes": []}}}, - account=account, - dependencies=[dependency], + assert result.status == ImportStatus.FAILED + assert result.error == "File size exceeds the limit of 10MB" + + +@pytest.mark.parametrize( + ("raw_content", "size_limit", "error"), + [ + (b"\xff\xff", 1, "File size exceeds the limit of 10MB"), + (b"\xff", 2, "utf-8"), + (b"", 2, "Empty content from url"), + ], +) +def test_import_validates_url_bytes_before_parsing( + monkeypatch: pytest.MonkeyPatch, + service: RagPipelineDslService, + raw_content: bytes, + size_limit: int, + error: str, +) -> None: + response = MagicMock() + response.content = raw_content + monkeypatch.setattr(module, "DSL_MAX_SIZE", size_limit) + monkeypatch.setattr(module.remote_fetcher, "make_request", Mock(return_value=response)) + + result = service.import_rag_pipeline( + account=_account(), + import_mode=ImportMode.YAML_URL.value, + yaml_url="https://example.com/pipeline.yml", ) + assert result.status == ImportStatus.FAILED + assert error in result.error + + +def test_import_supplies_default_version_and_kind( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService +) -> None: + monkeypatch.setattr(module.KnowledgeConfiguration, "model_validate", Mock(return_value=_knowledge_configuration())) + yaml_content = f""" +rag_pipeline: + name: Imported +workflow: + graph: + nodes: + - data: + type: {KNOWLEDGE_INDEX_NODE_TYPE} +""" + + result = service.import_rag_pipeline( + account=_account(), + import_mode=ImportMode.YAML_CONTENT.value, + yaml_content=yaml_content, + ) + + assert result.status == ImportStatus.COMPLETED + assert result.imported_dsl_version == module.CURRENT_DSL_VERSION + + +def test_confirm_import_rejects_non_serialized_pending_data( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService +) -> None: + monkeypatch.setattr(module.redis_client, "get", Mock(return_value=object())) + + result = service.confirm_import(import_id="import-1", account=_account()) + + assert result.status == ImportStatus.FAILED + assert result.error == "Invalid import information" + + +def test_import_pending_version_stores_redis(monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService) -> None: + setex = Mock() + monkeypatch.setattr(module.redis_client, "setex", setex) + result = service.import_rag_pipeline( + account=_account(), import_mode=ImportMode.YAML_CONTENT.value, yaml_content=_valid_dsl(version="1.0.0") + ) + assert result.status == ImportStatus.PENDING + assert setex.call_args.args[0] == f"app_import_info:{result.id}" + pending = RagPipelinePendingData.model_validate_json(setex.call_args.args[2]) + assert pending.tenant_id == "tenant-1" + assert pending.account_id == "account-1" setex.assert_called_once() -def test_extract_dependencies_from_workflow_graph_knowledge_index_without_embedding_provider( - mocker: MockerFixture, +def test_import_creates_real_pipeline_dataset_binding_and_workflow_without_commit( + monkeypatch: pytest.MonkeyPatch, + service: RagPipelineDslService, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], ) -> None: - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", - return_value="dep", + monkeypatch.setattr(module.KnowledgeConfiguration, "model_validate", Mock(return_value=_knowledge_configuration())) + result = service.import_rag_pipeline( + account=_account(), import_mode=ImportMode.YAML_CONTENT.value, yaml_content=_valid_dsl() ) - knowledge = Mock() - knowledge.indexing_technique = "high_quality" - knowledge.embedding_model_provider = None - knowledge.retrieval_model.reranking_mode = "reranking_model" - knowledge.retrieval_model.reranking_enable = False - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate", return_value=knowledge - ) - service = RagPipelineDslService(session=Mock()) + assert result.status == ImportStatus.COMPLETED + assert sqlite_session.in_transaction() + pipeline = sqlite_session.get(Pipeline, result.pipeline_id) + dataset = sqlite_session.get(Dataset, result.dataset_id) + assert pipeline is not None + assert dataset is not None + assert dataset.pipeline_id == pipeline.id + assert dataset.collection_binding_id is not None + assert sqlite_session.scalar(select(Workflow).where(Workflow.app_id == pipeline.id)) is not None + with sqlite_session_factory() as observer: + assert observer.get(Pipeline, pipeline.id) is None + assert observer.get(Dataset, dataset.id) is None - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": KNOWLEDGE_INDEX_NODE_TYPE}}]} + sqlite_session.commit() + with sqlite_session_factory() as observer: + assert observer.get(Pipeline, pipeline.id) is not None + assert observer.get(Dataset, dataset.id) is not None + + +def test_import_pipeline_id_is_tenant_scoped(service: RagPipelineDslService, sqlite_session: Session) -> None: + foreign = _pipeline(sqlite_session, tenant_id="tenant-2") + result = service.import_rag_pipeline( + account=_account(tenant_id="tenant-1"), + import_mode=ImportMode.YAML_CONTENT.value, + yaml_content=_valid_dsl(), + pipeline_id=foreign.id, + ) + assert result.status == ImportStatus.FAILED + assert result.error == "Pipeline not found" + + +def test_import_failure_leaves_rollback_to_caller( + service: RagPipelineDslService, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + yaml_content = """ +version: 0.1.0 +kind: rag_pipeline +rag_pipeline: + name: Invalid +workflow: + graph: + nodes: + - data: + type: start +""" + + result = service.import_rag_pipeline( + account=_account(), + import_mode=ImportMode.YAML_CONTENT.value, + yaml_content=yaml_content, ) - assert result == [] + assert result.status == ImportStatus.FAILED + assert "Knowledge Index node" in result.error + assert sqlite_session.in_transaction() + assert sqlite_session.scalar(select(Pipeline).where(Pipeline.name == "Invalid")) is not None + with sqlite_session_factory() as observer: + assert observer.scalar(select(Pipeline).where(Pipeline.name == "Invalid")) is None + + sqlite_session.rollback() + assert sqlite_session.scalar(select(Pipeline).where(Pipeline.name == "Invalid")) is None -def test_extract_dependencies_from_workflow_graph_multiple_reranking_without_model(mocker: MockerFixture) -> None: - retrieval = Mock() - retrieval.retrieval_mode = "multiple" - retrieval.multiple_retrieval_config.reranking_mode = "reranking_model" - retrieval.multiple_retrieval_config.reranking_model = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeRetrievalNodeData.model_validate", - return_value=retrieval, - ) - service = RagPipelineDslService(session=Mock()) +def test_import_rejects_chunk_structure_change_for_published_pipeline( + monkeypatch: pytest.MonkeyPatch, + service: RagPipelineDslService, + sqlite_session: Session, +) -> None: + pipeline = _pipeline(sqlite_session, published=True) + _dataset(sqlite_session, pipeline, chunk_structure="hierarchical_model") + _workflow(sqlite_session, pipeline) + monkeypatch.setattr(module.KnowledgeConfiguration, "model_validate", Mock(return_value=_knowledge_configuration())) - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL}}]} + result = service.import_rag_pipeline( + account=_account(), + import_mode=ImportMode.YAML_CONTENT.value, + yaml_content=_valid_dsl(), + pipeline_id=pipeline.id, ) - assert result == [] + assert result.status == ImportStatus.FAILED + assert result.error == "Chunk structure is not compatible with the published pipeline" + sqlite_session.rollback() -def test_extract_dependencies_from_workflow_graph_multiple_weighted_without_weights(mocker: MockerFixture) -> None: - retrieval = Mock() - retrieval.retrieval_mode = "multiple" - retrieval.multiple_retrieval_config.reranking_mode = "weighted_score" - retrieval.multiple_retrieval_config.weights = None - mocker.patch( - "services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeRetrievalNodeData.model_validate", - return_value=retrieval, +def test_create_or_update_pipeline_flushes_caller_transaction_and_updates_existing( + service: RagPipelineDslService, sqlite_session: Session +) -> None: + data: dict[str, Any] = { + "rag_pipeline": {"name": "New", "description": "description"}, + "workflow": {"graph": {"nodes": [], "edges": []}}, + } + created = service._create_or_update_pipeline(pipeline=None, data=data, account=_account()) + assert created in sqlite_session + assert created.workflow_id is not None + assert sqlite_session.in_transaction() + sqlite_session.commit() + + updated = service._create_or_update_pipeline( + pipeline=created, + data={ + "rag_pipeline": {"name": "Updated", "description": "changed"}, + "workflow": {"graph": {"nodes": [{"id": "node"}], "edges": []}}, + }, + account=_account(), ) - service = RagPipelineDslService(session=Mock()) + assert updated.id == created.id + assert updated.name == "Updated" + workflow = sqlite_session.get(Workflow, created.workflow_id) + assert workflow is not None + assert workflow.graph_dict["nodes"] == [{"id": "node"}] - result = service._extract_dependencies_from_workflow_graph( - {"nodes": [{"data": {"type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL}}]} + +def test_create_pipeline_flush_failure_is_rolled_back_by_caller( + service: RagPipelineDslService, sqlite_session: Session, sqlite_engine: Engine +) -> None: + with _raise_on_workflow_insert(sqlite_engine), pytest.raises(RuntimeError, match="forced workflow INSERT"): + service._create_or_update_pipeline( + pipeline=None, + data={ + "rag_pipeline": {"name": "Broken"}, + "workflow": {"graph": {"nodes": [], "edges": []}}, + }, + account=_account(), + ) + sqlite_session.rollback() + assert sqlite_session.scalar(select(Pipeline)) is None + assert sqlite_session.scalar(select(Workflow)) is None + + +def test_confirm_import_updates_tenant_pipeline_and_dataset( + monkeypatch: pytest.MonkeyPatch, + service: RagPipelineDslService, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + pipeline = _pipeline(sqlite_session) + dataset = _dataset(sqlite_session, pipeline) + _workflow(sqlite_session, pipeline) + pending = RagPipelinePendingData( + tenant_id="tenant-1", + account_id="account-1", + import_mode=ImportMode.YAML_CONTENT.value, + yaml_content=_valid_dsl(name="Confirmed"), + pipeline_id=pipeline.id, ) + redis_key = "app_import_info:import-1" + monkeypatch.setattr( + module.redis_client, + "get", + Mock(side_effect=lambda key: pending.model_dump_json() if key == redis_key else None), + ) + delete = Mock() + monkeypatch.setattr(module.redis_client, "delete", delete) + monkeypatch.setattr(module.KnowledgeConfiguration, "model_validate", Mock(return_value=_knowledge_configuration())) + for foreign_account in (_account(tenant_id="tenant-2"), _account(account_id="account-2")): + assert service.confirm_import(import_id="import-1", account=foreign_account).status == ImportStatus.FAILED + delete.assert_not_called() + assert pipeline.name == "Pipeline" - assert result == [] + result = service.confirm_import(import_id="import-1", account=_account()) + assert result.status == ImportStatus.COMPLETED + assert result.pipeline_id == pipeline.id + assert result.dataset_id == dataset.id + persisted_pipeline = sqlite_session.get(Pipeline, pipeline.id) + assert persisted_pipeline is not None + assert persisted_pipeline.name == "Confirmed" + assert sqlite_session.in_transaction() + with sqlite_session_factory() as observer: + observed_pipeline = observer.get(Pipeline, pipeline.id) + assert observed_pipeline is not None + assert observed_pipeline.name == "Pipeline" + + sqlite_session.commit() + with sqlite_session_factory() as observer: + observed_pipeline = observer.get(Pipeline, pipeline.id) + assert observed_pipeline is not None + assert observed_pipeline.name == "Confirmed" + + delete.assert_called_once_with(redis_key) + + +def test_export_reads_real_dataset_and_workflow_and_filters_credentials( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService, sqlite_session: Session +) -> None: + pipeline = _pipeline(sqlite_session) + with pytest.raises(ValueError, match="Missing dataset"): + service.export_rag_pipeline_dsl(pipeline) + _dataset(sqlite_session, pipeline) + with pytest.raises(ValueError, match="Missing draft workflow"): + service.export_rag_pipeline_dsl(pipeline) + _workflow( + sqlite_session, + pipeline, + graph={ + "nodes": [ + { + "data": { + "type": BuiltinNodeTypes.TOOL, + "credential_id": "secret", + "provider_id": "provider", + } + }, + { + "data": { + "type": BuiltinNodeTypes.KNOWLEDGE_RETRIEVAL, + "dataset_ids": ["dataset-1"], + } + }, + { + "data": { + "type": BuiltinNodeTypes.AGENT, + "agent_parameters": { + "tools": {"value": [{"credential_id": "agent-secret"}]}, + }, + } + }, + ], + "edges": [], + }, + ) + monkeypatch.setattr(module.DependenciesAnalysisService, "generate_dependencies", Mock(return_value=[])) + exported = yaml.safe_load(service.export_rag_pipeline_dsl(pipeline)) + assert exported["kind"] == "rag_pipeline" + nodes = exported["workflow"]["graph"]["nodes"] + assert "credential_id" not in nodes[0]["data"] + assert service.decrypt_dataset_id(nodes[1]["data"]["dataset_ids"][0], pipeline.tenant_id) == "dataset-1" + assert "credential_id" not in nodes[2]["data"]["agent_parameters"]["tools"]["value"][0] + + +def test_export_preserves_tool_and_agent_credentials_when_requested( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService, sqlite_session: Session +) -> None: + pipeline = _pipeline(sqlite_session) + _dataset(sqlite_session, pipeline) + _workflow( + sqlite_session, + pipeline, + graph={ + "nodes": [ + { + "data": { + "type": BuiltinNodeTypes.TOOL, + "credential_id": "tool-secret", + "provider_id": "provider", + } + }, + { + "data": { + "type": BuiltinNodeTypes.AGENT, + "agent_parameters": { + "tools": {"value": [{"credential_id": "agent-secret"}]}, + }, + } + }, + ], + "edges": [], + }, + ) + monkeypatch.setattr(module.DependenciesAnalysisService, "generate_dependencies", Mock(return_value=[])) + + exported = yaml.safe_load(service.export_rag_pipeline_dsl(pipeline, include_secret=True)) + + nodes = exported["workflow"]["graph"]["nodes"] + assert nodes[0]["data"]["credential_id"] == "tool-secret" + assert nodes[1]["data"]["agent_parameters"]["tools"]["value"][0]["credential_id"] == "agent-secret" + + +def test_create_dataset_name_is_tenant_scoped_and_ignores_template_rows( + monkeypatch: pytest.MonkeyPatch, service: RagPipelineDslService, sqlite_session: Session +) -> None: + foreign_pipeline = _pipeline(sqlite_session, tenant_id="tenant-2") + _dataset(sqlite_session, foreign_pipeline, name="Shared") + template = PipelineCustomizedTemplate( + tenant_id="tenant-1", + name="Shared", + description="template", + chunk_structure="text_model", + icon={}, + position=1, + yaml_content=_valid_dsl(), + install_count=0, + language="en-US", + created_by="account-1", + ) + sqlite_session.add(template) + sqlite_session.commit() + imported = Mock( + id="import-1", + dataset_id="dataset-1", + pipeline_id="pipeline-1", + status=ImportStatus.COMPLETED, + imported_dsl_version="0.1.0", + current_dsl_version="0.1.0", + error="", + ) + monkeypatch.setattr(service, "import_rag_pipeline", Mock(return_value=imported)) + monkeypatch.setattr(module, "current_user", _account()) + result = service.create_rag_pipeline_dataset( + "tenant-1", + RagPipelineDatasetCreateEntity( + name="Shared", + description="description", + yaml_content=_valid_dsl(), + icon_info=IconInfo(icon="📙"), + permission="only_me", + ), + ) + assert result["dataset_id"] == "dataset-1" + + local_pipeline = _pipeline(sqlite_session, tenant_id="tenant-1") + _dataset(sqlite_session, local_pipeline, name="Local") + with pytest.raises(ValueError, match="already exists"): + service.create_rag_pipeline_dataset( + "tenant-1", + RagPipelineDatasetCreateEntity( + name="Local", + description="description", + yaml_content=_valid_dsl(), + icon_info=IconInfo(icon="📙"), + permission="only_me", + ), + ) diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py index 32908e4cb09..085659aec4a 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_service.py @@ -3,20 +3,34 @@ import time from dataclasses import dataclass from datetime import datetime from types import SimpleNamespace -from unittest.mock import Mock import pytest from pytest_mock import MockerFixture +from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import InvokeFrom -from graphon.enums import WorkflowNodeExecutionStatus +from core.rag.index_processor.constant.index_type import IndexStructureType +from graphon.enums import ( + BuiltinNodeTypes, + ErrorStrategy, + WorkflowNodeExecutionMetadataKey, + WorkflowNodeExecutionStatus, +) from graphon.graph_events import NodeRunFailedEvent from graphon.node_events.base import NodeRunResult from models import Account, Tenant -from models.dataset import Dataset, Pipeline, PipelineCustomizedTemplate, PipelineRecommendedPlugin +from models.dataset import ( + Dataset, + Document, + DocumentPipelineExecutionLog, + Pipeline, + PipelineCustomizedTemplate, + PipelineRecommendedPlugin, +) +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus from models.workflow import Workflow -from services.dataset_ref_service import DatasetRefService from services.entities.knowledge_entities.rag_pipeline_entities import IconInfo, PipelineTemplateInfoEntity +from services.errors.rag_pipeline import RagPipelineResourceNotFoundError from services.rag_pipeline.rag_pipeline import RagPipelineService from services.workflow_ref_service import WorkflowRef @@ -24,26 +38,16 @@ from services.workflow_ref_service import WorkflowRef @dataclass class RagPipelineServiceTestContext: service: RagPipelineService - session: Mock - session_maker: Mock - - -def _make_mock_session_maker(mocker: MockerFixture, session: Mock) -> Mock: - session_context = mocker.MagicMock() - session_context.__enter__.return_value = session - session_context.__exit__.return_value = None - - transaction_context = mocker.MagicMock() - transaction_context.__enter__.return_value = session - transaction_context.__exit__.return_value = None - - session_maker = mocker.Mock(return_value=session_context) - session_maker.begin.return_value = transaction_context - return session_maker + session: Session + session_maker: sessionmaker[Session] @pytest.fixture -def rag_pipeline_service(mocker: MockerFixture) -> RagPipelineServiceTestContext: +def rag_pipeline_service( + mocker: MockerFixture, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> RagPipelineServiceTestContext: mocker.patch( "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository", return_value=MockRepo(), @@ -52,11 +56,13 @@ def rag_pipeline_service(mocker: MockerFixture) -> RagPipelineServiceTestContext "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_run_repository", return_value=MockRepo(), ) - session = mocker.Mock() - session_maker = _make_mock_session_maker(mocker, session) - mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=mocker.Mock())) - service = RagPipelineService(session=session, session_maker=session_maker) - return RagPipelineServiceTestContext(service=service, session=session, session_maker=session_maker) + mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=sqlite_session.get_bind())) + service = RagPipelineService(session=sqlite_session, session_maker=sqlite_session_factory) + return RagPipelineServiceTestContext( + service=service, + session=sqlite_session, + session_maker=sqlite_session_factory, + ) class MockRepo: @@ -79,13 +85,21 @@ def _make_pipeline( workflow_id: str | None = None, is_published: bool = False, ) -> Pipeline: - pipeline = Pipeline(tenant_id=tenant_id, name="Test Pipeline", description="test") + pipeline = Pipeline( + tenant_id=tenant_id, + name="Test Pipeline", + description="test", + workflow_id=workflow_id, + is_published=is_published, + ) pipeline.id = pipeline_id - pipeline.workflow_id = workflow_id - pipeline.is_published = is_published return pipeline +def _make_template_args(name: str = "New Template") -> dict[str, object]: + return {"name": name, "description": "Desc", "icon_info": {"icon": "star"}} + + def _make_workflow( *, workflow_id: str = "wf-1", @@ -94,13 +108,14 @@ def _make_workflow( graph: dict[str, object] | None = None, features: dict[str, object] | None = None, created_by: str = "u1", + version: str = Workflow.VERSION_DRAFT, ) -> Workflow: workflow = Workflow( id=workflow_id, tenant_id=tenant_id, app_id=app_id, type="workflow", - version="draft", + version=version, marked_name="", marked_comment="", graph=json.dumps(graph or {"nodes": []}), @@ -122,11 +137,32 @@ def _make_dataset(*, dataset_id: str = "d1", pipeline_id: str = "p1", tenant_id: tenant_id=tenant_id, name="Test Dataset", created_by="u1", + pipeline_id=pipeline_id, ) - dataset.pipeline_id = pipeline_id return dataset +def _make_document(*, document_id: str = "doc-1", dataset_id: str = "d1", tenant_id: str = "t1") -> Document: + return Document( + id=document_id, + tenant_id=tenant_id, + dataset_id=dataset_id, + position=1, + data_source_type=DataSourceType.LOCAL_FILE, + batch="batch", + name="Document", + created_from=DocumentCreatedFrom.API, + created_by="u1", + indexing_status=IndexingStatus.WAITING, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + + +def _persist(session: Session, *entities: object) -> None: + session.add_all(entities) + session.commit() + + def _make_failed_published_node_run() -> tuple[SimpleNamespace, NodeRunFailedEvent]: class FakeVariablePool: def __init__(self) -> None: @@ -185,9 +221,11 @@ def _make_recommended_plugin(plugin_id: str) -> PipelineRecommendedPlugin: return PipelineRecommendedPlugin(plugin_id=plugin_id, provider_name=plugin_id, type="tool", position=0, active=True) -def test_get_pipeline_templates_fallbacks_to_builtin_for_non_english_empty_result(mocker: MockerFixture) -> None: +def test_get_pipeline_templates_fallbacks_to_builtin_for_non_english_empty_result( + mocker: MockerFixture, sqlite_session: Session +) -> None: mocker.patch("services.rag_pipeline.rag_pipeline.dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE", "remote") - session = mocker.Mock() + session = sqlite_session remote_retrieval = mocker.Mock() remote_retrieval.get_pipeline_templates.return_value = {"pipeline_templates": []} @@ -206,8 +244,10 @@ def test_get_pipeline_templates_fallbacks_to_builtin_for_non_english_empty_resul builtin_retrieval.fetch_pipeline_templates_from_builtin.assert_called_once_with("en-US") -def test_get_pipeline_templates_customized_mode_uses_customized_factory(mocker: MockerFixture) -> None: - session = mocker.Mock() +def test_get_pipeline_templates_customized_mode_uses_customized_factory( + mocker: MockerFixture, sqlite_session: Session +) -> None: + session = sqlite_session retrieval = mocker.Mock() retrieval.get_pipeline_templates.return_value = {"pipeline_templates": [{"id": "custom-1"}]} @@ -222,21 +262,23 @@ def test_get_pipeline_templates_customized_mode_uses_customized_factory(mocker: @pytest.mark.parametrize("template_type", ["built-in", "customized"]) -def test_get_pipeline_template_detail_uses_expected_mode(mocker: MockerFixture, template_type: str) -> None: +def test_get_pipeline_template_detail_uses_expected_mode( + mocker: MockerFixture, template_type: str, sqlite_session: Session +) -> None: mocker.patch("services.rag_pipeline.rag_pipeline.dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE", "remote") - session = mocker.Mock() + session = sqlite_session retrieval = mocker.Mock() retrieval.get_pipeline_template_detail.return_value = {"id": "tpl-1"} factory_mock = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory_mock.get_pipeline_template_factory.return_value.return_value = retrieval - result = RagPipelineService.get_pipeline_template_detail("tpl-1", type=template_type, session=session) + result = RagPipelineService.get_pipeline_template_detail("tpl-1", "tenant-1", type=template_type, session=session) assert result == {"id": "tpl-1"} expected_mode = "remote" if template_type == "built-in" else "customized" factory_mock.get_pipeline_template_factory.assert_called_with(expected_mode) - retrieval.get_pipeline_template_detail.assert_called_once_with("tpl-1", session=session) + retrieval.get_pipeline_template_detail.assert_called_once_with("tpl-1", "tenant-1", session=session) def test_get_published_workflow_returns_none_when_pipeline_has_no_workflow_id( @@ -253,10 +295,8 @@ def test_get_all_published_workflow_returns_empty_for_unpublished_pipeline( rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: pipeline = _make_pipeline(workflow_id=None) - session = SimpleNamespace() - workflows, has_more = rag_pipeline_service.service.get_all_published_workflow( - session=session, + session=rag_pipeline_service.session, pipeline=pipeline, page=1, limit=20, @@ -271,12 +311,22 @@ def test_get_all_published_workflow_returns_empty_for_unpublished_pipeline( def test_get_all_published_workflow_applies_limit_and_has_more( rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - scalars_result = SimpleNamespace(all=lambda: ["wf1", "wf2", "wf3"]) - session = SimpleNamespace(scalars=lambda stmt: scalars_result) pipeline = _make_pipeline(pipeline_id="pipeline-1", workflow_id="wf-live") + workflows = [ + _make_workflow( + workflow_id=f"wf-{index}", + app_id="pipeline-1", + created_by="user-1", + ) + for index in range(1, 4) + ] + for index, workflow in enumerate(workflows, start=1): + workflow.version = f"2024-01-0{index}T00:00:00" + workflow.marked_name = f"release-{index}" + _persist(rag_pipeline_service.session, *workflows) - workflows, has_more = rag_pipeline_service.service.get_all_published_workflow( - session=session, + result, has_more = rag_pipeline_service.service.get_all_published_workflow( + session=rag_pipeline_service.session, pipeline=pipeline, page=1, limit=2, @@ -284,7 +334,7 @@ def test_get_all_published_workflow_applies_limit_and_has_more( named_only=True, ) - assert workflows == ["wf1", "wf2"] + assert [workflow.id for workflow in result] == ["wf-3", "wf-2"] assert has_more is True @@ -292,12 +342,8 @@ def test_get_all_published_workflow_applies_limit_and_has_more( def test_sync_draft_workflow_creates_new_when_none_exists( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - session = rag_pipeline_service.session - session.get.return_value = _make_pipeline(workflow_id=None) - session.scalar.return_value = None - pipeline = _make_pipeline(workflow_id=None) account = _make_account() @@ -316,14 +362,12 @@ def test_sync_draft_workflow_creates_new_when_none_exists( def test_sync_draft_workflow_raises_on_hash_mismatch( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: from services.errors.app import WorkflowHashNotEqualError existing_wf = _make_workflow(graph={"nodes": [{"id": "old"}]}) - session = rag_pipeline_service.session - session.get.return_value = _make_pipeline() - session.scalar.return_value = existing_wf + _persist(rag_pipeline_service.session, existing_wf) pipeline = _make_pipeline() account = _make_account() @@ -341,20 +385,10 @@ def test_sync_draft_workflow_raises_on_hash_mismatch( def test_sync_draft_workflow_updates_existing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - existing_wf = SimpleNamespace( - unique_hash="hash-1", - graph=None, - updated_by=None, - updated_at=None, - environment_variables=None, - conversation_variables=None, - rag_pipeline_variables=None, - ) - session = rag_pipeline_service.session - session.get.return_value = _make_pipeline() - session.scalar.return_value = existing_wf + existing_wf = _make_workflow(graph={"nodes": [{"id": "old"}]}) + _persist(rag_pipeline_service.session, existing_wf) pipeline = _make_pipeline() account = _make_account() @@ -362,16 +396,16 @@ def test_sync_draft_workflow_updates_existing( result = rag_pipeline_service.service.sync_draft_workflow( pipeline=pipeline, graph={"nodes": [{"id": "n1"}]}, - unique_hash="hash-1", + unique_hash=existing_wf.unique_hash, account=account, - environment_variables=["env1"], - conversation_variables=["conv1"], - rag_pipeline_variables=["rp1"], + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], ) assert result is existing_wf assert result.updated_by == "u1" - assert result.environment_variables == ["env1"] + assert result.graph_dict == {"nodes": [{"id": "n1"}]} # --- get_default_block_config --- @@ -384,7 +418,6 @@ def test_get_default_block_config_returns_config_for_valid_type( fake_node_class.get_default_config.return_value = {"type": "start", "config": {}} # Use a simpler approach: test with a known valid node type - from graphon.enums import BuiltinNodeTypes mocker.patch( "services.rag_pipeline.rag_pipeline.get_node_type_classes_mapping", @@ -407,17 +440,16 @@ def test_get_default_block_config_returns_none_for_unmapped_type( def test_update_workflow_updates_allowed_fields( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - workflow = SimpleNamespace( - id="wf-1", marked_name="", marked_comment="", updated_by=None, updated_at=None, disallowed="original" - ) + workflow = _make_workflow() + _persist(rag_pipeline_service.session, workflow) workflow_ref = WorkflowRef(tenant_id="t1", owner_id="pipeline-1", workflow_id="wf-1") - session = mocker.Mock() - session.scalar.return_value = workflow + workflow.app_id = "pipeline-1" + rag_pipeline_service.session.commit() result = rag_pipeline_service.service.update_workflow( - session=session, + session=rag_pipeline_service.session, account_id="u1", data={"marked_name": "v1", "marked_comment": "release", "disallowed": "hacked"}, workflow_ref=workflow_ref, @@ -425,18 +457,15 @@ def test_update_workflow_updates_allowed_fields( assert result.marked_name == "v1" assert result.marked_comment == "release" - assert result.disallowed == "original" # non-allowed field not updated + assert not hasattr(result, "disallowed") assert result.updated_by == "u1" def test_update_workflow_returns_none_when_not_found( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - session = mocker.Mock() - session.scalar.return_value = None - result = rag_pipeline_service.service.update_workflow( - session=session, + session=rag_pipeline_service.session, account_id="u1", data={"marked_name": "v1"}, workflow_ref=WorkflowRef(tenant_id="t1", owner_id="pipeline-1", workflow_id="wf-missing"), @@ -446,32 +475,22 @@ def test_update_workflow_returns_none_when_not_found( def test_update_workflow_with_ref_scopes_lookup_to_pipeline( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - workflow = SimpleNamespace( - id="wf-1", marked_name="", marked_comment="", updated_by=None, updated_at=None, disallowed="original" - ) + workflow = _make_workflow(app_id="pipeline-1") + other_owner_workflow = _make_workflow(workflow_id="wf-other", app_id="pipeline-2") + _persist(rag_pipeline_service.session, workflow, other_owner_workflow) workflow_ref = WorkflowRef(tenant_id="t1", owner_id="pipeline-1", workflow_id="wf-1") - session = mocker.Mock() - session.scalar.return_value = workflow result = rag_pipeline_service.service.update_workflow( - session=session, + session=rag_pipeline_service.session, account_id="u1", data={"marked_name": "v1"}, workflow_ref=workflow_ref, ) - stmt = session.scalar.call_args.args[0] - compiled = stmt.compile() - statement = str(compiled) - assert "workflows.id" in statement - assert "workflows.tenant_id" in statement - assert "workflows.app_id" in statement - assert "wf-1" in compiled.params.values() - assert "t1" in compiled.params.values() - assert "pipeline-1" in compiled.params.values() assert result is workflow + assert other_owner_workflow.marked_name == "" # --- get_rag_pipeline_paginate_workflow_runs --- @@ -522,19 +541,17 @@ def test_get_rag_pipeline_workflow_run_delegates( def test_is_workflow_exist_returns_true_when_draft_exists( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - rag_pipeline_service.session.scalar.return_value = 1 + _persist(rag_pipeline_service.session, _make_workflow()) pipeline = _make_pipeline() assert rag_pipeline_service.service.is_workflow_exist(pipeline) is True def test_is_workflow_exist_returns_false_when_no_draft( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - rag_pipeline_service.session.scalar.return_value = 0 - pipeline = _make_pipeline() assert rag_pipeline_service.service.is_workflow_exist(pipeline) is False @@ -543,16 +560,7 @@ def test_is_workflow_exist_returns_false_when_no_draft( def test_publish_workflow_success(mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext) -> None: - # Don't import Workflow from rag_pipeline to avoid confusion during patching - - # 1. Mock select to bypass SQLAlchemy validation - mock_select = mocker.patch("services.rag_pipeline.rag_pipeline.select") - - # 2. Setup draft workflow mock - draft_wf = mocker.Mock() - draft_wf.id = "wf-draft" - draft_wf.unique_hash = "hash-1" - draft_wf.graph = { + graph = { "nodes": [ { "data": { @@ -566,48 +574,52 @@ def test_publish_workflow_success(mocker: MockerFixture, rag_pipeline_service: R } ] } - draft_wf.environment_variables = [] - draft_wf.conversation_variables = [] - draft_wf.rag_pipeline_variables = [] - draft_wf.type = "workflow" - draft_wf.features = {} - - # 3. Setup pipeline and account - pipeline = mocker.Mock() - pipeline.id = "p1" - pipeline.tenant_id = "t1" - pipeline.workflow_id = "wf-old-published" - - account = mocker.Mock() - account.id = "u1" - - # 4. Mock Workflow class and its .new() method - mock_workflow_class = mocker.patch("services.rag_pipeline.rag_pipeline.Workflow") - new_wf = mocker.Mock() - new_wf.id = "wf-published-new" - new_wf.graph_dict = draft_wf.graph - mock_workflow_class.new.return_value = new_wf - - # 5. Mock DatasetService + draft_workflow = _make_workflow(workflow_id="wf-draft", graph=graph) + pipeline = _make_pipeline(workflow_id="wf-old-published") + dataset = _make_dataset() + _persist(rag_pipeline_service.session, draft_workflow, pipeline, dataset) + mocker.patch("services.rag_pipeline.rag_pipeline.KnowledgeConfiguration.model_validate", return_value=mocker.Mock()) mock_dataset_service_class = mocker.patch("services.dataset_service.DatasetService") - # 6. Mock session and dataset lookup - mock_session = mocker.Mock() - mock_session.scalar.return_value = draft_wf + result = rag_pipeline_service.service.publish_workflow( + session=rag_pipeline_service.session, pipeline=pipeline, account=_make_account() + ) - dataset = mocker.Mock() - dataset.retrieval_model_dict = {} - pipeline.retrieve_dataset.return_value = dataset - - # 7. Run test - result = rag_pipeline_service.service.publish_workflow(session=mock_session, pipeline=pipeline, account=account) - - # 8. Assertions - assert result == new_wf - # Note: dataset settings are updated via DatasetService now, so we can verify the call + assert result.app_id == pipeline.id + assert result.graph_dict == graph mock_dataset_service_class.update_rag_pipeline_dataset_settings.assert_called_once() +def test_publish_workflow_rejects_missing_llm_environment_reference( + rag_pipeline_service: RagPipelineServiceTestContext, +) -> None: + draft_workflow = _make_workflow( + workflow_id="wf-draft", + tenant_id="tenant", + app_id="pipeline", + graph={ + "nodes": [ + { + "id": "llm-node", + "data": { + "type": "llm", + "model": {"provider": "provider", "name": "model", "mode": "chat"}, + "model_selector": ["env", "missing_model"], + }, + } + ] + }, + ) + _persist(rag_pipeline_service.session, draft_workflow) + + with pytest.raises(ValueError, match="missing_model.*not found"): + rag_pipeline_service.service.publish_workflow( + session=rag_pipeline_service.session, + pipeline=_make_pipeline(pipeline_id="pipeline", tenant_id="tenant"), + account=_make_account(account_id="account", tenant_id="tenant"), + ) + + # --- run_datasource_workflow_node --- @@ -767,7 +779,6 @@ def test_run_datasource_node_preview_online_document( def test_handle_node_run_result_success( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - from graphon.enums import WorkflowNodeExecutionMetadataKey, WorkflowNodeExecutionStatus from graphon.graph_events import NodeRunSucceededEvent from graphon.node_events.base import NodeRunResult @@ -865,37 +876,28 @@ def test_get_second_step_parameters_success( def test_publish_customized_pipeline_template_success( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - # 1. Setup mocks pipeline = _make_pipeline(workflow_id="wf-1", is_published=True) - workflow = _make_workflow(workflow_id="wf-1") - session = rag_pipeline_service.session - session.get.side_effect = [pipeline, workflow] - session.scalar.side_effect = [None, 5] - - # Mock retrieve_dataset dataset = _make_dataset() dataset.chunk_structure = "paragraph" - pipeline.retrieve_dataset = mocker.Mock(return_value=dataset) + _persist(session, pipeline, workflow, dataset) # Mock RagPipelineDslService mock_dsl_service = mocker.Mock() - mock_dsl_service.export_rag_pipeline_dsl.return_value = {"dsl": "content"} + mock_dsl_service.export_rag_pipeline_dsl.return_value = "dsl: content" mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.RagPipelineDslService", return_value=mock_dsl_service) account = _make_account(account_id="user-123") - # 2. Run test - args = {"name": "New Template", "description": "Desc", "icon_info": {"icon": "star"}, "tags": ["tag1"]} - rag_pipeline_service.service.publish_customized_pipeline_template("p1", args, account, "t1", session=session) + rag_pipeline_service.service.publish_customized_pipeline_template( + pipeline, dataset, _make_template_args(), account, session=session + ) - # 3. Assertions - # Verify a new template was added to session or similar? - # Since we can't easily check the session inside the context manager with Mock, - # we just check that no error was raised and DSL was exported. - pipeline.retrieve_dataset.assert_called_once_with(session=session) mock_dsl_service.export_rag_pipeline_dsl.assert_called_once_with(pipeline=pipeline, include_secret=True) + templates = session.query(PipelineCustomizedTemplate).all() + assert len(templates) == 1 + assert templates[0].name == "New Template" # --- get_datasource_plugins --- @@ -909,24 +911,24 @@ def test_get_datasource_plugins_success( pipeline = _make_pipeline(workflow_id="wf-1") - workflow = mocker.Mock() - workflow.graph_dict = { - "nodes": [ - { - "id": "node-1", - "data": { - "type": "datasource", - "plugin_id": "p-1", - "provider_name": "notion", - "provider_type": "online_document", - "title": "Notion", - }, - } - ] - } - workflow.rag_pipeline_variables = [] - - rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, workflow] + workflow = _make_workflow( + workflow_id="wf-1", + graph={ + "nodes": [ + { + "id": "node-1", + "data": { + "type": "datasource", + "plugin_id": "p-1", + "provider_name": "notion", + "provider_type": "online_document", + "title": "Notion", + }, + } + ] + }, + ) + _persist(rag_pipeline_service.session, dataset, pipeline, workflow) # Mock DatasourceProviderService mock_provider_service = mocker.Mock() @@ -950,24 +952,25 @@ def test_get_datasource_plugins_success( def test_retry_error_document_success( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - from models.dataset import Document, DocumentPipelineExecutionLog, Pipeline - - # 1. Setup mocks dataset = mocker.Mock() - document = mocker.Mock(spec=Document) - document.id = "doc-1" + document = SimpleNamespace(id="doc-1") - log = mocker.Mock(spec=DocumentPipelineExecutionLog) - log.pipeline_id = "p-1" - log.datasource_info = "{}" # Ensure it's a string if it's used as JSON later + log = DocumentPipelineExecutionLog( + pipeline_id="p-1", + document_id="doc-1", + datasource_type="upload_file", + datasource_info="{}", + datasource_node_id="node-id", + input_data={}, + created_by="account-id", + ) + # Ensure it's a string if it's used as JSON later - pipeline = mocker.Mock(spec=Pipeline) + pipeline = Pipeline(tenant_id="tenant-id", name="Test Pipeline", workflow_id="wf-1") pipeline.id = "p-1" - workflow = mocker.Mock() - - rag_pipeline_service.session.scalar.side_effect = [log, workflow] - rag_pipeline_service.session.get.return_value = pipeline + workflow = _make_workflow(workflow_id="wf-1", tenant_id="tenant-id", app_id="p-1") + _persist(rag_pipeline_service.session, log, pipeline, workflow) # Mock PipelineGenerator mock_gen_instance = mocker.Mock() @@ -991,9 +994,11 @@ def test_set_datasource_variables_success( from models.dataset import Pipeline # 1. Setup mocks - pipeline = mocker.Mock(spec=Pipeline) + pipeline = Pipeline( + tenant_id="t1", + name="Test Pipeline", + ) pipeline.id = "p-1" - pipeline.tenant_id = "t1" draft_wf = mocker.Mock() draft_wf.id = "wf-1" @@ -1045,8 +1050,7 @@ def test_get_draft_workflow_success(rag_pipeline_service: RagPipelineServiceTest pipeline = _make_pipeline() workflow = _make_workflow() - - rag_pipeline_service.session.scalar.return_value = workflow + _persist(rag_pipeline_service.session, workflow) # 2. Run test result = rag_pipeline_service.service.get_draft_workflow(pipeline) @@ -1060,8 +1064,7 @@ def test_get_published_workflow_success(rag_pipeline_service: RagPipelineService pipeline = _make_pipeline(workflow_id="wf-pub") workflow = _make_workflow(workflow_id="wf-pub") - - rag_pipeline_service.session.scalar.return_value = workflow + _persist(rag_pipeline_service.session, workflow) # 2. Run test result = rag_pipeline_service.service.get_published_workflow(pipeline) @@ -1079,7 +1082,6 @@ def test_get_default_block_configs_success(rag_pipeline_service: RagPipelineServ def test_get_default_block_config_success(rag_pipeline_service: RagPipelineServiceTestContext) -> None: - from graphon.enums import BuiltinNodeTypes result = rag_pipeline_service.service.get_default_block_config(BuiltinNodeTypes.LLM) assert result is not None @@ -1087,21 +1089,20 @@ def test_get_default_block_config_success(rag_pipeline_service: RagPipelineServi def test_publish_workflow_raises_when_draft_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - session = mocker.Mock() - session.scalar.return_value = None pipeline = _make_pipeline() account = _make_account() with pytest.raises(ValueError, match="No valid workflow found"): - rag_pipeline_service.service.publish_workflow(session=session, pipeline=pipeline, account=account) + rag_pipeline_service.service.publish_workflow( + session=rag_pipeline_service.session, pipeline=pipeline, account=account + ) def test_get_default_block_config_returns_none_when_mapped_type_missing( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - from graphon.enums import BuiltinNodeTypes mocker.patch("services.rag_pipeline.rag_pipeline.get_node_type_classes_mapping", return_value={}) @@ -1111,7 +1112,6 @@ def test_get_default_block_config_returns_none_when_mapped_type_missing( def test_get_default_block_config_injects_http_request_filter( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - from graphon.enums import BuiltinNodeTypes fake_node_cls = mocker.Mock() fake_node_cls.get_default_config.return_value = {"type": "http-request"} @@ -1138,6 +1138,64 @@ def test_run_draft_workflow_node_raises_when_workflow_missing( rag_pipeline_service.service.run_draft_workflow_node(pipeline, "node-1", {}, account) +def test_run_draft_workflow_node_seeds_llm_environment_variable( + mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +) -> None: + from factories import variable_factory + + pipeline = _make_pipeline() + account = _make_account() + llm_environment_variable = variable_factory.build_environment_variable_from_mapping( + { + "id": "env-1", + "name": "for_summarize", + "value_type": "llm", + "value": { + "provider": "langgenius/openai/openai", + "name": "gpt-4o", + "mode": "chat", + }, + "description": "Shared summarization model", + } + ) + draft_workflow = mocker.Mock(id="wf-1", environment_variables=[llm_environment_variable]) + draft_workflow.get_node_config_by_id.return_value = {"id": "node-1"} + draft_workflow.get_enclosing_node_type_and_id.return_value = None + mocker.patch.object(rag_pipeline_service.service, "get_draft_workflow", return_value=draft_workflow) + + execution = SimpleNamespace(id="exec-1", node_id="node-1", node_type="llm", process_data={}, outputs={}) + handle_node_run_result = mocker.patch.object( + rag_pipeline_service.service, "_handle_node_run_result", return_value=execution + ) + single_step_run = mocker.patch("services.rag_pipeline.rag_pipeline.WorkflowEntry.single_step_run") + + repo = mocker.Mock() + mocker.patch( + "services.rag_pipeline.rag_pipeline.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", + return_value=repo, + ) + rag_pipeline_service.service._node_execution_service_repo = mocker.Mock( + get_execution_by_id=mocker.Mock(return_value="db") + ) + mocker.patch("services.rag_pipeline.rag_pipeline.DraftVariableSaver", return_value=mocker.Mock()) + session_ctx = mocker.MagicMock() + session_ctx.begin.return_value = mocker.MagicMock() + mocker.patch("services.rag_pipeline.rag_pipeline.Session", return_value=session_ctx) + + rag_pipeline_service.service.run_draft_workflow_node(pipeline, "node-1", {}, account) + + getter = handle_node_run_result.call_args.kwargs["getter"] + getter() + variable_pool = single_step_run.call_args.kwargs["variable_pool"] + llm_model = variable_pool.get(["env", "for_summarize"]) + assert llm_model is not None + assert llm_model.value == { + "provider": "langgenius/openai/openai", + "name": "gpt-4o", + "mode": "chat", + } + + def test_run_draft_workflow_node_saves_execution_and_variables( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: @@ -1296,7 +1354,6 @@ def test_handle_node_run_result_default_value_strategy( ) -> None: from datetime import datetime - from graphon.enums import BuiltinNodeTypes, ErrorStrategy, WorkflowNodeExecutionStatus from graphon.graph_events import NodeRunFailedEvent from graphon.node_events.base import NodeRunResult @@ -1421,10 +1478,8 @@ def test_get_rag_pipeline_workflow_run_node_executions_assembles_configured_repo def test_get_recommended_plugins_returns_empty_when_no_active_plugins( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - rag_pipeline_service.session.scalars.return_value.all.return_value = [] - result = rag_pipeline_service.service.get_recommended_plugins("all", _make_account(), "t1") assert result == { @@ -1438,7 +1493,7 @@ def test_get_recommended_plugins_returns_installed_and_uninstalled( ) -> None: plugin_a = _make_recommended_plugin("plugin-a") plugin_b = _make_recommended_plugin("plugin-b") - rag_pipeline_service.session.scalars.return_value.all.return_value = [plugin_a, plugin_b] + _persist(rag_pipeline_service.session, plugin_a, plugin_b) mocker.patch( "services.rag_pipeline.rag_pipeline.BuiltinToolManageService.list_builtin_tools", return_value=[SimpleNamespace(plugin_id="plugin-a", to_dict=lambda: {"plugin_id": "plugin-a"})], @@ -1448,7 +1503,7 @@ def test_get_recommended_plugins_returns_installed_and_uninstalled( return_value=[{"plugin_id": "plugin-b", "name": "Plugin B"}], ) - result = rag_pipeline_service.service.get_recommended_plugins("custom", _make_account(), "t1") + result = rag_pipeline_service.service.get_recommended_plugins("tool", _make_account(), "t1") assert result["installed_recommended_plugins"] == [{"plugin_id": "plugin-a"}] assert result["uninstalled_recommended_plugins"] == [{"plugin_id": "plugin-b", "name": "Plugin B"}] @@ -1482,7 +1537,6 @@ def test_set_datasource_variables_raises_when_node_id_missing( def test_get_default_block_configs_skips_empty_configs( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - from graphon.enums import BuiltinNodeTypes http_node = mocker.Mock() http_node.get_default_config.return_value = {"type": "http-request"} @@ -1677,10 +1731,8 @@ def test_get_second_step_parameters_filters_first_step_variables( def test_retry_error_document_raises_when_execution_log_not_found( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - rag_pipeline_service.session.scalar.return_value = None - with pytest.raises(ValueError, match="Document pipeline execution log not found"): rag_pipeline_service.service.retry_error_document( SimpleNamespace(), SimpleNamespace(id="doc-1"), SimpleNamespace(id="u1") @@ -1688,11 +1740,11 @@ def test_retry_error_document_raises_when_execution_log_not_found( def test_get_datasource_plugins_raises_when_workflow_not_found( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - dataset = SimpleNamespace(pipeline_id="p1") - pipeline = SimpleNamespace(id="p1", tenant_id="t1", workflow_id="wf-1") - rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, None] + dataset = _make_dataset() + pipeline = _make_pipeline(workflow_id="wf-1") + _persist(rag_pipeline_service.session, dataset, pipeline) with pytest.raises(ValueError, match="Pipeline or workflow not found"): rag_pipeline_service.service.get_datasource_plugins("t1", "d1", True) @@ -1722,14 +1774,16 @@ def test_handle_node_run_result_raises_when_no_terminal_event( def test_handle_node_run_result_marks_document_error_for_published_invoke( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: node_instance, event = _make_failed_published_node_run() - document = SimpleNamespace(indexing_status="waiting", error=None) - rag_pipeline_service.session.scalar.return_value = _make_dataset( - dataset_id="dataset-1", pipeline_id="pipeline-1", tenant_id="t1" + dataset = _make_dataset(dataset_id="dataset-1", pipeline_id="pipeline-1", tenant_id="t1") + document = _make_document(document_id="doc-1", dataset_id="dataset-1", tenant_id="t1") + _persist( + rag_pipeline_service.session, + dataset, + document, ) - get_document_by_ref = mocker.patch.object(DatasetRefService, "get_document_by_ref", return_value=document) result = rag_pipeline_service.service._handle_node_run_result( getter=lambda: (node_instance, iter([event])), @@ -1739,32 +1793,17 @@ def test_handle_node_run_result_marks_document_error_for_published_invoke( ) assert result.status == WorkflowNodeExecutionStatus.FAILED - stmt = rag_pipeline_service.session.scalar.call_args.args[0] - compiled = stmt.compile() - statement = str(compiled) - assert "datasets.id" in statement - assert "datasets.tenant_id" in statement - assert "datasets.pipeline_id" in statement - assert "t1" in compiled.params.values() - assert "dataset-1" in compiled.params.values() - assert "pipeline-1" in compiled.params.values() - document_ref = get_document_by_ref.call_args.args[0] - assert document_ref.dataset.tenant_id == "t1" - assert document_ref.dataset.dataset_id == "dataset-1" - assert document_ref.document_id == "doc-1" - get_document_by_ref.assert_called_once_with(document_ref, session=rag_pipeline_service.session) - assert document.indexing_status == "error" - assert document.error == "boom" - rag_pipeline_service.session.add.assert_called_once_with(document) - rag_pipeline_service.session.commit.assert_called_once_with() + rag_pipeline_service.session.expire_all() + updated_document = rag_pipeline_service.session.get(Document, "doc-1") + assert updated_document is not None + assert updated_document.indexing_status == "error" + assert updated_document.error == "boom" def test_handle_node_run_result_does_not_write_when_pipeline_dataset_is_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: node_instance, event = _make_failed_published_node_run() - rag_pipeline_service.session.scalar.return_value = None - get_document_by_ref = mocker.patch.object(DatasetRefService, "get_document_by_ref") result = rag_pipeline_service.service._handle_node_run_result( getter=lambda: (node_instance, iter([event])), @@ -1774,9 +1813,7 @@ def test_handle_node_run_result_does_not_write_when_pipeline_dataset_is_missing( ) assert result.status == WorkflowNodeExecutionStatus.FAILED - get_document_by_ref.assert_not_called() - rag_pipeline_service.session.add.assert_not_called() - rag_pipeline_service.session.commit.assert_not_called() + assert rag_pipeline_service.session.query(Document).count() == 0 def test_run_datasource_node_preview_raises_for_unsupported_provider( @@ -1817,54 +1854,39 @@ def test_run_datasource_node_preview_raises_for_unsupported_provider( ) -def test_publish_customized_pipeline_template_raises_for_missing_pipeline( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext -) -> None: - session = mocker.Mock() - session.get.return_value = None - - with pytest.raises(ValueError, match="Pipeline not found"): - rag_pipeline_service.service.publish_customized_pipeline_template( - "p1", {}, _make_account(), "t1", session=session - ) - - def test_publish_customized_pipeline_template_raises_for_missing_workflow_id( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: pipeline = _make_pipeline(workflow_id=None) - session = mocker.Mock() - session.get.return_value = pipeline + _persist(rag_pipeline_service.session, pipeline) - with pytest.raises(ValueError, match="Pipeline workflow not found"): + with pytest.raises(RagPipelineResourceNotFoundError, match="Pipeline workflow not found"): rag_pipeline_service.service.publish_customized_pipeline_template( - "p1", {"name": "template-name"}, _make_account(), "t1", session=session + pipeline, _make_dataset(), _make_template_args(), _make_account(), session=rag_pipeline_service.session ) def test_get_pipeline_raises_when_dataset_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - rag_pipeline_service.session.scalar.return_value = None - with pytest.raises(ValueError, match="Dataset not found"): rag_pipeline_service.service.get_pipeline("t1", "d1") def test_get_pipeline_raises_when_pipeline_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - dataset = SimpleNamespace(pipeline_id="p1") - rag_pipeline_service.session.scalar.side_effect = [dataset, None] + _persist(rag_pipeline_service.session, _make_dataset()) with pytest.raises(ValueError, match="Pipeline not found"): rag_pipeline_service.service.get_pipeline("t1", "d1") -def test_init_uses_default_sessionmaker_when_none(mocker: MockerFixture) -> None: - default_session_maker = mocker.Mock() - mocker.patch("services.rag_pipeline.rag_pipeline.sessionmaker", return_value=default_session_maker) - mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=mocker.Mock(), session=mocker.Mock())) +def test_init_uses_default_sessionmaker_when_none( + mocker: MockerFixture, sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + mocker.patch("services.rag_pipeline.rag_pipeline.sessionmaker", return_value=sqlite_session_factory) + mocker.patch("services.rag_pipeline.rag_pipeline.db", SimpleNamespace(engine=sqlite_session.get_bind())) create_exec_repo = mocker.patch( "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_node_execution_repository" ) @@ -1872,22 +1894,20 @@ def test_init_uses_default_sessionmaker_when_none(mocker: MockerFixture) -> None "services.rag_pipeline.rag_pipeline.DifyAPIRepositoryFactory.create_api_workflow_run_repository" ) - RagPipelineService(session=mocker.Mock(), session_maker=None) + RagPipelineService(session=sqlite_session, session_maker=None) - create_exec_repo.assert_called_once_with(default_session_maker) - create_run_repo.assert_called_once_with(default_session_maker) + create_exec_repo.assert_called_once_with(sqlite_session_factory) + create_run_repo.assert_called_once_with(sqlite_session_factory) -def test_get_pipeline_templates_builtin_en_us_no_fallback(mocker: MockerFixture) -> None: +def test_get_pipeline_templates_builtin_en_us_no_fallback(mocker: MockerFixture, sqlite_session: Session) -> None: mocker.patch("services.rag_pipeline.rag_pipeline.dify_config.HOSTED_FETCH_PIPELINE_TEMPLATES_MODE", "remote") - session = mocker.Mock() + session = sqlite_session retrieval = mocker.Mock() retrieval.get_pipeline_templates.return_value = {"pipeline_templates": []} factory = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory") factory.get_pipeline_template_factory.return_value.return_value = retrieval builtin = factory.get_built_in_pipeline_template_retrieval.return_value - session = mocker.Mock() - result = RagPipelineService.get_pipeline_templates(type="built-in", language="en-US", session=session) assert result == {"pipeline_templates": []} @@ -1895,28 +1915,33 @@ def test_get_pipeline_templates_builtin_en_us_no_fallback(mocker: MockerFixture) builtin.fetch_pipeline_templates_from_builtin.assert_not_called() -def test_update_customized_pipeline_template_commits_when_name_empty(mocker: MockerFixture) -> None: +def test_update_customized_pipeline_template_commits_when_name_empty(sqlite_session: Session) -> None: template = _make_customized_template() - session = mocker.Mock() - session.scalar.return_value = template + template.id = "tpl-1" + _persist(sqlite_session, template) info = PipelineTemplateInfoEntity(name="", description="updated", icon_info=IconInfo(icon="i")) result = RagPipelineService.update_customized_pipeline_template( - "tpl-1", info, _make_account(), "t1", session=session + "tpl-1", info, _make_account(), "t1", session=sqlite_session ) assert result.description == "updated" - session.commit.assert_called_once() + sqlite_session.expire_all() + updated_template = sqlite_session.get(PipelineCustomizedTemplate, "tpl-1") + assert updated_template is not None + assert updated_template.description == "updated" def test_get_all_published_workflow_without_filters_has_no_more( rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - session = SimpleNamespace(scalars=lambda stmt: SimpleNamespace(all=lambda: ["wf1"])) pipeline = _make_pipeline(workflow_id="wf-live") + workflow = _make_workflow(workflow_id="wf1") + workflow.version = "2024-01-01T00:00:00" + _persist(rag_pipeline_service.session, workflow) workflows, has_more = rag_pipeline_service.service.get_all_published_workflow( - session=session, + session=rag_pipeline_service.session, pipeline=pipeline, page=1, limit=2, @@ -1924,40 +1949,30 @@ def test_get_all_published_workflow_without_filters_has_no_more( named_only=False, ) - assert workflows == ["wf1"] + assert workflows == [workflow] assert has_more is False def test_publish_workflow_skips_dataset_update_for_non_knowledge_nodes( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - draft = SimpleNamespace( - type="workflow", - graph={"nodes": [{"data": {"type": "start"}}]}, - features={}, - environment_variables=[], - conversation_variables=[], - rag_pipeline_variables=[], - ) - session = mocker.Mock() - session.scalar.return_value = draft - published = SimpleNamespace(graph_dict={"nodes": [{"data": {"type": "start"}}]}) - mocker.patch("services.rag_pipeline.rag_pipeline.select") - mocker.patch("services.rag_pipeline.rag_pipeline.Workflow.new", return_value=published) + draft = _make_workflow(graph={"nodes": [{"data": {"type": "start"}}]}) + _persist(rag_pipeline_service.session, draft) + dataset_service = mocker.patch("services.dataset_service.DatasetService") result = rag_pipeline_service.service.publish_workflow( - session=session, - pipeline=SimpleNamespace(id="p1", tenant_id="t1", is_published=False, retrieve_dataset=lambda session: None), - account=SimpleNamespace(id="u1"), + session=rag_pipeline_service.session, + pipeline=_make_pipeline(), + account=_make_account(), ) - assert result is published + assert result.graph_dict == {"nodes": [{"data": {"type": "start"}}]} + dataset_service.update_rag_pipeline_dataset_settings.assert_not_called() def test_get_default_block_config_returns_none_when_default_empty( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - from graphon.enums import BuiltinNodeTypes node_cls = mocker.Mock() node_cls.get_default_config.return_value = None @@ -2136,39 +2151,75 @@ def test_run_free_workflow_node_delegates_to_handle_result( handle.assert_called_once() -def test_publish_customized_pipeline_template_raises_when_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext +@pytest.mark.parametrize(("workflow_tenant_id", "workflow_app_id"), [("t2", "p1"), ("t1", "p2")]) +def test_publish_customized_pipeline_template_rejects_unowned_workflow_before_export( + mocker: MockerFixture, + rag_pipeline_service: RagPipelineServiceTestContext, + workflow_tenant_id: str, + workflow_app_id: str, ) -> None: pipeline = _make_pipeline(workflow_id="wf-1") - session = mocker.Mock() - session.get.side_effect = [pipeline, None] + workflow = _make_workflow(workflow_id="wf-1", tenant_id=workflow_tenant_id, app_id=workflow_app_id) + _persist(rag_pipeline_service.session, pipeline, workflow) + dsl_service = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.RagPipelineDslService") - with pytest.raises(ValueError, match="Workflow not found"): + with pytest.raises(RagPipelineResourceNotFoundError, match="Workflow not found"): rag_pipeline_service.service.publish_customized_pipeline_template( - "p1", {}, _make_account(), "t1", session=session + pipeline, + _make_dataset(), + _make_template_args(), + _make_account(), + session=rag_pipeline_service.session, ) + dsl_service.assert_not_called() + assert rag_pipeline_service.session.query(PipelineCustomizedTemplate).count() == 0 -def test_publish_customized_pipeline_template_raises_when_dataset_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + +def test_pipeline_retrieve_dataset_rejects_unowned_dataset( + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: pipeline = _make_pipeline(workflow_id="wf-1") - workflow = _make_workflow(workflow_id="wf-1") + other_tenant_dataset = _make_dataset(tenant_id="t2") + _persist(rag_pipeline_service.session, pipeline, other_tenant_dataset) + + assert pipeline.retrieve_dataset(session=rag_pipeline_service.session) is None + + +@pytest.mark.parametrize( + ("draft_tenant_id", "draft_app_id"), + [(None, None), ("t2", "p1"), ("t1", "p2")], +) +def test_publish_customized_pipeline_template_rejects_missing_or_unowned_draft_before_side_effects( + mocker: MockerFixture, + rag_pipeline_service: RagPipelineServiceTestContext, + draft_tenant_id: str | None, + draft_app_id: str | None, +) -> None: session = rag_pipeline_service.session - session.get.side_effect = [pipeline, workflow] - pipeline.retrieve_dataset = mocker.Mock(return_value=None) + pipeline = _make_pipeline(workflow_id="wf-published") + published_workflow = _make_workflow(workflow_id="wf-published", version="published") + dataset = _make_dataset() + resources = [pipeline, published_workflow, dataset] + if draft_tenant_id and draft_app_id: + resources.append(_make_workflow(workflow_id="wf-draft", tenant_id=draft_tenant_id, app_id=draft_app_id)) + _persist(session, *resources) + dsl_service = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.RagPipelineDslService") - with pytest.raises(ValueError, match="Dataset not found"): + with pytest.raises(RagPipelineResourceNotFoundError, match="Draft workflow not found"): rag_pipeline_service.service.publish_customized_pipeline_template( - "p1", {}, _make_account(), "t1", session=session + pipeline, dataset, _make_template_args(), _make_account(), session=session ) + dsl_service.assert_not_called() + assert session.query(PipelineCustomizedTemplate).count() == 0 + def test_get_recommended_plugins_skips_manifest_when_missing( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: plugin = _make_recommended_plugin("plugin-a") - rag_pipeline_service.session.scalars.return_value.all.return_value = [plugin] + _persist(rag_pipeline_service.session, plugin) mocker.patch("services.rag_pipeline.rag_pipeline.BuiltinToolManageService.list_builtin_tools", return_value=[]) mocker.patch("services.rag_pipeline.rag_pipeline.marketplace.batch_fetch_plugin_by_ids", return_value=[]) @@ -2179,11 +2230,18 @@ def test_get_recommended_plugins_skips_manifest_when_missing( def test_retry_error_document_raises_when_pipeline_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - exec_log = SimpleNamespace(pipeline_id="p1") - rag_pipeline_service.session.scalar.return_value = exec_log - rag_pipeline_service.session.get.return_value = None + exec_log = DocumentPipelineExecutionLog( + pipeline_id="p1", + document_id="doc-1", + datasource_type="local_file", + datasource_info="{}", + datasource_node_id="node-1", + input_data={}, + created_by="u1", + ) + _persist(rag_pipeline_service.session, exec_log) with pytest.raises(ValueError, match="Pipeline not found"): rag_pipeline_service.service.retry_error_document( @@ -2192,12 +2250,19 @@ def test_retry_error_document_raises_when_pipeline_missing( def test_retry_error_document_raises_when_workflow_missing( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: - exec_log = SimpleNamespace(pipeline_id="p1") - pipeline = SimpleNamespace(id="p1", tenant_id="t1", workflow_id="wf-1") - rag_pipeline_service.session.scalar.side_effect = [exec_log, None] - rag_pipeline_service.session.get.return_value = pipeline + exec_log = DocumentPipelineExecutionLog( + pipeline_id="p1", + document_id="doc-1", + datasource_type="local_file", + datasource_info="{}", + datasource_node_id="node-1", + input_data={}, + created_by="u1", + ) + pipeline = _make_pipeline(workflow_id="wf-1") + _persist(rag_pipeline_service.session, exec_log, pipeline) with pytest.raises(ValueError, match="Workflow not found"): rag_pipeline_service.service.retry_error_document( @@ -2206,14 +2271,12 @@ def test_retry_error_document_raises_when_workflow_missing( def test_get_datasource_plugins_returns_empty_for_non_datasource_nodes( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: dataset = _make_dataset() pipeline = _make_pipeline(workflow_id="wf-1") - workflow = SimpleNamespace( - graph_dict={"nodes": [{"id": "n1", "data": {"type": "start"}}]}, rag_pipeline_variables=[] - ) - rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, workflow] + workflow = _make_workflow(graph={"nodes": [{"id": "n1", "data": {"type": "start"}}]}) + _persist(rag_pipeline_service.session, dataset, pipeline, workflow) assert rag_pipeline_service.service.get_datasource_plugins("t1", "d1", True) == [] @@ -2221,27 +2284,14 @@ def test_get_datasource_plugins_returns_empty_for_non_datasource_nodes( def test_publish_workflow_raises_when_knowledge_index_dataset_missing( mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext ) -> None: - draft = SimpleNamespace( - type="workflow", - graph={"nodes": [{"data": {"type": "knowledge-index"}}]}, - features={}, - environment_variables=[], - conversation_variables=[], - rag_pipeline_variables=[], - ) - session = mocker.Mock() - session.scalar.return_value = draft - mocker.patch("services.rag_pipeline.rag_pipeline.select") - mocker.patch( - "services.rag_pipeline.rag_pipeline.Workflow.new", - return_value=SimpleNamespace(graph_dict={"nodes": [{"data": {"type": "knowledge-index"}}]}), - ) + draft = _make_workflow(graph={"nodes": [{"data": {"type": "knowledge-index"}}]}) + _persist(rag_pipeline_service.session, draft) mocker.patch("services.rag_pipeline.rag_pipeline.KnowledgeConfiguration.model_validate", return_value=mocker.Mock()) - pipeline = SimpleNamespace(id="p1", tenant_id="t1", is_published=False, retrieve_dataset=lambda session: None) + pipeline = _make_pipeline() with pytest.raises(ValueError, match="Dataset not found"): rag_pipeline_service.service.publish_workflow( - session=session, pipeline=pipeline, account=SimpleNamespace(id="u1") + session=rag_pipeline_service.session, pipeline=pipeline, account=_make_account() ) @@ -2386,11 +2436,13 @@ def test_get_datasource_plugins_handles_empty_datasource_data_and_non_published( ) -> None: dataset = _make_dataset() pipeline = _make_pipeline() - workflow = SimpleNamespace( - graph_dict={"nodes": [{"id": "n1", "data": {"type": "datasource", "datasource_parameters": {}}}]}, - rag_pipeline_variables=[{"variable": "v1", "belong_to_node_id": "shared"}], + workflow = _make_workflow( + graph={"nodes": [{"id": "n1", "data": {"type": "datasource", "datasource_parameters": {}}}]}, ) - rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, workflow] + workflow.rag_pipeline_variables = [ + {"variable": "v1", "belong_to_node_id": "shared", "type": "text-input", "label": "V1"} + ] + _persist(rag_pipeline_service.session, dataset, pipeline, workflow) mocker.patch( "services.rag_pipeline.rag_pipeline.DatasourceProviderService.list_datasource_credentials", return_value=[] ) @@ -2405,8 +2457,8 @@ def test_get_datasource_plugins_extracts_user_inputs_and_credentials( ) -> None: dataset = _make_dataset() pipeline = _make_pipeline(workflow_id="wf-1") - workflow = SimpleNamespace( - graph_dict={ + workflow = _make_workflow( + graph={ "nodes": [ { "id": "n1", @@ -2424,13 +2476,13 @@ def test_get_datasource_plugins_extracts_user_inputs_and_credentials( } ] }, - rag_pipeline_variables=[ - {"variable": "v1", "belong_to_node_id": "shared"}, - {"variable": "v2", "belong_to_node_id": "shared"}, - {"variable": "v3", "belong_to_node_id": "shared"}, - ], ) - rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline, workflow] + workflow.rag_pipeline_variables = [ + {"variable": "v1", "belong_to_node_id": "shared", "type": "text-input", "label": "V1"}, + {"variable": "v2", "belong_to_node_id": "shared", "type": "text-input", "label": "V2"}, + {"variable": "v3", "belong_to_node_id": "shared", "type": "text-input", "label": "V3"}, + ] + _persist(rag_pipeline_service.session, dataset, pipeline, workflow) mocker.patch( "services.rag_pipeline.rag_pipeline.DatasourceProviderService.list_datasource_credentials", return_value=[{"id": "c1", "name": "Cred", "type": "api", "is_default": True}], @@ -2444,25 +2496,23 @@ def test_get_datasource_plugins_extracts_user_inputs_and_credentials( def test_get_pipeline_returns_pipeline_when_found( - mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext + rag_pipeline_service: RagPipelineServiceTestContext, ) -> None: dataset = _make_dataset() pipeline = _make_pipeline() - rag_pipeline_service.session.scalar.side_effect = [dataset, pipeline] + _persist(rag_pipeline_service.session, dataset, pipeline) result = rag_pipeline_service.service.get_pipeline("t1", "d1") assert result is pipeline -def test_get_pipeline_by_id_uses_provided_session() -> None: +def test_get_pipeline_by_id_uses_provided_session(sqlite_session: Session) -> None: pipeline = _make_pipeline() - session = Mock() - session.scalar.return_value = pipeline + sqlite_session.add(pipeline) + sqlite_session.flush() - result = RagPipelineService.get_pipeline_by_id("p1", "t1", session=session) + result = RagPipelineService.get_pipeline_by_id("p1", "t1", session=sqlite_session) assert result is pipeline - statement = session.scalar.call_args.args[0] - where_clauses = statement.whereclause.clauses - assert [clause.right.value for clause in where_clauses] == ["p1", "t1"] + assert RagPipelineService.get_pipeline_by_id("p1", "other-tenant", session=sqlite_session) is None diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_task_proxy.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_task_proxy.py index 281de58abad..e9f5bcab1d2 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_task_proxy.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_task_proxy.py @@ -4,6 +4,7 @@ from unittest.mock import Mock import pytest from pytest_mock import MockerFixture +from enums import CloudPlan from services.rag_pipeline.rag_pipeline_task_proxy import RagPipelineTaskProxy @@ -52,8 +53,6 @@ def test_dispatch_billing_sandbox_uses_default_tenant_queue(mocker: MockerFixtur upload_mock = mocker.patch.object(proxy, "_upload_invoke_entities", return_value="file-1") send_mock = mocker.patch.object(proxy, "_send_to_default_tenant_queue") - from enums.cloud_plan import CloudPlan - features = SimpleNamespace( billing=SimpleNamespace(enabled=True, subscription=SimpleNamespace(plan=CloudPlan.SANDBOX)) ) @@ -69,8 +68,6 @@ def test_dispatch_billing_non_sandbox_uses_priority_tenant_queue(mocker: MockerF upload_mock = mocker.patch.object(proxy, "_upload_invoke_entities", return_value="file-1") send_mock = mocker.patch.object(proxy, "_send_to_priority_tenant_queue") - from enums.cloud_plan import CloudPlan - features = SimpleNamespace( billing=SimpleNamespace(enabled=True, subscription=SimpleNamespace(plan=CloudPlan.PROFESSIONAL)) ) diff --git a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py index 4ee1a5831a0..8c4485b290c 100644 --- a/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py +++ b/api/tests/unit_tests/services/rag_pipeline/test_rag_pipeline_transform_service.py @@ -5,12 +5,66 @@ from typing import cast import pytest from pytest_mock import MockerFixture +from sqlalchemy import select +from sqlalchemy.orm import Session -from models.dataset import Dataset +from extensions.storage.storage_type import StorageType +from models.dataset import Dataset, Document, DocumentPipelineExecutionLog, Pipeline +from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom +from models.model import UploadFile from services.entities.knowledge_entities.rag_pipeline_entities import KnowledgeConfiguration +from services.errors.rag_pipeline import RagPipelineResourceNotFoundError from services.rag_pipeline.rag_pipeline_transform_service import RagPipelineTransformService +def _dataset(**overrides: object) -> Dataset: + values = { + "id": "dataset-1", + "tenant_id": "tenant-1", + "name": "Dataset", + "description": "desc", + "created_by": "user-1", + "provider": "vendor", + } + values.update(overrides) + return Dataset(**values) + + +def _document(**overrides: object) -> Document: + values = { + "id": "document-1", + "tenant_id": "tenant-1", + "dataset_id": "dataset-1", + "position": 1, + "data_source_type": DataSourceType.UPLOAD_FILE, + "data_source_info": None, + "batch": "batch-1", + "name": "Document", + "created_from": DocumentCreatedFrom.WEB, + "created_by": "user-1", + } + values.update(overrides) + return Document(**values) + + +def _upload_file(*, file_id: str = "file-1", tenant_id: str = "tenant-1") -> UploadFile: + upload_file = UploadFile( + tenant_id=tenant_id, + storage_type=StorageType.LOCAL, + key="files/f.txt", + name="f.txt", + size=10, + extension="txt", + mime_type="text/plain", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + created_at=datetime.now(UTC).replace(tzinfo=None), + used=False, + ) + upload_file.id = file_id + return upload_file + + @pytest.mark.parametrize( ("doc_form", "datasource_type", "indexing_technique"), [ @@ -92,129 +146,108 @@ def test_deal_dependencies_installs_missing_marketplace_plugins(mocker: MockerFi install_mock.assert_called_once_with("tenant-1", ["missing-plugin:1.0.0"]) -def test_transform_to_empty_pipeline_updates_dataset_and_commits(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [(Dataset, Pipeline)], indirect=True) +def test_transform_to_empty_pipeline_updates_dataset_and_commits(sqlite_session: Session) -> None: service = RagPipelineTransformService() - mocker.patch( - "services.rag_pipeline.rag_pipeline_transform_service.current_user", - SimpleNamespace(id="user-1"), - ) - class FakePipeline: - def __init__(self, **kwargs): - self.id = "pipeline-1" - self.tenant_id = kwargs["tenant_id"] - self.name = kwargs["name"] - self.description = kwargs["description"] - self.created_by = kwargs["created_by"] + dataset = _dataset() + sqlite_session.add(dataset) + sqlite_session.commit() - mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.Pipeline", FakePipeline) - session_mock = mocker.Mock() - add_mock = session_mock.add - flush_mock = session_mock.flush - commit_mock = session_mock.commit + result = service._transform_to_empty_pipeline(dataset, account_id="user-1", session=sqlite_session) - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - name="Dataset", - description="desc", - pipeline_id=None, - runtime_mode="general", - updated_by=None, - updated_at=None, - ) - - result = service._transform_to_empty_pipeline(cast(Dataset, dataset), session=session_mock) - - assert result == {"pipeline_id": "pipeline-1", "dataset_id": "dataset-1", "status": "success"} - assert dataset.pipeline_id == "pipeline-1" + pipeline = sqlite_session.get(Pipeline, result["pipeline_id"]) + assert pipeline is not None + assert pipeline.name == "Dataset" + assert result == {"pipeline_id": pipeline.id, "dataset_id": "dataset-1", "status": "success"} + assert dataset.pipeline_id == pipeline.id assert dataset.runtime_mode == "rag_pipeline" assert dataset.updated_by == "user-1" - add_mock.assert_called() - flush_mock.assert_called_once() - commit_mock.assert_called_once() # --- transform_dataset --- -def test_transform_dataset_returns_early_when_pipeline_exists(mocker: MockerFixture) -> None: +def test_transform_dataset_returns_early_when_pipeline_exists(sqlite_session: Session) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace( - id="d1", - pipeline_id="p1", - runtime_mode="rag_pipeline", - ) - session_mock = mocker.Mock() - session_mock.get.return_value = dataset + dataset = _dataset(id="d1", pipeline_id="p1", runtime_mode="rag_pipeline") + pipeline = Pipeline(tenant_id="tenant-1", name="Pipeline", description="") + pipeline.id = "p1" + sqlite_session.add_all([dataset, pipeline]) - result = service.transform_dataset("d1", session_mock) + result = service.transform_dataset(dataset, "user-1", sqlite_session) assert result == {"pipeline_id": "p1", "dataset_id": "d1", "status": "success"} -def test_transform_dataset_raises_for_dataset_not_found(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("pipeline_tenant_id", [None, "tenant-2"]) +def test_transform_dataset_rejects_missing_or_foreign_pipeline_before_side_effects( + mocker: MockerFixture, + sqlite_session: Session, + pipeline_tenant_id: str | None, +) -> None: service = RagPipelineTransformService() - session_mock = mocker.Mock() - session_mock.get.return_value = None - with pytest.raises(ValueError, match="Dataset not found"): - service.transform_dataset("d1", session_mock) - - -def test_transform_dataset_raises_for_external_dataset(mocker: MockerFixture) -> None: - service = RagPipelineTransformService() - dataset = SimpleNamespace( - id="d1", - pipeline_id=None, - runtime_mode=None, - provider="external", + dataset = _dataset(id="d1", pipeline_id="p1", runtime_mode="rag_pipeline") + sqlite_session.add(dataset) + if pipeline_tenant_id is not None: + pipeline = Pipeline(tenant_id=pipeline_tenant_id, name="Pipeline", description="") + pipeline.id = "p1" + sqlite_session.add(pipeline) + install_plugins = mocker.patch( + "services.rag_pipeline.rag_pipeline_transform_service.PluginService.install_from_marketplace_pkg" ) - session_mock = mocker.Mock() - session_mock.get.return_value = dataset + create_pipeline = mocker.patch.object(service, "_create_pipeline") + transform_empty = mocker.patch.object(service, "_transform_to_empty_pipeline") + + with pytest.raises(RagPipelineResourceNotFoundError, match="Pipeline not found"): + service.transform_dataset(dataset, "user-1", sqlite_session) + + install_plugins.assert_not_called() + create_pipeline.assert_not_called() + transform_empty.assert_not_called() + + +@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) +def test_transform_dataset_raises_for_external_dataset(sqlite_session: Session) -> None: + service = RagPipelineTransformService() + dataset = _dataset(id="d1", provider="external") + sqlite_session.add(dataset) + sqlite_session.commit() with pytest.raises(ValueError, match="External dataset is not supported"): - service.transform_dataset("d1", session_mock) + service.transform_dataset(dataset, "user-1", sqlite_session) -def test_transform_dataset_calls_empty_pipeline_when_no_datasource(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) +def test_transform_dataset_calls_empty_pipeline_when_no_datasource( + mocker: MockerFixture, sqlite_session: Session +) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace( - id="d1", - pipeline_id=None, - runtime_mode=None, - provider="vendor", - data_source_type=None, - indexing_technique=None, - ) - session_mock = mocker.Mock() - session_mock.get.return_value = dataset + dataset = _dataset(id="d1", data_source_type=None, indexing_technique=None) + sqlite_session.add(dataset) + sqlite_session.commit() empty_result = {"pipeline_id": "p-empty", "dataset_id": "d1", "status": "success"} mocker.patch.object(service, "_transform_to_empty_pipeline", return_value=empty_result) - result = service.transform_dataset("d1", session_mock) + result = service.transform_dataset(dataset, "user-1", sqlite_session) assert result == empty_result -def test_transform_dataset_calls_empty_pipeline_when_no_doc_form(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [(Dataset, Document)], indirect=True) +def test_transform_dataset_calls_empty_pipeline_when_no_doc_form( + mocker: MockerFixture, sqlite_session: Session +) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace( - id="d1", - pipeline_id=None, - runtime_mode=None, - provider="vendor", - data_source_type="upload_file", - indexing_technique="high_quality", - doc_form=None, - ) - session_mock = mocker.Mock() - session_mock.get.return_value = dataset + dataset = _dataset(id="d1", data_source_type="upload_file", indexing_technique="high_quality", chunk_structure=None) + sqlite_session.add(dataset) + sqlite_session.commit() empty_result = {"pipeline_id": "p-empty", "dataset_id": "d1", "status": "success"} mocker.patch.object(service, "_transform_to_empty_pipeline", return_value=empty_result) - result = service.transform_dataset("d1", session_mock) + result = service.transform_dataset(dataset, "user-1", sqlite_session) assert result == empty_result @@ -274,78 +307,65 @@ def test_deal_knowledge_index_high_quality_sets_embedding(mocker: MockerFixture) # --- _deal_document_data --- -def test_deal_document_data_notion(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog)], indirect=True) +def test_deal_document_data_notion(sqlite_session: Session) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace(id="d1", pipeline_id="p1") - doc = SimpleNamespace( + dataset = _dataset(id="d1", pipeline_id="p1") + doc = _document( id="doc1", dataset_id="d1", data_source_type="notion_import", - data_source_info_dict={ - "notion_workspace_id": "ws1", - "notion_page_id": "page1", - "notion_page_icon": "icon1", - "type": "page", - "last_edited_time": 12345, - }, + data_source_info=( + '{"notion_workspace_id":"ws1","notion_page_id":"page1","notion_page_icon":"icon1",' + '"type":"page","last_edited_time":12345}' + ), name="Notion Doc", - created_by="u1", - created_at=datetime.now(UTC).replace(tzinfo=None), - data_source_info=None, ) + sqlite_session.add(doc) + sqlite_session.commit() - scalars_mock = mocker.Mock() - scalars_mock.all.return_value = [doc] - session_mock = mocker.Mock() - session_mock.scalars.return_value = scalars_mock - add_mock = session_mock.add - - service._deal_document_data(cast(Dataset, dataset), session_mock) + service._deal_document_data(dataset, sqlite_session) + sqlite_session.flush() assert doc.data_source_type == "online_document" assert "page1" in doc.data_source_info - assert add_mock.call_count == 2 # document + log + log = sqlite_session.scalar(select(DocumentPipelineExecutionLog)) + assert log is not None + assert log.document_id == doc.id @pytest.mark.parametrize(("provider", "node_id"), [("firecrawl", "1752565402678"), ("jinareader", "1752491761974")]) -def test_deal_document_data_website(mocker: MockerFixture, provider: str, node_id: str) -> None: +@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog)], indirect=True) +def test_deal_document_data_website(sqlite_session: Session, provider: str, node_id: str) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace(id="d1", pipeline_id="p1") - doc = SimpleNamespace( + dataset = _dataset(id="d1", pipeline_id="p1") + doc = _document( id="doc1", dataset_id="d1", data_source_type="website_crawl", - data_source_info_dict={ - "url": "https://example.com", - "provider": provider, - }, + data_source_info=f'{{"url":"https://example.com","provider":"{provider}"}}', name="Web Doc", - created_by="u1", - created_at=datetime.now(UTC).replace(tzinfo=None), - data_source_info=None, ) + sqlite_session.add(doc) + sqlite_session.commit() - scalars_mock = mocker.Mock() - scalars_mock.all.return_value = [doc] - session_mock = mocker.Mock() - session_mock.scalars.return_value = scalars_mock - add_mock = session_mock.add - - service._deal_document_data(cast(Dataset, dataset), session_mock) + service._deal_document_data(dataset, sqlite_session) + sqlite_session.flush() assert doc.data_source_type == "website_crawl" assert "example.com" in doc.data_source_info - # Check if correct node id was used in log - log = add_mock.call_args_list[1][0][0] + log = sqlite_session.scalar(select(DocumentPipelineExecutionLog)) + assert log is not None assert log.datasource_node_id == node_id # --- transform_dataset complex flow --- -def test_transform_dataset_full_flow(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) +def test_transform_dataset_full_flow(mocker: MockerFixture, sqlite_session: Session) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace( + dataset = _dataset( id="d1", tenant_id="t1", name="D", @@ -355,38 +375,40 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None: provider="vendor", data_source_type="upload_file", indexing_technique="high_quality", - doc_form="text_model", + chunk_structure="text_model", retrieval_model={"search_method": "semantic_search", "top_k": 3}, embedding_model="m1", embedding_model_provider="p1", summary_index_setting=None, - chunk_structure=None, ) - session_mock = mocker.Mock() - session_mock.get.return_value = dataset + sqlite_session.add(dataset) + sqlite_session.commit() mocker.patch.object(service, "_deal_dependencies") mocker.patch.object(service, "_deal_document_data") - session_mock.commit = mocker.Mock() - - # Mock current_user to have the same tenant_id as dataset - mock_current_user = SimpleNamespace(current_tenant_id="t1") - mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.current_user", mock_current_user) pipeline = SimpleNamespace(id="p-new") - mocker.patch.object(service, "_create_pipeline", return_value=pipeline) + create_pipeline = mocker.patch.object(service, "_create_pipeline", return_value=pipeline) - result = service.transform_dataset("d1", session_mock) + result = service.transform_dataset(dataset, "user-1", sqlite_session) assert result["pipeline_id"] == "p-new" assert dataset.runtime_mode == "rag_pipeline" assert dataset.chunk_structure == "text_model" + assert create_pipeline.call_args.kwargs == { + "tenant_id": "t1", + "account_id": "user-1", + "session": sqlite_session, + } -def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) +def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create( + mocker: MockerFixture, sqlite_session: Session +) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace( + dataset = _dataset( id="d1", tenant_id="t1", name="D", @@ -396,22 +418,25 @@ def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create provider="vendor", data_source_type="upload_file", indexing_technique="high_quality", - doc_form="unsupported", + chunk_structure="unsupported", retrieval_model=None, ) - session_mock = mocker.Mock() - session_mock.get.return_value = dataset + sqlite_session.add(dataset) + sqlite_session.commit() mocker.patch.object(service, "_get_transform_yaml", return_value={"workflow": {"graph": {"nodes": []}}}) mocker.patch.object(service, "_deal_dependencies") mocker.patch.object(service, "_create_pipeline", return_value=SimpleNamespace(id="p-new")) with pytest.raises(ValueError, match="Unsupported doc form"): - service.transform_dataset("d1", session_mock) + service.transform_dataset(dataset, "user-1", sqlite_session) -def test_transform_dataset_raises_when_transform_yaml_missing_workflow(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True) +def test_transform_dataset_raises_when_transform_yaml_missing_workflow( + mocker: MockerFixture, sqlite_session: Session +) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace( + dataset = _dataset( id="d1", tenant_id="t1", name="D", @@ -421,53 +446,82 @@ def test_transform_dataset_raises_when_transform_yaml_missing_workflow(mocker: M provider="vendor", data_source_type="upload_file", indexing_technique="high_quality", - doc_form="text_model", + chunk_structure="text_model", retrieval_model=None, ) - session_mock = mocker.Mock() - session_mock.get.return_value = dataset + sqlite_session.add(dataset) + sqlite_session.commit() mocker.patch.object(service, "_get_transform_yaml", return_value={}) mocker.patch.object(service, "_deal_dependencies") with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"): - service.transform_dataset("d1", session_mock) + service.transform_dataset(dataset, "user-1", sqlite_session) -def test_create_pipeline_raises_when_workflow_data_missing(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_create_pipeline_raises_when_workflow_data_missing(sqlite_session: Session) -> None: service = RagPipelineTransformService() - session = mocker.Mock() with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"): - service._create_pipeline({"rag_pipeline": {"name": "N"}}, session=session) + service._create_pipeline( + {"rag_pipeline": {"name": "N"}}, + tenant_id="tenant-1", + account_id="user-1", + session=sqlite_session, + ) -def test_deal_document_data_upload_file_with_existing_file(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog, UploadFile)], indirect=True) +def test_deal_document_data_upload_file_with_existing_file(sqlite_session: Session) -> None: service = RagPipelineTransformService() - dataset = SimpleNamespace(id="d1", pipeline_id="p1") - document = SimpleNamespace( + dataset = _dataset(id="d1", pipeline_id="p1") + document = _document( id="doc-1", dataset_id="d1", data_source_type="upload_file", - data_source_info_dict={"upload_file_id": "file-1"}, + data_source_info='{"upload_file_id":"file-1"}', name="Doc", - created_by="u1", - created_at=datetime.now(UTC).replace(tzinfo=None), - data_source_info=None, ) - upload_file = SimpleNamespace(name="f.txt", size=10, extension="txt", mime_type="text/plain") + sqlite_session.add_all([document, _upload_file()]) + sqlite_session.commit() - scalars_mock = mocker.Mock() - scalars_mock.all.return_value = [document] - session_mock = mocker.Mock() - session_mock.scalars.return_value = scalars_mock - session_mock.get.return_value = upload_file - add_mock = session_mock.add - - service._deal_document_data(cast(Dataset, dataset), session_mock) + service._deal_document_data(dataset, sqlite_session) + sqlite_session.flush() assert document.data_source_type == "local_file" assert "real_file_id" in document.data_source_info - assert add_mock.call_count >= 2 + log = sqlite_session.scalar(select(DocumentPipelineExecutionLog)) + assert log is not None + assert log.document_id == document.id + + +@pytest.mark.parametrize( + ("document_tenant_id", "upload_file_tenant_id"), + [("tenant-2", "tenant-1"), ("tenant-1", "tenant-2")], +) +@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog, UploadFile)], indirect=True) +def test_deal_document_data_scopes_documents_and_upload_files_to_dataset_tenant( + sqlite_session: Session, + document_tenant_id: str, + upload_file_tenant_id: str, +) -> None: + service = RagPipelineTransformService() + dataset = _dataset(id="d1", tenant_id="tenant-1", pipeline_id="p1") + document = _document( + id="doc-1", + tenant_id=document_tenant_id, + dataset_id="d1", + data_source_type="upload_file", + data_source_info='{"upload_file_id":"file-1"}', + ) + sqlite_session.add_all([document, _upload_file(tenant_id=upload_file_tenant_id)]) + sqlite_session.commit() + + service._deal_document_data(dataset, sqlite_session) + sqlite_session.flush() + + assert document.data_source_type == DataSourceType.UPLOAD_FILE + assert sqlite_session.scalar(select(DocumentPipelineExecutionLog)) is None def _make_service(): diff --git a/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py b/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py index 366ceb14f28..fd50db9f384 100644 --- a/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py +++ b/api/tests/unit_tests/services/recommend_app/test_remote_retrieval.py @@ -7,7 +7,14 @@ from sqlalchemy import Engine from sqlalchemy.orm import Session from services.recommend_app.recommend_app_type import RecommendAppType -from services.recommend_app.remote.remote_retrieval import RemoteRecommendAppRetrieval +from services.recommend_app.remote.remote_retrieval import RemoteRecommendAppRetrieval, clear_remote_fetch_cache + + +@pytest.fixture(autouse=True) +def _clear_remote_fetch_cache_between_tests(): + clear_remote_fetch_cache() + yield + clear_remote_fetch_cache() @pytest.fixture @@ -249,3 +256,63 @@ class TestFetchFromDifyOfficial: with pytest.raises(ValueError, match="fetch learn dify apps failed"): RemoteRecommendAppRetrieval.fetch_learn_dify_apps_from_dify_official("en-US") + + @patch("services.recommend_app.remote.remote_retrieval.dify_config") + @patch("services.recommend_app.remote.remote_retrieval.httpx.get") + def test_apps_uses_cache_for_repeated_requests(self, mock_get, mock_config): + mock_config.HOSTED_FETCH_APP_TEMPLATES_REMOTE_DOMAIN = "https://example.com" + mock_config.HOSTED_FETCH_APP_TEMPLATES_CACHE_TTL = 600 + mock_response = MagicMock(status_code=200) + mock_response.json.return_value = {"recommended_apps": [{"id": "app-1"}]} + mock_get.return_value = mock_response + + first = RemoteRecommendAppRetrieval.fetch_recommended_apps_from_dify_official("en-US") + second = RemoteRecommendAppRetrieval.fetch_recommended_apps_from_dify_official("en-US") + + assert first == second == {"recommended_apps": [{"id": "app-1"}]} + mock_get.assert_called_once() + + @patch("services.recommend_app.remote.remote_retrieval.dify_config") + @patch("services.recommend_app.remote.remote_retrieval.httpx.get") + def test_apps_does_not_cache_failed_responses(self, mock_get, mock_config): + mock_config.HOSTED_FETCH_APP_TEMPLATES_REMOTE_DOMAIN = "https://example.com" + mock_config.HOSTED_FETCH_APP_TEMPLATES_CACHE_TTL = 600 + mock_get.return_value = MagicMock(status_code=500) + + with pytest.raises(ValueError, match="fetch recommended apps failed"): + RemoteRecommendAppRetrieval.fetch_recommended_apps_from_dify_official("en-US") + with pytest.raises(ValueError, match="fetch recommended apps failed"): + RemoteRecommendAppRetrieval.fetch_recommended_apps_from_dify_official("en-US") + + assert mock_get.call_count == 2 + + @patch("services.recommend_app.remote.remote_retrieval.dify_config") + @patch("services.recommend_app.remote.remote_retrieval.httpx.get") + def test_apps_skips_cache_when_ttl_disabled(self, mock_get, mock_config): + mock_config.HOSTED_FETCH_APP_TEMPLATES_REMOTE_DOMAIN = "https://example.com" + mock_config.HOSTED_FETCH_APP_TEMPLATES_CACHE_TTL = 0 + mock_response = MagicMock(status_code=200) + mock_response.json.return_value = {"recommended_apps": []} + mock_get.return_value = mock_response + + RemoteRecommendAppRetrieval.fetch_recommended_apps_from_dify_official("en-US") + RemoteRecommendAppRetrieval.fetch_recommended_apps_from_dify_official("en-US") + + assert mock_get.call_count == 2 + + @patch("services.recommend_app.remote.remote_retrieval.dify_config") + @patch("services.recommend_app.remote.remote_retrieval.httpx.get") + def test_apps_cache_isolated_by_origin_header(self, mock_get, mock_config): + mock_config.HOSTED_FETCH_APP_TEMPLATES_REMOTE_DOMAIN = "https://example.com" + mock_config.HOSTED_FETCH_APP_TEMPLATES_CACHE_TTL = 600 + mock_response = MagicMock(status_code=200) + mock_response.json.return_value = {"recommended_apps": []} + mock_get.return_value = mock_response + + flask_app = Flask(__name__) + with flask_app.test_request_context(headers={"Origin": "https://cloud-a.example.com"}): + RemoteRecommendAppRetrieval.fetch_recommended_apps_from_dify_official("en-US") + with flask_app.test_request_context(headers={"Origin": "https://cloud-b.example.com"}): + RemoteRecommendAppRetrieval.fetch_recommended_apps_from_dify_official("en-US") + + assert mock_get.call_count == 2 diff --git a/api/tests/unit_tests/services/retention/test_messages_clean_policy.py b/api/tests/unit_tests/services/retention/test_messages_clean_policy.py index 79c079c683a..6179cead0cc 100644 --- a/api/tests/unit_tests/services/retention/test_messages_clean_policy.py +++ b/api/tests/unit_tests/services/retention/test_messages_clean_policy.py @@ -1,6 +1,7 @@ import datetime from unittest.mock import MagicMock, patch +from enums import DeploymentEdition from services.retention.conversation.messages_clean_policy import ( BillingDisabledPolicy, BillingSandboxPolicy, @@ -115,19 +116,19 @@ class TestBillingSandboxPolicy: class TestCreateMessageCleanPolicy: - def test_billing_disabled_returns_disabled_policy(self): + def test_non_cloud_edition_returns_disabled_policy(self): with patch(f"{MODULE}.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY policy = create_message_clean_policy() assert isinstance(policy, BillingDisabledPolicy) - def test_billing_enabled_returns_sandbox_policy(self): + def test_cloud_edition_returns_sandbox_policy(self): with ( patch(f"{MODULE}.dify_config") as cfg, patch(f"{MODULE}.BillingService") as bs, ): - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD bs.get_expired_subscription_cleanup_whitelist.return_value = ["wl1"] bs.get_plan_bulk_with_cache = MagicMock() policy = create_message_clean_policy(graceful_period_days=30) diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_archive_bundle_index.py b/api/tests/unit_tests/services/retention/workflow_run/test_archive_bundle_index.py index 0f17f96f571..6dd0098c863 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_archive_bundle_index.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_archive_bundle_index.py @@ -3,6 +3,10 @@ import json from typing import cast from unittest.mock import MagicMock +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker + +from models.workflow import WorkflowRunArchiveBundle from services.retention.workflow_run.archive_bundle_index import ( ARCHIVE_BUNDLE_ROOT_PREFIX, ArchiveBundleManifest, @@ -90,12 +94,11 @@ def test_decode_and_calculate_archive_bundle_index_values() -> None: assert values.archived_at == datetime.datetime(2026, 6, 25, 8, 0) -def test_upsert_archive_bundle_index_inserts_new_bundle() -> None: - session = MagicMock() - session.scalar.return_value = None +def test_upsert_archive_bundle_index_inserts_new_bundle(sqlite_session: Session) -> None: data = _manifest_bytes() - bundle = upsert_archive_bundle_index_from_manifest(session, decode_archive_bundle_manifest(data), len(data)) + bundle = upsert_archive_bundle_index_from_manifest(sqlite_session, decode_archive_bundle_manifest(data), len(data)) + sqlite_session.flush() assert bundle.tenant_id == TENANT_ID assert bundle.year == 2025 @@ -103,34 +106,42 @@ def test_upsert_archive_bundle_index_inserts_new_bundle() -> None: assert bundle.workflow_run_count == 2 assert bundle.row_count == 5 assert bundle.archive_bytes == len(data) + 300 - session.add.assert_called_once_with(bundle) + assert sqlite_session.get(WorkflowRunArchiveBundle, bundle.id) is bundle -def test_upsert_archive_bundle_index_updates_existing_bundle() -> None: - existing = MagicMock() - session = MagicMock() - session.scalar.return_value = existing +def test_upsert_archive_bundle_index_updates_existing_bundle(sqlite_session: Session) -> None: + existing = WorkflowRunArchiveBundle( + tenant_id=TENANT_ID, + year=2025, + month=3, + shard="00-of-01", + bundle_id=BUNDLE_ID, + workflow_run_count=1, + row_count=1, + archive_bytes=1, + archived_at=datetime.datetime(2025, 3, 1), + ) + sqlite_session.add(existing) + sqlite_session.flush() data = _manifest_bytes() - bundle = upsert_archive_bundle_index_from_manifest(session, decode_archive_bundle_manifest(data), len(data)) + bundle = upsert_archive_bundle_index_from_manifest(sqlite_session, decode_archive_bundle_manifest(data), len(data)) assert bundle is existing assert existing.workflow_run_count == 2 assert existing.row_count == 5 assert existing.archive_bytes == len(data) + 300 assert existing.archived_at == datetime.datetime(2026, 6, 25, 8, 0) - session.add.assert_not_called() + assert sqlite_session.query(WorkflowRunArchiveBundle).count() == 1 -def test_backfill_lists_tenant_month_prefix_and_upserts_bundle_index() -> None: +def test_backfill_lists_tenant_month_prefix_and_upserts_bundle_index( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: storage = FakeArchiveStorage({MANIFEST_KEY: _manifest_bytes()}) - session = MagicMock() - session.scalar.return_value = None - session_factory = MagicMock() - session_factory.return_value.__enter__.return_value = session backfill = WorkflowRunArchiveBundleIndexBackfill( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, session_factory), + session_factory=sqlite_session_factory, ) summary = backfill.run(tenant_ids=[TENANT_ID], year=2025, month=3) @@ -142,11 +153,13 @@ def test_backfill_lists_tenant_month_prefix_and_upserts_bundle_index() -> None: assert summary.bundles_processed == 1 assert summary.bundles_upserted == 1 assert summary.bundles_failed == 0 - session.add.assert_called_once() - session.commit.assert_called_once() + sqlite_session.expire_all() + assert sqlite_session.scalar(select(WorkflowRunArchiveBundle)) is not None -def test_backfill_dry_run_filters_by_year_month_without_database_write() -> None: +def test_backfill_dry_run_filters_by_year_month_without_database_write( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: other_month_prefix = OBJECT_PREFIX.replace("month=03", "month=04") storage = FakeArchiveStorage( { @@ -156,10 +169,9 @@ def test_backfill_dry_run_filters_by_year_month_without_database_write() -> None ), } ) - session_factory = MagicMock() backfill = WorkflowRunArchiveBundleIndexBackfill( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, session_factory), + session_factory=sqlite_session_factory, ) summary = backfill.run(tenant_prefixes=["1"], year=2025, month=3, dry_run=True) @@ -169,4 +181,4 @@ def test_backfill_dry_run_filters_by_year_month_without_database_write() -> None assert summary.bundles_processed == 1 assert summary.bundles_upserted == 0 assert summary.archive_bytes > 0 - session_factory.assert_not_called() + assert sqlite_session.scalar(select(WorkflowRunArchiveBundle)) is None diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_archive_download_preparation.py b/api/tests/unit_tests/services/retention/workflow_run/test_archive_download_preparation.py index 2bf96dc9701..d83a7bc978e 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_archive_download_preparation.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_archive_download_preparation.py @@ -4,7 +4,6 @@ import io import json import zipfile from contextlib import nullcontext -from types import SimpleNamespace from typing import cast from unittest.mock import MagicMock @@ -83,7 +82,17 @@ def _object_prefix(bundle_id: str = BUNDLE_ID) -> str: def _bundle(bundle_id: str = BUNDLE_ID) -> WorkflowRunArchiveBundle: - return cast(WorkflowRunArchiveBundle, SimpleNamespace(shard=SHARD, bundle_id=bundle_id)) + return WorkflowRunArchiveBundle( + tenant_id=TENANT_ID, + year=2025, + month=3, + shard=SHARD, + bundle_id=bundle_id, + workflow_run_count=1, + row_count=1, + archive_bytes=100, + archived_at=datetime.datetime(2026, 6, 25, 8), + ) def _task(bundle_refs: list[tuple[str, str]] | None = None) -> WorkflowRunArchiveDownloadTask: @@ -145,16 +154,16 @@ def _preparer( archive_storage: FakeArchiveStorage | None = None, download_storage: FakeArchiveStorage | None = None, cache: FakeTaskCache, + session_factory: sessionmaker[Session], bundles: list[WorkflowRunArchiveBundle] | None = None, ) -> WorkflowRunArchiveDownloadPreparer: source_storage = archive_storage or storage target_storage = download_storage or storage assert source_storage is not None assert target_storage is not None - session = MagicMock() - session.scalars.return_value = bundles or [_bundle()] - session_factory = MagicMock() - session_factory.return_value.__enter__.return_value = session + with session_factory() as session: + session.add_all([_bundle()] if bundles is None else bundles) + session.commit() return WorkflowRunArchiveDownloadPreparer( archive_storage=cast(ArchiveStorage, source_storage), download_storage=cast(ArchiveStorage, target_storage), @@ -169,7 +178,9 @@ def _parquet_bytes(records: list[dict[str, object]]) -> bytes: return buffer.getvalue() -def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready() -> None: +def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready( + sqlite_session_factory: sessionmaker[Session], +) -> None: bundle_refs = [(SHARD, "bundle-a"), (SHARD, "bundle-b")] task = _task(bundle_refs) first_bundle_payloads = { @@ -206,6 +217,7 @@ def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready() archive_storage=archive_storage, download_storage=download_storage, cache=cache, + session_factory=sqlite_session_factory, bundles=[_bundle("bundle-a"), _bundle("bundle-b")], ) @@ -232,7 +244,9 @@ def test_prepare_workflow_run_archive_download_builds_csv_zip_and_marks_ready() assert '"run-b","failed","safe"' in workflow_runs_csv -def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch() -> None: +def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch( + sqlite_session_factory: sessionmaker[Session], +) -> None: task = _task() table_payloads = {"workflow_runs": _parquet_bytes([{"id": "run-a", "status": "succeeded"}])} manifest_data = json.loads(_manifest_bytes(table_payloads).decode("utf-8")) @@ -244,7 +258,7 @@ def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch } ) cache = FakeTaskCache(task) - preparer = _preparer(storage=storage, cache=cache) + preparer = _preparer(storage=storage, cache=cache, session_factory=sqlite_session_factory) result = preparer.prepare(tenant_id=TENANT_ID, download_id=task.download_id) @@ -254,11 +268,18 @@ def test_prepare_workflow_run_archive_download_marks_failed_on_checksum_mismatch assert storage.put_objects == {} -def test_prepare_workflow_run_archive_download_skips_duplicate_worker() -> None: +def test_prepare_workflow_run_archive_download_skips_duplicate_worker( + sqlite_session_factory: sessionmaker[Session], +) -> None: task = _task().model_copy(update={"celery_task_id": "celery-task-1"}) storage = FakeArchiveStorage({}) cache = FakeTaskCache(task) - preparer = _preparer(storage=storage, cache=cache, bundles=[]) + preparer = _preparer( + storage=storage, + cache=cache, + session_factory=sqlite_session_factory, + bundles=[], + ) nested_results: list[WorkflowRunArchiveDownloadTask | None] = [] preparer._get_task_bundles = MagicMock(return_value=[]) @@ -277,13 +298,15 @@ def test_prepare_workflow_run_archive_download_skips_duplicate_worker() -> None: preparer._build_zip_payload.assert_called_once() -def test_failed_worker_cannot_overwrite_ready_task() -> None: +def test_failed_worker_cannot_overwrite_ready_task( + sqlite_session_factory: sessionmaker[Session], +) -> None: processing_task = _task().model_copy( update={"status": WorkflowRunArchiveDownloadStatus.PROCESSING, "celery_task_id": "celery-task-1"} ) ready_task = processing_task.model_copy(update={"status": WorkflowRunArchiveDownloadStatus.READY}) cache = FakeTaskCache(ready_task) - preparer = _preparer(storage=FakeArchiveStorage({}), cache=cache) + preparer = _preparer(storage=FakeArchiveStorage({}), cache=cache, session_factory=sqlite_session_factory) result = preparer._mark_failed(processing_task, error="late failure") diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_archive_log_service.py b/api/tests/unit_tests/services/retention/workflow_run/test_archive_log_service.py index 6fdf9d8358e..8817c71744f 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_archive_log_service.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_archive_log_service.py @@ -1,10 +1,9 @@ import datetime from contextlib import nullcontext -from types import SimpleNamespace from typing import cast -from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from models.workflow import WorkflowRunArchiveBundle from services.retention.workflow_run.archive_download_task_cache import ( @@ -63,18 +62,16 @@ def _bundle( row_count: int = 9, archived_at: datetime.datetime | None = None, ) -> WorkflowRunArchiveBundle: - return cast( - WorkflowRunArchiveBundle, - SimpleNamespace( - year=year, - month=month, - shard=shard, - bundle_id=bundle_id, - workflow_run_count=workflow_run_count, - row_count=row_count, - archive_bytes=archive_bytes, - archived_at=archived_at or datetime.datetime(2026, 6, 25, 8, 0), - ), + return WorkflowRunArchiveBundle( + tenant_id="tenant-1", + year=year, + month=month, + shard=shard, + bundle_id=bundle_id, + workflow_run_count=workflow_run_count, + row_count=row_count, + archive_bytes=archive_bytes, + archived_at=archived_at or datetime.datetime(2026, 6, 25, 8, 0), ) @@ -89,10 +86,9 @@ def _fake_dispatcher(dispatched_tasks: list[WorkflowRunArchiveDownloadTask]) -> return dispatch -def test_list_workflow_run_archives_aggregates_month_rows() -> None: +def test_list_workflow_run_archives_aggregates_month_rows(sqlite_session: Session) -> None: latest = datetime.datetime(2026, 6, 25, 8, 0) previous = datetime.datetime(2026, 6, 24, 8, 0) - session = MagicMock() march_download_id = build_archive_download_id( tenant_id="tenant-1", year=2025, @@ -117,40 +113,45 @@ def test_list_workflow_run_archives_aggregates_month_rows() -> None: } ) cache = FakeTaskCache(tasks_by_download_id={march_download_id: ready_task}) - session.scalars.return_value = [ - _bundle( - year=2025, - month=3, - shard="00-of-01", - bundle_id="bundle-a", - workflow_run_count=40, - row_count=360, - archive_bytes=1024, - archived_at=previous, - ), - _bundle( - year=2025, - month=3, - shard="00-of-01", - bundle_id="bundle-b", - workflow_run_count=60, - row_count=540, - archive_bytes=3072, - archived_at=latest, - ), - _bundle( - year=2025, - month=2, - shard="00-of-01", - bundle_id="bundle-c", - workflow_run_count=20, - row_count=180, - archive_bytes=1024, - archived_at=previous, - ), - ] + sqlite_session.add_all( + [ + _bundle( + year=2025, + month=3, + shard="00-of-01", + bundle_id="bundle-a", + workflow_run_count=40, + row_count=360, + archive_bytes=1024, + archived_at=previous, + ), + _bundle( + year=2025, + month=3, + shard="00-of-01", + bundle_id="bundle-b", + workflow_run_count=60, + row_count=540, + archive_bytes=3072, + archived_at=latest, + ), + _bundle( + year=2025, + month=2, + shard="00-of-01", + bundle_id="bundle-c", + workflow_run_count=20, + row_count=180, + archive_bytes=1024, + archived_at=previous, + ), + ] + ) + sqlite_session.flush() - result = list_workflow_run_archives(session, "tenant-1", cache=cast(WorkflowRunArchiveDownloadTaskCache, cache)) + result = list_workflow_run_archives( + sqlite_session, "tenant-1", cache=cast(WorkflowRunArchiveDownloadTaskCache, cache) + ) assert result.summary.archived_month_count == 2 assert result.summary.workflow_run_count == 120 @@ -165,17 +166,19 @@ def test_list_workflow_run_archives_aggregates_month_rows() -> None: assert result.months[1].download_task is None -def test_create_workflow_run_archive_download_task_creates_stable_pending_task() -> None: - session = MagicMock() - session.scalars.return_value = [ - _bundle(shard="01-of-02", bundle_id="bundle-b", archive_bytes=2048), - _bundle(shard="00-of-02", bundle_id="bundle-a", archive_bytes=1024), - ] +def test_create_workflow_run_archive_download_task_creates_stable_pending_task(sqlite_session: Session) -> None: + sqlite_session.add_all( + [ + _bundle(shard="01-of-02", bundle_id="bundle-b", archive_bytes=2048), + _bundle(shard="00-of-02", bundle_id="bundle-a", archive_bytes=1024), + ] + ) + sqlite_session.flush() cache = FakeTaskCache() dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = [] task = create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, @@ -188,22 +191,24 @@ def test_create_workflow_run_archive_download_task_creates_stable_pending_task() tenant_id="tenant-1", year=2025, month=3, - bundle_refs=[("01-of-02", "bundle-b"), ("00-of-02", "bundle-a")], + bundle_refs=[("00-of-02", "bundle-a"), ("01-of-02", "bundle-b")], ) assert task.requested_by == "account-1" - assert task.bundle_ids == ["bundle-b", "bundle-a"] + assert task.bundle_ids == ["bundle-a", "bundle-b"] assert [(ref.shard, ref.bundle_id) for ref in task.bundle_refs] == [ - ("01-of-02", "bundle-b"), ("00-of-02", "bundle-a"), + ("01-of-02", "bundle-b"), ] assert task.archive_bytes == 3072 assert cache.saved_task == dispatched_tasks[0] assert task.celery_task_id is not None -def test_create_workflow_run_archive_download_task_returns_existing_task_when_cache_key_exists() -> None: - session = MagicMock() - session.scalars.return_value = [_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)] +def test_create_workflow_run_archive_download_task_returns_existing_task_when_cache_key_exists( + sqlite_session: Session, +) -> None: + sqlite_session.add(_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)) + sqlite_session.flush() existing_task = build_pending_archive_download_task( tenant_id="tenant-1", requested_by="account-1", @@ -217,7 +222,7 @@ def test_create_workflow_run_archive_download_task_returns_existing_task_when_ca dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = [] task = create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, @@ -231,9 +236,11 @@ def test_create_workflow_run_archive_download_task_returns_existing_task_when_ca assert dispatched_tasks == [] -def test_create_workflow_run_archive_download_task_retries_failed_cached_task() -> None: - session = MagicMock() - session.scalars.return_value = [_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)] +def test_create_workflow_run_archive_download_task_retries_failed_cached_task( + sqlite_session: Session, +) -> None: + sqlite_session.add(_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)) + sqlite_session.flush() existing_task = build_pending_archive_download_task( tenant_id="tenant-1", requested_by="account-1", @@ -253,7 +260,7 @@ def test_create_workflow_run_archive_download_task_retries_failed_cached_task() dispatched_tasks: list[WorkflowRunArchiveDownloadTask] = [] task = create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, @@ -269,9 +276,11 @@ def test_create_workflow_run_archive_download_task_retries_failed_cached_task() @pytest.mark.parametrize("retry_failed_task", [False, True]) -def test_create_workflow_run_archive_download_task_claims_dispatch_once(retry_failed_task: bool) -> None: - session = MagicMock() - session.scalars.return_value = [_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)] +def test_create_workflow_run_archive_download_task_claims_dispatch_once( + retry_failed_task: bool, sqlite_session: Session +) -> None: + sqlite_session.add(_bundle(shard="00-of-01", bundle_id="bundle-a", archive_bytes=1024)) + sqlite_session.flush() download_id = build_archive_download_id( tenant_id="tenant-1", year=2025, @@ -301,7 +310,7 @@ def test_create_workflow_run_archive_download_task_claims_dispatch_once(retry_fa dispatched_tasks.append(task) concurrent_results.append( create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-2", year=2025, @@ -313,7 +322,7 @@ def test_create_workflow_run_archive_download_task_claims_dispatch_once(retry_fa return task result = create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, @@ -326,13 +335,10 @@ def test_create_workflow_run_archive_download_task_claims_dispatch_once(retry_fa assert concurrent_results == [result] -def test_create_workflow_run_archive_download_task_rejects_missing_month() -> None: - session = MagicMock() - session.scalars.return_value = [] - +def test_create_workflow_run_archive_download_task_rejects_missing_month(sqlite_session: Session) -> None: with pytest.raises(WorkflowRunArchiveNotFoundError): create_workflow_run_archive_download_task( - session, + sqlite_session, tenant_id="tenant-1", requested_by="account-1", year=2025, diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_bundle_archive_maintenance.py b/api/tests/unit_tests/services/retention/workflow_run/test_bundle_archive_maintenance.py index 5fefb149f3a..82b616f3b5e 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_bundle_archive_maintenance.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_bundle_archive_maintenance.py @@ -1,10 +1,14 @@ +import datetime import json from types import SimpleNamespace from typing import Any, cast from unittest.mock import MagicMock, call, patch import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker +from models.workflow import WorkflowRunArchiveBundle from services.retention.workflow_run.bundle_archive_maintenance import ( ARCHIVED_TABLES, ArchiveBundleCatalogEntry, @@ -35,14 +39,16 @@ def _table_records( return records -def _catalog_entry(*, catalog_id: str = CATALOG_ID, shard: str = "00-of-01") -> ArchiveBundleCatalogEntry: +def _catalog_entry( + *, catalog_id: str = CATALOG_ID, shard: str = "00-of-01", bundle_id: str = BUNDLE_ID +) -> ArchiveBundleCatalogEntry: return ArchiveBundleCatalogEntry( catalog_id=catalog_id, tenant_id=TENANT_ID, year=2025, month=3, shard=shard, - bundle_id=BUNDLE_ID, + bundle_id=bundle_id, workflow_run_count=0, row_count=0, archive_bytes=0, @@ -93,10 +99,25 @@ def _manifest( ).encode() -def _session_factory(session: MagicMock) -> MagicMock: - factory = MagicMock() - factory.return_value.__enter__.return_value = session - return factory +def _bundle_model(entry: ArchiveBundleCatalogEntry) -> WorkflowRunArchiveBundle: + bundle = WorkflowRunArchiveBundle( + tenant_id=entry.tenant_id, + year=entry.year, + month=entry.month, + shard=entry.shard, + bundle_id=entry.bundle_id, + workflow_run_count=entry.workflow_run_count, + row_count=entry.row_count, + archive_bytes=entry.archive_bytes, + archived_at=datetime.datetime(2026, 1, 1), + ) + bundle.id = entry.catalog_id + return bundle + + +def _persist_catalog(session: Session, entry: ArchiveBundleCatalogEntry) -> None: + session.add(_bundle_model(entry)) + session.flush() def _bundle_reference( @@ -141,26 +162,17 @@ def _sample_archive_records() -> dict[str, list[dict[str, Any]]]: ) -def test_catalog_discovery_is_ordered_and_limited_before_storage_io() -> None: - entry = _catalog_entry() - bundle = SimpleNamespace( - id=entry.catalog_id, - tenant_id=entry.tenant_id, - year=entry.year, - month=entry.month, - shard=entry.shard, - bundle_id=entry.bundle_id, - workflow_run_count=entry.workflow_run_count, - row_count=entry.row_count, - archive_bytes=entry.archive_bytes, - ) - session = MagicMock() - session.get.return_value = bundle - session.scalars.return_value = [bundle] +def test_catalog_discovery_is_ordered_and_limited_before_storage_io( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + cursor = _catalog_entry() + entry = _catalog_entry(catalog_id="019f63b7-5ca4-7681-9ce0-800283608f40", bundle_id="bundle-b") + sqlite_session.add_all([_bundle_model(cursor), _bundle_model(entry)]) + sqlite_session.commit() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) entries = maintenance._list_catalog_entries( @@ -171,36 +183,29 @@ def test_catalog_discovery_is_ordered_and_limited_before_storage_io() -> None: limit=2, ) - statement = session.scalars.call_args.args[0] - rendered = str(statement) - assert "workflow_run_archive_bundles.year" in rendered - assert "workflow_run_archive_bundles.month" in rendered - assert "workflow_run_archive_bundles.id >" in rendered - assert "ORDER BY workflow_run_archive_bundles.id ASC" in rendered - assert "LIMIT" in rendered assert entries == [entry] storage.list_objects.assert_not_called() -def test_catalog_discovery_filters_and_validates_the_requested_shard() -> None: - entry = _catalog_entry(shard="03-of-16") - bundle = SimpleNamespace( - id=entry.catalog_id, - tenant_id=entry.tenant_id, - year=entry.year, - month=entry.month, - shard=entry.shard, - bundle_id=entry.bundle_id, - workflow_run_count=entry.workflow_run_count, - row_count=entry.row_count, - archive_bytes=entry.archive_bytes, +def test_catalog_discovery_filters_and_validates_the_requested_shard( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + cursor = _catalog_entry(shard="03-of-16") + entry = _catalog_entry( + catalog_id="019f63b7-5ca4-7681-9ce0-800283608f40", + shard="03-of-16", + bundle_id="bundle-b", ) - session = MagicMock() - session.get.return_value = bundle - session.scalars.return_value = [bundle] + wrong_cursor = _catalog_entry( + catalog_id="019f63b7-5ca4-7681-9ce0-800283608f41", + shard="04-of-16", + bundle_id="bundle-c", + ) + sqlite_session.add_all([_bundle_model(cursor), _bundle_model(entry), _bundle_model(wrong_cursor)]) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) entries = maintenance._list_catalog_entries( @@ -212,34 +217,27 @@ def test_catalog_discovery_filters_and_validates_the_requested_shard() -> None: shard="03-of-16", ) - statement = session.scalars.call_args.args[0] - rendered = str(statement) - assert "workflow_run_archive_bundles.shard =" in rendered assert entries == [entry] - session.get.return_value = SimpleNamespace( - year=2025, - month=3, - tenant_id=TENANT_ID, - shard="04-of-16", - ) with pytest.raises(ValueError, match="requested archive shard"): maintenance._list_catalog_entries( tenant_ids=None, target_year=2025, target_month=3, - after_catalog_id=CATALOG_ID, + after_catalog_id=wrong_cursor.catalog_id, limit=2, shard="03-of-16", ) -def test_catalog_shard_preflight_rejects_mixed_layout_before_delete() -> None: - session = MagicMock() - session.scalars.return_value = ["00-of-01"] +def test_catalog_shard_preflight_rejects_mixed_layout_before_delete( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + sqlite_session.add(_bundle_model(_catalog_entry(shard="00-of-01"))) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) with pytest.raises(ValueError, match=r"unexpected shards.*00-of-01"): @@ -249,19 +247,15 @@ def test_catalog_shard_preflight_rejects_mixed_layout_before_delete() -> None: shard_total=16, ) - statement = session.scalars.call_args.args[0] - rendered = str(statement) - assert "workflow_run_archive_bundles.year" in rendered - assert "workflow_run_archive_bundles.month" in rendered - assert "workflow_run_archive_bundles.shard NOT IN" in rendered - -def test_catalog_shard_preflight_accepts_an_expected_subset() -> None: - session = MagicMock() - session.scalars.return_value = [] +def test_catalog_shard_preflight_accepts_an_expected_subset( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + sqlite_session.add(_bundle_model(_catalog_entry(shard="03-of-16"))) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) maintenance.validate_catalog_shards( @@ -271,12 +265,17 @@ def test_catalog_shard_preflight_accepts_an_expected_subset() -> None: ) -def test_catalog_shard_preflight_uses_requested_tenant_scope() -> None: - session = MagicMock() - session.scalars.return_value = [] +def test_catalog_shard_preflight_uses_requested_tenant_scope( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + other_entry = _catalog_entry(shard="00-of-01") + other_bundle = _bundle_model(other_entry) + other_bundle.tenant_id = "other-tenant" + sqlite_session.add(other_bundle) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) maintenance.validate_catalog_shards( @@ -286,9 +285,6 @@ def test_catalog_shard_preflight_uses_requested_tenant_scope() -> None: tenant_ids=[TENANT_ID], ) - statement = session.scalars.call_args.args[0] - assert "workflow_run_archive_bundles.tenant_id IN" in str(statement) - @pytest.mark.parametrize( ("cursor_bundle", "tenant_ids", "error_message"), @@ -310,12 +306,20 @@ def test_catalog_discovery_rejects_cursor_outside_requested_scope( cursor_bundle: SimpleNamespace | None, tenant_ids: list[str] | None, error_message: str, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, ) -> None: - session = MagicMock() - session.get.return_value = cursor_bundle + if cursor_bundle is not None: + entry = _catalog_entry() + stored = _bundle_model(entry) + stored.year = cursor_bundle.year + stored.month = cursor_bundle.month + stored.tenant_id = cursor_bundle.tenant_id + sqlite_session.add(stored) + sqlite_session.commit() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) with pytest.raises(ValueError, match=error_message): @@ -327,43 +331,38 @@ def test_catalog_discovery_rejects_cursor_outside_requested_scope( limit=1, ) - session.scalars.assert_not_called() - -def test_catalog_manifest_identity_mismatch_fails_closed() -> None: +def test_catalog_manifest_identity_mismatch_fails_closed( + unbound_session_factory: sessionmaker[Session], +) -> None: entry = _catalog_entry() storage = MagicMock() storage.get_object.return_value = _manifest(entry, bundle_id="other-bundle") maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(MagicMock())), + session_factory=unbound_session_factory, ) with pytest.raises(ValueError, match="identity does not match catalog"): maintenance._build_bundle_reference(cast(MagicMock, storage), entry) -def test_bundle_maintenance_locks_the_existing_catalog_row() -> None: +def test_bundle_maintenance_locks_the_existing_catalog_row(sqlite_session: Session) -> None: entry = _catalog_entry() - session = MagicMock() - session.scalar.return_value = entry.catalog_id + sqlite_session.add(_bundle_model(entry)) + sqlite_session.flush() - WorkflowRunBundleArchiveMaintenance._lock_catalog_entry(session, entry) - - statement = session.scalar.call_args.args[0] - rendered = str(statement) - assert "workflow_run_archive_bundles.id" in rendered - assert "workflow_run_archive_bundles.tenant_id" in rendered - assert "FOR UPDATE" in rendered + WorkflowRunBundleArchiveMaintenance._lock_catalog_entry(sqlite_session, entry) -def test_failure_and_dry_run_do_not_return_a_persistable_cursor() -> None: +def test_failure_and_dry_run_do_not_return_a_persistable_cursor( + sqlite_session_factory: sessionmaker[Session], +) -> None: entry = _catalog_entry() - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) with ( @@ -385,7 +384,7 @@ def test_failure_and_dry_run_do_not_return_a_persistable_cursor() -> None: dry_run = WorkflowRunBundleArchiveMaintenance( dry_run=True, storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) bundle_ref = BundleReference( catalog=entry, @@ -468,12 +467,13 @@ def test_live_archive_subset_rejects_content_mismatch() -> None: ) -def test_live_bundle_scope_includes_archived_ids_and_indirect_children() -> None: +def test_live_bundle_scope_includes_archived_ids_and_indirect_children( + unbound_session: Session, unbound_session_factory: sessionmaker[Session] +) -> None: archive_records = _sample_archive_records() manifest = _bundle_reference(_catalog_entry(), table_records=archive_records).manifest - session = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=unbound_session_factory, ) def select_live_parent_ids(_session, model, _run_ids): @@ -488,7 +488,7 @@ def test_live_bundle_scope_includes_archived_ids_and_indirect_children() -> None patch.object(maintenance, "_load_records_by_column", return_value=[]) as load_records, ): maintenance._load_live_bundle_records( - session, + unbound_session, manifest, archive_records, lock=True, @@ -511,8 +511,11 @@ def test_live_bundle_scope_includes_archived_ids_and_indirect_children() -> None assert ("workflow_app_logs", "id", {"app-log-1"}, True) in queries -def test_delete_bundle_accepts_matching_partial_rows_and_deletes_only_that_subset() -> None: +def test_delete_bundle_accepts_matching_partial_rows_and_deletes_only_that_subset( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() + _persist_catalog(sqlite_session, entry) archive_records = _sample_archive_records() bundle_ref = _bundle_reference(entry, table_records=archive_records) partial_records = _table_records( @@ -520,11 +523,13 @@ def test_delete_bundle_accepts_matching_partial_rows_and_deletes_only_that_subse workflow_pause_reasons=archive_records["workflow_pause_reasons"], ) expected_deleted_counts = {table_name: len(partial_records[table_name]) for table_name in ARCHIVED_TABLES} - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + transaction_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: transaction_events.append("commit")) + event.listen(sqlite_session, "after_rollback", lambda _session: transaction_events.append("rollback")) with ( patch.object(maintenance, "_is_restore_started", return_value=False), @@ -548,24 +553,28 @@ def test_delete_bundle_accepts_matching_partial_rows_and_deletes_only_that_subse patch.object(maintenance, "_mark_deleted") as mark_deleted, patch.object(maintenance, "_delete_marker"), ): - result = maintenance._delete_bundle(session, storage, bundle_ref) + result = maintenance._delete_bundle(sqlite_session, storage, bundle_ref) assert result.success - delete_bundle_rows.assert_called_once_with(session, partial_records) - session.commit.assert_called_once_with() - session.rollback.assert_not_called() + delete_bundle_rows.assert_called_once_with(sqlite_session, partial_records) + assert transaction_events == ["commit"] mark_deleted.assert_called_once_with(storage, bundle_ref.object_prefix) -def test_delete_bundle_marks_an_already_absent_source_without_deleting_rows() -> None: +def test_delete_bundle_marks_an_already_absent_source_without_deleting_rows( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() + _persist_catalog(sqlite_session, entry) archive_records = _sample_archive_records() bundle_ref = _bundle_reference(entry, table_records=archive_records) - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + transaction_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: transaction_events.append("commit")) + event.listen(sqlite_session, "after_rollback", lambda _session: transaction_events.append("rollback")) with ( patch.object(maintenance, "_is_restore_started", return_value=False), @@ -580,28 +589,31 @@ def test_delete_bundle_marks_an_already_absent_source_without_deleting_rows() -> patch.object(maintenance, "_mark_deleted") as mark_deleted, patch.object(maintenance, "_delete_marker") as delete_marker, ): - result = maintenance._delete_bundle(session, storage, bundle_ref) + result = maintenance._delete_bundle(sqlite_session, storage, bundle_ref) assert result.success delete_bundle_rows.assert_not_called() - session.commit.assert_not_called() - session.rollback.assert_not_called() + assert transaction_events == [] mark_deleted.assert_called_once_with(storage, bundle_ref.object_prefix) assert delete_marker.call_count == 2 -def test_delete_bundle_with_deleted_marker_rejects_remaining_orphan_children() -> None: +def test_delete_bundle_with_deleted_marker_rejects_remaining_orphan_children( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() + _persist_catalog(sqlite_session, entry) archive_records = _sample_archive_records() bundle_ref = _bundle_reference(entry, table_records=archive_records) live_records = _table_records( workflow_node_execution_offload=archive_records["workflow_node_execution_offload"], ) - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + rollback_events: list[str] = [] + event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback")) with ( patch.object(maintenance, "_is_restore_started", return_value=False), @@ -614,42 +626,45 @@ def test_delete_bundle_with_deleted_marker_rejects_remaining_orphan_children() - patch.object(maintenance, "_load_live_bundle_records", return_value=live_records), patch.object(maintenance, "_delete_bundle_rows") as delete_bundle_rows, ): - result = maintenance._delete_bundle(session, storage, bundle_ref) + result = maintenance._delete_bundle(sqlite_session, storage, bundle_ref) assert not result.success assert "Live rows exist for bundle with deleted marker" in result.error delete_bundle_rows.assert_not_called() - session.commit.assert_not_called() - session.rollback.assert_called_once_with() + assert rollback_events == ["rollback"] -def test_delete_bundle_rejects_an_in_progress_restore() -> None: +def test_delete_bundle_rejects_an_in_progress_restore( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() + _persist_catalog(sqlite_session, entry) bundle_ref = _bundle_reference(entry) - session = MagicMock() storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + rollback_events: list[str] = [] + event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback")) with ( patch.object(maintenance, "_is_restore_started", return_value=True), patch.object(maintenance, "_validate_archive_object") as validate_archive, ): - result = maintenance._delete_bundle(session, storage, bundle_ref) + result = maintenance._delete_bundle(sqlite_session, storage, bundle_ref) assert not result.success assert "reconcile restore before delete" in result.error validate_archive.assert_not_called() - session.commit.assert_not_called() - session.rollback.assert_called_once_with() + assert rollback_events == ["rollback"] -def test_delete_bundle_rows_use_only_verified_primary_keys() -> None: +def test_delete_bundle_rows_use_only_verified_primary_keys( + unbound_session: Session, unbound_session_factory: sessionmaker[Session] +) -> None: live_records = _sample_archive_records() - session = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=unbound_session_factory, ) with patch.object( @@ -657,7 +672,7 @@ def test_delete_bundle_rows_use_only_verified_primary_keys() -> None: "_delete_by_column", side_effect=lambda _session, _model, _column, values: len(values), ) as delete_by_column: - deleted_counts = maintenance._delete_bundle_rows(session, live_records) + deleted_counts = maintenance._delete_bundle_rows(unbound_session, live_records) expected_table_order = [ "workflow_pause_reasons", @@ -673,33 +688,39 @@ def test_delete_bundle_rows_use_only_verified_primary_keys() -> None: assert deleted_counts == {table_name: len(live_records[table_name]) for table_name in ARCHIVED_TABLES} -def test_restore_does_not_skip_an_interrupted_delete_without_deleted_marker() -> None: +def test_restore_does_not_skip_an_interrupted_delete_without_deleted_marker( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() - session = MagicMock() + _persist_catalog(sqlite_session, entry) maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + commit_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: commit_events.append("commit")) with ( patch.object(maintenance, "_is_deleted", return_value=False), patch.object(maintenance, "_is_delete_started", return_value=True), patch.object(maintenance, "_validate_live_counts") as validate_live_counts, ): - result = maintenance._restore_bundle(session, MagicMock(), _bundle_reference(entry)) + result = maintenance._restore_bundle(sqlite_session, MagicMock(), _bundle_reference(entry)) assert not result.success assert "reconcile delete first" in result.error validate_live_counts.assert_not_called() - session.commit.assert_not_called() + assert commit_events == [] -def test_restore_does_not_skip_missing_source_rows_without_deleted_marker() -> None: +def test_restore_does_not_skip_missing_source_rows_without_deleted_marker( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() - session = MagicMock() + _persist_catalog(sqlite_session, entry) maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, MagicMock()), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) bundle_ref = _bundle_reference(entry) @@ -712,22 +733,25 @@ def test_restore_does_not_skip_missing_source_rows_without_deleted_marker() -> N side_effect=ValueError("source rows are missing"), ) as validate_live_counts, ): - result = maintenance._restore_bundle(session, MagicMock(), bundle_ref) + result = maintenance._restore_bundle(sqlite_session, MagicMock(), bundle_ref) assert not result.success assert "source rows are missing" in result.error - validate_live_counts.assert_called_once_with(session, bundle_ref.manifest, expected_present=True) - session.commit.assert_not_called() + validate_live_counts.assert_called_once_with(sqlite_session, bundle_ref.manifest, expected_present=True) -def test_restore_reconciles_a_started_marker_after_the_source_commit() -> None: +def test_restore_reconciles_a_started_marker_after_the_source_commit( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: entry = _catalog_entry() - session = MagicMock() + _persist_catalog(sqlite_session, entry) storage = MagicMock() maintenance = WorkflowRunBundleArchiveMaintenance( storage=cast(MagicMock, storage), - session_factory=cast(MagicMock, _session_factory(session)), + session_factory=sqlite_session_factory, ) + commit_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: commit_events.append("commit")) bundle_ref = _bundle_reference(entry) with ( @@ -737,12 +761,12 @@ def test_restore_reconciles_a_started_marker_after_the_source_commit() -> None: patch.object(maintenance, "_validate_live_counts") as validate_live_counts, patch.object(maintenance, "_mark_restored") as mark_restored, ): - result = maintenance._restore_bundle(session, storage, bundle_ref) + result = maintenance._restore_bundle(sqlite_session, storage, bundle_ref) assert result.success - validate_live_counts.assert_called_once_with(session, bundle_ref.manifest, expected_present=True) + validate_live_counts.assert_called_once_with(sqlite_session, bundle_ref.manifest, expected_present=True) mark_restored.assert_called_once_with(storage, bundle_ref.object_prefix) - session.commit.assert_not_called() + assert commit_events == [] def test_mark_restored_clears_stale_delete_marker_before_releasing_restore_fence() -> None: diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_clear_free_plan_expired_workflow_run_logs.py b/api/tests/unit_tests/services/retention/workflow_run/test_clear_free_plan_expired_workflow_run_logs.py index 524dcd5952f..e8cc8a0c2a7 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_clear_free_plan_expired_workflow_run_logs.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_clear_free_plan_expired_workflow_run_logs.py @@ -8,6 +8,7 @@ from unittest.mock import MagicMock, patch import pytest from sqlalchemy.orm import Session +from enums import CloudPlan, DeploymentEdition from repositories.api_workflow_run_repository import WorkflowRunCleanupRef from services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs import WorkflowRunCleanup @@ -29,7 +30,7 @@ def mock_repo(): def cleanup(mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY yield WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) @@ -42,7 +43,7 @@ class TestWorkflowRunCleanupInit: def test_only_start_from_raises(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY with pytest.raises(ValueError, match="both set or both omitted"): WorkflowRunCleanup( days=30, @@ -54,7 +55,7 @@ class TestWorkflowRunCleanupInit: def test_only_end_before_raises(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY with pytest.raises(ValueError, match="both set or both omitted"): WorkflowRunCleanup( days=30, @@ -66,7 +67,7 @@ class TestWorkflowRunCleanupInit: def test_end_before_not_greater_than_start_raises(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY with pytest.raises(ValueError, match="end_before must be greater than start_from"): WorkflowRunCleanup( days=30, @@ -80,7 +81,7 @@ class TestWorkflowRunCleanupInit: dt = datetime.datetime(2024, 1, 1) with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY with pytest.raises(ValueError): WorkflowRunCleanup( days=30, @@ -93,21 +94,21 @@ class TestWorkflowRunCleanupInit: def test_zero_batch_size_raises(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY with pytest.raises(ValueError, match="batch_size must be greater than 0"): WorkflowRunCleanup(days=30, batch_size=0, workflow_run_repo=mock_repo) def test_negative_batch_size_raises(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY with pytest.raises(ValueError): WorkflowRunCleanup(days=30, batch_size=-1, workflow_run_repo=mock_repo) def test_valid_window_init(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 7 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY start = datetime.datetime(2024, 1, 1) end = datetime.datetime(2024, 6, 1) c = WorkflowRunCleanup( @@ -123,7 +124,7 @@ class TestWorkflowRunCleanupInit: def test_default_task_label_is_custom(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY c = WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) assert c._metrics._base_attributes["task_label"] == "custom" @@ -216,17 +217,17 @@ class TestIsWithinGracePeriod: class TestGetCleanupWhitelist: - def test_billing_disabled_returns_empty(self, cleanup): + def test_non_cloud_edition_returns_empty(self, cleanup): cleanup._cleanup_whitelist = None with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY result = cleanup._get_cleanup_whitelist() assert result == set() - def test_billing_enabled_fetches_whitelist(self, mock_repo): + def test_cloud_edition_fetches_whitelist(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD c = WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) with patch( "services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.BillingService" @@ -243,7 +244,7 @@ class TestGetCleanupWhitelist: def test_billing_service_error_returns_empty(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD c = WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) with patch( "services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.BillingService" @@ -259,27 +260,25 @@ class TestGetCleanupWhitelist: class TestFilterFreeTenants: - def test_billing_disabled_all_tenants_free(self, cleanup): + def test_non_cloud_edition_treats_all_tenants_as_free(self, cleanup): result = cleanup._filter_free_tenants(["t1", "t2"]) assert result == {"t1", "t2"} def test_empty_tenants_returns_empty(self, cleanup): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD result = cleanup._filter_free_tenants([]) assert result == set() def test_whitelisted_tenant_excluded(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD c = WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) c._cleanup_whitelist = {"t1"} with patch( "services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.BillingService" ) as bs: - from enums.cloud_plan import CloudPlan - bs.get_plan_bulk_with_cache.return_value = { "t1": {"plan": CloudPlan.SANDBOX, "expiration_date": -1}, "t2": {"plan": CloudPlan.SANDBOX, "expiration_date": -1}, @@ -291,7 +290,7 @@ class TestFilterFreeTenants: def test_paid_tenant_excluded(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD c = WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) c._cleanup_whitelist = set() with patch( @@ -306,7 +305,7 @@ class TestFilterFreeTenants: def test_missing_billing_info_treats_as_non_free(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD c = WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) c._cleanup_whitelist = set() with patch( @@ -319,7 +318,7 @@ class TestFilterFreeTenants: def test_billing_bulk_error_treats_as_non_free(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = True + cfg.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD c = WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) c._cleanup_whitelist = set() with patch( @@ -336,17 +335,17 @@ class TestFilterFreeTenants: class TestRunDeleteMode: - def _make_cleanup(self, mock_repo, billing_enabled=False): + def _make_cleanup(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = billing_enabled + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY return WorkflowRunCleanup(days=30, batch_size=10, workflow_run_repo=mock_repo) def test_no_rows_stops_immediately(self, mock_repo): mock_repo.get_cleanup_refs_batch_by_time_range.return_value = [] c = self._make_cleanup(mock_repo) with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY c.run() mock_repo.delete_runs_with_related_by_ids.assert_not_called() @@ -354,10 +353,10 @@ class TestRunDeleteMode: ref = make_ref("t1") mock_repo.get_cleanup_refs_batch_by_time_range.side_effect = [[ref], []] c = self._make_cleanup(mock_repo) - # billing disabled -> all free; but let's override _filter_free_tenants to return empty + # Override the non-Cloud default to exercise the no-deletion path. c._filter_free_tenants = MagicMock(return_value=set()) with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY c.run() mock_repo.delete_runs_with_related_by_ids.assert_not_called() @@ -375,7 +374,7 @@ class TestRunDeleteMode: } c = self._make_cleanup(mock_repo) with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.time.sleep"): c.run() mock_repo.delete_runs_with_related_by_ids.assert_called_once() @@ -386,7 +385,7 @@ class TestRunDeleteMode: mock_repo.delete_runs_with_related_by_ids.side_effect = RuntimeError("db error") c = self._make_cleanup(mock_repo) with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY with pytest.raises(RuntimeError): c.run() @@ -394,7 +393,7 @@ class TestRunDeleteMode: mock_repo.get_cleanup_refs_batch_by_time_range.return_value = [] with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY c = WorkflowRunCleanup( days=30, batch_size=10, @@ -414,7 +413,7 @@ class TestRunDryRunMode: def _make_dry_cleanup(self, mock_repo): with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY return WorkflowRunCleanup( days=30, batch_size=10, @@ -436,7 +435,7 @@ class TestRunDryRunMode: } c = self._make_dry_cleanup(mock_repo) with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY c.run() mock_repo.delete_runs_with_related_by_ids.assert_not_called() mock_repo.count_runs_with_related_by_ids.assert_called_once() @@ -445,7 +444,7 @@ class TestRunDryRunMode: mock_repo.get_cleanup_refs_batch_by_time_range.return_value = [] with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: cfg.SANDBOX_EXPIRED_RECORDS_CLEAN_GRACEFUL_PERIOD = 0 - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY c = WorkflowRunCleanup( days=30, batch_size=10, @@ -462,7 +461,7 @@ class TestRunDryRunMode: c = self._make_dry_cleanup(mock_repo) c._filter_free_tenants = MagicMock(return_value=set()) with patch("services.retention.workflow_run.clear_free_plan_expired_workflow_run_logs.dify_config") as cfg: - cfg.BILLING_ENABLED = False + cfg.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY c.run() mock_repo.count_runs_with_related_by_ids.assert_not_called() diff --git a/api/tests/unit_tests/services/retention/workflow_run/test_restore_archived_workflow_run.py b/api/tests/unit_tests/services/retention/workflow_run/test_restore_archived_workflow_run.py index 4768215210e..62c042ff2cf 100644 --- a/api/tests/unit_tests/services/retention/workflow_run/test_restore_archived_workflow_run.py +++ b/api/tests/unit_tests/services/retention/workflow_run/test_restore_archived_workflow_run.py @@ -10,14 +10,19 @@ import io import json import logging import zipfile +from collections.abc import Iterator +from dataclasses import dataclass from datetime import datetime -from unittest.mock import Mock, create_autospec, patch +from unittest.mock import Mock, patch import pytest from pydantic import ValidationError -from sqlalchemy import Column, Integer, MetaData, String, Table +from sqlalchemy import Column, Engine, Integer, MetaData, String, Table, delete, event, func, select +from sqlalchemy.dialects.sqlite import insert as sqlite_insert +from sqlalchemy.orm import Session, sessionmaker from libs.archive_storage import ArchiveStorageNotConfiguredError +from models.enums import CreatorUserRole from models.trigger import WorkflowTriggerLog from models.workflow import ( WorkflowAppLog, @@ -28,6 +33,7 @@ from models.workflow import ( WorkflowPauseReason, WorkflowRun, ) +from services.retention.workflow_run import restore_archived_workflow_run as restore_module from services.retention.workflow_run.restore_archived_workflow_run import ( SCHEMA_MAPPERS, TABLE_MODELS, @@ -36,24 +42,49 @@ from services.retention.workflow_run.restore_archived_workflow_run import ( ) +@dataclass(frozen=True) +class Database: + """Explicit SQLite engine, caller session, and real service-owned session factory.""" + + engine: Engine + session: Session + session_maker: sessionmaker[Session] + + +@pytest.fixture +def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Database]: + WorkflowRun.metadata.create_all( + sqlite_engine, + tables=[WorkflowRun.__table__, WorkflowAppLog.__table__, WorkflowArchiveLog.__table__], + ) + session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + with session_maker() as session: + database = Database(engine=sqlite_engine, session=session, session_maker=session_maker) + monkeypatch.setattr(restore_module, "db", database) + # Production constructs PostgreSQL's equivalent statement; SQLite's + # dialect keeps the conflict behavior executable in these tests. + monkeypatch.setattr(restore_module, "pg_insert", sqlite_insert) + yield database + + class WorkflowRunRestoreTestDataFactory: """ - Factory for creating test data and mock objects. + Factory for creating persisted-model-compatible test data. Provides reusable methods to create consistent mock objects for testing workflow run restore operations. """ @staticmethod - def create_workflow_run_mock( + def create_workflow_run( run_id: str = "run-123", tenant_id: str = "tenant-123", app_id: str = "app-123", created_at: datetime | None = None, **kwargs, - ) -> Mock: + ) -> WorkflowRun: """ - Create a mock WorkflowRun object. + Create a concrete WorkflowRun object. Args: run_id: Unique identifier for the workflow run @@ -63,27 +94,44 @@ class WorkflowRunRestoreTestDataFactory: **kwargs: Additional attributes to set on the mock Returns: - Mock WorkflowRun object with specified attributes + WorkflowRun object with specified attributes """ - run = create_autospec(WorkflowRun, instance=True) - run.id = run_id - run.tenant_id = tenant_id - run.app_id = app_id - run.created_at = created_at or datetime(2024, 1, 1, 12, 0, 0) - for key, value in kwargs.items(): - setattr(run, key, value) + attrs = { + "id": run_id, + "tenant_id": tenant_id, + "app_id": app_id, + "workflow_id": "workflow-123", + "type": "workflow", + "triggered_from": "app-run", + "version": "1", + "graph": None, + "inputs": None, + "status": "succeeded", + "outputs": "{}", + "error": None, + "elapsed_time": 0, + "total_tokens": 0, + "total_steps": 0, + "created_by_role": CreatorUserRole.ACCOUNT, + "created_by": "user-123", + "created_at": created_at or datetime(2024, 1, 1, 12, 0, 0), + "finished_at": None, + "exceptions_count": 0, + } + attrs.update(kwargs) + run = WorkflowRun(**attrs) return run @staticmethod - def create_workflow_archive_log_mock( + def create_workflow_archive_log( run_id: str = "run-123", tenant_id: str = "tenant-123", app_id: str = "app-123", created_at: datetime | None = None, **kwargs, - ) -> Mock: + ) -> WorkflowArchiveLog: """ - Create a mock WorkflowArchiveLog object. + Create a concrete WorkflowArchiveLog object. Args: run_id: Unique identifier for the workflow run @@ -93,16 +141,32 @@ class WorkflowRunRestoreTestDataFactory: **kwargs: Additional attributes to set on the mock Returns: - Mock WorkflowArchiveLog object with specified attributes + WorkflowArchiveLog object with specified attributes """ - archive_log = create_autospec(WorkflowArchiveLog, instance=True) - archive_log.workflow_run_id = run_id - archive_log.tenant_id = tenant_id - archive_log.app_id = app_id - archive_log.run_created_at = created_at or datetime(2024, 1, 1, 12, 0, 0) - for key, value in kwargs.items(): - setattr(archive_log, key, value) - return archive_log + attrs = { + "tenant_id": tenant_id, + "app_id": app_id, + "workflow_id": "workflow-123", + "workflow_run_id": run_id, + "created_by_role": CreatorUserRole.ACCOUNT, + "created_by": "user-123", + "log_id": None, + "log_created_at": None, + "log_created_from": None, + "run_version": "1", + "run_status": "succeeded", + "run_triggered_from": "app-run", + "run_error": None, + "run_elapsed_time": 0, + "run_total_tokens": 0, + "run_total_steps": 0, + "run_created_at": created_at or datetime(2024, 1, 1, 12, 0, 0), + "run_finished_at": None, + "run_exceptions_count": 0, + "trigger_metadata": None, + } + attrs.update(kwargs) + return WorkflowArchiveLog(**attrs) @staticmethod def create_archive_zip_mock( @@ -137,7 +201,7 @@ class WorkflowRunRestoreTestDataFactory: "app_id": "app-123", "workflow_id": "workflow-123", "type": "workflow", - "triggered_from": "app", + "triggered_from": "app-run", "version": "1", "status": "succeeded", "created_by_role": "account", @@ -151,7 +215,7 @@ class WorkflowRunRestoreTestDataFactory: "app_id": "app-123", "workflow_id": "workflow-123", "workflow_run_id": "run-123", - "created_from": "app", + "created_from": "service-api", "created_by_role": "account", "created_by": "user-123", }, @@ -161,7 +225,7 @@ class WorkflowRunRestoreTestDataFactory: "app_id": "app-123", "workflow_id": "workflow-123", "workflow_run_id": "run-123", - "created_from": "app", + "created_from": "service-api", "created_by_role": "account", "created_by": "user-123", }, @@ -225,14 +289,10 @@ class TestGetWorkflowRunRepo: """Tests for WorkflowRunRestore._get_workflow_run_repo method.""" @patch("services.retention.workflow_run.restore_archived_workflow_run.DifyAPIRepositoryFactory") - @patch("services.retention.workflow_run.restore_archived_workflow_run.sessionmaker") - @patch("services.retention.workflow_run.restore_archived_workflow_run.db") - def test_first_call_creates_repo(self, mock_db, mock_sessionmaker, mock_factory): + def test_first_call_creates_repo(self, mock_factory, database: Database): """First call should create and cache repository.""" restore = WorkflowRunRestore() - mock_session = Mock() - mock_sessionmaker.return_value = mock_session mock_repo = Mock() mock_factory.create_api_workflow_run_repository.return_value = mock_repo @@ -240,8 +300,9 @@ class TestGetWorkflowRunRepo: assert result is mock_repo assert restore.workflow_run_repo is mock_repo - mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) - mock_factory.create_api_workflow_run_repository.assert_called_once_with(mock_session) + session_maker = mock_factory.create_api_workflow_run_repository.call_args.args[0] + assert isinstance(session_maker, sessionmaker) + assert session_maker.kw["bind"] is database.engine def test_cached_repo_returned(self): """Subsequent calls should return cached repository.""" @@ -492,47 +553,27 @@ class TestGetModelColumnInfo: class TestRestoreTableRecords: """Tests for WorkflowRunRestore._restore_table_records method.""" - @patch("services.retention.workflow_run.restore_archived_workflow_run.TABLE_MODELS") - def test_unknown_table_returns_zero(self, mock_table_models, caplog: pytest.LogCaptureFixture): + def test_unknown_table_returns_zero(self, database: Database, caplog: pytest.LogCaptureFixture): """Should return 0 for unknown table.""" restore = WorkflowRunRestore() - mock_table_models.get.return_value = None - - mock_session = Mock() records = [{"id": "test"}] caplog.set_level(logging.WARNING, logger="services.retention.workflow_run.restore_archived_workflow_run") - result = restore._restore_table_records(mock_session, "unknown_table", records, schema_version="1.0") + result = restore._restore_table_records(database.session, "unknown_table", records, schema_version="1.0") assert result == 0 assert "Unknown table: unknown_table" in caplog.messages - def test_empty_records_returns_zero(self): + def test_empty_records_returns_zero(self, database: Database): """Should return 0 for empty records list.""" restore = WorkflowRunRestore() - mock_session = Mock() - - result = restore._restore_table_records(mock_session, "workflow_runs", [], schema_version="1.0") + result = restore._restore_table_records(database.session, "workflow_runs", [], schema_version="1.0") assert result == 0 - @patch("services.retention.workflow_run.restore_archived_workflow_run.pg_insert") - @patch("services.retention.workflow_run.restore_archived_workflow_run.cast") - def test_successful_restore(self, mock_cast, mock_pg_insert): + def test_successful_restore(self, database: Database): """Should successfully restore records.""" restore = WorkflowRunRestore() - # Mock session and execution - mock_session = Mock() - mock_result = Mock() - mock_result.rowcount = 2 - mock_session.execute.return_value = mock_result - mock_cast.return_value = mock_result - - # Mock insert statement - mock_stmt = Mock() - mock_stmt.on_conflict_do_nothing.return_value = mock_stmt - mock_pg_insert.return_value = mock_stmt - records = [ { "id": "test1", @@ -540,7 +581,7 @@ class TestRestoreTableRecords: "app_id": "app-123", "workflow_id": "workflow-123", "type": "workflow", - "triggered_from": "app", + "triggered_from": "app-run", "version": "1", "status": "succeeded", "created_by_role": "account", @@ -552,7 +593,7 @@ class TestRestoreTableRecords: "app_id": "app-123", "workflow_id": "workflow-123", "type": "workflow", - "triggered_from": "app", + "triggered_from": "app-run", "version": "1", "status": "succeeded", "created_by_role": "account", @@ -560,38 +601,20 @@ class TestRestoreTableRecords: }, ] - result = restore._restore_table_records(mock_session, "workflow_runs", records, schema_version="1.0") + result = restore._restore_table_records(database.session, "workflow_runs", records, schema_version="1.0") assert result == 2 - mock_session.execute.assert_called_once() + assert database.session.scalar(select(func.count(WorkflowRun.id))) == 2 + assert restore._restore_table_records(database.session, "workflow_runs", records, schema_version="1.0") == 0 - def test_missing_required_columns_raises_error(self): + def test_missing_required_columns_raises_error(self, database: Database): """Should raise ValueError for missing required columns.""" restore = WorkflowRunRestore() - mock_session = Mock() - # Use a dedicated mock model to isolate required-column validation behavior. - mock_model = Mock() + records = [{"id": "test"}] - # Mock a required column - required_column = Mock() - required_column.key = "required_field" - required_column.nullable = False - required_column.default = None - required_column.server_default = None - required_column.autoincrement = False - required_column.type = Mock() - - # Mock the __table__ attribute properly - mock_table = Mock() - mock_table.columns = [required_column] - mock_model.__table__ = mock_table - - records = [{"name": "test"}] # Missing required 'required_field' - - with patch.dict(TABLE_MODELS, {"test_table": mock_model}): - with pytest.raises(ValueError, match="Missing required columns for test_table"): - restore._restore_table_records(mock_session, "test_table", records, schema_version="1.0") + with pytest.raises(ValueError, match="Missing required columns for workflow_runs"): + restore._restore_table_records(database.session, "workflow_runs", records, schema_version="1.0") # --------------------------------------------------------------------------- @@ -603,38 +626,38 @@ class TestRestoreFromRun: """Tests for WorkflowRunRestore._restore_from_run method.""" @patch("services.retention.workflow_run.restore_archived_workflow_run.get_archive_storage") - def test_archive_storage_not_configured(self, mock_get_storage): + def test_archive_storage_not_configured(self, mock_get_storage, database: Database): """Should handle ArchiveStorageNotConfiguredError.""" restore = WorkflowRunRestore() mock_get_storage.side_effect = ArchiveStorageNotConfiguredError("Storage not configured") - run = WorkflowRunRestoreTestDataFactory.create_workflow_run_mock() + run = WorkflowRunRestoreTestDataFactory.create_workflow_run() with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - result = restore._restore_from_run(run, session_maker=lambda: Mock()) + result = restore._restore_from_run(run, session_maker=database.session_maker) assert result.success is False assert "Storage not configured" in result.error assert result.elapsed_time > 0 @patch("services.retention.workflow_run.restore_archived_workflow_run.get_archive_storage") - def test_archive_bundle_not_found(self, mock_get_storage): + def test_archive_bundle_not_found(self, mock_get_storage, database: Database): """Should handle FileNotFoundError when archive bundle is missing.""" restore = WorkflowRunRestore() mock_storage = Mock() mock_storage.get_object.side_effect = FileNotFoundError("Bundle not found") mock_get_storage.return_value = mock_storage - run = WorkflowRunRestoreTestDataFactory.create_workflow_run_mock() + run = WorkflowRunRestoreTestDataFactory.create_workflow_run() with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - result = restore._restore_from_run(run, session_maker=lambda: Mock()) + result = restore._restore_from_run(run, session_maker=database.session_maker) assert result.success is False assert "Archive bundle not found" in result.error @patch("services.retention.workflow_run.restore_archived_workflow_run.get_archive_storage") - def test_dry_run_mode(self, mock_get_storage): + def test_dry_run_mode(self, mock_get_storage, database: Database): """Should handle dry run mode correctly.""" restore = WorkflowRunRestore(dry_run=True) @@ -644,23 +667,16 @@ class TestRestoreFromRun: mock_storage.get_object.return_value = archive_data mock_get_storage.return_value = mock_storage - run = WorkflowRunRestoreTestDataFactory.create_workflow_run_mock() + run = WorkflowRunRestoreTestDataFactory.create_workflow_run() - # Create a proper mock session with context manager support - mock_session = Mock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - - result = restore._restore_from_run(run, session_maker=lambda: mock_session) + result = restore._restore_from_run(run, session_maker=database.session_maker) assert result.success is True assert result.restored_counts["workflow_runs"] == 1 assert result.restored_counts["workflow_app_logs"] == 2 @patch("services.retention.workflow_run.restore_archived_workflow_run.get_archive_storage") - @patch("services.retention.workflow_run.restore_archived_workflow_run.pg_insert") - @patch("services.retention.workflow_run.restore_archived_workflow_run.cast") - def test_successful_restore(self, mock_cast, mock_pg_insert, mock_get_storage): + def test_successful_restore(self, mock_get_storage, database: Database): """Should successfully restore from archive.""" restore = WorkflowRunRestore() @@ -670,53 +686,57 @@ class TestRestoreFromRun: mock_storage.get_object.return_value = archive_data mock_get_storage.return_value = mock_storage - # Mock session with context manager support - mock_session = Mock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - - def session_maker(): - return mock_session - - # Mock database execution to return integer counts - mock_result_workflow_runs = Mock() - mock_result_workflow_runs.rowcount = 1 - mock_result_app_logs = Mock() - mock_result_app_logs.rowcount = 2 - - # Configure session.execute to return different results based on the table - def mock_execute(stmt): - if "workflow_runs" in str(stmt): - return mock_result_workflow_runs - else: - return mock_result_app_logs - - mock_session.execute.side_effect = mock_execute - mock_cast.return_value = mock_result_workflow_runs - - # Mock insert statement - mock_stmt = Mock() - mock_stmt.on_conflict_do_nothing.return_value = mock_stmt - mock_pg_insert.return_value = mock_stmt - - run = WorkflowRunRestoreTestDataFactory.create_workflow_run_mock() + run = WorkflowRunRestoreTestDataFactory.create_workflow_run() + archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log() + database.session.add(archive_log) + database.session.commit() # Mock repository methods with patch.object(restore, "_get_workflow_run_repo") as mock_get_repo: mock_repo = Mock() + mock_repo.delete_archive_log_by_run_id.side_effect = lambda session, run_id: session.execute( + delete(WorkflowArchiveLog).where(WorkflowArchiveLog.workflow_run_id == run_id) + ) mock_get_repo.return_value = mock_repo with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - result = restore._restore_from_run(run, session_maker=session_maker) + result = restore._restore_from_run(run, session_maker=database.session_maker) assert result.success is True assert result.restored_counts["workflow_runs"] == 1 - assert result.restored_counts["workflow_app_logs"] >= 1 # Just check it's restored - mock_session.commit.assert_called_once() - mock_repo.delete_archive_log_by_run_id.assert_called_once_with(mock_session, run.id) + assert result.restored_counts["workflow_app_logs"] == 2 + database.session.expire_all() + assert database.session.scalar(select(func.count(WorkflowRun.id))) == 1 + assert database.session.scalar(select(func.count(WorkflowAppLog.id))) == 2 + assert database.session.scalar(select(func.count(WorkflowArchiveLog.id))) == 0 @patch("services.retention.workflow_run.restore_archived_workflow_run.get_archive_storage") - def test_invalid_archive_bundle(self, mock_get_storage): + def test_insert_failure_rolls_back_all_tables(self, mock_get_storage, database: Database): + """A later table failure must roll back earlier restored rows.""" + restore = WorkflowRunRestore() + mock_storage = Mock() + mock_storage.get_object.return_value = WorkflowRunRestoreTestDataFactory.create_archive_zip_mock() + mock_get_storage.return_value = mock_storage + run = WorkflowRunRestoreTestDataFactory.create_workflow_run() + + def fail_app_log_insert(_connection, _cursor, statement, _parameters, _context, _executemany): + if statement.startswith("INSERT INTO workflow_app_logs"): + raise RuntimeError("forced app-log insert failure") + + event.listen(database.engine, "before_cursor_execute", fail_app_log_insert) + try: + with patch("services.retention.workflow_run.restore_archived_workflow_run.click"): + result = restore._restore_from_run(run, session_maker=database.session_maker) + finally: + event.remove(database.engine, "before_cursor_execute", fail_app_log_insert) + + assert result.success is False + assert result.error == "forced app-log insert failure" + assert database.session.scalar(select(func.count(WorkflowRun.id))) == 0 + assert database.session.scalar(select(func.count(WorkflowAppLog.id))) == 0 + + @patch("services.retention.workflow_run.restore_archived_workflow_run.get_archive_storage") + def test_invalid_archive_bundle(self, mock_get_storage, database: Database): """Should handle invalid archive bundle.""" restore = WorkflowRunRestore() @@ -725,22 +745,17 @@ class TestRestoreFromRun: mock_storage.get_object.return_value = b"invalid zip data" mock_get_storage.return_value = mock_storage - run = WorkflowRunRestoreTestDataFactory.create_workflow_run_mock() - - # Create proper mock session - mock_session = Mock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) + run = WorkflowRunRestoreTestDataFactory.create_workflow_run() with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - result = restore._restore_from_run(run, session_maker=lambda: mock_session) + result = restore._restore_from_run(run, session_maker=database.session_maker) assert result.success is False # The error message comes from zipfile.BadZipFile which says "File is not a zip file" assert "File is not a zip file" in result.error @patch("services.retention.workflow_run.restore_archived_workflow_run.get_archive_storage") - def test_workflow_archive_log_input(self, mock_get_storage): + def test_workflow_archive_log_input(self, mock_get_storage, database: Database): """Should handle WorkflowArchiveLog input correctly.""" restore = WorkflowRunRestore(dry_run=True) @@ -750,14 +765,11 @@ class TestRestoreFromRun: mock_storage.get_object.return_value = archive_data mock_get_storage.return_value = mock_storage - archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log_mock() + archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log() + database.session.add(archive_log) + database.session.commit() - # Create proper mock session - mock_session = Mock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - - result = restore._restore_from_run(archive_log, session_maker=lambda: mock_session) + result = restore._restore_from_run(archive_log, session_maker=database.session_maker) assert result.success is True assert result.run_id == archive_log.workflow_run_id @@ -772,39 +784,29 @@ class TestRestoreFromRun: class TestRestoreBatch: """Tests for WorkflowRunRestore.restore_batch method.""" - @patch("services.retention.workflow_run.restore_archived_workflow_run.sessionmaker") - def test_empty_tenant_ids_returns_empty(self, mock_sessionmaker): + def test_empty_tenant_ids_returns_empty(self, database: Database): """Should return empty list when tenant_ids is empty list.""" restore = WorkflowRunRestore() - # Mock db.engine to avoid SQLAlchemy issues - with patch("services.retention.workflow_run.restore_archived_workflow_run.db") as mock_db: - mock_db.engine = Mock() - result = restore.restore_batch( - tenant_ids=[], - start_date=datetime(2024, 1, 1), - end_date=datetime(2024, 1, 2), - ) + result = restore.restore_batch( + tenant_ids=[], + start_date=datetime(2024, 1, 1), + end_date=datetime(2024, 1, 2), + ) assert result == [] @patch("services.retention.workflow_run.restore_archived_workflow_run.ThreadPoolExecutor") - def test_successful_batch_restore(self, mock_executor): + def test_successful_batch_restore(self, mock_executor, database: Database): """Should successfully restore batch of workflow runs.""" restore = WorkflowRunRestore(workers=2) - # Mock session that supports context manager protocol - mock_session = Mock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - - # Mock session factory that returns context manager sessions - mock_session_factory = Mock(return_value=mock_session) - # Mock repository and archive logs mock_repo = Mock() - archive_log1 = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log_mock("run-1") - archive_log2 = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log_mock("run-2") + archive_log1 = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log("run-1") + archive_log2 = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log("run-2") + database.session.add_all([archive_log1, archive_log2]) + database.session.commit() mock_repo.get_archived_logs_by_time_range.return_value = [archive_log1, archive_log2] # Mock restore results @@ -821,38 +823,25 @@ class TestRestoreBatch: with patch.object(restore, "_get_workflow_run_repo", return_value=mock_repo): with patch.object(restore, "_restore_from_run", side_effect=[result1, result2]): with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - # Mock sessionmaker and db.engine to avoid SQLAlchemy issues - with patch( - "services.retention.workflow_run.restore_archived_workflow_run.sessionmaker" - ) as mock_sessionmaker: - mock_sessionmaker.return_value = mock_session_factory - with patch("services.retention.workflow_run.restore_archived_workflow_run.db") as mock_db: - mock_db.engine = Mock() - results = restore.restore_batch( - tenant_ids=["tenant-1"], - start_date=datetime(2024, 1, 1), - end_date=datetime(2024, 1, 2), - ) + results = restore.restore_batch( + tenant_ids=["tenant-1"], + start_date=datetime(2024, 1, 1), + end_date=datetime(2024, 1, 2), + ) assert len(results) == 2 assert results[0].run_id == "run-1" assert results[1].run_id == "run-2" @patch("services.retention.workflow_run.restore_archived_workflow_run.ThreadPoolExecutor") - def test_dry_run_batch_restore(self, mock_executor): + def test_dry_run_batch_restore(self, mock_executor, database: Database): """Should handle dry run mode for batch restore.""" restore = WorkflowRunRestore(dry_run=True) - # Mock session that supports context manager protocol - mock_session = Mock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - - # Mock session factory that returns context manager sessions - mock_session_factory = Mock(return_value=mock_session) - mock_repo = Mock() - archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log_mock() + archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log() + database.session.add(archive_log) + database.session.commit() mock_repo.get_archived_logs_by_time_range.return_value = [archive_log] result = RestoreResult(run_id="run-1", tenant_id="tenant-1", success=True, restored_counts={"workflow_runs": 1}) @@ -867,18 +856,11 @@ class TestRestoreBatch: with patch.object(restore, "_get_workflow_run_repo", return_value=mock_repo): with patch.object(restore, "_restore_from_run", return_value=result): with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - # Mock sessionmaker and db.engine to avoid SQLAlchemy issues - with patch( - "services.retention.workflow_run.restore_archived_workflow_run.sessionmaker" - ) as mock_sessionmaker: - mock_sessionmaker.return_value = mock_session_factory - with patch("services.retention.workflow_run.restore_archived_workflow_run.db") as mock_db: - mock_db.engine = Mock() - results = restore.restore_batch( - tenant_ids=["tenant-1"], - start_date=datetime(2024, 1, 1), - end_date=datetime(2024, 1, 2), - ) + results = restore.restore_batch( + tenant_ids=["tenant-1"], + start_date=datetime(2024, 1, 1), + end_date=datetime(2024, 1, 2), + ) assert len(results) == 1 assert results[0].success is True @@ -907,16 +889,14 @@ class TestRestoreByRunId: assert "not found" in result.error assert result.run_id == "nonexistent-run" - @patch("services.retention.workflow_run.restore_archived_workflow_run.sessionmaker") - def test_successful_restore_by_id(self, mock_sessionmaker): + def test_successful_restore_by_id(self, database: Database): """Should successfully restore by run ID.""" restore = WorkflowRunRestore() - mock_session = Mock() - mock_sessionmaker.return_value = mock_session - mock_repo = Mock() - archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log_mock() + archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log() + database.session.add(archive_log) + database.session.commit() mock_repo.get_archived_log_by_run_id.return_value = archive_log result = RestoreResult(run_id="run-1", tenant_id="tenant-1", success=True, restored_counts={}) @@ -924,24 +904,19 @@ class TestRestoreByRunId: with patch.object(restore, "_get_workflow_run_repo", return_value=mock_repo): with patch.object(restore, "_restore_from_run", return_value=result): with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - # Mock db.engine to avoid SQLAlchemy issues - with patch("services.retention.workflow_run.restore_archived_workflow_run.db") as mock_db: - mock_db.engine = Mock() - actual_result = restore.restore_by_run_id("run-1") + actual_result = restore.restore_by_run_id("run-1") assert actual_result.success is True assert actual_result.run_id == "run-1" - @patch("services.retention.workflow_run.restore_archived_workflow_run.sessionmaker") - def test_dry_run_restore_by_id(self, mock_sessionmaker): + def test_dry_run_restore_by_id(self, database: Database): """Should handle dry run mode for restore by ID.""" restore = WorkflowRunRestore(dry_run=True) - mock_session = Mock() - mock_sessionmaker.return_value = mock_session - mock_repo = Mock() - archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log_mock() + archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log() + database.session.add(archive_log) + database.session.commit() mock_repo.get_archived_log_by_run_id.return_value = archive_log result = RestoreResult(run_id="run-1", tenant_id="tenant-1", success=True, restored_counts={"workflow_runs": 1}) @@ -949,10 +924,7 @@ class TestRestoreByRunId: with patch.object(restore, "_get_workflow_run_repo", return_value=mock_repo): with patch.object(restore, "_restore_from_run", return_value=result): with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - # Mock db.engine to avoid SQLAlchemy issues - with patch("services.retention.workflow_run.restore_archived_workflow_run.db") as mock_db: - mock_db.engine = Mock() - actual_result = restore.restore_by_run_id("run-1") + actual_result = restore.restore_by_run_id("run-1") assert actual_result.success is True assert actual_result.run_id == "run-1" @@ -1038,8 +1010,7 @@ class TestIntegration: """Integration tests combining multiple components.""" @patch("services.retention.workflow_run.restore_archived_workflow_run.get_archive_storage") - @patch("services.retention.workflow_run.restore_archived_workflow_run.ThreadPoolExecutor") - def test_full_restore_flow(self, mock_executor, mock_get_storage): + def test_full_restore_flow(self, mock_get_storage, database: Database): """Test complete restore flow with all components.""" restore = WorkflowRunRestore(workers=1) @@ -1059,7 +1030,7 @@ class TestIntegration: "app_id": "app-123", "workflow_id": "workflow-123", "type": "workflow", - "triggered_from": "app", + "triggered_from": "app-run", "version": "1", "status": "succeeded", "created_by_role": "account", @@ -1072,48 +1043,20 @@ class TestIntegration: mock_storage.get_object.return_value = archive_data mock_get_storage.return_value = mock_storage - # Mock session that supports context manager protocol - mock_session = Mock() - mock_session.__enter__ = Mock(return_value=mock_session) - mock_session.__exit__ = Mock(return_value=None) - - # Mock session factory that returns context manager sessions - mock_session_factory = Mock(return_value=mock_session) - - mock_result = Mock() - mock_result.rowcount = 1 - mock_session.execute.return_value = mock_result - # Mock repository mock_repo = Mock() - archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log_mock() + archive_log = WorkflowRunRestoreTestDataFactory.create_workflow_archive_log() + database.session.add(archive_log) + database.session.commit() mock_repo.get_archived_log_by_run_id.return_value = archive_log - - # Mock ThreadPoolExecutor (not actually used in restore_by_run_id but needed for patch) - mock_executor_instance = Mock() - mock_executor_instance.__enter__ = Mock(return_value=mock_executor_instance) - mock_executor_instance.__exit__ = Mock(return_value=None) - mock_executor_instance.map = Mock(return_value=[]) - mock_executor.return_value = mock_executor_instance + mock_repo.delete_archive_log_by_run_id.side_effect = lambda session, run_id: session.execute( + delete(WorkflowArchiveLog).where(WorkflowArchiveLog.workflow_run_id == run_id) + ) with patch.object(restore, "_get_workflow_run_repo", return_value=mock_repo): - with patch("services.retention.workflow_run.restore_archived_workflow_run.pg_insert") as mock_insert: - mock_stmt = Mock() - mock_stmt.on_conflict_do_nothing.return_value = mock_stmt - mock_insert.return_value = mock_stmt - - with patch("services.retention.workflow_run.restore_archived_workflow_run.cast") as mock_cast: - mock_cast.return_value = mock_result - - with patch("services.retention.workflow_run.restore_archived_workflow_run.click") as mock_click: - # Mock sessionmaker and db.engine to avoid SQLAlchemy issues - with patch( - "services.retention.workflow_run.restore_archived_workflow_run.sessionmaker" - ) as mock_sessionmaker: - mock_sessionmaker.return_value = mock_session_factory - with patch("services.retention.workflow_run.restore_archived_workflow_run.db") as mock_db: - mock_db.engine = Mock() - result = restore.restore_by_run_id("run-123") + with patch("services.retention.workflow_run.restore_archived_workflow_run.click"): + result = restore.restore_by_run_id("run-123") assert result.success is True assert result.restored_counts.get("workflow_runs") == 1 + assert database.session.scalar(select(func.count(WorkflowRun.id))) == 1 diff --git a/api/tests/unit_tests/services/test_account_service.py b/api/tests/unit_tests/services/test_account_service.py index ad40dab358c..1e01951a49c 100644 --- a/api/tests/unit_tests/services/test_account_service.py +++ b/api/tests/unit_tests/services/test_account_service.py @@ -1,25 +1,26 @@ import json -from collections.abc import Iterator +from collections.abc import Iterator, Sequence from datetime import datetime, timedelta from unittest.mock import MagicMock, patch +from uuid import UUID import pytest from sqlalchemy import event, select -from sqlalchemy.orm import Session +from sqlalchemy.engine import Connection +from sqlalchemy.engine.interfaces import DBAPICursor, ExecutionContext +from sqlalchemy.orm import Session, sessionmaker from configs import dify_config +from enums import DeploymentEdition from models.account import ( Account, - AccountIntegrate, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, - TenantPluginAutoUpgradeStrategy, TenantStatus, ) -from models.dataset import Dataset -from models.model import App, DifySetup +from models.model import DifySetup from services.account_service import AccountService, RegisterService, TenantService from services.enterprise.rbac_service import MembersInRole, Paginated from services.errors.account import ( @@ -31,6 +32,8 @@ from services.errors.account import ( NoPermissionError, ) +type _MockDependencies = dict[str, MagicMock] + class TestAccountAssociatedDataFactory: """Factory class for creating test data and mock objects for account service tests.""" @@ -40,65 +43,30 @@ class TestAccountAssociatedDataFactory: account_id: str = "user-123", email: str = "test@example.com", name: str = "Test User", - status: str = "active", + status: AccountStatus | str = AccountStatus.ACTIVE, password: str = "hashed_password", password_salt: str = "salt", interface_language: str = "en-US", interface_theme: str = "light", timezone: str = "UTC", - **kwargs, - ) -> MagicMock: - """Create a mock account with specified attributes.""" - account = MagicMock(spec=Account) + ) -> Account: + """Create an account with specified attributes.""" + account = Account( + name=name, + email=email, + password=password, + password_salt=password_salt, + interface_language=interface_language, + interface_theme=interface_theme, + timezone=timezone, + status=AccountStatus(status), + initialized_at=None, + ) account.id = account_id - account.email = email - account.name = name - account.status = status - account.password = password - account.password_salt = password_salt - account.interface_language = interface_language - account.interface_theme = interface_theme - account.timezone = timezone # Set last_active_at to a datetime object that's older than 10 minutes account.last_active_at = datetime.now() - timedelta(minutes=15) - account.initialized_at = None - for key, value in kwargs.items(): - setattr(account, key, value) return account - @staticmethod - def create_tenant_join_mock( - tenant_id: str = "tenant-456", - account_id: str = "user-123", - current: bool = True, - role: str = "normal", - **kwargs, - ) -> MagicMock: - """Create a mock tenant account join record.""" - tenant_join = MagicMock() - tenant_join.tenant_id = tenant_id - tenant_join.account_id = account_id - tenant_join.current = current - tenant_join.role = role - tenant_join.last_opened_at = kwargs.pop("last_opened_at", None) - for key, value in kwargs.items(): - setattr(tenant_join, key, value) - return tenant_join - - @staticmethod - def create_feature_service_mock(allow_register: bool = True): - """Create a mock feature service.""" - mock_service = MagicMock() - mock_service.get_system_features.return_value.is_allow_register = allow_register - return mock_service - - @staticmethod - def create_billing_service_mock(email_frozen: bool = False): - """Create a mock billing service.""" - mock_service = MagicMock() - mock_service.is_email_in_freeze.return_value = email_frozen - return mock_service - class TestAccountService: """ @@ -114,23 +82,7 @@ class TestAccountService: """ @pytest.fixture - def sqlite_session(self, sqlite_engine) -> Iterator[Session]: - """SQLite session with the account/workspace tables these service tests touch.""" - tables = [ - model.metadata.tables[model.__tablename__] - for model in ( - Account, - Tenant, - TenantAccountJoin, - TenantPluginAutoUpgradeStrategy, - ) - ] - Account.metadata.create_all(sqlite_engine, tables=tables) - with Session(sqlite_engine, expire_on_commit=False) as session: - yield session - - @pytest.fixture - def mock_password_dependencies(self): + def mock_password_dependencies(self) -> Iterator[_MockDependencies]: """Mock setup for password-related functions.""" with ( patch("services.account_service.compare_password") as mock_compare_password, @@ -144,7 +96,7 @@ class TestAccountService: } @pytest.fixture - def mock_external_service_dependencies(self): + def mock_external_service_dependencies(self) -> Iterator[_MockDependencies]: """Mock setup for external service dependencies.""" with ( patch("services.account_service.FeatureService") as mock_feature_service, @@ -157,14 +109,9 @@ class TestAccountService: "passport_service": mock_passport_service, } - def _assert_exception_raised(self, exception_type, callable_func, *args, **kwargs): - """Helper method to verify that specific exception is raised.""" - with pytest.raises(exception_type): - callable_func(*args, **kwargs) - # ==================== Authentication Tests ==================== - def test_authenticate_success(self, sqlite_session: Session, mock_password_dependencies): + def test_authenticate_success(self, sqlite_session: Session, mock_password_dependencies: _MockDependencies) -> None: """Test successful authentication with correct email and password.""" account = Account( name="Test User", @@ -181,17 +128,12 @@ class TestAccountService: assert result is account - def test_authenticate_account_not_found(self, sqlite_session: Session): + def test_authenticate_account_not_found(self, sqlite_session: Session) -> None: """Test authentication when account does not exist.""" - self._assert_exception_raised( - AccountPasswordError, - AccountService.authenticate, - "notfound@example.com", - "password", - session=sqlite_session, - ) + with pytest.raises(AccountPasswordError): + AccountService.authenticate("notfound@example.com", "password", session=sqlite_session) - def test_authenticate_account_banned(self, sqlite_session: Session): + def test_authenticate_account_banned(self, sqlite_session: Session) -> None: """Test authentication when account is banned.""" account = Account( name="Banned User", @@ -203,15 +145,12 @@ class TestAccountService: sqlite_session.add(account) sqlite_session.commit() - self._assert_exception_raised( - AccountLoginError, - AccountService.authenticate, - "banned@example.com", - "password", - session=sqlite_session, - ) + with pytest.raises(AccountLoginError): + AccountService.authenticate("banned@example.com", "password", session=sqlite_session) - def test_authenticate_password_error(self, sqlite_session: Session, mock_password_dependencies): + def test_authenticate_password_error( + self, sqlite_session: Session, mock_password_dependencies: _MockDependencies + ) -> None: """Test authentication with wrong password.""" account = Account( name="Test User", @@ -224,39 +163,46 @@ class TestAccountService: mock_password_dependencies["compare_password"].return_value = False - self._assert_exception_raised( - AccountPasswordError, - AccountService.authenticate, - "test@example.com", - "wrongpassword", - session=sqlite_session, - ) + with pytest.raises(AccountPasswordError): + AccountService.authenticate("test@example.com", "wrongpassword", session=sqlite_session) - def test_authenticate_pending_account_activates(self, sqlite_session: Session, mock_password_dependencies): + def test_authenticate_pending_account_activates( + self, + sqlite_session_factory: sessionmaker[Session], + mock_password_dependencies: _MockDependencies, + ) -> None: """Test authentication for a pending account, which should activate on login.""" - account = Account( - name="Pending User", - email="pending@example.com", - password="hashed_password", - password_salt="salt", - status=AccountStatus.PENDING, - ) - sqlite_session.add(account) - sqlite_session.commit() + with sqlite_session_factory() as service_session: + account = Account( + name="Pending User", + email="pending@example.com", + password="hashed_password", + password_salt="salt", + status=AccountStatus.PENDING, + ) + service_session.add(account) + service_session.commit() + account_id = account.id - mock_password_dependencies["compare_password"].return_value = True + mock_password_dependencies["compare_password"].return_value = True - result = AccountService.authenticate("pending@example.com", "password", session=sqlite_session) + result = AccountService.authenticate("pending@example.com", "password", session=service_session) + assert result.id == account_id - assert result is account - assert account.status == AccountStatus.ACTIVE - assert account.initialized_at is not None + with sqlite_session_factory() as assertion_session: + persisted_account = assertion_session.get(Account, account_id) + assert persisted_account is not None + assert persisted_account.status == AccountStatus.ACTIVE + assert persisted_account.initialized_at is not None # ==================== Account Creation Tests ==================== def test_create_account_success( - self, sqlite_session: Session, mock_password_dependencies, mock_external_service_dependencies - ): + self, + sqlite_session_factory: sessionmaker[Session], + mock_password_dependencies: _MockDependencies, + mock_external_service_dependencies: _MockDependencies, + ) -> None: """Test successful account creation with all required parameters.""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True @@ -264,142 +210,186 @@ class TestAccountService: mock_password_dependencies["hash_password"].return_value = b"hashed_password" # Execute test - result = AccountService.create_account( - email="test@example.com", - name="Test User", - interface_language="en-US", - password="password123", - interface_theme="light", - session=sqlite_session, - ) + with sqlite_session_factory() as service_session: + result = AccountService.create_account( + email="test@example.com", + name="Test User", + interface_language="en-US", + password="password123", + interface_theme="light", + session=service_session, + ) + account_id = result.id - # Verify results - assert result.email == "test@example.com" - assert result.name == "Test User" - assert result.interface_language == "en-US" - assert result.interface_theme == "light" - assert result.password is not None - assert result.password_salt is not None - assert result.timezone == "America/New_York" + assert result.email == "test@example.com" + assert result.name == "Test User" + assert result.interface_language == "en-US" + assert result.interface_theme == "light" + assert result.password is not None + assert result.password_salt is not None + assert result.timezone == "America/New_York" - persisted_account = sqlite_session.scalar(select(Account).where(Account.email == "test@example.com")) - assert persisted_account is result + with sqlite_session_factory() as assertion_session: + persisted_account = assertion_session.get(Account, account_id) + assert persisted_account is not None + assert persisted_account.email == "test@example.com" + assert persisted_account.name == "Test User" + assert persisted_account.interface_language == "en-US" + assert persisted_account.interface_theme == "light" + assert persisted_account.password is not None + assert persisted_account.password_salt is not None + assert persisted_account.timezone == "America/New_York" def test_create_account_uses_explicit_timezone( - self, sqlite_session: Session, mock_password_dependencies, mock_external_service_dependencies - ): + self, + sqlite_session_factory: sessionmaker[Session], + mock_password_dependencies: _MockDependencies, + mock_external_service_dependencies: _MockDependencies, + ) -> None: """Test account creation prefers explicit browser timezone.""" mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False mock_password_dependencies["hash_password"].return_value = b"hashed_password" - result = AccountService.create_account( - email="test@example.com", - name="Test User", - interface_language="en-US", - password="password123", - timezone="Asia/Shanghai", - session=sqlite_session, - ) + with sqlite_session_factory() as service_session: + result = AccountService.create_account( + email="test@example.com", + name="Test User", + interface_language="en-US", + password="password123", + timezone="Asia/Shanghai", + session=service_session, + ) + account_id = result.id + assert result.timezone == "Asia/Shanghai" - assert result.timezone == "Asia/Shanghai" - persisted_account = sqlite_session.scalar(select(Account).where(Account.email == "test@example.com")) - assert persisted_account is result - assert persisted_account.timezone == "Asia/Shanghai" + with sqlite_session_factory() as assertion_session: + persisted_account = assertion_session.get(Account, account_id) + assert persisted_account is not None + assert persisted_account.timezone == "Asia/Shanghai" - def test_create_account_registration_disabled(self, sqlite_session: Session, mock_external_service_dependencies): + def test_create_account_registration_disabled( + self, unbound_session: Session, mock_external_service_dependencies: _MockDependencies + ) -> None: """Test account creation when registration is disabled.""" + from controllers.console.error import AccountNotFound + # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = False # Execute test and verify exception - self._assert_exception_raised( - Exception, # AccountNotFound - AccountService.create_account, - email="test@example.com", - name="Test User", - interface_language="en-US", - session=sqlite_session, - ) + with pytest.raises(AccountNotFound): + AccountService.create_account( + email="test@example.com", + name="Test User", + interface_language="en-US", + session=unbound_session, + ) - def test_create_account_email_frozen(self, sqlite_session: Session, mock_external_service_dependencies): + def test_create_account_email_frozen( + self, unbound_session: Session, mock_external_service_dependencies: _MockDependencies + ) -> None: """Test account creation with frozen email address.""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = True - with patch("services.account_service.dify_config.BILLING_ENABLED", True): - self._assert_exception_raised( - AccountRegisterError, - AccountService.create_account, - email="frozen@example.com", - name="Test User", - interface_language="en-US", - session=sqlite_session, - ) + with patch("services.account_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD): + with pytest.raises(AccountRegisterError): + AccountService.create_account( + email="frozen@example.com", + name="Test User", + interface_language="en-US", + session=unbound_session, + ) - def test_create_account_without_password(self, sqlite_session: Session, mock_external_service_dependencies): + def test_create_account_without_password( + self, + sqlite_session_factory: sessionmaker[Session], + mock_external_service_dependencies: _MockDependencies, + ) -> None: """Test account creation without password (for invite-based registration).""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False # Execute test - result = AccountService.create_account( - email="test@example.com", - name="Test User", - interface_language="zh-CN", - password=None, - interface_theme="dark", - session=sqlite_session, - ) + with sqlite_session_factory() as service_session: + result = AccountService.create_account( + email="test@example.com", + name="Test User", + interface_language="zh-CN", + password=None, + interface_theme="dark", + session=service_session, + ) + account_id = result.id - # Verify results - assert result.email == "test@example.com" - assert result.name == "Test User" - assert result.interface_language == "zh-CN" - assert result.interface_theme == "dark" - assert result.password is None - assert result.password_salt is None - assert result.timezone is not None + assert result.email == "test@example.com" + assert result.name == "Test User" + assert result.interface_language == "zh-CN" + assert result.interface_theme == "dark" + assert result.password is None + assert result.password_salt is None + assert result.timezone is not None - persisted_account = sqlite_session.scalar(select(Account).where(Account.email == "test@example.com")) - assert persisted_account is result + with sqlite_session_factory() as assertion_session: + persisted_account = assertion_session.get(Account, account_id) + assert persisted_account is not None + assert persisted_account.email == "test@example.com" + assert persisted_account.name == "Test User" + assert persisted_account.interface_language == "zh-CN" + assert persisted_account.interface_theme == "dark" + assert persisted_account.password is None + assert persisted_account.password_salt is None + assert persisted_account.timezone is not None # ==================== Password Management Tests ==================== - def test_update_account_password_success(self, sqlite_session: Session, mock_password_dependencies): + def test_update_account_password_success( + self, + sqlite_session_factory: sessionmaker[Session], + mock_password_dependencies: _MockDependencies, + ) -> None: """Test successful password update with correct current password and valid new password.""" - account = Account( - name="Test User", - email="test@example.com", - password="hashed_password", - password_salt="salt", - ) - sqlite_session.add(account) - sqlite_session.commit() - mock_password_dependencies["compare_password"].return_value = True - mock_password_dependencies["valid_password"].return_value = None - mock_password_dependencies["hash_password"].return_value = b"new_hashed_password" + with sqlite_session_factory() as service_session: + account = Account( + name="Test User", + email="test@example.com", + password="hashed_password", + password_salt="salt", + ) + service_session.add(account) + service_session.commit() + account_id = account.id - result = AccountService.update_account_password( - account, "old_password", "new_password123", session=sqlite_session - ) + mock_password_dependencies["compare_password"].return_value = True + mock_password_dependencies["valid_password"].return_value = None + mock_password_dependencies["hash_password"].return_value = b"new_hashed_password" - assert result is account - assert account.password is not None - assert account.password_salt is not None + result = AccountService.update_account_password( + account, + "old_password", + "new_password123", + session=service_session, + ) + assert result is account - # Verify password validation was called mock_password_dependencies["compare_password"].assert_called_once_with( "old_password", "hashed_password", "salt" ) mock_password_dependencies["valid_password"].assert_called_once_with("new_password123") - # Verify database operations + with sqlite_session_factory() as assertion_session: + persisted_account = assertion_session.get(Account, account_id) + assert persisted_account is not None + assert persisted_account.password is not None + assert persisted_account.password != "hashed_password" + assert persisted_account.password_salt is not None + assert persisted_account.password_salt != "salt" def test_update_account_password_current_password_incorrect( - self, sqlite_session: Session, mock_password_dependencies - ): + self, unbound_session: Session, mock_password_dependencies: _MockDependencies + ) -> None: """Test password update with incorrect current password.""" # Setup test data mock_account = Account( @@ -411,21 +401,22 @@ class TestAccountService: mock_password_dependencies["compare_password"].return_value = False # Execute test and verify exception - self._assert_exception_raised( - CurrentPasswordIncorrectError, - AccountService.update_account_password, - mock_account, - "wrong_password", - "new_password123", - session=sqlite_session, - ) + with pytest.raises(CurrentPasswordIncorrectError): + AccountService.update_account_password( + mock_account, + "wrong_password", + "new_password123", + session=unbound_session, + ) # Verify password comparison was called mock_password_dependencies["compare_password"].assert_called_once_with( "wrong_password", "hashed_password", "salt" ) - def test_update_account_password_invalid_new_password(self, sqlite_session: Session, mock_password_dependencies): + def test_update_account_password_invalid_new_password( + self, unbound_session: Session, mock_password_dependencies: _MockDependencies + ) -> None: """Test password update with invalid new password.""" # Setup test data mock_account = Account( @@ -438,21 +429,20 @@ class TestAccountService: mock_password_dependencies["valid_password"].side_effect = ValueError("Password too short") # Execute test and verify exception - self._assert_exception_raised( - ValueError, - AccountService.update_account_password, - mock_account, - "old_password", - "short", - session=sqlite_session, - ) + with pytest.raises(ValueError): + AccountService.update_account_password( + mock_account, + "old_password", + "short", + session=unbound_session, + ) # Verify password validation was called mock_password_dependencies["valid_password"].assert_called_once_with("short") # ==================== User Loading Tests ==================== - def test_load_user_success(self, sqlite_session: Session): + def test_load_user_success(self, sqlite_session: Session) -> None: """Test successful user loading with current tenant.""" account = Account(name="Test User", email="test@example.com") tenant = Tenant(name="Test Workspace") @@ -467,68 +457,115 @@ class TestAccountService: sqlite_session.add(tenant_join) sqlite_session.commit() - with ( - patch.object(Account, "set_tenant_id_with_session") as mock_set_tenant_id, - patch.object(AccountService, "_refresh_account_last_active") as mock_refresh_last_active, - ): + with patch.object(AccountService, "_refresh_account_last_active") as mock_refresh_last_active: result = AccountService.load_user(account.id, sqlite_session) assert result is account - mock_set_tenant_id.assert_called_once_with(tenant.id, session=sqlite_session) + assert result.current_tenant_id == tenant.id mock_refresh_last_active.assert_called_once_with(account, sqlite_session) - def test_load_user_not_found(self, sqlite_session: Session): + def test_load_user_not_found(self, sqlite_session: Session) -> None: """Test user loading when user does not exist.""" result = AccountService.load_user("non-existent-user", sqlite_session) assert result is None - def test_load_user_banned(self, sqlite_session: Session): + def test_load_user_banned(self, sqlite_session: Session) -> None: """Test user loading when user is banned.""" + from werkzeug.exceptions import Unauthorized + account = Account(name="Banned User", email="banned@example.com", status=AccountStatus.BANNED) sqlite_session.add(account) sqlite_session.commit() - self._assert_exception_raised( - Exception, # Unauthorized - AccountService.load_user, - account.id, - sqlite_session, - ) + with pytest.raises(Unauthorized): + AccountService.load_user(account.id, sqlite_session) - def test_load_user_no_current_tenant(self, sqlite_session: Session): + def test_load_user_no_current_tenant(self, sqlite_session_factory: sessionmaker[Session]) -> None: """Test user loading when user has no current tenant but has available tenants.""" + with sqlite_session_factory() as service_session: + account = Account(name="Test User", email="test@example.com") + tenant = Tenant(name="Test Workspace") + service_session.add_all([account, tenant]) + service_session.flush() + available_tenant_join = TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + role=TenantAccountRole.NORMAL, + current=False, + ) + service_session.add(available_tenant_join) + service_session.commit() + account_id = account.id + tenant_id = tenant.id + tenant_join_id = available_tenant_join.id + + mock_now = datetime(2026, 6, 5, 11, 0, 0) + with ( + patch.object(Account, "set_tenant_id_with_session") as mock_set_tenant_id, + patch("services.account_service.naive_utc_now", return_value=mock_now), + patch.object(AccountService, "_refresh_account_last_active") as mock_refresh_last_active, + ): + result = AccountService.load_user(account_id, service_session) + assert result is not None + mock_set_tenant_id.assert_called_once_with(tenant_id, session=service_session) + mock_refresh_last_active.assert_called_once_with(result, service_session) + + with sqlite_session_factory() as assertion_session: + persisted_tenant_join = assertion_session.get(TenantAccountJoin, tenant_join_id) + assert persisted_tenant_join is not None + assert persisted_tenant_join.current is True + assert persisted_tenant_join.last_opened_at == mock_now + + def test_load_user_switches_from_archived_current_tenant(self, sqlite_session: Session) -> None: account = Account(name="Test User", email="test@example.com") - tenant = Tenant(name="Test Workspace") - sqlite_session.add_all([account, tenant]) + archived_tenant = Tenant(name="Archived Workspace", status=TenantStatus.ARCHIVE) + available_tenant = Tenant(name="Available Workspace") + sqlite_session.add_all([account, archived_tenant, available_tenant]) sqlite_session.flush() - available_tenant_join = TenantAccountJoin( - tenant_id=tenant.id, + archived_join = TenantAccountJoin( + tenant_id=archived_tenant.id, + account_id=account.id, + role=TenantAccountRole.NORMAL, + current=True, + ) + available_join = TenantAccountJoin( + tenant_id=available_tenant.id, account_id=account.id, role=TenantAccountRole.NORMAL, current=False, ) - sqlite_session.add(available_tenant_join) + sqlite_session.add_all([archived_join, available_join]) sqlite_session.commit() - with ( - patch.object(Account, "set_tenant_id_with_session") as mock_set_tenant_id, - patch("services.account_service.naive_utc_now") as mock_naive_utc_now, - patch.object(AccountService, "_refresh_account_last_active") as mock_refresh_last_active, - ): - mock_now = datetime.now() - mock_naive_utc_now.return_value = mock_now - + with patch.object(AccountService, "_refresh_account_last_active"): result = AccountService.load_user(account.id, sqlite_session) - assert result is account - assert available_tenant_join.current is True - assert available_tenant_join.last_opened_at == mock_now - mock_set_tenant_id.assert_called_once_with(tenant.id, session=sqlite_session) + assert result is account + assert result.current_tenant_id == available_tenant.id + assert archived_join.current is False + assert available_join.current is True - mock_refresh_last_active.assert_called_once_with(account, sqlite_session) + def test_load_user_returns_none_without_normal_tenant(self, sqlite_session: Session) -> None: + account = Account(name="Test User", email="test@example.com") + archived_tenant = Tenant(name="Archived Workspace", status=TenantStatus.ARCHIVE) + sqlite_session.add_all([account, archived_tenant]) + sqlite_session.flush() + archived_join = TenantAccountJoin( + tenant_id=archived_tenant.id, + account_id=account.id, + role=TenantAccountRole.NORMAL, + current=True, + ) + sqlite_session.add(archived_join) + sqlite_session.commit() - def test_load_user_keeps_tenant_accessible_with_expiring_session(self, sqlite_session: Session): + result = AccountService.load_user(account.id, sqlite_session) + + assert result is None + assert archived_join.current is False + + def test_load_user_keeps_tenant_accessible_with_expiring_session(self, sqlite_session: Session) -> None: account = Account(name="Test User", email="test@example.com") tenant = Tenant(name="Test Workspace") sqlite_session.add_all([account, tenant]) @@ -552,7 +589,7 @@ class TestAccountService: assert result is not None assert result.current_tenant_id == tenant_id - def test_load_user_no_tenants(self, sqlite_session: Session): + def test_load_user_no_tenants(self, sqlite_session: Session) -> None: """Test user loading when user has no tenants at all.""" account = Account(name="Test User", email="test@example.com") sqlite_session.add(account) @@ -562,78 +599,93 @@ class TestAccountService: assert result is None - def test_refresh_account_last_active_uses_redis_gate_and_conditional_update(self, sqlite_session: Session): + def test_refresh_account_last_active_uses_redis_gate_and_conditional_update( + self, sqlite_session_factory: sessionmaker[Session] + ) -> None: """Test last-active refresh is gated in Redis and conditionally written to DB.""" now = datetime(2026, 6, 2, 2, 45, 49) - account = Account(name="Test User", email="test@example.com") - sqlite_session.add(account) - sqlite_session.commit() - account.last_active_at = now - timedelta(minutes=15) - sqlite_session.commit() + with sqlite_session_factory() as service_session: + account = Account(name="Test User", email="test@example.com") + account.last_active_at = now - timedelta(minutes=15) + service_session.add(account) + service_session.commit() + account_id = account.id - with ( - patch("services.account_service.naive_utc_now", return_value=now), - patch("services.account_service.redis_client") as mock_redis_client, - ): - mock_redis_client.set.return_value = True + with ( + patch("services.account_service.naive_utc_now", return_value=now), + patch("services.account_service.redis_client") as mock_redis_client, + ): + mock_redis_client.set.return_value = True - AccountService._refresh_account_last_active(account, sqlite_session) + AccountService._refresh_account_last_active(account, service_session) mock_redis_client.set.assert_called_once_with( - f"account_last_active_refresh:{account.id}", + f"account_last_active_refresh:{account_id}", 1, ex=600, nx=True, ) - sqlite_session.refresh(account) - assert account.last_active_at == now + with sqlite_session_factory() as assertion_session: + persisted_account = assertion_session.get(Account, account_id) + assert persisted_account is not None + assert persisted_account.last_active_at == now - def test_refresh_account_last_active_skips_db_when_redis_gate_exists(self, sqlite_session: Session): + def test_refresh_account_last_active_skips_db_when_redis_gate_exists( + self, sqlite_session_factory: sessionmaker[Session] + ) -> None: """Test concurrent refresh attempts do not enqueue duplicate DB updates.""" now = datetime(2026, 6, 2, 2, 45, 49) original_last_active_at = now - timedelta(minutes=15) - account = Account(name="Test User", email="test@example.com") - sqlite_session.add(account) - sqlite_session.commit() - account.last_active_at = original_last_active_at - sqlite_session.commit() + with sqlite_session_factory() as service_session: + account = Account(name="Test User", email="test@example.com") + account.last_active_at = original_last_active_at + service_session.add(account) + service_session.commit() + account_id = account.id - with ( - patch("services.account_service.naive_utc_now", return_value=now), - patch("services.account_service.redis_client") as mock_redis_client, - ): - mock_redis_client.set.return_value = None + with ( + patch("services.account_service.naive_utc_now", return_value=now), + patch("services.account_service.redis_client") as mock_redis_client, + ): + mock_redis_client.set.return_value = None - AccountService._refresh_account_last_active(account, sqlite_session) + AccountService._refresh_account_last_active(account, service_session) mock_redis_client.set.assert_called_once_with( - f"account_last_active_refresh:{account.id}", + f"account_last_active_refresh:{account_id}", 1, ex=600, nx=True, ) - sqlite_session.refresh(account) - assert account.last_active_at == original_last_active_at + with sqlite_session_factory() as assertion_session: + persisted_account = assertion_session.get(Account, account_id) + assert persisted_account is not None + assert persisted_account.last_active_at == original_last_active_at - def test_refresh_account_last_active_skips_recent_account(self, sqlite_session: Session): + def test_refresh_account_last_active_skips_recent_account( + self, sqlite_session_factory: sessionmaker[Session] + ) -> None: """Test recent activity does not touch Redis or DB.""" now = datetime(2026, 6, 2, 2, 45, 49) original_last_active_at = now - timedelta(minutes=5) - account = Account(name="Test User", email="test@example.com") - sqlite_session.add(account) - sqlite_session.commit() - account.last_active_at = original_last_active_at - sqlite_session.commit() + with sqlite_session_factory() as service_session: + account = Account(name="Test User", email="test@example.com") + account.last_active_at = original_last_active_at + service_session.add(account) + service_session.commit() + account_id = account.id - with ( - patch("services.account_service.naive_utc_now", return_value=now), - patch("services.account_service.redis_client") as mock_redis_client, - ): - AccountService._refresh_account_last_active(account, sqlite_session) + with ( + patch("services.account_service.naive_utc_now", return_value=now), + patch("services.account_service.redis_client") as mock_redis_client, + ): + AccountService._refresh_account_last_active(account, service_session) mock_redis_client.set.assert_not_called() - sqlite_session.refresh(account) - assert account.last_active_at == original_last_active_at + with sqlite_session_factory() as assertion_session: + persisted_account = assertion_session.get(Account, account_id) + assert persisted_account is not None + assert persisted_account.last_active_at == original_last_active_at class TestTenantService: @@ -649,13 +701,13 @@ class TestTenantService: """ @pytest.fixture - def mock_rsa_dependencies(self): + def mock_rsa_dependencies(self) -> Iterator[MagicMock]: """Mock setup for RSA-related functions.""" with patch("services.account_service.generate_key_pair") as mock_generate_key_pair: yield mock_generate_key_pair @pytest.fixture - def mock_external_service_dependencies(self): + def mock_external_service_dependencies(self) -> Iterator[_MockDependencies]: """Mock setup for external service dependencies.""" with ( patch("services.account_service.FeatureService") as mock_feature_service, @@ -666,11 +718,6 @@ class TestTenantService: "billing_service": mock_billing_service, } - def _assert_exception_raised(self, exception_type, callable_func, *args, **kwargs): - """Helper method to verify that specific exception is raised.""" - with pytest.raises(exception_type): - callable_func(*args, **kwargs) - def _add_tenant_account_join( self, sqlite_session: Session, @@ -690,8 +737,7 @@ class TestTenantService: sqlite_session.add(tenant_account_join) return tenant_account_join - @pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True) - def test_iter_member_account_id_batches_uses_offset_limit(self, sqlite_session: Session): + def test_iter_member_account_id_batches_uses_offset_limit(self, sqlite_session: Session) -> None: tenant_id = "00000000-0000-0000-0000-000000000001" account_ids = [ "00000000-0000-0000-0000-000000000011", @@ -714,9 +760,19 @@ class TestTenantService: pagination_parameters: list[tuple[int, int]] = [] - def record_sql(_conn, _cursor, statement, parameters, _context, _executemany): + def record_sql( + _conn: Connection, + _cursor: DBAPICursor, + statement: str, + parameters: Sequence[object], + _context: ExecutionContext | None, + _executemany: bool, + ) -> None: if "FROM tenant_account_joins" in statement: - pagination_parameters.append((parameters[-2], parameters[-1])) + limit, offset = parameters[-2:] + assert isinstance(limit, int) + assert isinstance(offset, int) + pagination_parameters.append((limit, offset)) bind = sqlite_session.get_bind() event.listen(bind, "before_cursor_execute", record_sql) @@ -732,8 +788,7 @@ class TestTenantService: # Backs the auth pipeline's `load_workspace_role`: None => non-member # (pipeline maps to 404), otherwise the caller's role (out-of-set role => 403). - @pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True) - def test_get_account_role_in_tenant_returns_role_for_member(self, sqlite_session: Session): + def test_get_account_role_in_tenant_returns_role_for_member(self, sqlite_session: Session) -> None: """A row in TenantAccountJoin yields the caller's role.""" sqlite_session.add( TenantAccountJoin(tenant_id="tenant-1", account_id="account-1", role=TenantAccountRole.ADMIN) @@ -744,33 +799,18 @@ class TestTenantService: assert role == TenantAccountRole.ADMIN - @pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True) - def test_get_account_role_in_tenant_returns_none_for_non_member(self, sqlite_session: Session): + def test_get_account_role_in_tenant_returns_none_for_non_member(self, sqlite_session: Session) -> None: """No join row => None, so the gate cannot leak the workspace's existence.""" role = TenantService.get_account_role_in_tenant("account-1", "tenant-1", session=sqlite_session) assert role is None - @pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) - def test_get_account_role_in_tenant_short_circuits_empty_account_id(self, sqlite_session: Session): + def test_get_account_role_in_tenant_short_circuits_empty_account_id(self, unbound_session: Session) -> None: """None/empty account_id (SSO bearer, missing identity) returns None without ever touching the session.""" - statements: list[str] = [] + assert TenantService.get_account_role_in_tenant(None, "tenant-1", session=unbound_session) is None - def record_sql(conn, cursor, statement, parameters, context, executemany): - statements.append(statement) - - bind = sqlite_session.get_bind() - event.listen(bind, "before_cursor_execute", record_sql) - try: - assert TenantService.get_account_role_in_tenant(None, "tenant-1", session=sqlite_session) is None - finally: - event.remove(bind, "before_cursor_execute", record_sql) - - assert statements == [] - - @pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True) - def test_get_account_role_in_tenant_query_is_scoped(self, sqlite_session: Session): + def test_get_account_role_in_tenant_query_is_scoped(self, sqlite_session: Session) -> None: """The lookup must filter on BOTH tenant_id and account_id.""" account_id = "11111111-1111-1111-1111-111111111111" tenant_id = "22222222-2222-2222-2222-222222222222" @@ -792,14 +832,12 @@ class TestTenantService: # ==================== Tenant Creation Tests ==================== - @pytest.mark.parametrize( - "sqlite_session", - [(Tenant, TenantAccountJoin, TenantPluginAutoUpgradeStrategy)], - indirect=True, - ) def test_create_owner_tenant_if_not_exist_new_user( - self, sqlite_session: Session, mock_rsa_dependencies, mock_external_service_dependencies - ): + self, + sqlite_session_factory: sessionmaker[Session], + mock_rsa_dependencies: MagicMock, + mock_external_service_dependencies: _MockDependencies, + ) -> None: """Creating an owner workspace persists both the tenant and owner membership.""" mock_account = TestAccountAssociatedDataFactory.create_account_mock() @@ -813,227 +851,302 @@ class TestTenantService: patch("services.credit_pool_service.CreditPoolService.create_default_pool"), patch("services.account_service.tenant_was_created.send") as mock_tenant_was_created, ): - TenantService.create_owner_tenant_if_not_exist(mock_account, session=sqlite_session) + with sqlite_session_factory() as service_session: + TenantService.create_owner_tenant_if_not_exist(mock_account, session=service_session) + tenant = service_session.scalar(select(Tenant).where(Tenant.name == "Test User's Workspace")) + assert tenant is not None + tenant_id = tenant.id + assert mock_account.current_tenant_id == tenant.id + mock_tenant_was_created.assert_called_once_with(tenant) - tenant = sqlite_session.scalar(select(Tenant).where(Tenant.name == "Test User's Workspace")) - assert tenant is not None - assert tenant.encrypt_public_key == "mock_public_key" + mock_rsa_dependencies.assert_called_once_with(tenant_id) - tenant_account_join = sqlite_session.scalar( - select(TenantAccountJoin).where( - TenantAccountJoin.tenant_id == tenant.id, - TenantAccountJoin.account_id == "user-123", + with sqlite_session_factory() as assertion_session: + tenant = assertion_session.get(Tenant, tenant_id) + assert tenant is not None + assert tenant.encrypt_public_key == "mock_public_key" + + tenant_account_join = assertion_session.scalar( + select(TenantAccountJoin).where( + TenantAccountJoin.tenant_id == tenant_id, + TenantAccountJoin.account_id == "user-123", + ) ) - ) - assert tenant_account_join is not None - assert tenant_account_join.role == TenantAccountRole.OWNER - mock_account.set_current_tenant_with_session.assert_called_once_with(tenant, session=sqlite_session) - mock_tenant_was_created.assert_called_once_with(tenant) - mock_rsa_dependencies.assert_called_once_with(tenant.id) + assert tenant_account_join is not None + assert tenant_account_join.role == TenantAccountRole.OWNER # ==================== Member Management Tests ==================== - @pytest.mark.parametrize("sqlite_session", [(Account, Tenant, TenantAccountJoin)], indirect=True) - def test_create_tenant_member_success(self, sqlite_session: Session): + def test_create_tenant_member_success(self, sqlite_session_factory: sessionmaker[Session]) -> None: """Creating a member persists and returns the tenant/account join row.""" - tenant = Tenant(name="Test Workspace") - account = Account(name="Test User", email="test@example.com") - sqlite_session.add_all([tenant, account]) - sqlite_session.commit() + with sqlite_session_factory() as service_session: + tenant = Tenant(name="Test Workspace") + account = Account(name="Test User", email="test@example.com") + service_session.add_all([tenant, account]) + service_session.flush() + tenant_id = tenant.id + account_id = account.id + service_session.commit() - result = TenantService.create_tenant_member(tenant, account, sqlite_session, "normal") + result = TenantService.create_tenant_member(tenant, account, service_session, "normal") + tenant_account_join_id = result.id - assert result.tenant_id == tenant.id - assert result.account_id == account.id - assert result.role == TenantAccountRole.NORMAL - - persisted_tenant_account_join = sqlite_session.scalar( - select(TenantAccountJoin).where( - TenantAccountJoin.tenant_id == tenant.id, - TenantAccountJoin.account_id == account.id, + with sqlite_session_factory() as assertion_session: + persisted_tenant_account_join = assertion_session.get( + TenantAccountJoin, + tenant_account_join_id, ) - ) - assert persisted_tenant_account_join is result + assert persisted_tenant_account_join is not None + assert persisted_tenant_account_join.tenant_id == tenant_id + assert persisted_tenant_account_join.account_id == account_id + assert persisted_tenant_account_join.role == TenantAccountRole.NORMAL # ==================== Member Removal Tests ==================== - @pytest.mark.parametrize("sqlite_session", [(Account, Tenant, TenantAccountJoin, App, Dataset)], indirect=True) - def test_remove_pending_member_deletes_orphaned_account(self, sqlite_session: Session): + def test_remove_pending_member_deletes_orphaned_account( + self, sqlite_session_factory: sessionmaker[Session] + ) -> None: """Test that removing a pending member with no other workspaces deletes the account.""" - tenant = Tenant(name="Test Workspace") - operator = Account(name="Operator", email="operator@example.com") - pending_member = Account(name="Pending Member", email="pending@example.com", status=AccountStatus.PENDING) - sqlite_session.add_all([tenant, operator, pending_member]) - sqlite_session.flush() - self._add_tenant_account_join(sqlite_session, tenant, operator.id, TenantAccountRole.OWNER) - member_join = self._add_tenant_account_join(sqlite_session, tenant, pending_member.id, TenantAccountRole.NORMAL) - sqlite_session.commit() - - with ( - patch("services.account_service.dify_config.BILLING_ENABLED", False), - patch("services.enterprise.account_deletion_sync.sync_workspace_member_removal") as mock_sync, - ): - mock_sync.return_value = True - - TenantService.remove_member_from_tenant(tenant, pending_member, operator, session=sqlite_session) - - mock_sync.assert_called_once_with( - workspace_id=tenant.id, - member_id=pending_member.id, - source="workspace_member_removed", + with sqlite_session_factory() as service_session: + tenant = Tenant(name="Test Workspace") + operator = Account(name="Operator", email="operator@example.com") + pending_member = Account(name="Pending Member", email="pending@example.com", status=AccountStatus.PENDING) + service_session.add_all([tenant, operator, pending_member]) + service_session.flush() + self._add_tenant_account_join(service_session, tenant, operator.id, TenantAccountRole.OWNER) + member_join = self._add_tenant_account_join( + service_session, + tenant, + pending_member.id, + TenantAccountRole.NORMAL, ) + service_session.flush() + tenant_id = tenant.id + member_id = pending_member.id + member_join_id = member_join.id + service_session.commit() - assert sqlite_session.get(TenantAccountJoin, member_join.id) is None - assert sqlite_session.get(Account, pending_member.id) is None + with ( + patch("services.account_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), + patch("services.enterprise.account_deletion_sync.sync_workspace_member_removal") as mock_sync, + ): + mock_sync.return_value = True - @pytest.mark.parametrize("sqlite_session", [(Account, Tenant, TenantAccountJoin, App, Dataset)], indirect=True) - def test_remove_pending_member_keeps_account_with_other_workspaces(self, sqlite_session: Session): + TenantService.remove_member_from_tenant( + tenant, + pending_member, + operator, + session=service_session, + ) + + mock_sync.assert_called_once_with( + workspace_id=tenant_id, + member_id=member_id, + source="workspace_member_removed", + ) + + with sqlite_session_factory() as assertion_session: + assert assertion_session.get(TenantAccountJoin, member_join_id) is None + assert assertion_session.get(Account, member_id) is None + + def test_remove_pending_member_keeps_account_with_other_workspaces( + self, sqlite_session_factory: sessionmaker[Session] + ) -> None: """Test that removing a pending member who belongs to other workspaces preserves the account.""" - tenant = Tenant(name="Test Workspace") - other_tenant = Tenant(name="Other Workspace") - operator = Account(name="Operator", email="operator@example.com") - pending_member = Account(name="Pending Member", email="pending@example.com", status=AccountStatus.PENDING) - sqlite_session.add_all([tenant, other_tenant, operator, pending_member]) - sqlite_session.flush() - self._add_tenant_account_join(sqlite_session, tenant, operator.id, TenantAccountRole.OWNER) - member_join = self._add_tenant_account_join(sqlite_session, tenant, pending_member.id, TenantAccountRole.NORMAL) - self._add_tenant_account_join(sqlite_session, other_tenant, pending_member.id, TenantAccountRole.NORMAL) - sqlite_session.commit() - - with ( - patch("services.account_service.dify_config.BILLING_ENABLED", False), - patch("services.enterprise.account_deletion_sync.sync_workspace_member_removal") as mock_sync, - ): - mock_sync.return_value = True - - TenantService.remove_member_from_tenant(tenant, pending_member, operator, session=sqlite_session) - - mock_sync.assert_called_once_with( - workspace_id=tenant.id, - member_id=pending_member.id, - source="workspace_member_removed", + with sqlite_session_factory() as service_session: + tenant = Tenant(name="Test Workspace") + other_tenant = Tenant(name="Other Workspace") + operator = Account(name="Operator", email="operator@example.com") + pending_member = Account(name="Pending Member", email="pending@example.com", status=AccountStatus.PENDING) + service_session.add_all([tenant, other_tenant, operator, pending_member]) + service_session.flush() + self._add_tenant_account_join(service_session, tenant, operator.id, TenantAccountRole.OWNER) + member_join = self._add_tenant_account_join( + service_session, + tenant, + pending_member.id, + TenantAccountRole.NORMAL, ) + other_member_join = self._add_tenant_account_join( + service_session, + other_tenant, + pending_member.id, + TenantAccountRole.NORMAL, + ) + service_session.flush() + tenant_id = tenant.id + member_id = pending_member.id + member_join_id = member_join.id + other_member_join_id = other_member_join.id + service_session.commit() - assert sqlite_session.get(TenantAccountJoin, member_join.id) is None - assert sqlite_session.get(Account, pending_member.id) is pending_member + with ( + patch("services.account_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), + patch("services.enterprise.account_deletion_sync.sync_workspace_member_removal") as mock_sync, + ): + mock_sync.return_value = True - @pytest.mark.parametrize("sqlite_session", [(Account, Tenant, TenantAccountJoin, App, Dataset)], indirect=True) - def test_remove_active_member_preserves_account(self, sqlite_session: Session): + TenantService.remove_member_from_tenant( + tenant, + pending_member, + operator, + session=service_session, + ) + + mock_sync.assert_called_once_with( + workspace_id=tenant_id, + member_id=member_id, + source="workspace_member_removed", + ) + + with sqlite_session_factory() as assertion_session: + assert assertion_session.get(TenantAccountJoin, member_join_id) is None + assert assertion_session.get(TenantAccountJoin, other_member_join_id) is not None + assert assertion_session.get(Account, member_id) is not None + + def test_remove_active_member_preserves_account(self, sqlite_session_factory: sessionmaker[Session]) -> None: """Test that removing an active member never deletes the account, even with no other workspaces.""" - tenant = Tenant(name="Test Workspace") - operator = Account(name="Operator", email="operator@example.com") - active_member = Account(name="Active Member", email="active@example.com", status=AccountStatus.ACTIVE) - sqlite_session.add_all([tenant, operator, active_member]) - sqlite_session.flush() - self._add_tenant_account_join(sqlite_session, tenant, operator.id, TenantAccountRole.OWNER) - member_join = self._add_tenant_account_join(sqlite_session, tenant, active_member.id, TenantAccountRole.NORMAL) - sqlite_session.commit() - - with ( - patch("services.account_service.dify_config.BILLING_ENABLED", False), - patch("services.enterprise.account_deletion_sync.sync_workspace_member_removal") as mock_sync, - ): - mock_sync.return_value = True - - TenantService.remove_member_from_tenant(tenant, active_member, operator, session=sqlite_session) - - mock_sync.assert_called_once_with( - workspace_id=tenant.id, - member_id=active_member.id, - source="workspace_member_removed", + with sqlite_session_factory() as service_session: + tenant = Tenant(name="Test Workspace") + operator = Account(name="Operator", email="operator@example.com") + active_member = Account(name="Active Member", email="active@example.com", status=AccountStatus.ACTIVE) + service_session.add_all([tenant, operator, active_member]) + service_session.flush() + self._add_tenant_account_join(service_session, tenant, operator.id, TenantAccountRole.OWNER) + member_join = self._add_tenant_account_join( + service_session, + tenant, + active_member.id, + TenantAccountRole.NORMAL, ) + service_session.flush() + tenant_id = tenant.id + member_id = active_member.id + member_join_id = member_join.id + service_session.commit() - assert sqlite_session.get(TenantAccountJoin, member_join.id) is None - assert sqlite_session.get(Account, active_member.id) is active_member + with ( + patch("services.account_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), + patch("services.enterprise.account_deletion_sync.sync_workspace_member_removal") as mock_sync, + ): + mock_sync.return_value = True + + TenantService.remove_member_from_tenant( + tenant, + active_member, + operator, + session=service_session, + ) + + mock_sync.assert_called_once_with( + workspace_id=tenant_id, + member_id=member_id, + source="workspace_member_removed", + ) + + with sqlite_session_factory() as assertion_session: + assert assertion_session.get(TenantAccountJoin, member_join_id) is None + assert assertion_session.get(Account, member_id) is not None # ==================== Tenant Switching Tests ==================== - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_switch_tenant_success(self, sqlite_session: Session): + def test_switch_tenant_success(self, sqlite_session_factory: sessionmaker[Session]) -> None: """Test successful tenant switching.""" - mock_account = TestAccountAssociatedDataFactory.create_account_mock() - tenant = Tenant(name="Target Workspace") - other_tenant = Tenant(name="Other Workspace") - sqlite_session.add_all([tenant, other_tenant]) - sqlite_session.flush() - tenant_join = self._add_tenant_account_join( - sqlite_session, tenant, mock_account.id, TenantAccountRole.NORMAL, current=False - ) - other_tenant_join = self._add_tenant_account_join( - sqlite_session, other_tenant, mock_account.id, TenantAccountRole.NORMAL, current=True - ) - sqlite_session.commit() + with sqlite_session_factory() as service_session: + account = Account(name="Test User", email="test@example.com") + tenant = Tenant(name="Target Workspace") + other_tenant = Tenant(name="Other Workspace") + service_session.add_all([account, tenant, other_tenant]) + service_session.flush() + tenant_join = self._add_tenant_account_join( + service_session, tenant, account.id, TenantAccountRole.NORMAL, current=False + ) + other_tenant_join = self._add_tenant_account_join( + service_session, other_tenant, account.id, TenantAccountRole.NORMAL, current=True + ) + tenant_id = tenant.id + tenant_join_id = tenant_join.id + other_tenant_join_id = other_tenant_join.id + service_session.commit() - with patch("services.account_service.naive_utc_now") as mock_naive_utc_now: mock_now = datetime(2026, 6, 5, 11, 0, 0) - mock_naive_utc_now.return_value = mock_now + with patch("services.account_service.naive_utc_now", return_value=mock_now): + TenantService.switch_tenant(account, tenant_id, session=service_session) - TenantService.switch_tenant(mock_account, tenant.id, session=sqlite_session) + assert account.current_tenant_id == tenant_id - assert tenant_join.current is True - assert tenant_join.last_opened_at == mock_now - assert other_tenant_join.current is False - mock_account.set_tenant_id_with_session.assert_called_once_with(tenant.id, session=sqlite_session) + with sqlite_session_factory() as assertion_session: + tenant_join = assertion_session.get(TenantAccountJoin, tenant_join_id) + other_tenant_join = assertion_session.get(TenantAccountJoin, other_tenant_join_id) + assert tenant_join is not None + assert tenant_join.current is True + assert tenant_join.last_opened_at == mock_now + assert other_tenant_join is not None + assert other_tenant_join.current is False - def test_switch_tenant_commits_changes(self): - account = TestAccountAssociatedDataFactory.create_account_mock() - tenant_join = TestAccountAssociatedDataFactory.create_tenant_join_mock( - tenant_id="tenant-456", account_id="user-123", current=False - ) - session = MagicMock() - session.scalar.return_value = tenant_join - - TenantService.switch_tenant(account, "tenant-456", session=session) - - session.commit.assert_called_once_with() - - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_switch_tenant_no_tenant_id(self, sqlite_session: Session): - """Test tenant switching without providing tenant ID.""" - # Setup test data + def test_switch_tenant_no_tenant_id(self, unbound_session: Session) -> None: mock_account = TestAccountAssociatedDataFactory.create_account_mock() - # Execute test and verify exception - self._assert_exception_raised( - ValueError, TenantService.switch_tenant, mock_account, None, session=sqlite_session - ) + with pytest.raises(ValueError): + TenantService.switch_tenant(mock_account, None, session=unbound_session) # ==================== Role Management Tests ==================== - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_update_member_role_success(self, sqlite_session: Session): + def test_update_member_role_success(self, sqlite_session_factory: sessionmaker[Session]) -> None: """Test successful member role update.""" - tenant = Tenant(name="Test Workspace") - sqlite_session.add(tenant) - sqlite_session.flush() - mock_member = TestAccountAssociatedDataFactory.create_account_mock(account_id="member-789") - mock_operator = TestAccountAssociatedDataFactory.create_account_mock(account_id="operator-123") - target_join = self._add_tenant_account_join(sqlite_session, tenant, mock_member.id, TenantAccountRole.NORMAL) - self._add_tenant_account_join(sqlite_session, tenant, mock_operator.id, TenantAccountRole.OWNER) - sqlite_session.commit() + with sqlite_session_factory() as service_session: + tenant = Tenant(name="Test Workspace") + service_session.add(tenant) + service_session.flush() + target_join = self._add_tenant_account_join( + service_session, + tenant, + "member-789", + TenantAccountRole.NORMAL, + ) + self._add_tenant_account_join( + service_session, + tenant, + "operator-123", + TenantAccountRole.OWNER, + ) + service_session.flush() + target_join_id = target_join.id + service_session.commit() - TenantService.update_member_role(tenant, mock_member, "admin", mock_operator, session=sqlite_session) + mock_member = TestAccountAssociatedDataFactory.create_account_mock(account_id="member-789") + mock_operator = TestAccountAssociatedDataFactory.create_account_mock(account_id="operator-123") - assert target_join.role == TenantAccountRole.ADMIN + TenantService.update_member_role( + tenant, + mock_member, + "admin", + mock_operator, + session=service_session, + ) + + with sqlite_session_factory() as assertion_session: + persisted_target_join = assertion_session.get(TenantAccountJoin, target_join_id) + assert persisted_target_join is not None + assert persisted_target_join.role == TenantAccountRole.ADMIN - @pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True) def test_create_owner_tenant_rbac_enabled_assigns_owner_role( - self, sqlite_session: Session, mock_external_service_dependencies - ): + self, sqlite_session: Session, mock_external_service_dependencies: _MockDependencies + ) -> None: mock_account = TestAccountAssociatedDataFactory.create_account_mock(account_id="user-rbac", name="RBAC User") mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True mock_external_service_dependencies[ "feature_service" ].get_license.return_value.workspaces.is_available.return_value = True - mock_tenant = MagicMock() + mock_tenant = Tenant(name="RBAC User's Workspace") mock_tenant.id = "tenant-rbac" - mock_tenant.name = "RBAC User's Workspace" + sqlite_session.add(mock_tenant) + sqlite_session.flush() with ( patch("services.account_service.dify_config.RBAC_ENABLED", True), patch("services.account_service.TenantService.create_tenant", return_value=mock_tenant), - patch("services.account_service.TenantService.create_tenant_member"), patch( "services.account_service.AccountService._resolve_legacy_role_id", return_value="rbac-owner-id", @@ -1051,24 +1164,36 @@ class TestTenantService: session=sqlite_session, ) - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_admin_can_update_admin_member_role(self, sqlite_session: Session): + def test_admin_can_update_admin_member_role(self, sqlite_session_factory: sessionmaker[Session]) -> None: """Test admin can update another non-owner member, including an admin.""" - tenant = Tenant(name="Test Workspace") - sqlite_session.add(tenant) - sqlite_session.flush() - mock_member = TestAccountAssociatedDataFactory.create_account_mock(account_id="member-789") - mock_operator = TestAccountAssociatedDataFactory.create_account_mock(account_id="operator-123") - target_join = self._add_tenant_account_join(sqlite_session, tenant, mock_member.id, TenantAccountRole.ADMIN) - self._add_tenant_account_join(sqlite_session, tenant, mock_operator.id, TenantAccountRole.ADMIN) - sqlite_session.commit() + with sqlite_session_factory() as service_session: + tenant = Tenant(name="Test Workspace") + service_session.add(tenant) + service_session.flush() + mock_member = TestAccountAssociatedDataFactory.create_account_mock(account_id="member-789") + mock_operator = TestAccountAssociatedDataFactory.create_account_mock(account_id="operator-123") + target_join = self._add_tenant_account_join( + service_session, tenant, mock_member.id, TenantAccountRole.ADMIN + ) + self._add_tenant_account_join(service_session, tenant, mock_operator.id, TenantAccountRole.ADMIN) + service_session.flush() + target_join_id = target_join.id + service_session.commit() - TenantService.update_member_role(tenant, mock_member, "editor", mock_operator, session=sqlite_session) + TenantService.update_member_role( + tenant, + mock_member, + "editor", + mock_operator, + session=service_session, + ) - assert target_join.role == TenantAccountRole.EDITOR + with sqlite_session_factory() as assertion_session: + persisted_target_join = assertion_session.get(TenantAccountJoin, target_join_id) + assert persisted_target_join is not None + assert persisted_target_join.role == TenantAccountRole.EDITOR - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_admin_cannot_update_owner_member_role(self, sqlite_session: Session): + def test_admin_cannot_update_owner_member_role(self, sqlite_session: Session) -> None: """Test admin cannot update an owner member.""" tenant = Tenant(name="Test Workspace") sqlite_session.add(tenant) @@ -1082,8 +1207,7 @@ class TestTenantService: with pytest.raises(NoPermissionError): TenantService.update_member_role(tenant, mock_member, "editor", mock_operator, session=sqlite_session) - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_admin_cannot_promote_member_to_owner(self, sqlite_session: Session): + def test_admin_cannot_promote_member_to_owner(self, sqlite_session: Session) -> None: """Test admin cannot promote a non-owner member to owner.""" tenant = Tenant(name="Test Workspace") sqlite_session.add(tenant) @@ -1099,8 +1223,7 @@ class TestTenantService: # ==================== Permission Check Tests ==================== - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_check_member_permission_success(self, sqlite_session: Session): + def test_check_member_permission_success(self, sqlite_session: Session) -> None: """Test successful member permission check.""" tenant = Tenant(name="Test Workspace") sqlite_session.add(tenant) @@ -1112,8 +1235,7 @@ class TestTenantService: TenantService.check_member_permission(tenant, mock_operator, mock_member, "add", session=sqlite_session) - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_check_member_permission_operate_self(self, sqlite_session: Session): + def test_check_member_permission_operate_self(self, unbound_session: Session) -> None: """Test member permission check when operator tries to operate self.""" # Setup test data mock_tenant = MagicMock() @@ -1123,18 +1245,16 @@ class TestTenantService: # Execute test and verify exception from services.errors.account import CannotOperateSelfError - self._assert_exception_raised( - CannotOperateSelfError, - TenantService.check_member_permission, - mock_tenant, - mock_operator, - mock_operator, # Same as operator - "add", - session=sqlite_session, - ) + with pytest.raises(CannotOperateSelfError): + TenantService.check_member_permission( + mock_tenant, + mock_operator, + mock_operator, # Same as operator + "add", + session=unbound_session, + ) - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_admin_can_remove_non_owner_member(self, sqlite_session: Session): + def test_admin_can_remove_non_owner_member(self, sqlite_session: Session) -> None: """Test admin can remove a non-owner member.""" tenant = Tenant(name="Test Workspace") sqlite_session.add(tenant) @@ -1147,8 +1267,7 @@ class TestTenantService: TenantService.check_member_permission(tenant, mock_operator, mock_member, "remove", session=sqlite_session) - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_admin_cannot_remove_owner_member(self, sqlite_session: Session): + def test_admin_cannot_remove_owner_member(self, sqlite_session: Session) -> None: """Test admin cannot remove an owner member.""" tenant = Tenant(name="Test Workspace") sqlite_session.add(tenant) @@ -1162,8 +1281,7 @@ class TestTenantService: with pytest.raises(NoPermissionError): TenantService.check_member_permission(tenant, mock_operator, mock_member, "remove", session=sqlite_session) - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_rbac_member_can_remove_non_owner_member(self, sqlite_session: Session): + def test_rbac_member_can_remove_non_owner_member(self, sqlite_session: Session) -> None: """Test RBAC workspace.member.manage allows removing a non-owner member.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-456" @@ -1182,8 +1300,7 @@ class TestTenantService: mock_tenant, mock_operator, mock_member, "remove", session=sqlite_session ) - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_rbac_member_cannot_remove_without_permission(self, sqlite_session: Session): + def test_rbac_member_cannot_remove_without_permission(self, sqlite_session: Session) -> None: """Test RBAC permission check rejects removal without workspace.member.manage.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-456" @@ -1202,8 +1319,7 @@ class TestTenantService: mock_tenant, mock_operator, mock_member, "remove", session=sqlite_session ) - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_rbac_member_cannot_remove_owner_member(self, sqlite_session: Session): + def test_rbac_member_cannot_remove_owner_member(self, sqlite_session: Session) -> None: """Test RBAC permission check rejects removing an owner member.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-456" @@ -1223,8 +1339,7 @@ class TestTenantService: mock_tenant, mock_operator, mock_member, "remove", session=sqlite_session ) - @pytest.mark.parametrize("sqlite_session", [()], indirect=True) - def test_get_rbac_workspace_owner_account_id(self, sqlite_session: Session): + def test_get_rbac_workspace_owner_account_id(self, sqlite_session: Session) -> None: mock_roles = Paginated[MembersInRole](data=[MembersInRole(account_id="owner-account")]) mock_rbac_roles = MagicMock() mock_rbac_roles.members.return_value = mock_roles @@ -1264,31 +1379,13 @@ class TestRegisterService: """ @pytest.fixture - def sqlite_session(self, sqlite_engine) -> Iterator[Session]: - """SQLite session with the account/workspace tables registration flows touch.""" - tables = [ - model.metadata.tables[model.__tablename__] - for model in ( - Account, - AccountIntegrate, - Tenant, - TenantAccountJoin, - TenantPluginAutoUpgradeStrategy, - DifySetup, - ) - ] - Account.metadata.create_all(sqlite_engine, tables=tables) - with Session(sqlite_engine, expire_on_commit=False) as session: - yield session - - @pytest.fixture - def mock_redis_dependencies(self): + def mock_redis_dependencies(self) -> Iterator[MagicMock]: """Mock setup for Redis-related functions.""" with patch("services.account_service.redis_client") as mock_redis: yield mock_redis @pytest.fixture - def mock_external_service_dependencies(self): + def mock_external_service_dependencies(self) -> Iterator[_MockDependencies]: """Mock setup for external service dependencies.""" with ( patch("services.account_service.FeatureService") as mock_feature_service, @@ -1302,19 +1399,18 @@ class TestRegisterService: } @pytest.fixture - def mock_task_dependencies(self): + def mock_task_dependencies(self) -> Iterator[MagicMock]: """Mock setup for task dependencies.""" with patch("services.account_service.send_invite_member_mail_task") as mock_send_mail: yield mock_send_mail - def _assert_exception_raised(self, exception_type, callable_func, *args, **kwargs): - """Helper method to verify that specific exception is raised.""" - with pytest.raises(exception_type): - callable_func(*args, **kwargs) - # ==================== Setup Tests ==================== - def test_setup_success(self, sqlite_session: Session, mock_external_service_dependencies): + def test_setup_success( + self, + sqlite_session_factory: sessionmaker[Session], + mock_external_service_dependencies: _MockDependencies, + ) -> None: """Test successful system setup.""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True @@ -1325,58 +1421,113 @@ class TestRegisterService: with patch("services.account_service.AccountService.create_account") as mock_create_account: mock_create_account.return_value = mock_account - with patch("services.account_service.TenantService.create_owner_tenant_if_not_exist") as mock_create_tenant: + with ( + patch("services.account_service.TenantService.create_owner_tenant_if_not_exist") as mock_create_tenant, + patch("services.account_service.CommunityTelemetryService.report_install") as mock_report_install, + ): + with sqlite_session_factory() as service_session: + RegisterService.setup( + "admin@example.com", + "Admin User", + "password123", + "192.168.1.1", + "en-US", + session=service_session, + ) + + mock_create_account.assert_called_once_with( + email="admin@example.com", + name="Admin User", + interface_language="en-US", + password="password123", + is_setup=True, + session=service_session, + ) + mock_create_tenant.assert_called_once_with( + account=mock_account, + is_setup=True, + session=service_session, + ) + mock_report_install.assert_called_once_with(session=service_session) + + with sqlite_session_factory() as assertion_session: + dify_setup = assertion_session.scalar(select(DifySetup)) + assert dify_setup is not None + assert dify_setup.instance_id is not None + assert str(UUID(dify_setup.instance_id)) == dify_setup.instance_id + assert dify_setup.install_reported_at is None + assert dify_setup.last_heartbeat_at is None + + def test_setup_succeeds_when_telemetry_install_report_fails( + self, + sqlite_session_factory: sessionmaker[Session], + mock_external_service_dependencies: _MockDependencies, + ) -> None: + mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True + mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False + mock_account = TestAccountAssociatedDataFactory.create_account_mock() + + with ( + patch("services.account_service.AccountService.create_account", return_value=mock_account), + patch("services.account_service.TenantService.create_owner_tenant_if_not_exist"), + patch( + "services.account_service.CommunityTelemetryService.report_install", + side_effect=RuntimeError("telemetry unavailable"), + ), + ): + with sqlite_session_factory() as service_session: RegisterService.setup( "admin@example.com", "Admin User", "password123", "192.168.1.1", "en-US", - session=sqlite_session, + session=service_session, ) - mock_create_account.assert_called_once_with( - email="admin@example.com", - name="Admin User", - interface_language="en-US", - password="password123", - is_setup=True, - session=sqlite_session, - ) - mock_create_tenant.assert_called_once_with(account=mock_account, is_setup=True, session=sqlite_session) - assert sqlite_session.scalar(select(DifySetup)) is not None + with sqlite_session_factory() as assertion_session: + assert assertion_session.scalar(select(DifySetup)) is not None - def test_setup_failure_rollback(self, sqlite_session: Session, mock_external_service_dependencies): - """Test setup failure with proper rollback.""" - # Setup mocks to simulate failure + def test_setup_failure_cleans_partially_persisted_account( + self, + sqlite_session_factory: sessionmaker[Session], + mock_external_service_dependencies: _MockDependencies, + ) -> None: mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True + mock_external_service_dependencies[ + "feature_service" + ].get_license.return_value.seats.is_available.return_value = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False - # Mock AccountService.create_account to raise exception - with patch("services.account_service.AccountService.create_account") as mock_create_account: - mock_create_account.side_effect = Exception("Database error") + with patch( + "services.account_service.TenantService.create_owner_tenant_if_not_exist", + side_effect=RuntimeError("tenant creation failed"), + ): + with sqlite_session_factory() as service_session: + with pytest.raises(ValueError, match="Setup failed: tenant creation failed"): + RegisterService.setup( + "admin@example.com", + "Admin User", + "password123", + "192.168.1.1", + "en-US", + session=service_session, + ) - # Execute test and verify exception - self._assert_exception_raised( - ValueError, - RegisterService.setup, - "admin@example.com", - "Admin User", - "password123", - "192.168.1.1", - "en-US", - session=sqlite_session, - ) - - assert sqlite_session.scalar(select(DifySetup)) is None + with sqlite_session_factory() as assertion_session: + assert assertion_session.scalar(select(Account).where(Account.email == "admin@example.com")) is None + assert assertion_session.scalar(select(DifySetup)) is None # ==================== Registration Tests ==================== - def test_create_account_and_tenant_calls_default_workspace_join_when_enterprise_enabled( - self, sqlite_session: Session, mock_external_service_dependencies, monkeypatch: pytest.MonkeyPatch - ): - """Enterprise-only side effect should be invoked when ENTERPRISE_ENABLED is True.""" - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", True, raising=False) + def test_create_account_and_tenant_calls_default_workspace_join_for_enterprise_edition( + self, + sqlite_session: Session, + mock_external_service_dependencies: _MockDependencies, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """Enterprise-only side effect should be invoked for the ENTERPRISE edition.""" + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE, raising=False) mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False @@ -1402,13 +1553,16 @@ class TestRegisterService: assert result == mock_account mock_create_workspace.assert_called_once_with(account=mock_account, session=sqlite_session) - mock_join_default_workspace.assert_called_once_with(str(mock_account.id)) + mock_join_default_workspace.assert_called_once_with(mock_account.id) - def test_create_account_and_tenant_does_not_call_default_workspace_join_when_enterprise_disabled( - self, sqlite_session: Session, mock_external_service_dependencies, monkeypatch: pytest.MonkeyPatch - ): - """Enterprise-only side effect should not be invoked when ENTERPRISE_ENABLED is False.""" - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", False, raising=False) + def test_create_account_and_tenant_skips_default_workspace_join_for_community_edition( + self, + sqlite_session: Session, + mock_external_service_dependencies: _MockDependencies, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """Enterprise-only side effect should not be invoked for the COMMUNITY edition.""" + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY, raising=False) mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False @@ -1436,12 +1590,15 @@ class TestRegisterService: mock_join_default_workspace.assert_not_called() def test_create_account_and_tenant_still_calls_default_workspace_join_when_workspace_creation_fails( - self, sqlite_session: Session, mock_external_service_dependencies, monkeypatch: pytest.MonkeyPatch - ): + self, + sqlite_session: Session, + mock_external_service_dependencies: _MockDependencies, + monkeypatch: pytest.MonkeyPatch, + ) -> None: """Default workspace join should still be attempted when personal workspace creation fails.""" from services.errors.workspace import WorkSpaceNotAllowedCreateError - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", True, raising=False) + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE, raising=False) mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False @@ -1466,9 +1623,11 @@ class TestRegisterService: session=sqlite_session, ) - mock_join_default_workspace.assert_called_once_with(str(mock_account.id)) + mock_join_default_workspace.assert_called_once_with(mock_account.id) - def test_register_success(self, sqlite_session: Session, mock_external_service_dependencies): + def test_register_success( + self, sqlite_session: Session, mock_external_service_dependencies: _MockDependencies + ) -> None: """Test successful account registration.""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True @@ -1510,11 +1669,14 @@ class TestRegisterService: ) mock_create_owner_tenant.assert_called_once_with(mock_account, session=sqlite_session) - def test_register_calls_default_workspace_join_when_enterprise_enabled( - self, sqlite_session: Session, mock_external_service_dependencies, monkeypatch: pytest.MonkeyPatch - ): + def test_register_calls_default_workspace_join_for_enterprise_edition( + self, + sqlite_session: Session, + mock_external_service_dependencies: _MockDependencies, + monkeypatch: pytest.MonkeyPatch, + ) -> None: """Enterprise-only side effect should be invoked after successful register commit.""" - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", True, raising=False) + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE, raising=False) mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False @@ -1539,13 +1701,16 @@ class TestRegisterService: ) assert result == mock_account - mock_join_default_workspace.assert_called_once_with(str(mock_account.id)) + mock_join_default_workspace.assert_called_once_with(mock_account.id) - def test_register_does_not_call_default_workspace_join_when_enterprise_disabled( - self, sqlite_session: Session, mock_external_service_dependencies, monkeypatch: pytest.MonkeyPatch - ): - """Enterprise-only side effect should not be invoked when ENTERPRISE_ENABLED is False.""" - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", False, raising=False) + def test_register_skips_default_workspace_join_for_community_edition( + self, + sqlite_session: Session, + mock_external_service_dependencies: _MockDependencies, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """Enterprise-only side effect should not be invoked for the COMMUNITY edition.""" + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY, raising=False) mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["billing_service"].is_email_in_freeze.return_value = False @@ -1572,12 +1737,15 @@ class TestRegisterService: mock_join_default_workspace.assert_not_called() def test_register_still_calls_default_workspace_join_when_personal_workspace_creation_fails( - self, sqlite_session: Session, mock_external_service_dependencies, monkeypatch: pytest.MonkeyPatch - ): + self, + sqlite_session: Session, + mock_external_service_dependencies: _MockDependencies, + monkeypatch: pytest.MonkeyPatch, + ) -> None: """Default workspace join should run even when personal workspace creation raises.""" from services.errors.workspace import WorkSpaceNotAllowedCreateError - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", True, raising=False) + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE, raising=False) mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True mock_external_service_dependencies[ @@ -1606,15 +1774,18 @@ class TestRegisterService: session=sqlite_session, ) - mock_join_default_workspace.assert_called_once_with(str(mock_account.id)) + mock_join_default_workspace.assert_called_once_with(mock_account.id) def test_register_still_calls_default_workspace_join_when_workspace_limit_exceeded( - self, sqlite_session: Session, mock_external_service_dependencies, monkeypatch: pytest.MonkeyPatch - ): + self, + sqlite_session: Session, + mock_external_service_dependencies: _MockDependencies, + monkeypatch: pytest.MonkeyPatch, + ) -> None: """Default workspace join should run before propagating workspace-limit registration failure.""" from services.errors.workspace import WorkspacesLimitExceededError - monkeypatch.setattr(dify_config, "ENTERPRISE_ENABLED", True, raising=False) + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE, raising=False) mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True mock_external_service_dependencies["feature_service"].is_workspace_creation_allowed.return_value = True mock_external_service_dependencies[ @@ -1643,9 +1814,11 @@ class TestRegisterService: session=sqlite_session, ) - mock_join_default_workspace.assert_called_once_with(str(mock_account.id)) + mock_join_default_workspace.assert_called_once_with(mock_account.id) - def test_register_with_oauth(self, sqlite_session: Session, mock_external_service_dependencies): + def test_register_with_oauth( + self, sqlite_session: Session, mock_external_service_dependencies: _MockDependencies + ) -> None: """Test account registration with OAuth integration.""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True @@ -1669,8 +1842,17 @@ class TestRegisterService: patch("services.account_service.TenantService.create_tenant_member") as mock_create_member, patch("services.account_service.tenant_was_created") as mock_event, ): - mock_tenant = MagicMock() + mock_tenant = Tenant(name="Test User's Workspace") + sqlite_session.add(mock_tenant) + sqlite_session.flush() mock_create_tenant.return_value = mock_tenant + mock_create_member.side_effect = lambda tenant, account, session, role: session.add( + TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + role=TenantAccountRole(role), + ) + ) # Execute test result = RegisterService.register( @@ -1687,7 +1869,9 @@ class TestRegisterService: assert result == mock_account mock_link_account.assert_called_once_with("google", "oauth123", mock_account, session=sqlite_session) - def test_register_with_pending_status(self, sqlite_session: Session, mock_external_service_dependencies): + def test_register_with_pending_status( + self, sqlite_session: Session, mock_external_service_dependencies: _MockDependencies + ) -> None: """Test account registration with pending status.""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True @@ -1708,8 +1892,17 @@ class TestRegisterService: patch("services.account_service.TenantService.create_tenant_member") as mock_create_member, patch("services.account_service.tenant_was_created") as mock_event, ): - mock_tenant = MagicMock() + mock_tenant = Tenant(name="Test User's Workspace") + sqlite_session.add(mock_tenant) + sqlite_session.flush() mock_create_tenant.return_value = mock_tenant + mock_create_member.side_effect = lambda tenant, account, session, role: session.add( + TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + role=TenantAccountRole(role), + ) + ) # Execute test with pending status from models.account import AccountStatus @@ -1727,7 +1920,9 @@ class TestRegisterService: assert result == mock_account assert result.status == "pending" - def test_register_workspace_not_allowed(self, sqlite_session: Session, mock_external_service_dependencies): + def test_register_workspace_not_allowed( + self, sqlite_session: Session, mock_external_service_dependencies: _MockDependencies + ) -> None: """Test registration when workspace creation is not allowed.""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True @@ -1748,19 +1943,20 @@ class TestRegisterService: with patch("services.account_service.TenantService.create_tenant") as mock_create_tenant: mock_create_tenant.side_effect = WorkSpaceNotAllowedCreateError() - self._assert_exception_raised( - AccountRegisterError, - RegisterService.register, - email="test@example.com", - name="Test User", - password="password123", - language="en-US", - session=sqlite_session, - ) + with pytest.raises(AccountRegisterError): + RegisterService.register( + email="test@example.com", + name="Test User", + password="password123", + language="en-US", + session=sqlite_session, + ) assert sqlite_session.scalar(select(Account).where(Account.email == "test@example.com")) is None - def test_register_general_exception(self, sqlite_session: Session, mock_external_service_dependencies): + def test_register_general_exception( + self, sqlite_session: Session, mock_external_service_dependencies: _MockDependencies + ) -> None: """Test registration with general exception handling.""" # Setup mocks mock_external_service_dependencies["feature_service"].get_system_features.return_value.is_allow_register = True @@ -1771,23 +1967,21 @@ class TestRegisterService: mock_create_account.side_effect = Exception("Unexpected error") # Execute test and verify exception - self._assert_exception_raised( - AccountRegisterError, - RegisterService.register, - email="test@example.com", - name="Test User", - password="password123", - language="en-US", - session=sqlite_session, - ) + with pytest.raises(AccountRegisterError): + RegisterService.register( + email="test@example.com", + name="Test User", + password="password123", + language="en-US", + session=sqlite_session, + ) assert sqlite_session.scalar(select(Account).where(Account.email == "test@example.com")) is None # ==================== Member Invitation Tests ==================== - def test_invite_new_member_new_account( - self, sqlite_session: Session, mock_redis_dependencies, mock_task_dependencies - ): + @pytest.mark.usefixtures("mock_task_dependencies") + def test_invite_new_member_new_account(self, sqlite_session: Session) -> None: """Test inviting a new member who doesn't have an account.""" # Setup test data mock_tenant = MagicMock() @@ -1839,8 +2033,8 @@ class TestRegisterService: mock_lookup.assert_called_once_with("newuser@example.com", session=sqlite_session) def test_invite_new_member_normalizes_new_account_email( - self, sqlite_session: Session, mock_redis_dependencies, mock_task_dependencies - ): + self, sqlite_session: Session, mock_task_dependencies: MagicMock + ) -> None: """Ensure inviting with mixed-case email normalizes before registering.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-456" @@ -1898,8 +2092,8 @@ class TestRegisterService: mock_task_dependencies.delay.assert_called_once() def test_invite_new_member_existing_account( - self, sqlite_session: Session, mock_redis_dependencies, mock_task_dependencies - ): + self, sqlite_session: Session, mock_task_dependencies: MagicMock + ) -> None: """Test inviting a pending account that is not in the tenant yet.""" # Setup test data mock_tenant = MagicMock() @@ -1943,8 +2137,8 @@ class TestRegisterService: mock_lookup.assert_called_once_with("existing@example.com", session=sqlite_session) def test_invite_existing_active_account_requires_acceptance_before_joining( - self, sqlite_session: Session, mock_redis_dependencies, mock_task_dependencies - ): + self, sqlite_session: Session, mock_task_dependencies: MagicMock + ) -> None: """Existing active accounts outside the tenant receive an invite without immediate membership.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-456" @@ -1987,7 +2181,7 @@ class TestRegisterService: ) mock_task_dependencies.delay.assert_called_once() - def test_invite_new_member_already_in_tenant(self, sqlite_session: Session, mock_redis_dependencies): + def test_invite_new_member_already_in_tenant(self, sqlite_session: Session) -> None: """Test inviting a member who is already in the tenant.""" # Setup test data mock_tenant = MagicMock() @@ -2013,40 +2207,37 @@ class TestRegisterService: ): mock_lookup.return_value = mock_existing_account # Execute test and verify exception - self._assert_exception_raised( - AccountAlreadyInTenantError, - RegisterService.invite_new_member, - tenant=mock_tenant, - email="existing@example.com", - language="en-US", - role="normal", - inviter=mock_inviter, - session=sqlite_session, - ) + with pytest.raises(AccountAlreadyInTenantError): + RegisterService.invite_new_member( + tenant=mock_tenant, + email="existing@example.com", + language="en-US", + role="normal", + inviter=mock_inviter, + session=sqlite_session, + ) mock_lookup.assert_called_once() - def test_invite_new_member_no_inviter(self, sqlite_session: Session): + def test_invite_new_member_no_inviter(self, unbound_session: Session) -> None: """Test inviting a member without providing an inviter.""" # Setup test data mock_tenant = MagicMock() # Execute test and verify exception - self._assert_exception_raised( - ValueError, - RegisterService.invite_new_member, - tenant=mock_tenant, - email="test@example.com", - language="en-US", - role="normal", - inviter=None, - session=sqlite_session, - ) + with pytest.raises(ValueError): + RegisterService.invite_new_member( + tenant=mock_tenant, + email="test@example.com", + language="en-US", + role="normal", + inviter=None, + session=unbound_session, + ) # ==================== RBAC Member Invitation Tests ==================== - def test_invite_new_member_rbac_enabled_new_account( - self, sqlite_session: Session, mock_redis_dependencies, mock_task_dependencies - ): + @pytest.mark.usefixtures("mock_task_dependencies") + def test_invite_new_member_rbac_enabled_new_account(self, sqlite_session: Session) -> None: """When RBAC is enabled, create the member join and replace RBAC member roles.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-789" @@ -2093,9 +2284,8 @@ class TestRegisterService: session=sqlite_session, ) - def test_invite_new_member_rbac_enabled_existing_account( - self, sqlite_session: Session, mock_redis_dependencies, mock_task_dependencies - ): + @pytest.mark.usefixtures("mock_task_dependencies") + def test_invite_new_member_rbac_enabled_existing_account(self, sqlite_session: Session) -> None: """When RBAC is enabled and account exists, create the member join and replace RBAC member roles.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-789" @@ -2142,8 +2332,8 @@ class TestRegisterService: ) def test_invite_new_member_rbac_enabled_existing_active_account_adds_role_before_signin_response( - self, sqlite_session: Session, mock_redis_dependencies, mock_task_dependencies - ): + self, sqlite_session: Session, mock_task_dependencies: MagicMock + ) -> None: """Existing active accounts still need an RBAC membership before the API returns the signin URL.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-789" @@ -2189,9 +2379,8 @@ class TestRegisterService: ) mock_task_dependencies.delay.assert_not_called() - def test_invite_new_member_rbac_disabled_uses_legacy_role( - self, sqlite_session: Session, mock_redis_dependencies, mock_task_dependencies - ): + @pytest.mark.usefixtures("mock_task_dependencies") + def test_invite_new_member_rbac_disabled_uses_legacy_role(self, sqlite_session: Session) -> None: """When RBAC is disabled, create_tenant_member should be called and MemberRoles.replace should NOT.""" mock_tenant = MagicMock() mock_tenant.id = "tenant-legacy" @@ -2232,7 +2421,7 @@ class TestRegisterService: # ==================== Token Management Tests ==================== - def test_generate_invite_token_success(self, mock_redis_dependencies): + def test_generate_invite_token_success(self, mock_redis_dependencies: MagicMock) -> None: """Test successful invite token generation.""" # Setup test data mock_tenant = MagicMock() @@ -2262,7 +2451,7 @@ class TestRegisterService: assert stored_data["role"] == "admin" assert stored_data["requires_setup"] is True - def test_is_valid_invite_token_valid(self, mock_redis_dependencies): + def test_is_valid_invite_token_valid(self, mock_redis_dependencies: MagicMock) -> None: """Test checking valid invite token.""" # Setup mock mock_redis_dependencies.get.return_value = b'{"test": "data"}' @@ -2274,7 +2463,7 @@ class TestRegisterService: assert result is True mock_redis_dependencies.get.assert_called_once_with("member_invite:token:valid-token") - def test_is_valid_invite_token_invalid(self, mock_redis_dependencies): + def test_is_valid_invite_token_invalid(self, mock_redis_dependencies: MagicMock) -> None: """Test checking invalid invite token.""" # Setup mock mock_redis_dependencies.get.return_value = None @@ -2286,7 +2475,7 @@ class TestRegisterService: assert result is False mock_redis_dependencies.get.assert_called_once_with("member_invite:token:invalid-token") - def test_revoke_token_with_workspace_and_email(self, mock_redis_dependencies): + def test_revoke_token_with_workspace_and_email(self, mock_redis_dependencies: MagicMock) -> None: """Test revoking token with workspace ID and email.""" # Execute test RegisterService.revoke_token("workspace-123", "test@example.com", "token-123") @@ -2298,7 +2487,7 @@ class TestRegisterService: # The email is hashed, so we check for the hash pattern instead assert "member_invite_token:" in call_args[0][0] - def test_revoke_token_without_workspace_and_email(self, mock_redis_dependencies): + def test_revoke_token_without_workspace_and_email(self, mock_redis_dependencies: MagicMock) -> None: """Test revoking token without workspace ID and email.""" # Execute test RegisterService.revoke_token("", "", "token-123") @@ -2308,7 +2497,7 @@ class TestRegisterService: # ==================== Invitation Validation Tests ==================== - def test_get_invitation_if_token_valid_success(self, sqlite_session: Session, mock_redis_dependencies): + def test_get_invitation_if_token_valid_success(self, sqlite_session: Session) -> None: """Test successful invitation validation.""" tenant = Tenant(name="Test Workspace") account = Account(name="Test User", email="test@example.com") @@ -2332,20 +2521,24 @@ class TestRegisterService: assert result["tenant"] is tenant assert result["data"] == invitation_data - def test_get_invitation_if_token_valid_no_token_data(self, sqlite_session: Session, mock_redis_dependencies): + def test_get_invitation_if_token_valid_no_token_data( + self, unbound_session: Session, mock_redis_dependencies: MagicMock + ) -> None: """Test invitation validation with no token data.""" # Setup mock mock_redis_dependencies.get.return_value = None # Execute test result = RegisterService.get_invitation_if_token_valid( - "tenant-456", "test@example.com", "token-123", session=sqlite_session + "tenant-456", "test@example.com", "token-123", session=unbound_session ) # Verify results assert result is None - def test_get_invitation_if_token_valid_tenant_not_found(self, sqlite_session: Session, mock_redis_dependencies): + def test_get_invitation_if_token_valid_tenant_not_found( + self, sqlite_session: Session, mock_redis_dependencies: MagicMock + ) -> None: """Test invitation validation when tenant is not found.""" # Setup mock Redis data invitation_data = { @@ -2362,7 +2555,9 @@ class TestRegisterService: # Verify results assert result is None - def test_get_invitation_if_token_valid_account_not_found(self, sqlite_session: Session, mock_redis_dependencies): + def test_get_invitation_if_token_valid_account_not_found( + self, sqlite_session: Session, mock_redis_dependencies: MagicMock + ) -> None: """Test invitation validation when account is not found.""" tenant = Tenant(name="Test Workspace") sqlite_session.add(tenant) @@ -2383,7 +2578,9 @@ class TestRegisterService: # Verify results assert result is None - def test_get_invitation_if_token_valid_account_id_mismatch(self, sqlite_session: Session, mock_redis_dependencies): + def test_get_invitation_if_token_valid_account_id_mismatch( + self, sqlite_session: Session, mock_redis_dependencies: MagicMock + ) -> None: """Test invitation validation when account ID doesn't match.""" tenant = Tenant(name="Test Workspace") account = Account(name="Test User", email="test@example.com") @@ -2405,7 +2602,7 @@ class TestRegisterService: # Verify results assert result is None - def test_get_invitation_with_case_fallback_returns_initial_match(self, sqlite_session: Session): + def test_get_invitation_with_case_fallback_returns_initial_match(self, sqlite_session: Session) -> None: """Fallback helper should return the initial invitation when present.""" invitation = {"workspace_id": "tenant-456"} with patch( @@ -2420,7 +2617,7 @@ class TestRegisterService: "tenant-456", "User@Test.com", "token-123", session=mock_get.call_args.kwargs["session"] ) - def test_get_invitation_with_case_fallback_retries_with_lowercase(self, sqlite_session: Session): + def test_get_invitation_with_case_fallback_retries_with_lowercase(self, sqlite_session: Session) -> None: """Fallback helper should retry with lowercase email when needed.""" invitation = {"workspace_id": "tenant-456"} with patch("services.account_service.RegisterService.get_invitation_if_token_valid") as mock_get: @@ -2437,7 +2634,7 @@ class TestRegisterService: # ==================== Helper Method Tests ==================== - def test_get_invitation_token_key(self): + def test_get_invitation_token_key(self) -> None: """Test the _get_invitation_token_key helper method.""" # Execute test result = RegisterService._get_invitation_token_key("test-token") @@ -2445,7 +2642,7 @@ class TestRegisterService: # Verify results assert result == "member_invite:token:test-token" - def test_get_invitation_by_token_with_workspace_and_email(self, mock_redis_dependencies): + def test_get_invitation_by_token_with_workspace_and_email(self, mock_redis_dependencies: MagicMock) -> None: """Test get_invitation_by_token with workspace ID and email.""" # Setup mock mock_redis_dependencies.get.return_value = b"user-123" @@ -2459,7 +2656,7 @@ class TestRegisterService: assert result["email"] == "test@example.com" assert result["workspace_id"] == "workspace-456" - def test_get_invitation_by_token_without_workspace_and_email(self, mock_redis_dependencies): + def test_get_invitation_by_token_without_workspace_and_email(self, mock_redis_dependencies: MagicMock) -> None: """Test get_invitation_by_token without workspace ID and email.""" # Setup mock invitation_data = { @@ -2476,7 +2673,7 @@ class TestRegisterService: assert result is not None assert result == invitation_data - def test_get_invitation_by_token_no_data(self, mock_redis_dependencies): + def test_get_invitation_by_token_no_data(self, mock_redis_dependencies: MagicMock) -> None: """Test get_invitation_by_token with no data.""" # Setup mock mock_redis_dependencies.get.return_value = None @@ -2514,28 +2711,31 @@ class TestSessionInjectedGetters: sqlite_session.add(tenant_account_join) return tenant_account_join - @pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) - def test_get_account_by_id_uses_passed_session_no_side_effects(self, sqlite_session: Session): + def test_get_account_by_id_uses_passed_session_no_side_effects( + self, sqlite_session_factory: sessionmaker[Session] + ) -> None: """``get_account_by_id`` must be a plain delegation to ``session.get(Account, ...)`` — no banned-status raise, no commit (those are the side-effects of ``load_user`` we explicitly want to skip). """ - account = Account(name="Alice", email="alice@example.com", status=AccountStatus.BANNED) - sqlite_session.add(account) - sqlite_session.commit() + with sqlite_session_factory.begin() as arrange_session: + account = Account(name="Alice", email="alice@example.com", status=AccountStatus.BANNED) + arrange_session.add(account) + arrange_session.flush() + account_id = account.id - result = AccountService.get_account_by_id(account.id, session=sqlite_session) + with sqlite_session_factory() as service_session: + result = AccountService.get_account_by_id(account_id, session=service_session) - assert result is account - assert account.status == AccountStatus.BANNED + assert result is not None + assert result.id == account_id + assert result.status == AccountStatus.BANNED - @pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) - def test_get_account_by_id_returns_none_for_unknown_account(self, sqlite_session: Session): + def test_get_account_by_id_returns_none_for_unknown_account(self, sqlite_session: Session) -> None: assert AccountService.get_account_by_id("missing", session=sqlite_session) is None - @pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) - def test_get_account_by_email_returns_scalar_or_none(self, sqlite_session: Session): + def test_get_account_by_email_returns_scalar_or_none(self, sqlite_session: Session) -> None: """Plain getter — case-sensitive equality (callers needing the case-insensitive existence check use :meth:`has_active_account_with_email`). @@ -2548,28 +2748,24 @@ class TestSessionInjectedGetters: assert AccountService.get_account_by_email("ALICE@example.com", session=sqlite_session) is None assert AccountService.get_account_by_email("ghost@example.com", session=sqlite_session) is None - @pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) - def test_account_belongs_to_tenant_short_circuits_on_falsy_account_id(self, sqlite_session: Session): + def test_account_belongs_to_tenant_short_circuits_on_falsy_account_id(self, unbound_session: Session) -> None: """SSO bearers with no ``account_id`` (and any other falsy id) must collapse to ``False`` before touching membership storage. """ - assert TenantService.account_belongs_to_tenant(None, "tenant-1", session=sqlite_session) is False - assert TenantService.account_belongs_to_tenant("", "tenant-1", session=sqlite_session) is False + assert TenantService.account_belongs_to_tenant(None, "tenant-1", session=unbound_session) is False + assert TenantService.account_belongs_to_tenant("", "tenant-1", session=unbound_session) is False - @pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True) - def test_account_belongs_to_tenant_true_when_join_row_exists(self, sqlite_session: Session): + def test_account_belongs_to_tenant_true_when_join_row_exists(self, sqlite_session: Session) -> None: sqlite_session.add(TenantAccountJoin(tenant_id="tenant-1", account_id="user-1", role=TenantAccountRole.NORMAL)) sqlite_session.commit() assert TenantService.account_belongs_to_tenant("user-1", "tenant-1", session=sqlite_session) is True assert TenantService.account_belongs_to_tenant("user-1", "other-tenant", session=sqlite_session) is False - @pytest.mark.parametrize("sqlite_session", [(TenantAccountJoin,)], indirect=True) - def test_account_belongs_to_tenant_false_when_no_join(self, sqlite_session: Session): + def test_account_belongs_to_tenant_false_when_no_join(self, sqlite_session: Session) -> None: assert TenantService.account_belongs_to_tenant("user-1", "tenant-1", session=sqlite_session) is False - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_get_account_memberships_returns_join_tenant_pairs(self, sqlite_session: Session): + def test_get_account_memberships_returns_join_tenant_pairs(self, sqlite_session: Session) -> None: """Returns every ``(TenantAccountJoin, Tenant)`` pair for an account.""" tenant = Tenant(name="Joined Workspace") other_tenant = Tenant(name="Other Workspace") @@ -2585,8 +2781,7 @@ class TestSessionInjectedGetters: assert out[0][0] is join assert out[0][1] is tenant - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_get_workspaces_for_account_uses_session_execute(self, sqlite_session: Session): + def test_get_workspaces_for_account_uses_session_execute(self, sqlite_session: Session) -> None: """The list endpoint orders by ``Tenant.created_at``; the helper returns ``(Tenant, TenantAccountJoin)`` rows in that order. """ @@ -2604,29 +2799,31 @@ class TestSessionInjectedGetters: assert [(row[0], row[1]) for row in out] == [(older_tenant, older_join), (newer_tenant, newer_join)] - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_get_tenant_by_id_is_plain_session_get(self, sqlite_session: Session): + def test_get_tenant_by_id_is_plain_session_get(self, sqlite_session_factory: sessionmaker[Session]) -> None: """``get_tenant_by_id`` must NOT apply a status filter — the openapi auth pipeline needs to map ``status == ARCHIVE`` to a 403, distinct from a 404 for "missing". """ - tenant = Tenant(name="Archived Workspace", status=TenantStatus.ARCHIVE) - sqlite_session.add(tenant) - sqlite_session.commit() + with sqlite_session_factory.begin() as arrange_session: + tenant = Tenant(name="Archived Workspace", status=TenantStatus.ARCHIVE) + arrange_session.add(tenant) + arrange_session.flush() + tenant_id = tenant.id - assert TenantService.get_tenant_by_id(tenant.id, session=sqlite_session) is tenant + with sqlite_session_factory() as service_session: + result = TenantService.get_tenant_by_id(tenant_id, session=service_session) + assert result is not None + assert result.id == tenant_id + assert result.status == TenantStatus.ARCHIVE - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_get_tenant_by_id_returns_none_when_missing(self, sqlite_session: Session): + def test_get_tenant_by_id_returns_none_when_missing(self, sqlite_session: Session) -> None: assert TenantService.get_tenant_by_id("missing", session=sqlite_session) is None - @pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) - def test_get_tenants_by_ids_short_circuits_on_empty_input(self, sqlite_session: Session): + def test_get_tenants_by_ids_short_circuits_on_empty_input(self, unbound_session: Session) -> None: """Empty id list must return before touching tenant storage.""" - assert TenantService.get_tenants_by_ids([], session=sqlite_session) == [] + assert TenantService.get_tenants_by_ids([], session=unbound_session) == [] - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_get_tenants_by_ids_returns_scalars(self, sqlite_session: Session): + def test_get_tenants_by_ids_returns_scalars(self, sqlite_session: Session) -> None: tenant_1 = Tenant(name="Workspace 1") tenant_2 = Tenant(name="Workspace 2") tenant_3 = Tenant(name="Workspace 3") @@ -2637,8 +2834,7 @@ class TestSessionInjectedGetters: assert {tenant.id for tenant in tenants} == {tenant_1.id, tenant_3.id} - @pytest.mark.parametrize("sqlite_session", [(Tenant,)], indirect=True) - def test_get_tenant_name_returns_scalar_or_none(self, sqlite_session: Session): + def test_get_tenant_name_returns_scalar_or_none(self, sqlite_session: Session) -> None: """Single-column lookup: ``session.execute(...).scalar_one_or_none()`` — used by openapi list endpoints to denormalise ``workspace_name`` onto each row. @@ -2650,8 +2846,7 @@ class TestSessionInjectedGetters: assert TenantService.get_tenant_name(tenant.id, session=sqlite_session) == "Acme Inc." assert TenantService.get_tenant_name("missing", session=sqlite_session) is None - @pytest.mark.parametrize("sqlite_session", [(Tenant, TenantAccountJoin)], indirect=True) - def test_find_workspace_for_account_returns_first_row_or_none(self, sqlite_session: Session): + def test_find_workspace_for_account_returns_first_row_or_none(self, sqlite_session: Session) -> None: """Per-id read returns ``session.execute(...).first()`` directly; callers map ``None`` → 404 to avoid leaking workspace IDs across tenants. @@ -2672,7 +2867,6 @@ class TestSessionInjectedGetters: assert TenantService.find_workspace_for_account("user-123", other_tenant.id, session=sqlite_session) is None -@pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) def test_get_account_by_email_with_case_fallback_uses_lowercase(sqlite_session: Session) -> None: account = Account(name="Case User", email="case@test.com") sqlite_session.add(account) @@ -2696,12 +2890,12 @@ class TestIsEmailSendIpLimit: redis_client.get.side_effect = lambda key: values.get(key) return redis_client - def test_frozen_ip_is_limited(self): + def test_frozen_ip_is_limited(self) -> None: redis_client = self._mock_redis(minute_count=0, hour_count=None, frozen=True) with patch("services.account_service.redis_client", redis_client): assert AccountService.is_email_send_ip_limit("1.2.3.4") is True - def test_first_strike_sets_ten_minute_window(self): + def test_first_strike_sets_ten_minute_window(self) -> None: redis_client = self._mock_redis(minute_count=999, hour_count=None) redis_client.set.return_value = True with ( @@ -2716,7 +2910,7 @@ class TestIsEmailSendIpLimit: redis_client.incr.assert_not_called() redis_client.expire.assert_not_called() - def test_first_strike_lost_claim_freezes_immediately(self): + def test_first_strike_lost_claim_freezes_immediately(self) -> None: redis_client = self._mock_redis(minute_count=999, hour_count=None) redis_client.set.return_value = None # another worker claimed the strike first with ( @@ -2727,7 +2921,7 @@ class TestIsEmailSendIpLimit: redis_client.setex.assert_called_once_with("email_send_ip_limit_freeze:1.2.3.4", 60 * 60, 1) - def test_second_strike_inside_window_freezes_for_an_hour(self): + def test_second_strike_inside_window_freezes_for_an_hour(self) -> None: redis_client = self._mock_redis(minute_count=999, hour_count=1) with ( patch("services.account_service.redis_client", redis_client), @@ -2737,7 +2931,7 @@ class TestIsEmailSendIpLimit: redis_client.setex.assert_called_once_with("email_send_ip_limit_freeze:1.2.3.4", 60 * 60, 1) - def test_under_limit_not_limited(self): + def test_under_limit_not_limited(self) -> None: redis_client = self._mock_redis(minute_count=0, hour_count=None) with ( patch("services.account_service.redis_client", redis_client), diff --git a/api/tests/unit_tests/services/test_agent_app_feature_service.py b/api/tests/unit_tests/services/test_agent_app_feature_service.py index 3d9337d79fb..c638aa2efaa 100644 --- a/api/tests/unit_tests/services/test_agent_app_feature_service.py +++ b/api/tests/unit_tests/services/test_agent_app_feature_service.py @@ -7,14 +7,16 @@ update_features persists those flags as a new app_model_config version without touching model / prompt / agent_mode. """ -from types import SimpleNamespace -from typing import Any - import pytest +from sqlalchemy.orm import Session +from models.account import Account +from models.model import App, AppMode, AppModelConfig from services.agent_app_feature_service import AgentAppFeatureConfigService TENANT_ID = "11111111-1111-1111-1111-111111111111" +APP_ID = "22222222-2222-2222-2222-222222222222" +ACCOUNT_ID = "33333333-3333-3333-3333-333333333333" class TestValidateFeatures: @@ -71,45 +73,47 @@ class TestValidateFeatures: AgentAppFeatureConfigService.validate_features(TENANT_ID, {"suggested_questions": "nope"}) -class _FakeWriteSession: - def __init__(self) -> None: - self.added: list[Any] = [] - self.flushed = 0 - self.committed = 0 - - def add(self, obj: Any) -> None: - self.added.append(obj) - - def flush(self) -> None: - self.flushed += 1 - - def commit(self) -> None: - self.committed += 1 - - class TestUpdateFeatures: - def test_persists_new_app_model_config_version(self): - session = _FakeWriteSession() - app_model = SimpleNamespace( - tenant_id=TENANT_ID, id="app-1", app_model_config_id=None, updated_by=None, updated_at=None + @pytest.mark.parametrize("sqlite_session", [(Account, App, AppModelConfig)], indirect=True) + def test_persists_new_app_model_config_version(self, sqlite_session: Session): + app_model = App( + id=APP_ID, + tenant_id=TENANT_ID, + name="Agent App", + description="", + mode=AppMode.AGENT, + enable_site=True, + enable_api=True, + max_active_requests=0, ) - account = SimpleNamespace(id="acct-1") + account = Account(name="Test User", email="test@example.com") + account.id = ACCOUNT_ID + sqlite_session.add_all([account, app_model]) + sqlite_session.commit() new_config = AgentAppFeatureConfigService.update_features( - app_model=app_model, # type: ignore[arg-type] - account=account, # type: ignore[arg-type] + app_model=app_model, + account=account, config={"opening_statement": "Hi!", "suggested_questions_after_answer": {"enabled": True}}, - session=session, + session=sqlite_session, ) + assert not sqlite_session.in_transaction() # New row carries the features but no Soul-owned model/prompt/agent_mode. - assert new_config.app_id == "app-1" + assert new_config.app_id == APP_ID assert new_config.opening_statement == "Hi!" assert new_config.model is None assert new_config.agent_mode is None # App is repointed at the new version and the write is committed. assert app_model.app_model_config_id == new_config.id - assert app_model.updated_by == "acct-1" - assert new_config in session.added - assert session.flushed == 1 - assert session.committed == 1 + assert app_model.updated_by == ACCOUNT_ID + sqlite_session.expunge_all() + persisted_config = sqlite_session.get(AppModelConfig, new_config.id) + persisted_app = sqlite_session.get(App, APP_ID) + assert persisted_config is not None + assert persisted_config.opening_statement == "Hi!" + assert persisted_config.model is None + assert persisted_config.agent_mode is None + assert persisted_app is not None + assert persisted_app.app_model_config_id == new_config.id + assert persisted_app.updated_by == ACCOUNT_ID diff --git a/api/tests/unit_tests/services/test_agent_app_sandbox_service.py b/api/tests/unit_tests/services/test_agent_app_sandbox_service.py index f3052a4ab49..c2418ecd09a 100644 --- a/api/tests/unit_tests/services/test_agent_app_sandbox_service.py +++ b/api/tests/unit_tests/services/test_agent_app_sandbox_service.py @@ -1,470 +1,1031 @@ -"""Unit tests for the Agent App / workflow sandbox services.""" - -from __future__ import annotations - -from collections.abc import Generator +import json +from contextlib import contextmanager, nullcontext from datetime import datetime +from types import SimpleNamespace +from typing import cast +from unittest.mock import MagicMock import pytest -from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot -from agenton.layers.base import LifecycleState -from dify_agent.protocol import RuntimeLayerSpec, SandboxListResponse, SandboxReadResponse, SandboxUploadResponse -from sqlalchemy import delete +from dify_agent.client import Client +from dify_agent.protocol import BindingFileDownloadResponse, BindingFileListResponse, BindingFileReadResponse +from sqlalchemy.orm import Session, sessionmaker -from core.app.apps.agent_app.session_store import AgentAppSessionScope, StoredAgentAppSession -from core.db.session_factory import session_factory -from models.agent import AgentRuntimeSession, AgentRuntimeSessionOwnerType, AgentRuntimeSessionStatus +from graphon.enums import WorkflowNodeExecutionStatus +from models.agent import ( + Agent, + AgentConfigDraft, + AgentConfigDraftType, + AgentConfigVersionKind, + AgentKind, + AgentScope, + AgentSource, + AgentStatus, + AgentWorkingResourceStatus, + AgentWorkspace, + AgentWorkspaceBinding, + AgentWorkspaceOwnerType, +) +from models.agent_config_entities import AgentSoulConfig +from models.enums import ConversationFromSource, CreatorUserRole +from models.model import App, AppMode, Conversation, IconType +from models.workflow import WorkflowNodeExecutionModel, WorkflowNodeExecutionTriggeredFrom +from services import agent_app_sandbox_service as sandbox_module from services.agent_app_sandbox_service import ( AgentAppSandboxService, + AgentSandboxDownload, AgentSandboxInspectorError, - AgentSandboxUploadDownload, WorkflowAgentSandboxService, - _default_client_factory, - _upload_download_response, ) -def _snapshot( - *, - session_id: str = "abc1234", - shell_runtime_state: dict[str, str] | None = None, -) -> CompositorSessionSnapshot: - runtime_state = ( - {"session_id": session_id, "workspace_cwd": f"~/workspace/{session_id}"} - if shell_runtime_state is None - else shell_runtime_state +def _add_normal_conversation(session: Session, *, binding_id: str) -> Conversation: + session.add( + App( + id="app-1", + tenant_id="tenant-1", + name="Agent App", + description="", + mode=AppMode.AGENT, + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + max_active_requests=0, + ) ) - return CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot(name="execution_context", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={}), - LayerSessionSnapshot( - name="shell", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state=runtime_state, + conversation = Conversation( + id="conversation-1", + app_id="app-1", + mode=AppMode.AGENT, + name="Conversation", + from_source=ConversationFromSource.CONSOLE, + from_account_id="account-1", + is_deleted=False, + agent_workspace_binding_id=binding_id, + ) + conversation._inputs = {} + session.add(conversation) + return conversation + + +def _add_conversation_bindings(session: Session) -> tuple[AgentWorkspaceBinding, AgentWorkspaceBinding]: + workspace = AgentWorkspace( + id="workspace-1", + tenant_id="tenant-1", + app_id="app-1", + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id="conversation-1", + owner_scope_key="root", + backend_workspace_ref="workspace-ref", + status=AgentWorkingResourceStatus.ACTIVE, + active_guard=1, + ) + expected = AgentWorkspaceBinding( + id="binding-expected", + tenant_id="tenant-1", + app_id="app-1", + workspace_id=workspace.id, + agent_id="agent-1", + base_home_snapshot_id="home-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + backend_binding_ref="binding-expected-ref", + status=AgentWorkingResourceStatus.ACTIVE, + updated_at=datetime(2026, 7, 23, 10), + ) + other = AgentWorkspaceBinding( + id="binding-other", + tenant_id="tenant-1", + app_id="app-1", + workspace_id=workspace.id, + agent_id="agent-1", + base_home_snapshot_id="home-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + backend_binding_ref="binding-other-ref", + status=AgentWorkingResourceStatus.ACTIVE, + updated_at=datetime(2026, 7, 23, 11), + ) + session.add_all([workspace, expected, other]) + return expected, other + + +def _use_session(monkeypatch: pytest.MonkeyPatch, session: Session) -> None: + monkeypatch.setattr( + "services.agent_app_sandbox_service.session_factory.create_session", + lambda: nullcontext(session), + ) + + +def _add_app(session: Session, *, app_id: str, tenant_id: str) -> None: + session.add( + App( + id=app_id, + tenant_id=tenant_id, + name=f"App {app_id}", + description="", + mode=AppMode.AGENT, + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + max_active_requests=0, + ) + ) + + +def _add_binding( + session: Session, + *, + binding_id: str, + workspace_id: str, + tenant_id: str = "tenant-1", + app_id: str = "app-1", + agent_id: str = "agent-1", + owner_type: AgentWorkspaceOwnerType, + owner_id: str, + owner_scope_key: str = "root", + status: AgentWorkingResourceStatus = AgentWorkingResourceStatus.ACTIVE, + agent_config_version_id: str | None = None, + agent_config_version_kind: AgentConfigVersionKind = AgentConfigVersionKind.SNAPSHOT, +) -> AgentWorkspaceBinding: + active_guard = 1 if status is AgentWorkingResourceStatus.ACTIVE else None + workspace = AgentWorkspace( + id=workspace_id, + tenant_id=tenant_id, + app_id=app_id, + owner_type=owner_type, + owner_id=owner_id, + owner_scope_key=owner_scope_key, + backend_workspace_ref=f"{workspace_id}-ref", + status=status, + active_guard=active_guard, + ) + binding = AgentWorkspaceBinding( + id=binding_id, + tenant_id=tenant_id, + app_id=app_id, + workspace_id=workspace_id, + agent_id=agent_id, + base_home_snapshot_id=None, + agent_config_version_id=agent_config_version_id or f"{binding_id}-config", + agent_config_version_kind=agent_config_version_kind, + backend_binding_ref=f"{binding_id}-ref", + status=status, + ) + session.add_all([workspace, binding]) + return binding + + +def _add_build_draft_caller( + session: Session, + *, + parent_app_id: str = "app-1", + backing_app_id: str | None = None, + runtime_app_id: str = "app-1", +) -> AgentWorkspaceBinding: + session.add( + Agent( + id="agent-1", + tenant_id="tenant-1", + name="Agent", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.WORKFLOW_ONLY if backing_app_id else AgentScope.ROSTER, + source=AgentSource.WORKFLOW if backing_app_id else AgentSource.AGENT_APP, + app_id=parent_app_id, + backing_app_id=backing_app_id, + status=AgentStatus.ACTIVE, + ) + ) + binding = _add_binding( + session, + binding_id="binding-build", + workspace_id="workspace-build", + app_id=runtime_app_id, + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id="build-1", + agent_config_version_id="config-1", + agent_config_version_kind=AgentConfigVersionKind.DRAFT, + ) + session.add( + AgentConfigDraft( + id="build-1", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + agent_workspace_binding_id=binding.id, + config_snapshot=AgentSoulConfig(), + ) + ) + return binding + + +def _download_client() -> MagicMock: + client = MagicMock() + client.download_binding_file_sync.return_value = BindingFileDownloadResponse(reference="dify-file-ref:canonical") + return client + + +def _stub_download_response(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + sandbox_module, + "_download_response", + lambda **_kwargs: AgentSandboxDownload(url="https://files.example/report.txt"), + ) + + +@pytest.mark.parametrize( + "sqlite_session", + [(AgentWorkspace, AgentWorkspaceBinding, App, Conversation)], + indirect=True, +) +def test_agent_app_file_browsing_uses_conversation_pointer( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + expected, _ = _add_conversation_bindings(sqlite_session) + _add_normal_conversation(sqlite_session, binding_id=expected.id) + sqlite_session.commit() + _use_session(monkeypatch, sqlite_session) + client = MagicMock() + response = BindingFileListResponse(path=".", entries=[], truncated=False) + client.list_binding_files_sync.return_value = response + + result = AgentAppSandboxService(client_factory=lambda: nullcontext(client)).list_files( + tenant_id="tenant-1", + app_id="app-1", + agent_id="agent-1", + caller_type="conversation", + caller_id="conversation-1", + account_id="account-1", + path=".", + ) + + assert result is response + client.list_binding_files_sync.assert_called_once_with(expected.backend_binding_ref, ".") + + +@pytest.mark.parametrize( + "sqlite_session", + [(AgentWorkspace, AgentWorkspaceBinding, App, Conversation)], + indirect=True, +) +def test_agent_app_file_browsing_rejects_other_account( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: + expected, _ = _add_conversation_bindings(sqlite_session) + _add_normal_conversation(sqlite_session, binding_id=expected.id) + sqlite_session.commit() + _use_session(monkeypatch, sqlite_session) + client = MagicMock() + + with pytest.raises(AgentSandboxInspectorError) as exc_info: + AgentAppSandboxService(client_factory=lambda: nullcontext(client)).list_files( + tenant_id="tenant-1", + app_id="app-1", + agent_id="agent-1", + caller_type="conversation", + caller_id="conversation-1", + account_id="account-2", + path=".", + ) + + assert exc_info.value.code == "no_active_binding" + client.list_binding_files_sync.assert_not_called() + + +@pytest.mark.parametrize( + "sqlite_session", + [(AgentWorkspace, AgentWorkspaceBinding, App, Conversation)], + indirect=True, +) +def test_agent_conversation_download_resolves_only_exact_active_owner_chain( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + _add_app(sqlite_session, app_id="app-1", tenant_id="tenant-1") + _add_app(sqlite_session, app_id="app-other", tenant_id="tenant-other") + valid = _add_binding( + sqlite_session, + binding_id="binding-valid", + workspace_id="workspace-valid", + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id="conversation-valid", + ) + wrong_owner = _add_binding( + sqlite_session, + binding_id="binding-wrong-owner", + workspace_id="workspace-wrong-owner", + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id="conversation-not-the-caller", + ) + retired = _add_binding( + sqlite_session, + binding_id="binding-retired", + workspace_id="workspace-retired", + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id="conversation-retired", + status=AgentWorkingResourceStatus.RETIRED, + ) + cross_tenant = _add_binding( + sqlite_session, + binding_id="binding-cross-tenant", + workspace_id="workspace-cross-tenant", + tenant_id="tenant-other", + app_id="app-other", + owner_type=AgentWorkspaceOwnerType.CONVERSATION, + owner_id="conversation-cross-tenant", + ) + conversations = [ + Conversation( + id="conversation-valid", + app_id="app-1", + mode=AppMode.AGENT, + name="Valid", + from_source=ConversationFromSource.CONSOLE, + from_account_id="account-1", + is_deleted=False, + agent_workspace_binding_id=valid.id, + ), + Conversation( + id="conversation-wrong-owner", + app_id="app-1", + mode=AppMode.AGENT, + name="Wrong owner", + from_source=ConversationFromSource.CONSOLE, + from_account_id="account-1", + is_deleted=False, + agent_workspace_binding_id=wrong_owner.id, + ), + Conversation( + id="conversation-retired", + app_id="app-1", + mode=AppMode.AGENT, + name="Retired", + from_source=ConversationFromSource.CONSOLE, + from_account_id="account-1", + is_deleted=False, + agent_workspace_binding_id=retired.id, + ), + Conversation( + id="conversation-cross-tenant", + app_id="app-other", + mode=AppMode.AGENT, + name="Cross tenant", + from_source=ConversationFromSource.CONSOLE, + from_account_id="account-other", + is_deleted=False, + agent_workspace_binding_id=cross_tenant.id, + ), + ] + for conversation in conversations: + conversation._inputs = {} + sqlite_session.add_all(conversations) + sqlite_session.commit() + _use_session(monkeypatch, sqlite_session) + _stub_download_response(monkeypatch) + client = _download_client() + service = AgentAppSandboxService(client_factory=lambda: nullcontext(cast(Client, client))) + + result = service.download_file( + tenant_id="tenant-1", + app_id="app-1", + agent_id="agent-1", + caller_type="conversation", + caller_id="conversation-valid", + account_id="account-1", + path="report.txt", + ) + + assert result.url == "https://files.example/report.txt" + request = client.download_binding_file_sync.call_args.args[0] + assert request.backend_binding_ref == "binding-valid-ref" + client.download_binding_file_sync.reset_mock() + + rejected_locators = [ + {"account_id": "account-other"}, + {"app_id": "app-other"}, + {"caller_id": "conversation-wrong-owner"}, + {"caller_id": "conversation-retired"}, + { + "tenant_id": "tenant-1", + "app_id": "app-other", + "caller_id": "conversation-cross-tenant", + "account_id": "account-other", + }, + ] + for override in rejected_locators: + locator = { + "tenant_id": "tenant-1", + "app_id": "app-1", + "agent_id": "agent-1", + "caller_type": "conversation", + "caller_id": "conversation-valid", + "account_id": "account-1", + "path": "report.txt", + } + locator.update(override) + with pytest.raises(AgentSandboxInspectorError, match="active Agent Workspace Binding"): + service.download_file(**locator) # type: ignore[arg-type] + + client.download_binding_file_sync.assert_not_called() + + +@pytest.mark.parametrize( + "sqlite_session", + [(Agent, AgentConfigDraft, AgentWorkspace, AgentWorkspaceBinding)], + indirect=True, +) +def test_agent_build_draft_download_resolves_only_exact_active_owner_chain( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + sqlite_session.add_all( + [ + Agent( + id="agent-1", + tenant_id="tenant-1", + name="Agent", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + app_id="app-1", + status=AgentStatus.ACTIVE, + ), + Agent( + id="agent-cross-tenant", + tenant_id="tenant-other", + name="Other tenant Agent", + description="", + agent_kind=AgentKind.DIFY_AGENT, + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + app_id="app-other", + status=AgentStatus.ACTIVE, ), ] ) - - -def _runtime_layer_specs() -> list[RuntimeLayerSpec]: - return [ - RuntimeLayerSpec(name="execution_context", type="dify.execution_context", config={"tenant_id": "tenant-1"}), - RuntimeLayerSpec(name="shell", type="dify.shell", deps={"execution_context": "execution_context"}, config={}), - ] - - -class FakeStore: - def __init__(self, session: StoredAgentAppSession | None) -> None: - self.session = session - self.scope: tuple[str, str, str] | None = None - - def load_active_session_for_conversation(self, *, tenant_id: str, app_id: str, conversation_id: str): - self.scope = (tenant_id, app_id, conversation_id) - return self.session - - -class FakeClient: - def __init__(self) -> None: - self.calls: list[tuple[str, str]] = [] - self.locators: list[object] = [] - - def list_sandbox_files_sync(self, locator, path: str) -> SandboxListResponse: - self.locators.append(locator) - self.calls.append(("list", path)) - return SandboxListResponse(path=path, entries=[], truncated=False) - - def read_sandbox_file_sync(self, locator, path: str, max_bytes: int = 262144) -> SandboxReadResponse: - del max_bytes - self.locators.append(locator) - self.calls.append(("read", path)) - return SandboxReadResponse(path=path, size=5, truncated=False, binary=False, text="hello") - - def upload_sandbox_file_sync(self, locator, path: str) -> SandboxUploadResponse: - self.locators.append(locator) - self.calls.append(("upload", path)) - return SandboxUploadResponse( - path=path, - file={ - "transfer_method": "tool_file", - "reference": "dify-file-ref:file-1", - "download_url": "https://files.example/report.txt?token=1", - }, - ) - - -def _stored_session( - *, - session_snapshot: CompositorSessionSnapshot | None = None, -) -> StoredAgentAppSession: - return StoredAgentAppSession( - scope=AgentAppSessionScope( + valid = _add_binding( + sqlite_session, + binding_id="binding-build-valid", + workspace_id="workspace-build-valid", + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id="draft-valid", + ) + wrong_owner = _add_binding( + sqlite_session, + binding_id="binding-build-wrong-owner", + workspace_id="workspace-build-wrong-owner", + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id="draft-not-the-caller", + ) + retired = _add_binding( + sqlite_session, + binding_id="binding-build-retired", + workspace_id="workspace-build-retired", + owner_type=AgentWorkspaceOwnerType.BUILD_DRAFT, + owner_id="draft-retired", + status=AgentWorkingResourceStatus.RETIRED, + ) + drafts = [ + AgentConfigDraft( + id="draft-valid", tenant_id="tenant-1", - app_id="app-1", - conversation_id="conv-1", agent_id="agent-1", - agent_config_snapshot_id="snapshot-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-1", + draft_owner_key="account-1", + agent_workspace_binding_id=valid.id, + config_snapshot=AgentSoulConfig(), ), - session_snapshot=_snapshot() if session_snapshot is None else session_snapshot, - backend_run_id="run-1", - runtime_layer_specs=_runtime_layer_specs(), + AgentConfigDraft( + id="draft-wrong-owner", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-2", + draft_owner_key="account-2", + agent_workspace_binding_id=wrong_owner.id, + config_snapshot=AgentSoulConfig(), + ), + AgentConfigDraft( + id="draft-retired", + tenant_id="tenant-1", + agent_id="agent-1", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-3", + draft_owner_key="account-3", + agent_workspace_binding_id=retired.id, + config_snapshot=AgentSoulConfig(), + ), + AgentConfigDraft( + id="draft-cross-tenant", + tenant_id="tenant-other", + agent_id="agent-cross-tenant", + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id="account-other", + draft_owner_key="account-other", + agent_workspace_binding_id=None, + config_snapshot=AgentSoulConfig(), + ), + ] + sqlite_session.add_all(drafts) + sqlite_session.commit() + _use_session(monkeypatch, sqlite_session) + _stub_download_response(monkeypatch) + client = _download_client() + service = AgentAppSandboxService(client_factory=lambda: nullcontext(cast(Client, client))) + + result = service.download_file( + tenant_id="tenant-1", + app_id="app-1", + agent_id="agent-1", + caller_type="build_draft", + caller_id="draft-valid", + account_id="account-1", + path="report.txt", ) + assert result.url == "https://files.example/report.txt" + assert client.download_binding_file_sync.call_args.args[0].backend_binding_ref == "binding-build-valid-ref" + client.download_binding_file_sync.reset_mock() -def test_agent_app_sandbox_service_get_info_returns_metadata() -> None: - store = FakeStore(_stored_session()) - client = FakeClient() - service = AgentAppSandboxService(session_store=store, client_factory=lambda: client) # type: ignore[arg-type] - - result = service.get_info(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1") - - assert result.session_id == "abc1234" - assert result.workspace_cwd == "~/workspace/abc1234" - assert client.calls == [] - assert store.scope == ("tenant-1", "app-1", "conv-1") - - -def test_agent_app_sandbox_service_builds_locator_and_proxies() -> None: - store = FakeStore(_stored_session()) - client = FakeClient() - service = AgentAppSandboxService(session_store=store, client_factory=lambda: client) # type: ignore[arg-type] - - result = service.list_files(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1", path=".") - - assert result.path == "." - assert client.calls == [("list", ".")] - assert store.scope == ("tenant-1", "app-1", "conv-1") - - -def test_agent_app_sandbox_service_upload_returns_download_url(monkeypatch: pytest.MonkeyPatch) -> None: - store = FakeStore(_stored_session()) - client = FakeClient() - captured: dict[str, object] = {} - - def fake_upload_download_response(*, tenant_id: str, file_mapping: dict[str, object]) -> AgentSandboxUploadDownload: - captured["tenant_id"] = tenant_id - captured["file_mapping"] = file_mapping - return AgentSandboxUploadDownload(url="https://files.example/report.txt?token=1&as_attachment=true") - - monkeypatch.setattr("services.agent_app_sandbox_service._upload_download_response", fake_upload_download_response) - service = AgentAppSandboxService(session_store=store, client_factory=lambda: client) # type: ignore[arg-type] - - result = service.upload_file(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1", path="report.txt") - - assert result.url == "https://files.example/report.txt?token=1&as_attachment=true" - assert client.calls == [("upload", "report.txt")] - assert store.scope == ("tenant-1", "app-1", "conv-1") - assert captured == { - "tenant_id": "tenant-1", - "file_mapping": { - "transfer_method": "tool_file", - "reference": "dify-file-ref:file-1", - "download_url": "https://files.example/report.txt?token=1", + rejected_locators = [ + {"account_id": "account-other"}, + {"app_id": "app-other"}, + {"caller_id": "draft-wrong-owner", "account_id": "account-2"}, + {"caller_id": "draft-retired", "account_id": "account-3"}, + { + "tenant_id": "tenant-1", + "app_id": "app-other", + "agent_id": "agent-cross-tenant", + "caller_id": "draft-cross-tenant", + "account_id": "account-other", }, - } + ] + for override in rejected_locators: + locator = { + "tenant_id": "tenant-1", + "app_id": "app-1", + "agent_id": "agent-1", + "caller_type": "build_draft", + "caller_id": "draft-valid", + "account_id": "account-1", + "path": "report.txt", + } + locator.update(override) + with pytest.raises(AgentSandboxInspectorError, match="active Agent Workspace Binding"): + service.download_file(**locator) # type: ignore[arg-type] + + client.download_binding_file_sync.assert_not_called() -def test_agent_app_sandbox_service_raises_when_no_active_session() -> None: - service = AgentAppSandboxService(session_store=FakeStore(None), client_factory=lambda: FakeClient()) # type: ignore[arg-type] - - with pytest.raises(AgentSandboxInspectorError) as exc_info: - service.get_info(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1") - - assert exc_info.value.code == "no_active_session" - assert exc_info.value.status_code == 404 - - -def test_agent_app_sandbox_service_raises_when_shell_workspace_metadata_missing() -> None: - broken_session = _stored_session(session_snapshot=_snapshot(shell_runtime_state={"session_id": "abc1234"})) - service = AgentAppSandboxService(session_store=FakeStore(broken_session), client_factory=lambda: FakeClient()) # type: ignore[arg-type] - - with pytest.raises(AgentSandboxInspectorError) as exc_info: - service.get_info(tenant_id="tenant-1", app_id="app-1", conversation_id="conv-1") - - assert exc_info.value.code == "no_sandbox" - assert exc_info.value.status_code == 404 - - -def test_default_client_factory_requires_agent_backend_base_url(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr("services.agent_app_sandbox_service.dify_config.AGENT_BACKEND_BASE_URL", "") - - with pytest.raises(AgentSandboxInspectorError) as exc_info: - _default_client_factory() - - assert exc_info.value.code == "inspector_unavailable" - assert exc_info.value.status_code == 503 - - -@pytest.fixture -def _runtime_session_table() -> Generator[None, None, None]: - engine = session_factory.get_session_maker().kw["bind"] - AgentRuntimeSession.__table__.create(bind=engine, checkfirst=True) - yield - with session_factory.create_session() as session: - session.execute(delete(AgentRuntimeSession)) - session.commit() - AgentRuntimeSession.__table__.drop(bind=engine, checkfirst=True) - - -def _insert_workflow_session( +def _workflow_execution( *, - runtime_layer_specs: str | None = None, + execution_id: str, + tenant_id: str = "tenant-1", + app_id: str = "app-1", workflow_run_id: str = "run-1", node_id: str = "node-1", - node_execution_id: str = "node-exec-1", - binding_id: str = "binding-1", - backend_run_id: str = "backend-run-1", - updated_at: datetime | None = None, - session_id: str = "abc1234", -) -> None: - default_runtime_layer_specs = ( - '[{"name":"execution_context","type":"dify.execution_context","config":{"tenant_id":"tenant-1"}},' - '{"name":"shell","type":"dify.shell","deps":{"execution_context":"execution_context"},"config":{}}]' + binding_id: str | None, + workflow_agent_binding_id: str | None = "workflow-binding-1", + created_by: str = "historical-account", +) -> WorkflowNodeExecutionModel: + return WorkflowNodeExecutionModel( + id=execution_id, + tenant_id=tenant_id, + app_id=app_id, + workflow_id="workflow-1", + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + workflow_run_id=workflow_run_id, + index=1, + node_id=node_id, + node_type="agent", + title=node_id, + agent_workspace_binding_id=binding_id, + inputs=None, + process_data=json.dumps( + {"workflow_agent_binding_id": workflow_agent_binding_id} if workflow_agent_binding_id is not None else {} + ), + outputs=None, + status=WorkflowNodeExecutionStatus.SUCCEEDED, + error=None, + created_by_role=CreatorUserRole.ACCOUNT, + created_by=created_by, ) - with session_factory.create_session() as session: - row = AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.WORKFLOW_RUN, - workflow_id="workflow-1", - workflow_run_id=workflow_run_id, - node_id=node_id, - node_execution_id=node_execution_id, - binding_id=binding_id, - agent_id="agent-1", - agent_config_snapshot_id="snapshot-1", - backend_run_id=backend_run_id, - session_snapshot=_snapshot(session_id=session_id).model_dump_json(), - composition_layer_specs=runtime_layer_specs or default_runtime_layer_specs, - status=AgentRuntimeSessionStatus.ACTIVE, - ) - if updated_at is not None: - row.updated_at = updated_at - session.add(row) - session.commit() -@pytest.mark.usefixtures("_runtime_session_table") -def test_workflow_sandbox_service_resolves_locator_and_returns_download_url( +@pytest.mark.parametrize( + "sqlite_session", + [(WorkflowNodeExecutionModel, AgentWorkspace, AgentWorkspaceBinding)], + indirect=True, +) +def test_workflow_download_resolves_only_exact_active_owner_chain( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: - _insert_workflow_session() - client = FakeClient() - captured: dict[str, object] = {} + valid = _add_binding( + sqlite_session, + binding_id="binding-workflow-valid", + workspace_id="workspace-workflow-valid", + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id="run-1", + owner_scope_key="node-1:workflow-binding-1", + ) + wrong_owner = _add_binding( + sqlite_session, + binding_id="binding-workflow-wrong-owner", + workspace_id="workspace-workflow-wrong-owner", + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id="run-not-the-caller", + owner_scope_key="node-1:workflow-binding-1", + ) + retired = AgentWorkspaceBinding( + id="binding-workflow-retired", + tenant_id="tenant-1", + app_id="app-1", + workspace_id=valid.workspace_id, + agent_id="agent-1", + base_home_snapshot_id=None, + agent_config_version_id="binding-workflow-retired-config", + agent_config_version_kind=AgentConfigVersionKind.SNAPSHOT, + backend_binding_ref="binding-workflow-retired-ref", + status=AgentWorkingResourceStatus.RETIRED, + ) + sqlite_session.add(retired) + sqlite_session.add_all( + [ + _workflow_execution(execution_id="execution-valid", binding_id=valid.id), + _workflow_execution( + execution_id="execution-cross-tenant", + tenant_id="tenant-other", + binding_id=valid.id, + ), + _workflow_execution(execution_id="execution-wrong-app", app_id="app-other", binding_id=valid.id), + _workflow_execution(execution_id="execution-wrong-run", workflow_run_id="run-other", binding_id=valid.id), + _workflow_execution(execution_id="execution-wrong-node", node_id="node-other", binding_id=valid.id), + _workflow_execution(execution_id="execution-wrong-owner", binding_id=wrong_owner.id), + _workflow_execution(execution_id="execution-retired", binding_id=retired.id), + ] + ) + sqlite_session.commit() + persisted_execution = sqlite_session.get(WorkflowNodeExecutionModel, "execution-valid") + assert persisted_execution is not None + assert persisted_execution.created_by == "historical-account" + _use_session(monkeypatch, sqlite_session) + client = _download_client() + request_download = MagicMock( + return_value=SimpleNamespace(download_uri="/files/tools/report.txt?timestamp=1&sign=2") + ) + monkeypatch.setattr( + sandbox_module, + "FileRequestService", + lambda: SimpleNamespace(request_download=request_download), + ) + monkeypatch.setattr(sandbox_module.dify_config, "FILES_URL", "https://files.example") + service = WorkflowAgentSandboxService(client_factory=lambda: nullcontext(cast(Client, client))) - def fake_upload_download_response(*, tenant_id: str, file_mapping: dict[str, object]) -> AgentSandboxUploadDownload: - captured["tenant_id"] = tenant_id - captured["file_mapping"] = file_mapping - return AgentSandboxUploadDownload(url="https://files.example/report.txt?token=1&as_attachment=true") - - monkeypatch.setattr("services.agent_app_sandbox_service._upload_download_response", fake_upload_download_response) - service = WorkflowAgentSandboxService(client_factory=lambda: client) # type: ignore[arg-type] - - result = service.upload_file( + result = service.download_file( tenant_id="tenant-1", app_id="app-1", workflow_run_id="run-1", node_id="node-1", - node_execution_id="node-exec-1", + node_execution_id="execution-valid", + account_id="authenticated-account", path="report.txt", - session=session_factory.create_session(), ) - assert result.url == "https://files.example/report.txt?token=1&as_attachment=true" - assert client.calls == [("upload", "report.txt")] - assert captured == { - "tenant_id": "tenant-1", - "file_mapping": { - "transfer_method": "tool_file", - "reference": "dify-file-ref:file-1", - "download_url": "https://files.example/report.txt?token=1", - }, - } - - -def test_upload_download_response_resolves_signed_external_url(monkeypatch: pytest.MonkeyPatch) -> None: - built_file = object() - built_with: dict[str, object] = {} - - def fake_build_from_mapping(*, mapping: dict[str, object], tenant_id: str, access_controller: object) -> object: - built_with["mapping"] = mapping - built_with["tenant_id"] = tenant_id - built_with["access_controller"] = access_controller - return built_file - - class FakeRuntime: - def __init__(self, *, file_access_controller: object) -> None: - self.file_access_controller = file_access_controller - - def resolve_file_url(self, *, file: object, for_external: bool) -> str: - assert file is built_file - assert for_external is True - return "https://files.example/files/tools/tool-file.txt?timestamp=1&nonce=2&sign=3" - - monkeypatch.setattr("services.agent_app_sandbox_service.file_factory.build_from_mapping", fake_build_from_mapping) - monkeypatch.setattr("services.agent_app_sandbox_service.DifyWorkflowFileRuntime", FakeRuntime) - - result = _upload_download_response( + assert result.url == "https://files.example/files/tools/report.txt?timestamp=1&sign=2&as_attachment=true" + request = client.download_binding_file_sync.call_args.args[0] + assert request.backend_binding_ref == "binding-workflow-valid-ref" + assert request.execution_context.user_id == "authenticated-account" + assert request.execution_context.user_from == "account" + request_download.assert_called_once_with( tenant_id="tenant-1", - file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:file-1"}, + user_id="authenticated-account", + user_from="account", + invoke_from="debugger", + file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:canonical"}, ) + client.download_binding_file_sync.reset_mock() - assert result.url == ( - "https://files.example/files/tools/tool-file.txt?timestamp=1&nonce=2&sign=3&as_attachment=true" - ) - assert built_with["mapping"] == {"transfer_method": "tool_file", "reference": "dify-file-ref:file-1"} - assert built_with["tenant_id"] == "tenant-1" - assert built_with["access_controller"] is not None + for node_execution_id in ( + "execution-cross-tenant", + "execution-wrong-app", + "execution-wrong-run", + "execution-wrong-node", + "execution-wrong-owner", + "execution-retired", + ): + with pytest.raises(AgentSandboxInspectorError, match="active Workspace Binding"): + service.download_file( + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + node_id="node-1", + node_execution_id=node_execution_id, + account_id="authenticated-account", + path="report.txt", + ) + + client.download_binding_file_sync.assert_not_called() + assert request_download.call_count == 1 -def test_upload_download_response_maps_resolution_failure_to_inspector_error( +@pytest.mark.parametrize( + ("parent_app_id", "backing_app_id", "runtime_app_id"), + [ + ("app-1", None, "app-1"), + ("workflow-app-1", "runtime-app-1", "runtime-app-1"), + ], +) +def test_agent_app_file_browsing_uses_build_draft_caller( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + parent_app_id: str, + backing_app_id: str | None, + runtime_app_id: str, ) -> None: - def fake_build_from_mapping(*, mapping: dict[str, object], tenant_id: str, access_controller: object) -> object: - del mapping, tenant_id, access_controller - raise ValueError("missing tool file") - - monkeypatch.setattr("services.agent_app_sandbox_service.file_factory.build_from_mapping", fake_build_from_mapping) - - with pytest.raises(AgentSandboxInspectorError) as exc_info: - _upload_download_response( - tenant_id="tenant-1", - file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:file-1"}, - ) - - assert exc_info.value.code == "sandbox_upload_download_unavailable" - assert exc_info.value.status_code == 502 - - -def test_upload_download_response_maps_missing_url_to_inspector_error( - monkeypatch: pytest.MonkeyPatch, -) -> None: - built_file = object() - - def fake_build_from_mapping(*, mapping: dict[str, object], tenant_id: str, access_controller: object) -> object: - del mapping, tenant_id, access_controller - return built_file - - class FakeRuntime: - def __init__(self, *, file_access_controller: object) -> None: - self.file_access_controller = file_access_controller - - def resolve_file_url(self, *, file: object, for_external: bool) -> None: - assert file is built_file - assert for_external is True - - monkeypatch.setattr("services.agent_app_sandbox_service.file_factory.build_from_mapping", fake_build_from_mapping) - monkeypatch.setattr("services.agent_app_sandbox_service.DifyWorkflowFileRuntime", FakeRuntime) - - with pytest.raises(AgentSandboxInspectorError) as exc_info: - _upload_download_response( - tenant_id="tenant-1", - file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:file-1"}, - ) - - assert exc_info.value.code == "sandbox_upload_download_unavailable" - assert exc_info.value.status_code == 502 - - -@pytest.mark.usefixtures("_runtime_session_table") -def test_workflow_sandbox_service_filters_by_node_execution_id() -> None: - _insert_workflow_session( - node_execution_id="node-exec-1", - binding_id="binding-1", - backend_run_id="run-a", - session_id="abc1234", + _add_build_draft_caller( + sqlite_session, + parent_app_id=parent_app_id, + backing_app_id=backing_app_id, + runtime_app_id=runtime_app_id, ) - _insert_workflow_session( - node_execution_id="node-exec-2", - binding_id="binding-2", - backend_run_id="run-b", - session_id="def5678", - ) - client = FakeClient() - service = WorkflowAgentSandboxService(client_factory=lambda: client) # type: ignore[arg-type] + sqlite_session.commit() + _use_session(monkeypatch, sqlite_session) + client = MagicMock() + response = BindingFileListResponse(path=".", entries=[], truncated=False) + client.list_binding_files_sync.return_value = response - result = service.read_file( + result = AgentAppSandboxService(client_factory=lambda: nullcontext(client)).list_files( tenant_id="tenant-1", - app_id="app-1", - workflow_run_id="run-1", - node_id="node-1", - node_execution_id="node-exec-2", - path="out.txt", - session=session_factory.create_session(), - ) - - assert result.text == "hello" - assert client.calls == [("read", "out.txt")] - assert client.locators[0].session_snapshot.layers[1].runtime_state["session_id"] == "def5678" - - -@pytest.mark.usefixtures("_runtime_session_table") -def test_workflow_sandbox_service_uses_latest_active_session_when_execution_id_omitted() -> None: - _insert_workflow_session( - node_execution_id="node-exec-1", - binding_id="binding-1", - backend_run_id="run-older", - updated_at=datetime(2026, 1, 1, 0, 0, 0), - session_id="abc1234", - ) - _insert_workflow_session( - node_execution_id="node-exec-2", - binding_id="binding-2", - backend_run_id="run-newer", - updated_at=datetime(2026, 1, 1, 0, 0, 1), - session_id="def5678", - ) - client = FakeClient() - service = WorkflowAgentSandboxService(client_factory=lambda: client) # type: ignore[arg-type] - - result = service.list_files( - tenant_id="tenant-1", - app_id="app-1", - workflow_run_id="run-1", - node_id="node-1", - node_execution_id=None, + app_id=runtime_app_id, + agent_id="agent-1", + caller_type="build_draft", + caller_id="build-1", + account_id="account-1", path=".", - session=session_factory.create_session(), ) - assert result.path == "." - assert client.calls == [("list", ".")] - assert client.locators[0].session_snapshot.layers[1].runtime_state["session_id"] == "def5678" + assert result is response + client.list_binding_files_sync.assert_called_once_with("binding-build-ref", ".") -@pytest.mark.usefixtures("_runtime_session_table") -def test_workflow_sandbox_service_raises_when_no_active_session() -> None: - service = WorkflowAgentSandboxService(client_factory=lambda: FakeClient()) # type: ignore[arg-type] +def test_workflow_file_access_uses_node_execution_pointer( + sqlite_session: Session, +) -> None: + binding = _add_binding( + sqlite_session, + binding_id="binding-workflow", + workspace_id="workspace-workflow", + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id="run-1", + owner_scope_key="node-1:workflow-binding-1", + ) + sqlite_session.add(_workflow_execution(execution_id="execution-1", binding_id=binding.id)) + sqlite_session.commit() + client = MagicMock() + response = BindingFileReadResponse(path="report.txt", size=2, truncated=False, binary=False, text="ok") + client.read_binding_file_sync.return_value = response + + result = WorkflowAgentSandboxService(client_factory=lambda: nullcontext(client)).read_file( + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + node_id="node-1", + node_execution_id="execution-1", + path="report.txt", + session=sqlite_session, + ) + + assert result is response + client.read_binding_file_sync.assert_called_once_with("binding-workflow-ref", "report.txt") + assert not sqlite_session.in_transaction() + + +def test_workflow_download_uses_authenticated_account_and_trusted_file_request( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + binding = _add_binding( + sqlite_session, + binding_id="binding-workflow", + workspace_id="workspace-workflow", + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id="run-1", + owner_scope_key="node-1:workflow-binding-1", + ) + sqlite_session.add(_workflow_execution(execution_id="execution-1", binding_id=binding.id)) + sqlite_session.commit() + events: list[str] = [] + + @contextmanager + def session_scope(): + with sqlite_session_factory() as service_session: + yield service_session + assert not service_session.in_transaction() + events.append("session-exit") + + monkeypatch.setattr(sandbox_module.session_factory, "create_session", session_scope) + client = MagicMock() + client.download_binding_file_sync.side_effect = lambda _request: ( + events.append("client-download") or BindingFileDownloadResponse(reference="dify-file-ref:canonical") + ) + request_download = MagicMock( + side_effect=lambda **_kwargs: ( + events.append("file-request") or SimpleNamespace(download_uri="/files/tools/report.txt?timestamp=1&sign=2") + ) + ) + monkeypatch.setattr( + "services.agent_app_sandbox_service.FileRequestService", + lambda: SimpleNamespace(request_download=request_download), + ) + monkeypatch.setattr("services.agent_app_sandbox_service.dify_config.FILES_URL", "https://files.example") + + result = WorkflowAgentSandboxService(client_factory=lambda: nullcontext(client)).download_file( + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + node_id="node-1", + node_execution_id="execution-1", + account_id="account-1", + path="report.txt", + ) + + request = client.download_binding_file_sync.call_args.args[0] + assert request.execution_context.user_id == "account-1" + assert request.execution_context.user_from == "account" + assert request.execution_context.node_execution_id == "execution-1" + assert events == ["session-exit", "client-download", "file-request"] + request_download.assert_called_once_with( + tenant_id="tenant-1", + user_id="account-1", + user_from="account", + invoke_from="debugger", + file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:canonical"}, + ) + assert result.url == "https://files.example/files/tools/report.txt?timestamp=1&sign=2&as_attachment=true" + + +def test_agent_app_download_uses_complete_account_context_after_session_exit( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: + events: list[str] = [] + _add_build_draft_caller(sqlite_session) + sqlite_session.commit() + + @contextmanager + def session_scope(): + with sqlite_session_factory() as service_session: + yield service_session + events.append("session-exit") + + monkeypatch.setattr(sandbox_module.session_factory, "create_session", session_scope) + client = MagicMock() + client.download_binding_file_sync.side_effect = lambda _request: ( + events.append("client-download") or BindingFileDownloadResponse(reference="dify-file-ref:canonical") + ) + request_download = MagicMock( + side_effect=lambda **_kwargs: ( + events.append("file-request") or SimpleNamespace(download_uri="/files/tools/report.txt?timestamp=1&sign=2") + ) + ) + monkeypatch.setattr( + sandbox_module, + "FileRequestService", + lambda: SimpleNamespace(request_download=request_download), + ) + monkeypatch.setattr(sandbox_module.dify_config, "FILES_URL", "https://files.example") + + result = AgentAppSandboxService(client_factory=lambda: nullcontext(client)).download_file( + tenant_id="tenant-1", + app_id="app-1", + agent_id="agent-1", + caller_type="build_draft", + caller_id="build-1", + account_id="account-1", + path="report.txt", + ) + + request = client.download_binding_file_sync.call_args.args[0] + assert request.backend_binding_ref == "binding-build-ref" + assert request.path == "report.txt" + assert request.execution_context.model_dump(exclude_none=True) == { + "tenant_id": "tenant-1", + "user_id": "account-1", + "user_from": "account", + "app_id": "app-1", + "agent_id": "agent-1", + "agent_config_version_id": "config-1", + "agent_config_version_kind": "draft", + "agent_mode": "agent_app", + "invoke_from": "debugger", + } + assert events == ["session-exit", "client-download", "file-request"] + request_download.assert_called_once_with( + tenant_id="tenant-1", + user_id="account-1", + user_from="account", + invoke_from="debugger", + file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:canonical"}, + ) + assert result.url == "https://files.example/files/tools/report.txt?timestamp=1&sign=2&as_attachment=true" + + +def test_file_request_rejection_maps_to_download_unavailable(monkeypatch: pytest.MonkeyPatch) -> None: + request_download = MagicMock(side_effect=ValueError("reference is not accessible")) + monkeypatch.setattr( + sandbox_module, + "FileRequestService", + lambda: SimpleNamespace(request_download=request_download), + ) with pytest.raises(AgentSandboxInspectorError) as exc_info: - service.list_files( + sandbox_module._download_response( + tenant_id="tenant-1", + account_id="account-1", + reference="dify-file-ref:untrusted", + ) + + assert exc_info.value.code == "binding_file_download_unavailable" + assert exc_info.value.status_code == 502 + + +@pytest.mark.parametrize( + ("binding_id", "workflow_agent_binding_id"), + [ + pytest.param(None, "workflow-binding-1", id="missing-binding-pointer"), + pytest.param("binding-workflow", None, id="missing-process-data"), + ], +) +def test_workflow_download_rejects_missing_binding_metadata_before_network( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + binding_id: str | None, + workflow_agent_binding_id: str | None, +) -> None: + sqlite_session.add( + _workflow_execution( + execution_id="execution-1", + binding_id=binding_id, + workflow_agent_binding_id=workflow_agent_binding_id, + ) + ) + sqlite_session.commit() + _use_session(monkeypatch, sqlite_session) + client = MagicMock() + + with pytest.raises(AgentSandboxInspectorError) as exc_info: + WorkflowAgentSandboxService(client_factory=lambda: nullcontext(client)).download_file( tenant_id="tenant-1", app_id="app-1", workflow_run_id="run-1", node_id="node-1", - node_execution_id=None, - path=".", - session=session_factory.create_session(), + node_execution_id="execution-1", + account_id="account-1", + path="report.txt", ) - assert exc_info.value.code == "no_active_session" - assert exc_info.value.status_code == 404 + assert exc_info.value.code == "no_active_binding" + client.download_binding_file_sync.assert_not_called() -@pytest.mark.usefixtures("_runtime_session_table") -def test_workflow_sandbox_service_raises_when_runtime_specs_missing() -> None: - _insert_workflow_session(runtime_layer_specs="[]") - service = WorkflowAgentSandboxService(client_factory=lambda: FakeClient()) # type: ignore[arg-type] +def test_workflow_download_rejects_non_active_or_mismatched_binding_before_network( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, +) -> None: + binding = _add_binding( + sqlite_session, + binding_id="binding-workflow", + workspace_id="workspace-workflow", + owner_type=AgentWorkspaceOwnerType.WORKFLOW_RUN, + owner_id="run-other", + owner_scope_key="node-1:workflow-binding-1", + ) + sqlite_session.add(_workflow_execution(execution_id="execution-1", binding_id=binding.id)) + sqlite_session.commit() + _use_session(monkeypatch, sqlite_session) + client = MagicMock() with pytest.raises(AgentSandboxInspectorError) as exc_info: - service.list_files( + WorkflowAgentSandboxService(client_factory=lambda: nullcontext(client)).download_file( tenant_id="tenant-1", app_id="app-1", workflow_run_id="run-1", node_id="node-1", - node_execution_id=None, - path=".", - session=session_factory.create_session(), + node_execution_id="execution-1", + account_id="account-1", + path="report.txt", ) - assert exc_info.value.code == "no_sandbox" + assert exc_info.value.code == "no_active_binding" + client.download_binding_file_sync.assert_not_called() diff --git a/api/tests/unit_tests/services/test_agent_config_service.py b/api/tests/unit_tests/services/test_agent_config_service.py index 4b1313be7c6..a62faf893e5 100644 --- a/api/tests/unit_tests/services/test_agent_config_service.py +++ b/api/tests/unit_tests/services/test_agent_config_service.py @@ -4,11 +4,22 @@ from __future__ import annotations import io import zipfile +from datetime import datetime from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest +from sqlalchemy.orm import Session, sessionmaker +from extensions.storage.storage_type import StorageType +from models.agent import ( + Agent, + AgentConfigDraft, + AgentConfigDraftType, + AgentConfigSnapshot, + AgentScope, + AgentSource, +) from models.agent_config_entities import ( AgentConfigFileRefConfig, AgentConfigSkillRefConfig, @@ -16,27 +27,35 @@ from models.agent_config_entities import ( AgentFileRefConfig, AgentSoulConfig, ) +from models.enums import CreatorUserRole +from models.model import UploadFile +from models.tools import ToolFile from services.agent.skill_package_service import SkillPackageError from services.agent_config_service import ( AgentConfigService, AgentConfigServiceError, AgentConfigTarget, AgentConfigVersionKind, + ConfigDownloadRequest, ConfigPushPayload, ConfigPushSkillItem, ) MODULE = "services.agent_config_service" -TENANT = "tenant-1" -AGENT = "agent-1" -USER = "user-1" +TENANT = "11111111-1111-1111-1111-111111111111" +OTHER_TENANT = "22222222-2222-2222-2222-222222222222" +AGENT = "33333333-3333-3333-3333-333333333333" +USER = "44444444-4444-4444-4444-444444444444" +END_USER = "55555555-5555-5555-5555-555555555555" +SNAPSHOT = "66666666-6666-6666-6666-666666666666" +DRAFT = "77777777-7777-7777-7777-777777777777" +BUILD_DRAFT = "88888888-8888-8888-8888-888888888888" +TOOL_FILE = "99999999-9999-9999-9999-999999999999" +SKILL_FILE = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" +UPLOAD_FILE = "bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb" +NORMALIZED_SKILL_FILE = "cccccccc-cccc-cccc-cccc-cccccccccccc" - -def _session_cm(session: MagicMock) -> MagicMock: - context_manager = MagicMock() - context_manager.__enter__.return_value = session - context_manager.__exit__.return_value = None - return context_manager +AGENT_CONFIG_TABLES = (Agent, AgentConfigDraft, AgentConfigSnapshot) def _soul(**updates) -> AgentSoulConfig: @@ -45,6 +64,62 @@ def _soul(**updates) -> AgentSoulConfig: return AgentSoulConfig.model_validate(payload) +def _agent(*, tenant_id: str = TENANT) -> Agent: + return Agent( + id=AGENT, + tenant_id=tenant_id, + name="Config Agent", + scope=AgentScope.ROSTER, + source=AgentSource.ROSTER, + ) + + +def _draft( + *, + version_id: str = DRAFT, + draft_type: AgentConfigDraftType = AgentConfigDraftType.DRAFT, + account_id: str | None = None, + soul: AgentSoulConfig | None = None, +) -> AgentConfigDraft: + return AgentConfigDraft( + id=version_id, + tenant_id=TENANT, + agent_id=AGENT, + draft_type=draft_type, + account_id=account_id, + draft_owner_key=account_id or "", + config_snapshot=soul or _soul(), + ) + + +def _snapshot(*, soul: AgentSoulConfig | None = None) -> AgentConfigSnapshot: + return AgentConfigSnapshot( + id=SNAPSHOT, + tenant_id=TENANT, + agent_id=AGENT, + version=1, + config_snapshot=soul or _soul(), + ) + + +def _service(sqlite_session: Session) -> AgentConfigService: + """Bind service-owned sessions to the current test's isolated SQLite engine.""" + + return AgentConfigService( + session_factory=sessionmaker(bind=sqlite_session.get_bind(), expire_on_commit=False), + ) + + +def _persist_target( + sqlite_session: Session, + version: AgentConfigDraft | AgentConfigSnapshot, + *, + tenant_id: str = TENANT, +) -> None: + sqlite_session.add_all([_agent(tenant_id=tenant_id), version]) + sqlite_session.commit() + + def _version(*, version_id: str = "version-1", snapshot: AgentSoulConfig | None = None) -> SimpleNamespace: agent_soul = snapshot or _soul() return SimpleNamespace( @@ -83,301 +158,475 @@ def _zip_bytes(members: dict[str, bytes]) -> bytes: @pytest.mark.parametrize( - ("kind", "user_id", "version_row", "expected_writable"), + ("kind", "user_id", "version_id", "expected_writable"), [ - (AgentConfigVersionKind.SNAPSHOT, None, _version(version_id="snapshot-1"), False), - (AgentConfigVersionKind.DRAFT, USER, _version(version_id="draft-1"), False), - (AgentConfigVersionKind.BUILD_DRAFT, USER, _version(version_id="build-draft-1"), True), + (AgentConfigVersionKind.SNAPSHOT, None, SNAPSHOT, False), + (AgentConfigVersionKind.DRAFT, USER, DRAFT, False), + (AgentConfigVersionKind.BUILD_DRAFT, USER, BUILD_DRAFT, True), ], ) +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) def test_resolve_target_supports_snapshot_draft_and_build_draft( kind: AgentConfigVersionKind, user_id: str | None, - version_row: SimpleNamespace, + version_id: str, expected_writable: bool, + sqlite_session: Session, ) -> None: - session = MagicMock() - session.scalar.side_effect = [AGENT, version_row] - service = AgentConfigService() - - with patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)): - target = service.resolve_target( - tenant_id=TENANT, - agent_id=AGENT, - config_version_id=version_row.id, - config_version_kind=kind, - user_id=user_id, + if kind == AgentConfigVersionKind.SNAPSHOT: + version = _snapshot() + else: + version = _draft( + version_id=version_id, + draft_type=( + AgentConfigDraftType.DEBUG_BUILD + if kind == AgentConfigVersionKind.BUILD_DRAFT + else AgentConfigDraftType.DRAFT + ), + account_id=USER if kind == AgentConfigVersionKind.BUILD_DRAFT else None, ) + _persist_target(sqlite_session, version) + + target = _service(sqlite_session).resolve_target( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=version_id, + config_version_kind=kind, + user_id=user_id, + ) assert target.agent_id == AGENT - assert target.version_id == version_row.id + assert target.version_id == version_id assert target.kind == kind assert target.writable is expected_writable + assert target.agent_soul == _soul() -def test_resolve_target_requires_user_for_build_draft() -> None: - session = MagicMock() - session.scalar.side_effect = [AGENT] - service = AgentConfigService() +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) +def test_resolve_target_requires_user_for_build_draft(sqlite_session: Session) -> None: + _persist_target( + sqlite_session, + _draft( + version_id=BUILD_DRAFT, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id=USER, + ), + ) - with patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)): - with pytest.raises(AgentConfigServiceError, match="user_id is required") as exc_info: - service.resolve_target( - tenant_id=TENANT, - agent_id=AGENT, - config_version_id="build-draft-1", - config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, - ) + with pytest.raises(AgentConfigServiceError, match="user_id is required") as exc_info: + _service(sqlite_session).resolve_target( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=BUILD_DRAFT, + config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + ) assert exc_info.value.code == "missing_user_id" +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) +def test_resolve_target_hides_build_draft_from_another_user(sqlite_session: Session) -> None: + _persist_target( + sqlite_session, + _draft( + version_id=BUILD_DRAFT, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id=USER, + ), + ) + + with pytest.raises(AgentConfigServiceError) as exc_info: + _service(sqlite_session).resolve_target( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=BUILD_DRAFT, + config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + user_id=END_USER, + ) + + assert (exc_info.value.code, exc_info.value.status_code) == ("config_version_not_found", 404) + + @pytest.mark.parametrize( - ("first_scalar", "expected_code"), + ("agent_tenant_id", "expected_code"), [ - (None, "agent_not_found"), - (AGENT, "config_version_not_found"), + (OTHER_TENANT, "agent_not_found"), + (TENANT, "config_version_not_found"), ], ) -def test_resolve_target_maps_missing_agent_and_version(first_scalar: str | None, expected_code: str) -> None: - session = MagicMock() - if first_scalar is None: - session.scalar.return_value = None - else: - session.scalar.side_effect = [first_scalar, None] - service = AgentConfigService() +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) +def test_resolve_target_maps_missing_agent_and_version( + agent_tenant_id: str, + expected_code: str, + sqlite_session: Session, +) -> None: + sqlite_session.add(_agent(tenant_id=agent_tenant_id)) + sqlite_session.commit() - with patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)): - with pytest.raises(AgentConfigServiceError) as exc_info: - service.resolve_target( - tenant_id=TENANT, - agent_id=AGENT, - config_version_id="missing", - config_version_kind=AgentConfigVersionKind.SNAPSHOT, - user_id=USER, - ) + with pytest.raises(AgentConfigServiceError) as exc_info: + _service(sqlite_session).resolve_target( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=SNAPSHOT, + config_version_kind=AgentConfigVersionKind.SNAPSHOT, + user_id=USER, + ) assert exc_info.value.code == expected_code -def test_push_rejects_non_build_draft_writes() -> None: - session = MagicMock() - service = AgentConfigService() +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) +def test_push_rejects_non_build_draft_writes(sqlite_session: Session) -> None: + _persist_target(sqlite_session, _draft(soul=_soul(config_note="before"))) - with ( - patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)), - patch.object( - service, - "_resolve_target_in_session", - return_value=_target(kind=AgentConfigVersionKind.DRAFT, writable=False), - ), - ): - with pytest.raises(AgentConfigServiceError, match="build drafts") as exc_info: - service.push( - tenant_id=TENANT, - agent_id=AGENT, - user_id=USER, - config_version_id="draft-1", - config_version_kind=AgentConfigVersionKind.DRAFT, - payload=ConfigPushPayload(note="ignored"), - ) - - assert exc_info.value.code == "config_not_writable" - session.commit.assert_not_called() - - -def test_push_for_console_allows_shared_draft_mutations() -> None: - session = MagicMock() - service = AgentConfigService() - target = _target(kind=AgentConfigVersionKind.DRAFT, writable=False, soul=_soul(config_note="before")) - - with ( - patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)), - patch.object(service, "_resolve_target_in_session", return_value=target), - ): - manifest = service.push_for_console( + with pytest.raises(AgentConfigServiceError, match="build drafts") as exc_info: + _service(sqlite_session).push( tenant_id=TENANT, agent_id=AGENT, user_id=USER, - config_version_id="draft-1", + config_version_id=DRAFT, config_version_kind=AgentConfigVersionKind.DRAFT, - payload=ConfigPushPayload(note="after"), + payload=ConfigPushPayload(note="ignored"), ) + assert exc_info.value.code == "config_not_writable" + sqlite_session.expire_all() + persisted = sqlite_session.get(AgentConfigDraft, DRAFT) + assert persisted is not None + assert persisted.config_snapshot.config_note == "before" + + +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) +def test_push_for_console_allows_shared_draft_mutations(sqlite_session: Session) -> None: + _persist_target(sqlite_session, _draft(soul=_soul(config_note="before"))) + + manifest = _service(sqlite_session).push_for_console( + tenant_id=TENANT, + agent_id=AGENT, + user_id=USER, + config_version_id=DRAFT, + config_version_kind=AgentConfigVersionKind.DRAFT, + payload=ConfigPushPayload(note="after"), + ) + assert manifest["note"] == "after" - assert target.version.config_snapshot.config_note == "after" - session.commit.assert_called_once() + sqlite_session.expire_all() + persisted = sqlite_session.get(AgentConfigDraft, DRAFT) + assert persisted is not None + assert persisted.config_snapshot.config_note == "after" -def test_push_accepts_tenant_scoped_tool_file_sources_from_different_upload_owner() -> None: - session = MagicMock() - service = AgentConfigService() - target = _target( - kind=AgentConfigVersionKind.BUILD_DRAFT, - writable=True, - soul=_soul( - config_skills=[{"name": "alpha", "file_id": "", "is_missing": True}], - config_files=[{"name": "guide.txt", "file_kind": "tool_file", "file_id": "", "is_missing": True}], +@pytest.mark.parametrize("sqlite_session", [(*AGENT_CONFIG_TABLES, ToolFile)], indirect=True) +def test_push_accepts_tenant_scoped_tool_file_sources_from_different_upload_owner( + sqlite_session: Session, +) -> None: + _persist_target( + sqlite_session, + _draft( + version_id=BUILD_DRAFT, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id=USER, ), ) - file_source = SimpleNamespace( - id="tool-file-file", + file_source = ToolFile( tenant_id=TENANT, - user_id="end-user-1", + user_id=END_USER, + conversation_id=None, size=7, mimetype="text/plain", file_key="file-key", name="guide.txt", ) - skill_source = SimpleNamespace( - id="tool-file-skill", + file_source.id = TOOL_FILE + skill_source = ToolFile( tenant_id=TENANT, - user_id="end-user-1", + user_id=END_USER, + conversation_id=None, size=123, mimetype="application/zip", file_key="skill-key", name="alpha.zip", ) + skill_source.id = SKILL_FILE + sqlite_session.add_all([file_source, skill_source]) + sqlite_session.commit() skill_ref = AgentConfigSkillRefConfig( name="alpha", description="Alpha skill", - file_id="normalized-skill-file", + file_id=NORMALIZED_SKILL_FILE, size=321, mime_type="application/zip", ) + service = _service(sqlite_session) with ( - patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)), - patch.object(service, "_resolve_target_in_session", return_value=target), - patch.object(service, "_require_tool_file_source", side_effect=[file_source, skill_source]) as require_source, patch(f"{MODULE}.storage.load_once", return_value=b"skill-archive"), - patch.object(service._skill_normalizer, "normalize", return_value=(skill_ref, object())), + patch.object( + service._skill_normalizer, + "normalize", + return_value=(skill_ref, object()), + ), ): manifest = service.push( tenant_id=TENANT, agent_id=AGENT, user_id=USER, - config_version_id="build-draft-1", + config_version_id=BUILD_DRAFT, config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, payload=ConfigPushPayload.model_validate( { - "files": [{"name": "guide.txt", "file_ref": {"kind": "tool_file", "id": "tool-file-file"}}], - "skills": [{"name": "alpha", "file_ref": {"kind": "tool_file", "id": "tool-file-skill"}}], + "files": [{"name": "guide.txt", "file_ref": {"kind": "tool_file", "id": TOOL_FILE}}], + "skills": [{"name": "alpha", "file_ref": {"kind": "tool_file", "id": SKILL_FILE}}], } ), ) - assert [call.args for call in require_source.call_args_list] == [(session,), (session,)] - assert [call.kwargs for call in require_source.call_args_list] == [ - {"tenant_id": TENANT, "file_id": "tool-file-file"}, - {"tenant_id": TENANT, "file_id": "tool-file-skill"}, - ] files = manifest["files"] skills = manifest["skills"] assert isinstance(files, dict) assert isinstance(skills, dict) - assert files["items"][0]["file_id"] == "tool-file-file" - assert files["items"][0]["is_missing"] is False - assert skills["items"][0]["file_id"] == "normalized-skill-file" - assert skills["items"][0]["is_missing"] is False - session.commit.assert_called_once() + assert files["items"][0]["file_id"] == TOOL_FILE + assert skills["items"][0]["file_id"] == NORMALIZED_SKILL_FILE + sqlite_session.expire_all() + persisted = sqlite_session.get(AgentConfigDraft, BUILD_DRAFT) + assert persisted is not None + assert persisted.config_snapshot.config_files[0].file_id == TOOL_FILE + assert persisted.config_snapshot.config_skills[0].file_id == NORMALIZED_SKILL_FILE + persisted_source = sqlite_session.get(ToolFile, TOOL_FILE) + assert persisted_source is not None + assert persisted_source.user_id == END_USER -def test_push_file_for_console_rejects_snapshot_writes() -> None: - session = MagicMock() - service = AgentConfigService() - - with ( - patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)), - patch.object( - service, - "_resolve_target_in_session", - return_value=_target(kind=AgentConfigVersionKind.SNAPSHOT, writable=False), +@pytest.mark.parametrize("sqlite_session", [(*AGENT_CONFIG_TABLES, ToolFile)], indirect=True) +def test_request_download_signs_config_tool_files_without_rechecking_end_user_owner( + sqlite_session: Session, +) -> None: + soul = _soul( + config_files=[AgentConfigFileRefConfig(name="guide.txt", file_kind="tool_file", file_id=TOOL_FILE)], + config_skills=[AgentConfigSkillRefConfig(name="alpha", file_id=SKILL_FILE)], + ) + _persist_target( + sqlite_session, + _draft( + version_id=BUILD_DRAFT, + draft_type=AgentConfigDraftType.DEBUG_BUILD, + account_id=USER, + soul=soul, ), - ): - with pytest.raises(AgentConfigServiceError, match="editable drafts") as exc_info: - service.push_file_for_console( - tenant_id=TENANT, - agent_id=AGENT, - user_id=USER, - config_version_id="snapshot-1", - config_version_kind=AgentConfigVersionKind.SNAPSHOT, - upload_file_id="upload-1", - ) - - assert exc_info.value.code == "config_not_writable" - - -def test_push_file_for_console_uses_service_owned_upload_lookup_and_naming() -> None: - session = MagicMock() - service = AgentConfigService() - target = _target( - kind=AgentConfigVersionKind.DRAFT, - writable=False, - soul=_soul(config_files=[{"name": "guide.txt", "file_kind": "upload_file", "file_id": "", "is_missing": True}]), ) - upload_file = SimpleNamespace( - id="upload-1", - name="guide.txt", + guide = ToolFile( + tenant_id=TENANT, + user_id=END_USER, + conversation_id=None, size=7, - hash="sha256:abc", - mime_type="text/plain", + mimetype="text/plain", + file_key="tool-files/guide.txt", + name="guide.txt", + ) + guide.id = TOOL_FILE + skill = ToolFile( + tenant_id=TENANT, + user_id=END_USER, + conversation_id=None, + size=123, + mimetype="application/zip", + file_key="tool-files/alpha.zip", + name="alpha.zip", + ) + skill.id = SKILL_FILE + sqlite_session.add_all([guide, skill]) + sqlite_session.commit() + service = _service(sqlite_session) + + file_download = service.request_download( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=BUILD_DRAFT, + config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + kind="file", + name="guide.txt", + user_id=USER, + ) + skill_download = service.request_download( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=BUILD_DRAFT, + config_version_kind=AgentConfigVersionKind.BUILD_DRAFT, + kind="skill", + name="alpha", + user_id=USER, ) - with ( - patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)), - patch.object(service, "_resolve_target_in_session", return_value=target), - patch.object(service, "_require_console_upload_file_source", return_value=upload_file), - ): - response = service.push_file_for_console( + assert (file_download.filename, file_download.mime_type, file_download.size) == ("guide.txt", "text/plain", 7) + assert file_download.download_uri.startswith(f"/files/tools/{TOOL_FILE}.txt?") + assert "as_attachment=true" in file_download.download_uri + assert (skill_download.filename, skill_download.mime_type, skill_download.size) == ( + "alpha.zip", + "application/zip", + 123, + ) + assert skill_download.download_uri.startswith(f"/files/tools/{SKILL_FILE}.zip?") + + +@pytest.mark.parametrize("sqlite_session", [(*AGENT_CONFIG_TABLES, UploadFile)], indirect=True) +def test_request_download_signs_tenant_scoped_config_upload_file(sqlite_session: Session) -> None: + soul = _soul( + config_files=[AgentConfigFileRefConfig(name="guide.txt", file_kind="upload_file", file_id=UPLOAD_FILE)] + ) + _persist_target(sqlite_session, _draft(soul=soul)) + upload_file = UploadFile( + tenant_id=TENANT, + storage_type=StorageType.LOCAL, + key="uploads/guide.txt", + name="source-name.txt", + size=7, + extension="txt", + mime_type="text/plain", + created_by_role=CreatorUserRole.END_USER, + created_by=END_USER, + created_at=datetime(2025, 1, 1), + used=False, + ) + upload_file.id = UPLOAD_FILE + sqlite_session.add(upload_file) + sqlite_session.commit() + + result = _service(sqlite_session).request_download( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=DRAFT, + config_version_kind=AgentConfigVersionKind.DRAFT, + kind="file", + name="guide.txt", + user_id=USER, + ) + + assert (result.filename, result.mime_type, result.size) == ("guide.txt", "text/plain", 7) + assert result.download_uri.startswith(f"/files/{UPLOAD_FILE}/file-preview?") + assert "as_attachment=true" in result.download_uri + + +@pytest.mark.parametrize("sqlite_session", [(*AGENT_CONFIG_TABLES, ToolFile)], indirect=True) +def test_request_download_rejects_config_source_from_another_tenant(sqlite_session: Session) -> None: + soul = _soul(config_files=[AgentConfigFileRefConfig(name="guide.txt", file_kind="tool_file", file_id=TOOL_FILE)]) + _persist_target(sqlite_session, _draft(soul=soul)) + source = ToolFile( + tenant_id=OTHER_TENANT, + user_id=END_USER, + conversation_id=None, + size=7, + mimetype="text/plain", + file_key="tool-files/guide.txt", + name="guide.txt", + ) + source.id = TOOL_FILE + sqlite_session.add(source) + sqlite_session.commit() + + with pytest.raises(AgentConfigServiceError) as exc_info: + _service(sqlite_session).request_download( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=DRAFT, + config_version_kind=AgentConfigVersionKind.DRAFT, + kind="file", + name="guide.txt", + user_id=USER, + ) + + assert (exc_info.value.code, exc_info.value.status_code) == ("config_file_not_found", 404) + + +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) +def test_push_file_for_console_rejects_snapshot_writes(sqlite_session: Session) -> None: + _persist_target(sqlite_session, _snapshot()) + + with pytest.raises(AgentConfigServiceError, match="editable drafts") as exc_info: + _service(sqlite_session).push_file_for_console( tenant_id=TENANT, agent_id=AGENT, user_id=USER, - config_version_id="draft-1", - config_version_kind=AgentConfigVersionKind.DRAFT, - upload_file_id="upload-1", + config_version_id=SNAPSHOT, + config_version_kind=AgentConfigVersionKind.SNAPSHOT, + upload_file_id=UPLOAD_FILE, ) + assert exc_info.value.code == "config_not_writable" + sqlite_session.expire_all() + persisted = sqlite_session.get(AgentConfigSnapshot, SNAPSHOT) + assert persisted is not None + assert persisted.config_snapshot.config_files == [] + + +@pytest.mark.parametrize("sqlite_session", [(*AGENT_CONFIG_TABLES, UploadFile)], indirect=True) +def test_push_file_for_console_uses_service_owned_upload_lookup_and_naming(sqlite_session: Session) -> None: + _persist_target(sqlite_session, _draft()) + upload_file = UploadFile( + tenant_id=TENANT, + storage_type=StorageType.LOCAL, + key="uploads/guide.txt", + name="guide.txt", + size=7, + extension="txt", + mime_type="text/plain", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=USER, + created_at=datetime(2025, 1, 1), + used=False, + hash="sha256:abc", + ) + upload_file.id = UPLOAD_FILE + sqlite_session.add(upload_file) + sqlite_session.commit() + + response = _service(sqlite_session).push_file_for_console( + tenant_id=TENANT, + agent_id=AGENT, + user_id=USER, + config_version_id=DRAFT, + config_version_kind=AgentConfigVersionKind.DRAFT, + upload_file_id=UPLOAD_FILE, + ) + assert response == { "file": { "id": "guide.txt", "name": "guide.txt", - "file_id": "upload-1", + "file_id": UPLOAD_FILE, "is_missing": False, "size": 7, "hash": "sha256:abc", "mime_type": "text/plain", }, "config_version": { - "id": "version-1", + "id": DRAFT, "kind": "draft", "writable": True, }, } - session.commit.assert_called_once() + sqlite_session.expire_all() + persisted = sqlite_session.get(AgentConfigDraft, DRAFT) + assert persisted is not None + assert persisted.config_snapshot.config_files[0].file_id == UPLOAD_FILE -def test_upload_skill_for_console_maps_package_validation_failures() -> None: - session = MagicMock() - service = AgentConfigService() - target = _target(kind=AgentConfigVersionKind.DRAFT, writable=False) +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) +def test_upload_skill_for_console_maps_package_validation_failures(sqlite_session: Session) -> None: + _persist_target(sqlite_session, _draft()) + service = _service(sqlite_session) message = "skill package must contain exactly one skill; multiple skill folders in one archive are not supported" - with ( - patch(f"{MODULE}.session_factory.create_session", return_value=_session_cm(session)), - patch.object(service, "_resolve_target_in_session", return_value=target), - patch.object( - service._skill_normalizer, - "normalize", - side_effect=SkillPackageError("files_outside_skill_root", message, status_code=400), - ), + with patch.object( + service._skill_normalizer, + "normalize", + side_effect=SkillPackageError("files_outside_skill_root", message, status_code=400), ): with pytest.raises(AgentConfigServiceError, match="exactly one skill") as exc_info: service.upload_skill_for_console( tenant_id=TENANT, agent_id=AGENT, user_id=USER, - config_version_id="draft-1", + config_version_id=DRAFT, config_version_kind=AgentConfigVersionKind.DRAFT, content=b"bad-archive", filename="skills.zip", @@ -386,15 +635,19 @@ def test_upload_skill_for_console_maps_package_validation_failures() -> None: assert exc_info.value.code == "files_outside_skill_root" assert exc_info.value.message == message assert exc_info.value.status_code == 400 - session.commit.assert_not_called() + sqlite_session.expire_all() + persisted = sqlite_session.get(AgentConfigDraft, DRAFT) + assert persisted is not None + assert persisted.config_snapshot.config_skills == [] -def test_apply_skill_updates_rejects_non_tool_file_refs() -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_apply_skill_updates_rejects_non_tool_file_refs(sqlite_session: Session) -> None: service = AgentConfigService() with pytest.raises(AgentConfigServiceError, match="tool files") as exc_info: service._apply_skill_updates( - MagicMock(), + sqlite_session, tenant_id=TENANT, user_id=USER, current=[], @@ -415,7 +668,12 @@ def test_apply_skill_updates_rejects_non_tool_file_refs() -> None: ("invalid_archive", "stored tool file is not a valid skill archive"), ], ) -def test_apply_skill_updates_maps_normalizer_failures(error_code: str, message: str) -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_apply_skill_updates_maps_normalizer_failures( + error_code: str, + message: str, + sqlite_session: Session, +) -> None: service = AgentConfigService() tool_file = SimpleNamespace(name="alpha.zip", file_key="tool-files/alpha.zip") @@ -430,7 +688,7 @@ def test_apply_skill_updates_maps_normalizer_failures(error_code: str, message: ): with pytest.raises(AgentConfigServiceError, match=message) as exc_info: service._apply_skill_updates( - MagicMock(), + sqlite_session, tenant_id=TENANT, user_id=USER, current=[], @@ -549,37 +807,46 @@ def test_manifest_uses_items_shape_without_download_urls() -> None: } -def test_manifest_preserves_missing_config_assets_and_pull_rejects_them() -> None: +@pytest.mark.parametrize("sqlite_session", [AGENT_CONFIG_TABLES], indirect=True) +def test_manifest_preserves_missing_config_assets_and_download_rejects_them(sqlite_session: Session) -> None: soul = _soul( config_skills=[{"name": "alpha", "file_id": "", "is_missing": True}], config_files=[{"name": "guide.txt", "file_kind": "upload_file", "file_id": "", "is_missing": True}], ) - target = _target(kind=AgentConfigVersionKind.DRAFT, writable=False, soul=soul) - service = AgentConfigService() + _persist_target(sqlite_session, _draft(soul=soul)) + service = _service(sqlite_session) + target = service.resolve_target( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=DRAFT, + config_version_kind=AgentConfigVersionKind.DRAFT, + user_id=USER, + ) manifest = service._manifest_for_target(target) assert manifest["skills"]["items"][0]["is_missing"] is True # type: ignore[index] assert manifest["files"]["items"][0]["is_missing"] is True # type: ignore[index] - with patch.object(service, "resolve_target", return_value=target): - with pytest.raises(AgentConfigServiceError) as skill_error: - service.pull_skill( - tenant_id=TENANT, - agent_id=AGENT, - config_version_id="draft-1", - config_version_kind=AgentConfigVersionKind.DRAFT, - name="alpha", - user_id=USER, - ) - with pytest.raises(AgentConfigServiceError) as file_error: - service.pull_file( - tenant_id=TENANT, - agent_id=AGENT, - config_version_id="draft-1", - config_version_kind=AgentConfigVersionKind.DRAFT, - name="guide.txt", - user_id=USER, - ) + with pytest.raises(AgentConfigServiceError) as skill_error: + service.request_download( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=DRAFT, + config_version_kind=AgentConfigVersionKind.DRAFT, + kind="skill", + name="alpha", + user_id=USER, + ) + with pytest.raises(AgentConfigServiceError) as file_error: + service.request_download( + tenant_id=TENANT, + agent_id=AGENT, + config_version_id=DRAFT, + config_version_kind=AgentConfigVersionKind.DRAFT, + kind="file", + name="guide.txt", + user_id=USER, + ) assert (skill_error.value.code, skill_error.value.status_code) == ("config_skill_missing", 409) assert (file_error.value.code, file_error.value.status_code) == ("config_file_missing", 409) @@ -726,24 +993,19 @@ def test_resolve_skill_file_member_path_requires_existing_member() -> None: assert exc_info.value.status_code == 404 -def test_download_url_helpers_use_shared_url_resolution() -> None: +def test_download_url_helpers_bind_shared_download_request_to_console_origin() -> None: service = AgentConfigService() - target = _target( - kind=AgentConfigVersionKind.BUILD_DRAFT, - writable=True, - soul=_soul( - config_skills=[AgentConfigSkillRefConfig(name="alpha", file_id="tool-file-1")], - config_files=[AgentConfigFileRefConfig(name="guide.txt", file_kind="upload_file", file_id="upload-file-1")], - ), - ) with ( - patch.object(service, "resolve_target", return_value=target), patch.object( service, - "_resolve_download_url", - side_effect=["https://example.com/alpha.zip", "https://example.com/guide.txt"], + "request_download", + side_effect=[ + ConfigDownloadRequest("alpha.zip", "application/zip", 10, "/files/alpha.zip?sign=1"), + ConfigDownloadRequest("guide.txt", "text/plain", 20, "/files/guide.txt?sign=2"), + ], ), + patch(f"{MODULE}.dify_config.FILES_URL", "https://example.com"), ): assert ( service.download_skill_url( @@ -754,7 +1016,7 @@ def test_download_url_helpers_use_shared_url_resolution() -> None: name="alpha", user_id=USER, ) - == "https://example.com/alpha.zip" + == "https://example.com/files/alpha.zip?sign=1" ) assert ( service.download_file_url( @@ -765,5 +1027,5 @@ def test_download_url_helpers_use_shared_url_resolution() -> None: name="guide.txt", user_id=USER, ) - == "https://example.com/guide.txt" + == "https://example.com/files/guide.txt?sign=2" ) diff --git a/api/tests/unit_tests/services/test_agent_file_request_service.py b/api/tests/unit_tests/services/test_agent_file_request_service.py deleted file mode 100644 index c39e3f82558..00000000000 --- a/api/tests/unit_tests/services/test_agent_file_request_service.py +++ /dev/null @@ -1,105 +0,0 @@ -"""Unit tests for the Agent Files download-request service (ENG-592).""" - -from __future__ import annotations - -from types import SimpleNamespace -from unittest.mock import patch - -import pytest - -from services.agent_file_request_service import AgentFileDownloadRequestService, FileDownloadRequestError - -_MOD = "services.agent_file_request_service" - - -def _fake_file() -> SimpleNamespace: - return SimpleNamespace(filename="report.pdf", mime_type="application/pdf", size=12) - - -def test_resolve_returns_metadata_and_internal_url(): - with ( - patch(f"{_MOD}.file_factory.build_from_mapping", return_value=_fake_file()) as build, - patch(f"{_MOD}.DifyWorkflowFileRuntime") as runtime_cls, - ): - runtime_cls.return_value.resolve_file_url.return_value = "http://internal/files/x?sign=1" - data = AgentFileDownloadRequestService.resolve( - tenant_id="tenant-1", - user_id="user-1", - user_from="account", - invoke_from="service-api", - file_mapping={"transfer_method": "tool_file", "reference": "tool-file-1"}, - ) - - assert data == { - "filename": "report.pdf", - "mime_type": "application/pdf", - "size": 12, - "download_url": "http://internal/files/x?sign=1", - } - assert build.call_args.kwargs["tenant_id"] == "tenant-1" - # Sandbox/agent backend consumes the URL -> must be internal, not external. - assert runtime_cls.return_value.resolve_file_url.call_args.kwargs["for_external"] is False - - -@pytest.mark.parametrize( - ("user_from", "invoke_from", "code"), - [ - ("bogus", "service-api", "invalid_access_context"), - ("account", "not-a-source", "invalid_access_context"), - ], -) -def test_invalid_access_context_rejected(user_from: str, invoke_from: str, code: str): - with pytest.raises(FileDownloadRequestError) as exc_info: - AgentFileDownloadRequestService.resolve( - tenant_id="t", - user_id="u", - user_from=user_from, - invoke_from=invoke_from, - file_mapping={"transfer_method": "tool_file", "reference": "x"}, - ) - assert exc_info.value.status_code == 400 - assert exc_info.value.code == code - - -def test_missing_transfer_method_rejected(): - with pytest.raises(FileDownloadRequestError) as exc_info: - AgentFileDownloadRequestService.resolve( - tenant_id="t", - user_id="u", - user_from="account", - invoke_from="service-api", - file_mapping={}, - ) - assert exc_info.value.status_code == 400 - assert exc_info.value.code == "invalid_file_mapping" - - -def test_inaccessible_file_maps_to_404(): - with patch(f"{_MOD}.file_factory.build_from_mapping", side_effect=ValueError("ToolFile x not found")): - with pytest.raises(FileDownloadRequestError) as exc_info: - AgentFileDownloadRequestService.resolve( - tenant_id="t", - user_id="u", - user_from="end-user", - invoke_from="web-app", - file_mapping={"transfer_method": "tool_file", "reference": "x"}, - ) - assert exc_info.value.status_code == 404 - assert exc_info.value.code == "file_not_accessible" - - -def test_unresolved_url_maps_to_502(): - with ( - patch(f"{_MOD}.file_factory.build_from_mapping", return_value=_fake_file()), - patch(f"{_MOD}.DifyWorkflowFileRuntime") as runtime_cls, - ): - runtime_cls.return_value.resolve_file_url.return_value = None - with pytest.raises(FileDownloadRequestError) as exc_info: - AgentFileDownloadRequestService.resolve( - tenant_id="t", - user_id="u", - user_from="account", - invoke_from="service-api", - file_mapping={"transfer_method": "tool_file", "reference": "x"}, - ) - assert exc_info.value.status_code == 502 diff --git a/api/tests/unit_tests/services/test_annotation_service.py b/api/tests/unit_tests/services/test_annotation_service.py index d729b1c5a80..4814490c9d5 100644 --- a/api/tests/unit_tests/services/test_annotation_service.py +++ b/api/tests/unit_tests/services/test_annotation_service.py @@ -19,19 +19,20 @@ from services.annotation_service import AppAnnotationService from services.app_ref_service import AnnotationRef, AppRef -def _make_app(app_id: str = "app-1", tenant_id: str = "tenant-1") -> MagicMock: - app = MagicMock(spec=App) - app.id = app_id - app.tenant_id = tenant_id - app.status = "normal" +def _make_app(app_id: str = "app-1", tenant_id: str = "tenant-1") -> App: + app = App( + id=app_id, + tenant_id=tenant_id, + status="normal", + ) return app -def _make_app_ref(app: MagicMock) -> AppRef: +def _make_app_ref(app: App) -> AppRef: return AppRef(tenant_id=app.tenant_id, app_id=app.id) -def _make_annotation_ref(app: MagicMock, annotation_id: str = "ann-1") -> AnnotationRef: +def _make_annotation_ref(app: App, annotation_id: str = "ann-1") -> AnnotationRef: return AnnotationRef(app=AppRef(tenant_id=app.tenant_id, app_id=app.id), annotation_id=annotation_id) @@ -41,35 +42,36 @@ def _make_user(user_id: str = "user-1") -> MagicMock: return user -def _make_message(message_id: str = "msg-1", app_id: str = "app-1") -> MagicMock: - message = MagicMock(spec=Message) - message.id = message_id - message.app_id = app_id - message.conversation_id = "conv-1" - message.query = "default-question" - message.annotation = None +def _make_message(message_id: str = "msg-1", app_id: str = "app-1") -> Message: + message = Message( + id=message_id, + app_id=app_id, + conversation_id="conv-1", + query="default-question", + ) return message -def _make_annotation(annotation_id: str = "ann-1", app_id: str = "app-1") -> MagicMock: - annotation = MagicMock(spec=MessageAnnotation) +def _make_annotation(annotation_id: str = "ann-1", app_id: str = "app-1") -> MessageAnnotation: + annotation = MessageAnnotation( + app_id=app_id, + question="", + content="", + account_id="account-id", + ) annotation.id = annotation_id - annotation.app_id = app_id - annotation.content = "" - annotation.question = "" - annotation.question_text = "" return annotation -def _make_setting(setting_id: str = "setting-1", with_detail: bool = True) -> MagicMock: - setting = MagicMock(spec=AppAnnotationSetting) +def _make_setting(setting_id: str = "setting-1") -> AppAnnotationSetting: + setting = AppAnnotationSetting( + app_id="app-id", + score_threshold=0.5, + collection_binding_id="collection-1", + created_user_id="account-id", + updated_user_id="account-id", + ) setting.id = setting_id - setting.score_threshold = 0.5 - setting.collection_binding_id = "collection-1" - if with_detail: - setting.collection_binding_detail = SimpleNamespace(provider_name="provider-a", model_name="model-a") - else: - setting.collection_binding_detail = None return setting @@ -188,7 +190,6 @@ class TestAppAnnotationServiceUpInsert: tenant_id = "tenant-1" app = _make_app() message = _make_message(message_id="msg-1", app_id=app.id) - message.annotation = None with ( patch("services.annotation_service.current_account_with_tenant", return_value=(current_user, tenant_id)), @@ -595,7 +596,7 @@ class TestAppAnnotationServiceDirectManipulation: tenant_id = "tenant-1" app = _make_app() annotation = _make_annotation("ann-1") - annotation.question_text = "q1" + annotation.question = "q1" setting = _make_setting() with ( @@ -631,8 +632,28 @@ class TestAppAnnotationServiceDirectManipulation: tenant_id = "tenant-1" app = _make_app() annotation = _make_annotation("ann-1") - history1 = MagicMock(spec=AppAnnotationHitHistory) - history2 = MagicMock(spec=AppAnnotationHitHistory) + history1 = AppAnnotationHitHistory( + app_id="app-id", + annotation_id="annotation-id", + source="hit-testing", + question="question", + account_id="account-id", + score=0.0, + message_id="message-id", + annotation_question="question", + annotation_content="content", + ) + history2 = AppAnnotationHitHistory( + app_id="app-id", + annotation_id="annotation-id", + source="hit-testing", + question="question", + account_id="account-id", + score=0.0, + message_id="message-id", + annotation_question="question", + annotation_content="content", + ) setting = _make_setting() with ( @@ -1183,9 +1204,8 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Arrange tenant_id = "tenant-1" app = _make_app() - setting = _make_setting(with_detail=True) - detail = setting.collection_binding_detail - setting.collection_binding_detail = None + setting = _make_setting() + detail = SimpleNamespace(provider_name="provider-a", model_name="model-a") with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), @@ -1224,8 +1244,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: # Arrange tenant_id = "tenant-1" app = _make_app() - setting = _make_setting(with_detail=False) - setting.collection_binding_detail = SimpleNamespace(provider_name="wrong", model_name="wrong") + setting = _make_setting() with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), @@ -1266,9 +1285,8 @@ class TestAppAnnotationServiceHitHistoryAndSettings: tenant_id = "tenant-1" current_user = _make_user("user-1") app = _make_app() - setting = _make_setting(with_detail=True) - detail = setting.collection_binding_detail - setting.collection_binding_detail = None + setting = _make_setting() + detail = SimpleNamespace(provider_name="provider-a", model_name="model-a") args = {"score_threshold": 0.8} with ( @@ -1297,8 +1315,7 @@ class TestAppAnnotationServiceHitHistoryAndSettings: tenant_id = "tenant-1" current_user = _make_user("user-1") app = _make_app() - setting = _make_setting(with_detail=False) - setting.collection_binding_detail = SimpleNamespace(provider_name="wrong", model_name="wrong") + setting = _make_setting() args = {"score_threshold": 0.7} with ( @@ -1365,7 +1382,17 @@ class TestAppAnnotationServiceClearAll: setting = _make_setting() annotation1 = _make_annotation("ann-1") annotation2 = _make_annotation("ann-2") - history = MagicMock(spec=AppAnnotationHitHistory) + history = AppAnnotationHitHistory( + app_id="app-id", + annotation_id="annotation-id", + source="hit-testing", + question="question", + account_id="account-id", + score=0.0, + message_id="message-id", + annotation_question="question", + annotation_content="content", + ) with ( patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)), diff --git a/api/tests/unit_tests/services/test_app_dsl_service.py b/api/tests/unit_tests/services/test_app_dsl_service.py index 534fd41a33b..7afbb0fa042 100644 --- a/api/tests/unit_tests/services/test_app_dsl_service.py +++ b/api/tests/unit_tests/services/test_app_dsl_service.py @@ -3,41 +3,116 @@ from typing import cast from unittest.mock import Mock import pytest -from sqlalchemy.orm import Session +import yaml +from sqlalchemy import event, select +from sqlalchemy.orm import Session, sessionmaker from core.rbac import RBACPermission +from core.workflow.llm_environment_variable import LLMEnvironmentVariable from models import App, AppMode -from models.model import AppModelConfig, IconType -from services.app_dsl_service import AppDslService +from models.model import AppModelConfig, AppModelConfigDict, IconType +from models.workflow import Workflow +from services.app_dsl_service import AppDslService, PendingData from services.entities.dsl_entities import ImportStatus from services.errors.account import NoPermissionError +from services.errors.app import WorkflowNotFoundError + + +def test_extract_workflow_dependencies_uses_llm_environment_variable_provider(monkeypatch: pytest.MonkeyPatch) -> None: + workflow = SimpleNamespace( + graph_dict={ + "nodes": [ + { + "id": "llm-node", + "data": { + "type": "llm", + "title": "LLM", + "model": {"provider": "old-provider", "name": "old-model", "mode": "chat"}, + "model_selector": ["env", "shared_model"], + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + }, + } + ] + }, + environment_variables=[ + LLMEnvironmentVariable( + name="shared_model", + value={"provider": "new-provider", "name": "new-model", "mode": "chat"}, + ) + ], + ) + analyze_dependency = Mock(side_effect=lambda provider: provider) + monkeypatch.setattr( + "services.app_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", + analyze_dependency, + ) + + result = AppDslService._extract_dependencies_from_workflow(cast(Workflow, workflow)) + + assert result == ["new-provider"] + analyze_dependency.assert_called_once_with("new-provider") + + +@pytest.mark.parametrize("model_selector", [[], ["env", "missing_model"]]) +def test_extract_workflow_dependencies_tolerates_unresolved_llm_environment_reference( + monkeypatch: pytest.MonkeyPatch, model_selector: list[str] +) -> None: + workflow = SimpleNamespace( + graph_dict={ + "nodes": [ + { + "id": "llm-node", + "data": { + "type": "llm", + "title": "LLM", + "model": {"provider": "old-provider", "name": "old-model", "mode": "chat"}, + "model_selector": model_selector, + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + }, + } + ] + }, + environment_variables=[], + ) + analyze_dependency = Mock(side_effect=lambda provider: provider) + monkeypatch.setattr( + "services.app_dsl_service.DependenciesAnalysisService.analyze_model_provider_dependency", + analyze_dependency, + ) + + result = AppDslService._extract_dependencies_from_workflow(cast(Workflow, workflow)) + + assert result == ["old-provider"] + analyze_dependency.assert_called_once_with("old-provider") -@pytest.mark.parametrize("sqlite_session", [()], indirect=True) def test_import_app_rejects_oversized_yaml_content_before_parsing( - monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + monkeypatch: pytest.MonkeyPatch, unbound_session: Session ) -> None: monkeypatch.setattr("services.app_dsl_service.DSL_MAX_SIZE", 3) - service = AppDslService(session=sqlite_session) + service = AppDslService(session=unbound_session) account = Mock(current_tenant_id="tenant-1") result = service.import_app(account=account, import_mode="yaml-content", yaml_content="你你") assert result.status == ImportStatus.FAILED assert result.error == "File size exceeds the limit of 10MB" - assert not sqlite_session.in_transaction() + assert not unbound_session.in_transaction() -@pytest.mark.parametrize("sqlite_session", [()], indirect=True) def test_import_app_rejects_oversized_yaml_url_bytes_before_decode( - monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + monkeypatch: pytest.MonkeyPatch, unbound_session: Session ) -> None: monkeypatch.setattr("services.app_dsl_service.DSL_MAX_SIZE", 1) response = Mock() response.raise_for_status.return_value = None response.content = b"\xff\xff" monkeypatch.setattr("services.app_dsl_service.remote_fetcher.make_request", Mock(return_value=response)) - service = AppDslService(session=sqlite_session) + service = AppDslService(session=unbound_session) result = service.import_app( account=Mock(current_tenant_id="tenant-1"), @@ -47,18 +122,17 @@ def test_import_app_rejects_oversized_yaml_url_bytes_before_decode( assert result.status == ImportStatus.FAILED assert result.error == "File size exceeds the limit of 10MB" - assert not sqlite_session.in_transaction() + assert not unbound_session.in_transaction() -@pytest.mark.parametrize("sqlite_session", [()], indirect=True) def test_import_app_returns_decode_error_for_invalid_yaml_url_bytes( - monkeypatch: pytest.MonkeyPatch, sqlite_session: Session + monkeypatch: pytest.MonkeyPatch, unbound_session: Session ) -> None: response = Mock() response.raise_for_status.return_value = None response.content = b"\xff" monkeypatch.setattr("services.app_dsl_service.remote_fetcher.make_request", Mock(return_value=response)) - service = AppDslService(session=sqlite_session) + service = AppDslService(session=unbound_session) result = service.import_app( account=Mock(current_tenant_id="tenant-1"), @@ -68,19 +142,99 @@ def test_import_app_returns_decode_error_for_invalid_yaml_url_bytes( assert result.status == ImportStatus.FAILED assert "utf-8" in result.error - assert not sqlite_session.in_transaction() + assert not unbound_session.in_transaction() -def test_create_or_update_app_loads_existing_model_config_with_service_session() -> None: - session = Mock() - session.get.return_value = Mock() - service = AppDslService(session=session) +def test_pending_import_is_scoped_to_its_owner(monkeypatch: pytest.MonkeyPatch, unbound_session: Session) -> None: + pending_imports: dict[str, str] = {} + monkeypatch.setattr( + "services.app_dsl_service.redis_client.setex", + lambda key, _expiry, value: pending_imports.__setitem__(key, value), + ) + service = AppDslService(session=unbound_session) + creator = Mock(id="account-1", current_tenant_id="tenant-1") + + pending = service.import_app( + account=creator, + import_mode="yaml-content", + yaml_content="version: 99.0.0\nkind: app\napp: {name: Test, mode: workflow}\n", + ) + + redis_key = f"app_import_info:{pending.id}" + assert pending.status == ImportStatus.PENDING + assert redis_key in pending_imports + pending_data = PendingData.model_validate_json(pending_imports[redis_key]) + assert pending_data.tenant_id == "tenant-1" + assert pending_data.account_id == "account-1" + + monkeypatch.setattr("services.app_dsl_service.redis_client.get", pending_imports.get) + monkeypatch.setattr("services.app_dsl_service.redis_client.delete", pending_imports.pop) + monkeypatch.setattr( + service, + "_create_or_update_app", + Mock(return_value=Mock(id="app-1", mode=AppMode.WORKFLOW)), + ) + + for other_account in ( + Mock(id="account-1", current_tenant_id="tenant-2"), + Mock(id="account-2", current_tenant_id="tenant-1"), + ): + assert service.confirm_import(import_id=pending.id, account=other_account).status == ImportStatus.FAILED + + assert service.confirm_import(import_id=pending.id, account=creator).status == ImportStatus.COMPLETED + assert redis_key not in pending_imports + + +@pytest.mark.parametrize( + ("tenant_id", "account_id", "expected"), + [ + ("tenant-1", "account-1", True), + (None, "account-1", False), + ("tenant-1", None, False), + ("tenant-2", "account-1", False), + ("tenant-1", "account-2", False), + ], +) +def test_pending_import_owner_access( + tenant_id: str | None, + account_id: str | None, + expected: bool, +) -> None: + pending = PendingData( + tenant_id=tenant_id, + account_id=account_id, + import_mode="yaml-content", + yaml_content="", + ) + + assert pending.is_accessible_by(tenant_id="tenant-1", account_id="account-1") is expected + + +def test_pending_import_owner_access_accepts_legacy_json() -> None: + pending = PendingData.model_validate_json('{"import_mode":"yaml-content","yaml_content":""}') + + assert pending.is_accessible_by(tenant_id="tenant-1", account_id="account-1") + assert not pending.is_accessible_by(tenant_id=None, account_id="account-1") + + +def test_create_or_update_app_loads_existing_model_config_with_service_session( + sqlite_session_factory: sessionmaker[Session], +) -> None: + with sqlite_session_factory() as arrange_session: + app_model_config = AppModelConfig( + app_id="11111111-1111-1111-1111-111111111111", + created_by="22222222-2222-2222-2222-222222222222", + updated_by="22222222-2222-2222-2222-222222222222", + ) + arrange_session.add(app_model_config) + arrange_session.commit() + app_model_config_id = app_model_config.id app = cast( App, SimpleNamespace( - id="app-1", - tenant_id="tenant-1", - app_model_config_id="config-1", + id="11111111-1111-1111-1111-111111111111", + tenant_id="33333333-3333-3333-3333-333333333333", + app_model_config_id=app_model_config_id, name="Existing app", description="", icon_type=IconType.EMOJI, @@ -89,29 +243,39 @@ def test_create_or_update_app_loads_existing_model_config_with_service_session() ), ) - result = service._create_or_update_app( - app=app, - data={"app": {"mode": AppMode.CHAT}, "model_config": {"model": {}}}, - account=Mock(id="account-1"), - ) + with sqlite_session_factory() as service_session: + result = AppDslService(session=service_session)._create_or_update_app( + app=app, + data={"app": {"mode": AppMode.CHAT}, "model_config": {"model": {}}}, + account=Mock(id="account-1"), + ) - assert result is app - session.get.assert_called_once_with(AppModelConfig, "config-1") + assert result is app + assert app.app_model_config_id == app_model_config_id + configs = list(service_session.scalars(select(AppModelConfig))) + assert [config.id for config in configs] == [app_model_config_id] -def test_create_or_update_app_flushes_new_model_config_before_signal(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_or_update_app_flushes_new_model_config_before_signal( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: events: list[str] = [] - session = Mock() - session.add.side_effect = lambda _config: events.append("add") - session.flush.side_effect = lambda: events.append("flush") + + def record_flush(_session: Session, _flush_context: object) -> None: + events.append("flush") + + def record_signal(*_args: object, **_kwargs: object) -> None: + events.append("signal") + + event.listen(sqlite_session, "after_flush", record_flush) signal = Mock() - signal.send.side_effect = lambda *_args, **_kwargs: events.append("signal") + signal.send.side_effect = record_signal monkeypatch.setattr("services.app_dsl_service.app_model_config_was_updated", signal) app = cast( App, SimpleNamespace( - id="app-1", - tenant_id="tenant-1", + id="11111111-1111-1111-1111-111111111111", + tenant_id="33333333-3333-3333-3333-333333333333", app_model_config_id=None, name="Existing app", description="", @@ -121,25 +285,37 @@ def test_create_or_update_app_flushes_new_model_config_before_signal(monkeypatch ), ) - AppDslService(session=session)._create_or_update_app( - app=app, - data={"app": {"mode": AppMode.CHAT}, "model_config": {"model": {}}}, - account=Mock(id="account-1"), - ) + try: + AppDslService(session=sqlite_session)._create_or_update_app( + app=app, + data={"app": {"mode": AppMode.CHAT}, "model_config": {"model": {}}}, + account=Mock(id="22222222-2222-2222-2222-222222222222"), + ) + finally: + event.remove(sqlite_session, "after_flush", record_flush) - assert events == ["add", "flush", "signal"] - assert signal.send.call_args.kwargs["session"] is session - session.commit.assert_not_called() + assert events == ["flush", "signal"] + assert signal.send.call_args.kwargs["session"] is sqlite_session + assert app.app_model_config_id is not None + assert sqlite_session.get(AppModelConfig, app.app_model_config_id) is not None + assert sqlite_session.in_transaction() def test_export_dsl_loads_model_config_and_annotation_reply_with_request_session( monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], ) -> None: - model_config = {"model": {}, "agent_mode": {"tools": []}} - app_model_config = Mock(app_id="app-1") - app_model_config.to_dict.return_value = model_config - session = Mock() - session.get.return_value = app_model_config + model_config = cast(AppModelConfigDict, {"model": {}, "agent_mode": {"tools": []}}) + with sqlite_session_factory() as arrange_session: + app_model_config = AppModelConfig( + app_id="11111111-1111-1111-1111-111111111111", + created_by="22222222-2222-2222-2222-222222222222", + updated_by="22222222-2222-2222-2222-222222222222", + ).from_model_config_dict(model_config) + arrange_session.add(app_model_config) + arrange_session.commit() + app_model_config_id = app_model_config.id + app_id = app_model_config.app_id annotation_reply = {"enabled": False} load_annotation_reply_config = Mock(return_value=annotation_reply) monkeypatch.setattr("services.app_dsl_service.load_annotation_reply_config", load_annotation_reply_config) @@ -150,9 +326,9 @@ def test_export_dsl_loads_model_config_and_annotation_reply_with_request_session app = cast( App, SimpleNamespace( - id="app-1", - tenant_id="tenant-1", - app_model_config_id="config-1", + id="11111111-1111-1111-1111-111111111111", + tenant_id="33333333-3333-3333-3333-333333333333", + app_model_config_id=app_model_config_id, mode=AppMode.CHAT, name="Chat app", icon_type=IconType.EMOJI, @@ -163,11 +339,13 @@ def test_export_dsl_loads_model_config_and_annotation_reply_with_request_session ), ) - AppDslService.export_dsl(app, session=session) + with sqlite_session_factory() as service_session: + exported = AppDslService.export_dsl(app, session=service_session) - session.get.assert_called_once_with(AppModelConfig, "config-1") - load_annotation_reply_config.assert_called_once_with(session, "app-1") - app_model_config.to_dict.assert_called_once_with(annotation_reply=annotation_reply) + export_data = yaml.safe_load(exported) + assert export_data["model_config"]["model"] == {} + assert export_data["model_config"]["annotation_reply"] == annotation_reply + load_annotation_reply_config.assert_called_once_with(service_session, app_id) def test_ensure_agent_manage_permission_noops_when_rbac_disabled(monkeypatch: pytest.MonkeyPatch) -> None: @@ -198,11 +376,12 @@ def test_ensure_agent_manage_permission_rejects_without_agent_manage(monkeypatch AppDslService._ensure_agent_manage_permission(Mock(id="account-1", current_tenant_id="tenant-1")) -def test_create_or_update_app_gates_agent_mode_before_creation(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_or_update_app_gates_agent_mode_before_creation( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: monkeypatch.setattr("services.app_dsl_service.dify_config.RBAC_ENABLED", True) monkeypatch.setattr("services.app_dsl_service.RBACService.CheckAccess.check", Mock(return_value=False)) - session = Mock() - service = AppDslService(session=session) + service = AppDslService(session=unbound_session) with pytest.raises(NoPermissionError): service._create_or_update_app( @@ -211,14 +390,15 @@ def test_create_or_update_app_gates_agent_mode_before_creation(monkeypatch: pyte account=Mock(id="account-1", current_tenant_id="tenant-1"), ) - session.add.assert_not_called() - session.flush.assert_not_called() + assert not unbound_session.in_transaction() -def test_import_app_reraises_permission_denial_instead_of_failed_result(monkeypatch: pytest.MonkeyPatch) -> None: +def test_import_app_reraises_permission_denial_instead_of_failed_result( + monkeypatch: pytest.MonkeyPatch, unbound_session: Session +) -> None: monkeypatch.setattr("services.app_dsl_service.dify_config.RBAC_ENABLED", True) monkeypatch.setattr("services.app_dsl_service.RBACService.CheckAccess.check", Mock(return_value=False)) - service = AppDslService(session=Mock()) + service = AppDslService(session=unbound_session) with pytest.raises(NoPermissionError): service.import_app( @@ -226,3 +406,22 @@ def test_import_app_reraises_permission_denial_instead_of_failed_result(monkeypa import_mode="yaml-content", yaml_content="app:\n mode: agent\n name: Denied agent\n", ) + + assert not unbound_session.in_transaction() + + +def test_append_workflow_export_data_reports_missing_selected_workflow(monkeypatch: pytest.MonkeyPatch) -> None: + workflow_id = "11111111-1111-4111-8111-111111111111" + workflow_service = Mock() + workflow_service.get_draft_workflow.return_value = None + monkeypatch.setattr("services.app_dsl_service.WorkflowService", Mock(return_value=workflow_service)) + app = cast(App, SimpleNamespace(id="app-1", tenant_id="tenant-1")) + + with pytest.raises(WorkflowNotFoundError, match=f"Workflow version not found. Workflow ID: {workflow_id}"): + AppDslService._append_workflow_export_data( + export_data={}, + app_model=app, + include_secret=False, + session=Mock(), + workflow_id=workflow_id, + ) diff --git a/api/tests/unit_tests/services/test_app_generate_service.py b/api/tests/unit_tests/services/test_app_generate_service.py index 507977f287f..ff8eceed664 100644 --- a/api/tests/unit_tests/services/test_app_generate_service.py +++ b/api/tests/unit_tests/services/test_app_generate_service.py @@ -24,7 +24,7 @@ from pytest_mock import MockerFixture import services.app_generate_service as ags_module from core.app.entities.app_invoke_entities import InvokeFrom -from enums.quota_type import QuotaType +from enums import DeploymentEdition, QuotaType from models.model import AppMode from services.app_generate_service import AppGenerateService from services.errors.app import WorkflowIdFormatError, WorkflowNotFoundError @@ -217,7 +217,7 @@ class TestGenerate: @pytest.fixture(autouse=True) def _common(self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(ags_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) mocker.patch("services.app_generate_service.RateLimit", _DummyRateLimit) # Prevent AppExecutionParams.new from touching real models via isinstance mocker.patch( @@ -486,8 +486,8 @@ class TestGenerateBilling: _noop_rate_limit_context, ) - def test_billing_enabled_consumes_quota(self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", True) + def test_cloud_edition_consumes_quota(self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(ags_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) quota_charge = MagicMock() reserve_mock = mocker.patch( "services.app_generate_service.QuotaService.reserve", @@ -519,7 +519,7 @@ class TestGenerateBilling: from services.errors.app import QuotaExceededError from services.errors.llm import InvokeRateLimitError - monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(ags_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) mocker.patch( "services.app_generate_service.QuotaService.reserve", side_effect=QuotaExceededError(feature="workflow", tenant_id="t", required=1), @@ -536,7 +536,7 @@ class TestGenerateBilling: ) def test_exception_refunds_quota_and_exits_rate_limit(self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(ags_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) quota_charge = MagicMock() mocker.patch( "services.app_generate_service.QuotaService.reserve", @@ -566,7 +566,7 @@ class TestGenerateBilling: self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch ): """For non-streaming (blocking) calls, rate_limit.exit should be called in finally.""" - monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(ags_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) exit_calls: list[str] = [] @@ -596,7 +596,7 @@ class TestGenerateBilling: assert exit_calls == ["dummy-request-id"] def test_blocking_failure_exits_rate_limit_once(self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(ags_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) quota_charge = MagicMock() mocker.patch( "services.app_generate_service.QuotaService.reserve", @@ -628,7 +628,7 @@ class TestGenerateBilling: assert exit_calls == ["dummy-request-id"] def test_streaming_failure_exits_rate_limit_once(self, mocker: MockerFixture, monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(ags_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(ags_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) quota_charge = MagicMock() mocker.patch( "services.app_generate_service.QuotaService.reserve", diff --git a/api/tests/unit_tests/services/test_app_service.py b/api/tests/unit_tests/services/test_app_service.py index 9707855a6d7..171c9a94a00 100644 --- a/api/tests/unit_tests/services/test_app_service.py +++ b/api/tests/unit_tests/services/test_app_service.py @@ -1,27 +1,98 @@ from __future__ import annotations from collections.abc import Callable +from datetime import datetime from types import SimpleNamespace -from typing import cast from unittest.mock import MagicMock, patch +from uuid import uuid4 import pytest -from sqlalchemy.exc import IntegrityError +from sqlalchemy import event +from sqlalchemy.orm import Session +from enums import DeploymentEdition from graphon.model_runtime.entities.model_entities import ModelType -from models import Account -from models.model import App, AppMode, AppModelConfig +from models import Account, Tenant +from models.account import TenantAccountJoin, TenantAccountRole +from models.agent import Agent, AgentIconType, AgentScope, AgentSource, AgentStatus +from models.model import App, AppMode, AppModelConfig, IconType from models.workflow import Workflow -from services.agent.errors import AgentNameConflictError -from services.app_service import AppService, CreateAppParams +from services.agent.errors import AgentAccessNotReadyError, AgentNameConflictError +from services.app_service import AppListParams, AppService, CreateAppParams + + +def _persist_account(session: Session) -> Account: + tenant = Tenant(name="App Service Workspace") + account = Account(name="Test Account", email=f"app-service-{uuid4()}@example.com") + membership = TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + current=True, + role=TenantAccountRole.OWNER, + ) + account._current_tenant = tenant + session.add_all([tenant, account, membership]) + session.commit() + return account + + +def _persist_app(session: Session, *, tenant_id: str, name: str = "Visible App") -> App: + app = App( + id=str(uuid4()), + tenant_id=tenant_id, + name=name, + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="chat", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + ) + session.add(app) + session.commit() + return app + + +def _persist_agent_app(session: Session, *, app_name: str = "Old", agent_name: str = "Old") -> tuple[App, Agent]: + tenant_id = str(uuid4()) + creator_id = str(uuid4()) + app = App( + id=str(uuid4()), + tenant_id=tenant_id, + name=app_name, + description="old", + mode=AppMode.AGENT, + icon_type=IconType.EMOJI, + icon="robot", + icon_background="#fff", + enable_site=False, + enable_api=False, + created_by=creator_id, + ) + agent = Agent( + tenant_id=tenant_id, + name=agent_name, + description="old", + role="research assistant", + scope=AgentScope.ROSTER, + source=AgentSource.AGENT_APP, + status=AgentStatus.ACTIVE, + icon_type=AgentIconType.EMOJI, + icon="robot", + icon_background="#fff", + app_id=app.id, + created_by=creator_id, + ) + session.add_all([app, agent]) + session.commit() + return app, agent class TestCreateAppTransactionBoundary: - def test_commits_database_state_before_external_side_effects(self) -> None: - session = MagicMock() - account = MagicMock(spec=Account, id="account-1", current_tenant_id="tenant-1") + def test_commits_database_state_before_external_side_effects(self, sqlite_session: Session) -> None: + account = _persist_account(sqlite_session) phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") + event.listen(sqlite_session, "after_commit", lambda _session: phase_events.append("commit")) with ( patch( @@ -36,20 +107,20 @@ class TestCreateAppTransactionBoundary: "services.app_service.FeatureService.get_system_features", return_value=SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ), - patch("services.app_service.dify_config.BILLING_ENABLED", False), + patch("services.app_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), ): - AppService().create_app( - "tenant-1", + app = AppService().create_app( + account.current_tenant_id, CreateAppParams(name="Workflow", mode=AppMode.WORKFLOW.value), account, - session=session, + session=sqlite_session, ) assert phase_events == ["commit", "signal", "commit", "external"] + assert sqlite_session.get(App, app.id) is app - def test_falls_back_when_default_model_schema_is_unavailable(self) -> None: - session = MagicMock() - account = MagicMock(spec=Account, id="account-1", current_tenant_id="tenant-1") + def test_falls_back_when_default_model_schema_is_unavailable(self, sqlite_session: Session) -> None: + account = _persist_account(sqlite_session) model_type_instance = MagicMock() model_type_instance.get_model_schema.side_effect = ValueError("Base model unknown-model not found") model_instance = SimpleNamespace( @@ -61,9 +132,6 @@ class TestCreateAppTransactionBoundary: model_manager = MagicMock() model_manager.get_default_model_instance.return_value = model_instance model_manager.get_default_provider_model_name.return_value = ("openai", "gpt-4o") - added_objects: list[object] = [] - session.add.side_effect = added_objects.append - with ( patch("services.app_service.ModelManager.for_tenant", return_value=model_manager), patch("services.app_service.app_was_created.send"), @@ -72,16 +140,17 @@ class TestCreateAppTransactionBoundary: "services.app_service.FeatureService.get_system_features", return_value=SimpleNamespace(webapp_auth=SimpleNamespace(enabled=False)), ), - patch("services.app_service.dify_config.BILLING_ENABLED", False), + patch("services.app_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), ): app = AppService().create_app( - "tenant-1", + account.current_tenant_id, CreateAppParams(name="Chat", mode=AppMode.CHAT.value), account, - session=session, + session=sqlite_session, ) - app_model_config = next(obj for obj in added_objects if isinstance(obj, AppModelConfig)) + app_model_config = sqlite_session.get(AppModelConfig, app.app_model_config_id) + assert app_model_config is not None assert app.mode == AppMode.CHAT assert app_model_config.model_dict == { "provider": "openai", @@ -90,7 +159,7 @@ class TestCreateAppTransactionBoundary: "completion_params": {}, } model_manager.get_default_provider_model_name.assert_called_once_with( - tenant_id="tenant-1", model_type=ModelType.LLM + tenant_id=account.current_tenant_id, model_type=ModelType.LLM ) @@ -98,21 +167,44 @@ class TestCreateAppTransactionBoundary: "update_status", [AppService.update_app_site_status, AppService.update_app_api_status], ) -def test_app_status_updates_commit_before_signal(update_status: Callable[..., App]) -> None: - app = cast(App, SimpleNamespace(enable_site=False, enable_api=False)) - session = MagicMock() +def test_app_status_updates_commit_before_signal(update_status: Callable[..., App], sqlite_session: Session) -> None: + account = _persist_account(sqlite_session) + app = _persist_app(sqlite_session, tenant_id=account.current_tenant_id or "") phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") + event.listen(sqlite_session, "after_commit", lambda _session: phase_events.append("commit")) with ( - patch("services.app_service.current_user", SimpleNamespace(id="account-1")), + patch("services.app_service.current_user", account), patch("services.app_service.app_was_updated.send", side_effect=lambda *_args: phase_events.append("signal")), ): - update_status(AppService(), app, True, session=session) + update_status(AppService(), app, True, session=sqlite_session) assert phase_events == ["commit", "signal"] +@pytest.mark.parametrize( + "update_status", + [ + AppService.update_app_site_status, + AppService.update_app_api_status, + ], +) +def test_unpublished_agent_app_access_cannot_be_enabled( + update_status: Callable[..., App], sqlite_session: Session +) -> None: + app, _ = _persist_agent_app(sqlite_session) + commits: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: commits.append("commit")) + + with patch("services.app_service.agent_has_workflow_callable_active_snapshot", return_value=False): + with pytest.raises(AgentAccessNotReadyError): + update_status(AppService(), app, True, session=sqlite_session) + + assert app.enable_site is False + assert app.enable_api is False + assert commits == [] + + class TestOpenapiVisibilityHelpers: """Coverage for the session-injected, openapi-visibility-scoped ``AppService`` getters used by ``/openapi/v1/apps*``. These helpers @@ -120,170 +212,252 @@ class TestOpenapiVisibilityHelpers: gate passes" check so the controller can stay free of SQL. """ - def test_get_app_by_id_is_plain_session_get(self): + def test_get_app_by_id_is_plain_session_get(self, sqlite_session: Session): """``get_app_by_id`` must NOT apply status / visibility filters — callers (e.g. the openapi auth pipeline) need to differentiate 404 (missing) from 403 (``enable_api`` off) and would lose that signal if the helper coalesced both into ``None``. """ - mock_session = MagicMock() - sentinel_app = MagicMock(spec=App) - sentinel_app.status = "archived" # explicitly NOT "normal" - mock_session.get.return_value = sentinel_app + sentinel_app = _persist_app(sqlite_session, tenant_id=str(uuid4())) + sentinel_app.status = "archived" # type: ignore[assignment] - assert AppService.get_app_by_id("app-uuid", mock_session) is sentinel_app - mock_session.get.assert_called_once_with(App, "app-uuid") + assert AppService.get_app_by_id(sentinel_app.id, sqlite_session) is sentinel_app - def test_get_app_by_id_returns_none_when_missing(self): - mock_session = MagicMock() - mock_session.get.return_value = None + def test_get_app_by_id_returns_none_when_missing(self, sqlite_session: Session): + assert AppService.get_app_by_id(str(uuid4()), sqlite_session) is None - assert AppService.get_app_by_id("missing", mock_session) is None - - def test_get_visible_app_by_id_returns_app_when_visible(self): - mock_session = MagicMock() - app = MagicMock(spec=App) - app.status = "normal" - mock_session.get.return_value = app + def test_get_visible_app_by_id_returns_app_when_visible(self, sqlite_session: Session): + app = _persist_app(sqlite_session, tenant_id=str(uuid4())) with patch("services.app_service.is_openapi_visible", return_value=True): - assert AppService.get_visible_app_by_id("app-uuid", mock_session) is app + assert AppService.get_visible_app_by_id(app.id, sqlite_session) is app - mock_session.get.assert_called_once_with(App, "app-uuid") + def test_get_visible_app_by_id_returns_none_when_row_missing(self, sqlite_session: Session): + assert AppService.get_visible_app_by_id(str(uuid4()), sqlite_session) is None - def test_get_visible_app_by_id_returns_none_when_row_missing(self): - mock_session = MagicMock() - mock_session.get.return_value = None - - assert AppService.get_visible_app_by_id("missing", mock_session) is None - - def test_get_visible_app_by_id_returns_none_when_status_not_normal(self): + def test_get_visible_app_by_id_returns_none_when_status_not_normal(self, sqlite_session: Session): """Soft-deleted/archived rows must not surface on the openapi surface — the helper hides them by returning ``None``. """ - mock_session = MagicMock() - app = MagicMock(spec=App) - app.status = "archived" - mock_session.get.return_value = app + app = _persist_app(sqlite_session, tenant_id=str(uuid4())) + app.status = "archived" # type: ignore[assignment] with patch("services.app_service.is_openapi_visible", return_value=True): - assert AppService.get_visible_app_by_id("app-uuid", mock_session) is None + assert AppService.get_visible_app_by_id(app.id, sqlite_session) is None - def test_get_visible_app_by_id_returns_none_when_visibility_gate_rejects(self): + def test_get_visible_app_by_id_returns_none_when_visibility_gate_rejects(self, sqlite_session: Session): """``is_openapi_visible`` is the per-row counterpart to ``apply_openapi_gate`` — when it returns False the helper must treat the row as invisible (not "found but unauthorized"). """ - mock_session = MagicMock() - app = MagicMock(spec=App) - app.status = "normal" - mock_session.get.return_value = app + app = _persist_app(sqlite_session, tenant_id=str(uuid4())) with patch("services.app_service.is_openapi_visible", return_value=False): - assert AppService.get_visible_app_by_id("app-uuid", mock_session) is None + assert AppService.get_visible_app_by_id(app.id, sqlite_session) is None - def test_find_visible_apps_by_name_returns_scalars_through_visibility_gate(self): + def test_find_visible_apps_by_name_returns_scalars_through_visibility_gate(self, sqlite_session: Session): """Tenant-scoped name lookup. The helper passes the SELECT through ``apply_openapi_gate`` and materialises ``.scalars()`` into a list so the controller can branch on length (404 / single / 409). """ - mock_session = MagicMock() - rows = [MagicMock(spec=App), MagicMock(spec=App)] - mock_session.execute.return_value.scalars.return_value = iter(rows) + tenant_id = str(uuid4()) + rows = [ + _persist_app(sqlite_session, tenant_id=tenant_id, name="my-app"), + _persist_app(sqlite_session, tenant_id=tenant_id, name="my-app"), + ] + _persist_app(sqlite_session, tenant_id=str(uuid4()), name="my-app") with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q) as gate: - out = AppService.find_visible_apps_by_name(name="my-app", tenant_id="tenant-1", session=mock_session) + out = AppService.find_visible_apps_by_name(name="my-app", tenant_id=tenant_id, session=sqlite_session) - assert out == rows + assert {app.id for app in out} == {app.id for app in rows} # Visibility gate must wrap the SELECT exactly once. gate.assert_called_once() - mock_session.execute.assert_called_once() - - def test_find_visible_apps_by_name_returns_empty_list_on_no_match(self): - mock_session = MagicMock() - mock_session.execute.return_value.scalars.return_value = iter([]) + def test_find_visible_apps_by_name_returns_empty_list_on_no_match(self, sqlite_session: Session): with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q): - out = AppService.find_visible_apps_by_name(name="nope", tenant_id="tenant-1", session=mock_session) + out = AppService.find_visible_apps_by_name(name="nope", tenant_id=str(uuid4()), session=sqlite_session) assert out == [] - def test_find_visible_apps_by_ids_short_circuits_on_empty_input(self): + def test_find_visible_apps_by_ids_short_circuits_on_empty_input(self, unbound_session: Session): """Empty id list must not emit ``WHERE id IN ()`` — Postgres rejects empty IN lists and the call is a guaranteed no-op anyway. The helper returns ``[]`` without touching the session. """ - mock_session = MagicMock() + assert AppService.find_visible_apps_by_ids([], unbound_session) == [] - assert AppService.find_visible_apps_by_ids([], mock_session) == [] - mock_session.execute.assert_not_called() - - def test_find_visible_apps_by_ids_passes_through_visibility_gate(self): + def test_find_visible_apps_by_ids_passes_through_visibility_gate(self, sqlite_session: Session): """Bulk fetch routes through ``apply_openapi_gate`` exactly once and materialises the scalar rows. **No** status filter is applied here — the EE permitted-external pipeline filters non-normal hits in Python so its page count stays anchored. """ - mock_session = MagicMock() - rows = [MagicMock(spec=App), MagicMock(spec=App)] - mock_session.execute.return_value.scalars.return_value.all.return_value = rows + rows = [_persist_app(sqlite_session, tenant_id=str(uuid4())) for _ in range(2)] with patch("services.app_service.apply_openapi_gate", side_effect=lambda q: q) as gate: - out = AppService.find_visible_apps_by_ids(["a", "b"], mock_session) + out = AppService.find_visible_apps_by_ids([app.id for app in rows], sqlite_session) - assert out == rows + assert {app.id for app in out} == {app.id for app in rows} gate.assert_called_once() - mock_session.execute.assert_called_once() + + +def test_get_recent_apps_uses_one_tenant_scoped_projection_query(sqlite_session: Session) -> None: + tenant_id = str(uuid4()) + other_tenant_id = str(uuid4()) + account = Account(name="Recent Apps Author", email="recent-apps@example.com") + sqlite_session.add(account) + sqlite_session.flush() + + def create_app(*, name: str, tenant_id: str, updated_at: datetime, mode: AppMode = AppMode.CHAT) -> App: + app = App( + id=str(uuid4()), + tenant_id=tenant_id, + name=name, + description="", + mode=mode, + icon_type=IconType.EMOJI, + icon="🚀", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + created_by=account.id, + maintainer=account.id, + created_at=updated_at, + updated_at=updated_at, + use_icon_as_answer_icon=False, + ) + return app + + newest = create_app(name="Newest", tenant_id=tenant_id, updated_at=datetime(2026, 7, 3)) + legacy_agent = AppModelConfig( + app_id=newest.id, + agent_mode='{"enabled": true, "strategy": "react"}', + ) + newest.app_model_config_id = legacy_agent.id + second = create_app( + name="Second", + tenant_id=tenant_id, + updated_at=datetime(2026, 7, 2), + mode=AppMode.WORKFLOW, + ) + second.icon_type = None + second.icon = None + second.icon_background = None + second.created_by = None + second.maintainer = None + channel = create_app( + name="Channel", + tenant_id=tenant_id, + updated_at=datetime(2026, 7, 5), + mode=AppMode.CHANNEL, + ) + rag_pipeline = create_app( + name="RAG Pipeline", + tenant_id=tenant_id, + updated_at=datetime(2026, 7, 4), + mode=AppMode.RAG_PIPELINE, + ) + oldest = create_app(name="Oldest", tenant_id=tenant_id, updated_at=datetime(2026, 7, 1)) + foreign = create_app(name="Foreign", tenant_id=other_tenant_id, updated_at=datetime(2026, 7, 4)) + sqlite_session.add_all([newest, legacy_agent, second, channel, rag_pipeline, oldest, foreign]) + sqlite_session.commit() + + statements: list[str] = [] + bind = sqlite_session.get_bind() + + def record_sql(_conn, _cursor, statement, _parameters, _context, _executemany) -> None: + statements.append(statement) + + event.listen(bind, "before_cursor_execute", record_sql) + try: + recent_apps = AppService().get_recent_apps( + account.id, + tenant_id, + AppListParams(limit=2), + sqlite_session, + ) + finally: + event.remove(bind, "before_cursor_execute", record_sql) + + assert [(app.name, app.mode, app.icon_type, app.author_name, app.maintainer) for app in recent_apps] == [ + ("Newest", AppMode.CHAT, IconType.EMOJI, "Recent Apps Author", account.id), + ("Second", AppMode.WORKFLOW, None, None, None), + ] + select_statements = [statement for statement in statements if statement.lstrip().upper().startswith("SELECT")] + assert len(select_statements) == 1 + assert "count(" not in select_statements[0].lower() + assert "app_model_configs" not in select_statements[0].lower() class TestAppMeta: - def test_loads_workflow_with_caller_session(self): - session = MagicMock() - session.get.return_value = SimpleNamespace(graph_dict={"nodes": []}) - app = cast(App, SimpleNamespace(mode=AppMode.WORKFLOW, workflow_id="workflow-1")) + def test_loads_workflow_with_caller_session(self, sqlite_session: Session): + tenant_id = str(uuid4()) + app = _persist_app(sqlite_session, tenant_id=tenant_id) + app.mode = AppMode.WORKFLOW + workflow = Workflow( + id=str(uuid4()), + tenant_id=tenant_id, + app_id=app.id, + type="workflow", + version="draft", + graph='{"nodes": []}', + features="{}", + created_by=str(uuid4()), + ) + app.workflow_id = workflow.id + sqlite_session.add(workflow) + sqlite_session.commit() - assert AppService().get_app_meta(app, session=session) == {"tool_icons": {}} + assert AppService().get_app_meta(app, session=sqlite_session) == {"tool_icons": {}} - session.get.assert_called_once_with(Workflow, "workflow-1") + def test_loads_app_model_config_with_caller_session(self, sqlite_session: Session): + app = _persist_app(sqlite_session, tenant_id=str(uuid4())) + config = AppModelConfig(app_id=app.id, agent_mode='{"tools": []}') + sqlite_session.add(config) + sqlite_session.flush() + app.app_model_config_id = config.id + sqlite_session.commit() - def test_loads_app_model_config_with_caller_session(self): - session = MagicMock() - session.get.return_value = SimpleNamespace(agent_mode_dict={"tools": []}) - app = cast(App, SimpleNamespace(mode=AppMode.CHAT, app_model_config_id="config-1")) - - assert AppService().get_app_meta(app, session=session) == {"tool_icons": {}} - - session.get.assert_called_once_with(AppModelConfig, "config-1") + assert AppService().get_app_meta(app, session=sqlite_session) == {"tool_icons": {}} class TestGetApp: - def test_legacy_agent_detection_uses_caller_session(self): - session = MagicMock() - app = MagicMock(spec=App) - app.mode = AppMode.CHAT - app.is_agent_with_session.return_value = False - account = MagicMock(spec=Account) - account.current_tenant_id = "tenant-1" + def test_legacy_agent_detection_uses_caller_session(self, unbound_session: Session): + app = App( + mode=AppMode.CHAT, + ) + account = Account(name="Test Account", email="test@example.com") + account._current_tenant = Tenant(name="Test Tenant") + account._current_tenant.id = "tenant-1" - with patch("services.app_service.current_user", account): - assert AppService().get_app(app, session=session) is app + with ( + patch.object(App, "is_agent_with_session", return_value=False) as is_agent, + patch.object(App, "app_model_config_with_session") as get_model_config, + patch("services.app_service.current_user", account), + ): + assert AppService().get_app(app, session=unbound_session) is app - app.is_agent_with_session.assert_called_once_with(session=session) - app.app_model_config_with_session.assert_not_called() + is_agent.assert_called_once_with(session=unbound_session) + get_model_config.assert_not_called() - def test_agent_model_config_uses_caller_session(self): - session = MagicMock() - app = MagicMock(spec=App) - app.mode = AppMode.AGENT_CHAT - app.app_model_config_with_session.return_value = None - account = MagicMock(spec=Account) - account.current_tenant_id = "tenant-1" + def test_agent_model_config_uses_caller_session(self, unbound_session: Session): + app = App( + mode=AppMode.AGENT_CHAT, + ) + account = Account(name="Test Account", email="test@example.com") + account._current_tenant = Tenant(name="Test Tenant") + account._current_tenant.id = "tenant-1" - with patch("services.app_service.current_user", account): - assert AppService().get_app(app, session=session) is app + with ( + patch.object(App, "is_agent_with_session") as is_agent, + patch.object(App, "app_model_config_with_session", return_value=None) as get_model_config, + patch("services.app_service.current_user", account), + ): + assert AppService().get_app(app, session=unbound_session) is app - app.is_agent_with_session.assert_not_called() - app.app_model_config_with_session.assert_called_once_with(session=session) + is_agent.assert_not_called() + get_model_config.assert_called_once_with(session=unbound_session) class TestAgentAppType: @@ -298,6 +472,8 @@ class TestAgentAppType: # Runtime config comes from the Agent Soul, so no model_config is seeded. assert "model_config" not in default_app_templates[AppMode.AGENT] assert default_app_templates[AppMode.AGENT]["app"]["mode"] == AppMode.AGENT + assert default_app_templates[AppMode.AGENT]["app"]["enable_site"] is False + assert default_app_templates[AppMode.AGENT]["app"]["enable_api"] is False def test_create_app_params_accepts_agent_mode(self): from services.app_service import CreateAppParams @@ -309,47 +485,21 @@ class TestAgentAppType: """Non-agent apps short-circuit without touching the DB.""" from models.model import App, AppMode - app = App() - app.mode = AppMode.CHAT + app = App( + mode=AppMode.CHAT, + ) assert app.bound_agent_id is None - def test_update_agent_app_syncs_backing_agent_identity(self): - from models.agent import AgentIconType - from models.model import AppMode, IconType - from services.app_service import AppService - - app = SimpleNamespace( - id="app-1", - tenant_id="tenant-1", - mode=AppMode.AGENT, - name="Old", - description="old", - role="draft", - icon_type=IconType.EMOJI, - icon="robot", - icon_background="#fff", - use_icon_as_answer_icon=False, - max_active_requests=None, - created_by="account-1", - ) - backing_agent = SimpleNamespace( - name="Old", - description="old", - role="draft", - icon_type=AgentIconType.EMOJI, - icon="robot", - icon_background="#fff", - updated_by=None, - updated_at=None, - ) + def test_update_agent_app_syncs_backing_agent_identity(self, sqlite_session: Session): + app, backing_agent = _persist_agent_app(sqlite_session) + account_id = str(uuid4()) with ( - patch("services.app_service.db") as mock_db, - patch("services.app_service.current_user", SimpleNamespace(id="account-2")), + patch("services.app_service.current_user", SimpleNamespace(id=account_id)), + patch("services.app_service.app_was_updated.send"), ): - mock_db.session.scalar.return_value = backing_agent updated_app = AppService().update_app( - app, # type: ignore[arg-type] + app, { "name": "Iris", "description": "agent app", @@ -360,7 +510,7 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, - session=mock_db.session, + session=sqlite_session, ) assert updated_app.name == "Iris" @@ -370,46 +520,18 @@ class TestAgentAppType: assert backing_agent.icon_type == AgentIconType.IMAGE assert backing_agent.icon == "file-id" assert backing_agent.icon_background == "#123456" - assert backing_agent.updated_by == "account-2" + assert backing_agent.updated_by == account_id assert backing_agent.updated_at == updated_app.updated_at - def test_update_agent_app_preserves_role_when_args_omit_it(self): - from models.agent import AgentIconType - from models.model import AppMode, IconType - from services.app_service import AppService - - app = SimpleNamespace( - id="app-1", - tenant_id="tenant-1", - mode=AppMode.AGENT, - name="Old", - description="old", - role="draft", - icon_type=IconType.EMOJI, - icon="robot", - icon_background="#fff", - use_icon_as_answer_icon=False, - max_active_requests=None, - created_by="account-1", - ) - backing_agent = SimpleNamespace( - name="Old", - description="old", - role="research assistant", - icon_type=AgentIconType.EMOJI, - icon="robot", - icon_background="#fff", - updated_by=None, - updated_at=None, - ) + def test_update_agent_app_preserves_role_when_args_omit_it(self, sqlite_session: Session): + app, backing_agent = _persist_agent_app(sqlite_session) with ( - patch("services.app_service.db") as mock_db, - patch("services.app_service.current_user", SimpleNamespace(id="account-2")), + patch("services.app_service.current_user", SimpleNamespace(id=str(uuid4()))), + patch("services.app_service.app_was_updated.send"), ): - mock_db.session.scalar.return_value = backing_agent AppService().update_app( - app, # type: ignore[arg-type] + app, { "name": "Iris", "description": "agent app", @@ -419,48 +541,20 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, - session=mock_db.session, + session=sqlite_session, ) assert backing_agent.role == "research assistant" - def test_update_agent_app_clears_role_when_args_set_empty_string(self): - from models.agent import AgentIconType - from models.model import AppMode, IconType - from services.app_service import AppService - - app = SimpleNamespace( - id="app-1", - tenant_id="tenant-1", - mode=AppMode.AGENT, - name="Old", - description="old", - role="draft", - icon_type=IconType.EMOJI, - icon="robot", - icon_background="#fff", - use_icon_as_answer_icon=False, - max_active_requests=None, - created_by="account-1", - ) - backing_agent = SimpleNamespace( - name="Old", - description="old", - role="research assistant", - icon_type=AgentIconType.EMOJI, - icon="robot", - icon_background="#fff", - updated_by=None, - updated_at=None, - ) + def test_update_agent_app_clears_role_when_args_set_empty_string(self, sqlite_session: Session): + app, backing_agent = _persist_agent_app(sqlite_session) with ( - patch("services.app_service.db") as mock_db, - patch("services.app_service.current_user", SimpleNamespace(id="account-2")), + patch("services.app_service.current_user", SimpleNamespace(id=str(uuid4()))), + patch("services.app_service.app_was_updated.send"), ): - mock_db.session.scalar.return_value = backing_agent AppService().update_app( - app, # type: ignore[arg-type] + app, { "name": "Iris", "description": "agent app", @@ -471,50 +565,34 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, - session=mock_db.session, + session=sqlite_session, ) assert backing_agent.role == "" - def test_update_agent_app_duplicate_name_rolls_back_and_raises_conflict(self): - from models.agent import AgentIconType - from models.model import AppMode, IconType - from services.app_service import AppService - - app = SimpleNamespace( - id="app-1", - tenant_id="tenant-1", - mode=AppMode.AGENT, - name="Old", - description="old", - role="draft", - icon_type=IconType.EMOJI, - icon="robot", - icon_background="#fff", - use_icon_as_answer_icon=False, - max_active_requests=None, - created_by="account-1", - ) - backing_agent = SimpleNamespace( - name="Old", - description="old", - role="research assistant", - icon_type=AgentIconType.EMOJI, - icon="robot", - icon_background="#fff", - updated_by=None, - updated_at=None, + def test_update_agent_app_duplicate_name_rolls_back_and_raises_conflict(self, sqlite_session: Session): + app, backing_agent = _persist_agent_app(sqlite_session) + existing = Agent( + tenant_id=app.tenant_id, + name="Existing Agent", + description="existing", + role="", + scope=AgentScope.ROSTER, + source=AgentSource.ROSTER, + status=AgentStatus.ACTIVE, ) + sqlite_session.add(existing) + sqlite_session.commit() + rollback_events: list[str] = [] + event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback")) with ( - patch("services.app_service.db") as mock_db, - patch("services.app_service.current_user", SimpleNamespace(id="account-2")), + patch("services.app_service.current_user", SimpleNamespace(id=str(uuid4()))), + patch("services.app_service.app_was_updated.send"), ): - mock_db.session.scalar.return_value = backing_agent - mock_db.session.commit.side_effect = IntegrityError("duplicate", None, None) with pytest.raises(AgentNameConflictError): AppService().update_app( - app, # type: ignore[arg-type] + app, { "name": "Existing Agent", "description": "agent app", @@ -525,32 +603,124 @@ class TestAgentAppType: "use_icon_as_answer_icon": False, "max_active_requests": 0, }, - session=mock_db.session, + session=sqlite_session, ) - mock_db.session.rollback.assert_called_once() + assert rollback_events == ["rollback"] + sqlite_session.expire_all() + assert sqlite_session.get(Agent, backing_agent.id).name == "Old" # type: ignore[union-attr] - def test_delete_agent_app_archives_backing_agent(self): - from models.agent import AgentStatus - from models.model import AppMode - from services.app_service import AppService - - app = SimpleNamespace(id="app-1", tenant_id="tenant-1", mode=AppMode.AGENT) - backing_agent = SimpleNamespace(status=AgentStatus.ACTIVE, archived_by=None, archived_at=None) + def test_delete_agent_app_archives_backing_agent(self, sqlite_session: Session): + app, backing_agent = _persist_agent_app(sqlite_session) + workflow_agents = [ + Agent( + tenant_id=app.tenant_id, + name=f"Workflow Agent {index}", + description="", + role="", + scope=AgentScope.WORKFLOW_ONLY, + source=AgentSource.WORKFLOW, + status=AgentStatus.ACTIVE, + app_id=app.id, + workflow_id=str(uuid4()), + workflow_node_id=f"node-{index}", + ) + for index in range(2) + ] + sqlite_session.add_all(workflow_agents) + sqlite_session.commit() + account_id = str(uuid4()) + events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: events.append("commit")) with ( - patch("services.app_service.db") as mock_db, - patch("services.app_service.current_user", SimpleNamespace(id="account-2")), + patch("services.app_service.current_user", SimpleNamespace(id=account_id)), + patch("services.app_service.app_was_deleted.send"), patch("services.app_service.BillingService"), patch("services.app_service.EnterpriseService"), patch("services.app_service.FeatureService"), patch("services.app_service.dify_config"), patch("services.app_service.remove_app_and_related_data_task"), + patch( + "services.app_service.AgentHomeSnapshotService.retire_all_for_agent", + return_value=["home-1"], + ) as mock_retire_homes, + patch( + "services.app_service.AgentWorkspaceService.retire_all_for_app", + side_effect=lambda **_kwargs: events.append("retire-app-workspaces") or ["workspace-1"], + ) as mock_retire_workspaces, + patch( + "services.app_service.WorkflowAgentRetirementService.retire_unowned", + side_effect=lambda **_kwargs: ( + events.append("retire-workflow-agents") or (["workflow-binding-1"], ["workflow-home-1"]) + ), + ) as mock_workflow_retirement, + patch( + "services.app_service.enqueue_agent_resource_collection", + side_effect=lambda **_kwargs: events.append("enqueue"), + ) as mock_enqueue_collection, ): - mock_db.session.scalar.return_value = backing_agent - AppService().delete_app(app, session=mock_db.session) # type: ignore[arg-type] + AppService().delete_app(app, session=sqlite_session) - assert backing_agent.status == AgentStatus.ARCHIVED - assert backing_agent.archived_by == "account-2" - assert backing_agent.archived_at is not None - mock_db.session.delete.assert_called_once_with(app) + assert events == ["retire-app-workspaces", "commit", "retire-workflow-agents", "enqueue"] + sqlite_session.expire_all() + persisted_agent = sqlite_session.get(Agent, backing_agent.id) + assert persisted_agent is not None + assert sqlite_session.get(App, app.id) is None + assert persisted_agent.status == AgentStatus.ARCHIVED + assert persisted_agent.archived_by == account_id + assert persisted_agent.archived_at is not None + mock_workflow_retirement.assert_called_once_with( + tenant_id=app.tenant_id, + agent_ids=[agent.id for agent in workflow_agents], + account_id=account_id, + ) + mock_retire_workspaces.assert_called_once_with( + session=sqlite_session, + tenant_id=app.tenant_id, + app_id=app.id, + ) + mock_retire_homes.assert_called_once_with( + session=sqlite_session, + tenant_id=app.tenant_id, + agent_id=backing_agent.id, + ) + mock_enqueue_collection.assert_called_once_with( + tenant_id=app.tenant_id, + workspace_ids=["workspace-1"], + binding_ids=["workflow-binding-1"], + home_snapshot_ids=["home-1", "workflow-home-1"], + ) + + def test_delete_app_commit_failure_does_not_retire_workflow_agents_or_enqueue( + self, sqlite_session: Session, monkeypatch: pytest.MonkeyPatch + ): + app = _persist_app(sqlite_session, tenant_id=str(uuid4())) + app.mode = AppMode.WORKFLOW + workflow_agent = Agent( + tenant_id=app.tenant_id, + name="Workflow Agent", + description="", + role="", + scope=AgentScope.WORKFLOW_ONLY, + source=AgentSource.WORKFLOW, + status=AgentStatus.ACTIVE, + app_id=app.id, + workflow_id=str(uuid4()), + workflow_node_id="node-1", + ) + sqlite_session.add(workflow_agent) + sqlite_session.commit() + monkeypatch.setattr(sqlite_session, "commit", MagicMock(side_effect=RuntimeError("commit failed"))) + with ( + patch("services.app_service.current_user", SimpleNamespace(id=str(uuid4()))), + patch("services.app_service.app_was_deleted.send"), + patch("services.app_service.AgentWorkspaceService.retire_all_for_app", return_value=["workspace-1"]), + patch("services.app_service.WorkflowAgentRetirementService.retire_unowned") as retire_unowned, + patch("services.app_service.enqueue_agent_resource_collection") as enqueue_collection, + ): + with pytest.raises(RuntimeError, match="commit failed"): + AppService().delete_app(app, session=sqlite_session) + + retire_unowned.assert_not_called() + enqueue_collection.assert_not_called() diff --git a/api/tests/unit_tests/services/test_archive_workflow_run_logs.py b/api/tests/unit_tests/services/test_archive_workflow_run_logs.py index a21f1de769c..5177167ed0e 100644 --- a/api/tests/unit_tests/services/test_archive_workflow_run_logs.py +++ b/api/tests/unit_tests/services/test_archive_workflow_run_logs.py @@ -9,6 +9,8 @@ This module contains tests for: from datetime import datetime from unittest.mock import MagicMock, patch +from enums import DeploymentEdition + class TestWorkflowRunArchiver: """Tests for the WorkflowRunArchiver class.""" @@ -19,7 +21,7 @@ class TestWorkflowRunArchiver: """Test archiver can be initialized with various options.""" from services.retention.workflow_run.archive_paid_plan_workflow_run import WorkflowRunArchiver - mock_config.BILLING_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY archiver = WorkflowRunArchiver( days=90, diff --git a/api/tests/unit_tests/services/test_audio_service.py b/api/tests/unit_tests/services/test_audio_service.py index 0722faad651..1dd1277c2fb 100644 --- a/api/tests/unit_tests/services/test_audio_service.py +++ b/api/tests/unit_tests/services/test_audio_service.py @@ -224,6 +224,10 @@ def factory(): class TestAudioServiceASR: """Test speech-to-text (ASR) operations.""" + @pytest.fixture(autouse=True) + def _bind_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + @patch("services.audio_service.ModelManager.for_tenant", autospec=True) def test_transcript_asr_success_chat_mode(self, mock_model_manager_class, factory: AudioServiceTestDataFactory): """Test successful ASR transcription in CHAT mode.""" @@ -242,7 +246,7 @@ class TestAudioServiceASR: mock_model_manager.get_default_model_instance.return_value = mock_model_instance # Act - result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock(), end_user="user-123") + result = AudioService.transcript_asr(app_model=app, file=file, session=self.session, end_user="user-123") # Assert assert result == {"text": "Transcribed text"} @@ -269,7 +273,7 @@ class TestAudioServiceASR: mock_model_manager.get_default_model_instance.return_value = mock_model_instance # Act - result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) + result = AudioService.transcript_asr(app_model=app, file=file, session=self.session) # Assert assert result == {"text": "Workflow transcribed text"} @@ -290,7 +294,7 @@ class TestAudioServiceASR: mock_model_instance.invoke_speech2text.return_value = "Published Agent transcript" mock_model_manager_class.return_value.get_default_model_instance.return_value = mock_model_instance - result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock(), end_user="end-user-1") + result = AudioService.transcript_asr(app_model=app, file=file, session=self.session, end_user="end-user-1") assert result == {"text": "Published Agent transcript"} mock_roster_service_class.return_value.get_published_agent_soul_for_app.assert_called_once_with( @@ -314,7 +318,7 @@ class TestAudioServiceASR: mock_model_instance.invoke_speech2text.return_value = "Legacy Agent transcript" mock_model_manager_class.return_value.get_default_model_instance.return_value = mock_model_instance - result = AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) + result = AudioService.transcript_asr(app_model=app, file=file, session=self.session) assert result == {"text": "Legacy Agent transcript"} @@ -333,7 +337,7 @@ class TestAudioServiceASR: app_model=app, agent_soul=agent_soul, file=file, - session=MagicMock(), + session=self.session, end_user="account-1", ) @@ -354,7 +358,7 @@ class TestAudioServiceASR: file = factory.create_file_storage_mock() with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=MagicMock()) + AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=self.session) @patch("services.audio_service.ModelManager.for_tenant", autospec=True) def test_transcript_agent_asr_preserves_legacy_feature_fallback( @@ -372,7 +376,7 @@ class TestAudioServiceASR: app_model=app, agent_soul=AgentSoulConfig(), file=file, - session=MagicMock(), + session=self.session, ) assert result == {"text": "Legacy feature transcript"} @@ -385,7 +389,7 @@ class TestAudioServiceASR: agent_soul = AgentSoulConfig.model_validate({"app_features": {"speech_to_text": {"enabled": False}}}) with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=MagicMock()) + AudioService.transcript_agent_asr(app_model=app, agent_soul=agent_soul, file=file, session=self.session) def test_transcript_asr_raises_error_when_feature_disabled_chat_mode(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error when speech-to-text is disabled in CHAT mode.""" @@ -399,7 +403,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) + AudioService.transcript_asr(app_model=app, file=file, session=self.session) def test_transcript_asr_raises_error_when_feature_disabled_workflow_mode( self, factory: AudioServiceTestDataFactory @@ -415,7 +419,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) + AudioService.transcript_asr(app_model=app, file=file, session=self.session) def test_transcript_asr_raises_error_when_workflow_missing(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error when workflow is missing in WORKFLOW mode.""" @@ -428,7 +432,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(SpeechToTextDisabledServiceError): - AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) + AudioService.transcript_asr(app_model=app, file=file, session=self.session) def test_transcript_asr_raises_error_when_no_file_uploaded(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error when no file is uploaded.""" @@ -441,7 +445,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(NoAudioUploadedServiceError): - AudioService.transcript_asr(app_model=app, file=None, session=MagicMock()) + AudioService.transcript_asr(app_model=app, file=None, session=self.session) def test_transcript_asr_raises_error_for_unsupported_audio_type(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error for unsupported audio file types.""" @@ -455,7 +459,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(UnsupportedAudioTypeServiceError): - AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) + AudioService.transcript_asr(app_model=app, file=file, session=self.session) def test_transcript_asr_raises_error_for_large_file(self, factory: AudioServiceTestDataFactory): """Test that ASR raises error when file exceeds size limit (30MB).""" @@ -471,7 +475,7 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(AudioTooLargeServiceError, match="Audio size larger than 30 mb"): - AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) + AudioService.transcript_asr(app_model=app, file=file, session=self.session) @patch("services.audio_service.ModelManager.for_tenant", autospec=True) def test_transcript_asr_raises_error_when_no_model_instance( @@ -492,10 +496,9 @@ class TestAudioServiceASR: # Act & Assert with pytest.raises(ProviderNotSupportSpeechToTextServiceError): - AudioService.transcript_asr(app_model=app, file=file, session=MagicMock()) + AudioService.transcript_asr(app_model=app, file=file, session=self.session) -@pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) class TestAudioServiceTTS: """Test text-to-speech (TTS) operations.""" diff --git a/api/tests/unit_tests/services/test_batch_indexing_base.py b/api/tests/unit_tests/services/test_batch_indexing_base.py index 8d07fc30509..17420379427 100644 --- a/api/tests/unit_tests/services/test_batch_indexing_base.py +++ b/api/tests/unit_tests/services/test_batch_indexing_base.py @@ -5,7 +5,7 @@ from unittest.mock import MagicMock, patch import pytest from core.entities.document_task import DocumentTask -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from services.document_indexing_proxy.batch_indexing_base import BatchDocumentIndexingProxy # --------------------------------------------------------------------------- @@ -154,6 +154,7 @@ class TestSendToTenantQueue: """When get_task_key() is truthy, tasks must be pushed via push_tasks().""" # Arrange proxy = make_proxy() + # pyrefly: ignore [missing-attribute] proxy._tenant_isolated_task_queue.get_task_key.return_value = "existing-key" task_func = MagicMock() @@ -169,6 +170,7 @@ class TestSendToTenantQueue: """When a key already exists, task_func.delay must never be called.""" # Arrange proxy = make_proxy() + # pyrefly: ignore [missing-attribute] proxy._tenant_isolated_task_queue.get_task_key.return_value = "existing-key" task_func = MagicMock() @@ -182,6 +184,7 @@ class TestSendToTenantQueue: """When a key already exists, set_task_waiting_time must never be called.""" # Arrange proxy = make_proxy() + # pyrefly: ignore [missing-attribute] proxy._tenant_isolated_task_queue.get_task_key.return_value = "existing-key" task_func = MagicMock() @@ -196,6 +199,7 @@ class TestSendToTenantQueue: """Verify the serialised payload matches asdict(DocumentTask(...)).""" # Arrange proxy = make_proxy(document_ids=["doc-x"]) + # pyrefly: ignore [missing-attribute] proxy._tenant_isolated_task_queue.get_task_key.return_value = "k" task_func = MagicMock() @@ -219,6 +223,7 @@ class TestSendToTenantQueue: """When get_task_key() is falsy, set_task_waiting_time and task_func.delay are invoked.""" # Arrange proxy = make_proxy() + # pyrefly: ignore [missing-attribute] proxy._tenant_isolated_task_queue.get_task_key.return_value = None task_func = MagicMock() @@ -238,6 +243,7 @@ class TestSendToTenantQueue: """When get_task_key() is falsy, push_tasks must never be called.""" # Arrange proxy = make_proxy() + # pyrefly: ignore [missing-attribute] proxy._tenant_isolated_task_queue.get_task_key.return_value = None task_func = MagicMock() @@ -253,6 +259,7 @@ class TestSendToTenantQueue: """Verify that any falsy return from get_task_key() triggers the init branch.""" # Arrange proxy = make_proxy() + # pyrefly: ignore [missing-attribute] proxy._tenant_isolated_task_queue.get_task_key.return_value = falsy_key task_func = MagicMock() @@ -278,6 +285,7 @@ class TestDispatchRouting: """Sandbox plan routes to normal priority queue with tenant isolation.""" # Arrange proxy = make_proxy() + # pyrefly: ignore [missing-attribute] proxy._tenant_isolated_task_queue.get_task_key.return_value = None with patch("services.document_indexing_proxy.base.FeatureService.get_features") as mock_features: diff --git a/api/tests/unit_tests/services/test_billing_service.py b/api/tests/unit_tests/services/test_billing_service.py index a8d405ae8a3..b6d3a5c150e 100644 --- a/api/tests/unit_tests/services/test_billing_service.py +++ b/api/tests/unit_tests/services/test_billing_service.py @@ -23,7 +23,7 @@ import pytest from sqlalchemy.orm import Session from werkzeug.exceptions import InternalServerError -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from models import Account, Tenant, TenantAccountJoin, TenantAccountRole from services.billing_service import BillingService @@ -462,6 +462,37 @@ class TestBillingServiceSubscriptionInfo: params={"tenant_id": tenant_id}, ) + def test_get_vector_space_preserves_unknown_usage(self, mock_send_request): + tenant_id = "tenant-123" + expected_response = {"size": 0.0, "limit": 50, "usage_unknown": True} + mock_send_request.return_value = expected_response + + result = BillingService.get_vector_space(tenant_id) + + assert result == expected_response + + def test_get_info_preserves_unknown_vector_space_usage(self, mock_send_request): + tenant_id = "tenant-123" + expected_response = { + "enabled": True, + "subscription": {"plan": "sandbox", "interval": "", "education": False}, + "members": {"size": 1, "limit": 1}, + "apps": {"size": 1, "limit": 10}, + "vector_space": {"size": 0.0, "limit": 50, "usage_unknown": True}, + "knowledge_rate_limit": {"limit": 10}, + "documents_upload_quota": {"size": 1, "limit": 50}, + "annotation_quota_limit": {"size": 0, "limit": 10}, + "docs_processing": "standard", + "can_replace_logo": False, + "model_load_balancing_enabled": False, + "knowledge_pipeline_publish_enabled": False, + } + mock_send_request.return_value = expected_response + + result = BillingService.get_info(tenant_id) + + assert result["vector_space"]["usage_unknown"] is True + def test_get_vector_space_bypasses_cache(self, mock_send_request): tenant_id = "tenant-123" mock_send_request.return_value = {"size": 4096, "limit": 20480} @@ -1989,6 +2020,8 @@ class TestBillingServiceSubscriptionInfoDataType: if "vector_space" in result: assert isinstance(result["vector_space"]["size"], float) assert isinstance(result["vector_space"]["limit"], int) + if "usage_unknown" in result["vector_space"]: + assert isinstance(result["vector_space"]["usage_unknown"], bool) assert isinstance(result["knowledge_rate_limit"]["limit"], int) @@ -2055,3 +2088,17 @@ class TestBillingServiceSubscriptionInfoDataType: with pytest.raises(ValidationError): BillingService.get_info("tenant-type-test") + + +def test_pooled_billing_client_carries_bounded_timeout() -> None: + """Regression for #39874: the pooled billing client must carry a + read/connect timeout so a stalled Stripe / cloud-billing proxy + fails fast instead of pinning a worker. Same shape as the + JinaReader / WaterCrawl hardening that landed in PR #39860 and #39824. + """ + import services.billing_service as billing_service_module + + client = billing_service_module._http_client + assert client.timeout is not None + assert client.timeout.read == 30.0 + assert client.timeout.connect == 5.0 diff --git a/api/tests/unit_tests/services/test_clear_free_plan_expired_workflow_run_logs.py b/api/tests/unit_tests/services/test_clear_free_plan_expired_workflow_run_logs.py index 60488beb248..c50696a5f98 100644 --- a/api/tests/unit_tests/services/test_clear_free_plan_expired_workflow_run_logs.py +++ b/api/tests/unit_tests/services/test_clear_free_plan_expired_workflow_run_logs.py @@ -3,6 +3,7 @@ from typing import Any import pytest +from enums import DeploymentEdition from repositories.api_workflow_run_repository import WorkflowRunCleanupRef from services.billing_service import SubscriptionPlan from services.retention.workflow_run import clear_free_plan_expired_workflow_run_logs as cleanup_module @@ -126,10 +127,10 @@ def create_cleanup( return WorkflowRunCleanup(workflow_run_repo=repo, **kwargs) -def test_filter_free_tenants_billing_disabled(monkeypatch: pytest.MonkeyPatch) -> None: +def test_filter_free_tenants_outside_cloud_edition(monkeypatch: pytest.MonkeyPatch) -> None: cleanup = create_cleanup(monkeypatch, repo=FakeRepo([]), days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) def fail_bulk(_: list[str]) -> dict[str, SubscriptionPlan]: raise RuntimeError("should not call") @@ -145,7 +146,7 @@ def test_filter_free_tenants_billing_disabled(monkeypatch: pytest.MonkeyPatch) - def test_filter_free_tenants_bulk_mixed(monkeypatch: pytest.MonkeyPatch) -> None: cleanup = create_cleanup(monkeypatch, repo=FakeRepo([]), days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr( cleanup_module.BillingService, "get_plan_bulk_with_cache", @@ -165,7 +166,7 @@ def test_filter_free_tenants_bulk_mixed(monkeypatch: pytest.MonkeyPatch) -> None def test_filter_free_tenants_respects_grace_period(monkeypatch: pytest.MonkeyPatch) -> None: cleanup = create_cleanup(monkeypatch, repo=FakeRepo([]), days=30, batch_size=10, grace_period_days=45) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) now = datetime.datetime.now(datetime.UTC) within_grace_ts = int((now - datetime.timedelta(days=10)).timestamp()) outside_grace_ts = int((now - datetime.timedelta(days=90)).timestamp()) @@ -192,7 +193,7 @@ def test_filter_free_tenants_skips_cleanup_whitelist(monkeypatch: pytest.MonkeyP whitelist={"tenant_whitelist"}, ) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr( cleanup_module.BillingService, "get_plan_bulk_with_cache", @@ -213,7 +214,7 @@ def test_filter_free_tenants_skips_cleanup_whitelist(monkeypatch: pytest.MonkeyP def test_filter_free_tenants_bulk_failure(monkeypatch: pytest.MonkeyPatch) -> None: cleanup = create_cleanup(monkeypatch, repo=FakeRepo([]), days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr( cleanup_module.BillingService, "get_plan_bulk_with_cache", @@ -237,7 +238,7 @@ def test_run_deletes_only_free_tenants(monkeypatch: pytest.MonkeyPatch) -> None: ) cleanup = create_cleanup(monkeypatch, repo=repo, days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr( cleanup_module.BillingService, "get_plan_bulk_with_cache", @@ -267,7 +268,7 @@ def test_run_filters_candidate_tenants_before_target_query(monkeypatch: pytest.M ) cleanup = create_cleanup(monkeypatch, repo=repo, days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) billing_calls: list[list[str]] = [] def fake_bulk(tenant_ids: list[str]) -> dict[str, SubscriptionPlan]: @@ -291,7 +292,7 @@ def test_run_skips_when_no_free_tenants(monkeypatch: pytest.MonkeyPatch) -> None repo = FakeRepo(batches=[[make_ref("run-paid", "t_paid", cutoff)]]) cleanup = create_cleanup(monkeypatch, repo=repo, days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr( cleanup_module.BillingService, "get_plan_bulk_with_cache", @@ -309,7 +310,7 @@ def test_run_paid_only_records_skipped_metrics(monkeypatch: pytest.MonkeyPatch) repo = FakeRepo(batches=[[make_ref("run-paid", "t_paid", cutoff)]]) cleanup = create_cleanup(monkeypatch, repo=repo, days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr( cleanup_module.BillingService, "get_plan_bulk_with_cache", @@ -341,7 +342,7 @@ def test_run_target_query_is_bounded_by_candidate_high_water(monkeypatch: pytest ) cleanup = create_cleanup(monkeypatch, repo=repo, days=30, batch_size=2) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) cleanup.run() @@ -371,7 +372,7 @@ def test_run_records_metrics_on_success(monkeypatch: pytest.MonkeyPatch) -> None }, ) cleanup = create_cleanup(monkeypatch, repo=repo, days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) batch_calls: list[dict[str, object]] = [] completion_calls: list[dict[str, object]] = [] @@ -399,7 +400,7 @@ def test_run_records_failed_metrics(monkeypatch: pytest.MonkeyPatch) -> None: cutoff = datetime.datetime.now() repo = FailingRepo(batches=[[make_ref("run-free", "t_free", cutoff)]]) cleanup = create_cleanup(monkeypatch, repo=repo, days=30, batch_size=10) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) completion_calls: list[dict[str, object]] = [] monkeypatch.setattr(cleanup._metrics, "record_completion", lambda **kwargs: completion_calls.append(kwargs)) @@ -427,7 +428,7 @@ def test_run_dry_run_skips_deletions(monkeypatch: pytest.MonkeyPatch, capsys: py ) cleanup = create_cleanup(monkeypatch, repo=repo, days=30, batch_size=10, dry_run=True) - monkeypatch.setattr(cleanup_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(cleanup_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) cleanup.run() diff --git a/api/tests/unit_tests/services/test_clear_free_plan_tenant_expired_logs.py b/api/tests/unit_tests/services/test_clear_free_plan_tenant_expired_logs.py index 863d0b3aef6..d49cfa3bd35 100644 --- a/api/tests/unit_tests/services/test_clear_free_plan_tenant_expired_logs.py +++ b/api/tests/unit_tests/services/test_clear_free_plan_tenant_expired_logs.py @@ -11,7 +11,7 @@ from sqlalchemy import event from sqlalchemy.engine import Engine from sqlalchemy.orm import Session -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from graphon.file import FileTransferMethod, FileType from models.account import Tenant from models.enums import ( @@ -471,7 +471,7 @@ def test_process_with_tenant_ids_filters_by_plan_and_logs_errors( sqlite_session.commit() _configure_process_boundaries(monkeypatch, sqlite_engine) monkeypatch.setattr(service_module.click, "echo", MagicMock()) - monkeypatch.setattr(service_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) def fake_get_info(tenant_id: str) -> dict[str, dict[str, str]]: if tenant_id == "tenant-sandbox": @@ -524,7 +524,7 @@ def test_process_without_tenant_ids_batches_and_scales_interval( monkeypatch.setattr(service_module.datetime, "datetime", FixedDateTime) _configure_process_boundaries(monkeypatch, sqlite_engine) monkeypatch.setattr(service_module.click, "echo", lambda *_args, **_kwargs: None) - monkeypatch.setattr(service_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) process_tenant = MagicMock() monkeypatch.setattr(ClearFreePlanTenantExpiredLogs, "process_tenant", process_tenant) statements: list[str] = [] @@ -564,7 +564,7 @@ def test_process_with_tenant_ids_emits_progress_every_100( sqlite_session.add_all([_create_tenant(tenant_id) for tenant_id in tenant_ids]) sqlite_session.commit() _configure_process_boundaries(monkeypatch, sqlite_engine) - monkeypatch.setattr(service_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) echo = MagicMock() monkeypatch.setattr(service_module.click, "echo", echo) monkeypatch.setattr(ClearFreePlanTenantExpiredLogs, "process_tenant", MagicMock()) @@ -598,7 +598,7 @@ def test_process_without_tenant_ids_all_intervals_too_many_uses_min_interval( monkeypatch.setattr(service_module.datetime, "datetime", FixedDateTime) _configure_process_boundaries(monkeypatch, sqlite_engine) monkeypatch.setattr(service_module.click, "echo", lambda *_args, **_kwargs: None) - monkeypatch.setattr(service_module.dify_config, "BILLING_ENABLED", False) + monkeypatch.setattr(service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) process_tenant = MagicMock() monkeypatch.setattr(ClearFreePlanTenantExpiredLogs, "process_tenant", process_tenant) statements: list[str] = [] diff --git a/api/tests/unit_tests/services/test_conversation_service.py b/api/tests/unit_tests/services/test_conversation_service.py index 90727705061..641731c4384 100644 --- a/api/tests/unit_tests/services/test_conversation_service.py +++ b/api/tests/unit_tests/services/test_conversation_service.py @@ -8,10 +8,10 @@ in-memory SQLite sessions with persisted ORM rows. """ import json -from unittest.mock import patch +from unittest.mock import MagicMock, Mock, patch import pytest -from sqlalchemy import asc, desc +from sqlalchemy import asc, desc, event from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom @@ -19,6 +19,8 @@ from libs.datetime_utils import naive_utc_now from models import Account, ConversationVariable from models.enums import AppStatus, ConversationFromSource, ConversationStatus from models.model import App, AppMode, Conversation +from services import conversation_service +from services.agent.workspace_service import AgentWorkspaceService from services.conversation_service import ConversationService TENANT_ID = "11111111-1111-1111-1111-111111111111" @@ -148,7 +150,62 @@ class ConversationServiceTestDataFactory: return conversation -@pytest.mark.parametrize("sqlite_session", [(Conversation,)], indirect=True) +def test_delete_retires_then_commits_before_enqueue(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: + app = ConversationServiceTestDataFactory.create_app() + conversation = ConversationServiceTestDataFactory.create_conversation() + conversation.agent_workspace_binding_id = "conversation-binding-1" + sqlite_session.add(conversation) + sqlite_session.flush() + events: list[str] = [] + get_binding = MagicMock(return_value=Mock(id="conversation-binding-1")) + retire_binding = MagicMock(side_effect=lambda **_kwargs: events.append("retire") or "conversation-binding-1") + monkeypatch.setattr(ConversationService, "get_conversation", MagicMock(return_value=conversation)) + monkeypatch.setattr(AgentWorkspaceService, "get_active_binding", get_binding) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", retire_binding) + event.listen(sqlite_session, "after_commit", lambda _session: events.append("commit")) + monkeypatch.setattr( + conversation_service, + "enqueue_agent_resource_collection", + MagicMock(side_effect=lambda **_kwargs: events.append("enqueue")), + ) + monkeypatch.setattr(conversation_service.delete_conversation_related_data, "delay", MagicMock()) + + ConversationService.delete(app, conversation.id, None, session=sqlite_session) + + assert events == ["retire", "commit", "enqueue"] + assert get_binding.call_args.kwargs["binding_id"] == "conversation-binding-1" + assert retire_binding.call_args.kwargs["binding_id"] == "conversation-binding-1" + + +def test_delete_commit_failure_does_not_enqueue(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: + app = ConversationServiceTestDataFactory.create_app() + conversation = ConversationServiceTestDataFactory.create_conversation() + conversation.agent_workspace_binding_id = "binding-1" + sqlite_session.add(conversation) + sqlite_session.flush() + rollback_events: list[str] = [] + event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback")) + monkeypatch.setattr(sqlite_session, "commit", MagicMock(side_effect=RuntimeError("commit failed"))) + monkeypatch.setattr(ConversationService, "get_conversation", MagicMock(return_value=conversation)) + monkeypatch.setattr( + AgentWorkspaceService, + "get_active_binding", + MagicMock(return_value=Mock(id="binding-1")), + ) + monkeypatch.setattr(AgentWorkspaceService, "retire_binding", MagicMock(return_value="binding-1")) + enqueue_collection = MagicMock() + delete_related = MagicMock() + monkeypatch.setattr(conversation_service, "enqueue_agent_resource_collection", enqueue_collection) + monkeypatch.setattr(conversation_service.delete_conversation_related_data, "delay", delete_related) + + with pytest.raises(RuntimeError, match="commit failed"): + ConversationService.delete(app, conversation.id, None, session=sqlite_session) + + assert rollback_events == ["rollback"] + enqueue_collection.assert_not_called() + delete_related.assert_not_called() + + class TestConversationServicePagination: """Test conversation pagination operations.""" diff --git a/api/tests/unit_tests/services/test_credential_permission_service.py b/api/tests/unit_tests/services/test_credential_permission_service.py index d4e8596b14c..cc416969bc9 100644 --- a/api/tests/unit_tests/services/test_credential_permission_service.py +++ b/api/tests/unit_tests/services/test_credential_permission_service.py @@ -5,42 +5,66 @@ and admin bypass behavior. """ from types import SimpleNamespace -from unittest.mock import MagicMock +from typing import cast from uuid import uuid4 import pytest from sqlalchemy import select from sqlalchemy.orm import Session +from core.plugin.entities.plugin_daemon import CredentialType as TriggerCredentialType +from models.account import Account from models.credential_permission import CredentialPermission, CredentialType +from models.enums import PermissionEnum +from models.trigger import TriggerSubscription from services.credential_permission_service import CredentialPermissionService -pytestmark = [ - pytest.mark.usefixtures("sqlite_session"), - pytest.mark.parametrize("sqlite_session", [(CredentialPermission,)], indirect=True), -] - @pytest.fixture -def tenant_id(): +def tenant_id() -> str: return str(uuid4()) @pytest.fixture -def user_id(): +def user_id() -> str: return str(uuid4()) @pytest.fixture -def other_user_id(): +def other_user_id() -> str: return str(uuid4()) @pytest.fixture -def credential_id(): +def credential_id() -> str: return str(uuid4()) +def _subscription( + *, + tenant_id: str, + owner_id: str, + name: str, + visibility: PermissionEnum, +) -> TriggerSubscription: + return TriggerSubscription( + name=name, + tenant_id=tenant_id, + user_id=owner_id, + provider_id="test/provider", + endpoint_id=f"{name}-endpoint", + parameters={}, + properties={}, + credentials={}, + credential_type=TriggerCredentialType.API_KEY, + visibility=visibility, + ) + + +def _user(user_id: str, *, is_admin: bool) -> Account: + return cast(Account, SimpleNamespace(id=user_id, is_admin_or_owner=is_admin)) + + class TestGetPartialMemberList: def test_returns_empty_when_no_permissions( self, sqlite_session: Session, credential_id: str, tenant_id: str, user_id: str @@ -89,50 +113,88 @@ class TestGetPartialMemberList: class TestApplyVisibilityFilter: - """Test the visibility filter logic using mock model columns.""" + def test_admin_does_not_bypass_personal_visibility( + self, + sqlite_session: Session, + tenant_id: str, + user_id: str, + other_user_id: str, + ) -> None: + private_subscription = _subscription( + tenant_id=tenant_id, + owner_id=other_user_id, + name="private", + visibility=PermissionEnum.ONLY_ME, + ) + sqlite_session.add(private_subscription) + sqlite_session.commit() - def _make_mock_columns(self): - """Create mock model columns for testing.""" - model_id = MagicMock(name="id_column") - model_user_id = MagicMock(name="user_id_column") - model_visibility = MagicMock(name="visibility_column") - return model_id, model_user_id, model_visibility - - def _make_user(self, user_id: str, is_admin: bool): - return SimpleNamespace(id=user_id, is_admin_or_owner=is_admin) - - def test_admin_gets_filtered_too(self, user_id): - """Admin should NOT bypass visibility — personal credentials are private regardless of role.""" - from models.trigger import TriggerSubscription - - query = select(TriggerSubscription) - result = CredentialPermissionService.apply_visibility_filter( - query, + query = CredentialPermissionService.apply_visibility_filter( + select(TriggerSubscription).where(TriggerSubscription.tenant_id == tenant_id), model_id_column=TriggerSubscription.id, model_user_id_column=TriggerSubscription.user_id, model_visibility_column=TriggerSubscription.visibility, credential_type=CredentialType.TRIGGER_SUBSCRIPTION, - user=self._make_user(user_id, is_admin=True), + user=_user(user_id, is_admin=True), ) - # No admin bypass: query should have WHERE clause - compiled = str(result.compile(compile_kwargs={"literal_binds": True})) - assert "WHERE" in compiled - def test_non_admin_adds_filter_on_real_model(self, user_id): - """Non-admin should get a filtered query when using real SQLAlchemy columns.""" - from models.trigger import TriggerSubscription + assert sqlite_session.scalars(query).all() == [] - query = select(TriggerSubscription) - result = CredentialPermissionService.apply_visibility_filter( - query, + def test_non_admin_sees_team_owned_and_partial_member_subscriptions( + self, + sqlite_session: Session, + tenant_id: str, + user_id: str, + other_user_id: str, + ) -> None: + team_subscription = _subscription( + tenant_id=tenant_id, + owner_id=other_user_id, + name="team", + visibility=PermissionEnum.ALL_TEAM, + ) + owned_subscription = _subscription( + tenant_id=tenant_id, + owner_id=user_id, + name="owned", + visibility=PermissionEnum.ONLY_ME, + ) + shared_subscription = _subscription( + tenant_id=tenant_id, + owner_id=other_user_id, + name="shared", + visibility=PermissionEnum.PARTIAL_TEAM, + ) + private_subscription = _subscription( + tenant_id=tenant_id, + owner_id=other_user_id, + name="private", + visibility=PermissionEnum.ONLY_ME, + ) + sqlite_session.add_all( + [ + team_subscription, + owned_subscription, + shared_subscription, + private_subscription, + CredentialPermission( + credential_id=shared_subscription.id, + credential_type=CredentialType.TRIGGER_SUBSCRIPTION, + account_id=user_id, + tenant_id=tenant_id, + ), + ] + ) + sqlite_session.commit() + + query = CredentialPermissionService.apply_visibility_filter( + select(TriggerSubscription).where(TriggerSubscription.tenant_id == tenant_id), model_id_column=TriggerSubscription.id, model_user_id_column=TriggerSubscription.user_id, model_visibility_column=TriggerSubscription.visibility, credential_type=CredentialType.TRIGGER_SUBSCRIPTION, - user=self._make_user(user_id, is_admin=False), + user=_user(user_id, is_admin=False), ) - # The compiled SQL should include a WHERE clause referencing user_id and visibility - compiled = str(result.compile(compile_kwargs={"literal_binds": True})) - assert "WHERE" in compiled - assert "visibility" in compiled - assert "user_id" in compiled + + visible_ids = {subscription.id for subscription in sqlite_session.scalars(query)} + assert visible_ids == {team_subscription.id, owned_subscription.id, shared_subscription.id} diff --git a/api/tests/unit_tests/services/test_credit_pool_service.py b/api/tests/unit_tests/services/test_credit_pool_service.py index caf5d6ade02..304411c8f94 100644 --- a/api/tests/unit_tests/services/test_credit_pool_service.py +++ b/api/tests/unit_tests/services/test_credit_pool_service.py @@ -1,14 +1,16 @@ +"""Credit-pool accounting tests backed by real SQLite sessions.""" + from collections.abc import Generator from types import SimpleNamespace from unittest.mock import ANY, MagicMock, patch from uuid import uuid4 import pytest -from sqlalchemy import create_engine, select -from sqlalchemy.engine import Engine -from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy import select +from sqlalchemy.orm import Session from core.errors.error import QuotaExceededError +from enums import DeploymentEdition from models import TenantCreditPool from models.enums import ProviderQuotaType from services.credit_pool_service import ( @@ -16,36 +18,25 @@ from services.credit_pool_service import ( CREDIT_POOL_TENANT_LOCK_TIMEOUT_SECONDS, FEATURE_KEY_CREDIT_POOL, CreditPoolBalance, + CreditPoolReservationState, CreditPoolService, ) -def _create_engine_with_pool(*, quota_limit: int, quota_used: int) -> tuple[Engine, str, str]: - engine = create_engine("sqlite:///:memory:") - TenantCreditPool.__table__.create(engine) - tenant_id = str(uuid4()) - pool_id = str(uuid4()) - with engine.begin() as connection: - connection.execute( - TenantCreditPool.__table__.insert(), - { - "id": pool_id, - "tenant_id": tenant_id, - "pool_type": ProviderQuotaType.TRIAL, - "quota_limit": quota_limit, - "quota_used": quota_used, - }, - ) - return engine, tenant_id, pool_id +def _create_pool(session: Session, *, quota_limit: int, quota_used: int) -> TenantCreditPool: + pool = TenantCreditPool( + tenant_id=str(uuid4()), + pool_type=ProviderQuotaType.TRIAL, + quota_limit=quota_limit, + quota_used=quota_used, + ) + session.add(pool) + session.commit() + return pool -def _make_session(engine: Engine) -> Session: - return sessionmaker(bind=engine, expire_on_commit=False)() - - -def _get_quota_used(*, engine: Engine, pool_id: str) -> int | None: - with engine.connect() as connection: - return connection.scalar(select(TenantCreditPool.quota_used).where(TenantCreditPool.id == pool_id)) +def _get_quota_used(*, session: Session, pool_id: str) -> int | None: + return session.scalar(select(TenantCreditPool.quota_used).where(TenantCreditPool.id == pool_id)) def _make_redis_lock() -> MagicMock: @@ -56,18 +47,20 @@ def _make_redis_lock() -> MagicMock: @pytest.fixture(autouse=True) def _disable_billing_quota_by_default() -> Generator[None, None, None]: - with patch("services.credit_pool_service.dify_config.BILLING_ENABLED", False): + with patch("services.credit_pool_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY): yield -def test_get_pool_uses_provided_session() -> None: - engine, tenant_id, _ = _create_engine_with_pool(quota_limit=10, quota_used=2) - - with _make_session(engine) as session: - pool = CreditPoolService.get_pool(tenant_id=tenant_id, pool_type=ProviderQuotaType.TRIAL, session=session) +def test_get_pool_uses_provided_session(sqlite_session: Session) -> None: + persisted_pool = _create_pool(sqlite_session, quota_limit=10, quota_used=2) + pool = CreditPoolService.get_pool( + tenant_id=persisted_pool.tenant_id, + pool_type=ProviderQuotaType.TRIAL, + session=sqlite_session, + ) assert pool is not None - assert pool.tenant_id == tenant_id + assert pool.tenant_id == persisted_pool.tenant_id assert pool.quota_used == 2 @@ -78,132 +71,127 @@ def test_credit_pool_balance_unlimited_remaining_and_sufficiency() -> None: assert pool.has_sufficient_credits(10_000) -def test_check_and_deduct_credits_deducts_exact_amount_when_sufficient() -> None: - engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) +def test_check_and_deduct_credits_deducts_exact_amount_when_sufficient(sqlite_session: Session) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=2) - with _make_session(engine) as session: - deducted_credits = CreditPoolService.check_and_deduct_credits( - tenant_id=tenant_id, credits_required=3, session=session - ) + deducted_credits = CreditPoolService.check_and_deduct_credits( + tenant_id=pool.tenant_id, credits_required=3, session=sqlite_session + ) assert deducted_credits == 3 - assert _get_quota_used(engine=engine, pool_id=pool_id) == 5 + assert sqlite_session.in_transaction() is False + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 5 -def test_check_and_deduct_credits_returns_zero_for_non_positive_request() -> None: +def test_check_and_deduct_credits_returns_zero_for_non_positive_request(sqlite_session: Session) -> None: assert ( - CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=0, session=MagicMock()) == 0 + CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=0, session=sqlite_session) + == 0 ) -def test_check_and_deduct_credits_raises_when_pool_is_missing() -> None: - engine = create_engine("sqlite:///:memory:") - TenantCreditPool.__table__.create(engine) - - with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="Credit pool not found"): - CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=1, session=session) +def test_check_and_deduct_credits_raises_when_pool_is_missing(sqlite_session: Session) -> None: + with pytest.raises(QuotaExceededError, match="Credit pool not found"): + CreditPoolService.check_and_deduct_credits(tenant_id=str(uuid4()), credits_required=1, session=sqlite_session) -def test_check_and_deduct_credits_raises_when_pool_is_empty() -> None: - engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=10) +def test_check_and_deduct_credits_raises_when_pool_is_empty(sqlite_session: Session) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=10) - with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="No credits remaining"): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1, session=session) + with pytest.raises(QuotaExceededError, match="No credits remaining"): + CreditPoolService.check_and_deduct_credits(tenant_id=pool.tenant_id, credits_required=1, session=sqlite_session) - assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 10 -def test_check_and_deduct_credits_raises_without_partial_deduction_when_insufficient() -> None: - engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=9) +def test_check_and_deduct_credits_raises_without_partial_deduction_when_insufficient( + sqlite_session: Session, +) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=9) - with _make_session(engine) as session, pytest.raises(QuotaExceededError, match="Insufficient credits remaining"): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=3, session=session) + with pytest.raises(QuotaExceededError, match="Insufficient credits remaining"): + CreditPoolService.check_and_deduct_credits(tenant_id=pool.tenant_id, credits_required=3, session=sqlite_session) - assert _get_quota_used(engine=engine, pool_id=pool_id) == 9 + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 9 -def test_check_and_deduct_credits_wraps_unexpected_deduction_errors() -> None: - engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) +def test_check_and_deduct_credits_wraps_unexpected_deduction_errors(sqlite_session: Session) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=2) with ( - _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.check_and_deduct_credits(tenant_id=tenant_id, credits_required=1, session=session) + CreditPoolService.check_and_deduct_credits(tenant_id=pool.tenant_id, credits_required=1, session=sqlite_session) - assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 2 -def test_deduct_credits_capped_returns_zero_for_non_positive_request() -> None: - assert CreditPoolService.deduct_credits_capped(tenant_id=str(uuid4()), credits_required=0, session=MagicMock()) == 0 +def test_deduct_credits_capped_returns_zero_for_non_positive_request(sqlite_session: Session) -> None: + assert ( + CreditPoolService.deduct_credits_capped(tenant_id=str(uuid4()), credits_required=0, session=sqlite_session) == 0 + ) -def test_deduct_credits_capped_returns_zero_when_pool_is_missing() -> None: - engine = create_engine("sqlite:///:memory:") - TenantCreditPool.__table__.create(engine) - - with _make_session(engine) as session: - deducted_credits = CreditPoolService.deduct_credits_capped( - tenant_id=str(uuid4()), credits_required=1, session=session - ) +def test_deduct_credits_capped_returns_zero_when_pool_is_missing(sqlite_session: Session) -> None: + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=str(uuid4()), credits_required=1, session=sqlite_session + ) assert deducted_credits == 0 -def test_deduct_credits_capped_returns_zero_when_pool_is_empty() -> None: - engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=10) +def test_deduct_credits_capped_returns_zero_when_pool_is_empty(sqlite_session: Session) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=10) - with _make_session(engine) as session: - deducted_credits = CreditPoolService.deduct_credits_capped( - tenant_id=tenant_id, credits_required=1, session=session - ) + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=pool.tenant_id, credits_required=1, session=sqlite_session + ) assert deducted_credits == 0 - assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 10 -def test_deduct_credits_capped_deducts_only_remaining_balance_when_insufficient() -> None: - engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=9) +def test_deduct_credits_capped_deducts_only_remaining_balance_when_insufficient( + sqlite_session: Session, +) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=9) - with _make_session(engine) as session: - deducted_credits = CreditPoolService.deduct_credits_capped( - tenant_id=tenant_id, credits_required=3, session=session - ) + deducted_credits = CreditPoolService.deduct_credits_capped( + tenant_id=pool.tenant_id, credits_required=3, session=sqlite_session + ) assert deducted_credits == 1 - assert _get_quota_used(engine=engine, pool_id=pool_id) == 10 + assert sqlite_session.in_transaction() is False + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 10 -def test_deduct_credits_capped_wraps_unexpected_deduction_errors() -> None: - engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) +def test_deduct_credits_capped_wraps_unexpected_deduction_errors(sqlite_session: Session) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=2) with ( - _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=RuntimeError("database unavailable")), pytest.raises(QuotaExceededError, match="Failed to deduct credits"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1, session=session) + CreditPoolService.deduct_credits_capped(tenant_id=pool.tenant_id, credits_required=1, session=sqlite_session) - assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 2 -def test_deduct_credits_capped_reraises_quota_exceeded_errors() -> None: - engine, tenant_id, pool_id = _create_engine_with_pool(quota_limit=10, quota_used=2) +def test_deduct_credits_capped_reraises_quota_exceeded_errors(sqlite_session: Session) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=2) with ( - _make_session(engine) as session, patch.object(CreditPoolService, "_get_locked_pool", side_effect=QuotaExceededError("quota unavailable")), pytest.raises(QuotaExceededError, match="quota unavailable"), ): - CreditPoolService.deduct_credits_capped(tenant_id=tenant_id, credits_required=1, session=session) + CreditPoolService.deduct_credits_capped(tenant_id=pool.tenant_id, credits_required=1, session=sqlite_session) - assert _get_quota_used(engine=engine, pool_id=pool_id) == 2 + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 2 -def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction() -> None: +def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction(sqlite_session: Session) -> None: tenant_id = "tenant-1" - session = MagicMock() pool = SimpleNamespace(remaining_credits=10, quota_used=2) redis_lock = _make_redis_lock() @@ -215,7 +203,7 @@ def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction() - tenant_id=tenant_id, credits_required=3, pool_type=ProviderQuotaType.TRIAL, - session=session, + session=sqlite_session, ) assert result == 3 @@ -227,12 +215,11 @@ def test_check_and_deduct_credits_uses_tenant_redis_lock_before_db_deduction() - ) redis_lock.acquire.assert_called_once_with(blocking=True) redis_lock.release.assert_called_once_with() - get_locked_pool.assert_called_once_with(session=session, tenant_id=tenant_id, pool_type="trial") + get_locked_pool.assert_called_once_with(session=sqlite_session, tenant_id=tenant_id, pool_type="trial") -def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> None: +def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction(sqlite_session: Session) -> None: tenant_id = "tenant-1" - session = MagicMock() pool = SimpleNamespace(remaining_credits=2, quota_used=8) redis_lock = _make_redis_lock() @@ -244,7 +231,7 @@ def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> N tenant_id=tenant_id, credits_required=5, pool_type=ProviderQuotaType.PAID, - session=session, + session=sqlite_session, ) assert result == 2 @@ -256,13 +243,13 @@ def test_deduct_credits_capped_uses_tenant_redis_lock_before_db_deduction() -> N ) redis_lock.acquire.assert_called_once_with(blocking=True) redis_lock.release.assert_called_once_with() - get_locked_pool.assert_called_once_with(session=session, tenant_id=tenant_id, pool_type="paid") + get_locked_pool.assert_called_once_with(session=sqlite_session, tenant_id=tenant_id, pool_type="paid") def test_get_pool_uses_billing_quota_balance_when_enabled() -> None: tenant_id = "tenant-1" with ( - patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.credit_pool_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("services.billing_service.BillingService.quota_get_balance") as quota_get_balance, ): quota_get_balance.return_value = { @@ -287,10 +274,94 @@ def test_get_pool_uses_billing_quota_balance_when_enabled() -> None: ) +def test_reserve_credits_commits_billing_reservation_once() -> None: + with ( + patch.object(CreditPoolService, "_use_billing_quota", return_value=True), + patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, + patch("services.billing_service.BillingService.quota_commit") as quota_commit, + patch("services.billing_service.BillingService.quota_release") as quota_release, + ): + quota_reserve.return_value = {"reservation_id": "reservation-1", "available": 7, "reserved": 3} + + reservation = CreditPoolService.reserve_credits( + tenant_id="tenant-1", + credits_required=3, + pool_type=ProviderQuotaType.TRIAL, + request_id="request-1", + meta={"source": "test"}, + ) + reservation.commit() + reservation.commit() + reservation.release() + + assert reservation.state == CreditPoolReservationState.COMMITTED + quota_reserve.assert_called_once_with( + tenant_id="tenant-1", + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + request_id="request-1", + amount=3, + meta={"source": "test"}, + ) + quota_commit.assert_called_once_with( + tenant_id="tenant-1", + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + reservation_id="reservation-1", + actual_amount=3, + meta={"source": "test", "request_id": "request-1"}, + ) + quota_release.assert_not_called() + + +def test_reserve_credits_releases_billing_reservation() -> None: + with ( + patch.object(CreditPoolService, "_use_billing_quota", return_value=True), + patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, + patch("services.billing_service.BillingService.quota_release") as quota_release, + ): + quota_reserve.return_value = {"reservation_id": "reservation-1", "available": 7, "reserved": 3} + + reservation = CreditPoolService.reserve_credits( + tenant_id="tenant-1", + credits_required=3, + request_id="request-1", + ) + reservation.release() + reservation.release() + + assert reservation.state == CreditPoolReservationState.RELEASED + quota_release.assert_called_once_with( + tenant_id="tenant-1", + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + reservation_id="reservation-1", + ) + + +def test_reserve_credits_database_fallback_restores_released_amount(sqlite_session: Session) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=2) + redis_lock = _make_redis_lock() + + with patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock): + reservation = CreditPoolService.reserve_credits( + tenant_id=pool.tenant_id, + credits_required=3, + request_id="request-1", + session_factory=lambda: sqlite_session, + ) + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 5 + + reservation.release() + + assert reservation.state == CreditPoolReservationState.RELEASED + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 2 + + def test_check_and_deduct_credits_uses_billing_reserve_and_commit_when_enabled() -> None: tenant_id = "tenant-1" with ( - patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.credit_pool_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, patch("services.billing_service.BillingService.quota_commit") as quota_commit, patch("services.billing_service.BillingService.quota_release") as quota_release, @@ -325,7 +396,7 @@ def test_check_and_deduct_credits_uses_billing_reserve_and_commit_when_enabled() def test_check_and_deduct_credits_raises_when_billing_reserve_is_insufficient() -> None: with ( - patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.credit_pool_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, ): quota_reserve.return_value = {"reservation_id": "", "available": 1, "reserved": 0} @@ -336,7 +407,7 @@ def test_check_and_deduct_credits_raises_when_billing_reserve_is_insufficient() def test_check_and_deduct_credits_releases_billing_reservation_when_commit_fails() -> None: with ( - patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.credit_pool_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, patch("services.billing_service.BillingService.quota_commit", side_effect=RuntimeError("commit failed")), patch("services.billing_service.BillingService.quota_release") as quota_release, @@ -358,7 +429,7 @@ def test_check_and_deduct_credits_logs_when_billing_release_fails( caplog: pytest.LogCaptureFixture, ) -> None: with ( - patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.credit_pool_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, patch("services.billing_service.BillingService.quota_commit", side_effect=RuntimeError("commit failed")), patch( @@ -384,7 +455,7 @@ def test_check_and_deduct_credits_logs_when_billing_release_fails( def test_deduct_credits_capped_uses_billing_consume_capped_when_enabled() -> None: tenant_id = "tenant-1" with ( - patch("services.credit_pool_service.dify_config.BILLING_ENABLED", True), + patch("services.credit_pool_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("services.billing_service.BillingService.quota_consume_capped") as quota_consume_capped, ): quota_consume_capped.return_value = { @@ -419,28 +490,28 @@ def test_deduct_credits_capped_uses_billing_consume_capped_when_enabled() -> Non CreditPoolService.deduct_credits_capped, ], ) -def test_non_positive_credit_request_skips_tenant_redis_lock(deduct_method) -> None: +def test_non_positive_credit_request_skips_tenant_redis_lock( + deduct_method, + sqlite_session: Session, +) -> None: with patch("services.credit_pool_service.redis_client.lock") as lock: - result = deduct_method(tenant_id="tenant-1", credits_required=0, session=MagicMock()) + result = deduct_method(tenant_id="tenant-1", credits_required=0, session=sqlite_session) assert result == 0 lock.assert_not_called() -def test_check_and_deduct_credits_wraps_redis_lock_errors_without_querying_db() -> None: - session = MagicMock() +def test_check_and_deduct_credits_wraps_redis_lock_errors_without_querying_db(sqlite_session: Session) -> None: + with patch("services.credit_pool_service.redis_client.lock", side_effect=RuntimeError("redis unavailable")): + with pytest.raises(QuotaExceededError, match="Failed to deduct credits"): + CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=1, session=sqlite_session) - with ( - patch("services.credit_pool_service.redis_client.lock", side_effect=RuntimeError("redis unavailable")), - pytest.raises(QuotaExceededError, match="Failed to deduct credits"), - ): - CreditPoolService.check_and_deduct_credits(tenant_id="tenant-1", credits_required=1, session=session) - - session.scalar.assert_not_called() + assert sqlite_session.in_transaction() is False -def test_deduct_credits_capped_ignores_release_errors_after_successful_deduction() -> None: - session = MagicMock() +def test_deduct_credits_capped_ignores_release_errors_after_successful_deduction( + sqlite_session: Session, +) -> None: pool = SimpleNamespace(remaining_credits=3, quota_used=7) redis_lock = _make_redis_lock() redis_lock.release.side_effect = RuntimeError("release failed") @@ -449,7 +520,9 @@ def test_deduct_credits_capped_ignores_release_errors_after_successful_deduction patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock), patch.object(CreditPoolService, "_get_locked_pool", return_value=pool), ): - result = CreditPoolService.deduct_credits_capped(tenant_id="tenant-1", credits_required=2, session=session) + result = CreditPoolService.deduct_credits_capped( + tenant_id="tenant-1", credits_required=2, session=sqlite_session + ) assert result == 2 assert pool.quota_used == 9 diff --git a/api/tests/unit_tests/services/test_dataset_service_dataset.py b/api/tests/unit_tests/services/test_dataset_service_dataset.py index 1fcc81039ca..5487da2aa22 100644 --- a/api/tests/unit_tests/services/test_dataset_service_dataset.py +++ b/api/tests/unit_tests/services/test_dataset_service_dataset.py @@ -1,5 +1,9 @@ """Unit tests for DatasetService and dataset-related collaborators.""" +from sqlalchemy.orm import Session + +from models.dataset import DatasetPermission + from .dataset_service_test_helpers import ( DatasetNameDuplicateError, DatasetPermissionEnum, @@ -34,11 +38,10 @@ class TestDatasetServiceValidation: def test_check_doc_form_allows_matching_or_missing_dataset_doc_form(self, dataset_doc_form, incoming_doc_form): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(doc_form=dataset_doc_form) session = MagicMock() + session.scalar.return_value = None DatasetService.check_doc_form(dataset, incoming_doc_form, session=session) - dataset.get_doc_form.assert_called_once_with(session=session) - def test_check_doc_form_rejects_mismatched_doc_form(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(doc_form="qa_model") session = MagicMock() @@ -46,7 +49,28 @@ class TestDatasetServiceValidation: with pytest.raises(ValueError, match="doc_form is different"): DatasetService.check_doc_form(dataset, "text_model", session=session) - dataset.get_doc_form.assert_called_once_with(session=session) + @pytest.mark.parametrize("operator_check", [False, True]) + def test_dataset_permission_checks_ignore_foreign_tenant_binding( + self, sqlite_session: Session, operator_check: bool + ) -> None: + dataset = DatasetServiceUnitDataFactory.create_dataset_mock( + dataset_id="dataset-1", + tenant_id="tenant-1", + permission=DatasetPermissionEnum.PARTIAL_TEAM, + maintainer="owner-1", + ) + user = DatasetServiceUnitDataFactory.create_user_mock( + user_id="user-1", + tenant_id="tenant-1", + role=TenantAccountRole.NORMAL, + ) + sqlite_session.add(DatasetPermission(dataset_id=dataset.id, account_id=user.id, tenant_id="tenant-2")) + + with pytest.raises(NoPermissionError): + if operator_check: + DatasetService.check_dataset_operator_permission(user, dataset, session=sqlite_session) + else: + DatasetService.check_dataset_permission(dataset, user, sqlite_session) def test_check_dataset_model_setting_skips_non_high_quality_datasets(self): dataset = DatasetServiceUnitDataFactory.create_dataset_mock(indexing_technique="economy") diff --git a/api/tests/unit_tests/services/test_dataset_service_document.py b/api/tests/unit_tests/services/test_dataset_service_document.py index 26a3ac08d5f..7093053d1cb 100644 --- a/api/tests/unit_tests/services/test_dataset_service_document.py +++ b/api/tests/unit_tests/services/test_dataset_service_document.py @@ -44,6 +44,44 @@ from .dataset_service_test_helpers import ( ) +class _RetryFlagLock: + def __init__(self, store: "_RetryFlagStore", key: str): + self.store = store + self.key = key + self.token = f"owner-{store.next_token}" + store.next_token += 1 + + def acquire(self, *, blocking: bool): + assert blocking is False + if self.key in self.store.values: + if self.store.replacement_on_conflict: + replacement_key, replacement_value = self.store.replacement_on_conflict + self.store.values[replacement_key] = replacement_value + return False + self.store.values[self.key] = self.token + return True + + def release(self): + if self.store.values.get(self.key) == self.token: + self.store.values.pop(self.key) + + +class _RetryFlagStore: + def __init__( + self, + values: dict[str, str] | None = None, + replacement_on_conflict: tuple[str, str] | None = None, + ): + self.values = values or {} + self.replacement_on_conflict = replacement_on_conflict + self.next_token = 1 + + def lock(self, key: str, *, timeout: int, thread_local: bool): + assert timeout == 600 + assert thread_local is False + return _RetryFlagLock(self, key) + + class TestDocumentServiceDisplayStatus: """Unit tests for DocumentService display-status helpers.""" @@ -120,10 +158,12 @@ class TestDocumentServiceMutations: def test_delete_documents_limits_query_and_cleanup_to_dataset_ref(self): session = MagicMock() - dataset = _make_dataset(dataset_id="dataset-1", tenant_id="tenant-1") - dataset.doc_form = "paragraph_index" + dataset = _make_dataset( + dataset_id="dataset-1", + tenant_id="tenant-1", + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) document = _make_document(document_id="doc-1", dataset_id=dataset.id, tenant_id=dataset.tenant_id) - document.data_source_info_dict = {} with ( patch("services.dataset_service.batch_clean_document_task") as clean_task, @@ -214,33 +254,114 @@ class TestDocumentServiceMutations: with pytest.raises(DocumentIndexingError): DocumentService.recover_document(document, session) - def test_retry_document_raises_when_retry_flag_is_already_set(self): + def test_retry_document_raises_when_retry_flag_is_already_set(self, rename_account_context): document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1") session = MagicMock() with patch("services.dataset_service.redis_client") as mock_redis: - mock_redis.get.return_value = "1" + mock_redis.lock.return_value.acquire.return_value = False with pytest.raises(ValueError, match="being retried"): DocumentService.retry_document("dataset-1", [document], session) + def test_retry_document_leaves_batch_unchanged_when_later_document_is_already_being_retried( + self, rename_account_context + ): + first_document = DatasetServiceUnitDataFactory.create_document_mock( + document_id="doc-1", indexing_status="error" + ) + second_document = DatasetServiceUnitDataFactory.create_document_mock( + document_id="doc-2", indexing_status="error" + ) + first_retry_key = "document_doc-1_is_retried" + second_retry_key = "document_doc-2_is_retried" + retry_flags = _RetryFlagStore({second_retry_key: "other-request"}) + session = MagicMock() + + with ( + patch("services.dataset_service.redis_client", retry_flags), + patch("services.dataset_service.retry_document_indexing_task") as retry_task, + ): + with pytest.raises(ValueError, match="being retried"): + DocumentService.retry_document( + "dataset-1", + [first_document, second_document], + session, + ) + + assert first_document.indexing_status == "error" + assert second_document.indexing_status == "error" + assert first_retry_key not in retry_flags.values + assert retry_flags.values[second_retry_key] == "other-request" + retry_task.delay.assert_not_called() + + def test_retry_document_does_not_release_a_retry_flag_reacquired_by_another_request(self, rename_account_context): + first_retry_key = "document_doc-1_is_retried" + second_retry_key = "document_doc-2_is_retried" + retry_flags = _RetryFlagStore( + {second_retry_key: "other-request"}, + replacement_on_conflict=(first_retry_key, "new-owner"), + ) + documents = [ + DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1", indexing_status="error"), + DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-2", indexing_status="error"), + ] + + with patch("services.dataset_service.redis_client", retry_flags): + with pytest.raises(ValueError, match="being retried"): + DocumentService.retry_document("dataset-1", documents, MagicMock()) + + assert retry_flags.values[first_retry_key] == "new-owner" + assert retry_flags.values[second_retry_key] == "other-request" + + def test_retry_document_releases_flags_when_status_commit_fails(self, rename_account_context): + retry_flags = _RetryFlagStore() + document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1", indexing_status="error") + session = MagicMock() + session.commit.side_effect = RuntimeError("database unavailable") + + with ( + patch("services.dataset_service.redis_client", retry_flags), + patch("services.dataset_service.retry_document_indexing_task") as retry_task, + ): + with pytest.raises(RuntimeError, match="database unavailable"): + DocumentService.retry_document("dataset-1", [document], session) + + assert retry_flags.values == {} + session.rollback.assert_called_once_with() + retry_task.delay.assert_not_called() + def test_sync_website_document_raises_when_sync_flag_exists(self): - document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1") + dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") + document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1", dataset_id=dataset.id) session = MagicMock() with patch("services.dataset_service.redis_client") as mock_redis: mock_redis.get.return_value = "1" with pytest.raises(ValueError, match="being synced"): - DocumentService.sync_website_document("dataset-1", document, session) + DocumentService.sync_website_document(dataset, document, session) + + def test_sync_website_document_rejects_document_outside_dataset(self): + dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") + document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1", dataset_id="dataset-2") + + with ( + pytest.raises(ValueError, match="Document not found"), + patch("services.dataset_service.redis_client") as mock_redis, + ): + DocumentService.sync_website_document(dataset, document, MagicMock()) + + mock_redis.get.assert_not_called() def test_sync_website_document_updates_status_sets_cache_and_dispatches_task(self): session = MagicMock() + dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1") document = DatasetServiceUnitDataFactory.create_document_mock( document_id="doc-1", + dataset_id=dataset.id, data_source_info_dict={"mode": "crawl"}, ) - document.data_source_info = "{}" with ( patch("services.dataset_service.redis_client") as mock_redis, @@ -248,14 +369,14 @@ class TestDocumentServiceMutations: ): mock_redis.get.return_value = None - DocumentService.sync_website_document("dataset-1", document, session) + DocumentService.sync_website_document(dataset, document, session) assert document.indexing_status == "waiting" assert '"mode": "scrape"' in document.data_source_info session.add.assert_called_once_with(document) session.commit.assert_called_once() mock_redis.setex.assert_called_once_with("document_doc-1_is_sync", 600, 1) - sync_task.delay.assert_called_once_with("dataset-1", "doc-1") + sync_task.delay.assert_called_once_with(dataset.id, document.id) class TestDocumentServiceSaveDocumentWithoutDatasetId: @@ -663,13 +784,14 @@ class TestDocumentServiceCreateValidation: ], ) def test_data_source_args_validate_requires_source_specific_info(self, data_source_type, field_name, message): - info_list = SimpleNamespace( - data_source_type=data_source_type, - file_info_list=object(), - notion_info_list=object(), - website_info_list=object(), - ) - setattr(info_list, field_name, None) + info_values = { + "data_source_type": data_source_type, + "file_info_list": object(), + "notion_info_list": object(), + "website_info_list": object(), + } + info_values[field_name] = None + info_list = SimpleNamespace(**info_values) knowledge_config = SimpleNamespace(data_source=SimpleNamespace(info_list=info_list)) with pytest.raises(ValueError, match=message): @@ -865,7 +987,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: ) def test_save_document_with_dataset_id_requires_existing_process_rule_for_custom_mode(self, account_context): - dataset = _make_dataset(latest_process_rule=None) + dataset = _make_dataset() knowledge_config = _make_upload_knowledge_config( file_ids=["file-1"], process_rule=ProcessRule(mode="custom"), @@ -873,6 +995,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: with patch("services.dataset_service.FeatureService.get_features", return_value=_make_features(enabled=False)): session = MagicMock() + session.scalar.return_value = None with pytest.raises(ValueError, match="No process rule found"): DocumentService.save_document_with_dataset_id( dataset, @@ -881,7 +1004,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: session=session, ) - dataset.get_latest_process_rule.assert_called_once_with(session=session) + session.scalar.assert_called_once() def test_save_document_with_dataset_id_rejects_invalid_indexing_technique(self, account_context): dataset = _make_dataset(indexing_technique=None) @@ -1022,6 +1145,7 @@ class TestDocumentServiceSaveDocumentWithDatasetId: patch.object(DocumentService, "build_document", return_value=created_document) as build_document, patch("services.dataset_service.clean_notion_document_task") as clean_task, patch("services.dataset_service.DocumentIndexingTaskProxy") as document_proxy_cls, + patch("services.dataset_service.uuid.uuid4", return_value="doc-new"), ): mock_redis.lock.return_value = _make_lock_context() session.scalars.return_value.all.return_value = [existing_keep, existing_remove] @@ -1791,7 +1915,7 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: self, account_context ): session = MagicMock() - dataset = _make_dataset(latest_process_rule=None) + dataset = _make_dataset() knowledge_config = _make_upload_knowledge_config(file_ids=["file-1"], process_rule=None) created_process_rule = SimpleNamespace(id="rule-fallback") created_document = _make_document(document_id="doc-created", name="file.txt") @@ -1809,6 +1933,7 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: mock_redis.lock.return_value = _make_lock_context() process_rule_cls.AUTOMATIC_RULES = DatasetProcessRule.AUTOMATIC_RULES process_rule_cls.return_value = created_process_rule + session.scalar.return_value = None session.scalars.return_value.all.side_effect = [[SimpleNamespace(id="file-1", name="file.txt")], []] DocumentService.save_document_with_dataset_id( @@ -1818,7 +1943,7 @@ class TestDocumentServiceSaveDocumentAdditionalBranches: session=session, ) - dataset.get_latest_process_rule.assert_called_once_with(session=session) + session.scalar.assert_called_once() assert process_rule_cls.call_args.kwargs == { "dataset_id": "dataset-1", "mode": "automatic", diff --git a/api/tests/unit_tests/services/test_dataset_service_segment.py b/api/tests/unit_tests/services/test_dataset_service_segment.py index 3e229d9a787..2e2ad0aa774 100644 --- a/api/tests/unit_tests/services/test_dataset_service_segment.py +++ b/api/tests/unit_tests/services/test_dataset_service_segment.py @@ -858,7 +858,7 @@ class TestSegmentServiceMutations: # scalars() for child_node_ids session.scalars.return_value.all.return_value = ["child-1"] - SegmentService.delete_segments(["segment-1", "segment-2"], document, dataset, session) + SegmentService.delete_segments(["segment-1", "segment-2", "foreign-segment"], document, dataset, session) assert document.word_count == 0 session.add.assert_called_once_with(document) @@ -869,6 +869,13 @@ class TestSegmentServiceMutations: ["segment-1", "segment-2"], ["child-1"], ) + delete_stmt = session.execute.call_args_list[1].args[0] + delete_sql = str(delete_stmt.compile(compile_kwargs={"literal_binds": True})) + assert "document_segments.id IN ('segment-1', 'segment-2')" in delete_sql + assert "document_segments.dataset_id = 'dataset-1'" in delete_sql + assert "document_segments.document_id = 'doc-1'" in delete_sql + assert "document_segments.tenant_id = 'tenant-1'" in delete_sql + assert "foreign-segment" not in delete_sql session.commit.assert_called() def test_update_segments_status_enables_only_segments_without_indexing_cache(self): diff --git a/api/tests/unit_tests/services/test_document_indexing_task_proxy.py b/api/tests/unit_tests/services/test_document_indexing_task_proxy.py index 082bb7aa865..c012b132cd9 100644 --- a/api/tests/unit_tests/services/test_document_indexing_task_proxy.py +++ b/api/tests/unit_tests/services/test_document_indexing_task_proxy.py @@ -2,7 +2,7 @@ from unittest.mock import Mock, patch from core.entities.document_task import DocumentTask from core.rag.pipeline.queue import TenantIsolatedTaskQueue -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from services.document_indexing_proxy.document_indexing_task_proxy import DocumentIndexingTaskProxy diff --git a/api/tests/unit_tests/services/test_duplicate_document_indexing_task_proxy.py b/api/tests/unit_tests/services/test_duplicate_document_indexing_task_proxy.py index e0370edec9e..2af411b0d76 100644 --- a/api/tests/unit_tests/services/test_duplicate_document_indexing_task_proxy.py +++ b/api/tests/unit_tests/services/test_duplicate_document_indexing_task_proxy.py @@ -2,7 +2,7 @@ from unittest.mock import Mock, patch from core.entities.document_task import DocumentTask from core.rag.pipeline.queue import TenantIsolatedTaskQueue -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from services.document_indexing_proxy.duplicate_document_indexing_task_proxy import ( DuplicateDocumentIndexingTaskProxy, ) diff --git a/api/tests/unit_tests/services/test_external_dataset_service.py b/api/tests/unit_tests/services/test_external_dataset_service.py index c2b8b392573..689af7300e6 100644 --- a/api/tests/unit_tests/services/test_external_dataset_service.py +++ b/api/tests/unit_tests/services/test_external_dataset_service.py @@ -1,23 +1,26 @@ -""" -Comprehensive unit tests for ExternalDatasetService. +"""Tests for external knowledge API, binding, dataset, and retrieval operations. -This test suite provides extensive coverage of external knowledge API and dataset operations. -Target: 1500+ lines of comprehensive test coverage. +Database-facing cases persist the minimum required ``TypeBase`` tables through +the shared in-memory SQLite fixture, including cross-tenant rows. HTTP calls, +pagination, and clocks remain mocked at their genuine I/O boundaries. """ import json import re from datetime import datetime from typing import Any -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import MagicMock, patch import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session from constants import HIDDEN_VALUE from models.dataset import Dataset, ExternalKnowledgeApis, ExternalKnowledgeBindings from services.entities.external_knowledge_entities.external_knowledge_entities import ( Authorization, AuthorizationConfig, + ExternalDatasetCreatePayload, ExternalKnowledgeApiSetting, ) from services.errors.dataset import DatasetNameDuplicateError @@ -26,7 +29,7 @@ from services.external_knowledge_service import ExternalDatasetService class ExternalDatasetServiceTestDataFactory: - """Factory for creating test data and mock objects.""" + """Build non-session value objects used by tests outside persistence paths.""" @staticmethod def create_external_knowledge_api_mock( @@ -34,28 +37,29 @@ class ExternalDatasetServiceTestDataFactory: tenant_id: str = "tenant-123", name: str = "Test API", settings: dict[str, Any] | None = None, - **kwargs, - ) -> Mock: - """Create a mock ExternalKnowledgeApis object.""" - api = Mock(spec=ExternalKnowledgeApis) + description: str = "Test description", + created_by: str = "user-123", + updated_by: str = "user-123", + created_at: datetime = datetime(2024, 1, 1, 12, 0), + updated_at: datetime = datetime(2024, 1, 1, 12, 0), + ) -> ExternalKnowledgeApis: + """Create an ExternalKnowledgeApis object.""" + api = ExternalKnowledgeApis( + name=name, + description=description, + tenant_id=tenant_id, + settings="{}", + created_by=created_by, + updated_by=updated_by, + ) api.id = api_id - api.tenant_id = tenant_id - api.name = name - api.description = kwargs.get("description", "Test description") if settings is None: settings = {"endpoint": "https://api.example.com", "api_key": "test-key-123"} api.settings = json.dumps(settings, ensure_ascii=False) - api.settings_dict = settings - api.created_by = kwargs.get("created_by", "user-123") - api.updated_by = kwargs.get("updated_by", "user-123") - api.created_at = kwargs.get("created_at", datetime(2024, 1, 1, 12, 0)) - api.updated_at = kwargs.get("updated_at", datetime(2024, 1, 1, 12, 0)) - - for key, value in kwargs.items(): - if key not in ["description", "created_by", "updated_by", "created_at", "updated_at"]: - setattr(api, key, value) + api.created_at = created_at + api.updated_at = updated_at return api @@ -65,23 +69,20 @@ class ExternalDatasetServiceTestDataFactory: tenant_id: str = "tenant-123", name: str = "Test Dataset", provider: str = "external", - **kwargs, - ) -> Mock: - """Create a mock Dataset object.""" - dataset = Mock(spec=Dataset) - dataset.id = dataset_id - dataset.tenant_id = tenant_id - dataset.name = name - dataset.provider = provider - dataset.description = kwargs.get("description", "") - dataset.retrieval_model = kwargs.get("retrieval_model", {}) - dataset.created_by = kwargs.get("created_by", "user-123") - - for key, value in kwargs.items(): - if key not in ["description", "retrieval_model", "created_by"]: - setattr(dataset, key, value) - - return dataset + description: str = "", + retrieval_model: dict[str, Any] | None = None, + created_by: str = "user-123", + ) -> Dataset: + """Create a Dataset object.""" + return Dataset( + id=dataset_id, + tenant_id=tenant_id, + name=name, + provider=provider, + description=description, + retrieval_model=retrieval_model or {}, + created_by=created_by, + ) @staticmethod def create_external_knowledge_binding_mock( @@ -90,20 +91,17 @@ class ExternalDatasetServiceTestDataFactory: dataset_id: str = "dataset-123", external_knowledge_api_id: str = "api-123", external_knowledge_id: str = "knowledge-123", - **kwargs, - ) -> Mock: - """Create a mock ExternalKnowledgeBindings object.""" - binding = Mock(spec=ExternalKnowledgeBindings) + created_by: str = "user-123", + ) -> ExternalKnowledgeBindings: + """Create an ExternalKnowledgeBindings object.""" + binding = ExternalKnowledgeBindings( + tenant_id=tenant_id, + external_knowledge_api_id=external_knowledge_api_id, + dataset_id=dataset_id, + external_knowledge_id=external_knowledge_id, + created_by=created_by, + ) binding.id = binding_id - binding.tenant_id = tenant_id - binding.dataset_id = dataset_id - binding.external_knowledge_api_id = external_knowledge_api_id - binding.external_knowledge_id = external_knowledge_id - binding.created_by = kwargs.get("created_by", "user-123") - - for key, value in kwargs.items(): - if key != "created_by": - setattr(binding, key, value) return binding @@ -140,6 +138,100 @@ def factory(): return ExternalDatasetServiceTestDataFactory() +def _make_external_knowledge_api( + *, + api_id: str = "api-123", + tenant_id: str = "tenant-123", + name: str = "Test API", + description: str = "Test description", + settings: dict[str, Any] | list[dict[str, Any]] | None = None, + created_by: str = "user-123", + updated_by: str = "user-123", +) -> ExternalKnowledgeApis: + """Build a real ExternalKnowledgeApis row for SQLite-backed service tests.""" + if settings is None: + settings = {"endpoint": "https://api.example.com", "api_key": "test-key-123"} + api = ExternalKnowledgeApis( + tenant_id=tenant_id, + created_by=created_by, + updated_by=updated_by, + name=name, + description=description, + settings=json.dumps(settings, ensure_ascii=False), + ) + api.id = api_id + return api + + +def _make_dataset( + *, + dataset_id: str = "dataset-123", + tenant_id: str = "tenant-123", + name: str = "Test Dataset", + provider: str = "external", + description: str = "", + retrieval_model: dict[str, Any] | None = None, + created_by: str = "user-123", +) -> Dataset: + """Build a real Dataset row with the fields required by ExternalDatasetService.""" + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name=name, + description=description, + provider=provider, + retrieval_model=retrieval_model or {}, + created_by=created_by, + maintainer=created_by, + ) + return dataset + + +def _make_external_knowledge_binding( + *, + binding_id: str = "binding-123", + tenant_id: str = "tenant-123", + dataset_id: str = "dataset-123", + external_knowledge_api_id: str = "api-123", + external_knowledge_id: str = "knowledge-123", + created_by: str = "user-123", +) -> ExternalKnowledgeBindings: + """Build a real ExternalKnowledgeBindings row for tenant-scoped lookup tests.""" + binding = ExternalKnowledgeBindings( + tenant_id=tenant_id, + dataset_id=dataset_id, + external_knowledge_api_id=external_knowledge_api_id, + external_knowledge_id=external_knowledge_id, + created_by=created_by, + ) + binding.id = binding_id + return binding + + +def _add_and_commit(session: Session, *objects: object) -> None: + """Persist rows so service methods exercise real SQLAlchemy queries.""" + session.add_all(objects) + session.commit() + + +def _seed_external_retrieval_dependencies( + session: Session, + *, + tenant_id: str = "tenant-123", + dataset_id: str = "dataset-123", + api_id: str = "api-123", +) -> tuple[ExternalKnowledgeBindings, ExternalKnowledgeApis]: + """Seed the binding and API template required by fetch_external_knowledge_retrieval.""" + binding = _make_external_knowledge_binding( + tenant_id=tenant_id, + dataset_id=dataset_id, + external_knowledge_api_id=api_id, + ) + api = _make_external_knowledge_api(api_id=api_id, tenant_id=tenant_id) + _add_and_commit(session, binding, api) + return binding, api + + class TestExternalDatasetServiceGetAPIs: """Test get_external_knowledge_apis operations - comprehensive coverage.""" @@ -434,13 +526,13 @@ class TestExternalDatasetServiceValidateAPIList: ExternalDatasetService.validate_api_list(api_settings) +@pytest.mark.parametrize("sqlite_session", [(ExternalKnowledgeApis,)], indirect=True) class TestExternalDatasetServiceCreateAPI: """Test create_external_knowledge_api operations.""" - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_success_full( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test successful creation with all fields.""" # Arrange @@ -453,7 +545,7 @@ class TestExternalDatasetServiceCreateAPI: } # Act - result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, session=mock_db.session) + result = ExternalDatasetService.create_external_knowledge_api(tenant_id, user_id, args, session=sqlite_session) # Assert assert result.name == "Test API" @@ -462,14 +554,12 @@ class TestExternalDatasetServiceCreateAPI: assert result.created_by == user_id assert result.updated_by == user_id mock_check.assert_called_once_with(args["settings"]) - mock_db.session.add.assert_called_once() - mock_db.session.flush.assert_called_once() - mock_db.session.commit.assert_not_called() + persisted_api = sqlite_session.get(ExternalKnowledgeApis, result.id) + assert persisted_api is result - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_minimal_fields( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test creation with minimal required fields.""" # Arrange @@ -480,16 +570,16 @@ class TestExternalDatasetServiceCreateAPI: # Act result = ExternalDatasetService.create_external_knowledge_api( - "tenant-123", "user-123", args, session=mock_db.session + "tenant-123", "user-123", args, session=sqlite_session ) # Assert assert result.name == "Minimal API" assert result.description == "" + assert sqlite_session.get(ExternalKnowledgeApis, result.id) is result - @patch("services.external_knowledge_service.db") def test_create_external_knowledge_api_missing_settings( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test creation fails when settings are missing.""" # Arrange @@ -497,26 +587,22 @@ class TestExternalDatasetServiceCreateAPI: # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api( - "tenant-123", "user-123", args, session=mock_db.session - ) + ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, session=sqlite_session) - @patch("services.external_knowledge_service.db") - def test_create_external_knowledge_api_none_settings(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_create_external_knowledge_api_none_settings( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test creation fails when settings are explicitly None.""" # Arrange args = {"name": "Test API", "settings": None} # Act & Assert with pytest.raises(ValueError, match="settings is required"): - ExternalDatasetService.create_external_knowledge_api( - "tenant-123", "user-123", args, session=mock_db.session - ) + ExternalDatasetService.create_external_knowledge_api("tenant-123", "user-123", args, session=sqlite_session) - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_settings_json_serialization( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test that settings are properly JSON serialized.""" # Arrange @@ -529,7 +615,7 @@ class TestExternalDatasetServiceCreateAPI: # Act result = ExternalDatasetService.create_external_knowledge_api( - "tenant-123", "user-123", args, session=mock_db.session + "tenant-123", "user-123", args, session=sqlite_session ) # Assert @@ -537,10 +623,9 @@ class TestExternalDatasetServiceCreateAPI: parsed_settings = json.loads(result.settings) assert parsed_settings == settings - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_unicode_handling( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test proper handling of Unicode characters in name and description.""" # Arrange @@ -552,17 +637,16 @@ class TestExternalDatasetServiceCreateAPI: # Act result = ExternalDatasetService.create_external_knowledge_api( - "tenant-123", "user-123", args, session=mock_db.session + "tenant-123", "user-123", args, session=sqlite_session ) # Assert assert result.name == "测试API" assert result.description == "テストの説明" - @patch("services.external_knowledge_service.db") @patch("services.external_knowledge_service.ExternalDatasetService.check_endpoint_and_api_key") def test_create_external_knowledge_api_long_description( - self, mock_check, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, mock_check, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test creation with very long description.""" # Arrange @@ -575,7 +659,7 @@ class TestExternalDatasetServiceCreateAPI: # Act result = ExternalDatasetService.create_external_knowledge_api( - "tenant-123", "user-123", args, session=mock_db.session + "tenant-123", "user-123", args, session=sqlite_session ) # Assert @@ -783,6 +867,43 @@ class TestExternalDatasetServiceCheckEndpoint: with pytest.raises(ValueError, match="Forbidden.*Authorization failed"): ExternalDatasetService.check_endpoint_and_api_key(settings) + @patch("services.external_knowledge_service.ssrf_proxy") + def test_check_endpoint_403_message_does_not_echo_api_key( + self, mock_proxy, factory: ExternalDatasetServiceTestDataFactory + ): + """Regression for #39888: the 403 error message must not contain the raw api_key. + + Before the fix, `external_knowledge_service.py:117` interpolated + `api_key` into the `ValueError` message, so the credential round-tripped + in the application log (via `current_app.logger.exception` in + `api/libs/external_api.py:94`) and in the 400 response body + (`{"code": "invalid_param", "message": str(e), ...}`). The 403 status + from the upstream provider was the only signal that the key was bad; + echoing it back is just a plaintext credential leak. + """ + # Arrange -- a real-looking key with a prefix that would be a high-signal + # substring to grep for in logs. + api_key = "sk-abcdefghijklmnop1234567890ABCDEF" + settings = {"endpoint": "https://api.example.com", "api_key": api_key} + + mock_response = MagicMock() + mock_response.status_code = 403 + mock_proxy.post.return_value = mock_response + + # Act + with pytest.raises(ValueError) as exc_info: + ExternalDatasetService.check_endpoint_and_api_key(settings) + + # Assert -- the message names the failure but does not include the key. + message = str(exc_info.value) + assert "Forbidden" in message + assert "Authorization failed" in message + assert api_key not in message + # Belt-and-braces: also check the prefix and a 6-char tail to catch + # regressions that only echo part of the key. + assert "sk-abcdef" not in message + assert "CDEF" not in message + @patch("services.external_knowledge_service.ssrf_proxy") def test_check_endpoint_other_4xx_codes_pass(self, mock_proxy, factory: ExternalDatasetServiceTestDataFactory): """Test that other 4xx codes don't raise exceptions.""" @@ -861,43 +982,42 @@ class TestExternalDatasetServiceCheckEndpoint: assert call_kwargs["headers"]["Authorization"] == "Bearer test-key-123" +@pytest.mark.parametrize("sqlite_session", [(ExternalKnowledgeApis,)], indirect=True) class TestExternalDatasetServiceGetAPI: """Test get_external_knowledge_api operations.""" - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_api_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_get_external_knowledge_api_success( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test successful retrieval of external knowledge API.""" # Arrange api_id = "api-123" - expected_api = factory.create_external_knowledge_api_mock(api_id=api_id) - - mock_db.session.scalar.return_value = expected_api + expected_api = _make_external_knowledge_api(api_id=api_id) + _add_and_commit(sqlite_session, expected_api) # Act tenant_id = "tenant-123" - result = ExternalDatasetService.get_external_knowledge_api(api_id, tenant_id, session=mock_db.session) + result = ExternalDatasetService.get_external_knowledge_api(api_id, tenant_id, session=sqlite_session) # Assert assert result.id == api_id - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_get_external_knowledge_api_not_found( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test error when API is not found.""" - # Arrange - mock_db.session.scalar.return_value = None - # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.get_external_knowledge_api("nonexistent-id", "tenant-123", session=mock_db.session) + ExternalDatasetService.get_external_knowledge_api("nonexistent-id", "tenant-123", session=sqlite_session) +@pytest.mark.parametrize("sqlite_session", [(ExternalKnowledgeApis,)], indirect=True) class TestExternalDatasetServiceUpdateAPI: """Test update_external_knowledge_api operations.""" @patch("services.external_knowledge_service.naive_utc_now") - @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_success_all_fields( - self, mock_db, mock_now, factory: ExternalDatasetServiceTestDataFactory + self, mock_now, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test successful update with all fields.""" # Arrange @@ -907,7 +1027,8 @@ class TestExternalDatasetServiceUpdateAPI: current_time = datetime(2024, 1, 2, 12, 0) mock_now.return_value = current_time - existing_api = factory.create_external_knowledge_api_mock(api_id=api_id, tenant_id=tenant_id) + existing_api = _make_external_knowledge_api(api_id=api_id, tenant_id=tenant_id) + _add_and_commit(sqlite_session, existing_api) args = { "name": "Updated API", @@ -915,11 +1036,9 @@ class TestExternalDatasetServiceUpdateAPI: "settings": {"endpoint": "https://new.example.com", "api_key": "new-key"}, } - mock_db.session.scalar.return_value = existing_api - # Act result = ExternalDatasetService.update_external_knowledge_api( - tenant_id, user_id, api_id, args, session=mock_db.session + tenant_id, user_id, api_id, args, session=sqlite_session ) # Assert @@ -927,193 +1046,207 @@ class TestExternalDatasetServiceUpdateAPI: assert result.description == "Updated description" assert result.updated_by == user_id assert result.updated_at == current_time - mock_db.session.flush.assert_called_once() - mock_db.session.commit.assert_not_called() + assert sqlite_session.get(ExternalKnowledgeApis, api_id) is result - @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_preserve_hidden_api_key( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test that hidden API key is preserved from existing settings.""" # Arrange api_id = "api-123" tenant_id = "tenant-123" - existing_api = factory.create_external_knowledge_api_mock( + existing_api = _make_external_knowledge_api( api_id=api_id, tenant_id=tenant_id, settings={"endpoint": "https://api.example.com", "api_key": "original-secret-key"}, ) + _add_and_commit(sqlite_session, existing_api) args = { "name": "Updated API", "settings": {"endpoint": "https://api.example.com", "api_key": HIDDEN_VALUE}, } - mock_db.session.scalar.return_value = existing_api - # Act result = ExternalDatasetService.update_external_knowledge_api( - tenant_id, "user-123", api_id, args, session=mock_db.session + tenant_id, "user-123", api_id, args, session=sqlite_session ) # Assert settings = json.loads(result.settings) assert settings["api_key"] == "original-secret-key" - @patch("services.external_knowledge_service.db") - def test_update_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_update_external_knowledge_api_not_found( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test error when API is not found.""" # Arrange - mock_db.session.scalar.return_value = None - args = {"name": "Updated API"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): ExternalDatasetService.update_external_knowledge_api( - "tenant-123", "user-123", "api-123", args, session=mock_db.session + "tenant-123", "user-123", "api-123", args, session=sqlite_session ) - @patch("services.external_knowledge_service.db") def test_update_external_knowledge_api_tenant_mismatch( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test error when tenant ID doesn't match.""" # Arrange - mock_db.session.scalar.return_value = None - + _add_and_commit(sqlite_session, _make_external_knowledge_api(api_id="api-123", tenant_id="tenant-123")) args = {"name": "Updated API"} # Act & Assert with pytest.raises(ValueError, match="api template not found"): ExternalDatasetService.update_external_knowledge_api( - "wrong-tenant", "user-123", "api-123", args, session=mock_db.session + "wrong-tenant", "user-123", "api-123", args, session=sqlite_session ) - @patch("services.external_knowledge_service.db") - def test_update_external_knowledge_api_name_only(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_update_external_knowledge_api_name_only( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test updating only the name field.""" # Arrange - existing_api = factory.create_external_knowledge_api_mock( + existing_api = _make_external_knowledge_api( description="Original description", settings={"endpoint": "https://api.example.com", "api_key": "key"}, ) + _add_and_commit(sqlite_session, existing_api) args = {"name": "New Name Only"} - mock_db.session.scalar.return_value = existing_api - # Act result = ExternalDatasetService.update_external_knowledge_api( - "tenant-123", "user-123", "api-123", args, session=mock_db.session + "tenant-123", "user-123", "api-123", args, session=sqlite_session ) # Assert assert result.name == "New Name Only" +@pytest.mark.parametrize("sqlite_session", [(ExternalKnowledgeApis,)], indirect=True) class TestExternalDatasetServiceDeleteAPI: """Test delete_external_knowledge_api operations.""" - @patch("services.external_knowledge_service.db") - def test_delete_external_knowledge_api_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_delete_external_knowledge_api_success( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test successful deletion of external knowledge API.""" # Arrange api_id = "api-123" tenant_id = "tenant-123" - existing_api = factory.create_external_knowledge_api_mock(api_id=api_id, tenant_id=tenant_id) - - mock_db.session.scalar.return_value = existing_api + existing_api = _make_external_knowledge_api(api_id=api_id, tenant_id=tenant_id) + _add_and_commit(sqlite_session, existing_api) # Act - ExternalDatasetService.delete_external_knowledge_api(tenant_id, api_id, session=mock_db.session) + ExternalDatasetService.delete_external_knowledge_api(tenant_id, api_id, session=sqlite_session) # Assert - mock_db.session.delete.assert_called_once_with(existing_api) - mock_db.session.flush.assert_called_once() - mock_db.session.commit.assert_not_called() + assert sqlite_session.get(ExternalKnowledgeApis, api_id) is None - @patch("services.external_knowledge_service.db") - def test_delete_external_knowledge_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_delete_external_knowledge_api_not_found( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test error when API is not found.""" - # Arrange - mock_db.session.scalar.return_value = None - # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api("tenant-123", "api-123", session=mock_db.session) + ExternalDatasetService.delete_external_knowledge_api("tenant-123", "api-123", session=sqlite_session) - @patch("services.external_knowledge_service.db") def test_delete_external_knowledge_api_tenant_mismatch( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test error when tenant ID doesn't match.""" # Arrange - mock_db.session.scalar.return_value = None + _add_and_commit(sqlite_session, _make_external_knowledge_api(api_id="api-123", tenant_id="tenant-123")) # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.delete_external_knowledge_api("wrong-tenant", "api-123", session=mock_db.session) + ExternalDatasetService.delete_external_knowledge_api("wrong-tenant", "api-123", session=sqlite_session) +@pytest.mark.parametrize("sqlite_session", [(ExternalKnowledgeBindings,)], indirect=True) class TestExternalDatasetServiceAPIUseCheck: """Test external_knowledge_api_use_check operations.""" - @patch("services.external_knowledge_service.db") def test_external_knowledge_api_use_check_in_use_single( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test API use check when API has one binding.""" # Arrange api_id = "api-123" tenant_id = "tenant-123" - mock_db.session.scalar.return_value = 1 + _add_and_commit( + sqlite_session, + _make_external_knowledge_binding(external_knowledge_api_id=api_id, tenant_id=tenant_id), + _make_external_knowledge_binding( + binding_id="binding-other", + external_knowledge_api_id=api_id, + tenant_id="other-tenant", + ), + ) # Act in_use, count = ExternalDatasetService.external_knowledge_api_use_check( - api_id, tenant_id, session=mock_db.session + api_id, tenant_id, session=sqlite_session ) # Assert assert in_use is True assert count == 1 - assert "tenant_id" in str(mock_db.session.scalar.call_args.args[0]) - @patch("services.external_knowledge_service.db") def test_external_knowledge_api_use_check_in_use_multiple( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test API use check with multiple bindings.""" # Arrange api_id = "api-123" tenant_id = "tenant-123" - mock_db.session.scalar.return_value = 10 + _add_and_commit( + sqlite_session, + *[ + _make_external_knowledge_binding( + binding_id=f"binding-{index}", + external_knowledge_api_id=api_id, + tenant_id=tenant_id, + dataset_id=f"dataset-{index}", + ) + for index in range(10) + ], + ) # Act in_use, count = ExternalDatasetService.external_knowledge_api_use_check( - api_id, tenant_id, session=mock_db.session + api_id, tenant_id, session=sqlite_session ) # Assert assert in_use is True assert count == 10 - @patch("services.external_knowledge_service.db") - def test_external_knowledge_api_use_check_not_in_use(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_external_knowledge_api_use_check_not_in_use( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test API use check when API is not in use.""" # Arrange api_id = "api-123" tenant_id = "tenant-123" - mock_db.session.scalar.return_value = 0 + _add_and_commit( + sqlite_session, + _make_external_knowledge_binding( + external_knowledge_api_id=api_id, + tenant_id="other-tenant", + ), + ) # Act in_use, count = ExternalDatasetService.external_knowledge_api_use_check( - api_id, tenant_id, session=mock_db.session + api_id, tenant_id, session=sqlite_session ) # Assert @@ -1121,48 +1254,47 @@ class TestExternalDatasetServiceAPIUseCheck: assert count == 0 +@pytest.mark.parametrize("sqlite_session", [(ExternalKnowledgeBindings,)], indirect=True) class TestExternalDatasetServiceGetBinding: """Test get_external_knowledge_binding_with_dataset_id operations.""" - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_binding_success(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_get_external_knowledge_binding_success( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test successful retrieval of external knowledge binding.""" # Arrange tenant_id = "tenant-123" dataset_id = "dataset-123" - expected_binding = factory.create_external_knowledge_binding_mock(tenant_id=tenant_id, dataset_id=dataset_id) - - mock_db.session.scalar.return_value = expected_binding + expected_binding = _make_external_knowledge_binding(tenant_id=tenant_id, dataset_id=dataset_id) + _add_and_commit(sqlite_session, expected_binding) # Act result = ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( - tenant_id, dataset_id, session=mock_db.session + tenant_id, dataset_id, session=sqlite_session ) # Assert assert result.dataset_id == dataset_id assert result.tenant_id == tenant_id - @patch("services.external_knowledge_service.db") - def test_get_external_knowledge_binding_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_get_external_knowledge_binding_not_found( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test error when binding is not found.""" - # Arrange - mock_db.session.scalar.return_value = None - # Act & Assert with pytest.raises(ValueError, match="external knowledge binding not found"): ExternalDatasetService.get_external_knowledge_binding_with_dataset_id( - "tenant-123", "dataset-123", session=mock_db.session + "tenant-123", "dataset-123", session=sqlite_session ) +@pytest.mark.parametrize("sqlite_session", [(ExternalKnowledgeApis,)], indirect=True) class TestExternalDatasetServiceDocumentValidate: """Test document_create_args_validate operations.""" - @patch("services.external_knowledge_service.db") def test_document_create_args_validate_success_all_params( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test successful validation with all required parameters.""" # Arrange @@ -1177,20 +1309,18 @@ class TestExternalDatasetServiceDocumentValidate: ] } - api = factory.create_external_knowledge_api_mock(api_id=api_id, settings=[settings]) - - mock_db.session.scalar.return_value = api + api = _make_external_knowledge_api(api_id=api_id, tenant_id=tenant_id, settings=[settings]) + _add_and_commit(sqlite_session, api) process_parameter = {"param1": "value1", "param2": "value2"} # Act & Assert - should not raise ExternalDatasetService.document_create_args_validate( - tenant_id, api_id, process_parameter, session=mock_db.session + tenant_id, api_id, process_parameter, session=sqlite_session ) - @patch("services.external_knowledge_service.db") def test_document_create_args_validate_missing_required_param( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test validation fails when required parameter is missing.""" # Arrange @@ -1199,45 +1329,39 @@ class TestExternalDatasetServiceDocumentValidate: settings = {"document_process_setting": [{"name": "required_param", "required": True}]} - api = factory.create_external_knowledge_api_mock(api_id=api_id, settings=[settings]) - - mock_db.session.scalar.return_value = api + api = _make_external_knowledge_api(api_id=api_id, tenant_id=tenant_id, settings=[settings]) + _add_and_commit(sqlite_session, api) process_parameter = {} # Act & Assert with pytest.raises(ValueError, match="required_param is required"): ExternalDatasetService.document_create_args_validate( - tenant_id, api_id, process_parameter, session=mock_db.session + tenant_id, api_id, process_parameter, session=sqlite_session ) - @patch("services.external_knowledge_service.db") - def test_document_create_args_validate_api_not_found(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_document_create_args_validate_api_not_found( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test validation fails when API is not found.""" - # Arrange - mock_db.session.scalar.return_value = None - # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}, session=mock_db.session) + ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}, session=sqlite_session) - @patch("services.external_knowledge_service.db") def test_document_create_args_validate_no_custom_parameters( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test validation succeeds when no custom parameters defined.""" # Arrange settings = {} - api = factory.create_external_knowledge_api_mock(settings=[settings]) - - mock_db.session.scalar.return_value = api + api = _make_external_knowledge_api(settings=[settings]) + _add_and_commit(sqlite_session, api) # Act & Assert - should not raise - ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}, session=mock_db.session) + ExternalDatasetService.document_create_args_validate("tenant-123", "api-123", {}, session=sqlite_session) - @patch("services.external_knowledge_service.db") def test_document_create_args_validate_optional_params_not_required( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test that optional parameters don't cause validation failure.""" # Arrange @@ -1248,15 +1372,14 @@ class TestExternalDatasetServiceDocumentValidate: ] } - api = factory.create_external_knowledge_api_mock(settings=[settings]) - - mock_db.session.scalar.return_value = api + api = _make_external_knowledge_api(settings=[settings]) + _add_and_commit(sqlite_session, api) process_parameter = {"required_param": "value"} # Act & Assert - should not raise ExternalDatasetService.document_create_args_validate( - "tenant-123", "api-123", process_parameter, session=mock_db.session + "tenant-123", "api-123", process_parameter, session=sqlite_session ) @@ -1544,107 +1667,120 @@ class TestExternalDatasetServiceGetSettings: assert result.params["key1"] == "value1" +@pytest.mark.parametrize( + "sqlite_session", + [(Dataset, ExternalKnowledgeApis, ExternalKnowledgeBindings)], + indirect=True, +) class TestExternalDatasetServiceCreateDataset: """Test create_external_dataset operations.""" - @patch("services.external_knowledge_service.db") - def test_create_external_dataset_success_full(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_create_external_dataset_success_full( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test successful creation of external dataset with all fields.""" # Arrange tenant_id = "tenant-123" user_id = "user-123" - args = { - "name": "Test External Dataset", - "description": "Comprehensive test description", - "external_knowledge_api_id": "api-123", - "external_knowledge_id": "knowledge-123", - "external_retrieval_model": {"top_k": 5, "score_threshold": 0.7}, - } + args = ExternalDatasetCreatePayload.model_validate( + { + "name": "Test External Dataset", + "description": "Comprehensive test description", + "external_knowledge_api_id": "api-123", + "external_knowledge_id": "knowledge-123", + "external_retrieval_model": {"top_k": 5, "score_threshold": 0.7}, + } + ) - api = factory.create_external_knowledge_api_mock(api_id="api-123") - - mock_db.session.scalar.side_effect = [None, api] + api = _make_external_knowledge_api(api_id="api-123", tenant_id=tenant_id) + _add_and_commit(sqlite_session, api) # Act - result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, session=mock_db.session) + result = ExternalDatasetService.create_external_dataset(tenant_id, user_id, args, session=sqlite_session) # Assert assert result.name == "Test External Dataset" assert result.description == "Comprehensive test description" assert result.provider == "external" assert result.created_by == user_id - mock_db.session.add.assert_called() - mock_db.session.flush.assert_called_once() - mock_db.session.commit.assert_called_once() + binding = sqlite_session.scalar( + select(ExternalKnowledgeBindings).where( + ExternalKnowledgeBindings.dataset_id == result.id, + ExternalKnowledgeBindings.tenant_id == tenant_id, + ) + ) + assert binding is not None + assert binding.external_knowledge_api_id == "api-123" + assert binding.external_knowledge_id == "knowledge-123" - @patch("services.external_knowledge_service.db") def test_create_external_dataset_duplicate_name_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test error when dataset name already exists.""" # Arrange - existing_dataset = factory.create_dataset_mock(name="Duplicate Dataset") + existing_dataset = _make_dataset(name="Duplicate Dataset") + _add_and_commit(sqlite_session, existing_dataset) - mock_db.session.scalar.return_value = existing_dataset - - args = {"name": "Duplicate Dataset"} + args = ExternalDatasetCreatePayload.model_validate( + { + "name": "Duplicate Dataset", + "external_knowledge_api_id": "api-123", + "external_knowledge_id": "knowledge-123", + } + ) # Act & Assert with pytest.raises(DatasetNameDuplicateError): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=sqlite_session) - @patch("services.external_knowledge_service.db") - def test_create_external_dataset_api_not_found_error(self, mock_db, factory: ExternalDatasetServiceTestDataFactory): + def test_create_external_dataset_api_not_found_error( + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session + ): """Test error when external knowledge API is not found.""" - # Arrange - mock_db.session.scalar.side_effect = [None, None] - - args = {"name": "Test Dataset", "external_knowledge_api_id": "nonexistent-api"} + args = ExternalDatasetCreatePayload.model_validate( + { + "name": "Test Dataset", + "external_knowledge_api_id": "nonexistent-api", + "external_knowledge_id": "knowledge-123", + } + ) # Act & Assert with pytest.raises(ValueError, match="api template not found"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) + ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=sqlite_session) - @patch("services.external_knowledge_service.db") def test_create_external_dataset_missing_knowledge_id_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test error when external_knowledge_id is missing.""" # Arrange - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [None, api] - - args = {"name": "Test Dataset", "external_knowledge_api_id": "api-123"} + api = _make_external_knowledge_api() + _add_and_commit(sqlite_session, api) # Act & Assert - with pytest.raises(ValueError, match="external_knowledge_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) + with pytest.raises(ValueError, match="external_knowledge_id"): + ExternalDatasetCreatePayload.model_validate( + {"name": "Test Dataset", "external_knowledge_api_id": "api-123"} + ) - @patch("services.external_knowledge_service.db") def test_create_external_dataset_missing_api_id_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test error when external_knowledge_api_id is missing.""" - # Arrange - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [None, api] - - args = {"name": "Test Dataset", "external_knowledge_id": "knowledge-123"} - # Act & Assert - with pytest.raises(ValueError, match="external_knowledge_api_id is required"): - ExternalDatasetService.create_external_dataset("tenant-123", "user-123", args, session=mock_db.session) + with pytest.raises(ValueError, match="external_knowledge_api_id"): + ExternalDatasetCreatePayload.model_validate( + {"name": "Test Dataset", "external_knowledge_id": "knowledge-123"} + ) +@pytest.mark.parametrize("sqlite_session", [(ExternalKnowledgeApis, ExternalKnowledgeBindings)], indirect=True) class TestExternalDatasetServiceFetchRetrieval: """Test fetch_external_knowledge_retrieval operations.""" @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_success_with_results( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test successful external knowledge retrieval with results.""" # Arrange @@ -1652,12 +1788,7 @@ class TestExternalDatasetServiceFetchRetrieval: dataset_id = "dataset-123" query = "test query for retrieval" - binding = factory.create_external_knowledge_binding_mock( - dataset_id=dataset_id, external_knowledge_api_id="api-123" - ) - api = factory.create_external_knowledge_api_mock(api_id="api-123") - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id) mock_response = MagicMock() mock_response.status_code = 200 @@ -1673,11 +1804,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - tenant_id, - dataset_id, - query, - external_retrieval_parameters, - session=mock_db.session, + tenant_id, dataset_id, query, external_retrieval_parameters, session=sqlite_session ) # Assert @@ -1685,46 +1812,38 @@ class TestExternalDatasetServiceFetchRetrieval: assert result[0]["content"] == "result 1" assert result[1]["score"] == 0.8 - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_binding_not_found_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test error when external knowledge binding is not found.""" - # Arrange - mock_db.session.scalar.return_value = None - # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external knowledge binding not found"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {}, session=mock_db.session + "tenant-123", "dataset-123", "query", {}, session=sqlite_session ) - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_cross_tenant_api_template_error( - self, mock_db, factory: ExternalDatasetServiceTestDataFactory + self, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test error when a binding points to an API template outside the dataset tenant.""" # Arrange - binding = factory.create_external_knowledge_binding_mock() - mock_db.session.scalar.side_effect = [binding, None] + binding = _make_external_knowledge_binding(tenant_id="tenant-123", external_knowledge_api_id="api-123") + cross_tenant_api = _make_external_knowledge_api(api_id="api-123", tenant_id="other-tenant") + _add_and_commit(sqlite_session, binding, cross_tenant_api) # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="external api template not found"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", "dataset-123", "query", {}, session=mock_db.session + "tenant-123", "dataset-123", "query", {}, session=sqlite_session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_empty_results( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test retrieval with empty results.""" # Arrange - binding = factory.create_external_knowledge_binding_mock() - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session) mock_response = MagicMock() mock_response.status_code = 200 @@ -1733,27 +1852,19 @@ class TestExternalDatasetServiceFetchRetrieval: # Act result = ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", - "dataset-123", - "query", - {"top_k": 5}, - session=mock_db.session, + "tenant-123", "dataset-123", "query", {"top_k": 5}, session=sqlite_session ) # Assert assert len(result) == 0 @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_with_score_threshold( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test retrieval with score threshold enabled.""" # Arrange - binding = factory.create_external_knowledge_binding_mock() - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session) mock_response = MagicMock() mock_response.status_code = 200 @@ -1772,7 +1883,7 @@ class TestExternalDatasetServiceFetchRetrieval: "dataset-123", "query", external_retrieval_parameters, - session=mock_db.session, + session=sqlite_session, ) # Assert @@ -1782,16 +1893,12 @@ class TestExternalDatasetServiceFetchRetrieval: assert call_args.params["retrieval_setting"]["score_threshold"] == 0.75 @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_non_200_status_raises_exception( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test that non-200 status code raises Exception with response text.""" # Arrange - binding = factory.create_external_knowledge_binding_mock() - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session) mock_response = MagicMock() mock_response.status_code = 500 @@ -1801,11 +1908,7 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match="Internal Server Error: Database connection failed"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", - "dataset-123", - "query", - {"top_k": 5}, - session=mock_db.session, + "tenant-123", "dataset-123", "query", {"top_k": 5}, session=sqlite_session ) @pytest.mark.parametrize( @@ -1822,21 +1925,20 @@ class TestExternalDatasetServiceFetchRetrieval: ], ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_various_error_status_codes( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory, status_code, error_message + self, + mock_process, + factory: ExternalDatasetServiceTestDataFactory, + sqlite_session: Session, + status_code, + error_message, ): """Test that various error status codes raise exceptions with response text.""" # Arrange tenant_id = "tenant-123" dataset_id = "dataset-123" - binding = factory.create_external_knowledge_binding_mock( - dataset_id=dataset_id, external_knowledge_api_id="api-123" - ) - api = factory.create_external_knowledge_api_mock(api_id="api-123") - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session, tenant_id=tenant_id, dataset_id=dataset_id) mock_response = MagicMock() mock_response.status_code = status_code @@ -1846,20 +1948,16 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError, match=re.escape(error_message)): ExternalDatasetService.fetch_external_knowledge_retrieval( - tenant_id, dataset_id, "query", {"top_k": 5}, session=mock_db.session + tenant_id, dataset_id, "query", {"top_k": 5}, session=sqlite_session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") def test_fetch_external_knowledge_retrieval_empty_response_text( - self, mock_db, mock_process, factory: ExternalDatasetServiceTestDataFactory + self, mock_process, factory: ExternalDatasetServiceTestDataFactory, sqlite_session: Session ): """Test exception with empty response text.""" # Arrange - binding = factory.create_external_knowledge_binding_mock() - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session) mock_response = MagicMock() mock_response.status_code = 503 @@ -1869,21 +1967,15 @@ class TestExternalDatasetServiceFetchRetrieval: # Act & Assert with pytest.raises(ExternalKnowledgeRetrievalError): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", - "dataset-123", - "query", - {"top_k": 5}, - session=mock_db.session, + "tenant-123", "dataset-123", "query", {"top_k": 5}, session=sqlite_session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") - def test_fetch_external_knowledge_retrieval_invalid_json_response(self, mock_db, mock_process, factory): + def test_fetch_external_knowledge_retrieval_invalid_json_response( + self, mock_process, factory, sqlite_session: Session + ): """Test malformed JSON success responses are normalized to external retrieval errors.""" - binding = factory.create_external_knowledge_binding_mock() - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session) mock_response = MagicMock() mock_response.status_code = 200 @@ -1892,21 +1984,15 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", - "dataset-123", - "query", - {"top_k": 5}, - session=mock_db.session, + "tenant-123", "dataset-123", "query", {"top_k": 5}, session=sqlite_session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") - def test_fetch_external_knowledge_retrieval_invalid_success_payload_shape(self, mock_db, mock_process, factory): + def test_fetch_external_knowledge_retrieval_invalid_success_payload_shape( + self, mock_process, factory, sqlite_session: Session + ): """Test malformed success payload shapes are normalized to external retrieval errors.""" - binding = factory.create_external_knowledge_binding_mock() - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session) mock_response = MagicMock() mock_response.status_code = 200 @@ -1915,21 +2001,15 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", - "dataset-123", - "query", - {"top_k": 5}, - session=mock_db.session, + "tenant-123", "dataset-123", "query", {"top_k": 5}, session=sqlite_session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") - def test_fetch_external_knowledge_retrieval_invalid_records_shape(self, mock_db, mock_process, factory): + def test_fetch_external_knowledge_retrieval_invalid_records_shape( + self, mock_process, factory, sqlite_session: Session + ): """Test non-list records payloads are normalized to external retrieval errors.""" - binding = factory.create_external_knowledge_binding_mock() - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session) mock_response = MagicMock() mock_response.status_code = 200 @@ -1938,28 +2018,18 @@ class TestExternalDatasetServiceFetchRetrieval: with pytest.raises(ExternalKnowledgeRetrievalError, match="invalid external knowledge response"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", - "dataset-123", - "query", - {"top_k": 5}, - session=mock_db.session, + "tenant-123", "dataset-123", "query", {"top_k": 5}, session=sqlite_session ) @patch("services.external_knowledge_service.ExternalDatasetService.process_external_api") - @patch("services.external_knowledge_service.db") - def test_fetch_external_knowledge_retrieval_wraps_transport_errors(self, mock_db, mock_process, factory): + def test_fetch_external_knowledge_retrieval_wraps_transport_errors( + self, mock_process, factory, sqlite_session: Session + ): """Test transport/runtime failures are normalized to external retrieval errors.""" - binding = factory.create_external_knowledge_binding_mock() - api = factory.create_external_knowledge_api_mock() - - mock_db.session.scalar.side_effect = [binding, api] + _seed_external_retrieval_dependencies(sqlite_session) mock_process.side_effect = RuntimeError("connection reset by peer") with pytest.raises(ExternalKnowledgeRetrievalError, match="connection reset by peer"): ExternalDatasetService.fetch_external_knowledge_retrieval( - "tenant-123", - "dataset-123", - "query", - {"top_k": 5}, - session=mock_db.session, + "tenant-123", "dataset-123", "query", {"top_k": 5}, session=sqlite_session ) diff --git a/api/tests/unit_tests/services/test_feature_entities.py b/api/tests/unit_tests/services/test_feature_entities.py new file mode 100644 index 00000000000..4b04d837036 --- /dev/null +++ b/api/tests/unit_tests/services/test_feature_entities.py @@ -0,0 +1,42 @@ +import pytest +from pydantic import ValidationError + +from enums import CloudPlan +from services.entities.feature_entities import LicenseLimitationModel, SubscriptionModel + + +def test_subscription_model_uses_the_cloud_plan_value_set() -> None: + subscription = SubscriptionModel(plan="team") + + assert subscription.plan is CloudPlan.TEAM + + with pytest.raises(ValidationError): + SubscriptionModel(plan="unknown") + + +@pytest.mark.parametrize( + ("enabled", "size", "limit", "required", "expected"), + [ + (False, 5, 10, 3, True), + (False, 5, 10, 10, True), + (True, 5, 0, 3, True), + (True, 5, 0, 100, True), + (True, 5, 10, 3, True), + (True, 5, 10, 5, True), + (True, 5, 10, 1, True), + (True, 8, 10, 3, False), + (True, 8, 10, 2, True), + (True, 8, 10, 1, True), + (True, 7, 10, 3, True), + ], +) +def test_license_limitation_availability( + enabled: bool, + size: int, + limit: int, + required: int, + expected: bool, +) -> None: + limitation = LicenseLimitationModel(enabled=enabled, size=size, limit=limit) + + assert limitation.is_available(required) is expected diff --git a/api/tests/unit_tests/services/test_feature_query_service.py b/api/tests/unit_tests/services/test_feature_query_service.py new file mode 100644 index 00000000000..febbfbc6af3 --- /dev/null +++ b/api/tests/unit_tests/services/test_feature_query_service.py @@ -0,0 +1,65 @@ +from unittest.mock import create_autospec + +import pytest + +from enums import DeploymentEdition +from machinery.context import RequestContext +from services.entities.feature_entities import ( + FeatureModel, + LicenseModel, + SystemFeatureModel, + VectorSpaceLimitationModel, +) +from services.feature_query_service import FeatureQueryGateway, FeatureQueryService + + +def _request_context(*, active_workspace_id: str | None = "workspace_123") -> RequestContext: + return RequestContext( + request_id="request_123", + trace_id=None, + account_id="account_123", + active_workspace_id=active_workspace_id, + ) + + +def test_workspace_queries_use_workspace_from_request_context() -> None: + gateway = create_autospec(FeatureQueryGateway, instance=True, spec_set=True) + features = FeatureModel() + vector_space = VectorSpaceLimitationModel(size=1, limit=5) + gateway.get_workspace_features.return_value = features + gateway.get_vector_space.return_value = vector_space + service = FeatureQueryService(features=gateway, trial_models=(), app_dsl_version="0.7.0") + context = _request_context() + + assert service.get_features(context) is features + assert service.get_vector_space(context) is vector_space + gateway.get_workspace_features.assert_called_once_with("workspace_123") + gateway.get_vector_space.assert_called_once_with("workspace_123") + + +def test_deployment_queries_delegate_without_request_context() -> None: + gateway = create_autospec(FeatureQueryGateway, instance=True, spec_set=True) + system_features = SystemFeatureModel(deployment_edition=DeploymentEdition.COMMUNITY) + license_model = LicenseModel() + gateway.get_public_system_features.return_value = system_features + gateway.get_license.return_value = license_model + service = FeatureQueryService( + features=gateway, + trial_models=["langgenius/openai/openai"], + app_dsl_version="0.6.0", + ) + + assert service.get_trial_models() == ["langgenius/openai/openai"] + assert service.get_app_dsl_version() == "0.6.0" + assert service.get_system_features() is system_features + assert service.get_license() is license_model + + +def test_workspace_queries_require_active_workspace() -> None: + gateway = create_autospec(FeatureQueryGateway, instance=True, spec_set=True) + service = FeatureQueryService(features=gateway, trial_models=(), app_dsl_version="0.7.0") + + with pytest.raises(RuntimeError, match="did not resolve an active workspace"): + service.get_features(_request_context(active_workspace_id=None)) + + gateway.get_workspace_features.assert_not_called() diff --git a/api/tests/unit_tests/services/test_feature_service_app_dsl_version.py b/api/tests/unit_tests/services/test_feature_service_app_dsl_version.py index b2ba664394e..9aa5301fbb0 100644 --- a/api/tests/unit_tests/services/test_feature_service_app_dsl_version.py +++ b/api/tests/unit_tests/services/test_feature_service_app_dsl_version.py @@ -1,4 +1,3 @@ -from constants.dsl_version import CURRENT_APP_DSL_VERSION from services.feature_service import FeatureService @@ -6,9 +5,3 @@ def test_get_system_features_excludes_app_dsl_version(): result = FeatureService.get_system_features().model_dump() assert "app_dsl_version" not in result - - -def test_get_app_dsl_version_returns_current_version(): - result = FeatureService.get_app_dsl_version() - - assert result == CURRENT_APP_DSL_VERSION diff --git a/api/tests/unit_tests/services/test_feature_service_deployment_edition.py b/api/tests/unit_tests/services/test_feature_service_deployment_edition.py index 966247c610c..8336b923ffb 100644 --- a/api/tests/unit_tests/services/test_feature_service_deployment_edition.py +++ b/api/tests/unit_tests/services/test_feature_service_deployment_edition.py @@ -1,8 +1,11 @@ +from unittest.mock import MagicMock + import pytest from pydantic import ValidationError -from enums.deployment_edition import DeploymentEdition -from services.feature_service import FeatureService, SystemFeatureModel +from enums import DeploymentEdition +from services.entities.feature_entities import SystemFeatureModel +from services.feature_service import FeatureService def test_system_feature_model_requires_deployment_edition() -> None: @@ -11,25 +14,29 @@ def test_system_feature_model_requires_deployment_edition() -> None: @pytest.mark.parametrize( - ("edition", "enterprise_enabled", "expected"), + "edition", [ - ("SELF_HOSTED", False, DeploymentEdition.COMMUNITY), - ("SELF_HOSTED", True, DeploymentEdition.ENTERPRISE), - ("CLOUD", False, DeploymentEdition.CLOUD), - ("CLOUD", True, DeploymentEdition.CLOUD), + DeploymentEdition.COMMUNITY, + DeploymentEdition.ENTERPRISE, + DeploymentEdition.CLOUD, ], ) -def test_get_system_features_resolves_deployment_edition( +def test_get_system_features_uses_configured_deployment_edition( monkeypatch: pytest.MonkeyPatch, - edition: str, - enterprise_enabled: bool, - expected: DeploymentEdition, + edition: DeploymentEdition, ) -> None: - monkeypatch.setattr("services.feature_service.dify_config.EDITION", edition) - monkeypatch.setattr("services.feature_service.dify_config.ENTERPRISE_ENABLED", enterprise_enabled) - monkeypatch.setattr("services.feature_service.FeatureService._fulfill_params_from_enterprise", lambda *_: None) + fulfill_from_enterprise = MagicMock() + monkeypatch.setattr("services.feature_service.dify_config.DEPLOYMENT_EDITION", edition) + monkeypatch.setattr( + "services.feature_service.FeatureService._fulfill_params_from_enterprise", + fulfill_from_enterprise, + ) result = FeatureService.get_system_features() - assert result.deployment_edition is expected - assert result.model_dump(mode="json")["deployment_edition"] == expected.value + assert result.deployment_edition is edition + assert result.model_dump(mode="json")["deployment_edition"] == edition.value + if edition is DeploymentEdition.ENTERPRISE: + fulfill_from_enterprise.assert_called_once_with(result) + else: + fulfill_from_enterprise.assert_not_called() diff --git a/api/tests/unit_tests/services/test_feature_service_enable_app_deploy.py b/api/tests/unit_tests/services/test_feature_service_enable_app_deploy.py index 0f0fa1e1dba..c9aa82a443f 100644 --- a/api/tests/unit_tests/services/test_feature_service_enable_app_deploy.py +++ b/api/tests/unit_tests/services/test_feature_service_enable_app_deploy.py @@ -1,8 +1,9 @@ import pytest -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from services import feature_service as feature_service_module -from services.feature_service import FeatureService, SystemFeatureModel +from services.entities.feature_entities import SystemFeatureModel +from services.feature_service import FeatureService @pytest.mark.parametrize( diff --git a/api/tests/unit_tests/services/test_feature_service_explore_banner.py b/api/tests/unit_tests/services/test_feature_service_explore_banner.py new file mode 100644 index 00000000000..fb709d532c2 --- /dev/null +++ b/api/tests/unit_tests/services/test_feature_service_explore_banner.py @@ -0,0 +1,30 @@ +import pytest + +from enums import DeploymentEdition +from services import feature_service as feature_service_module +from services.feature_service import FeatureService + + +@pytest.mark.parametrize( + ("edition", "configured", "expected"), + [ + (DeploymentEdition.CLOUD, True, True), + (DeploymentEdition.CLOUD, False, False), + (DeploymentEdition.COMMUNITY, True, False), + (DeploymentEdition.ENTERPRISE, True, False), + ], +) +def test_get_system_features_enables_explore_banner_only_for_cloud( + monkeypatch: pytest.MonkeyPatch, + edition: DeploymentEdition, + configured: bool, + expected: bool, +) -> None: + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", edition) + monkeypatch.setattr(feature_service_module.dify_config, "ENABLE_EXPLORE_BANNER", configured) + monkeypatch.setattr(FeatureService, "_fulfill_params_from_enterprise", lambda *_: None) + + result = FeatureService.get_system_features() + + assert FeatureService.is_explore_banner_enabled() is expected + assert result.enable_explore_banner is expected diff --git a/api/tests/unit_tests/services/test_feature_service_gateway.py b/api/tests/unit_tests/services/test_feature_service_gateway.py new file mode 100644 index 00000000000..253df785743 --- /dev/null +++ b/api/tests/unit_tests/services/test_feature_service_gateway.py @@ -0,0 +1,26 @@ +from pytest_mock import MockerFixture + +from enums import DeploymentEdition +from services.entities.feature_entities import FeatureModel, SystemFeatureModel +from services.feature_service import FeatureService +from services.feature_service_gateway import FeatureServiceGateway + + +def test_public_system_features_delegate_to_existing_service(mocker: MockerFixture) -> None: + system_features = SystemFeatureModel(deployment_edition=DeploymentEdition.COMMUNITY) + get_system_features = mocker.patch.object(FeatureService, "get_system_features", return_value=system_features) + + result = FeatureServiceGateway().get_public_system_features() + + assert result is system_features + get_system_features.assert_called_once_with() + + +def test_workspace_features_exclude_independently_queried_vector_space(mocker: MockerFixture) -> None: + features = FeatureModel(vector_space=None) + get_features = mocker.patch.object(FeatureService, "get_features", return_value=features) + + result = FeatureServiceGateway().get_workspace_features("workspace_123") + + assert result is features + get_features.assert_called_once_with("workspace_123", exclude_vector_space=True) diff --git a/api/tests/unit_tests/services/test_feature_service_human_input_email_delivery.py b/api/tests/unit_tests/services/test_feature_service_human_input_email_delivery.py index 8614d351f19..497e331a448 100644 --- a/api/tests/unit_tests/services/test_feature_service_human_input_email_delivery.py +++ b/api/tests/unit_tests/services/test_feature_service_human_input_email_delivery.py @@ -2,16 +2,16 @@ from dataclasses import dataclass import pytest -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from services import feature_service as feature_service_module -from services.feature_service import FeatureModel, FeatureService +from services.entities.feature_entities import FeatureModel +from services.feature_service import FeatureService @dataclass(frozen=True) class HumanInputEmailDeliveryCase: name: str - enterprise_enabled: bool - billing_enabled: bool + deployment_edition: DeploymentEdition tenant_id: str | None billing_feature_enabled: bool plan: str @@ -20,27 +20,24 @@ class HumanInputEmailDeliveryCase: CASES = [ HumanInputEmailDeliveryCase( - name="enterprise_enabled", - enterprise_enabled=True, - billing_enabled=True, + name="enterprise_edition", + deployment_edition=DeploymentEdition.ENTERPRISE, tenant_id=None, billing_feature_enabled=False, plan=CloudPlan.SANDBOX, expected=True, ), HumanInputEmailDeliveryCase( - name="billing_disabled", - enterprise_enabled=False, - billing_enabled=False, + name="community_edition", + deployment_edition=DeploymentEdition.COMMUNITY, tenant_id=None, billing_feature_enabled=False, plan=CloudPlan.SANDBOX, expected=True, ), HumanInputEmailDeliveryCase( - name="billing_enabled_requires_tenant", - enterprise_enabled=False, - billing_enabled=True, + name="cloud_edition_requires_tenant", + deployment_edition=DeploymentEdition.CLOUD, tenant_id=None, billing_feature_enabled=True, plan=CloudPlan.PROFESSIONAL, @@ -48,8 +45,7 @@ CASES = [ ), HumanInputEmailDeliveryCase( name="billing_feature_off", - enterprise_enabled=False, - billing_enabled=True, + deployment_edition=DeploymentEdition.CLOUD, tenant_id="tenant-1", billing_feature_enabled=False, plan=CloudPlan.PROFESSIONAL, @@ -57,8 +53,7 @@ CASES = [ ), HumanInputEmailDeliveryCase( name="professional_plan", - enterprise_enabled=False, - billing_enabled=True, + deployment_edition=DeploymentEdition.CLOUD, tenant_id="tenant-1", billing_feature_enabled=True, plan=CloudPlan.PROFESSIONAL, @@ -66,8 +61,7 @@ CASES = [ ), HumanInputEmailDeliveryCase( name="team_plan", - enterprise_enabled=False, - billing_enabled=True, + deployment_edition=DeploymentEdition.CLOUD, tenant_id="tenant-1", billing_feature_enabled=True, plan=CloudPlan.TEAM, @@ -75,8 +69,7 @@ CASES = [ ), HumanInputEmailDeliveryCase( name="sandbox_plan", - enterprise_enabled=False, - billing_enabled=True, + deployment_edition=DeploymentEdition.CLOUD, tenant_id="tenant-1", billing_feature_enabled=True, plan=CloudPlan.SANDBOX, @@ -90,8 +83,7 @@ def test_resolve_human_input_email_delivery_enabled_matrix( monkeypatch: pytest.MonkeyPatch, case: HumanInputEmailDeliveryCase, ): - monkeypatch.setattr(feature_service_module.dify_config, "ENTERPRISE_ENABLED", case.enterprise_enabled) - monkeypatch.setattr(feature_service_module.dify_config, "BILLING_ENABLED", case.billing_enabled) + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", case.deployment_edition) features = FeatureModel() features.billing.enabled = case.billing_feature_enabled features.billing.subscription.plan = case.plan @@ -105,7 +97,7 @@ def test_resolve_human_input_email_delivery_enabled_matrix( def test_get_vector_space_converts_billing_float_size(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(feature_service_module.dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr( feature_service_module.BillingService, "get_vector_space", @@ -116,3 +108,19 @@ def test_get_vector_space_converts_billing_float_size(monkeypatch: pytest.Monkey assert result.size == 5120 assert result.limit == 20480 + assert result.usage_unknown is False + + +def test_get_vector_space_preserves_unknown_usage(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) + monkeypatch.setattr( + feature_service_module.BillingService, + "get_vector_space", + lambda tenant_id: {"size": 0.0, "limit": 50, "usage_unknown": True}, + ) + + result = FeatureService.get_vector_space("tenant-1") + + assert result.size == 0 + assert result.limit == 50 + assert result.usage_unknown is True diff --git a/api/tests/unit_tests/services/test_feature_service_internal_policies.py b/api/tests/unit_tests/services/test_feature_service_internal_policies.py index 2ca475001e4..b2a326e1cdb 100644 --- a/api/tests/unit_tests/services/test_feature_service_internal_policies.py +++ b/api/tests/unit_tests/services/test_feature_service_internal_policies.py @@ -1,10 +1,11 @@ import pytest +from enums import DeploymentEdition from services.feature_service import FeatureService def test_workspace_creation_uses_environment_policy(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr("services.feature_service.dify_config.ENTERPRISE_ENABLED", False) + monkeypatch.setattr("services.feature_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) monkeypatch.setattr("services.feature_service.dify_config.ALLOW_CREATE_WORKSPACE", True) monkeypatch.setattr( "services.feature_service.EnterpriseService.get_info", @@ -15,7 +16,7 @@ def test_workspace_creation_uses_environment_policy(monkeypatch: pytest.MonkeyPa def test_workspace_creation_uses_enterprise_policy(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr("services.feature_service.dify_config.ENTERPRISE_ENABLED", True) + monkeypatch.setattr("services.feature_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) monkeypatch.setattr( "services.feature_service.EnterpriseService.get_info", lambda: {"IsAllowCreateWorkspace": False}, @@ -27,7 +28,7 @@ def test_workspace_creation_uses_enterprise_policy(monkeypatch: pytest.MonkeyPat def test_workspace_creation_keeps_environment_policy_when_enterprise_value_is_missing( monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr("services.feature_service.dify_config.ENTERPRISE_ENABLED", True) + monkeypatch.setattr("services.feature_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) monkeypatch.setattr("services.feature_service.dify_config.ALLOW_CREATE_WORKSPACE", True) monkeypatch.setattr("services.feature_service.EnterpriseService.get_info", lambda: {}) @@ -35,8 +36,8 @@ def test_workspace_creation_keeps_environment_policy_when_enterprise_value_is_mi def test_plugin_manager_is_enabled_only_for_enterprise(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr("services.feature_service.dify_config.ENTERPRISE_ENABLED", True) + monkeypatch.setattr("services.feature_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) assert FeatureService.is_plugin_manager_enabled() is True - monkeypatch.setattr("services.feature_service.dify_config.ENTERPRISE_ENABLED", False) + monkeypatch.setattr("services.feature_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) assert FeatureService.is_plugin_manager_enabled() is False diff --git a/api/tests/unit_tests/services/test_feature_service_knowledge_file_size_limit.py b/api/tests/unit_tests/services/test_feature_service_knowledge_file_size_limit.py new file mode 100644 index 00000000000..2bb9a6123c4 --- /dev/null +++ b/api/tests/unit_tests/services/test_feature_service_knowledge_file_size_limit.py @@ -0,0 +1,70 @@ +from unittest.mock import Mock + +import pytest + +from enums import CloudPlan, DeploymentEdition +from services import feature_service as feature_service_module +from services.feature_service import FeatureService + + +@pytest.mark.parametrize( + ("deployment_edition", "tenant_id", "billing_feature_enabled", "plan", "expected"), + [ + (DeploymentEdition.COMMUNITY, "tenant-1", True, CloudPlan.PROFESSIONAL, 15), + (DeploymentEdition.ENTERPRISE, "tenant-1", True, CloudPlan.PROFESSIONAL, 15), + (DeploymentEdition.CLOUD, None, True, CloudPlan.PROFESSIONAL, 15), + (DeploymentEdition.CLOUD, "tenant-1", False, CloudPlan.PROFESSIONAL, 15), + (DeploymentEdition.CLOUD, "tenant-1", True, CloudPlan.SANDBOX, 15), + (DeploymentEdition.CLOUD, "tenant-1", True, CloudPlan.PROFESSIONAL, 50), + (DeploymentEdition.CLOUD, "tenant-1", True, CloudPlan.TEAM, 50), + ], +) +def test_get_knowledge_file_size_limit( + monkeypatch: pytest.MonkeyPatch, + deployment_edition: DeploymentEdition, + tenant_id: str | None, + billing_feature_enabled: bool, + plan: CloudPlan, + expected: int, +) -> None: + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", deployment_edition) + monkeypatch.setattr(feature_service_module.dify_config, "UPLOAD_FILE_SIZE_LIMIT", 15) + monkeypatch.setattr( + feature_service_module.dify_config, + "KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN", + 50, + ) + get_info = Mock( + return_value={ + "enabled": billing_feature_enabled, + "subscription": {"plan": plan}, + } + ) + monkeypatch.setattr(feature_service_module.BillingService, "get_info", get_info) + + assert FeatureService.get_knowledge_file_size_limit(tenant_id) == expected + + if deployment_edition == DeploymentEdition.CLOUD and tenant_id: + get_info.assert_called_once_with(tenant_id, exclude_vector_space=True) + else: + get_info.assert_not_called() + + +def test_paid_knowledge_file_size_limit_never_reduces_default(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) + monkeypatch.setattr(feature_service_module.dify_config, "UPLOAD_FILE_SIZE_LIMIT", 100) + monkeypatch.setattr( + feature_service_module.dify_config, + "KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN", + 50, + ) + monkeypatch.setattr( + feature_service_module.BillingService, + "get_info", + lambda *_args, **_kwargs: { + "enabled": True, + "subscription": {"plan": CloudPlan.PROFESSIONAL}, + }, + ) + + assert FeatureService.get_knowledge_file_size_limit("tenant-1") == 100 diff --git a/api/tests/unit_tests/services/test_feature_service_knowledge_fs.py b/api/tests/unit_tests/services/test_feature_service_knowledge_fs.py index c151edf5406..e9b7d0e0761 100644 --- a/api/tests/unit_tests/services/test_feature_service_knowledge_fs.py +++ b/api/tests/unit_tests/services/test_feature_service_knowledge_fs.py @@ -1,7 +1,8 @@ import pytest -from enums.deployment_edition import DeploymentEdition -from services.feature_service import FeatureService, SystemFeatureModel +from enums import DeploymentEdition +from services.entities.feature_entities import SystemFeatureModel +from services.feature_service import FeatureService def test_system_feature_model_disables_knowledge_fs_by_default() -> None: diff --git a/api/tests/unit_tests/services/test_feature_service_learn_app.py b/api/tests/unit_tests/services/test_feature_service_learn_app.py index bc33bfabbc1..32525b8f3e4 100644 --- a/api/tests/unit_tests/services/test_feature_service_learn_app.py +++ b/api/tests/unit_tests/services/test_feature_service_learn_app.py @@ -1,8 +1,9 @@ import pytest -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from services import feature_service as feature_service_module -from services.feature_service import FeatureService, SystemFeatureModel +from services.entities.feature_entities import SystemFeatureModel +from services.feature_service import FeatureService def test_system_feature_model_defaults_enable_learn_app(): diff --git a/api/tests/unit_tests/services/test_feature_service_license_expiry_notice.py b/api/tests/unit_tests/services/test_feature_service_license_expiry_notice.py new file mode 100644 index 00000000000..84c130dda12 --- /dev/null +++ b/api/tests/unit_tests/services/test_feature_service_license_expiry_notice.py @@ -0,0 +1,54 @@ +import pytest + +from enums import DeploymentEdition +from services import feature_service as feature_service_module +from services.entities.feature_entities import LicenseModel, LicenseStatus +from services.feature_service import FeatureService + +_ENTERPRISE_INFO = {"License": {"status": LicenseStatus.EXPIRING, "expiredAt": "2026-12-31"}} + + +def test_license_model_defaults_license_expiry_notice_disabled() -> None: + """Without a license there is no expiry to announce, so the notice is off unless enabled explicitly.""" + assert LicenseModel().license_expiry_notice_enabled is False + + +@pytest.mark.parametrize("enabled", [True, False]) +def test_get_license_non_enterprise_ignores_expiry_notice_config( + monkeypatch: pytest.MonkeyPatch, enabled: bool +) -> None: + """Non-enterprise deployments have no license, so the env toggle never turns the notice on.""" + monkeypatch.setattr(feature_service_module.dify_config, "ENABLE_LICENSE_EXPIRY_NOTICE", enabled) + monkeypatch.setattr( + feature_service_module.dify_config, + "DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ) + + result = FeatureService.get_license() + + assert result.license_expiry_notice_enabled is False + + +@pytest.mark.parametrize("enabled", [True, False]) +def test_get_license_enterprise_reads_license_expiry_notice_enabled( + monkeypatch: pytest.MonkeyPatch, enabled: bool +) -> None: + """The enterprise-sourced license carries the env-resolved notice flag alongside its real status.""" + monkeypatch.setattr(feature_service_module.dify_config, "ENABLE_LICENSE_EXPIRY_NOTICE", enabled) + monkeypatch.setattr( + feature_service_module.dify_config, + "DEPLOYMENT_EDITION", + DeploymentEdition.ENTERPRISE, + ) + monkeypatch.setattr( + feature_service_module.EnterpriseService, + "get_info", + staticmethod(lambda: _ENTERPRISE_INFO), + ) + + result = FeatureService.get_license() + + assert result.status == LicenseStatus.EXPIRING + assert result.expired_at == "2026-12-31" + assert result.license_expiry_notice_enabled is enabled diff --git a/api/tests/unit_tests/services/test_feature_service_licensed_seats.py b/api/tests/unit_tests/services/test_feature_service_licensed_seats.py index fb1a5139118..1231c9c56ea 100644 --- a/api/tests/unit_tests/services/test_feature_service_licensed_seats.py +++ b/api/tests/unit_tests/services/test_feature_service_licensed_seats.py @@ -1,14 +1,16 @@ import pytest +from enums import DeploymentEdition from services import feature_service as feature_service_module -from services.feature_service import FeatureService, LicenseModel, LicenseStatus +from services.entities.feature_entities import LicenseModel, LicenseStatus +from services.feature_service import FeatureService _ENTERPRISE_INFO = {"License": {"licensedSeats": {"enabled": True, "limit": 3, "used": 1}}} def test_get_license_parses_licensed_seats(monkeypatch: pytest.MonkeyPatch): """The authenticated license accessor copies the licensed-seat quota out of the enterprise payload.""" - monkeypatch.setattr("services.feature_service.dify_config.ENTERPRISE_ENABLED", True) + monkeypatch.setattr("services.feature_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) monkeypatch.setattr( feature_service_module.EnterpriseService, "get_info", @@ -25,7 +27,7 @@ def test_get_license_parses_licensed_seats(monkeypatch: pytest.MonkeyPatch): def test_get_license_non_enterprise_is_unconstrained(monkeypatch: pytest.MonkeyPatch): """Non-enterprise deployments have no license; seat allocation is unconstrained.""" - monkeypatch.setattr("services.feature_service.dify_config.ENTERPRISE_ENABLED", False) + monkeypatch.setattr("services.feature_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) license_model = FeatureService.get_license() diff --git a/api/tests/unit_tests/services/test_feature_service_plugin_installation_permission.py b/api/tests/unit_tests/services/test_feature_service_plugin_installation_permission.py new file mode 100644 index 00000000000..8ae24b07eb9 --- /dev/null +++ b/api/tests/unit_tests/services/test_feature_service_plugin_installation_permission.py @@ -0,0 +1,93 @@ +import logging + +import pytest + +from enums import DeploymentEdition +from services import feature_service as feature_service_module +from services.entities.feature_entities import PluginInstallationScope, SystemFeatureModel +from services.feature_service import FeatureService + + +def test_get_plugin_installation_permission_defaults_to_all_for_non_enterprise( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + + permission = FeatureService.get_plugin_installation_permission() + + assert permission.plugin_installation_scope is PluginInstallationScope.ALL + assert permission.restrict_to_marketplace_only is False + + +def test_get_plugin_installation_permission_parses_enterprise_policy( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(feature_service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) + monkeypatch.setattr( + feature_service_module.EnterpriseService, + "get_info", + staticmethod( + lambda: { + "PluginInstallationPermission": { + "pluginInstallationScope": "official_only", + "restrictToMarketplaceOnly": True, + } + } + ), + ) + + permission = FeatureService.get_plugin_installation_permission() + + assert permission.plugin_installation_scope is PluginInstallationScope.OFFICIAL_ONLY + assert permission.restrict_to_marketplace_only is True + + +@pytest.mark.parametrize( + "invalid_permission", + [ + { + "pluginInstallationScope": "unknown-scope", + "restrictToMarketplaceOnly": False, + }, + { + "pluginInstallationScope": "all", + "restrictToMarketplaceOnly": "false", + }, + ], + ids=["unknown_scope", "non_boolean_marketplace_restriction"], +) +def test_invalid_enterprise_policy_denies_all_plugin_installations( + caplog: pytest.LogCaptureFixture, + invalid_permission: dict[str, object], +) -> None: + with caplog.at_level(logging.ERROR, logger="services.feature_service"): + permission = FeatureService._resolve_plugin_installation_permission( + {"PluginInstallationPermission": invalid_permission} + ) + + assert permission.plugin_installation_scope is PluginInstallationScope.NONE + assert permission.restrict_to_marketplace_only is True + assert "denying all plugin installations" in caplog.text + + +def test_system_features_exposes_only_validated_plugin_installation_policy( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + feature_service_module.EnterpriseService, + "get_info", + staticmethod( + lambda: { + "PluginInstallationPermission": { + "pluginInstallationScope": "unknown-scope", + "restrictToMarketplaceOnly": False, + } + } + ), + ) + features = SystemFeatureModel(deployment_edition=DeploymentEdition.ENTERPRISE) + + FeatureService._fulfill_params_from_enterprise(features) + + assert features.plugin_installation_permission.plugin_installation_scope is PluginInstallationScope.NONE + assert features.plugin_installation_permission.restrict_to_marketplace_only is True diff --git a/api/tests/unit_tests/services/test_feature_service_sso_protocol.py b/api/tests/unit_tests/services/test_feature_service_sso_protocol.py new file mode 100644 index 00000000000..0177239077c --- /dev/null +++ b/api/tests/unit_tests/services/test_feature_service_sso_protocol.py @@ -0,0 +1,88 @@ +import logging + +import pytest + +from enums import DeploymentEdition +from services import feature_service as feature_service_module +from services.entities.feature_entities import SSOProtocol, SystemFeatureModel +from services.feature_service import FeatureService + + +def test_system_features_exposes_valid_enterprise_sso_protocols( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + feature_service_module.EnterpriseService, + "get_info", + staticmethod( + lambda: { + "SSOEnforcedForSigninProtocol": "saml", + "WebAppAuth": {}, + "SSOEnforcedForWebProtocol": "oidc", + } + ), + ) + features = SystemFeatureModel(deployment_edition=DeploymentEdition.ENTERPRISE) + + FeatureService._fulfill_params_from_enterprise(features) + + assert features.sso_enforced_for_signin_protocol is SSOProtocol.SAML + assert features.webapp_auth.sso_config.protocol is SSOProtocol.OIDC + + +@pytest.mark.parametrize("empty_protocol", [None, "", " "]) +def test_system_features_normalizes_empty_enterprise_sso_protocols_to_none( + monkeypatch: pytest.MonkeyPatch, + empty_protocol: object, +) -> None: + monkeypatch.setattr( + feature_service_module.EnterpriseService, + "get_info", + staticmethod( + lambda: { + "SSOEnforcedForSigninProtocol": empty_protocol, + "WebAppAuth": {}, + "SSOEnforcedForWebProtocol": empty_protocol, + } + ), + ) + features = SystemFeatureModel(deployment_edition=DeploymentEdition.ENTERPRISE) + + FeatureService._fulfill_params_from_enterprise(features) + + assert features.sso_enforced_for_signin_protocol is None + assert features.webapp_auth.sso_config.protocol is None + + +@pytest.mark.parametrize("invalid_protocol", ["unknown", 42]) +def test_system_features_rejects_invalid_enterprise_sso_protocols( + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, + invalid_protocol: object, +) -> None: + monkeypatch.setattr( + feature_service_module.EnterpriseService, + "get_info", + staticmethod( + lambda: { + "SSOEnforcedForSigninProtocol": invalid_protocol, + "WebAppAuth": {}, + "SSOEnforcedForWebProtocol": invalid_protocol, + } + ), + ) + features = SystemFeatureModel(deployment_edition=DeploymentEdition.ENTERPRISE) + + with caplog.at_level(logging.ERROR, logger="services.feature_service"): + FeatureService._fulfill_params_from_enterprise(features) + + assert features.sso_enforced_for_signin_protocol is None + assert features.webapp_auth.sso_config.protocol is None + assert caplog.text.count("Invalid Enterprise SSO protocol") == 2 + + +def test_system_features_defaults_sso_protocols_to_none() -> None: + features = SystemFeatureModel(deployment_edition=DeploymentEdition.COMMUNITY) + + assert features.sso_enforced_for_signin_protocol is None + assert features.webapp_auth.sso_config.protocol is None diff --git a/api/tests/unit_tests/services/test_feature_service_trial_models.py b/api/tests/unit_tests/services/test_feature_service_trial_models.py index eac3d6235a1..599623aaba0 100644 --- a/api/tests/unit_tests/services/test_feature_service_trial_models.py +++ b/api/tests/unit_tests/services/test_feature_service_trial_models.py @@ -1,6 +1,6 @@ import pytest -from enums.hosted_provider import HostedTrialProvider +from enums import HostedTrialProvider from services import feature_service as feature_service_module from services.feature_service import FeatureService diff --git a/api/tests/unit_tests/services/test_feature_service_vector_space.py b/api/tests/unit_tests/services/test_feature_service_vector_space.py index dc0e7ddb548..54f98d4b6db 100644 --- a/api/tests/unit_tests/services/test_feature_service_vector_space.py +++ b/api/tests/unit_tests/services/test_feature_service_vector_space.py @@ -1,5 +1,9 @@ +from typing import cast from unittest.mock import patch +from enums import CloudPlan, DeploymentEdition +from services.billing_service import BillingInfo +from services.entities.feature_entities import LimitationModel from services.feature_service import FeatureService @@ -7,7 +11,7 @@ def test_get_features_exclude_vector_space_sets_vector_space_to_none(): tenant_id = "tenant-id" billing_info = { "enabled": True, - "subscription": {"plan": "pro", "interval": "monthly", "education": False}, + "subscription": {"plan": CloudPlan.PROFESSIONAL, "interval": "monthly", "education": False}, "members": {"size": 1, "limit": 10}, "apps": {"size": 2, "limit": 20}, "documents_upload_quota": {"size": 3, "limit": 100}, @@ -24,8 +28,7 @@ def test_get_features_exclude_vector_space_sets_vector_space_to_none(): patch("services.feature_service.BillingService.get_info", return_value=billing_info) as get_info, patch("services.feature_service.BillingService.get_quota_info", return_value={}), ): - mock_config.BILLING_ENABLED = True - mock_config.ENTERPRISE_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_config.CAN_REPLACE_LOGO = False mock_config.MODEL_LB_ENABLED = False mock_config.DATASET_OPERATOR_ENABLED = False @@ -35,3 +38,15 @@ def test_get_features_exclude_vector_space_sets_vector_space_to_none(): assert features.vector_space is None get_info.assert_called_once_with(tenant_id, exclude_vector_space=True) + + +def test_full_features_keep_treating_unknown_vector_usage_as_zero(): + vector_space = LimitationModel() + + FeatureService._fulfill_vector_space_from_billing_info( + vector_space, + cast(BillingInfo, {"vector_space": {"size": 0.0, "limit": 50, "usage_unknown": True}}), + ) + + assert vector_space.size == 0 + assert vector_space.limit == 50 diff --git a/api/tests/unit_tests/services/test_feature_service_webapp_public_access.py b/api/tests/unit_tests/services/test_feature_service_webapp_public_access.py index 510acd67332..00b4bb8ec43 100644 --- a/api/tests/unit_tests/services/test_feature_service_webapp_public_access.py +++ b/api/tests/unit_tests/services/test_feature_service_webapp_public_access.py @@ -1,7 +1,8 @@ import pytest -from enums.deployment_edition import DeploymentEdition -from services.feature_service import FeatureService, SystemFeatureModel +from enums import DeploymentEdition +from services.entities.feature_entities import SystemFeatureModel +from services.feature_service import FeatureService @pytest.mark.parametrize( diff --git a/api/tests/unit_tests/services/test_feedback_service.py b/api/tests/unit_tests/services/test_feedback_service.py index 36022936ba7..8e2899e4211 100644 --- a/api/tests/unit_tests/services/test_feedback_service.py +++ b/api/tests/unit_tests/services/test_feedback_service.py @@ -242,3 +242,23 @@ class TestFeedbackService: dislike_feedback = next(item for item in feedback_data if item["feedback_rating_raw"] == "dislike") assert like_feedback["feedback_rating"] == "👍" assert dislike_feedback["feedback_rating"] == "👎" + + +class TestEndDateBoundary: + """Verify that end_date includes feedback from the entire day (fix for #40050).""" + + def test_end_date_includes_entire_day(self): + from datetime import timedelta + + end_date = "2026-08-05" + end_dt = datetime.strptime(end_date, "%Y-%m-%d") + timedelta(days=1) + + assert end_dt == datetime(2026, 8, 6, 0, 0, 0) + assert datetime(2026, 8, 5, 14, 30, 0) < end_dt + assert datetime(2026, 8, 5, 23, 59, 59) < end_dt + assert not (datetime(2026, 8, 6, 0, 0, 1) < end_dt) + + def test_start_date_includes_midnight(self): + start_dt = datetime.strptime("2026-08-01", "%Y-%m-%d") + assert start_dt == datetime(2026, 8, 1, 0, 0, 0) + assert datetime(2026, 8, 1, 0, 0, 0) >= start_dt diff --git a/api/tests/unit_tests/services/test_file_request_service.py b/api/tests/unit_tests/services/test_file_request_service.py index 3e2ea0a258a..9acfa38a026 100644 --- a/api/tests/unit_tests/services/test_file_request_service.py +++ b/api/tests/unit_tests/services/test_file_request_service.py @@ -15,7 +15,7 @@ from services.file_request_service import FileRequestService ("end-user", "service-api", UserFrom.END_USER, InvokeFrom.SERVICE_API), ], ) -def test_request_download_url_builds_file_under_bound_scope( +def test_request_download_builds_file_under_bound_scope( user_from: UserFrom | str, invoke_from: InvokeFrom | str, expected_user_from: UserFrom, @@ -29,12 +29,9 @@ def test_request_download_url_builds_file_under_bound_scope( with ( patch("services.file_request_service.bind_file_access_scope", return_value=nullcontext()) as bind_scope, patch.object(service, "_build_file", return_value=fake_file) as build_file, - patch( - "services.file_request_service.file_helpers.resolve_file_url", - return_value="https://files.example.com/x", - ) as resolve_file_url, + patch.object(service._runtime, "resolve_file_uri", return_value="/files/tools/x?sign=1") as resolve_file_uri, ): - result = service.request_download_url( + result = service.request_download( tenant_id="tenant-1", user_id="user-1", user_from=user_from, @@ -52,48 +49,23 @@ def test_request_download_url_builds_file_under_bound_scope( build_file.assert_called_once_with( mapping={"transfer_method": "tool_file", "reference": reference}, tenant_id="tenant-1" ) - resolve_file_url.assert_called_once_with(fake_file, for_external=True) + resolve_file_uri.assert_called_once_with(file=fake_file) assert result.filename == "report.pdf" assert result.mime_type == "application/pdf" assert result.size == 123 - assert result.download_url == "https://files.example.com/x" + assert result.download_uri == "/files/tools/x?sign=1" -def test_request_download_url_supports_internal_download_urls() -> None: - fake_file = MagicMock(filename="report.pdf", mime_type="application/pdf", size=123) - service = FileRequestService(access_controller=MagicMock()) - - with ( - patch("services.file_request_service.bind_file_access_scope", return_value=nullcontext()), - patch.object(service, "_build_file", return_value=fake_file), - patch( - "services.file_request_service.file_helpers.resolve_file_url", - return_value="http://internal-files/report.pdf", - ) as resolve_file_url, - ): - result = service.request_download_url( - tenant_id="tenant-1", - user_id="user-1", - user_from="account", - invoke_from="debugger", - file_mapping={"transfer_method": "tool_file", "reference": "dify-file-ref:tool-file-1"}, - for_external=False, - ) - - resolve_file_url.assert_called_once_with(fake_file, for_external=False) - assert result.download_url == "http://internal-files/report.pdf" - - -def test_request_download_url_rejects_unsupported_files() -> None: +def test_request_download_rejects_unsupported_files() -> None: service = FileRequestService(access_controller=MagicMock()) with ( patch("services.file_request_service.bind_file_access_scope", return_value=nullcontext()), patch.object(service, "_build_file", return_value=MagicMock(filename="report.pdf", mime_type=None, size=1)), - patch("services.file_request_service.file_helpers.resolve_file_url", return_value=None), + patch.object(service._runtime, "resolve_file_uri", return_value=None), ): with pytest.raises(ValueError, match="file does not support signed download"): - service.request_download_url( + service.request_download( tenant_id="tenant-1", user_id="user-1", user_from="account", diff --git a/api/tests/unit_tests/services/test_file_service.py b/api/tests/unit_tests/services/test_file_service.py index 583509e20db..0994abecf65 100644 --- a/api/tests/unit_tests/services/test_file_service.py +++ b/api/tests/unit_tests/services/test_file_service.py @@ -1,6 +1,8 @@ import base64 import hashlib import os +from collections.abc import Iterator +from datetime import UTC, datetime from unittest.mock import MagicMock, patch import pytest @@ -9,6 +11,8 @@ from sqlalchemy.orm import Session, sessionmaker from werkzeug.exceptions import NotFound from configs import dify_config +from extensions.storage.storage_type import StorageType +from models.base import TypeBase from models.enums import CreatorUserRole from models.model import Account, EndUser, UploadFile from services.errors.file import BlockedFileExtensionError, FileTooLargeError, UnsupportedFileTypeError @@ -17,31 +21,54 @@ from services.file_service import FileService class TestFileService: @pytest.fixture - def mock_db_session(self): - session = MagicMock(spec=Session) - # Mock context manager behavior - session.__enter__.return_value = session - return session + def sqlite_session_maker(self, sqlite_engine: Engine) -> sessionmaker[Session]: + TypeBase.metadata.create_all(sqlite_engine, tables=[TypeBase.metadata.tables[UploadFile.__tablename__]]) + return sessionmaker(bind=sqlite_engine, expire_on_commit=False) @pytest.fixture - def mock_session_maker(self, mock_db_session): - maker = MagicMock(spec=sessionmaker) - maker.return_value = mock_db_session - return maker + def db_session(self, sqlite_session_maker: sessionmaker[Session]) -> Iterator[Session]: + with sqlite_session_maker() as session: + yield session @pytest.fixture - def file_service(self, mock_session_maker): - return FileService(session_factory=mock_session_maker) + def file_service(self, sqlite_session_maker: sessionmaker[Session]) -> FileService: + return FileService(session_factory=sqlite_session_maker) - def test_init_with_engine(self): - engine = MagicMock(spec=Engine) - service = FileService(session_factory=engine) + @staticmethod + def _persist_upload_file( + session: Session, + *, + file_id: str = "file_id", + tenant_id: str = "tenant_id", + extension: str = "txt", + mime_type: str = "text/plain", + key: str = "key", + ) -> UploadFile: + upload_file = UploadFile( + tenant_id=tenant_id, + storage_type=StorageType.LOCAL, + key=key, + name=f"test.{extension}", + size=10, + extension=extension, + mime_type=mime_type, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user_id", + created_at=datetime(2024, 1, 1, tzinfo=UTC), + used=False, + ) + upload_file.id = file_id + session.add(upload_file) + session.commit() + return upload_file + + def test_init_with_engine(self, sqlite_engine: Engine): + service = FileService(session_factory=sqlite_engine) assert isinstance(service._session_maker, sessionmaker) - def test_init_with_sessionmaker(self): - maker = MagicMock(spec=sessionmaker) - service = FileService(session_factory=maker) - assert service._session_maker == maker + def test_init_with_sessionmaker(self, sqlite_session_maker: sessionmaker[Session]): + service = FileService(session_factory=sqlite_session_maker) + assert service._session_maker == sqlite_session_maker def test_init_invalid_factory(self): with pytest.raises(AssertionError, match="must be a sessionmaker or an Engine."): @@ -52,14 +79,14 @@ class TestFileService: @patch("services.file_service.extract_tenant_id") @patch("services.file_service.file_helpers.get_signed_file_url") def test_upload_file_success( - self, mock_get_url, mock_tenant_id, mock_now, mock_storage, file_service: FileService, mock_db_session + self, mock_get_url, mock_tenant_id, mock_now, mock_storage, file_service: FileService, db_session: Session ): # Setup mock_tenant_id.return_value = "tenant_id" - mock_now.return_value = "2024-01-01" + mock_now.return_value = datetime(2024, 1, 1, tzinfo=UTC) mock_get_url.return_value = "http://signed-url" - user = MagicMock(spec=Account) + user = Account(name="Test Account", email="test@example.com") user.id = "user_id" content = b"file content" filename = "test.jpg" @@ -81,11 +108,26 @@ class TestFileService: assert result.source_url == "http://signed-url" mock_storage.save.assert_called_once() - mock_db_session.add.assert_called_once_with(result) - mock_db_session.commit.assert_called_once() + persisted = db_session.get(UploadFile, result.id) + assert persisted is not None + assert persisted.hash == result.hash + + @pytest.mark.parametrize("text", ["ASCII text", "包含多字节 UTF-8 文本 🚀"]) + def test_upload_text_uses_utf8_byte_length(self, text: str, file_service: FileService): + with patch("services.file_service.storage") as mock_storage: + result = file_service.upload_text( + text=text, + text_name="test.txt", + user_id="user_id", + tenant_id="tenant_id", + ) + + expected_content = text.encode("utf-8") + assert result.size == len(expected_content) + mock_storage.save.assert_called_once_with(result.key, expected_content) def test_upload_file_uses_explicit_resource_tenant(self, file_service: FileService): - user = MagicMock(spec=Account) + user = Account(name="Test Account", email="test@example.com") user.id = "user-id" with ( @@ -109,10 +151,10 @@ class TestFileService: with pytest.raises(ValueError, match="Filename contains invalid characters"): file_service.upload_file(filename="invalid/file.txt", content=b"", mimetype="text/plain", user=MagicMock()) - def test_upload_file_long_filename(self, file_service: FileService, mock_db_session): + def test_upload_file_long_filename(self, file_service: FileService, db_session: Session): # Setup long_name = "a" * 210 + ".txt" - user = MagicMock(spec=Account) + user = Account(name="Test Account", email="test@example.com") user.id = "user_id" with ( @@ -124,6 +166,7 @@ class TestFileService: result = file_service.upload_file(filename=long_name, content=b"test", mimetype="text/plain", user=user) assert len(result.name) <= 205 # 200 + . + extension assert result.name.endswith(".txt") + assert db_session.get(UploadFile, result.id) is not None def test_upload_file_blocked_extension(self, file_service): with patch.object(dify_config, "inner_UPLOAD_FILE_EXTENSION_BLACKLIST", "exe"): @@ -145,9 +188,10 @@ class TestFileService: with pytest.raises(FileTooLargeError): file_service.upload_file(filename="test.jpg", content=content, mimetype="image/jpeg", user=MagicMock()) - def test_upload_file_end_user(self, file_service: FileService, mock_db_session): - user = MagicMock(spec=EndUser) - user.id = "end_user_id" + def test_upload_file_end_user(self, file_service: FileService, db_session: Session): + user = EndUser( + id="end_user_id", + ) with ( patch("services.file_service.storage"), @@ -157,6 +201,7 @@ class TestFileService: mock_tenant.return_value = "tenant" result = file_service.upload_file(filename="test.txt", content=b"test", mimetype="text/plain", user=user) assert result.created_by_role == CreatorUserRole.END_USER + assert db_session.get(UploadFile, result.id) is not None def test_is_file_size_within_limit(self): with ( @@ -180,13 +225,35 @@ class TestFileService: # Default assert FileService.is_file_size_within_limit(extension="txt", file_size=5 * 1024 * 1024) is True assert FileService.is_file_size_within_limit(extension="pdf", file_size=6 * 1024 * 1024) is False + assert ( + FileService.is_file_size_within_limit( + extension="pdf", + file_size=6 * 1024 * 1024, + default_file_size_limit=7, + ) + is True + ) + assert ( + FileService.is_file_size_within_limit( + extension="pdf", + file_size=8 * 1024 * 1024, + default_file_size_limit=7, + ) + is False + ) - def test_get_file_base64_success(self, file_service: FileService, mock_db_session): - # Setup - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.key = "test_key" - mock_db_session.scalar.return_value = upload_file + # Media-specific limits are not affected by the knowledge document override. + assert ( + FileService.is_file_size_within_limit( + extension="jpg", + file_size=11 * 1024 * 1024, + default_file_size_limit=100, + ) + is False + ) + + def test_get_file_base64_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, key="test_key") with patch("services.file_service.storage") as mock_storage: mock_storage.load_once.return_value = b"test content" @@ -198,16 +265,17 @@ class TestFileService: assert result == base64.b64encode(b"test content").decode() mock_storage.load_once.assert_called_once_with("test_key") - def test_get_file_base64_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_file_base64_not_found(self, file_service: FileService): 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 + def test_get_file_presigned_url_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file( + db_session, + extension="png", + mime_type="image/png", + key="upload_files/tenant_id/icon.png", + ) with ( patch.object(dify_config, "FILES_ACCESS_TIMEOUT", 300), @@ -224,13 +292,11 @@ class TestFileService: 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 - + def test_get_file_presigned_url_not_found(self, file_service: FileService): 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): + def test_upload_text_success(self, file_service: FileService, db_session: Session): # Setup text = "sample text" text_name = "test.txt" @@ -249,21 +315,17 @@ class TestFileService: assert result.used is True assert result.extension == "txt" mock_storage.save.assert_called_once() - mock_db_session.add.assert_called_once() - mock_db_session.commit.assert_called_once() + assert db_session.get(UploadFile, result.id) is not None - def test_upload_text_long_name(self, file_service: FileService, mock_db_session): + def test_upload_text_long_name(self, file_service: FileService, db_session: Session): long_name = "a" * 210 with patch("services.file_service.storage"): result = file_service.upload_text("text", long_name, "user", "tenant") assert len(result.name) == 200 + assert db_session.get(UploadFile, result.id) is not None - def test_get_file_preview_success(self, file_service: FileService, mock_db_session): - # Setup - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "pdf" - mock_db_session.scalar.return_value = upload_file + def test_get_file_preview_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, extension="pdf", mime_type="application/pdf") with patch("services.file_service.ExtractProcessor.load_from_upload_file") as mock_extract: mock_extract.return_value = "Extracted text content" @@ -274,27 +336,17 @@ class TestFileService: # Assert assert result == "Extracted text content" - def test_get_file_preview_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_file_preview_not_found(self, file_service: FileService): with pytest.raises(NotFound, match="File not found"): file_service.get_file_preview("non_existent", "tenant_id") - def test_get_file_preview_unsupported_type(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "exe" - mock_db_session.scalar.return_value = upload_file + def test_get_file_preview_unsupported_type(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, extension="exe", mime_type="application/octet-stream") with pytest.raises(UnsupportedFileTypeError): file_service.get_file_preview("file_id", "tenant_id") - def test_get_image_preview_success(self, file_service: FileService, mock_db_session): - # Setup - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "jpg" - upload_file.mime_type = "image/jpeg" - upload_file.key = "key" - mock_db_session.scalar.return_value = upload_file + def test_get_image_preview_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, extension="jpg", mime_type="image/jpeg") with ( patch("services.file_service.file_helpers.verify_image_signature") as mock_verify, @@ -316,28 +368,21 @@ class TestFileService: with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_image_preview("file_id", "ts", "nonce", "sign") - def test_get_image_preview_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_image_preview_not_found(self, file_service: FileService): with patch("services.file_service.file_helpers.verify_image_signature") as mock_verify: mock_verify.return_value = True with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_image_preview("file_id", "ts", "nonce", "sign") - def test_get_image_preview_unsupported_type(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "txt" - mock_db_session.scalar.return_value = upload_file + def test_get_image_preview_unsupported_type(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session) with patch("services.file_service.file_helpers.verify_image_signature") as mock_verify: mock_verify.return_value = True with pytest.raises(UnsupportedFileTypeError): file_service.get_image_preview("file_id", "ts", "nonce", "sign") - def test_get_file_generator_by_file_id_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.key = "key" - mock_db_session.scalar.return_value = upload_file + def test_get_file_generator_by_file_id_success(self, file_service: FileService, db_session: Session): + upload_file = self._persist_upload_file(db_session) with ( patch("services.file_service.file_helpers.verify_file_signature") as mock_verify, @@ -348,7 +393,8 @@ class TestFileService: gen, file = file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign") assert list(gen) == [b"chunk"] - assert file == upload_file + assert file.id == upload_file.id + assert file.key == upload_file.key def test_get_file_generator_by_file_id_invalid_sig(self, file_service): with patch("services.file_service.file_helpers.verify_file_signature") as mock_verify: @@ -356,20 +402,14 @@ class TestFileService: with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign") - def test_get_file_generator_by_file_id_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_file_generator_by_file_id_not_found(self, file_service: FileService): with patch("services.file_service.file_helpers.verify_file_signature") as mock_verify: mock_verify.return_value = True with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_file_generator_by_file_id("file_id", "ts", "nonce", "sign") - def test_get_public_image_preview_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "png" - upload_file.mime_type = "image/png" - upload_file.key = "key" - mock_db_session.scalar.return_value = upload_file + def test_get_public_image_preview_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session, extension="png", mime_type="image/png") with patch("services.file_service.storage") as mock_storage: mock_storage.load.return_value = b"image content" @@ -377,66 +417,56 @@ class TestFileService: assert gen == b"image content" assert mime == "image/png" - def test_get_public_image_preview_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_public_image_preview_not_found(self, file_service: FileService): with pytest.raises(NotFound, match="File not found or signature is invalid"): file_service.get_public_image_preview("file_id") - def test_get_public_image_preview_unsupported_type(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.extension = "txt" - mock_db_session.scalar.return_value = upload_file + def test_get_public_image_preview_unsupported_type(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session) with pytest.raises(UnsupportedFileTypeError): file_service.get_public_image_preview("file_id") - def test_get_file_content_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.key = "key" - mock_db_session.scalar.return_value = upload_file + def test_get_file_content_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session) with patch("services.file_service.storage") as mock_storage: mock_storage.load.return_value = b"hello world" result = file_service.get_file_content("file_id") assert result == "hello world" - def test_get_file_content_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_get_file_content_not_found(self, file_service: FileService): with pytest.raises(NotFound, match="File not found"): file_service.get_file_content("file_id") - def test_delete_file_success(self, file_service: FileService, mock_db_session): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "file_id" - upload_file.key = "key" - # For session.scalar(select(...)) - mock_db_session.scalar.return_value = upload_file + def test_delete_file_success(self, file_service: FileService, db_session: Session): + self._persist_upload_file(db_session) with patch("services.file_service.storage") as mock_storage: file_service.delete_file("file_id") mock_storage.delete.assert_called_once_with("key") - mock_db_session.delete.assert_called_once_with(upload_file) + db_session.expire_all() + assert db_session.get(UploadFile, "file_id") is None - def test_delete_file_not_found(self, file_service: FileService, mock_db_session): - mock_db_session.scalar.return_value = None + def test_delete_file_not_found(self, file_service: FileService): file_service.delete_file("file_id") # Should return without doing anything - def test_get_upload_files_by_ids_empty(self): - session = MagicMock() - result = FileService.get_upload_files_by_ids("tenant_id", [], session=session) + def test_get_upload_files_by_ids_empty(self, db_session: Session): + result = FileService.get_upload_files_by_ids("tenant_id", [], session=db_session) assert result == {} - def test_get_upload_files_by_ids(self): - upload_file = MagicMock(spec=UploadFile) - upload_file.id = "550e8400-e29b-41d4-a716-446655440000" - upload_file.tenant_id = "tenant_id" - session = MagicMock() - session.scalars().all.return_value = [upload_file] + def test_get_upload_files_by_ids(self, db_session: Session): + upload_file = self._persist_upload_file(db_session, file_id="550e8400-e29b-41d4-a716-446655440000") + self._persist_upload_file( + db_session, + file_id="550e8400-e29b-41d4-a716-446655440001", + tenant_id="other-tenant", + ) result = FileService.get_upload_files_by_ids( - "tenant_id", ["550e8400-e29b-41d4-a716-446655440000"], session=session + "tenant_id", + ["550e8400-e29b-41d4-a716-446655440000", "550e8400-e29b-41d4-a716-446655440001"], + session=db_session, ) assert result["550e8400-e29b-41d4-a716-446655440000"] == upload_file @@ -453,10 +483,8 @@ class TestFileService: used.add("a (1).txt") assert FileService._dedupe_zip_entry_name("a.txt", used) == "a (2).txt" - def test_build_upload_files_zip_tempfile(self): - upload_file = MagicMock(spec=UploadFile) - upload_file.name = "test.txt" - upload_file.key = "key" + def test_build_upload_files_zip_tempfile(self, db_session: Session): + upload_file = self._persist_upload_file(db_session) with ( patch("services.file_service.storage") as mock_storage, diff --git a/api/tests/unit_tests/services/test_human_input_delivery_test_service.py b/api/tests/unit_tests/services/test_human_input_delivery_test_service.py index a3e0e6f618f..fb9ebf6e9be 100644 --- a/api/tests/unit_tests/services/test_human_input_delivery_test_service.py +++ b/api/tests/unit_tests/services/test_human_input_delivery_test_service.py @@ -22,7 +22,7 @@ from graphon.runtime import VariablePool from models.account import Account, TenantAccountJoin from models.engine import db from services import human_input_delivery_test_service as service_module -from services.feature_service import FeatureModel +from services.entities.feature_entities import FeatureModel from services.human_input_delivery_test_service import ( DeliveryTestContext, DeliveryTestEmailRecipient, diff --git a/api/tests/unit_tests/services/test_human_input_service.py b/api/tests/unit_tests/services/test_human_input_service.py index 55dd86129bb..58d5d84a8e3 100644 --- a/api/tests/unit_tests/services/test_human_input_service.py +++ b/api/tests/unit_tests/services/test_human_input_service.py @@ -1,10 +1,10 @@ import dataclasses import logging -from collections.abc import Iterator from datetime import datetime, timedelta from unittest.mock import MagicMock import pytest +from pydantic import JsonValue from pytest_mock import MockerFixture from sqlalchemy.engine import Engine from sqlalchemy.orm import Session, sessionmaker @@ -41,15 +41,8 @@ from services.human_input_service import ( ) -@pytest.fixture -def sqlite_session_factory(sqlite_engine: Engine) -> Iterator[tuple[sessionmaker[Session], Session]]: - factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) - with factory() as session: - yield factory, session - - -def _persist_app(sqlite_session: Session, mode: AppMode) -> App: - app = App( +def _make_app(mode: AppMode) -> App: + return App( id="app-id", tenant_id="tenant-id", name="Test App", @@ -60,13 +53,10 @@ def _persist_app(sqlite_session: Session, mode: AppMode) -> App: enable_api=True, max_active_requests=0, ) - sqlite_session.add(app) - sqlite_session.commit() - return app @pytest.fixture -def sample_form_record(): +def sample_form_record() -> HumanInputFormRecord: return HumanInputFormRecord( form_id="form-id", workflow_run_id="workflow-run-id", @@ -97,14 +87,11 @@ def sample_form_record(): ) -@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) def test_enqueue_resume_dispatches_task_for_workflow( mocker: MockerFixture, - sqlite_session_factory, - sqlite_session: Session, -): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) + sqlite_session_factory: sessionmaker[Session], +) -> None: + service = HumanInputService(sqlite_session_factory) workflow_run = MagicMock() workflow_run.app_id = "app-id" @@ -116,7 +103,8 @@ def test_enqueue_resume_dispatches_task_for_workflow( return_value=workflow_run_repo, ) - _persist_app(sqlite_session, AppMode.WORKFLOW) + with sqlite_session_factory.begin() as arrange_session: + arrange_session.add(_make_app(AppMode.WORKFLOW)) resume_task = mocker.patch("services.human_input_service.resume_app_execution") @@ -128,10 +116,11 @@ def test_enqueue_resume_dispatches_task_for_workflow( def test_ensure_form_active_respects_global_timeout( - monkeypatch, sample_form_record: HumanInputFormRecord, sqlite_session_factory -): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) + monkeypatch: pytest.MonkeyPatch, + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: + service = HumanInputService(unbound_session_factory) expired_record = dataclasses.replace( sample_form_record, created_at=naive_utc_now() - timedelta(hours=2), @@ -143,14 +132,11 @@ def test_ensure_form_active_respects_global_timeout( service.ensure_form_active(Form(expired_record)) -@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) def test_enqueue_resume_dispatches_task_for_advanced_chat( mocker: MockerFixture, - sqlite_session_factory, - sqlite_session: Session, -): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) + sqlite_session_factory: sessionmaker[Session], +) -> None: + service = HumanInputService(sqlite_session_factory) workflow_run = MagicMock() workflow_run.app_id = "app-id" @@ -162,7 +148,8 @@ def test_enqueue_resume_dispatches_task_for_advanced_chat( return_value=workflow_run_repo, ) - _persist_app(sqlite_session, AppMode.ADVANCED_CHAT) + with sqlite_session_factory.begin() as arrange_session: + arrange_session.add(_make_app(AppMode.ADVANCED_CHAT)) resume_task = mocker.patch("services.human_input_service.resume_app_execution") @@ -173,14 +160,11 @@ def test_enqueue_resume_dispatches_task_for_advanced_chat( assert call_kwargs["kwargs"]["payload"]["workflow_run_id"] == "workflow-run-id" -@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) def test_enqueue_resume_skips_unsupported_app_mode( mocker: MockerFixture, - sqlite_session_factory, - sqlite_session: Session, -): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) + sqlite_session_factory: sessionmaker[Session], +) -> None: + service = HumanInputService(sqlite_session_factory) workflow_run = MagicMock() workflow_run.app_id = "app-id" @@ -192,7 +176,8 @@ def test_enqueue_resume_skips_unsupported_app_mode( return_value=workflow_run_repo, ) - _persist_app(sqlite_session, AppMode.COMPLETION) + with sqlite_session_factory.begin() as arrange_session: + arrange_session.add(_make_app(AppMode.COMPLETION)) resume_task = mocker.patch("services.human_input_service.resume_app_execution") @@ -202,14 +187,14 @@ def test_enqueue_resume_skips_unsupported_app_mode( def test_get_form_definition_by_token_for_console_uses_repository( - sample_form_record: HumanInputFormRecord, sqlite_session_factory -): - session_factory, _ = sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) console_record = dataclasses.replace(sample_form_record, recipient_type=RecipientType.CONSOLE) repo.get_by_token.return_value = console_record - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) form = service.get_form_definition_by_token_for_console("token") repo.get_by_token.assert_called_once_with("token") @@ -245,9 +230,10 @@ def _build_resumption_context_state(*, options: list[str], workflow_run_id: str) def test_resolve_form_inputs_uses_runtime_select_options( - sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture -): - session_factory, _ = sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, +) -> None: configured_input = SelectInputConfig( output_variable_name="decision", option_source=StringListSource( @@ -272,7 +258,7 @@ def test_resolve_form_inputs_uses_runtime_select_options( "services.human_input_service.DifyAPIRepositoryFactory.create_api_workflow_run_repository", return_value=workflow_run_repo, ) - service = HumanInputService(session_factory) + service = HumanInputService(unbound_session_factory) resolved_inputs = service.resolve_form_inputs(Form(record)) @@ -284,13 +270,14 @@ def test_resolve_form_inputs_uses_runtime_select_options( def test_submit_form_by_token_calls_repository_and_enqueue( - sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture -): - session_factory, _ = sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = sample_form_record repo.mark_submitted.return_value = sample_form_record - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) enqueue_spy = mocker.patch.object(service, "enqueue_resume") service.submit_form_by_token( @@ -313,11 +300,12 @@ def test_submit_form_by_token_calls_repository_and_enqueue( def test_submit_form_by_token_enqueues_agent_app_resume_for_conversation_form( - sample_form_record, sqlite_session_factory, mocker: MockerFixture -): + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, +) -> None: # ENG-635: a conversation-owned (Agent v2 chat) form routes to the chat # resume, not the workflow resume. - session_factory, _ = sqlite_session_factory repo = MagicMock(spec=HumanInputFormSubmissionRepository) conversation_record = dataclasses.replace( sample_form_record, @@ -326,7 +314,7 @@ def test_submit_form_by_token_enqueues_agent_app_resume_for_conversation_form( ) repo.get_by_token.return_value = conversation_record repo.mark_submitted.return_value = conversation_record - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) workflow_enqueue_spy = mocker.patch.object(service, "enqueue_resume") chat_enqueue_spy = mocker.patch.object(service, "enqueue_agent_app_resume") @@ -343,9 +331,10 @@ def test_submit_form_by_token_enqueues_agent_app_resume_for_conversation_form( def test_submit_form_by_token_skips_enqueue_for_delivery_test( - sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture -): - session_factory, _ = sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) test_record = dataclasses.replace( sample_form_record, @@ -354,7 +343,7 @@ def test_submit_form_by_token_skips_enqueue_for_delivery_test( ) repo.get_by_token.return_value = test_record repo.mark_submitted.return_value = test_record - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) enqueue_spy = mocker.patch.object(service, "enqueue_resume") service.submit_form_by_token( @@ -368,13 +357,14 @@ def test_submit_form_by_token_skips_enqueue_for_delivery_test( def test_submit_form_by_token_passes_submission_user_id( - sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture -): - session_factory, _ = sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = sample_form_record repo.mark_submitted.return_value = sample_form_record - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) enqueue_spy = mocker.patch.object(service, "enqueue_resume") service.submit_form_by_token( @@ -391,11 +381,13 @@ def test_submit_form_by_token_passes_submission_user_id( enqueue_spy.assert_called_once_with(sample_form_record.workflow_run_id) -def test_submit_form_by_token_invalid_action(sample_form_record: HumanInputFormRecord, sqlite_session_factory): - session_factory, _ = sqlite_session_factory +def test_submit_form_by_token_invalid_action( + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = dataclasses.replace(sample_form_record) - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) with pytest.raises(InvalidFormDataError) as exc_info: service.submit_form_by_token( @@ -409,8 +401,10 @@ def test_submit_form_by_token_invalid_action(sample_form_record: HumanInputFormR repo.mark_submitted.assert_not_called() -def test_submit_form_by_token_missing_inputs(sample_form_record: HumanInputFormRecord, sqlite_session_factory): - session_factory, _ = sqlite_session_factory +def test_submit_form_by_token_missing_inputs( + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) definition_with_input = FormDefinition( @@ -422,7 +416,7 @@ def test_submit_form_by_token_missing_inputs(sample_form_record: HumanInputFormR ) form_with_input = dataclasses.replace(sample_form_record, definition=definition_with_input) repo.get_by_token.return_value = form_with_input - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) with pytest.raises(InvalidFormDataError) as exc_info: service.submit_form_by_token( @@ -436,42 +430,6 @@ def test_submit_form_by_token_missing_inputs(sample_form_record: HumanInputFormR repo.mark_submitted.assert_not_called() -def test_validate_human_input_submission_accepts_select_file_and_file_list(sqlite_session_factory): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) - definition = FormDefinition.model_validate( - { - "form_content": "Pick one and upload files", - "inputs": [ - { - "type": "select", - "output_variable_name": "decision", - "option_source": { - "type": "constant", - "value": ["approve", "reject"], - }, - }, - { - "type": "file", - "output_variable_name": "attachment", - "allowed_file_types": ["document"], - "allowed_file_upload_methods": ["remote_url"], - }, - { - "type": "file-list", - "output_variable_name": "attachments", - "allowed_file_types": ["document"], - "allowed_file_upload_methods": ["remote_url"], - "number_limits": 3, - }, - ], - "user_actions": [{"id": "submit", "title": "Submit"}], - "rendered_content": "

Pick one and upload files

", - "expiration_time": naive_utc_now() + timedelta(hours=1), - } - ) - - @pytest.mark.parametrize( ("input_definition", "submitted_value", "expected_message"), [ @@ -521,13 +479,12 @@ def test_validate_human_input_submission_accepts_select_file_and_file_list(sqlit ], ) def test_validate_human_input_submission_rejects_invalid_select_and_file_payloads( - sample_form_record, - sqlite_session_factory, - input_definition, - submitted_value, - expected_message, -): - session_factory, _ = sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + input_definition: dict[str, JsonValue], + submitted_value: JsonValue, + expected_message: str, +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) definition = FormDefinition.model_validate( { @@ -539,21 +496,21 @@ def test_validate_human_input_submission_rejects_invalid_select_and_file_payload } ) repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition) - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) with pytest.raises(InvalidFormDataError) as exc_info: service.submit_form_by_token( recipient_type=RecipientType.STANDALONE_WEB_APP, form_token="token", selected_action_id="submit", - form_data={input_definition["output_variable_name"]: submitted_value}, + form_data={definition.inputs[0].output_variable_name: submitted_value}, ) assert expected_message in str(exc_info.value) repo.mark_submitted.assert_not_called() -def test_form_properties(sample_form_record: HumanInputFormRecord): +def test_form_properties(sample_form_record: HumanInputFormRecord) -> None: form = Form(sample_form_record) assert form.id == "form-id" assert form.workflow_run_id == "workflow-run-id" @@ -567,74 +524,79 @@ def test_form_properties(sample_form_record: HumanInputFormRecord): assert isinstance(form.expiration_time, datetime) -def test_form_submitted_error_init(): +def test_form_submitted_error_init() -> None: error = FormSubmittedError(form_id="test-form") - assert "form_id=test-form" in error.description + assert error.description == "This form has already been submitted by another user, form_id=test-form" assert error.code == 412 -def test_human_input_service_init_with_engine(sqlite_engine: Engine): +def test_human_input_service_init_with_engine(sqlite_engine: Engine) -> None: service = HumanInputService(session_factory=sqlite_engine) assert isinstance(service._session_factory, sessionmaker) assert service._session_factory.kw["bind"] is sqlite_engine -def test_get_form_by_token_none(sqlite_session_factory): - session_factory, _ = sqlite_session_factory +def test_get_form_by_token_none(unbound_session_factory: sessionmaker[Session]) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = None - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) assert service.get_form_by_token("invalid") is None -def test_get_form_definition_by_token_mismatch(sample_form_record: HumanInputFormRecord, sqlite_session_factory): - session_factory, _ = sqlite_session_factory +def test_get_form_definition_by_token_mismatch( + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = sample_form_record - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) # RecipientType mismatch assert service.get_form_definition_by_token(RecipientType.CONSOLE, "token") is None -def test_get_form_definition_by_token_success(sample_form_record: HumanInputFormRecord, sqlite_session_factory): - session_factory, _ = sqlite_session_factory +def test_get_form_definition_by_token_success( + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = sample_form_record - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) form = service.get_form_definition_by_token(RecipientType.STANDALONE_WEB_APP, "token") assert form is not None assert form.id == sample_form_record.form_id def test_get_form_definition_by_token_for_console_mismatch( - sample_form_record: HumanInputFormRecord, sqlite_session_factory -): - session_factory, _ = sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = sample_form_record # is STANDALONE_WEB_APP - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) assert service.get_form_definition_by_token_for_console("token") is None -def test_submit_form_by_token_delivery_not_enabled(sqlite_session_factory): - session_factory, _ = sqlite_session_factory +def test_submit_form_by_token_delivery_not_enabled( + unbound_session_factory: sessionmaker[Session], +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = None - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) with pytest.raises(human_input_service_module.WebAppDeliveryNotEnabledError): service.submit_form_by_token(RecipientType.STANDALONE_WEB_APP, "token", "action", {}) def test_submit_form_by_token_no_workflow_run_id( - sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture -): - session_factory, _ = sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, +) -> None: repo = MagicMock(spec=HumanInputFormSubmissionRepository) repo.get_by_token.return_value = sample_form_record @@ -642,16 +604,18 @@ def test_submit_form_by_token_no_workflow_run_id( result_record = dataclasses.replace(sample_form_record, workflow_run_id=None) repo.mark_submitted.return_value = result_record - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) enqueue_spy = mocker.patch.object(service, "enqueue_resume") service.submit_form_by_token(RecipientType.STANDALONE_WEB_APP, "token", "submit", {}) enqueue_spy.assert_not_called() -def test_ensure_form_active_errors(sample_form_record: HumanInputFormRecord, sqlite_session_factory): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) +def test_ensure_form_active_errors( + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: + service = HumanInputService(unbound_session_factory) # Submitted submitted_record = dataclasses.replace(sample_form_record, submitted_at=naive_utc_now()) @@ -671,18 +635,22 @@ def test_ensure_form_active_errors(sample_form_record: HumanInputFormRecord, sql service.ensure_form_active(Form(expired_time_record)) -def test_ensure_not_submitted_raises(sample_form_record: HumanInputFormRecord, sqlite_session_factory): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) +def test_ensure_not_submitted_raises( + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: + service = HumanInputService(unbound_session_factory) submitted_record = dataclasses.replace(sample_form_record, submitted_at=naive_utc_now()) with pytest.raises(human_input_service_module.FormSubmittedError): service._ensure_not_submitted(Form(submitted_record)) -def test_enqueue_resume_workflow_not_found(mocker: MockerFixture, sqlite_session_factory): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) +def test_enqueue_resume_workflow_not_found( + mocker: MockerFixture, + unbound_session_factory: sessionmaker[Session], +) -> None: + service = HumanInputService(unbound_session_factory) workflow_run_repo = MagicMock() workflow_run_repo.get_workflow_run_by_id_without_tenant.return_value = None @@ -696,15 +664,12 @@ def test_enqueue_resume_workflow_not_found(mocker: MockerFixture, sqlite_session assert "WorkflowRun not found" in str(excinfo.value) -@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) def test_enqueue_resume_app_not_found( - mocker, - sqlite_session_factory, - sqlite_session: Session, + mocker: MockerFixture, + sqlite_session_factory: sessionmaker[Session], caplog: pytest.LogCaptureFixture, -): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) +) -> None: + service = HumanInputService(sqlite_session_factory) workflow_run = MagicMock() workflow_run.app_id = "app-id" @@ -715,26 +680,35 @@ def test_enqueue_resume_app_not_found( "services.human_input_service.DifyAPIRepositoryFactory.create_api_workflow_run_repository", return_value=workflow_run_repo, ) + resume_task = mocker.patch("services.human_input_service.resume_app_execution") with caplog.at_level(logging.ERROR, logger="services.human_input_service"): service.enqueue_resume("workflow-run-id") - assert any(r.levelno >= logging.ERROR for r in caplog.records) + + assert ( + "services.human_input_service", + logging.ERROR, + "App not found for WorkflowRun, workflow_run_id=workflow-run-id, app_id=app-id", + ) in caplog.record_tuples + resume_task.apply_async.assert_not_called() def test_is_globally_expired_zero_timeout( - monkeypatch: pytest.MonkeyPatch, sample_form_record: HumanInputFormRecord, sqlite_session_factory -): - session_factory, _ = sqlite_session_factory - service = HumanInputService(session_factory) + monkeypatch: pytest.MonkeyPatch, + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], +) -> None: + service = HumanInputService(unbound_session_factory) monkeypatch.setattr(human_input_service_module.dify_config, "HUMAN_INPUT_GLOBAL_TIMEOUT_SECONDS", 0) assert service._is_globally_expired(Form(sample_form_record)) is False def test_submit_form_by_token_normalizes_select_and_files( - sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, ) -> None: - session_factory, _ = sqlite_session_factory repo = MagicMock(spec=HumanInputFormSubmissionRepository) definition = FormDefinition( form_content="hello", @@ -753,7 +727,7 @@ def test_submit_form_by_token_normalizes_select_and_files( form_with_inputs = dataclasses.replace(sample_form_record, definition=definition) repo.get_by_token.return_value = form_with_inputs repo.mark_submitted.return_value = form_with_inputs - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) single_file = File( file_id="file-1", @@ -815,9 +789,9 @@ def test_submit_form_by_token_normalizes_select_and_files( def test_submit_form_by_token_invalid_select_value( - sample_form_record: HumanInputFormRecord, sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], ) -> None: - session_factory, _ = sqlite_session_factory repo = MagicMock(spec=HumanInputFormSubmissionRepository) definition = FormDefinition( form_content="hello", @@ -832,7 +806,7 @@ def test_submit_form_by_token_invalid_select_value( expiration_time=sample_form_record.expiration_time, ) repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition) - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) with pytest.raises(InvalidFormDataError, match="Invalid value for select input 'decision'"): service.submit_form_by_token( @@ -844,9 +818,9 @@ def test_submit_form_by_token_invalid_select_value( def test_submit_form_by_token_invalid_file_list_item( - sample_form_record: HumanInputFormRecord, sqlite_session_factory + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], ) -> None: - session_factory, _ = sqlite_session_factory repo = MagicMock(spec=HumanInputFormSubmissionRepository) definition = FormDefinition( form_content="hello", @@ -856,7 +830,7 @@ def test_submit_form_by_token_invalid_file_list_item( expiration_time=sample_form_record.expiration_time, ) repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition) - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) with pytest.raises( InvalidFormDataError, @@ -871,9 +845,10 @@ def test_submit_form_by_token_invalid_file_list_item( def test_submit_form_by_token_rejects_cross_tenant_file( - sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, ) -> None: - session_factory, _ = sqlite_session_factory repo = MagicMock(spec=HumanInputFormSubmissionRepository) definition = FormDefinition( form_content="hello", @@ -883,7 +858,7 @@ def test_submit_form_by_token_rejects_cross_tenant_file( expiration_time=sample_form_record.expiration_time, ) repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition) - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) mocker.patch("services.human_input_service.build_from_mapping", side_effect=ValueError("Invalid upload file")) with pytest.raises(InvalidFormDataError, match="Invalid value for file input 'attachment'"): @@ -904,9 +879,10 @@ def test_submit_form_by_token_rejects_cross_tenant_file( def test_submit_form_by_token_rejects_cross_tenant_file_list( - sample_form_record: HumanInputFormRecord, sqlite_session_factory, mocker: MockerFixture + sample_form_record: HumanInputFormRecord, + unbound_session_factory: sessionmaker[Session], + mocker: MockerFixture, ) -> None: - session_factory, _ = sqlite_session_factory repo = MagicMock(spec=HumanInputFormSubmissionRepository) definition = FormDefinition( form_content="hello", @@ -916,7 +892,7 @@ def test_submit_form_by_token_rejects_cross_tenant_file_list( expiration_time=sample_form_record.expiration_time, ) repo.get_by_token.return_value = dataclasses.replace(sample_form_record, definition=definition) - service = HumanInputService(session_factory, form_repository=repo) + service = HumanInputService(unbound_session_factory, form_repository=repo) mocker.patch("services.human_input_service.build_from_mappings", side_effect=ValueError("Invalid upload file")) with pytest.raises( diff --git a/api/tests/unit_tests/services/test_init_validation_service.py b/api/tests/unit_tests/services/test_init_validation_service.py new file mode 100644 index 00000000000..82fc787f78d --- /dev/null +++ b/api/tests/unit_tests/services/test_init_validation_service.py @@ -0,0 +1,101 @@ +"""Tests for initialization validation policy without Flask or persistence.""" + +from unittest.mock import Mock, create_autospec + +import pytest + +from services.init_validation_service import ( + AlreadyInitializedError, + InitValidationService, + InitValidationState, + InvalidInitializationPasswordError, +) + + +@pytest.fixture +def state() -> Mock: + return create_autospec(InitValidationState, instance=True, spec_set=True) + + +def test_status_is_valid_when_validation_is_not_required(state: Mock) -> None: + service = InitValidationService(state=state, validation_required=False, expected_password="") + + assert service.is_validated(session_validated=False) is True + state.is_setup.assert_not_called() + + +def test_status_is_valid_when_browser_session_was_validated(state: Mock) -> None: + service = InitValidationService(state=state, validation_required=True, expected_password="expected") + + assert service.is_validated(session_validated=True) is True + state.is_setup.assert_not_called() + + +@pytest.mark.parametrize("setup_exists", [False, True]) +def test_status_falls_back_to_persisted_setup_state(state: Mock, setup_exists: bool) -> None: + state.is_setup.return_value = setup_exists + service = InitValidationService(state=state, validation_required=True, expected_password="expected") + + assert service.is_validated(session_validated=False) is setup_exists + state.is_setup.assert_called_once_with() + + +def test_password_validation_rejects_an_initialized_installation(state: Mock) -> None: + state.has_tenants.return_value = True + service = InitValidationService(state=state, validation_required=True, expected_password="expected") + + with pytest.raises(AlreadyInitializedError): + service.validate_password("expected") + + +def test_initialized_installation_takes_precedence_over_a_password_mismatch(state: Mock) -> None: + state.has_tenants.return_value = True + service = InitValidationService(state=state, validation_required=True, expected_password="expected") + + with pytest.raises(AlreadyInitializedError): + service.validate_password("wrong") + + +def test_password_validation_rejects_a_mismatch(state: Mock) -> None: + state.has_tenants.return_value = False + service = InitValidationService(state=state, validation_required=True, expected_password="expected") + + with pytest.raises(InvalidInitializationPasswordError): + service.validate_password("wrong") + + +def test_password_validation_accepts_a_match(state: Mock) -> None: + state.has_tenants.return_value = False + service = InitValidationService(state=state, validation_required=True, expected_password="expected") + + service.validate_password("expected") + + state.has_tenants.assert_called_once_with() + + +@pytest.mark.parametrize("expected_password", ["", "expected"]) +def test_password_validation_rejects_an_empty_password(state: Mock, expected_password: str) -> None: + state.has_tenants.return_value = False + service = InitValidationService( + state=state, + validation_required=bool(expected_password), + expected_password=expected_password, + ) + + with pytest.raises(InvalidInitializationPasswordError): + service.validate_password("") + + +def test_password_validation_rejects_a_password_when_no_password_is_configured(state: Mock) -> None: + state.has_tenants.return_value = False + service = InitValidationService(state=state, validation_required=False, expected_password="") + + with pytest.raises(InvalidInitializationPasswordError): + service.validate_password("unexpected") + + +def test_password_validation_accepts_a_unicode_password(state: Mock) -> None: + state.has_tenants.return_value = False + service = InitValidationService(state=state, validation_required=True, expected_password="pässwörd-🔐") + + service.validate_password("pässwörd-🔐") diff --git a/api/tests/unit_tests/services/test_message_service.py b/api/tests/unit_tests/services/test_message_service.py index dfdaf7b40e2..afc49c8865c 100644 --- a/api/tests/unit_tests/services/test_message_service.py +++ b/api/tests/unit_tests/services/test_message_service.py @@ -1,13 +1,35 @@ import json +from collections.abc import Iterator from datetime import datetime -from unittest.mock import MagicMock, patch +from decimal import Decimal +from unittest.mock import MagicMock import pytest +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, scoped_session +import models.model as model_module +import services.message_service as service_module +from core.app.entities.app_invoke_entities import InvokeFrom from graphon.model_runtime.entities.model_entities import ModelType -from libs.infinite_scroll_pagination import InfiniteScrollPagination -from models.enums import FeedbackFromSource, FeedbackRating -from models.model import App, AppMode, EndUser, Message +from models.account import Account, AccountStatus +from models.enums import ( + ConversationFromSource, + EndUserType, + FeedbackFromSource, + FeedbackRating, +) +from models.model import ( + App, + AppAnnotationSetting, + AppMode, + AppModelConfig, + Conversation, + EndUser, + Message, + MessageFeedback, +) +from repositories.sqlalchemy_execution_extra_content_repository import SQLAlchemyExecutionExtraContentRepository from services.errors.message import ( FirstMessageNotExistsError, LastMessageNotExistsError, @@ -16,1247 +38,819 @@ from services.errors.message import ( ) from services.message_service import MessageService, attach_message_extra_contents +SQLITE_MODELS = (Conversation, Message, MessageFeedback, AppModelConfig, AppAnnotationSetting) +pytestmark = [ + pytest.mark.usefixtures("sqlite_session"), + pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True), +] -class TestMessageServiceFactory: - """Factory class for creating test data and mock objects for message service tests.""" + +class _DatabaseBinding: + """Expose the SQLite engine and shared session through the production DB interface.""" + + engine: Engine + session: scoped_session[Session] + + def __init__(self, engine: Engine, session: scoped_session[Session]) -> None: + self.engine = engine + self.session = session + + +class MessageServiceTestDataFactory: + """Create real service inputs and persistent message-domain rows.""" @staticmethod - def create_app_mock( + def create_app( app_id: str = "app-123", - mode: str = AppMode.ADVANCED_CHAT.value, - name: str = "Test App", - ) -> MagicMock: - """Create a mock App object.""" - app = MagicMock(spec=App) - app.id = app_id - app.mode = mode - app.name = name - return app + mode: AppMode = AppMode.ADVANCED_CHAT, + tenant_id: str = "tenant-123", + ) -> App: + return App( + id=app_id, + tenant_id=tenant_id, + name="Test App", + description="", + mode=mode, + enable_site=True, + enable_api=True, + max_active_requests=0, + ) @staticmethod - def create_end_user_mock( - user_id: str = "user-456", - session_id: str = "session-789", - ) -> MagicMock: - """Create a mock EndUser object.""" - user = MagicMock(spec=EndUser) - user.id = user_id - user.session_id = session_id - return user + def create_end_user(user_id: str = "user-456") -> EndUser: + return EndUser( + id=user_id, + tenant_id="tenant-123", + app_id="app-123", + type=EndUserType.SERVICE_API, + session_id="session-789", + ) @staticmethod - def create_conversation_mock( + def create_account(user_id: str = "account-123") -> Account: + account = Account(name="Admin", email="admin@example.com", status=AccountStatus.ACTIVE) + account.id = user_id + return account + + @staticmethod + def create_conversation( conversation_id: str = "conv-001", app_id: str = "app-123", - ) -> MagicMock: - """Create a mock Conversation object.""" - conversation = MagicMock() - conversation.id = conversation_id - conversation.app_id = app_id + *, + app_model_config_id: str | None = None, + override_model_configs: str | None = None, + ) -> Conversation: + conversation = Conversation( + id=conversation_id, + app_id=app_id, + app_model_config_id=app_model_config_id, + override_model_configs=override_model_configs, + mode=AppMode.CHAT, + name="Test conversation", + status="normal", + from_source=ConversationFromSource.API, + from_end_user_id="user-456", + ) + conversation._inputs = {} return conversation @staticmethod - def create_message_mock( + def create_message( message_id: str = "msg-001", conversation_id: str = "conv-001", - query: str = "What is AI?", - answer: str = "AI stands for Artificial Intelligence.", + app_id: str = "app-123", + *, created_at: datetime | None = None, - ) -> MagicMock: - """Create a mock Message object.""" - message = MagicMock(spec=Message) - message.id = message_id - message.conversation_id = conversation_id - message.query = query - message.answer = answer - message.created_at = created_at or datetime.now() - message.user_feedback_with_session.return_value = None - message.admin_feedback_with_session.return_value = None + from_source: ConversationFromSource = ConversationFromSource.API, + from_end_user_id: str | None = "user-456", + from_account_id: str | None = None, + ) -> Message: + message = Message( + id=message_id, + app_id=app_id, + conversation_id=conversation_id, + query="What is AI?", + message={"role": "user", "content": "What is AI?"}, + answer="AI stands for Artificial Intelligence.", + message_unit_price=Decimal("0.0001"), + answer_unit_price=Decimal("0.0002"), + currency="USD", + from_source=from_source, + from_end_user_id=from_end_user_id, + from_account_id=from_account_id, + ) + message._inputs = {} + timestamp = created_at or datetime.now() + message.created_at = timestamp + message.updated_at = timestamp return message + @staticmethod + def create_feedback( + feedback_id: str, + message: Message, + *, + source: FeedbackFromSource, + rating: FeedbackRating = FeedbackRating.LIKE, + ) -> MessageFeedback: + feedback = MessageFeedback( + app_id=message.app_id, + conversation_id=message.conversation_id, + message_id=message.id, + rating=rating, + from_source=source, + from_end_user_id="user-456" if source == FeedbackFromSource.USER else None, + from_account_id="account-123" if source == FeedbackFromSource.ADMIN else None, + ) + feedback.id = feedback_id + return feedback + + +@pytest.fixture +def factory() -> MessageServiceTestDataFactory: + return MessageServiceTestDataFactory() + + +@pytest.fixture(autouse=True) +def database_boundaries( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, +) -> Iterator[None]: + """Bind global model properties and service-owned factories to the shared SQLite session.""" + sessions = scoped_session(lambda: sqlite_session) + database = _DatabaseBinding(engine=sqlite_engine, session=sessions) + monkeypatch.setattr(service_module, "db", database) + monkeypatch.setattr(model_module, "db", database) + try: + yield + finally: + sessions.remove() + + +@pytest.fixture +def empty_extra_content_repository(monkeypatch: pytest.MonkeyPatch) -> MagicMock: + repository = MagicMock() + repository.get_by_message_ids.side_effect = lambda message_ids: [[] for _ in message_ids] + monkeypatch.setattr(service_module, "_create_execution_extra_content_repository", lambda: repository) + return repository + + +def _persist(session: Session, *records: object) -> None: + session.add_all(records) + session.commit() + + +def _patch_conversation(monkeypatch: pytest.MonkeyPatch, conversation: Conversation) -> MagicMock: + get_conversation = MagicMock(return_value=conversation) + monkeypatch.setattr(service_module.ConversationService, "get_conversation", get_conversation) + return get_conversation + class TestMessageServicePaginationByFirstId: - """ - Unit tests for MessageService.pagination_by_first_id method. + """Verify cursor pagination using persisted message timestamps and IDs.""" - This test suite covers: - - Basic pagination with and without first_id - - Order handling (asc/desc) - - Edge cases (no user, no conversation, invalid first_id) - - Has_more flag logic - """ - - @pytest.fixture - def factory(self): - """Provide test data factory.""" - return TestMessageServiceFactory() - - # Test 01: No user provided - def test_pagination_by_first_id_no_user(self, factory: TestMessageServiceFactory): - """Test pagination returns empty result when no user is provided.""" - # Arrange - app = factory.create_app_mock() - - # Act + @pytest.mark.parametrize(("user", "conversation_id"), [(None, "conv-001"), ("end_user", "")]) + def test_early_return( + self, + user: str | None, + conversation_id: str, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: result = MessageService.pagination_by_first_id( - app_model=app, - user=None, - conversation_id="conv-001", + app_model=factory.create_app(), + user=factory.create_end_user() if user else None, + conversation_id=conversation_id, first_id=None, limit=10, - session=MagicMock(), + session=sqlite_session, ) - # Assert - assert isinstance(result, InfiniteScrollPagination) assert result.data == [] assert result.limit == 10 assert result.has_more is False - # Test 02: No conversation_id provided - def test_pagination_by_first_id_no_conversation(self, factory: TestMessageServiceFactory): - """Test pagination returns empty result when no conversation_id is provided.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - - # Act - result = MessageService.pagination_by_first_id( - app_model=app, - user=user, - conversation_id="", - first_id=None, - limit=10, - session=MagicMock(), - ) - - # Assert - assert isinstance(result, InfiniteScrollPagination) - assert result.data == [] - assert result.limit == 10 - assert result.has_more is False - - # Test 03: Basic pagination without first_id (desc order) - @patch("services.message_service._create_execution_extra_content_repository") - @patch("services.message_service.db") - @patch("services.message_service.ConversationService") - def test_pagination_by_first_id_without_first_id_desc( - self, mock_conversation_service, mock_db, mock_create_repo, factory: TestMessageServiceFactory - ): - """Test basic pagination without first_id in descending order.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - conversation = factory.create_conversation_mock() - - mock_conversation_service.get_conversation.return_value = conversation - - # Create 5 messages + @pytest.mark.parametrize( + ("order", "expected_ids"), + [ + ("desc", ["msg-004", "msg-003", "msg-002", "msg-001", "msg-000"]), + ("asc", ["msg-000", "msg-001", "msg-002", "msg-003", "msg-004"]), + ], + ) + def test_orders_persisted_messages( + self, + order: str, + expected_ids: list[str], + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + empty_extra_content_repository: MagicMock, + ) -> None: + conversation = factory.create_conversation() messages = [ - factory.create_message_mock( - message_id=f"msg-{i:03d}", - created_at=datetime(2024, 1, 1, 12, i), - ) - for i in range(5) + factory.create_message(f"msg-{index:03d}", created_at=datetime(2024, 1, 1, 12, index)) for index in range(5) ] + _persist(sqlite_session, conversation, *messages) + _patch_conversation(monkeypatch, conversation) - mock_db.session.scalars.return_value.all.return_value = messages - - # Act result = MessageService.pagination_by_first_id( - app_model=app, - user=user, - conversation_id="conv-001", + app_model=factory.create_app(), + user=factory.create_end_user(), + conversation_id=conversation.id, first_id=None, limit=10, - order="desc", - session=mock_db.session, + order=order, + session=sqlite_session, ) - # Assert - assert len(result.data) == 5 + assert [message.id for message in result.data] == expected_ids assert result.has_more is False - assert result.limit == 10 - # Messages should remain in desc order (not reversed) - assert result.data[0].id == "msg-000" - # Test 04: Basic pagination without first_id (asc order) - @patch("services.message_service._create_execution_extra_content_repository") - @patch("services.message_service.db") - @patch("services.message_service.ConversationService") - def test_pagination_by_first_id_without_first_id_asc( - self, mock_conversation_service, mock_db, mock_create_repo, factory: TestMessageServiceFactory - ): - """Test basic pagination without first_id in ascending order.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - conversation = factory.create_conversation_mock() - - mock_conversation_service.get_conversation.return_value = conversation - - # Create 5 messages (returned in desc order from DB) + def test_first_id_excludes_cursor_and_newer_messages( + self, + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + empty_extra_content_repository: MagicMock, + ) -> None: + conversation = factory.create_conversation() messages = [ - factory.create_message_mock( - message_id=f"msg-{i:03d}", - created_at=datetime(2024, 1, 1, 12, 4 - i), # Descending timestamps - ) - for i in range(5) + factory.create_message(f"msg-{index:03d}", created_at=datetime(2024, 1, 1, 12, index)) for index in range(7) ] + _persist(sqlite_session, conversation, *messages) + _patch_conversation(monkeypatch, conversation) - mock_db.session.scalars.return_value.all.return_value = messages - - # Act result = MessageService.pagination_by_first_id( - app_model=app, - user=user, - conversation_id="conv-001", - first_id=None, - limit=10, - order="asc", - session=mock_db.session, - ) - - # Assert - assert len(result.data) == 5 - assert result.has_more is False - # Messages should be reversed to asc order - assert result.data[0].id == "msg-004" - assert result.data[4].id == "msg-000" - - # Test 05: Pagination with first_id - @patch("services.message_service._create_execution_extra_content_repository") - @patch("services.message_service.db") - @patch("services.message_service.ConversationService") - def test_pagination_by_first_id_with_first_id( - self, mock_conversation_service, mock_db, mock_create_repo, factory: TestMessageServiceFactory - ): - """Test pagination with first_id to get messages before a specific message.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - conversation = factory.create_conversation_mock() - - mock_conversation_service.get_conversation.return_value = conversation - - first_message = factory.create_message_mock( - message_id="msg-005", - created_at=datetime(2024, 1, 1, 12, 5), - ) - - # Messages before first_message - history_messages = [ - factory.create_message_mock( - message_id=f"msg-{i:03d}", - created_at=datetime(2024, 1, 1, 12, i), - ) - for i in range(5) - ] - - mock_db.session.scalar.return_value = first_message - mock_db.session.scalars.return_value.all.return_value = history_messages - - # Act - result = MessageService.pagination_by_first_id( - app_model=app, - user=user, - conversation_id="conv-001", + app_model=factory.create_app(), + user=factory.create_end_user(), + conversation_id=conversation.id, first_id="msg-005", limit=10, order="desc", - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert len(result.data) == 5 - assert result.has_more is False + assert [message.id for message in result.data] == [f"msg-{index:03d}" for index in range(4, -1, -1)] - # Test 06: First message not found - @patch("services.message_service.db") - @patch("services.message_service.ConversationService") - def test_pagination_by_first_id_first_message_not_exists( - self, mock_conversation_service, mock_db, factory: TestMessageServiceFactory - ): - """Test error handling when first_id doesn't exist.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - conversation = factory.create_conversation_mock() + def test_missing_first_id_raises( + self, + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: + conversation = factory.create_conversation() + _persist(sqlite_session, conversation) + _patch_conversation(monkeypatch, conversation) - mock_conversation_service.get_conversation.return_value = conversation - - mock_db.session.scalar.return_value = None # Message not found - - # Act & Assert with pytest.raises(FirstMessageNotExistsError): MessageService.pagination_by_first_id( - app_model=app, - user=user, - conversation_id="conv-001", - first_id="nonexistent-msg", + app_model=factory.create_app(), + user=factory.create_end_user(), + conversation_id=conversation.id, + first_id="missing", limit=10, - session=mock_db.session, + session=sqlite_session, ) - # Test 07: Has_more flag when results exceed limit - @patch("services.message_service._create_execution_extra_content_repository") - @patch("services.message_service.db") - @patch("services.message_service.ConversationService") - def test_pagination_by_first_id_has_more_true( - self, mock_conversation_service, mock_db, mock_create_repo, factory: TestMessageServiceFactory - ): - """Test has_more flag is True when results exceed limit.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - conversation = factory.create_conversation_mock() - - mock_conversation_service.get_conversation.return_value = conversation - - # Create limit+1 messages (11 messages for limit=10) + def test_has_more_trims_oldest_extra_row( + self, + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + empty_extra_content_repository: MagicMock, + ) -> None: + conversation = factory.create_conversation() messages = [ - factory.create_message_mock( - message_id=f"msg-{i:03d}", - created_at=datetime(2024, 1, 1, 12, i), - ) - for i in range(11) + factory.create_message(f"msg-{index:03d}", created_at=datetime(2024, 1, 1, 12, index)) + for index in range(11) ] + _persist(sqlite_session, conversation, *messages) + _patch_conversation(monkeypatch, conversation) - mock_db.session.scalars.return_value.all.return_value = messages - - # Act result = MessageService.pagination_by_first_id( - app_model=app, - user=user, - conversation_id="conv-001", + app_model=factory.create_app(), + user=factory.create_end_user(), + conversation_id=conversation.id, first_id=None, limit=10, - session=mock_db.session, + order="desc", + session=sqlite_session, ) - # Assert - assert len(result.data) == 10 # Last message trimmed + assert len(result.data) == 10 assert result.has_more is True - assert result.limit == 10 + assert result.data[-1].id == "msg-001" - # Test 08: Empty conversation - @patch("services.message_service.db") - @patch("services.message_service.ConversationService") - def test_pagination_by_first_id_empty_conversation( - self, mock_conversation_service, mock_db, factory: TestMessageServiceFactory - ): - """Test pagination with conversation that has no messages.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - conversation = factory.create_conversation_mock() + def test_empty_conversation( + self, + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + empty_extra_content_repository: MagicMock, + ) -> None: + conversation = factory.create_conversation() + _persist(sqlite_session, conversation) + _patch_conversation(monkeypatch, conversation) - mock_conversation_service.get_conversation.return_value = conversation - - mock_db.session.scalars.return_value.all.return_value = [] - - # Act result = MessageService.pagination_by_first_id( - app_model=app, - user=user, - conversation_id="conv-001", + app_model=factory.create_app(), + user=factory.create_end_user(), + conversation_id=conversation.id, first_id=None, limit=10, - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert len(result.data) == 0 + assert result.data == [] assert result.has_more is False - assert result.limit == 10 class TestMessageServicePaginationByLastId: - """ - Unit tests for MessageService.pagination_by_last_id method. + """Verify reverse cursor, conversation, and include-ID filtering.""" - This test suite covers: - - Basic pagination with and without last_id - - Conversation filtering - - Include_ids filtering - - Edge cases (no user, invalid last_id) - """ - - @pytest.fixture - def factory(self): - """Provide test data factory.""" - return TestMessageServiceFactory() - - # Test 09: No user provided - def test_pagination_by_last_id_no_user(self, factory: TestMessageServiceFactory): - """Test pagination returns empty result when no user is provided.""" - # Arrange - app = factory.create_app_mock() - - # Act + def test_no_user(self, factory: MessageServiceTestDataFactory, sqlite_session: Session) -> None: result = MessageService.pagination_by_last_id( - app_model=app, - user=None, - last_id=None, - limit=10, - session=MagicMock(), + app_model=factory.create_app(), user=None, last_id=None, limit=10, session=sqlite_session ) - - # Assert - assert isinstance(result, InfiniteScrollPagination) assert result.data == [] assert result.limit == 10 assert result.has_more is False - # Test 10: Basic pagination without last_id - @patch("services.message_service.db") - def test_pagination_by_last_id_without_last_id(self, mock_db, factory: TestMessageServiceFactory): - """Test basic pagination without last_id.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - + def test_without_last_id(self, factory: MessageServiceTestDataFactory, sqlite_session: Session) -> None: messages = [ - factory.create_message_mock( - message_id=f"msg-{i:03d}", - created_at=datetime(2024, 1, 1, 12, i), - ) - for i in range(5) + factory.create_message(f"msg-{index:03d}", created_at=datetime(2024, 1, 1, 12, index)) for index in range(5) ] + _persist(sqlite_session, *messages) - mock_db.session.scalars.return_value.all.return_value = messages - - # Act result = MessageService.pagination_by_last_id( - app_model=app, - user=user, + app_model=factory.create_app(), + user=factory.create_end_user(), last_id=None, limit=10, - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert len(result.data) == 5 + assert [message.id for message in result.data] == [f"msg-{index:03d}" for index in range(4, -1, -1)] assert result.has_more is False - assert result.limit == 10 - # Test 11: Pagination with last_id - @patch("services.message_service.db") - def test_pagination_by_last_id_with_last_id(self, mock_db, factory: TestMessageServiceFactory): - """Test pagination with last_id to get messages after a specific message.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - - last_message = factory.create_message_mock( - message_id="msg-005", - created_at=datetime(2024, 1, 1, 12, 5), - ) - - # Messages after last_message - new_messages = [ - factory.create_message_mock( - message_id=f"msg-{i:03d}", - created_at=datetime(2024, 1, 1, 12, i), - ) - for i in range(6, 10) + def test_last_id_returns_older_rows(self, factory: MessageServiceTestDataFactory, sqlite_session: Session) -> None: + messages = [ + factory.create_message(f"msg-{index:03d}", created_at=datetime(2024, 1, 1, 12, index)) for index in range(7) ] + _persist(sqlite_session, *messages) - mock_db.session.scalar.return_value = last_message - mock_db.session.scalars.return_value.all.return_value = new_messages - - # Act result = MessageService.pagination_by_last_id( - app_model=app, - user=user, + app_model=factory.create_app(), + user=factory.create_end_user(), last_id="msg-005", limit=10, - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert len(result.data) == 4 - assert result.has_more is False + assert [message.id for message in result.data] == [f"msg-{index:03d}" for index in range(4, -1, -1)] - # Test 12: Last message not found - @patch("services.message_service.db") - def test_pagination_by_last_id_last_message_not_exists(self, mock_db, factory: TestMessageServiceFactory): - """Test error handling when last_id doesn't exist.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - - mock_db.session.scalar.return_value = None # Message not found - - # Act & Assert + def test_missing_last_id_raises(self, factory: MessageServiceTestDataFactory, sqlite_session: Session) -> None: with pytest.raises(LastMessageNotExistsError): MessageService.pagination_by_last_id( - app_model=app, - user=user, - last_id="nonexistent-msg", + app_model=factory.create_app(), + user=factory.create_end_user(), + last_id="missing", limit=10, - session=mock_db.session, + session=sqlite_session, ) - # Test 13: Pagination with conversation_id filter - @patch("services.message_service.ConversationService") - @patch("services.message_service.db") - def test_pagination_by_last_id_with_conversation_filter( - self, mock_db, mock_conversation_service, factory: TestMessageServiceFactory - ): - """Test pagination filtered by conversation_id.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - conversation = factory.create_conversation_mock(conversation_id="conv-001") + def test_conversation_filter( + self, + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: + conversation = factory.create_conversation() + other_conversation = factory.create_conversation("conv-002") + matching = factory.create_message("matching", conversation_id=conversation.id) + excluded = factory.create_message("excluded", conversation_id=other_conversation.id) + _persist(sqlite_session, conversation, other_conversation, matching, excluded) + get_conversation = _patch_conversation(monkeypatch, conversation) - mock_conversation_service.get_conversation.return_value = conversation - - messages = [ - factory.create_message_mock( - message_id=f"msg-{i:03d}", - conversation_id="conv-001", - created_at=datetime(2024, 1, 1, 12, i), - ) - for i in range(5) - ] - - mock_db.session.scalars.return_value.all.return_value = messages - - # Act result = MessageService.pagination_by_last_id( - app_model=app, - user=user, + app_model=factory.create_app(), + user=factory.create_end_user(), last_id=None, limit=10, - conversation_id="conv-001", - session=mock_db.session, + conversation_id=conversation.id, + session=sqlite_session, ) - # Assert - assert len(result.data) == 5 - assert result.has_more is False - mock_conversation_service.get_conversation.assert_called_once() + assert [message.id for message in result.data] == [matching.id] + get_conversation.assert_called_once() - # Test 14: Pagination with include_ids filter - @patch("services.message_service.db") - def test_pagination_by_last_id_with_include_ids(self, mock_db, factory: TestMessageServiceFactory): - """Test pagination filtered by include_ids.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - - # Only messages with IDs in include_ids should be returned + def test_include_ids_filter(self, factory: MessageServiceTestDataFactory, sqlite_session: Session) -> None: messages = [ - factory.create_message_mock(message_id="msg-001"), - factory.create_message_mock(message_id="msg-003"), + factory.create_message(f"msg-{index:03d}", created_at=datetime(2024, 1, 1, 12, index)) for index in range(4) ] + _persist(sqlite_session, *messages) - mock_db.session.scalars.return_value.all.return_value = messages - - # Act result = MessageService.pagination_by_last_id( - app_model=app, - user=user, + app_model=factory.create_app(), + user=factory.create_end_user(), last_id=None, limit=10, include_ids=["msg-001", "msg-003"], - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert len(result.data) == 2 - assert result.data[0].id == "msg-001" - assert result.data[1].id == "msg-003" + assert [message.id for message in result.data] == ["msg-003", "msg-001"] - # Test 15: Has_more flag when results exceed limit - @patch("services.message_service.db") - def test_pagination_by_last_id_has_more_true(self, mock_db, factory: TestMessageServiceFactory): - """Test has_more flag is True when results exceed limit.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - - # Create limit+1 messages (11 messages for limit=10) + def test_has_more(self, factory: MessageServiceTestDataFactory, sqlite_session: Session) -> None: messages = [ - factory.create_message_mock( - message_id=f"msg-{i:03d}", - created_at=datetime(2024, 1, 1, 12, i), - ) - for i in range(11) + factory.create_message(f"msg-{index:03d}", created_at=datetime(2024, 1, 1, 12, index)) + for index in range(11) ] + _persist(sqlite_session, *messages) - mock_db.session.scalars.return_value.all.return_value = messages - - # Act result = MessageService.pagination_by_last_id( - app_model=app, - user=user, + app_model=factory.create_app(), + user=factory.create_end_user(), last_id=None, limit=10, - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert len(result.data) == 10 # Last message trimmed + assert len(result.data) == 10 assert result.has_more is True - assert result.limit == 10 class TestMessageServiceUtilities: - """Unit tests for MessageService module-level utility functions.""" - - @pytest.fixture - def factory(self): - """Provide test data factory.""" - return TestMessageServiceFactory() - - # Test 16: attach_message_extra_contents with empty list - def test_attach_message_extra_contents_empty(self): - """Test attach_message_extra_contents with empty list does nothing.""" - # Act & Assert (should not raise error) + def test_attach_message_extra_contents_empty(self) -> None: attach_message_extra_contents([]) - # Test 17: attach_message_extra_contents with messages - @patch("services.message_service._create_execution_extra_content_repository") - def test_attach_message_extra_contents_with_messages(self, mock_create_repo, factory: TestMessageServiceFactory): - """Test attach_message_extra_contents correctly attaches content.""" - # Arrange - messages = [factory.create_message_mock(message_id="msg-1"), factory.create_message_mock(message_id="msg-2")] + def test_attach_message_extra_contents( + self, + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + ) -> None: + messages = [factory.create_message("msg-1"), factory.create_message("msg-2")] + content_one = MagicMock() + content_one.model_dump.return_value = {"key": "value1"} + content_two = MagicMock() + content_two.model_dump.return_value = {"key": "value2"} + repository = MagicMock() + repository.get_by_message_ids.return_value = [[content_one], [content_two]] + monkeypatch.setattr(service_module, "_create_execution_extra_content_repository", lambda: repository) - mock_repo = MagicMock() - mock_create_repo.return_value = mock_repo - - # Mock extra content models - mock_content1 = MagicMock() - mock_content1.model_dump.return_value = {"key": "value1"} - mock_content2 = MagicMock() - mock_content2.model_dump.return_value = {"key": "value2"} - - mock_repo.get_by_message_ids.return_value = [[mock_content1], [mock_content2]] - - # Act attach_message_extra_contents(messages) - # Assert - mock_repo.get_by_message_ids.assert_called_once_with(["msg-1", "msg-2"]) - messages[0].set_extra_contents.assert_called_once_with([{"key": "value1"}]) - messages[1].set_extra_contents.assert_called_once_with([{"key": "value2"}]) + assert messages[0].extra_contents == [{"key": "value1"}] + assert messages[1].extra_contents == [{"key": "value2"}] - # Test 18: attach_message_extra_contents with index out of bounds - @patch("services.message_service._create_execution_extra_content_repository") - def test_attach_message_extra_contents_index_out_of_bounds( - self, mock_create_repo, factory: TestMessageServiceFactory - ): - """Test attach_message_extra_contents handles missing content lists.""" - # Arrange - messages = [factory.create_message_mock(message_id="msg-1")] + def test_attach_message_extra_contents_missing_list( + self, + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + ) -> None: + message = factory.create_message("msg-1") + repository = MagicMock() + repository.get_by_message_ids.return_value = [] + monkeypatch.setattr(service_module, "_create_execution_extra_content_repository", lambda: repository) - mock_repo = MagicMock() - mock_create_repo.return_value = mock_repo - mock_repo.get_by_message_ids.return_value = [] # Empty returned list + attach_message_extra_contents([message]) - # Act - attach_message_extra_contents(messages) + assert message.extra_contents == [] - # Assert - messages[0].set_extra_contents.assert_called_once_with([]) + def test_create_execution_extra_content_repository_uses_sqlite_factory(self, sqlite_engine: Engine) -> None: + repository = service_module._create_execution_extra_content_repository() - # Test 19: _create_execution_extra_content_repository - @patch("services.message_service.db") - @patch("services.message_service.sessionmaker") - @patch("services.message_service.SQLAlchemyExecutionExtraContentRepository") - def test_create_execution_extra_content_repository(self, mock_repo_class, mock_sessionmaker, mock_db): - """Test _create_execution_extra_content_repository creates expected repository.""" - from services.message_service import _create_execution_extra_content_repository - - # Act - _create_execution_extra_content_repository() - - # Assert - mock_sessionmaker.assert_called_once() - mock_repo_class.assert_called_once() + assert isinstance(repository, SQLAlchemyExecutionExtraContentRepository) + assert repository._session_maker.kw["bind"] is sqlite_engine + with repository._session_maker() as session: + assert isinstance(session, Session) class TestMessageServiceGetMessage: - """Unit tests for MessageService.get_message method.""" + @pytest.mark.parametrize("actor", ["end_user", "account"]) + def test_identity_scoped_success( + self, + actor: str, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: + if actor == "end_user": + user: Account | EndUser = factory.create_end_user("end-user-123") + message = factory.create_message( + "msg-123", from_end_user_id=user.id, from_account_id=None, from_source=ConversationFromSource.API + ) + else: + user = factory.create_account("account-123") + message = factory.create_message( + "msg-123", + from_end_user_id=None, + from_account_id=user.id, + from_source=ConversationFromSource.CONSOLE, + ) + distractor = factory.create_message("wrong-app", app_id="app-456") + _persist(sqlite_session, message, distractor) - @pytest.fixture - def factory(self): - """Provide test data factory.""" - return TestMessageServiceFactory() + result = MessageService.get_message( + app_model=factory.create_app(), user=user, message_id=message.id, session=sqlite_session + ) - # Test 20: get_message success for EndUser - @patch("services.message_service.db") - def test_get_message_end_user_success(self, mock_db, factory: TestMessageServiceFactory): - """Test get_message returns message for EndUser.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock(user_id="end-user-123") - message = factory.create_message_mock() + assert result.id == message.id - mock_db.session.scalar.return_value = message - - # Act, - result = MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) - - # Assert - assert result == message - - # Test 21: get_message success for Account (Admin) - @patch("services.message_service.db") - def test_get_message_account_success(self, mock_db, factory: TestMessageServiceFactory): - """Test get_message returns message for Account.""" - # Arrange - from models import Account - - app = factory.create_app_mock() - user = MagicMock(spec=Account) - user.id = "account-123" - message = factory.create_message_mock() - - mock_db.session.scalar.return_value = message - - # Act, - result = MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) - - # Assert - assert result == message - - # Test 22: get_message not found - @patch("services.message_service.db") - def test_get_message_not_found(self, mock_db, factory: TestMessageServiceFactory): - """Test get_message raises MessageNotExistsError when not found.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - - mock_db.session.scalar.return_value = None - - # Act & Assert + def test_not_found(self, factory: MessageServiceTestDataFactory, sqlite_session: Session) -> None: with pytest.raises(MessageNotExistsError): - MessageService.get_message(app_model=app, user=user, message_id="msg-123", session=mock_db.session) + MessageService.get_message( + app_model=factory.create_app(), + user=factory.create_end_user(), + message_id="missing", + session=sqlite_session, + ) class TestMessageServiceFeedback: - """Unit tests for MessageService feedback-related methods.""" + def test_create_new_end_user_feedback( + self, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + sqlite_engine: Engine, + ) -> None: + user = factory.create_end_user() + message = factory.create_message("msg-123") + _persist(sqlite_session, message) - @pytest.fixture - def factory(self): - """Provide test data factory.""" - return TestMessageServiceFactory() - - # Test 23: create_feedback - new feedback for EndUser - @patch("services.message_service.db") - @patch.object(MessageService, "get_message") - def test_create_feedback_new_end_user(self, mock_get_message, mock_db, factory: TestMessageServiceFactory): - """Test creating new feedback for an end user.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - message = factory.create_message_mock() - message.user_feedback = None - message.user_feedback_with_session.return_value = None - mock_get_message.return_value = message - - # Act - result = MessageService.create_feedback( - app_model=app, - message_id="msg-123", + feedback = MessageService.create_feedback( + app_model=factory.create_app(), + message_id=message.id, user=user, rating=FeedbackRating.LIKE, content="Good answer", - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert result.rating == FeedbackRating.LIKE - assert result.content == "Good answer" - assert result.from_source == FeedbackFromSource.USER - mock_db.session.add.assert_called_once() - mock_db.session.commit.assert_called_once() + with Session(sqlite_engine) as verification_session: + persisted = verification_session.get(MessageFeedback, feedback.id) + assert persisted is not None + assert persisted.rating == FeedbackRating.LIKE + assert persisted.content == "Good answer" + assert persisted.from_source == FeedbackFromSource.USER - # Test 24: create_feedback - update feedback for Account - @patch("services.message_service.db") - @patch.object(MessageService, "get_message") - def test_create_feedback_update_account(self, mock_get_message, mock_db, factory: TestMessageServiceFactory): - """Test updating existing feedback for an account.""" - # Arrange - from models import Account, MessageFeedback + def test_update_account_feedback( + self, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + sqlite_engine: Engine, + ) -> None: + user = factory.create_account() + message = factory.create_message( + "msg-123", + from_source=ConversationFromSource.CONSOLE, + from_end_user_id=None, + from_account_id=user.id, + ) + feedback = factory.create_feedback("feedback-1", message, source=FeedbackFromSource.ADMIN) + _persist(sqlite_session, message, feedback) - app = factory.create_app_mock() - user = MagicMock(spec=Account) - user.id = "account-123" - message = factory.create_message_mock() - feedback = MagicMock(spec=MessageFeedback) - message.admin_feedback = feedback - message.admin_feedback_with_session.return_value = feedback - mock_get_message.return_value = message - - # Act result = MessageService.create_feedback( - app_model=app, - message_id="msg-123", + app_model=factory.create_app(), + message_id=message.id, user=user, rating=FeedbackRating.DISLIKE, content="Bad answer", - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert result == feedback - assert feedback.rating == FeedbackRating.DISLIKE - assert feedback.content == "Bad answer" - mock_db.session.commit.assert_called_once() + assert result.id == feedback.id + with Session(sqlite_engine) as verification_session: + persisted = verification_session.get(MessageFeedback, feedback.id) + assert persisted is not None + assert persisted.rating == FeedbackRating.DISLIKE + assert persisted.content == "Bad answer" - # Test 25: create_feedback - delete feedback (rating is None) - @patch("services.message_service.db") - @patch.object(MessageService, "get_message") - def test_create_feedback_delete(self, mock_get_message, mock_db, factory: TestMessageServiceFactory): - """Test deleting feedback by passing rating=None.""" - # Arrange - app = factory.create_app_mock() - user = factory.create_end_user_mock() - message = factory.create_message_mock() - feedback = MagicMock() - message.user_feedback = feedback - message.user_feedback_with_session.return_value = feedback - mock_get_message.return_value = message + def test_delete_feedback( + self, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + sqlite_engine: Engine, + ) -> None: + user = factory.create_end_user() + message = factory.create_message("msg-123") + feedback = factory.create_feedback("feedback-1", message, source=FeedbackFromSource.USER) + _persist(sqlite_session, message, feedback) - # Act - result = MessageService.create_feedback( - app_model=app, - message_id="msg-123", + MessageService.create_feedback( + app_model=factory.create_app(), + message_id=message.id, user=user, rating=None, content=None, - session=mock_db.session, + session=sqlite_session, ) - # Assert - assert result == feedback - mock_db.session.delete.assert_called_once_with(feedback) - mock_db.session.commit.assert_called_once() + with Session(sqlite_engine) as verification_session: + assert verification_session.get(MessageFeedback, feedback.id) is None - # Test 26: get_all_messages_feedbacks - @patch("services.message_service.db") - def test_get_all_messages_feedbacks(self, mock_db, factory: TestMessageServiceFactory): - """Test get_all_messages_feedbacks returns list of dicts.""" - # Arrange - app = factory.create_app_mock() - feedback = MagicMock() - feedback.to_dict.return_value = {"id": "fb-1"} + def test_get_all_feedbacks_is_app_scoped_and_paginated( + self, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: + message = factory.create_message("msg-123") + newest = factory.create_feedback("feedback-new", message, source=FeedbackFromSource.USER) + oldest = factory.create_feedback("feedback-old", message, source=FeedbackFromSource.USER) + other_message = factory.create_message("other-msg", app_id="app-456") + other_app = factory.create_feedback("feedback-other", other_message, source=FeedbackFromSource.USER) + newest.created_at = datetime(2024, 1, 2) + oldest.created_at = datetime(2024, 1, 1) + other_app.created_at = datetime(2024, 1, 3) + _persist(sqlite_session, newest, oldest, other_app) - mock_db.session.scalars.return_value.all.return_value = [feedback] + result = MessageService.get_all_messages_feedbacks( + app_model=factory.create_app(), page=1, limit=1, session=sqlite_session + ) - # Act, - result = MessageService.get_all_messages_feedbacks(app_model=app, page=1, limit=10, session=mock_db.session) - - # Assert - assert result == [{"id": "fb-1"}] + assert [record["id"] for record in result] == [newest.id] class TestMessageServiceSuggestedQuestions: - """Unit tests for MessageService.get_suggested_questions_after_answer method.""" + @staticmethod + def _chat_boundaries( + monkeypatch: pytest.MonkeyPatch, + conversation: Conversation, + ) -> tuple[MagicMock, MagicMock, MagicMock]: + message = MagicMock() + message.conversation_id = conversation.id + monkeypatch.setattr(service_module.MessageService, "get_message", MagicMock(return_value=message)) + monkeypatch.setattr( + service_module.ConversationService, "get_conversation", MagicMock(return_value=conversation) + ) + model_manager = MagicMock() + monkeypatch.setattr(service_module.ModelManager, "for_tenant", MagicMock(return_value=model_manager)) + memory = MagicMock() + memory.return_value.get_history_prompt_text.return_value = "histories" + monkeypatch.setattr(service_module, "TokenBufferMemory", memory) + llm_generator = MagicMock() + llm_generator.generate_suggested_questions_after_answer.return_value = ["Q1?"] + monkeypatch.setattr(service_module, "LLMGenerator", llm_generator) + monkeypatch.setattr(service_module, "TraceQueueManager", MagicMock()) + return model_manager, memory, llm_generator - @pytest.fixture - def factory(self): - """Provide test data factory.""" - return TestMessageServiceFactory() - - # Test 27: get_suggested_questions_after_answer - user is None - def test_get_suggested_questions_user_none(self, factory: TestMessageServiceFactory): - app = factory.create_app_mock() + def test_user_none(self, factory: MessageServiceTestDataFactory, sqlite_session: Session) -> None: with pytest.raises(ValueError, match="user cannot be None"): MessageService.get_suggested_questions_after_answer( - app_model=app, + app_model=factory.create_app(), user=None, message_id="msg-123", - invoke_from=MagicMock(), - session=MagicMock(), + invoke_from=InvokeFrom.WEB_APP, + session=sqlite_session, ) - # Test 28: get_suggested_questions_after_answer - Advanced Chat success - @patch("services.message_service.ModelManager.for_tenant") - @patch("services.message_service.WorkflowService") - @patch("services.message_service.AdvancedChatAppConfigManager") - @patch("services.message_service.TokenBufferMemory") - @patch("services.message_service.LLMGenerator") - @patch("services.message_service.TraceQueueManager") - @patch.object(MessageService, "get_message") - @patch("services.message_service.ConversationService") - def test_get_suggested_questions_advanced_chat_success( + def test_advanced_chat_success( self, - mock_conversation_service, - mock_get_message, - mock_trace_manager, - mock_llm_gen, - mock_memory, - mock_config_manager, - mock_workflow_service, - mock_model_manager, - factory: TestMessageServiceFactory, - ): - """Test successful suggested questions generation in Advanced Chat mode.""" - from core.app.entities.app_invoke_entities import InvokeFrom - - # Arrange - app = factory.create_app_mock(mode=AppMode.ADVANCED_CHAT.value) - user = factory.create_end_user_mock() - message = factory.create_message_mock() - mock_get_message.return_value = message - + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: + conversation = factory.create_conversation() + _, _, llm_generator = self._chat_boundaries(monkeypatch, conversation) workflow = MagicMock() - mock_workflow_service.return_value.get_published_workflow.return_value = workflow + workflow.features_dict = {"suggested_questions_after_answer": {"enabled": True}} + workflow_service = MagicMock() + workflow_service.return_value.get_published_workflow.return_value = workflow + monkeypatch.setattr(service_module, "WorkflowService", workflow_service) + app_config_manager = MagicMock() + app_config_manager.get_app_config.return_value.additional_features.suggested_questions_after_answer = True + monkeypatch.setattr(service_module, "AdvancedChatAppConfigManager", app_config_manager) - app_config = MagicMock() - app_config.additional_features.suggested_questions_after_answer = True - mock_config_manager.get_app_config.return_value = app_config - - mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] - - # Act result = MessageService.get_suggested_questions_after_answer( - app_model=app, - user=user, + app_model=factory.create_app(mode=AppMode.ADVANCED_CHAT), + user=factory.create_end_user(), message_id="msg-123", invoke_from=InvokeFrom.WEB_APP, - session=MagicMock(), + session=sqlite_session, ) - # Assert assert result == ["Q1?"] - mock_workflow_service.return_value.get_published_workflow.assert_called_once() - mock_llm_gen.generate_suggested_questions_after_answer.assert_called_once() + llm_generator.generate_suggested_questions_after_answer.assert_called_once() - # Test 29: get_suggested_questions_after_answer - Chat app success (no override) - @patch("services.message_service.db") - @patch("services.message_service.ModelManager.for_tenant") - @patch("services.message_service.TokenBufferMemory") - @patch("services.message_service.LLMGenerator") - @patch("services.message_service.TraceQueueManager") - @patch.object(MessageService, "get_message") - @patch("services.message_service.ConversationService") - def test_get_suggested_questions_chat_app_success( + @pytest.mark.parametrize( + ("config", "expected_prompt", "expected_model"), + [ + ({"enabled": True}, None, None), + ( + { + "enabled": True, + "prompt": "custom prompt", + "model": { + "provider": "openai", + "name": "gpt-4o-mini", + "completion_params": {"max_tokens": 2048, "temperature": 0.1}, + }, + }, + "custom prompt", + { + "provider": "openai", + "name": "gpt-4o-mini", + "completion_params": {"max_tokens": 2048, "temperature": 0.1}, + }, + ), + ( + {"enabled": True, "model": {"provider": "openai", "name": "invalid-model"}}, + None, + {"provider": "openai", "name": "invalid-model"}, + ), + ], + ) + def test_chat_app_uses_persisted_model_config( self, - mock_conversation_service: MagicMock, - mock_get_message: MagicMock, - mock_trace_manager: MagicMock, - mock_llm_gen: MagicMock, - mock_memory: MagicMock, - mock_model_manager: MagicMock, - mock_db: MagicMock, - factory: TestMessageServiceFactory, - ): - """Test successful suggested questions generation in basic Chat mode.""" - # Arrange - app = factory.create_app_mock(mode=AppMode.CHAT) - user = factory.create_end_user_mock() - message = factory.create_message_mock() - mock_get_message.return_value = message - - conversation = MagicMock() - conversation.override_model_configs = None - mock_conversation_service.get_conversation.return_value = conversation - - app_model_config = MagicMock() - app_model_config.suggested_questions_after_answer_dict = {"enabled": True} - app_model_config.model_dict = {"provider": "openai", "name": "gpt-4"} - - mock_db.session.scalar.return_value = app_model_config - - mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] - - # Act - result = MessageService.get_suggested_questions_after_answer( - app_model=app, - user=user, - message_id="msg-123", - invoke_from=MagicMock(), - session=mock_db.session, + config: dict[str, object], + expected_prompt: str | None, + expected_model: dict[str, object] | None, + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: + app_model_config = AppModelConfig( + app_id="app-123", + suggested_questions_after_answer=json.dumps(config), ) - - # Assert - assert result == ["Q1?"] - mock_llm_gen.generate_suggested_questions_after_answer.assert_called_once() - - @patch("services.message_service.db") - @patch("services.message_service.ModelManager.for_tenant") - @patch("services.message_service.TokenBufferMemory") - @patch("services.message_service.LLMGenerator") - @patch("services.message_service.TraceQueueManager") - @patch.object(MessageService, "get_message") - @patch("services.message_service.ConversationService") - def test_get_suggested_questions_chat_app_uses_frontend_model_and_prompt( - self, - mock_conversation_service: MagicMock, - mock_get_message: MagicMock, - mock_trace_manager: MagicMock, - mock_llm_gen: MagicMock, - mock_memory: MagicMock, - mock_model_manager: MagicMock, - mock_db: MagicMock, - factory: TestMessageServiceFactory, - ): - """Test suggested question generation uses frontend configured model and prompt.""" - from core.app.entities.app_invoke_entities import InvokeFrom - - app = factory.create_app_mock(mode=AppMode.CHAT) - app.tenant_id = "tenant-123" - user = factory.create_end_user_mock() - message = factory.create_message_mock() - mock_get_message.return_value = message - - conversation = MagicMock() - conversation.override_model_configs = None - mock_conversation_service.get_conversation.return_value = conversation - - app_model_config = MagicMock() - app_model_config.suggested_questions_after_answer_dict = { - "enabled": True, - "prompt": "custom prompt", - "model": { - "provider": "openai", - "name": "gpt-4o-mini", - "completion_params": {"max_tokens": 2048, "temperature": 0.1}, - }, - } - mock_db.session.scalar.return_value = app_model_config - - mock_memory.return_value.get_history_prompt_text.return_value = "histories" - mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] + app_model_config.id = "config-1" + conversation = factory.create_conversation(app_model_config_id=app_model_config.id) + _persist(sqlite_session, app_model_config) + model_manager, memory, llm_generator = self._chat_boundaries(monkeypatch, conversation) result = MessageService.get_suggested_questions_after_answer( - app_model=app, - user=user, + app_model=factory.create_app(mode=AppMode.CHAT), + user=factory.create_end_user(), message_id="msg-123", invoke_from=InvokeFrom.WEB_APP, - session=mock_db.session, + session=sqlite_session, ) assert result == ["Q1?"] - mock_model_manager.return_value.get_default_model_instance.assert_called_once_with( - tenant_id="tenant-123", - model_type=ModelType.LLM, + model_manager.get_default_model_instance.assert_called_once_with( + tenant_id="tenant-123", model_type=ModelType.LLM ) - mock_memory.assert_called_once_with( + memory.assert_called_once_with( conversation=conversation, - model_instance=mock_model_manager.return_value.get_default_model_instance.return_value, + model_instance=model_manager.get_default_model_instance.return_value, ) - mock_llm_gen.generate_suggested_questions_after_answer.assert_called_once_with( + llm_generator.generate_suggested_questions_after_answer.assert_called_once_with( tenant_id="tenant-123", histories="histories", - instruction_prompt="custom prompt", - model_config={ - "provider": "openai", - "name": "gpt-4o-mini", - "completion_params": {"max_tokens": 2048, "temperature": 0.1}, - }, + instruction_prompt=expected_prompt, + model_config=expected_model, ) - @patch("services.message_service.db") - @patch("services.message_service.ModelManager.for_tenant") - @patch("services.message_service.TokenBufferMemory") - @patch("services.message_service.LLMGenerator") - @patch("services.message_service.TraceQueueManager") - @patch.object(MessageService, "get_message") - @patch("services.message_service.ConversationService") - def test_get_suggested_questions_chat_app_invalid_frontend_model_fallback_to_default( + def test_chat_app_uses_compatible_override_model_config( self, - mock_conversation_service: MagicMock, - mock_get_message: MagicMock, - mock_trace_manager: MagicMock, - mock_llm_gen: MagicMock, - mock_memory: MagicMock, - mock_model_manager: MagicMock, - mock_db: MagicMock, - factory: TestMessageServiceFactory, - ): - """Test invalid frontend configured model falls back to tenant default model.""" - app = factory.create_app_mock(mode=AppMode.CHAT) - app.tenant_id = "tenant-123" - user = factory.create_end_user_mock() - message = factory.create_message_mock() - mock_get_message.return_value = message - - conversation = MagicMock() - conversation.override_model_configs = None - mock_conversation_service.get_conversation.return_value = conversation - - app_model_config = MagicMock() - app_model_config.suggested_questions_after_answer_dict = { - "enabled": True, - "model": {"provider": "openai", "name": "invalid-model"}, - } - mock_db.session.scalar.return_value = app_model_config - - mock_model_manager.return_value.get_model_instance.side_effect = ValueError("invalid model") - mock_memory.return_value.get_history_prompt_text.return_value = "histories" - mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] - - result = MessageService.get_suggested_questions_after_answer( - app_model=app, - user=user, - message_id="msg-123", - invoke_from=MagicMock(), - session=mock_db.session, - ) - - assert result == ["Q1?"] - mock_model_manager.return_value.get_default_model_instance.assert_called_once_with( - tenant_id="tenant-123", - model_type=ModelType.LLM, - ) - mock_model_manager.return_value.get_model_instance.assert_not_called() - - @patch("services.message_service.db") - @patch("services.message_service.ModelManager.for_tenant") - @patch("services.message_service.TokenBufferMemory") - @patch("services.message_service.LLMGenerator") - @patch("services.message_service.TraceQueueManager") - @patch.object(MessageService, "get_message") - @patch("services.message_service.ConversationService") - def test_get_suggested_questions_chat_app_uses_compatible_override_model_config( - self, - mock_conversation_service: MagicMock, - mock_get_message: MagicMock, - mock_trace_manager: MagicMock, - mock_llm_gen: MagicMock, - mock_memory: MagicMock, - mock_model_manager: MagicMock, - mock_db: MagicMock, - factory: TestMessageServiceFactory, - ): - """Test legacy override configs are normalized before suggested questions reads them.""" - app = factory.create_app_mock(mode=AppMode.CHAT) - app.tenant_id = "tenant-123" - user = factory.create_end_user_mock() - message = factory.create_message_mock() - mock_get_message.return_value = message - - conversation = MagicMock() - conversation.override_model_configs = json.dumps( - { - "speech_to_text": {"enabled": False}, - "text_to_speech": {"enabled": False}, - "retriever_resource": {"enabled": False}, - "model": {"provider": "openai", "name": "gpt-4o-mini", "mode": "chat"}, - "user_input_form": [], - "dataset_query_variable": "", - "pre_prompt": "", - "agent_mode": { - "enabled": False, - "max_iteration": 5, - "strategy": "function_call", - "tools": [], - }, - "prompt_type": "simple", - "chat_prompt_config": {}, - "completion_prompt_config": {}, - "dataset_configs": {"retrieval_model": "single", "datasets": {"datasets": []}}, - "file_upload": { - "image": { - "detail": "high", - "enabled": False, - "number_limits": 3, - "transfer_methods": ["remote_url", "local_file"], - } - }, - "suggested_questions_after_answer": { - "enabled": True, - "prompt": "legacy prompt", - }, - } - ) - conversation.model_config = { - "opening_statement": None, - "suggested_questions": [], - "suggested_questions_after_answer": { - "enabled": True, - "prompt": "legacy prompt", - }, - "speech_to_text": {"enabled": False}, - "text_to_speech": {"enabled": False}, - "retriever_resource": {"enabled": False}, - "annotation_reply": {"enabled": False}, - "more_like_this": {"enabled": False}, - "sensitive_word_avoidance": {"enabled": False, "type": "", "config": {}}, - "external_data_tools": [], + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: + override = { "model": {"provider": "openai", "name": "gpt-4o-mini", "mode": "chat"}, - "user_input_form": [], - "dataset_query_variable": "", - "pre_prompt": "", - "agent_mode": {"enabled": False, "strategy": "function_call", "tools": [], "prompt": None}, - "prompt_type": "simple", - "chat_prompt_config": {}, - "completion_prompt_config": {}, - "dataset_configs": {"retrieval_model": "single", "datasets": {"datasets": []}}, - "file_upload": { - "image": { - "detail": "high", - "enabled": False, - "number_limits": 3, - "transfer_methods": ["remote_url", "local_file"], - } - }, - "model_id": None, - "provider": None, + "suggested_questions_after_answer": {"enabled": True, "prompt": "legacy prompt"}, } - conversation.model_config_with_session.return_value = conversation.model_config - mock_conversation_service.get_conversation.return_value = conversation - - mock_memory.return_value.get_history_prompt_text.return_value = "histories" - mock_llm_gen.generate_suggested_questions_after_answer.return_value = ["Q1?"] + conversation = factory.create_conversation(override_model_configs=json.dumps(override)) + _, _, llm_generator = self._chat_boundaries(monkeypatch, conversation) result = MessageService.get_suggested_questions_after_answer( - app_model=app, - user=user, + app_model=factory.create_app(mode=AppMode.CHAT), + user=factory.create_end_user(), message_id="msg-123", - invoke_from=MagicMock(), - session=mock_db.session, + invoke_from=InvokeFrom.WEB_APP, + session=sqlite_session, ) assert result == ["Q1?"] - mock_db.session.scalar.assert_not_called() - mock_llm_gen.generate_suggested_questions_after_answer.assert_called_once_with( + llm_generator.generate_suggested_questions_after_answer.assert_called_once_with( tenant_id="tenant-123", histories="histories", instruction_prompt="legacy prompt", model_config=None, ) - # Test 30: get_suggested_questions_after_answer - Disabled Error - @patch("services.message_service.WorkflowService") - @patch("services.message_service.AdvancedChatAppConfigManager") - @patch.object(MessageService, "get_message") - @patch("services.message_service.ConversationService") - def test_get_suggested_questions_disabled_error( + def test_disabled_error( self, - mock_conversation_service, - mock_get_message, - mock_config_manager, - mock_workflow_service, - factory: TestMessageServiceFactory, - ): - """Test SuggestedQuestionsAfterAnswerDisabledError is raised when feature is disabled.""" - # Arrange - app = factory.create_app_mock(mode=AppMode.ADVANCED_CHAT.value) - user = factory.create_end_user_mock() - mock_get_message.return_value = factory.create_message_mock() - + monkeypatch: pytest.MonkeyPatch, + factory: MessageServiceTestDataFactory, + sqlite_session: Session, + ) -> None: + conversation = factory.create_conversation() + self._chat_boundaries(monkeypatch, conversation) workflow = MagicMock() - mock_workflow_service.return_value.get_published_workflow.return_value = workflow + workflow_service = MagicMock() + workflow_service.return_value.get_published_workflow.return_value = workflow + monkeypatch.setattr(service_module, "WorkflowService", workflow_service) + app_config_manager = MagicMock() + app_config_manager.get_app_config.return_value.additional_features.suggested_questions_after_answer = False + monkeypatch.setattr(service_module, "AdvancedChatAppConfigManager", app_config_manager) - app_config = MagicMock() - app_config.additional_features.suggested_questions_after_answer = False - mock_config_manager.get_app_config.return_value = app_config - - # Act & Assert with pytest.raises(SuggestedQuestionsAfterAnswerDisabledError): MessageService.get_suggested_questions_after_answer( - app_model=app, - user=user, + app_model=factory.create_app(mode=AppMode.ADVANCED_CHAT), + user=factory.create_end_user(), message_id="msg-123", - invoke_from=MagicMock(), - session=MagicMock(), + invoke_from=InvokeFrom.WEB_APP, + session=sqlite_session, ) diff --git a/api/tests/unit_tests/services/test_messages_clean_service.py b/api/tests/unit_tests/services/test_messages_clean_service.py index 73c096c749c..a7b0a68d2ec 100644 --- a/api/tests/unit_tests/services/test_messages_clean_service.py +++ b/api/tests/unit_tests/services/test_messages_clean_service.py @@ -4,7 +4,7 @@ from unittest.mock import MagicMock, patch import pytest -from enums.cloud_plan import CloudPlan +from enums import CloudPlan, DeploymentEdition from services.retention.conversation.messages_clean_policy import ( BillingDisabledPolicy, BillingSandboxPolicy, @@ -404,10 +404,10 @@ class TestCreateMessageCleanPolicy: """Unit tests for create_message_clean_policy factory function.""" @patch("services.retention.conversation.messages_clean_policy.dify_config") - def test_billing_disabled_returns_billing_disabled_policy(self, mock_config): - """Test that BILLING_ENABLED=False returns BillingDisabledPolicy.""" + def test_non_cloud_edition_returns_billing_disabled_policy(self, mock_config): + """Test that the Community edition returns BillingDisabledPolicy.""" # Arrange - mock_config.BILLING_ENABLED = False + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY # Act policy = create_message_clean_policy(graceful_period_days=21) @@ -417,10 +417,10 @@ class TestCreateMessageCleanPolicy: @patch("services.retention.conversation.messages_clean_policy.BillingService", autospec=True) @patch("services.retention.conversation.messages_clean_policy.dify_config") - def test_billing_enabled_policy_has_correct_internals(self, mock_config, mock_billing_service): + def test_cloud_edition_policy_has_correct_internals(self, mock_config, mock_billing_service): """Test that BillingSandboxPolicy is created with correct internal values.""" # Arrange - mock_config.BILLING_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD whitelist = ["tenant1", "tenant2"] mock_billing_service.get_expired_subscription_cleanup_whitelist.return_value = whitelist mock_plan_provider = MagicMock() diff --git a/api/tests/unit_tests/services/test_metadata_bug_complete.py b/api/tests/unit_tests/services/test_metadata_bug_complete.py index 00f16f75ac0..6b22281fb48 100644 --- a/api/tests/unit_tests/services/test_metadata_bug_complete.py +++ b/api/tests/unit_tests/services/test_metadata_bug_complete.py @@ -3,6 +3,7 @@ from typing import cast from unittest.mock import Mock import pytest +from sqlalchemy.orm import Session from models import Account, Tenant from services.entities.knowledge_entities.knowledge_entities import MetadataArgs @@ -38,7 +39,8 @@ class TestMetadataBugCompleteValidation: assert valid_args.type == "string" assert valid_args.name == "test_name" - def test_2_business_logic_layer_crashes_on_none(self) -> None: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_2_business_logic_layer_crashes_on_none(self, sqlite_session: Session) -> None: """Test Layer 2: Business logic crashes when None values slip through.""" # Create mock that bypasses Pydantic validation mock_metadata_args = Mock() @@ -48,15 +50,16 @@ class TestMetadataBugCompleteValidation: account = _make_account() # Should crash with TypeError with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.create_metadata("dataset-123", mock_metadata_args, account, "tenant-123", session=Mock()) + MetadataService.create_metadata( + "dataset-123", mock_metadata_args, account, "tenant-123", session=sqlite_session + ) # Test update method as well account = _make_account() none_name = cast(str, None) with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.update_metadata_name( - "dataset-123", "metadata-456", none_name, account, "tenant-123", session=Mock() - ) + MetadataService.update_metadata_name(Mock(), "metadata-456", none_name, account, session=sqlite_session) + assert not sqlite_session.in_transaction() def test_3_database_constraints_verification(self) -> None: """Test Layer 3: Verify database model has nullable=False constraints.""" @@ -91,7 +94,8 @@ class TestMetadataBugCompleteValidation: assert args.type == "string" assert args.name == "valid_name" - def test_6_simulated_buggy_behavior(self) -> None: + @pytest.mark.parametrize("sqlite_session", [()], indirect=True) + def test_6_simulated_buggy_behavior(self, sqlite_session: Session) -> None: """Test simulating the original buggy behavior by bypassing Pydantic validation.""" mock_metadata_args = Mock() mock_metadata_args.name = None @@ -99,7 +103,10 @@ class TestMetadataBugCompleteValidation: account = _make_account() with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.create_metadata("dataset-123", mock_metadata_args, account, "tenant-123", session=Mock()) + MetadataService.create_metadata( + "dataset-123", mock_metadata_args, account, "tenant-123", session=sqlite_session + ) + assert not sqlite_session.in_transaction() def test_7_end_to_end_validation_layers(self) -> None: """Test all validation layers work together correctly.""" diff --git a/api/tests/unit_tests/services/test_metadata_nullable_bug.py b/api/tests/unit_tests/services/test_metadata_nullable_bug.py index af898f3cb4a..4285c6f40d9 100644 --- a/api/tests/unit_tests/services/test_metadata_nullable_bug.py +++ b/api/tests/unit_tests/services/test_metadata_nullable_bug.py @@ -54,9 +54,7 @@ class TestMetadataNullableBug: none_name = cast(str, None) # This should crash with TypeError when calling len(None) with pytest.raises(TypeError, match="object of type 'NoneType' has no len"): - MetadataService.update_metadata_name( - "dataset-123", "metadata-456", none_name, account, "tenant-123", session=sqlite_session - ) + MetadataService.update_metadata_name(Mock(), "metadata-456", none_name, account, session=sqlite_session) assert not sqlite_session.in_transaction() def test_api_layer_now_uses_pydantic_validation(self) -> None: diff --git a/api/tests/unit_tests/services/test_metadata_service_session_boundary.py b/api/tests/unit_tests/services/test_metadata_service_session_boundary.py index 7832cdf435b..995919f4b0d 100644 --- a/api/tests/unit_tests/services/test_metadata_service_session_boundary.py +++ b/api/tests/unit_tests/services/test_metadata_service_session_boundary.py @@ -1,16 +1,30 @@ from datetime import datetime +from types import SimpleNamespace from unittest.mock import MagicMock, patch +import pytest +from sqlalchemy import event, select +from sqlalchemy.orm import Session + from core.rag.index_processor.constant.built_in_field import BuiltInField from models import Account +from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding, Document +from models.enums import DataSourceType, DocumentCreatedFrom from services.dataset_service import DocumentService from services.entities.knowledge_entities.knowledge_entities import ( DocumentMetadataOperation, MetadataArgs, + MetadataDetail, MetadataOperationData, ) +from services.errors.metadata import MetadataResourceNotFoundError from services.metadata_service import MetadataService +DOCUMENT_ID = "11111111-1111-1111-1111-111111111111" +FOREIGN_DOCUMENT_ID = "22222222-2222-2222-2222-222222222222" +METADATA_ID = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" +FOREIGN_METADATA_ID = "bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb" + def _account() -> Account: account = Account(name="User", email="user@example.com") @@ -18,57 +32,74 @@ def _account() -> Account: return account -def test_create_metadata_flushes_without_committing_caller_session() -> None: - session = MagicMock() - session.scalar.return_value = None +def test_create_metadata_flushes_without_committing_caller_session(sqlite_session: Session) -> None: + transaction_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: transaction_events.append("commit")) + event.listen(sqlite_session, "after_rollback", lambda _session: transaction_events.append("rollback")) metadata = MetadataService.create_metadata( "dataset-1", MetadataArgs(type="string", name="author"), _account(), "tenant-1", - session=session, + session=sqlite_session, ) assert metadata.name == "author" - session.flush.assert_called_once_with() - session.commit.assert_not_called() - session.rollback.assert_not_called() + assert sqlite_session.get(DatasetMetadata, metadata.id) is metadata + assert transaction_events == [] -def _document() -> MagicMock: - document = MagicMock() - document.id = "document-1" - document.name = "Document" - document.doc_metadata = {} - document.data_source_type = "upload_file" - document.upload_date = datetime(2026, 1, 1) - document.last_update_date = datetime(2026, 1, 2) - document.uploader = "global-session uploader" - document.get_uploader.return_value = "caller-session uploader" - return document +def _dataset(*, built_in_field_enabled: bool) -> Dataset: + return Dataset( + id="dataset-1", + tenant_id="tenant-1", + name="Dataset", + description="", + provider="vendor", + created_by="account-1", + built_in_field_enabled=built_in_field_enabled, + ) -def test_enable_built_in_field_uses_caller_session_for_uploader() -> None: - session = MagicMock() - dataset = MagicMock(id="dataset-1", built_in_field_enabled=False) +def _document() -> Document: + return Document( + id=DOCUMENT_ID, + tenant_id="tenant-1", + dataset_id="dataset-1", + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="Document", + created_from=DocumentCreatedFrom.API, + created_by="account-1", + created_at=datetime(2026, 1, 1), + updated_at=datetime(2026, 1, 2), + doc_metadata={}, + ) + + +def test_enable_built_in_field_uses_caller_session_for_uploader(sqlite_session: Session) -> None: + dataset = _dataset(built_in_field_enabled=False) document = _document() + sqlite_session.add_all([_account(), dataset, document]) + sqlite_session.commit() with ( patch.object(MetadataService, "knowledge_base_metadata_lock_check"), patch.object(DocumentService, "get_working_documents_by_dataset_id", return_value=[document]), patch("services.metadata_service.redis_client.delete"), ): - MetadataService.enable_built_in_field(dataset, session) + MetadataService.enable_built_in_field(dataset, sqlite_session) - assert document.doc_metadata[BuiltInField.uploader] == "caller-session uploader" - document.get_uploader.assert_called_once_with(session=session) + assert document.doc_metadata[BuiltInField.uploader] == "User" -def test_update_documents_metadata_uses_caller_session_for_uploader() -> None: - session = MagicMock() - dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", built_in_field_enabled=True) +def test_update_documents_metadata_uses_caller_session_for_uploader(sqlite_session: Session) -> None: + dataset = _dataset(built_in_field_enabled=True) document = _document() + sqlite_session.add_all([_account(), dataset, document]) + sqlite_session.commit() metadata_args = MetadataOperationData( operation_data=[ DocumentMetadataOperation(document_id=document.id, metadata_list=[], partial_update=False), @@ -77,31 +108,155 @@ def test_update_documents_metadata_uses_caller_session_for_uploader() -> None: with ( patch.object(MetadataService, "knowledge_base_metadata_lock_check"), - patch.object(DocumentService, "get_document", return_value=document), patch("services.metadata_service.redis_client.delete"), ): MetadataService.update_documents_metadata( dataset, metadata_args, _account(), - "tenant-1", - session=session, + session=sqlite_session, ) - assert document.doc_metadata[BuiltInField.uploader] == "caller-session uploader" - document.get_uploader.assert_called_once_with(session=session) + assert document.doc_metadata[BuiltInField.uploader] == "User" -def test_get_dataset_metadatas_uses_caller_session() -> None: +def test_update_documents_metadata_rejects_foreign_metadata_before_writes() -> None: session = MagicMock() - session.scalar.return_value = 2 - dataset = MagicMock(id="dataset-1", built_in_field_enabled=False) - dataset.get_doc_metadata.return_value = [{"id": "metadata-1", "name": "author", "type": "string"}] + session.scalars.return_value.all.return_value = [] + dataset = MagicMock(id="dataset-1", tenant_id="tenant-1") + metadata_args = MetadataOperationData( + operation_data=[ + DocumentMetadataOperation( + document_id=DOCUMENT_ID, + metadata_list=[MetadataDetail(id=FOREIGN_METADATA_ID, name="spoofed", value="value")], + partial_update=False, + ) + ] + ) - result = MetadataService.get_dataset_metadatas(dataset, session) + with ( + pytest.raises(MetadataResourceNotFoundError, match="Metadata not found"), + patch.object(MetadataService, "knowledge_base_metadata_lock_check") as lock_check, + ): + MetadataService.update_documents_metadata(dataset, metadata_args, _account(), session=session) + + lock_check.assert_not_called() + session.add.assert_not_called() + session.execute.assert_not_called() + session.commit.assert_not_called() + + +def test_update_documents_metadata_validates_all_documents_before_writes() -> None: + session = MagicMock() + metadata = SimpleNamespace(id=METADATA_ID, name="canonical") + session.scalars.return_value.all.side_effect = [[metadata], [DOCUMENT_ID]] + dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", built_in_field_enabled=False) + metadata_detail = MetadataDetail(id=metadata.id, name="spoofed", value="value") + metadata_args = MetadataOperationData( + operation_data=[ + DocumentMetadataOperation(document_id=DOCUMENT_ID, metadata_list=[metadata_detail], partial_update=False), + DocumentMetadataOperation( + document_id=FOREIGN_DOCUMENT_ID, metadata_list=[metadata_detail], partial_update=False + ), + ] + ) + + with ( + pytest.raises(MetadataResourceNotFoundError, match="Document not found"), + patch.object(MetadataService, "knowledge_base_metadata_lock_check") as lock_check, + patch("services.metadata_service.redis_client.delete"), + ): + MetadataService.update_documents_metadata(dataset, metadata_args, _account(), session=session) + + lock_check.assert_not_called() + session.add.assert_not_called() + session.execute.assert_not_called() + session.commit.assert_not_called() + + +def test_update_documents_metadata_uses_canonical_metadata_name() -> None: + session = MagicMock() + metadata = SimpleNamespace(id=METADATA_ID, name="canonical") + session.scalars.return_value.all.side_effect = [[metadata], [DOCUMENT_ID]] + dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", built_in_field_enabled=False) + document = _document() + session.scalar.return_value = document + metadata_args = MetadataOperationData( + operation_data=[ + DocumentMetadataOperation( + document_id=document.id, + metadata_list=[MetadataDetail(id=metadata.id, name="spoofed", value="value")], + partial_update=False, + ) + ] + ) + + with ( + patch.object(MetadataService, "knowledge_base_metadata_lock_check"), + patch("services.metadata_service.redis_client.delete"), + ): + MetadataService.update_documents_metadata(dataset, metadata_args, _account(), session=session) + + assert document.doc_metadata == {"canonical": "value"} + + +def test_metadata_operation_normalizes_uuid_ids() -> None: + operation = DocumentMetadataOperation( + document_id=DOCUMENT_ID.upper(), + metadata_list=[MetadataDetail(id=METADATA_ID.upper(), name="ignored", value="value")], + ) + + assert operation.document_id == DOCUMENT_ID + assert operation.metadata_list[0].id == METADATA_ID + + +def test_document_metadata_details_scopes_binding_to_document_owner() -> None: + session = MagicMock() + session.scalars.return_value.all.return_value = [] + document = MagicMock( + id="document-1", + tenant_id="tenant-1", + dataset_id="dataset-1", + doc_metadata={"canonical": "value"}, + ) + document.get_built_in_fields.return_value = [] + + assert Document.get_doc_metadata_details(document, session=session) == [] + + statement = str(session.scalars.call_args.args[0]) + assert "dataset_metadatas.tenant_id" in statement + assert "dataset_metadatas.dataset_id" in statement + assert "dataset_metadata_bindings.tenant_id" in statement + assert "dataset_metadata_bindings.dataset_id" in statement + assert "dataset_metadata_bindings.document_id" in statement + + +def test_get_dataset_metadatas_uses_caller_session(monkeypatch, sqlite_session: Session) -> None: + dataset = _dataset(built_in_field_enabled=False) + sqlite_session.add_all( + [ + DatasetMetadataBinding( + tenant_id="tenant-1", + dataset_id="dataset-1", + document_id=f"document-{index}", + metadata_id="metadata-1", + created_by="account-1", + ) + for index in range(2) + ] + ) + sqlite_session.commit() + + def get_doc_metadata(_dataset: Dataset, *, session: Session) -> list[dict[str, str]]: + assert session is sqlite_session + return [{"id": "metadata-1", "name": "author", "type": "string"}] + + monkeypatch.setattr(Dataset, "get_doc_metadata", get_doc_metadata) + + result = MetadataService.get_dataset_metadatas(dataset, sqlite_session) assert result == { "doc_metadata": [{"id": "metadata-1", "name": "author", "type": "string", "count": 2}], "built_in_field_enabled": False, } - dataset.get_doc_metadata.assert_called_once_with(session=session) + assert sqlite_session.scalar(select(DatasetMetadataBinding).limit(1)) is not None diff --git a/api/tests/unit_tests/services/test_model_load_balancing_service.py b/api/tests/unit_tests/services/test_model_load_balancing_service.py index 743e6e797a3..3d2c98102a0 100644 --- a/api/tests/unit_tests/services/test_model_load_balancing_service.py +++ b/api/tests/unit_tests/services/test_model_load_balancing_service.py @@ -1,28 +1,55 @@ +"""SQLite-backed tests for :mod:`services.model_load_balancing_service`.""" + from __future__ import annotations import json -from types import SimpleNamespace -from typing import Any, cast -from unittest.mock import MagicMock +from collections.abc import Iterator +from contextlib import contextmanager +from typing import cast +from unittest.mock import MagicMock, patch import pytest -from pytest_mock import MockerFixture +from sqlalchemy import Engine, event, select +from sqlalchemy.orm import Session from constants import HIDDEN_VALUE +from core.entities.provider_configuration import ProviderConfiguration, ProviderConfigurations +from core.entities.provider_entities import CustomConfiguration, CustomProviderConfiguration, SystemConfiguration +from core.plugin.impl.model_runtime_factory import PluginModelAssembly +from core.provider_manager import ProviderManager from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.provider_entities import ( + ConfigurateMethod, CredentialFormSchema, FieldModelSchema, FormType, ModelCredentialSchema, ProviderCredentialSchema, + ProviderEntity, ) -from models.provider import LoadBalancingModelConfig +from graphon.model_runtime.model_providers.model_provider_factory import ModelProviderFactory +from graphon.model_runtime.protocols.runtime import ModelRuntime +from models.engine import db +from models.enums import CredentialSourceType +from models.provider import ( + LoadBalancingModelConfig, + ProviderCredential, + ProviderModelCredential, + ProviderModelSetting, + ProviderType, +) +from models.provider_ids import ModelProviderID from services.model_load_balancing_service import ModelLoadBalancingService +pytestmark = pytest.mark.parametrize( + "sqlite_session", + [(LoadBalancingModelConfig, ProviderCredential, ProviderModelCredential, ProviderModelSetting)], + indirect=True, +) -def _build_provider_credential_schema() -> ProviderCredentialSchema: + +def _provider_schema() -> ProviderCredentialSchema: return ProviderCredentialSchema( credential_form_schemas=[ CredentialFormSchema(variable="api_key", label=I18nObject(en_US="API Key"), type=FormType.SECRET_INPUT) @@ -30,7 +57,7 @@ def _build_provider_credential_schema() -> ProviderCredentialSchema: ) -def _build_model_credential_schema() -> ModelCredentialSchema: +def _model_schema() -> ModelCredentialSchema: return ModelCredentialSchema( model=FieldModelSchema(label=I18nObject(en_US="Model")), credential_form_schemas=[ @@ -39,827 +66,442 @@ def _build_model_credential_schema() -> ModelCredentialSchema: ) -def _build_provider_configuration( - *, - custom_provider: bool = False, - load_balancing_enabled: bool | None = None, - model_schema: ModelCredentialSchema | None = None, - provider_schema: ProviderCredentialSchema | None = None, -) -> MagicMock: - provider_configuration = MagicMock() - provider_configuration.provider = SimpleNamespace( - provider="openai", - model_credential_schema=model_schema, - provider_credential_schema=provider_schema, +def _provider_configuration() -> ProviderConfiguration: + """Build a concrete provider configuration for service tests.""" + return ProviderConfiguration( + tenant_id="tenant-1", + provider=ProviderEntity( + provider="openai", + label=I18nObject(en_US="OpenAI"), + supported_model_types=[ModelType.LLM], + configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL], + provider_credential_schema=_provider_schema(), + ), + preferred_provider_type=ProviderType.SYSTEM, + using_provider_type=ProviderType.SYSTEM, + system_configuration=SystemConfiguration(enabled=False), + custom_configuration=CustomConfiguration(provider=None, models=[]), + model_settings=[], ) - provider_configuration.custom_configuration = SimpleNamespace(provider=custom_provider) - provider_configuration.extract_secret_variables.return_value = ["api_key"] - provider_configuration.obfuscated_credentials.side_effect = lambda credentials, credential_form_schemas: credentials - provider_configuration.get_provider_model_setting.return_value = ( - None if load_balancing_enabled is None else SimpleNamespace(load_balancing_enabled=load_balancing_enabled) - ) - return provider_configuration -def _load_balancing_model_config(**kwargs: Any) -> LoadBalancingModelConfig: - return cast(LoadBalancingModelConfig, SimpleNamespace(**kwargs)) +type ServiceFixture = tuple[ModelLoadBalancingService, MagicMock, ProviderConfiguration] @pytest.fixture -def service(mocker: MockerFixture) -> ModelLoadBalancingService: - # Arrange - provider_manager = MagicMock() - mocker.patch("services.model_load_balancing_service.create_plugin_provider_manager", return_value=provider_manager) - model_assembly = SimpleNamespace(provider_manager=provider_manager, model_provider_factory=MagicMock()) - mocker.patch("services.model_load_balancing_service.create_plugin_model_assembly", return_value=model_assembly) +def service(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> ServiceFixture: + configuration = _provider_configuration() + manager = MagicMock() + manager.get_configurations.return_value = {"openai": configuration} svc = ModelLoadBalancingService() - svc.provider_manager = provider_manager - svc.model_assembly = model_assembly - svc._get_provider_manager = lambda _tenant_id: provider_manager # type: ignore[method-assign] - return svc - - -@pytest.fixture -def mock_db() -> MagicMock: - # Arrange - mocked_db = MagicMock() - mocked_db.session = MagicMock() - return mocked_db - - -@pytest.mark.parametrize( - ("method_name", "expected_provider_method"), - [ - ("enable_model_load_balancing", "enable_model_load_balancing"), - ("disable_model_load_balancing", "disable_model_load_balancing"), - ], -) -def test_enable_disable_model_load_balancing_should_call_provider_configuration_method_when_provider_exists( - method_name: str, - expected_provider_method: str, - service: ModelLoadBalancingService, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - - # Act - getattr(service, method_name)("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM) - - # Assert - getattr(provider_configuration, expected_provider_method).assert_called_once_with( - model="gpt-4o-mini", model_type=ModelType.LLM + monkeypatch.setattr(svc, "_get_provider_manager", lambda _tenant_id: manager) + monkeypatch.setattr( + "services.model_load_balancing_service.create_plugin_provider_manager", lambda tenant_id: manager ) + monkeypatch.setattr( + "services.model_load_balancing_service.ProviderManager.invalidate_configurations_cache", MagicMock() + ) + monkeypatch.setattr("services.model_load_balancing_service.ProviderCredentialsCache", MagicMock()) + sqlite_engine = cast(Engine, sqlite_session.get_bind()) + monkeypatch.setattr(type(db), "engine", property(lambda _db: sqlite_engine)) + return svc, manager, configuration -@pytest.mark.parametrize( - ("method_name", "expected_provider_method"), - [ - ("enable_model_load_balancing", "enable_model_load_balancing"), - ("disable_model_load_balancing", "disable_model_load_balancing"), - ], -) -def test_enable_disable_model_load_balancing_uses_model_type_constructor_directly( - method_name: str, - expected_provider_method: str, - service: ModelLoadBalancingService, +def _config( + session: Session, + *, + tenant_id: str = "tenant-1", + provider: str = "openai", + model: str = "gpt-4o-mini", + name: str = "primary", + encrypted_config: str | None = '{"api_key":"encrypted"}', + credential_id: str | None = None, + source: CredentialSourceType | None = None, + enabled: bool = True, +) -> LoadBalancingModelConfig: + config = LoadBalancingModelConfig( + tenant_id=tenant_id, + provider_name=provider, + model_name=model, + model_type=ModelType.LLM, + name=name, + encrypted_config=encrypted_config, + credential_id=credential_id, + credential_source_type=source, + enabled=enabled, + ) + session.add(config) + session.commit() + return config + + +@contextmanager +def _raise_on_insert(engine: Engine) -> Iterator[None]: + def raise_error(_conn, _cursor, statement, _parameters, _context, _executemany): + if statement.lstrip().upper().startswith("INSERT") and "load_balancing_model_configs" in statement: + raise RuntimeError("forced INSERT") + + event.listen(engine, "before_cursor_execute", raise_error) + try: + yield + finally: + event.remove(engine, "before_cursor_execute", raise_error) + + +@pytest.mark.parametrize("enabled", [True, False]) +def test_enable_disable_persists_provider_model_setting( + enabled: bool, monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} + _config(sqlite_session, name="primary") + _config(sqlite_session, name="secondary") + configuration = _provider_configuration() + configurations = ProviderConfigurations(tenant_id="tenant-1") + configurations[str(ModelProviderID("openai"))] = configuration + manager = ProviderManager(cast(ModelRuntime, object())) + manager._configurations_cache["tenant-1"] = configurations + svc = ModelLoadBalancingService() + monkeypatch.setattr(svc, "_get_provider_manager", lambda _tenant_id: manager) + sqlite_engine = cast(Engine, sqlite_session.get_bind()) + monkeypatch.setattr(type(db), "engine", property(lambda _db: sqlite_engine)) - getattr(service, method_name)("tenant-1", "openai", "gpt-4o-mini", "text-generation") + if enabled: + svc.enable_model_load_balancing("tenant-1", "openai", "gpt-4o-mini", "llm") + else: + svc.disable_model_load_balancing("tenant-1", "openai", "gpt-4o-mini", "llm") - getattr(provider_configuration, expected_provider_method).assert_called_once_with( - model="gpt-4o-mini", model_type=ModelType.LLM + sqlite_session.expire_all() + model_setting = sqlite_session.scalar( + select(ProviderModelSetting).where( + ProviderModelSetting.tenant_id == "tenant-1", + ProviderModelSetting.provider_name == "openai", + ProviderModelSetting.model_name == "gpt-4o-mini", + ProviderModelSetting.model_type == ModelType.LLM, + ) ) + assert model_setting is not None + assert model_setting.load_balancing_enabled is enabled -@pytest.mark.parametrize( - "method_name", - ["enable_model_load_balancing", "disable_model_load_balancing"], -) -def test_enable_disable_model_load_balancing_should_raise_value_error_when_provider_missing( - method_name: str, - service: ModelLoadBalancingService, -) -> None: - # Arrange - service.provider_manager.get_configurations.return_value = {} - - # Act + Assert +def test_provider_missing_errors_use_runtime_boundary(service: ServiceFixture, sqlite_session: Session) -> None: + svc, manager, _ = service + manager.get_configurations.return_value = {} with pytest.raises(ValueError, match="Provider openai does not exist"): - getattr(service, method_name)("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM) - - -def test_get_load_balancing_configs_should_raise_value_error_when_provider_missing( - service: ModelLoadBalancingService, -) -> None: - # Arrange - service.provider_manager.get_configurations.return_value = {} - - # Act + Assert + svc.enable_model_load_balancing("tenant-1", "openai", "model", ModelType.LLM) with pytest.raises(ValueError, match="Provider openai does not exist"): - service.get_load_balancing_configs("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, session=MagicMock()) + svc.get_load_balancing_configs("tenant-1", "openai", "model", ModelType.LLM, session=sqlite_session) -def test_get_load_balancing_configs_should_insert_inherit_config_when_missing_for_custom_provider( - service: ModelLoadBalancingService, - mock_db: MagicMock, - mocker: MockerFixture, +def test_get_configs_inserts_inherit_and_filters_tenant_provider_and_source( + monkeypatch: pytest.MonkeyPatch, + service: ServiceFixture, + sqlite_session: Session, ) -> None: - # Arrange - provider_configuration = _build_provider_configuration( - custom_provider=True, - load_balancing_enabled=True, - provider_schema=_build_provider_credential_schema(), + svc, _, configuration = service + configuration.custom_configuration.provider = CustomProviderConfiguration(credentials={}) + sqlite_session.add( + ProviderModelSetting( + tenant_id="tenant-1", + provider_name="openai", + model_name="gpt-4o-mini", + model_type=ModelType.LLM, + load_balancing_enabled=True, + ) ) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - config = SimpleNamespace( - id="cfg-1", - name="primary", - encrypted_config=json.dumps({"api_key": "encrypted-key"}), - credential_id="cred-1", - enabled=True, + sqlite_session.commit() + matching = _config(sqlite_session, credential_id="cred-1", source=CredentialSourceType.PROVIDER, name="matching") + _config(sqlite_session, tenant_id="tenant-2", source=CredentialSourceType.PROVIDER, name="foreign-tenant") + _config(sqlite_session, provider="anthropic", source=CredentialSourceType.PROVIDER, name="foreign-provider") + _config(sqlite_session, source=CredentialSourceType.CUSTOM_MODEL, name="foreign-source") + monkeypatch.setattr( + "services.model_load_balancing_service.encrypter.get_decrypt_decoding", lambda _tenant: ("rsa", "cipher") ) - mock_db.session.scalars.return_value.all.return_value = [config] - mocker.patch( - "services.model_load_balancing_service.encrypter.get_decrypt_decoding", - return_value=("rsa", "cipher"), - ) - mocker.patch( + monkeypatch.setattr( "services.model_load_balancing_service.encrypter.decrypt_token_with_decoding", - return_value="plain-key", + lambda _value, _decoding: "plain", ) - mocker.patch( + monkeypatch.setattr( "services.model_load_balancing_service.LBModelManager.get_config_in_cooldown_and_ttl", - return_value=(False, 0), + lambda **_kwargs: (False, 0), ) - - # Act - is_enabled, configs = service.get_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - session=mock_db.session, - ) - - # Assert - assert is_enabled is True - assert len(configs) == 2 - assert configs[0]["name"] == "__inherit__" - assert configs[1]["name"] == "primary" - assert configs[1]["credentials"] == {"api_key": "plain-key"} - assert mock_db.session.add.call_count == 1 - assert mock_db.session.commit.call_count == 1 - - -def test_get_load_balancing_configs_should_reorder_existing_inherit_and_tolerate_json_or_decrypt_errors( - service: ModelLoadBalancingService, - mock_db: MagicMock, - mocker: MockerFixture, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration( - custom_provider=True, - load_balancing_enabled=None, - provider_schema=_build_provider_credential_schema(), - ) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - normal_config = SimpleNamespace( - id="cfg-1", - name="normal", - encrypted_config=json.dumps({"api_key": "bad-encrypted"}), - credential_id="cred-1", - enabled=True, - ) - inherit_config = SimpleNamespace( - id="cfg-2", - name="__inherit__", - encrypted_config="not-json", - credential_id=None, - enabled=False, - ) - mock_db.session.scalars.return_value.all.return_value = [ - normal_config, - inherit_config, - ] - mocker.patch( - "services.model_load_balancing_service.encrypter.get_decrypt_decoding", - return_value=("rsa", "cipher"), - ) - mocker.patch( - "services.model_load_balancing_service.encrypter.decrypt_token_with_decoding", - side_effect=ValueError("cannot decrypt"), - ) - mocker.patch( - "services.model_load_balancing_service.LBModelManager.get_config_in_cooldown_and_ttl", - return_value=(True, 15), - ) - - # Act - is_enabled, configs = service.get_load_balancing_configs( + enabled, configs = svc.get_load_balancing_configs( "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, config_from="predefined-model", - session=mock_db.session, + session=sqlite_session, ) + assert enabled is True + assert [config["name"] for config in configs] == ["__inherit__", "matching"] + assert configs[1]["id"] == matching.id + assert configs[1]["credentials"] == {"api_key": "*" * 20} + persisted = sqlite_session.scalar( + select(LoadBalancingModelConfig).where(LoadBalancingModelConfig.name == "__inherit__") + ) + assert persisted is not None + assert persisted.tenant_id == "tenant-1" - # Assert - assert is_enabled is False - assert configs[0]["name"] == "__inherit__" + +def test_get_configs_returns_empty_for_noncustom_provider( + monkeypatch: pytest.MonkeyPatch, + service: ServiceFixture, + sqlite_session: Session, +) -> None: + svc, _, _ = service + monkeypatch.setattr( + "services.model_load_balancing_service.encrypter.get_decrypt_decoding", lambda _tenant: ("rsa", "cipher") + ) + enabled, configs = svc.get_load_balancing_configs( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, session=sqlite_session + ) + assert enabled is False + assert configs == [] + + +def test_get_configs_reorders_existing_inherit_and_tolerates_bad_credentials( + monkeypatch: pytest.MonkeyPatch, + service: ServiceFixture, + sqlite_session: Session, +) -> None: + svc, _, configuration = service + configuration.custom_configuration.provider = CustomProviderConfiguration(credentials={}) + _config(sqlite_session, name="normal", encrypted_config='{"api_key":"bad"}') + _config(sqlite_session, name="__inherit__", encrypted_config="not-json", enabled=False) + monkeypatch.setattr( + "services.model_load_balancing_service.encrypter.get_decrypt_decoding", lambda _tenant: ("rsa", "cipher") + ) + monkeypatch.setattr( + "services.model_load_balancing_service.encrypter.decrypt_token_with_decoding", + MagicMock(side_effect=ValueError("cannot decrypt")), + ) + monkeypatch.setattr( + "services.model_load_balancing_service.LBModelManager.get_config_in_cooldown_and_ttl", + lambda **_kwargs: (True, 15), + ) + _, configs = svc.get_load_balancing_configs( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, session=sqlite_session + ) + assert [config["name"] for config in configs] == ["__inherit__", "normal"] assert configs[0]["credentials"] == {} - assert configs[1]["credentials"] == {"api_key": "bad-encrypted"} + assert configs[1]["credentials"] == {"api_key": "*" * 20} assert configs[1]["in_cooldown"] is True - assert configs[1]["ttl"] == 15 -def test_get_load_balancing_config_should_raise_value_error_when_provider_missing( - service: ModelLoadBalancingService, -) -> None: - # Arrange - service.provider_manager.get_configurations.return_value = {} - - # Act + Assert - with pytest.raises(ValueError, match="Provider openai does not exist"): - service.get_load_balancing_config( - "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=MagicMock() +def test_get_single_config_is_tenant_scoped_and_obfuscated(service: ServiceFixture, sqlite_session: Session) -> None: + svc, _, _ = service + config = _config(sqlite_session, encrypted_config='{"api_key":"secret"}') + assert svc.get_load_balancing_config( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, config.id, session=sqlite_session + ) == {"id": config.id, "name": "primary", "credentials": {"api_key": "*" * 20}, "enabled": True} + assert ( + svc.get_load_balancing_config( + "tenant-2", "openai", "gpt-4o-mini", ModelType.LLM, config.id, session=sqlite_session ) - - -def test_get_load_balancing_config_should_return_none_when_config_not_found( - service: ModelLoadBalancingService, - mock_db: MagicMock, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - mock_db.session.scalar.return_value = None - - # Act - result = service.get_load_balancing_config( - "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=mock_db.session + is None ) - # Assert - assert result is None - -def test_get_load_balancing_config_should_return_obfuscated_payload_when_config_exists( - service: ModelLoadBalancingService, - mock_db: MagicMock, +def test_init_inherit_config_persists_and_sql_failure_rolls_back( + service: ServiceFixture, + sqlite_session: Session, ) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - provider_configuration.obfuscated_credentials.side_effect = lambda credentials, credential_form_schemas: { - "masked": credentials.get("api_key", "") - } - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - config = SimpleNamespace(id="cfg-1", name="primary", encrypted_config="not-json", enabled=True) - mock_db.session.scalar.return_value = config - - # Act - result = service.get_load_balancing_config( - "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, "cfg-1", session=mock_db.session - ) - - # Assert - assert result == { - "id": "cfg-1", - "name": "primary", - "credentials": {"masked": ""}, - "enabled": True, - } + svc, _, _ = service + created = svc._init_inherit_config("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, sqlite_session) + assert sqlite_session.get(LoadBalancingModelConfig, created.id) is not None + sqlite_session.delete(created) + sqlite_session.commit() + sqlite_engine = cast(Engine, sqlite_session.get_bind()) + with _raise_on_insert(sqlite_engine), pytest.raises(RuntimeError, match="forced INSERT"): + svc._init_inherit_config("tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, sqlite_session) + sqlite_session.rollback() + assert sqlite_session.scalar(select(LoadBalancingModelConfig)) is None -def test_init_inherit_config_should_create_and_persist_inherit_configuration( - service: ModelLoadBalancingService, - mock_db: MagicMock, +@pytest.mark.parametrize( + ("configs", "message"), + [ + ("invalid", "Invalid load balancing configs"), + (["invalid"], "Invalid load balancing config"), + ([{"enabled": True}], "Invalid load balancing config name"), + ([{"name": "missing-enabled"}], "Invalid load balancing config enabled"), + ], +) +def test_update_configs_rejects_invalid_payloads( + configs, + message: str, + service: ServiceFixture, + sqlite_session: Session, ) -> None: - # Arrange - model_type = ModelType.LLM - - # Act - inherit_config = service._init_inherit_config( - "tenant-1", "openai", "gpt-4o-mini", model_type, session=mock_db.session - ) - - # Assert - assert inherit_config.tenant_id == "tenant-1" - assert inherit_config.provider_name == "openai" - assert inherit_config.model_name == "gpt-4o-mini" - assert inherit_config.model_type == "llm" - assert inherit_config.name == "__inherit__" - mock_db.session.add.assert_called_once_with(inherit_config) - mock_db.session.commit.assert_called_once() - - -def test_update_load_balancing_configs_should_raise_value_error_when_provider_missing( - service: ModelLoadBalancingService, -) -> None: - # Arrange - service.provider_manager.get_configurations.return_value = {} - - # Act + Assert - with pytest.raises(ValueError, match="Provider openai does not exist"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [], - "custom-model", - session=MagicMock(), + svc, _, _ = service + with pytest.raises(ValueError, match=message): + svc.update_load_balancing_configs( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, configs, "custom-model", sqlite_session ) -def test_update_load_balancing_configs_should_raise_value_error_when_configs_is_not_list( - service: ModelLoadBalancingService, +@pytest.mark.parametrize("existing", [True, False]) +def test_update_configs_rejects_invalid_credentials_for_existing_and_new_configs( + existing: bool, + service: ServiceFixture, + sqlite_session: Session, ) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - - # Act + Assert - with pytest.raises(ValueError, match="Invalid load balancing configs"): - service.update_load_balancing_configs( # type: ignore[arg-type] - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - cast(list[dict[str, object]], "invalid-configs"), - "custom-model", - session=MagicMock(), - ) - - -def test_update_load_balancing_configs_should_raise_value_error_when_config_item_is_not_dict( - service: ModelLoadBalancingService, - mock_db: MagicMock, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - mock_db.session.scalars.return_value.all.return_value = [] - - # Act + Assert - with pytest.raises(ValueError, match="Invalid load balancing config"): - service.update_load_balancing_configs( # type: ignore[list-item] - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - cast(list[dict[str, object]], ["bad-item"]), - "custom-model", - session=mock_db.session, - ) - - -def test_update_load_balancing_configs_should_raise_value_error_when_credential_id_not_found( - service: ModelLoadBalancingService, - mock_db: MagicMock, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - mock_db.session.scalars.return_value.all.return_value = [] - mock_db.session.scalar.return_value = None - - # Act + Assert - with pytest.raises(ValueError, match="Provider credential with id cred-1 not found"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"credential_id": "cred-1", "enabled": True}], - "predefined-model", - session=mock_db.session, - ) - - -def test_update_load_balancing_configs_should_raise_value_error_when_name_or_enabled_is_invalid( - service: ModelLoadBalancingService, - mock_db: MagicMock, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - mock_db.session.scalars.return_value.all.return_value = [] - - # Act + Assert - with pytest.raises(ValueError, match="Invalid load balancing config name"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"enabled": True}], - "custom-model", - session=mock_db.session, - ) - - with pytest.raises(ValueError, match="Invalid load balancing config enabled"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"name": "cfg-without-enabled"}], - "custom-model", - session=mock_db.session, - ) - - -def test_update_load_balancing_configs_should_raise_value_error_when_existing_config_id_is_invalid( - service: ModelLoadBalancingService, - mock_db: MagicMock, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - current_config = SimpleNamespace(id="cfg-1") - mock_db.session.scalars.return_value.all.return_value = [current_config] - - # Act + Assert - with pytest.raises(ValueError, match="Invalid load balancing config id: cfg-2"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"id": "cfg-2", "name": "invalid", "enabled": True}], - "custom-model", - session=mock_db.session, - ) - - -def test_update_load_balancing_configs_should_raise_value_error_when_credentials_are_invalid_for_update_or_create( - service: ModelLoadBalancingService, - mock_db: MagicMock, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - existing_config = SimpleNamespace(id="cfg-1", name="old", enabled=True, encrypted_config=None, updated_at=None) - mock_db.session.scalars.return_value.all.return_value = [existing_config] - - # Act + Assert - with pytest.raises(ValueError, match="Invalid load balancing config credentials"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"id": "cfg-1", "name": "new", "enabled": True, "credentials": "bad"}], - "custom-model", - session=mock_db.session, - ) + svc, _, _ = service + config = _config(sqlite_session) if existing else None + payload = {"name": "new", "enabled": True, "credentials": "bad"} + if config is not None: + payload["id"] = config.id with pytest.raises(ValueError, match="Invalid load balancing config credentials"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"name": "new-config", "enabled": True, "credentials": "bad"}], - "custom-model", - session=mock_db.session, + svc.update_load_balancing_configs( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, [payload], "custom-model", sqlite_session ) -def test_update_load_balancing_configs_should_update_existing_create_new_and_delete_removed_configs( - service: ModelLoadBalancingService, - mock_db: MagicMock, - mocker: MockerFixture, +def test_update_configs_updates_creates_and_deletes_persisted_rows( + monkeypatch: pytest.MonkeyPatch, + service: ServiceFixture, + sqlite_session: Session, ) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - existing_config_1 = SimpleNamespace( - id="cfg-1", - name="existing-one", - enabled=True, - encrypted_config=json.dumps({"api_key": "old"}), - updated_at=None, + svc, _, _ = service + keep = _config(sqlite_session, name="keep", encrypted_config='{"api_key":"old"}') + removed = _config(sqlite_session, name="remove") + monkeypatch.setattr( + svc, + "_custom_credentials_validate", + lambda **kwargs: {"api_key": f"enc-{kwargs['credentials']['api_key']}"}, ) - existing_config_2 = SimpleNamespace( - id="cfg-2", - name="existing-two", - enabled=True, - encrypted_config=None, - updated_at=None, - ) - mock_db.session.scalars.return_value.all.return_value = [existing_config_1, existing_config_2] - mocker.patch.object(service, "_custom_credentials_validate", return_value={"api_key": "encrypted"}) - mock_clear_cache = mocker.patch.object(service, "_clear_credentials_cache") - - # Act - service.update_load_balancing_configs( + svc.update_load_balancing_configs( "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, [ - {"id": "cfg-1", "name": "updated-name", "enabled": False, "credentials": {"api_key": "plain"}}, - {"name": "new-config", "enabled": True, "credentials": {"api_key": "plain"}}, + {"id": keep.id, "name": "updated", "enabled": False, "credentials": {"api_key": "new"}}, + {"name": "created", "enabled": True, "credentials": {"api_key": "fresh"}}, ], "custom-model", - session=mock_db.session, + sqlite_session, ) - - # Assert - assert existing_config_1.name == "updated-name" - assert existing_config_1.enabled is False - assert json.loads(existing_config_1.encrypted_config) == {"api_key": "encrypted"} - assert mock_db.session.add.call_count == 1 - mock_db.session.delete.assert_called_once_with(existing_config_2) - assert mock_db.session.commit.call_count >= 3 - mock_clear_cache.assert_any_call("tenant-1", "cfg-1") - mock_clear_cache.assert_any_call("tenant-1", "cfg-2") + sqlite_session.expire_all() + records = sqlite_session.scalars(select(LoadBalancingModelConfig)).all() + assert {record.name for record in records} == {"updated", "created"} + assert sqlite_session.get(LoadBalancingModelConfig, removed.id) is None + updated = sqlite_session.get(LoadBalancingModelConfig, keep.id) + assert updated is not None + assert updated.enabled is False + assert json.loads(updated.encrypted_config) == {"api_key": "enc-new"} -def test_update_load_balancing_configs_should_raise_value_error_for_invalid_new_config_name_or_missing_credentials( - service: ModelLoadBalancingService, - mock_db: MagicMock, +def test_update_configs_creates_from_tenant_scoped_provider_credential( + service: ServiceFixture, sqlite_session: Session ) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - mock_db.session.scalars.return_value.all.return_value = [] - - # Act + Assert - with pytest.raises(ValueError, match="Invalid load balancing config name"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"name": "__inherit__", "enabled": True, "credentials": {"api_key": "x"}}], - "custom-model", - session=mock_db.session, - ) - - with pytest.raises(ValueError, match="Invalid load balancing config credentials"): - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"name": "new", "enabled": True}], - "custom-model", - session=mock_db.session, - ) - - -def test_update_load_balancing_configs_should_create_from_existing_provider_credential_when_credential_id_provided( - service: ModelLoadBalancingService, - mock_db: MagicMock, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - mock_db.session.scalars.return_value.all.return_value = [] - credential_record = SimpleNamespace(credential_name="Main Credential", encrypted_config='{"api_key":"enc"}') - mock_db.session.scalar.return_value = credential_record - - # Act - service.update_load_balancing_configs( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - [{"credential_id": "cred-1", "enabled": True}], - "predefined-model", - session=mock_db.session, - ) - - # Assert - created_config = mock_db.session.add.call_args.args[0] - assert created_config.name == "Main Credential" - assert created_config.credential_id == "cred-1" - assert created_config.credential_source_type == "provider" - assert created_config.encrypted_config == '{"api_key":"enc"}' - mock_db.session.commit.assert_called() - - -def test_validate_load_balancing_credentials_should_raise_value_error_when_provider_missing( - service: ModelLoadBalancingService, -) -> None: - # Arrange - service.provider_manager.get_configurations.return_value = {} - - # Act + Assert - with pytest.raises(ValueError, match="Provider openai does not exist"): - service.validate_load_balancing_credentials( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - {"api_key": "plain"}, - session=MagicMock(), - ) - - -def test_validate_load_balancing_credentials_should_raise_value_error_when_config_id_is_invalid( - service: ModelLoadBalancingService, - mock_db: MagicMock, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - mock_db.session.scalar.return_value = None - - # Act + Assert - with pytest.raises(ValueError, match="Load balancing config cfg-1 does not exist"): - service.validate_load_balancing_credentials( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - {"api_key": "plain"}, - config_id="cfg-1", - session=mock_db.session, - ) - - -def test_validate_load_balancing_credentials_should_delegate_to_custom_validate_with_or_without_config( - service: ModelLoadBalancingService, - mock_db: MagicMock, - mocker: MockerFixture, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - service.provider_manager.get_configurations.return_value = {"openai": provider_configuration} - existing_config = SimpleNamespace(id="cfg-1") - mock_db.session.scalar.return_value = existing_config - mock_validate = mocker.patch.object(service, "_custom_credentials_validate") - - # Act - service.validate_load_balancing_credentials( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - {"api_key": "plain"}, - config_id="cfg-1", - session=mock_db.session, - ) - service.validate_load_balancing_credentials( - "tenant-1", - "openai", - "gpt-4o-mini", - ModelType.LLM, - {"api_key": "plain"}, - session=mock_db.session, - ) - - # Assert - assert mock_validate.call_count == 2 - assert mock_validate.call_args_list[0].kwargs["load_balancing_model_config"] is existing_config - assert mock_validate.call_args_list[1].kwargs["load_balancing_model_config"] is None - shared_model_provider_factory = service.model_assembly.model_provider_factory - assert mock_validate.call_args_list[0].kwargs["model_provider_factory"] is shared_model_provider_factory - assert mock_validate.call_args_list[1].kwargs["model_provider_factory"] is shared_model_provider_factory - - -def test_custom_credentials_validate_should_replace_hidden_secret_with_original_value_and_encrypt( - service: ModelLoadBalancingService, - mocker: MockerFixture, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - load_balancing_model_config = _load_balancing_model_config( - encrypted_config=json.dumps({"api_key": "old-encrypted-token"}) - ) - mocker.patch("services.model_load_balancing_service.encrypter.decrypt_token", return_value="old-plain-value") - mock_encrypt = mocker.patch( - "services.model_load_balancing_service.encrypter.encrypt_token", - side_effect=lambda tenant_id, value: f"enc:{value}", - ) - - # Act - result = service._custom_credentials_validate( + svc, _, _ = service + credential = ProviderCredential( tenant_id="tenant-1", - provider_configuration=provider_configuration, - model_type=ModelType.LLM, - model="gpt-4o-mini", - credentials={"api_key": HIDDEN_VALUE, "region": "us"}, - load_balancing_model_config=load_balancing_model_config, + provider_name="openai", + credential_name="Credential", + encrypted_config='{"api_key":"enc"}', + ) + foreign = ProviderCredential( + tenant_id="tenant-2", + provider_name="openai", + credential_name="Foreign", + encrypted_config="{}", + ) + sqlite_session.add_all([credential, foreign]) + sqlite_session.commit() + svc.update_load_balancing_configs( + "tenant-1", + "openai", + "gpt-4o-mini", + ModelType.LLM, + [{"credential_id": credential.id, "enabled": True}], + "predefined-model", + sqlite_session, + ) + created = sqlite_session.scalar(select(LoadBalancingModelConfig)) + assert created is not None + assert created.name == "Credential" + assert created.credential_id == credential.id + assert created.credential_source_type == CredentialSourceType.PROVIDER + with pytest.raises(ValueError, match="not found"): + svc.update_load_balancing_configs( + "tenant-1", + "openai", + "other-model", + ModelType.LLM, + [{"credential_id": foreign.id, "enabled": True}], + "predefined-model", + sqlite_session, + ) + + +def test_validate_credentials_uses_real_config_lookup( + monkeypatch: pytest.MonkeyPatch, + service: ServiceFixture, + sqlite_session: Session, +) -> None: + svc, manager, _ = service + config = _config(sqlite_session) + assembly = PluginModelAssembly(tenant_id="tenant-1") + assembly._provider_manager = manager + assembly._model_provider_factory = ModelProviderFactory(runtime=cast(ModelRuntime, object())) + monkeypatch.setattr( + "services.model_load_balancing_service.create_plugin_model_assembly", lambda **_kwargs: assembly + ) + validate = MagicMock() + monkeypatch.setattr(svc, "_custom_credentials_validate", validate) + svc.validate_load_balancing_credentials( + "tenant-1", + "openai", + "gpt-4o-mini", + ModelType.LLM, + {"api_key": "raw"}, + sqlite_session, + config.id, + ) + assert validate.call_args.kwargs["load_balancing_model_config"].id == config.id + with pytest.raises(ValueError, match="does not exist"): + svc.validate_load_balancing_credentials( + "tenant-1", "openai", "gpt-4o-mini", ModelType.LLM, {}, sqlite_session, "missing" + ) + + +def test_custom_credentials_validate_reuses_hidden_secret_and_encrypts( + monkeypatch: pytest.MonkeyPatch, + service: ServiceFixture, + sqlite_session: Session, +) -> None: + svc, _, configuration = service + config = _config(sqlite_session, encrypted_config='{"api_key":"old-encrypted"}') + monkeypatch.setattr("services.model_load_balancing_service.encrypter.decrypt_token", lambda *_args: "old-plain") + monkeypatch.setattr( + "services.model_load_balancing_service.encrypter.encrypt_token", lambda _tenant, value: f"enc-{value}" + ) + result = svc._custom_credentials_validate( + "tenant-1", + configuration, + ModelType.LLM, + "gpt-4o-mini", + {"api_key": HIDDEN_VALUE}, + config, validate=False, ) - - # Assert - assert result == {"api_key": "enc:old-plain-value", "region": "us"} - mock_encrypt.assert_called_once_with("tenant-1", "old-plain-value") + assert result == {"api_key": "enc-old-plain"} -def test_custom_credentials_validate_should_handle_invalid_original_json_and_validate_with_model_schema( - service: ModelLoadBalancingService, - mocker: MockerFixture, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(model_schema=_build_model_credential_schema()) - load_balancing_model_config = _load_balancing_model_config(encrypted_config="not-json") - mock_factory = MagicMock() - mock_factory.model_credentials_validate.return_value = {"api_key": "validated"} - mock_encrypt = mocker.patch( - "services.model_load_balancing_service.encrypter.encrypt_token", - side_effect=lambda tenant_id, value: f"enc:{value}", - ) - - # Act - result = service._custom_credentials_validate( - tenant_id="tenant-1", - provider_configuration=provider_configuration, - model_type=ModelType.LLM, - model="gpt-4o-mini", - credentials={"api_key": "plain"}, - load_balancing_model_config=load_balancing_model_config, - model_provider_factory=mock_factory, - validate=True, - ) - - # Assert - assert result == {"api_key": "enc:validated"} - mock_factory.model_credentials_validate.assert_called_once() - mock_factory.provider_credentials_validate.assert_not_called() - mock_encrypt.assert_called_once_with("tenant-1", "validated") - - -def test_custom_credentials_validate_should_validate_with_provider_schema_when_model_schema_absent( - service: ModelLoadBalancingService, - mocker: MockerFixture, -) -> None: - # Arrange - provider_configuration = _build_provider_configuration(provider_schema=_build_provider_credential_schema()) - mock_factory = MagicMock() - mock_factory.provider_credentials_validate.return_value = {"api_key": "provider-validated"} - mocker.patch( - "services.model_load_balancing_service.encrypter.encrypt_token", - side_effect=lambda tenant_id, value: f"enc:{value}", - ) - - # Act - result = service._custom_credentials_validate( - tenant_id="tenant-1", - provider_configuration=provider_configuration, - model_type=ModelType.LLM, - model="gpt-4o-mini", - credentials={"api_key": "plain"}, - model_provider_factory=mock_factory, - validate=True, - ) - - # Assert - assert result == {"api_key": "enc:provider-validated"} - mock_factory.provider_credentials_validate.assert_called_once() - mock_factory.model_credentials_validate.assert_not_called() - - -def test_get_credential_schema_should_return_model_schema_or_provider_schema_or_raise( - service: ModelLoadBalancingService, -) -> None: - # Arrange - model_schema = _build_model_credential_schema() - provider_schema = _build_provider_credential_schema() - provider_configuration_with_model = _build_provider_configuration(model_schema=model_schema) - provider_configuration_with_provider = _build_provider_configuration(provider_schema=provider_schema) - provider_configuration_without_schema = _build_provider_configuration() - - # Act - schema_from_model = service._get_credential_schema(provider_configuration_with_model) - schema_from_provider = service._get_credential_schema(provider_configuration_with_provider) - - # Assert - assert schema_from_model is model_schema - assert schema_from_provider is provider_schema - with pytest.raises(ValueError, match="No credential schema found"): - service._get_credential_schema(provider_configuration_without_schema) - - -def test_clear_credentials_cache_should_delete_load_balancing_cache_entry( - service: ModelLoadBalancingService, - mocker: MockerFixture, -) -> None: - # Arrange - mock_cache_instance = MagicMock() - mock_cache_cls = mocker.patch( - "services.model_load_balancing_service.ProviderCredentialsCache", - return_value=mock_cache_instance, - ) - - # Act - service._clear_credentials_cache("tenant-1", "cfg-1") - - # Assert - mock_cache_cls.assert_called_once() - assert mock_cache_cls.call_args.kwargs == { - "tenant_id": "tenant-1", - "identity_id": "cfg-1", - "cache_type": mocker.ANY, - } - assert mock_cache_cls.call_args.kwargs["cache_type"].name == "LOAD_BALANCING_MODEL" - mock_cache_instance.delete.assert_called_once() +def test_schema_selection_and_cache_boundary(service: ServiceFixture) -> None: + svc, _, configuration = service + provider_schema = configuration.provider.provider_credential_schema + assert svc._get_credential_schema(configuration) is provider_schema + configuration.provider.model_credential_schema = _model_schema() + assert isinstance(svc._get_credential_schema(configuration), ModelCredentialSchema) + configuration.provider.model_credential_schema = None + configuration.provider.provider_credential_schema = None + with pytest.raises(ValueError, match="No credential schema"): + svc._get_credential_schema(configuration) + with patch("services.model_load_balancing_service.ProviderCredentialsCache") as cache: + svc._clear_credentials_cache("tenant-1", "config-1") + cache.return_value.delete.assert_called_once() diff --git a/api/tests/unit_tests/services/test_model_provider_service.py b/api/tests/unit_tests/services/test_model_provider_service.py index a8a976f4b07..597d18f49d4 100644 --- a/api/tests/unit_tests/services/test_model_provider_service.py +++ b/api/tests/unit_tests/services/test_model_provider_service.py @@ -5,12 +5,17 @@ from unittest.mock import MagicMock import pytest from core.entities.model_entities import ModelStatus +from core.entities.provider_entities import CredentialConfiguration +from core.plugin.entities.plugin import PluginInstallationSource +from core.plugin.entities.plugin_daemon import PluginModelProviderBinding +from enums import DeploymentEdition from graphon.model_runtime.entities.common_entities import I18nObject from graphon.model_runtime.entities.model_entities import FetchFrom, ModelType, ParameterRule, ParameterType +from graphon.model_runtime.entities.provider_entities import ConfigurateMethod from models.provider import ProviderType from services import model_provider_service as service_module from services.errors.app_model_config import ProviderNotFoundError -from services.model_provider_service import ModelProviderService +from services.model_provider_service import ModelProviderService, _ProviderSummaryState def _create_service_with_mocked_manager() -> tuple[ModelProviderService, MagicMock]: @@ -59,6 +64,26 @@ def _build_provider_configuration( ) +def _build_model_provider_binding( + source: PluginInstallationSource, + *, + installation_id: str = "installation-1", + plugin_id: str = "langgenius/openai", + plugin_unique_identifier: str = "langgenius/openai:1.0.0@checksum", + verified: bool = True, +) -> PluginModelProviderBinding: + return PluginModelProviderBinding( + provider="openai", + installation_id=installation_id, + plugin_id=plugin_id, + plugin_unique_identifier=plugin_unique_identifier, + runtime_type="remote" if source == PluginInstallationSource.Remote else "local", + source=source, + version="1.0.0", + verified=verified, + ) + + class TestModelProviderServiceConfiguration: def test__get_provider_configuration_should_return_configuration_when_provider_exists(self) -> None: service, manager = _create_service_with_mocked_manager() @@ -96,6 +121,352 @@ class TestModelProviderServiceConfiguration: assert result[0].provider == "openai" assert result[0].custom_configuration.status.value == "no-configure" + def test_get_provider_summary_list_uses_lightweight_state_and_plugin_bindings( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + service = ModelProviderService() + provider = SimpleNamespace( + provider="langgenius/openai/openai", + label=I18nObject(en_US="OpenAI"), + description=I18nObject(en_US="OpenAI models"), + icon_small=I18nObject(en_US="icon.svg"), + icon_small_dark=I18nObject(en_US="icon-dark.svg"), + supported_model_types=[ModelType.LLM], + configurate_methods=[ConfigurateMethod.PREDEFINED_MODEL], + ) + binding = SimpleNamespace( + provider="openai", + plugin_id="langgenius/openai", + installation_id="installation-1", + plugin_unique_identifier="langgenius/openai:1.2.3@checksum", + runtime_type="local", + source=PluginInstallationSource.Marketplace, + version="1.2.3", + verified=True, + ) + state = _ProviderSummaryState( + has_custom_provider=True, + available_credentials=[ + CredentialConfiguration( + credential_id="credential-1", + credential_name="Production", + ), + CredentialConfiguration( + credential_id="credential-2", + credential_name="Backup", + ), + ], + has_custom_models=True, + current_credential_id="credential-1", + current_credential_name="Production", + current_credential_usable=True, + preferred_provider_type=ProviderType.CUSTOM, + ) + call_order: list[str] = [] + manager_constructor = MagicMock(side_effect=AssertionError("summary must not construct ProviderManager")) + monkeypatch.setattr(service, "_get_provider_manager", manager_constructor) + monkeypatch.setattr( + service_module.PluginService, + "list_model_provider_bindings", + MagicMock(side_effect=lambda *_args, **_kwargs: call_order.append("bindings") or [binding]), + ) + monkeypatch.setattr( + service_module.PluginService, + "fetch_plugin_model_providers", + MagicMock(side_effect=lambda *_args, **_kwargs: call_order.append("providers") or [provider]), + ) + monkeypatch.setattr( + service, + "_load_provider_summary_states", + MagicMock(return_value={provider.provider: state}), + ) + monkeypatch.setattr( + service_module.ext_hosting_provider.hosting_configuration, + "provider_map", + {provider.provider: SimpleNamespace(enabled=True, quotas=[SimpleNamespace()])}, + ) + + providers, plugins = service.get_provider_summary_list(tenant_id="tenant-1") + + assert len(providers) == 1 + assert providers[0].provider == provider.provider + assert providers[0].plugin_id == "langgenius/openai" + assert providers[0].is_configured is True + assert providers[0].custom_configuration.available_credentials == [ + CredentialConfiguration( + credential_id="credential-1", + credential_name="Production", + ), + CredentialConfiguration( + credential_id="credential-2", + credential_name="Backup", + ), + ] + assert providers[0].custom_configuration.has_custom_models is True + assert providers[0].custom_configuration.current_credential_name == "Production" + assert providers[0].custom_configuration.current_credential_usable is True + assert providers[0].system_configuration.enabled is True + assert plugins["langgenius/openai"].installation_id == "installation-1" + assert plugins["langgenius/openai"].version == "1.2.3" + assert call_order == ["bindings", "providers"] + manager_constructor.assert_not_called() + + def test_get_provider_summary_list_enables_system_only_for_verified_hosted_non_package_bindings( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + marketplace_binding = _build_model_provider_binding(PluginInstallationSource.Marketplace) + package_binding = _build_model_provider_binding( + PluginInstallationSource.Package, + installation_id="installation-package", + plugin_id="langgenius/package", + plugin_unique_identifier="langgenius/package:1.0.0@checksum", + ) + unhosted_binding = _build_model_provider_binding( + PluginInstallationSource.Marketplace, + installation_id="installation-unhosted", + plugin_id="langgenius/unhosted", + plugin_unique_identifier="langgenius/unhosted:1.0.0@checksum", + ) + remote_binding = _build_model_provider_binding( + PluginInstallationSource.Remote, + installation_id="installation-remote", + plugin_id="langgenius/remote", + plugin_unique_identifier="langgenius/remote:1.0.0@checksum", + ) + unverified_binding = _build_model_provider_binding( + PluginInstallationSource.Marketplace, + installation_id="installation-unverified", + plugin_id="langgenius/unverified", + plugin_unique_identifier="langgenius/unverified:1.0.0@checksum", + verified=False, + ) + bindings = [ + marketplace_binding, + package_binding, + unhosted_binding, + remote_binding, + unverified_binding, + ] + provider_entities = [ + SimpleNamespace( + provider=f"{binding.plugin_id}/openai", + label=I18nObject(en_US=binding.plugin_id), + description=None, + icon_small=None, + icon_small_dark=None, + supported_model_types=[ModelType.LLM], + configurate_methods=[], + ) + for binding in bindings + ] + monkeypatch.setattr( + service_module.ext_hosting_provider.hosting_configuration, + "provider_map", + { + "langgenius/openai/openai": SimpleNamespace(enabled=True, quotas=[SimpleNamespace()]), + "langgenius/package/openai": SimpleNamespace(enabled=True, quotas=[SimpleNamespace()]), + "langgenius/remote/openai": SimpleNamespace(enabled=True, quotas=[SimpleNamespace()]), + "langgenius/unverified/openai": SimpleNamespace(enabled=True, quotas=[SimpleNamespace()]), + }, + ) + monkeypatch.setattr( + service_module.PluginService, + "list_model_provider_bindings", + MagicMock(return_value=bindings), + ) + monkeypatch.setattr( + service_module.PluginService, + "fetch_plugin_model_providers", + MagicMock(return_value=provider_entities), + ) + monkeypatch.setattr(ModelProviderService, "_load_provider_summary_states", MagicMock(return_value={})) + monkeypatch.setattr(service_module, "is_filtered", MagicMock(return_value=False)) + + providers, _ = ModelProviderService().get_provider_summary_list("tenant-1") + + assert {provider.provider: provider.system_configuration.enabled for provider in providers} == { + "langgenius/openai/openai": True, + "langgenius/package/openai": False, + "langgenius/unhosted/openai": False, + "langgenius/remote/openai": True, + "langgenius/unverified/openai": False, + } + + def test_model_provider_binding_without_verified_field_fails_closed(self) -> None: + binding = PluginModelProviderBinding.model_validate( + { + "provider": "openai", + "installation_id": "installation-1", + "plugin_id": "langgenius/openai", + "plugin_unique_identifier": "langgenius/openai:1.0.0@checksum", + "runtime_type": "local", + "source": "marketplace", + "version": "1.0.0", + } + ) + + assert binding.verified is False + + def test_get_provider_summary_list_returns_all_unique_provider_metadata( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + service = ModelProviderService() + llm_provider = SimpleNamespace( + provider="langgenius/openai/openai", + label=I18nObject(en_US="OpenAI"), + description=None, + icon_small=None, + icon_small_dark=None, + supported_model_types=[ModelType.LLM], + configurate_methods=[], + ) + embedding_provider = SimpleNamespace( + provider="langgenius/embedding/embedding", + label=I18nObject(en_US="Embedding"), + description=None, + icon_small=None, + icon_small_dark=None, + supported_model_types=[ModelType.TEXT_EMBEDDING], + configurate_methods=[], + ) + llm_binding = SimpleNamespace( + provider="openai", + plugin_id="langgenius/openai", + installation_id="installation-openai", + plugin_unique_identifier="langgenius/openai:1.0.0@checksum", + runtime_type="local", + source=service_module.PluginInstallationSource.Marketplace, + version="1.0.0", + verified=False, + ) + embedding_binding = SimpleNamespace( + provider="embedding", + plugin_id="langgenius/embedding", + installation_id="installation-embedding", + plugin_unique_identifier="langgenius/embedding:1.0.0@checksum", + runtime_type="local", + source=service_module.PluginInstallationSource.Marketplace, + version="1.0.0", + verified=False, + ) + monkeypatch.setattr( + service_module.PluginService, + "list_model_provider_bindings", + MagicMock(return_value=[llm_binding, embedding_binding]), + ) + monkeypatch.setattr( + service_module.PluginService, + "fetch_plugin_model_providers", + MagicMock(return_value=[llm_provider, llm_provider, embedding_provider]), + ) + monkeypatch.setattr(service, "_load_provider_summary_states", MagicMock(return_value={})) + monkeypatch.setattr(service_module, "is_filtered", MagicMock(return_value=False)) + + providers, plugins = service.get_provider_summary_list(tenant_id="tenant-1") + + assert [provider.provider for provider in providers] == [ + "langgenius/openai/openai", + "langgenius/embedding/embedding", + ] + assert providers[0].is_configured is False + assert providers[0].custom_configuration.status.value == "no-configure" + assert providers[0].custom_configuration.has_custom_models is False + assert providers[0].custom_configuration.available_credentials == [] + assert set(plugins) == {"langgenius/openai", "langgenius/embedding"} + + def test_preferred_provider_fallback_uses_custom_presence_not_configuration_status( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.setattr(service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + state = _ProviderSummaryState(has_custom_provider=True) + + preferred_provider_type = ModelProviderService._get_preferred_provider_type( + state, + custom_present=True, + system_enabled=True, + ) + + assert preferred_provider_type == ProviderType.CUSTOM + + def test_load_provider_summary_states_reads_only_lightweight_columns(self, monkeypatch: pytest.MonkeyPatch) -> None: + canonical_provider = "langgenius/openai/openai" + session = MagicMock() + session.execute.side_effect = [ + SimpleNamespace( + all=lambda: [ + SimpleNamespace( + provider_name="openai", + credential_id="credential-legacy", + credential_provider_name="openai", + credential_name="Legacy", + ), + SimpleNamespace( + provider_name=canonical_provider, + credential_id="credential-current", + credential_provider_name=canonical_provider, + credential_name="Production", + ), + ] + ), + SimpleNamespace( + all=lambda: [ + SimpleNamespace( + id="credential-legacy", + provider_name="openai", + credential_name="Legacy", + ), + SimpleNamespace( + id="credential-current", + provider_name=canonical_provider, + credential_name="Production", + ), + ] + ), + SimpleNamespace(all=lambda: [SimpleNamespace(provider_name="openai")]), + SimpleNamespace( + all=lambda: [ + SimpleNamespace( + provider_name=canonical_provider, + preferred_provider_type=ProviderType.SYSTEM, + ) + ] + ), + ] + session_context = MagicMock() + session_context.__enter__.return_value = session + create_session = MagicMock(return_value=session_context) + monkeypatch.setattr(service_module.session_factory, "create_session", create_session) + + states = ModelProviderService._load_provider_summary_states("tenant-1") + + state = states[canonical_provider] + assert state.has_custom_provider is True + assert state.available_credentials == [ + CredentialConfiguration( + credential_id="credential-legacy", + credential_name="Legacy", + ), + CredentialConfiguration( + credential_id="credential-current", + credential_name="Production", + ), + ] + assert state.has_custom_models is True + assert state.current_credential_id == "credential-current" + assert state.current_credential_name == "Production" + assert state.current_credential_usable is True + assert state.preferred_provider_type == ProviderType.SYSTEM + + statements = [str(execute_call.args[0]) for execute_call in session.execute.call_args_list] + assert len(statements) == 4 + assert all("encrypted_config" not in statement for statement in statements) + assert "count(" not in statements[1].lower() + assert "provider_credentials.id" in statements[1] + assert "provider_credentials.credential_name" in statements[1] + assert "ORDER BY provider_credentials.created_at DESC, provider_credentials.id DESC" in statements[1] + assert "provider_model_credentials" in statements[2] + def test_get_models_by_provider_should_wrap_model_entities_with_tenant_context(self) -> None: service, manager = _create_service_with_mocked_manager() @@ -377,7 +748,7 @@ class TestModelProviderServiceDelegation: { "tenant_id": "tenant-1", "provider": "openai", - "model_type": "text-generation", + "model_type": "llm", "model": "gpt-4o", "credential_id": "cred-1", }, @@ -389,7 +760,7 @@ class TestModelProviderServiceDelegation: { "tenant_id": "tenant-1", "provider": "openai", - "model_type": "text-generation", + "model_type": "llm", "model": "gpt-4o", "credentials": {"api_key": "x"}, "credential_name": "cred-a", @@ -407,7 +778,7 @@ class TestModelProviderServiceDelegation: { "tenant_id": "tenant-1", "provider": "openai", - "model_type": "text-generation", + "model_type": "llm", "model": "gpt-4o", }, "delete_custom_model", diff --git a/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py b/api/tests/unit_tests/services/test_oauth_server_service.py similarity index 78% rename from api/tests/test_containers_integration_tests/services/test_oauth_server_service.py rename to api/tests/unit_tests/services/test_oauth_server_service.py index 7ea52ff6fef..ebc7b6501fe 100644 --- a/api/tests/test_containers_integration_tests/services/test_oauth_server_service.py +++ b/api/tests/unit_tests/services/test_oauth_server_service.py @@ -1,16 +1,20 @@ -"""Testcontainers integration tests for OAuthServerService.""" +"""Unit tests for OAuthServerService with SQLite-backed database access.""" from __future__ import annotations import uuid +from collections.abc import Iterator from typing import cast from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest +from flask import Flask +from sqlalchemy import Engine from sqlalchemy.orm import Session from werkzeug.exceptions import BadRequest +from models.engine import db from models.model import OAuthProviderApp from services.oauth_server import ( OAUTH_ACCESS_TOKEN_EXPIRES_IN, @@ -23,10 +27,24 @@ from services.oauth_server import ( ) -class TestOAuthServerServiceGetProviderApp: - """DB-backed tests for get_oauth_provider_app.""" +@pytest.fixture +def oauth_db() -> Iterator[Session]: + """Provide the production database extension with an isolated SQLite provider table.""" + app = Flask(__name__) + app.config["SQLALCHEMY_DATABASE_URI"] = "sqlite:///:memory:" + db.init_app(app) - def _create_oauth_provider_app(self, db_session_with_containers: Session, *, client_id: str) -> OAuthProviderApp: + with app.app_context(): + OAuthProviderApp.__table__.create(db.engine) + with Session(db.engine, expire_on_commit=False) as session: + yield session + + +class TestOAuthServerServiceGetProviderApp: + """Verify provider lookup against a real SQLAlchemy database.""" + + def test_get_oauth_provider_app_returns_app_when_exists(self, oauth_db: Session) -> None: + client_id = f"client-{uuid4()}" app = OAuthProviderApp( app_icon="icon.png", client_id=client_id, @@ -35,35 +53,30 @@ class TestOAuthServerServiceGetProviderApp: redirect_uris=["https://example.com/callback"], scope="read", ) - db_session_with_containers.add(app) - db_session_with_containers.commit() - return app - - def test_get_oauth_provider_app_returns_app_when_exists(self, db_session_with_containers: Session): - client_id = f"client-{uuid4()}" - created = self._create_oauth_provider_app(db_session_with_containers, client_id=client_id) + oauth_db.add(app) + oauth_db.commit() result = OAuthServerService.get_oauth_provider_app(client_id) assert result is not None assert result.client_id == client_id - assert result.id == created.id + assert result.id == app.id - def test_get_oauth_provider_app_returns_none_when_not_exists(self, db_session_with_containers: Session): + def test_get_oauth_provider_app_returns_none_when_not_exists(self, oauth_db: Session) -> None: result = OAuthServerService.get_oauth_provider_app(f"nonexistent-{uuid4()}") assert result is None class TestOAuthServerServiceTokenOperations: - """Redis-backed tests for token sign/validate operations.""" + """Verify Redis-backed token signing and validation branches.""" @pytest.fixture def mock_redis(self): with patch("services.oauth_server.redis_client") as mock: yield mock - def test_sign_authorization_code_stores_and_returns_code(self, mock_redis): + def test_sign_authorization_code_stores_and_returns_code(self, mock_redis) -> None: deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000111") with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid): code = OAuthServerService.sign_oauth_authorization_code("client-1", "user-1") @@ -75,7 +88,7 @@ class TestOAuthServerServiceTokenOperations: ex=600, ) - def test_sign_access_token_raises_bad_request_for_invalid_code(self, mock_redis): + def test_sign_access_token_raises_bad_request_for_invalid_code(self, mock_redis) -> None: mock_redis.get.return_value = None with pytest.raises(BadRequest, match="invalid code"): @@ -85,14 +98,13 @@ class TestOAuthServerServiceTokenOperations: client_id="client-1", ) - def test_sign_access_token_issues_tokens_for_valid_code(self, mock_redis): + def test_sign_access_token_issues_tokens_for_valid_code(self, mock_redis) -> None: token_uuids = [ uuid.UUID("00000000-0000-0000-0000-000000000201"), uuid.UUID("00000000-0000-0000-0000-000000000202"), ] with patch("services.oauth_server.uuid.uuid4", side_effect=token_uuids): mock_redis.get.return_value = b"user-1" - access_token, refresh_token = OAuthServerService.sign_oauth_access_token( grant_type=OAuthGrantType.AUTHORIZATION_CODE, code="code-1", @@ -114,7 +126,7 @@ class TestOAuthServerServiceTokenOperations: ex=OAUTH_REFRESH_TOKEN_EXPIRES_IN, ) - def test_sign_access_token_raises_bad_request_for_invalid_refresh_token(self, mock_redis): + def test_sign_access_token_raises_bad_request_for_invalid_refresh_token(self, mock_redis) -> None: mock_redis.get.return_value = None with pytest.raises(BadRequest, match="invalid refresh token"): @@ -124,11 +136,10 @@ class TestOAuthServerServiceTokenOperations: client_id="client-1", ) - def test_sign_access_token_issues_new_token_for_valid_refresh(self, mock_redis): + def test_sign_access_token_issues_new_token_for_valid_refresh(self, mock_redis) -> None: deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000301") with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid): mock_redis.get.return_value = b"user-1" - access_token, returned_refresh = OAuthServerService.sign_oauth_access_token( grant_type=OAuthGrantType.REFRESH_TOKEN, refresh_token="refresh-1", @@ -138,14 +149,14 @@ class TestOAuthServerServiceTokenOperations: assert access_token == str(deterministic_uuid) assert returned_refresh == "refresh-1" - def test_sign_access_token_returns_none_for_unknown_grant_type(self, mock_redis): + def test_sign_access_token_returns_none_for_unknown_grant_type(self, mock_redis) -> None: grant_type = cast(OAuthGrantType, "invalid-grant-type") result = OAuthServerService.sign_oauth_access_token(grant_type=grant_type, client_id="client-1") assert result is None - def test_sign_refresh_token_stores_with_expected_expiry(self, mock_redis): + def test_sign_refresh_token_stores_with_expected_expiry(self, mock_redis) -> None: deterministic_uuid = uuid.UUID("00000000-0000-0000-0000-000000000401") with patch("services.oauth_server.uuid.uuid4", return_value=deterministic_uuid): refresh_token = OAuthServerService._sign_oauth_refresh_token("client-2", "user-2") @@ -157,22 +168,21 @@ class TestOAuthServerServiceTokenOperations: ex=OAUTH_REFRESH_TOKEN_EXPIRES_IN, ) - def test_validate_access_token_returns_none_when_not_found(self, mock_redis, db_session_with_containers: Session): + def test_validate_access_token_returns_none_when_not_found(self, mock_redis, sqlite_engine: Engine) -> None: mock_redis.get.return_value = None - session = MagicMock() - result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", db_session_with_containers) + with Session(sqlite_engine) as session: + result = OAuthServerService.validate_oauth_access_token("client-1", "missing-token", session) assert result is None - def test_validate_access_token_loads_user_when_exists(self, mock_redis, db_session_with_containers: Session): + def test_validate_access_token_loads_user_when_exists(self, mock_redis, sqlite_engine: Engine) -> None: mock_redis.get.return_value = b"user-88" expected_user = MagicMock() - with patch("services.oauth_server.AccountService.load_user", return_value=expected_user) as mock_load: - result = OAuthServerService.validate_oauth_access_token( - "client-1", "access-token", db_session_with_containers - ) + with Session(sqlite_engine) as session: + with patch("services.oauth_server.AccountService.load_user", return_value=expected_user) as mock_load: + result = OAuthServerService.validate_oauth_access_token("client-1", "access-token", session) + mock_load.assert_called_once_with("user-88", session) assert result is expected_user - mock_load.assert_called_once_with("user-88", db_session_with_containers) diff --git a/api/tests/unit_tests/services/test_rag_pipeline_task_proxy.py b/api/tests/unit_tests/services/test_rag_pipeline_task_proxy.py index 305045cb6ee..cb72de5b53e 100644 --- a/api/tests/unit_tests/services/test_rag_pipeline_task_proxy.py +++ b/api/tests/unit_tests/services/test_rag_pipeline_task_proxy.py @@ -6,7 +6,7 @@ import pytest from core.app.entities.rag_pipeline_invoke_entities import RagPipelineInvokeEntity from core.rag.pipeline.queue import TenantIsolatedTaskQueue -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from services.rag_pipeline.rag_pipeline_task_proxy import RagPipelineTaskProxy diff --git a/api/tests/unit_tests/services/test_recommended_app_service.py b/api/tests/unit_tests/services/test_recommended_app_service.py index 827ded2d61a..98f4ca467ca 100644 --- a/api/tests/unit_tests/services/test_recommended_app_service.py +++ b/api/tests/unit_tests/services/test_recommended_app_service.py @@ -10,10 +10,9 @@ import pytest from sqlalchemy import select from sqlalchemy.orm import Session -from enums.deployment_edition import DeploymentEdition +from enums import DeploymentEdition from models.model import AccountTrialAppRecord, App, AppMode, TrialApp from services import recommended_app_service as service_module -from services.feature_service import SystemFeatureModel from services.recommended_app_service import RecommendedAppService pytestmark = pytest.mark.parametrize( @@ -49,6 +48,28 @@ class AppDetailKwargs(TypedDict, total=False): tools: list[str] +@pytest.mark.parametrize( + ("edition", "feature_enabled", "expected"), + [ + (DeploymentEdition.CLOUD, True, True), + (DeploymentEdition.CLOUD, False, False), + (DeploymentEdition.COMMUNITY, True, False), + (DeploymentEdition.ENTERPRISE, True, False), + ], +) +def test_trial_app_policy_is_cloud_only( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + edition: DeploymentEdition, + feature_enabled: bool, + expected: bool, +) -> None: + monkeypatch.setattr(service_module.dify_config, "DEPLOYMENT_EDITION", edition) + monkeypatch.setattr(service_module.dify_config, "ENABLE_TRIAL_APP", feature_enabled) + + assert RecommendedAppService.is_trial_app_enabled() is expected + + # ── Helpers ──────────────────────────────────────────────────────────── @@ -58,8 +79,8 @@ def _apps_response( ) -> AppsResponse: if recommended_apps is None: recommended_apps = [ - {"id": "app-1", "name": "Test App 1", "description": "d1", "category": "productivity"}, - {"id": "app-2", "name": "Test App 2", "description": "d2", "category": "communication"}, + {"app_id": "app-1", "name": "Test App 1", "description": "d1", "category": "productivity"}, + {"app_id": "app-2", "name": "Test App 2", "description": "d2", "category": "communication"}, ] if categories is None: categories = ["productivity", "communication", "utilities"] @@ -175,7 +196,7 @@ class TestRecommendedAppServiceGetApps: mock_config.HOSTED_FETCH_APP_TEMPLATES_MODE = "remote" empty_response = AppsResponse(recommended_apps=[], categories=[]) builtin_response = _apps_response( - recommended_apps=[{"id": "builtin-1", "name": "Builtin App", "category": "default"}] + recommended_apps=[{"app_id": "builtin-1", "name": "Builtin App", "category": "default"}] ) mock_remote_instance = MagicMock() @@ -189,7 +210,7 @@ class TestRecommendedAppServiceGetApps: result = RecommendedAppService.get_recommended_apps_and_categories("zh-CN", session=sqlite_session) assert result == builtin_response - assert result["recommended_apps"][0]["id"] == "builtin-1" + assert result["recommended_apps"][0]["app_id"] == "builtin-1" mock_builtin_instance.fetch_recommended_apps_from_builtin.assert_called_once_with("en-US") @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @@ -223,7 +244,7 @@ class TestRecommendedAppServiceGetApps: for language in ["en-US", "zh-CN", "ja-JP", "fr-FR"]: lang_response = _apps_response( - recommended_apps=[{"id": f"app-{language}", "name": f"App {language}", "category": "test"}] + recommended_apps=[{"app_id": f"app-{language}", "name": f"App {language}", "category": "test"}] ) mock_instance = MagicMock() mock_instance.get_recommended_apps_and_categories.return_value = lang_response @@ -231,7 +252,7 @@ class TestRecommendedAppServiceGetApps: result = RecommendedAppService.get_recommended_apps_and_categories(language, session=sqlite_session) - assert result["recommended_apps"][0]["id"] == f"app-{language}" + assert result["recommended_apps"][0]["app_id"] == f"app-{language}" mock_instance.get_recommended_apps_and_categories.assert_called_with(language, session=sqlite_session) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @@ -262,14 +283,14 @@ class TestRecommendedAppServiceGetApp: monkeypatch, result=RecommendedAppPayload(id=app.id), ) - feature_lookup = MagicMock(side_effect=AssertionError("get_app must not inspect trial features")) - monkeypatch.setattr(service_module.FeatureService, "get_system_features", feature_lookup) + trial_policy = MagicMock(side_effect=AssertionError("get_app must not inspect trial policy")) + monkeypatch.setattr(RecommendedAppService, "is_trial_app_enabled", trial_policy) result = RecommendedAppService.get_app(app.id, session=sqlite_session) assert result is app retrieval_instance.get_recommend_app_detail.assert_called_once_with(app.id, session=sqlite_session) - feature_lookup.assert_not_called() + trial_policy.assert_not_called() def test_returns_none_when_app_is_not_recommended( self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session @@ -288,21 +309,17 @@ class TestRecommendedAppServiceGetApp: class TestRecommendedAppServiceGetDetail: - @patch("services.recommended_app_service.FeatureService", autospec=True) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @patch("services.recommended_app_service.dify_config") def test_returns_retrieval_detail_when_trial_disabled( self, mock_config: MagicMock, mock_factory_class: MagicMock, - mock_feature_service: MagicMock, sqlite_session: Session, ) -> None: mock_config.HOSTED_FETCH_APP_TEMPLATES_MODE = "remote" - mock_feature_service.get_system_features.return_value = SystemFeatureModel( - deployment_edition=DeploymentEdition.COMMUNITY, - enable_trial_app=False, - ) + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY + mock_config.ENABLE_TRIAL_APP = True cases: list[tuple[str, RecommendedAppPayload]] = [ ( "complex-app", @@ -328,23 +345,20 @@ class TestRecommendedAppServiceGetDetail: result = RecommendedAppService.get_recommend_app_detail(app_id, session=sqlite_session) - assert result == expected + assert result is not None + assert result["can_trial"] is False mock_instance.get_recommend_app_detail.assert_called_once_with(app_id, session=sqlite_session) - @patch("services.recommended_app_service.FeatureService", autospec=True) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @patch("services.recommended_app_service.dify_config") def test_different_modes( self, mock_config: MagicMock, mock_factory_class: MagicMock, - mock_feature_service: MagicMock, sqlite_session: Session, ) -> None: - mock_feature_service.get_system_features.return_value = SystemFeatureModel( - deployment_edition=DeploymentEdition.COMMUNITY, - enable_trial_app=False, - ) + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY + mock_config.ENABLE_TRIAL_APP = True for mode in ["remote", "builtin", "db"]: mock_config.HOSTED_FETCH_APP_TEMPLATES_MODE = mode detail = _app_detail(app_id="test-app", name=f"App from {mode}") @@ -363,21 +377,17 @@ class TestRecommendedAppServiceGetDetail: class TestRecommendedAppServiceGetLearnDifyApps: - @patch("services.recommended_app_service.FeatureService", autospec=True) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @patch("services.recommended_app_service.dify_config") def test_uses_configured_retrieval_source( self, mock_config: MagicMock, mock_factory_class: MagicMock, - mock_feature_service: MagicMock, sqlite_session: Session, ) -> None: mock_config.HOSTED_FETCH_APP_TEMPLATES_MODE = "remote" - mock_feature_service.get_system_features.return_value = SystemFeatureModel( - deployment_edition=DeploymentEdition.COMMUNITY, - enable_trial_app=False, - ) + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY + mock_config.ENABLE_TRIAL_APP = True expected_app = RecommendedAppPayload(app_id="app-1", category="Workflow") mock_instance = MagicMock() mock_instance.get_learn_dify_apps.return_value = { @@ -388,7 +398,7 @@ class TestRecommendedAppServiceGetLearnDifyApps: result = RecommendedAppService.get_learn_dify_apps("en-US", session=sqlite_session) - assert result == {"recommended_apps": [expected_app]} + assert result == {"recommended_apps": [{**expected_app, "can_trial": False}]} mock_factory_class.get_recommend_app_factory.assert_called_once_with("remote") mock_instance.get_learn_dify_apps.assert_called_once_with("en-US", session=sqlite_session) @@ -409,23 +419,14 @@ class TestRecommendedAppServiceGetLearnDifyApps: "get_recommend_app_factory", MagicMock(return_value=mock_retrieval_factory), ) - monkeypatch.setattr( - service_module.FeatureService, - "get_system_features", - MagicMock( - return_value=SystemFeatureModel( - deployment_edition=DeploymentEdition.COMMUNITY, - enable_trial_app=True, - ) - ), - ) - can_trial_mock = MagicMock(return_value=True) - monkeypatch.setattr(RecommendedAppService, "_can_trial_app", can_trial_mock) + monkeypatch.setattr(RecommendedAppService, "is_trial_app_enabled", MagicMock(return_value=True)) + trial_app_ids = MagicMock(return_value={"app-1"}) + monkeypatch.setattr(RecommendedAppService, "_get_trial_app_ids", trial_app_ids) result = RecommendedAppService.get_learn_dify_apps("en-US", session=sqlite_session) assert result["recommended_apps"][0]["can_trial"] is True - can_trial_mock.assert_called_once_with(sqlite_session, "app-1") + trial_app_ids.assert_called_once_with(sqlite_session, ["app-1"]) # ── Integration tests: trial app features (real DB) ──────────────────── @@ -435,22 +436,20 @@ class TestRecommendedAppServiceTrialFeatures: def test_get_apps_should_not_query_trial_table_when_disabled( self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session ) -> None: - expected = AppsResponse(recommended_apps=[RecommendedAppPayload(app_id="app-1")], categories=["all"]) - retrieval_instance, builtin_instance = _mock_factory_for_apps(monkeypatch, mode="remote", result=expected) - monkeypatch.setattr( - service_module.FeatureService, - "get_system_features", - MagicMock( - return_value=SystemFeatureModel( - deployment_edition=DeploymentEdition.COMMUNITY, - enable_trial_app=False, - ) - ), + upstream_result = AppsResponse( + recommended_apps=[RecommendedAppPayload(app_id="app-1", can_trial=True)], categories=["all"] ) + retrieval_instance, builtin_instance = _mock_factory_for_apps( + monkeypatch, mode="remote", result=upstream_result + ) + monkeypatch.setattr(RecommendedAppService, "is_trial_app_enabled", MagicMock(return_value=False)) + trial_app_ids = MagicMock(side_effect=AssertionError("disabled trial must not query TrialApp")) + monkeypatch.setattr(RecommendedAppService, "_get_trial_app_ids", trial_app_ids) result = RecommendedAppService.get_recommended_apps_and_categories("en-US", session=sqlite_session) - assert result == expected + assert result["recommended_apps"][0]["can_trial"] is False + trial_app_ids.assert_not_called() retrieval_instance.get_recommended_apps_and_categories.assert_called_once_with("en-US", session=sqlite_session) builtin_instance.fetch_recommended_apps_from_builtin.assert_not_called() @@ -473,16 +472,7 @@ class TestRecommendedAppServiceTrialFeatures: _, builtin_instance = _mock_factory_for_apps( monkeypatch, mode="remote", result=remote_result, fallback_result=fallback_result ) - monkeypatch.setattr( - service_module.FeatureService, - "get_system_features", - MagicMock( - return_value=SystemFeatureModel( - deployment_edition=DeploymentEdition.COMMUNITY, - enable_trial_app=True, - ) - ), - ) + monkeypatch.setattr(RecommendedAppService, "is_trial_app_enabled", MagicMock(return_value=True)) result = RecommendedAppService.get_recommended_apps_and_categories("ja-JP", session=sqlite_session) @@ -514,16 +504,7 @@ class TestRecommendedAppServiceTrialFeatures: "get_recommend_app_factory", MagicMock(return_value=retrieval_factory), ) - monkeypatch.setattr( - service_module.FeatureService, - "get_system_features", - MagicMock( - return_value=SystemFeatureModel( - deployment_edition=DeploymentEdition.COMMUNITY, - enable_trial_app=True, - ) - ), - ) + monkeypatch.setattr(RecommendedAppService, "is_trial_app_enabled", MagicMock(return_value=True)) result = RecommendedAppService.get_recommend_app_detail(app_id, session=sqlite_session) assert result is not None @@ -532,26 +513,25 @@ class TestRecommendedAppServiceTrialFeatures: assert detail_result["id"] == app_id assert detail_result["can_trial"] is has_trial_app - @patch("services.recommended_app_service.FeatureService", autospec=True) @patch("services.recommended_app_service.RecommendAppRetrievalFactory", autospec=True) @patch("services.recommended_app_service.dify_config") def test_get_detail_returns_none_before_reading_trial_flag( self, mock_config: MagicMock, mock_factory_class: MagicMock, - mock_feature_service: MagicMock, sqlite_session: Session, ) -> None: mock_config.HOSTED_FETCH_APP_TEMPLATES_MODE = "remote" mock_instance = MagicMock() mock_instance.get_recommend_app_detail.return_value = None mock_factory_class.get_recommend_app_factory.return_value = MagicMock(return_value=mock_instance) - - result = RecommendedAppService.get_recommend_app_detail("nonexistent", session=sqlite_session) + trial_policy = MagicMock(side_effect=AssertionError("missing app must not inspect trial policy")) + with patch.object(RecommendedAppService, "is_trial_app_enabled", trial_policy): + result = RecommendedAppService.get_recommend_app_detail("nonexistent", session=sqlite_session) assert result is None mock_instance.get_recommend_app_detail.assert_called_once_with("nonexistent", session=sqlite_session) - mock_feature_service.get_system_features.assert_not_called() + trial_policy.assert_not_called() def test_add_trial_app_record_increments_count_for_existing(self, sqlite_session: Session) -> None: app_id = str(uuid.uuid4()) diff --git a/api/tests/unit_tests/services/test_schedule_service.py b/api/tests/unit_tests/services/test_schedule_service.py index aae37b885f0..4f5df2f8c8a 100644 --- a/api/tests/unit_tests/services/test_schedule_service.py +++ b/api/tests/unit_tests/services/test_schedule_service.py @@ -2,7 +2,7 @@ import json import unittest from datetime import UTC, datetime from typing import Any -from unittest.mock import MagicMock, Mock +from unittest.mock import MagicMock import pytest @@ -325,20 +325,23 @@ class TestExtractScheduleConfig(unittest.TestCase): def test_extract_schedule_config_with_cron_mode(self): """Test extracting schedule config in cron mode.""" - workflow = Mock(spec=Workflow) - workflow.graph_dict = { - "nodes": [ + workflow = Workflow( + graph=json.dumps( { - "id": "schedule-node", - "data": { - "type": "trigger-schedule", - "mode": "cron", - "cron_expression": "0 10 * * *", - "timezone": "America/New_York", - }, + "nodes": [ + { + "id": "schedule-node", + "data": { + "type": "trigger-schedule", + "mode": "cron", + "cron_expression": "0 10 * * *", + "timezone": "America/New_York", + }, + } + ] } - ] - } + ), + ) config = ScheduleService.extract_schedule_config(workflow) @@ -349,21 +352,24 @@ class TestExtractScheduleConfig(unittest.TestCase): def test_extract_schedule_config_with_visual_mode(self): """Test extracting schedule config in visual mode.""" - workflow = Mock(spec=Workflow) - workflow.graph_dict = { - "nodes": [ + workflow = Workflow( + graph=json.dumps( { - "id": "schedule-node", - "data": { - "type": "trigger-schedule", - "mode": "visual", - "frequency": "daily", - "visual_config": {"time": "10:30 AM"}, - "timezone": "UTC", - }, + "nodes": [ + { + "id": "schedule-node", + "data": { + "type": "trigger-schedule", + "mode": "visual", + "frequency": "daily", + "visual_config": {"time": "10:30 AM"}, + "timezone": "UTC", + }, + } + ] } - ] - } + ), + ) config = ScheduleService.extract_schedule_config(workflow) @@ -374,23 +380,27 @@ class TestExtractScheduleConfig(unittest.TestCase): def test_extract_schedule_config_no_schedule_node(self): """Test extracting config when no schedule node exists.""" - workflow = Mock(spec=Workflow) - workflow.graph_dict = { - "nodes": [ + workflow = Workflow( + graph=json.dumps( { - "id": "other-node", - "data": {"type": "llm"}, + "nodes": [ + { + "id": "other-node", + "data": {"type": "llm"}, + } + ] } - ] - } + ), + ) config = ScheduleService.extract_schedule_config(workflow) assert config is None def test_extract_schedule_config_invalid_graph(self): """Test extracting config with invalid graph data.""" - workflow = Mock(spec=Workflow) - workflow.graph_dict = None + workflow = Workflow( + graph="", + ) with pytest.raises(ScheduleConfigError, match="Workflow graph is empty"): ScheduleService.extract_schedule_config(workflow) @@ -496,9 +506,8 @@ class TestScheduleWithTimezone(unittest.TestCase): assert summer_next.hour == 14 -def _workflow(**kwargs: Any) -> Workflow: - graph_dict = kwargs.pop("graph_dict", {}) - workflow = Workflow.new( +def _workflow(*, graph_dict: dict[str, Any]) -> Workflow: + return Workflow.new( tenant_id="tenant-1", app_id="app-1", type=WorkflowType.WORKFLOW, @@ -510,9 +519,6 @@ def _workflow(**kwargs: Any) -> Workflow: conversation_variables=[], rag_pipeline_variables=[], ) - for key, value in kwargs.items(): - setattr(workflow, key, value) - return workflow def test_to_schedule_config_should_build_from_cron_mode() -> None: diff --git a/api/tests/unit_tests/services/test_schema_definition_service.py b/api/tests/unit_tests/services/test_schema_definition_service.py new file mode 100644 index 00000000000..c2f5793c6eb --- /dev/null +++ b/api/tests/unit_tests/services/test_schema_definition_service.py @@ -0,0 +1,57 @@ +from collections.abc import Mapping +from unittest.mock import Mock, create_autospec + +import pytest + +from services.schema_definition_service import ( + SchemaDefinitionService, + SchemaDefinitionSource, +) + + +@pytest.fixture +def source() -> Mock: + return create_autospec(SchemaDefinitionSource, instance=True, spec_set=True) + + +@pytest.fixture +def source_factory(source: Mock) -> Mock: + return Mock(return_value=source) + + +def test_list_returns_schema_definitions(source: Mock, source_factory: Mock) -> None: + definitions: list[Mapping[str, object]] = [ + { + "name": "conversation-variable", + "label": "Conversation variable", + "schema": {"type": "object"}, + } + ] + source.get_all_schema_definitions.return_value = definitions + service = SchemaDefinitionService(source_factory=source_factory) + + assert service.list() == tuple(definitions) + source_factory.assert_called_once_with() + source.get_all_schema_definitions.assert_called_once_with() + + +def test_list_returns_empty_tuple_when_source_query_fails( + source: Mock, + source_factory: Mock, + caplog: pytest.LogCaptureFixture, +) -> None: + source.get_all_schema_definitions.side_effect = RuntimeError("boom") + service = SchemaDefinitionService(source_factory=source_factory) + + assert service.list() == () + assert "Failed to get schema definitions from local registry" in caplog.text + assert "boom" in caplog.text + + +def test_list_returns_empty_tuple_when_source_construction_fails(caplog: pytest.LogCaptureFixture) -> None: + source_factory = Mock(side_effect=RuntimeError("construction failed")) + service = SchemaDefinitionService(source_factory=source_factory) + + assert service.list() == () + assert "Failed to get schema definitions from local registry" in caplog.text + assert "construction failed" in caplog.text diff --git a/api/tests/unit_tests/services/test_setup_adapters.py b/api/tests/unit_tests/services/test_setup_adapters.py new file mode 100644 index 00000000000..c99c97386c2 --- /dev/null +++ b/api/tests/unit_tests/services/test_setup_adapters.py @@ -0,0 +1,74 @@ +from contextlib import nullcontext +from unittest.mock import ANY, MagicMock, patch + +import pytest +from redis.exceptions import ConnectionError as RedisConnectionError +from redis.exceptions import LockError +from sqlalchemy.orm import Session, sessionmaker + +from extensions.ext_redis import RedisClientWrapper +from services.setup_adapters import RedisSetupLock, RegisterServiceAccountProvisioner +from services.setup_service import SetupInput + + +def test_provision_delegates_to_register_service_with_managed_session( + sqlite_session_factory: sessionmaker[Session], +) -> None: + provisioner = RegisterServiceAccountProvisioner(client=sqlite_session_factory) + setup = SetupInput( + email="admin@example.com", + name="Admin", + password="Passw0rd1", + ip_address="203.0.113.7", + language="en-US", + ) + + with patch("services.setup_adapters.RegisterService.setup") as register: + provisioner.provision(setup) + + register.assert_called_once_with( + email="admin@example.com", + name="Admin", + password="Passw0rd1", + ip_address="203.0.113.7", + language="en-US", + session=ANY, + ) + assert isinstance(register.call_args.kwargs["session"], Session) + + +def test_acquire_uses_bounded_distributed_lock() -> None: + redis = MagicMock(spec=RedisClientWrapper) + redis.lock.return_value = nullcontext() + lock = RedisSetupLock(client=redis) + + with lock.acquire(): + pass + + redis.lock.assert_called_once_with( + "setup:initialize", + timeout=300, + blocking_timeout=300, + ) + + +@pytest.mark.parametrize( + "error", + [ + pytest.param(LockError("lock acquisition timed out"), id="timeout"), + pytest.param(RedisConnectionError("redis unavailable"), id="connection"), + ], +) +def test_acquire_propagates_distributed_lock_failure(error: Exception) -> None: + redis = MagicMock(spec=RedisClientWrapper) + lock_context = MagicMock() + lock_context.__enter__.side_effect = error + redis.lock.return_value = lock_context + lock = RedisSetupLock(client=redis) + + with pytest.raises(type(error), match=str(error)) as raised: + with lock.acquire(): + pytest.fail("lock body must not run") + + assert raised.value is error + lock_context.__exit__.assert_not_called() diff --git a/api/tests/unit_tests/services/test_setup_service.py b/api/tests/unit_tests/services/test_setup_service.py new file mode 100644 index 00000000000..47ab9bf41bb --- /dev/null +++ b/api/tests/unit_tests/services/test_setup_service.py @@ -0,0 +1,249 @@ +from collections.abc import Generator +from contextlib import contextmanager, nullcontext +from dataclasses import dataclass +from datetime import datetime +from unittest.mock import Mock, create_autospec + +import pytest + +from services.setup_service import ( + InitializationValidationRequiredError, + SetupAccountProvisioner, + SetupAlreadyCompletedError, + SetupInput, + SetupLock, + SetupService, + SetupState, + SetupStatus, +) + + +@pytest.fixture +def state() -> Mock: + state = create_autospec(SetupState, instance=True, spec_set=True) + state.get_setup_at.return_value = None + state.has_tenants.return_value = False + return state + + +@pytest.fixture +def accounts() -> Mock: + return create_autospec(SetupAccountProvisioner, instance=True, spec_set=True) + + +@pytest.fixture +def lock() -> Mock: + lock = create_autospec(SetupLock, instance=True, spec_set=True) + lock.acquire.return_value = nullcontext() + return lock + + +@pytest.fixture +def setup_input() -> SetupInput: + return SetupInput( + email="Admin@Example.com", + name="Admin", + password="Passw0rd1", + ip_address="203.0.113.7", + language="en-US", + ) + + +@dataclass +class TrackingLock: + inside: bool = False + exited: bool = False + + @contextmanager + def acquire(self) -> Generator[None]: + self.inside = True + try: + yield + finally: + self.inside = False + self.exited = True + + +class FailingLock: + def __init__(self, error: Exception) -> None: + self._error = error + + @contextmanager + def acquire(self) -> Generator[None]: + raise self._error + yield + + +def test_cloud_status_is_finished_without_reading_persistence(state: Mock, accounts: Mock, lock: Mock) -> None: + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=False) + + assert service.get_status() == SetupStatus(completed=True) + state.get_setup_at.assert_not_called() + + +def test_self_hosted_status_is_not_started_without_setup(state: Mock, accounts: Mock, lock: Mock) -> None: + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + assert service.get_status() == SetupStatus(completed=False) + + +def test_self_hosted_status_includes_setup_time(state: Mock, accounts: Mock, lock: Mock) -> None: + setup_at = datetime(2026, 8, 6, 10, 30) + state.get_setup_at.return_value = setup_at + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + assert service.get_status() == SetupStatus(completed=True, setup_at=setup_at) + + +def test_initialize_rejects_existing_setup( + state: Mock, + accounts: Mock, + lock: Mock, + setup_input: SetupInput, +) -> None: + state.get_setup_at.return_value = datetime(2026, 8, 6, 10, 30) + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + with pytest.raises(SetupAlreadyCompletedError): + service.initialize(setup_input, initialization_validated=True) + + state.has_tenants.assert_not_called() + accounts.provision.assert_not_called() + + +def test_initialize_rejects_existing_tenant( + state: Mock, + accounts: Mock, + lock: Mock, + setup_input: SetupInput, +) -> None: + state.has_tenants.return_value = True + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + with pytest.raises(SetupAlreadyCompletedError): + service.initialize(setup_input, initialization_validated=True) + + accounts.provision.assert_not_called() + + +@pytest.mark.parametrize("persistent_state", ["setup", "tenant"]) +def test_initialize_prioritizes_existing_persistent_state_over_validation( + state: Mock, + accounts: Mock, + lock: Mock, + setup_input: SetupInput, + persistent_state: str, +) -> None: + if persistent_state == "setup": + state.get_setup_at.return_value = datetime(2026, 8, 6, 10, 30) + else: + state.has_tenants.return_value = True + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + with pytest.raises(SetupAlreadyCompletedError): + service.initialize(setup_input, initialization_validated=False) + + accounts.provision.assert_not_called() + + +def test_initialize_requires_initialization_validation( + state: Mock, + accounts: Mock, + lock: Mock, + setup_input: SetupInput, +) -> None: + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + with pytest.raises(InitializationValidationRequiredError): + service.initialize(setup_input, initialization_validated=False) + + accounts.provision.assert_not_called() + + +def test_initialize_normalizes_email_and_provisions_account( + state: Mock, + accounts: Mock, + lock: Mock, + setup_input: SetupInput, +) -> None: + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + service.initialize(setup_input, initialization_validated=True) + + accounts.provision.assert_called_once_with( + SetupInput( + email="admin@example.com", + name="Admin", + password="Passw0rd1", + ip_address="203.0.113.7", + language="en-US", + ) + ) + lock.acquire.assert_called_once_with() + + +def test_initialize_reads_and_writes_only_while_holding_lock( + state: Mock, + accounts: Mock, + setup_input: SetupInput, +) -> None: + lock = TrackingLock() + + def get_setup_at() -> None: + assert lock.inside + + def has_tenants() -> bool: + assert lock.inside + return False + + def provision(_setup: SetupInput) -> None: + assert lock.inside + + state.get_setup_at.side_effect = get_setup_at + state.has_tenants.side_effect = has_tenants + accounts.provision.side_effect = provision + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + service.initialize(setup_input, initialization_validated=True) + + assert lock.exited + + +def test_initialize_does_not_read_or_write_when_lock_acquisition_fails( + state: Mock, + accounts: Mock, + setup_input: SetupInput, +) -> None: + error = TimeoutError("lock acquisition timed out") + service = SetupService( + state=state, + accounts=accounts, + lock=FailingLock(error), + setup_required=True, + ) + + with pytest.raises(TimeoutError, match="lock acquisition timed out") as raised: + service.initialize(setup_input, initialization_validated=True) + + assert raised.value is error + state.get_setup_at.assert_not_called() + state.has_tenants.assert_not_called() + accounts.provision.assert_not_called() + + +def test_initialize_releases_lock_and_propagates_provision_failure( + state: Mock, + accounts: Mock, + setup_input: SetupInput, +) -> None: + lock = TrackingLock() + error = RuntimeError("provision failed") + accounts.provision.side_effect = error + service = SetupService(state=state, accounts=accounts, lock=lock, setup_required=True) + + with pytest.raises(RuntimeError, match="provision failed") as raised: + service.initialize(setup_input, initialization_validated=True) + + assert raised.value is error + assert lock.exited + assert not lock.inside diff --git a/api/tests/unit_tests/services/test_snippet_dsl_service.py b/api/tests/unit_tests/services/test_snippet_dsl_service.py index e527e5bd950..a9bbdc98495 100644 --- a/api/tests/unit_tests/services/test_snippet_dsl_service.py +++ b/api/tests/unit_tests/services/test_snippet_dsl_service.py @@ -253,7 +253,7 @@ workflow: assert result.error == "Snippet cannot contain the following node types: start" -def test_import_snippet_stores_pending_data_for_newer_dsl(monkeypatch): +def test_import_snippet_stores_pending_data_for_newer_dsl(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace(scalar=Mock(return_value=None))) setex = Mock() monkeypatch.setattr("services.snippet_dsl_service.redis_client.setex", setex) @@ -269,7 +269,7 @@ workflow: """ result = service.import_snippet( - account=SimpleNamespace(current_tenant_id="tenant-1"), + account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1"), import_mode=ImportMode.YAML_CONTENT.value, yaml_content=yaml_content, name="Override", @@ -278,7 +278,10 @@ workflow: assert result.status == ImportStatus.PENDING setex.assert_called_once() + assert setex.call_args.args[0] == f"snippet_import_info:{result.id}" pending = SnippetPendingData.model_validate_json(setex.call_args.args[2]) + assert pending.tenant_id == "tenant-1" + assert pending.account_id == "account-1" assert pending.name == "Override" assert pending.description == "Override description" @@ -307,7 +310,7 @@ workflow: assert result.error == "Snippet not found" -def test_import_snippet_passes_dependencies_to_create_or_update(monkeypatch): +def test_import_snippet_passes_dependencies_to_create_or_update(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace(scalar=Mock(return_value=None))) snippet = SimpleNamespace(id="snippet-1") create_or_update = Mock(return_value=snippet) @@ -339,7 +342,7 @@ workflow: assert dependencies[0].value.plugin_unique_identifier == "langgenius/openai:0.0.1" -def test_import_snippet_rolls_back_when_create_or_update_raises(monkeypatch): +def test_import_snippet_rolls_back_when_create_or_update_raises(monkeypatch: pytest.MonkeyPatch): session = SimpleNamespace(scalar=Mock(return_value=None), rollback=Mock()) service = SnippetDslService(session=session) monkeypatch.setattr(service, "_create_or_update_snippet", Mock(side_effect=RuntimeError("boom"))) @@ -355,27 +358,31 @@ def test_import_snippet_rolls_back_when_create_or_update_raises(monkeypatch): session.rollback.assert_called_once() -def test_confirm_import_returns_failed_when_pending_data_missing(monkeypatch): +def test_confirm_import_returns_failed_when_pending_data_missing(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace()) monkeypatch.setattr("services.snippet_dsl_service.redis_client.get", Mock(return_value=None)) - result = service.confirm_import(import_id="missing", account=SimpleNamespace(current_tenant_id="tenant-1")) + result = service.confirm_import( + import_id="missing", account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1") + ) assert result.status == ImportStatus.FAILED assert result.error == "Import information expired or does not exist" -def test_confirm_import_returns_failed_for_invalid_pending_payload(monkeypatch): +def test_confirm_import_returns_failed_for_invalid_pending_payload(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace()) monkeypatch.setattr("services.snippet_dsl_service.redis_client.get", Mock(return_value=object())) - result = service.confirm_import(import_id="bad", account=SimpleNamespace(current_tenant_id="tenant-1")) + result = service.confirm_import( + import_id="bad", account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1") + ) assert result.status == ImportStatus.FAILED assert result.error == "Invalid import information" -def test_confirm_import_creates_snippet_from_pending_data(monkeypatch): +def test_confirm_import_is_scoped_to_its_owner(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace(scalar=Mock(return_value=None))) account = SimpleNamespace(id="account-1", current_tenant_id="tenant-1") snippet = SimpleNamespace(id="snippet-new") @@ -391,6 +398,8 @@ workflow: edges: [] """ pending = SnippetPendingData( + tenant_id="tenant-1", + account_id="account-1", import_mode="yaml-content", yaml_content=yaml_content, name="Override name", @@ -399,10 +408,21 @@ workflow: ) create_or_update = Mock(return_value=snippet) monkeypatch.setattr(service, "_create_or_update_snippet", create_or_update) - monkeypatch.setattr("services.snippet_dsl_service.redis_client.get", Mock(return_value=pending.model_dump_json())) + redis_key = "snippet_import_info:import-1" + monkeypatch.setattr( + "services.snippet_dsl_service.redis_client.get", + Mock(side_effect=lambda key: pending.model_dump_json() if key == redis_key else None), + ) redis_delete = Mock() monkeypatch.setattr("services.snippet_dsl_service.redis_client.delete", redis_delete) + for other_account in ( + SimpleNamespace(id="account-1", current_tenant_id="tenant-2"), + SimpleNamespace(id="account-2", current_tenant_id="tenant-1"), + ): + assert service.confirm_import(import_id="import-1", account=other_account).status == ImportStatus.FAILED + + create_or_update.assert_not_called() result = service.confirm_import(import_id="import-1", account=account) assert result.status == ImportStatus.COMPLETED @@ -414,10 +434,10 @@ workflow: assert kwargs["account"] is account assert kwargs["name"] == "Override name" assert kwargs["description"] == "Override description" - redis_delete.assert_called_once_with("snippet_import_info:import-1") + redis_delete.assert_called_once_with(redis_key) -def test_confirm_import_returns_failed_for_non_mapping_yaml(monkeypatch): +def test_confirm_import_returns_failed_for_non_mapping_yaml(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace()) pending = SnippetPendingData( import_mode="yaml-content", @@ -426,13 +446,15 @@ def test_confirm_import_returns_failed_for_non_mapping_yaml(monkeypatch): ) monkeypatch.setattr("services.snippet_dsl_service.redis_client.get", Mock(return_value=pending.model_dump_json())) - result = service.confirm_import(import_id="import-1", account=SimpleNamespace(current_tenant_id="tenant-1")) + result = service.confirm_import( + import_id="import-1", account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1") + ) assert result.status == ImportStatus.FAILED assert result.error == "Invalid YAML format: expected a dictionary" -def test_confirm_import_returns_failed_when_create_or_update_raises(monkeypatch): +def test_confirm_import_returns_failed_when_create_or_update_raises(monkeypatch: pytest.MonkeyPatch): session = SimpleNamespace(scalar=Mock(return_value=None), rollback=Mock()) service = SnippetDslService(session=session) pending = SnippetPendingData( @@ -445,7 +467,7 @@ def test_confirm_import_returns_failed_when_create_or_update_raises(monkeypatch) result = service.confirm_import( import_id="import-1", - account=SimpleNamespace(current_tenant_id="tenant-1"), + account=SimpleNamespace(id="account-1", current_tenant_id="tenant-1"), ) assert result.status == ImportStatus.FAILED @@ -453,7 +475,7 @@ def test_confirm_import_returns_failed_when_create_or_update_raises(monkeypatch) session.rollback.assert_called_once() -def test_check_dependencies_returns_empty_without_draft_workflow(monkeypatch): +def test_check_dependencies_returns_empty_without_draft_workflow(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace(get_bind=Mock())) monkeypatch.setattr( "services.snippet_dsl_service.SnippetService", @@ -465,7 +487,7 @@ def test_check_dependencies_returns_empty_without_draft_workflow(monkeypatch): assert result.leaked_dependencies == [] -def test_check_dependencies_returns_generated_dependencies(monkeypatch): +def test_check_dependencies_returns_generated_dependencies(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace(get_bind=Mock())) workflow = SimpleNamespace(graph_dict={"nodes": []}) leaked_dependencies = [ @@ -489,9 +511,10 @@ def test_check_dependencies_returns_generated_dependencies(monkeypatch): assert result.leaked_dependencies[0].value.plugin_unique_identifier == "langgenius/openai:0.0.1" -def test_create_or_update_snippet_updates_existing_snippet_and_syncs_workflow(monkeypatch): +def test_create_or_update_snippet_updates_existing_snippet_and_syncs_workflow(monkeypatch: pytest.MonkeyPatch): snippet = SimpleNamespace( id="snippet-1", + tenant_id="tenant-1", name="Old", description="Old", type="node", @@ -508,6 +531,14 @@ def test_create_or_update_snippet_updates_existing_snippet_and_syncs_workflow(mo sync_draft_workflow=Mock(), ) monkeypatch.setattr("services.snippet_dsl_service.SnippetService", lambda *_args, **_kwargs: snippet_service) + monkeypatch.setattr( + "services.snippet_dsl_service.WorkflowAgentPublishService.sync_agent_bindings_for_draft", + Mock(return_value=set()), + ) + monkeypatch.setattr( + "services.snippet_dsl_service.WorkflowAgentPublishService.validate_agent_nodes_for_draft_sync", + Mock(), + ) result = service._create_or_update_snippet( snippet=snippet, @@ -532,11 +563,19 @@ def test_create_or_update_snippet_updates_existing_snippet_and_syncs_workflow(mo session.commit.assert_called_once() -def test_create_or_update_snippet_creates_new_snippet_and_flushes(monkeypatch): +def test_create_or_update_snippet_creates_new_snippet_and_flushes(monkeypatch: pytest.MonkeyPatch): session = SimpleNamespace(add=Mock(), flush=Mock(), commit=Mock(), get_bind=Mock()) service = SnippetDslService(session=session) snippet_service = SimpleNamespace(get_draft_workflow=Mock(return_value=None), sync_draft_workflow=Mock()) monkeypatch.setattr("services.snippet_dsl_service.SnippetService", lambda *_args, **_kwargs: snippet_service) + monkeypatch.setattr( + "services.snippet_dsl_service.WorkflowAgentPublishService.sync_agent_bindings_for_draft", + Mock(return_value=set()), + ) + monkeypatch.setattr( + "services.snippet_dsl_service.WorkflowAgentPublishService.validate_agent_nodes_for_draft_sync", + Mock(), + ) result = service._create_or_update_snippet( snippet=None, @@ -560,7 +599,7 @@ def test_create_or_update_snippet_creates_new_snippet_and_flushes(monkeypatch): session.commit.assert_called_once() -def test_export_snippet_dsl_raises_without_draft_workflow(monkeypatch): +def test_export_snippet_dsl_raises_without_draft_workflow(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace(get_bind=Mock())) monkeypatch.setattr( "services.snippet_dsl_service.SnippetService", @@ -571,7 +610,7 @@ def test_export_snippet_dsl_raises_without_draft_workflow(monkeypatch): service.export_snippet_dsl(SimpleNamespace()) -def test_export_snippet_dsl_returns_yaml(monkeypatch): +def test_export_snippet_dsl_returns_yaml(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace(get_bind=Mock())) workflow = SimpleNamespace( to_dict=Mock(return_value={"graph": {"nodes": []}}), @@ -601,7 +640,7 @@ def test_export_snippet_dsl_returns_yaml(monkeypatch): assert "input_fields:" in result -def test_append_workflow_export_data_filters_credentials_and_extracts_dependencies(monkeypatch): +def test_append_workflow_export_data_filters_credentials_and_extracts_dependencies(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace()) workflow_dict = { "graph": { @@ -659,7 +698,7 @@ def test_append_workflow_export_data_filters_credentials_and_extracts_dependenci assert "credential_id" not in nodes[2]["data"]["agent_parameters"]["tools"]["value"][0] -def test_append_workflow_export_data_rewrites_knowledge_dataset_ids(monkeypatch): +def test_append_workflow_export_data_rewrites_knowledge_dataset_ids(monkeypatch: pytest.MonkeyPatch): service = SnippetDslService(session=SimpleNamespace()) workflow_dict = { "graph": { diff --git a/api/tests/unit_tests/services/test_snippet_generate_service.py b/api/tests/unit_tests/services/test_snippet_generate_service.py index 63ccbdb351e..83568a382ae 100644 --- a/api/tests/unit_tests/services/test_snippet_generate_service.py +++ b/api/tests/unit_tests/services/test_snippet_generate_service.py @@ -4,6 +4,7 @@ from types import SimpleNamespace from unittest.mock import Mock import pytest +from sqlalchemy.orm import Session, sessionmaker from core.workflow.snippet_start import SNIPPET_VIRTUAL_START_NODE_ID from models.workflow import Workflow, WorkflowKind, WorkflowType @@ -73,7 +74,7 @@ def test_ensure_start_node_returns_workflow_when_start_already_exists(): assert result is workflow -def test_ensure_start_node_injects_virtual_start_for_root_candidates(monkeypatch): +def test_ensure_start_node_injects_virtual_start_for_root_candidates(monkeypatch: pytest.MonkeyPatch): graph = { "nodes": [ {"id": "llm-1", "data": {"type": "llm"}}, @@ -107,14 +108,14 @@ def test_ensure_start_node_injects_virtual_start_for_root_candidates(monkeypatch make_transient.assert_called_once_with(workflow) -def test_parse_files_returns_empty_when_upload_config_disabled(monkeypatch): +def test_parse_files_returns_empty_when_upload_config_disabled(monkeypatch: pytest.MonkeyPatch): workflow = _workflow({"nodes": [], "edges": []}) monkeypatch.setattr("services.snippet_generate_service.FileUploadConfigManager.convert", Mock(return_value=None)) assert SnippetGenerateService.parse_files(workflow, files=[{"id": "file-1"}]) == [] -def test_parse_files_delegates_to_file_factory(monkeypatch): +def test_parse_files_delegates_to_file_factory(monkeypatch: pytest.MonkeyPatch): workflow = _workflow({"nodes": [], "edges": []}) upload_config = SimpleNamespace(enabled=True) files = [SimpleNamespace(id="file-1")] @@ -130,7 +131,7 @@ def test_parse_files_delegates_to_file_factory(monkeypatch): build_from_mappings.assert_called_once() -def test_generate_raises_when_draft_workflow_missing(monkeypatch): +def test_generate_raises_when_draft_workflow_missing(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr( "services.snippet_generate_service.SnippetService", lambda *_args, **_kwargs: SimpleNamespace(get_draft_workflow=Mock(return_value=None)), @@ -145,7 +146,7 @@ def test_generate_raises_when_draft_workflow_missing(monkeypatch): ) -def test_generate_delegates_to_workflow_generator_and_filters_stream(monkeypatch): +def test_generate_delegates_to_workflow_generator_and_filters_stream(monkeypatch: pytest.MonkeyPatch): workflow = _workflow({"nodes": [{"id": "llm-1", "data": {"type": "llm"}}], "edges": []}) snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1", input_fields_list=[]) user = SimpleNamespace(id="user-1") @@ -186,7 +187,7 @@ def test_generate_delegates_to_workflow_generator_and_filters_stream(monkeypatch workflow_generator_class.convert_to_event_stream.assert_called_once() -def test_run_published_delegates_to_workflow_generator_non_streaming(monkeypatch): +def test_run_published_delegates_to_workflow_generator_non_streaming(monkeypatch: pytest.MonkeyPatch): workflow = _workflow({"nodes": [{"id": "llm-1", "data": {"type": "llm"}}], "edges": []}) snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1", input_fields_list=[]) user = SimpleNamespace(id="user-1") @@ -216,7 +217,7 @@ def test_run_published_delegates_to_workflow_generator_non_streaming(monkeypatch assert kwargs["call_depth"] == 0 -def test_ensure_start_node_for_worker_delegates(monkeypatch): +def test_ensure_start_node_for_worker_delegates(monkeypatch: pytest.MonkeyPatch): workflow = _workflow({"nodes": [], "edges": []}) snippet = SimpleNamespace(input_fields_list=[]) ensure_start_node = Mock(return_value=workflow) @@ -228,7 +229,7 @@ def test_ensure_start_node_for_worker_delegates(monkeypatch): ensure_start_node.assert_called_once_with(workflow, snippet) -def test_run_draft_node_delegates_to_workflow_service(monkeypatch): +def test_run_draft_node_delegates_to_workflow_service(monkeypatch: pytest.MonkeyPatch): workflow = _workflow({"nodes": [{"id": "llm-1", "data": {"type": "llm"}}], "edges": []}) snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") account = SimpleNamespace(id="account-1") @@ -262,7 +263,7 @@ def test_run_draft_node_delegates_to_workflow_service(monkeypatch): assert kwargs["files"] == [] -def test_run_draft_node_raises_when_draft_workflow_missing(monkeypatch): +def test_run_draft_node_raises_when_draft_workflow_missing(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr( "services.snippet_generate_service.SnippetService", lambda *_args, **_kwargs: SimpleNamespace(get_draft_workflow=Mock(return_value=None)), @@ -277,7 +278,10 @@ def test_run_draft_node_raises_when_draft_workflow_missing(monkeypatch): ) -def test_generate_single_iteration_delegates_to_workflow_generator(monkeypatch): +def test_generate_single_iteration_delegates_to_workflow_generator( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], +) -> None: workflow = _workflow({"nodes": [{"id": "iteration-1", "data": {"type": "iteration"}}], "edges": []}) snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") user = SimpleNamespace(id="user-1") @@ -292,13 +296,12 @@ def test_generate_single_iteration_delegates_to_workflow_generator(monkeypatch): ) monkeypatch.setattr("services.snippet_generate_service.WorkflowAppGenerator", workflow_generator_class) - session = Mock() result = SnippetGenerateService.generate_single_iteration( snippet=snippet, user=user, node_id="iteration-1", args={"inputs": {"items": [1]}}, - session_maker=_session_maker(session), + session_maker=sqlite_session_factory, ) assert list(result) == ["event"] @@ -309,11 +312,11 @@ def test_generate_single_iteration_delegates_to_workflow_generator(monkeypatch): assert kwargs["node_id"] == "iteration-1" assert kwargs["user"] is user assert kwargs["streaming"] is True - assert kwargs["session"] is session + assert isinstance(kwargs["session"], Session) workflow_generator_class.convert_to_event_stream.assert_called_once_with(response) -def test_generate_single_iteration_raises_when_draft_workflow_missing(monkeypatch): +def test_generate_single_iteration_raises_when_draft_workflow_missing(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr( "services.snippet_generate_service.SnippetService", lambda *_args, **_kwargs: SimpleNamespace(get_draft_workflow=Mock(return_value=None)), @@ -329,7 +332,10 @@ def test_generate_single_iteration_raises_when_draft_workflow_missing(monkeypatc ) -def test_generate_single_loop_delegates_to_workflow_generator(monkeypatch): +def test_generate_single_loop_delegates_to_workflow_generator( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], +) -> None: workflow = _workflow({"nodes": [{"id": "loop-1", "data": {"type": "loop"}}], "edges": []}) snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") user = SimpleNamespace(id="user-1") @@ -344,13 +350,12 @@ def test_generate_single_loop_delegates_to_workflow_generator(monkeypatch): ) monkeypatch.setattr("services.snippet_generate_service.WorkflowAppGenerator", workflow_generator_class) - session = Mock() result = SnippetGenerateService.generate_single_loop( snippet=snippet, user=user, node_id="loop-1", args=SimpleNamespace(inputs={"items": [1]}), - session_maker=_session_maker(session), + session_maker=sqlite_session_factory, ) assert list(result) == ["event"] @@ -361,11 +366,11 @@ def test_generate_single_loop_delegates_to_workflow_generator(monkeypatch): assert kwargs["node_id"] == "loop-1" assert kwargs["user"] is user assert kwargs["streaming"] is True - assert kwargs["session"] is session + assert isinstance(kwargs["session"], Session) workflow_generator_class.convert_to_event_stream.assert_called_once_with(response) -def test_generate_single_loop_raises_when_draft_workflow_missing(monkeypatch): +def test_generate_single_loop_raises_when_draft_workflow_missing(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr( "services.snippet_generate_service.SnippetService", lambda *_args, **_kwargs: SimpleNamespace(get_draft_workflow=Mock(return_value=None)), @@ -381,7 +386,7 @@ def test_generate_single_loop_raises_when_draft_workflow_missing(monkeypatch): ) -def test_run_published_raises_when_published_workflow_missing(monkeypatch): +def test_run_published_raises_when_published_workflow_missing(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr( "services.snippet_generate_service.SnippetService", lambda *_args, **_kwargs: SimpleNamespace(get_published_workflow=Mock(return_value=None)), diff --git a/api/tests/unit_tests/services/test_snippet_service.py b/api/tests/unit_tests/services/test_snippet_service.py index b3c6cdcbf6c..21a9deb824e 100644 --- a/api/tests/unit_tests/services/test_snippet_service.py +++ b/api/tests/unit_tests/services/test_snippet_service.py @@ -1,41 +1,33 @@ from __future__ import annotations import json +from datetime import datetime from types import SimpleNamespace from unittest.mock import Mock import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session, sessionmaker -from models.snippet import SnippetType -from models.workflow import Workflow, WorkflowKind, WorkflowType +from enums import DeploymentEdition +from extensions.storage.storage_type import StorageType +from graphon.variables.segments import StringSegment +from graphon.variables.types import SegmentType +from models.agent import Agent, AgentScope, AgentSource, AgentStatus +from models.enums import CreatorUserRole +from models.model import UploadFile +from models.snippet import CustomizedSnippet, SnippetType +from models.workflow import ( + Workflow, + WorkflowDraftVariable, + WorkflowDraftVariableFile, + WorkflowKind, + WorkflowType, +) from services.errors.app import IsDraftWorkflowError, WorkflowHashNotEqualError, WorkflowNotFoundError from services.snippet_service import SnippetService -class _SessionWithoutNameLookup: - def __init__(self) -> None: - self.add = Mock() - self.commit = Mock() - - def query(self, *args, **kwargs): - raise AssertionError("snippet name uniqueness lookup should not be used") - - -class _SessionContext: - def __init__(self, session) -> None: - self._session = session - - def __enter__(self): - return self._session - - def __exit__(self, *args) -> None: - return None - - -def _session_maker(session): - return lambda: _SessionContext(session) - - def _create_workflow(*, workflow_id: str, version: str, graph: dict, features: dict) -> Workflow: return Workflow( id=workflow_id, @@ -53,12 +45,26 @@ def _create_workflow(*, workflow_id: str, version: str, graph: dict, features: d ) -def test_create_snippet_allows_duplicate_names(monkeypatch: pytest.MonkeyPatch) -> None: - session = _SessionWithoutNameLookup() - account = SimpleNamespace(id="account-1") +def _snippet() -> CustomizedSnippet: + return CustomizedSnippet( + id="snippet-1", + tenant_id="tenant-1", + name="Snippet", + description="", + type=SnippetType.NODE, + created_by="account-1", + ) - service = SnippetService.__new__(SnippetService) - service._session_maker = _session_maker(session) + +def test_create_snippet_allows_duplicate_names( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + account = SimpleNamespace(id="account-1") + existing = _snippet() + existing.name = "shared name" + sqlite_session.add(existing) + sqlite_session.commit() + service = SnippetService(session_maker=sqlite_session_factory) snippet = service.create_snippet( tenant_id="tenant-1", @@ -71,8 +77,12 @@ def test_create_snippet_allows_duplicate_names(monkeypatch: pytest.MonkeyPatch) ) assert snippet.name == "shared name" - session.add.assert_called_once_with(snippet) - session.commit.assert_called_once() + stored = sqlite_session.scalars( + select(CustomizedSnippet).where( + CustomizedSnippet.tenant_id == "tenant-1", CustomizedSnippet.name == "shared name" + ) + ).all() + assert {item.id for item in stored} == {existing.id, snippet.id} def test_validate_snippet_graph_forbidden_nodes_ignores_malformed_nodes() -> None: @@ -93,30 +103,35 @@ def test_validate_snippet_graph_forbidden_nodes_raises_with_node_details() -> No SnippetService.validate_snippet_graph_forbidden_nodes({"nodes": [{"id": "start-1", "data": {"type": "start"}}]}) -def test_get_snippets_returns_empty_when_tag_filter_has_no_targets(monkeypatch: pytest.MonkeyPatch) -> None: - session = _SessionWithoutNameLookup() +def test_get_snippets_returns_empty_when_tag_filter_has_no_targets( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: get_target_ids = Mock(return_value=[]) monkeypatch.setattr("services.snippet_service.TagService.get_target_ids_by_tag_ids", get_target_ids) service = SnippetService.__new__(SnippetService) - result = service.get_snippets(tenant_id="tenant-1", session=session, tag_ids=["tag-1"]) + result = service.get_snippets(tenant_id="tenant-1", session=sqlite_session, tag_ids=["tag-1"]) assert result == ([], 0, False) - get_target_ids.assert_called_once_with("snippet", "tenant-1", ["tag-1"], session, match_all=True) + get_target_ids.assert_called_once_with("snippet", "tenant-1", ["tag-1"], sqlite_session, match_all=True) -def test_get_snippets_applies_filters_and_paginates(monkeypatch: pytest.MonkeyPatch) -> None: - snippets = [ - SimpleNamespace(id="snippet-1"), - SimpleNamespace(id="snippet-2"), - SimpleNamespace(id="snippet-3"), - ] - session = SimpleNamespace( - scalar=Mock(return_value=3), - scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=snippets))), - ) +def test_get_snippets_applies_filters_and_paginates(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: + snippets = [] + for index in range(3): + snippet = CustomizedSnippet( + id=f"snippet-{index + 1}", + tenant_id="tenant-1", + name=f"search {index}", + description="search result", + type=SnippetType.NODE, + created_by="account-1", + is_published=True, + ) + sqlite_session.add(snippet) + snippets.append(snippet) + sqlite_session.flush() service = SnippetService.__new__(SnippetService) - service._session_maker = _session_maker(session) get_target_ids = Mock(return_value=["snippet-1", "snippet-2", "snippet-3"]) monkeypatch.setattr( "services.snippet_service.TagService.get_target_ids_by_tag_ids", @@ -125,7 +140,7 @@ def test_get_snippets_applies_filters_and_paginates(monkeypatch: pytest.MonkeyPa result, total, has_more = service.get_snippets( tenant_id="tenant-1", - session=session, + session=sqlite_session, page=2, limit=2, keyword="search", @@ -134,26 +149,23 @@ def test_get_snippets_applies_filters_and_paginates(monkeypatch: pytest.MonkeyPa tag_ids=["tag-1"], ) - assert result == snippets[:2] + assert {snippet.id for snippet in result} <= {snippet.id for snippet in snippets} + assert len(result) == 1 assert total == 3 - assert has_more is True - get_target_ids.assert_called_once_with("snippet", "tenant-1", ["tag-1"], session, match_all=True) - session.scalar.assert_called_once() - session.scalars.assert_called_once() + assert has_more is False + get_target_ids.assert_called_once_with("snippet", "tenant-1", ["tag-1"], sqlite_session, match_all=True) -def test_update_snippet_allows_duplicate_names() -> None: - session = _SessionWithoutNameLookup() - snippet = SimpleNamespace( - id="snippet-1", - tenant_id="tenant-1", - name="old name", - description="", - icon_info=None, +def test_update_snippet_allows_duplicate_names(sqlite_session: Session) -> None: + snippet = _snippet() + other = CustomizedSnippet( + id="snippet-2", tenant_id="tenant-1", name="shared name", description="", type=SnippetType.NODE ) + sqlite_session.add_all([snippet, other]) + sqlite_session.flush() result = SnippetService.update_snippet( - session=session, + session=sqlite_session, snippet=snippet, account_id="account-1", data={"name": "shared name"}, @@ -161,21 +173,18 @@ def test_update_snippet_allows_duplicate_names() -> None: assert result is snippet assert snippet.name == "shared name" - session.add.assert_called_once_with(snippet) + sqlite_session.flush() + assert sqlite_session.get(CustomizedSnippet, snippet.id).name == "shared name" -def test_update_snippet_updates_optional_fields() -> None: - session = _SessionWithoutNameLookup() - snippet = SimpleNamespace( - id="snippet-1", - tenant_id="tenant-1", - name="old name", - description="old description", - icon_info=None, - ) +def test_update_snippet_updates_optional_fields(sqlite_session: Session) -> None: + snippet = _snippet() + snippet.description = "old description" + sqlite_session.add(snippet) + sqlite_session.flush() result = SnippetService.update_snippet( - session=session, + session=sqlite_session, snippet=snippet, account_id="account-1", data={"description": "new description", "icon_info": {"icon": "star"}}, @@ -185,22 +194,20 @@ def test_update_snippet_updates_optional_fields() -> None: assert snippet.description == "new description" assert snippet.icon_info == {"icon": "star"} assert snippet.updated_by == "account-1" - session.add.assert_called_once_with(snippet) + sqlite_session.flush() + stored = sqlite_session.get(CustomizedSnippet, snippet.id) + assert stored is not None + assert stored.description == "new description" -def test_sync_draft_workflow_creates_draft_and_updates_input_fields(monkeypatch: pytest.MonkeyPatch) -> None: - service = SnippetService.__new__(SnippetService) +def test_sync_draft_workflow_creates_draft_and_updates_input_fields( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, +) -> None: + service = SnippetService(session_maker=sqlite_session_factory) monkeypatch.setattr(service, "get_draft_workflow", Mock(return_value=None)) - session = Mock() - session.scalars.return_value.all.return_value = [] - service._session_maker = _session_maker(session) - snippet = SimpleNamespace( - id="snippet-1", - tenant_id="tenant-1", - input_fields=None, - updated_by=None, - updated_at=None, - ) + snippet = _snippet() account = SimpleNamespace(id="account-1") workflow = service.sync_draft_workflow( @@ -214,14 +221,18 @@ def test_sync_draft_workflow_creates_draft_and_updates_input_fields(monkeypatch: assert workflow.app_id == snippet.id assert workflow.kind == WorkflowKind.SNIPPET assert json.loads(snippet.input_fields) == [{"variable": "query"}] - session.add.assert_any_call(workflow) - session.add.assert_any_call(snippet) - session.commit.assert_called_once() + sqlite_session.expire_all() + stored_workflow = sqlite_session.scalar(select(Workflow).where(Workflow.id == workflow.id)) + stored_snippet = sqlite_session.get(CustomizedSnippet, snippet.id) + assert stored_workflow is not None + assert stored_snippet is not None + assert stored_snippet.input_fields_list == [{"variable": "query"}] -def test_sync_draft_workflow_raises_when_hash_mismatches() -> None: - service = SnippetService.__new__(SnippetService) - service._session_maker = _session_maker(SimpleNamespace(commit=Mock(), add=Mock())) +def test_sync_draft_workflow_raises_when_hash_mismatches( + sqlite_session_factory: sessionmaker[Session], +) -> None: + service = SnippetService(session_maker=sqlite_session_factory) service.get_draft_workflow = Mock(return_value=SimpleNamespace(unique_hash="server-hash")) with pytest.raises(WorkflowHashNotEqualError): @@ -233,8 +244,12 @@ def test_sync_draft_workflow_raises_when_hash_mismatches() -> None: ) -def test_sync_draft_workflow_updates_existing_draft_and_clears_variables(monkeypatch: pytest.MonkeyPatch) -> None: - service = SnippetService.__new__(SnippetService) +def test_sync_draft_workflow_updates_existing_draft_and_clears_variables( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, +) -> None: + service = SnippetService(session_maker=sqlite_session_factory) workflow = _create_workflow( workflow_id="workflow-1", version=Workflow.VERSION_DRAFT, @@ -242,19 +257,9 @@ def test_sync_draft_workflow_updates_existing_draft_and_clears_variables(monkeyp features={}, ) unique_hash = workflow.unique_hash - snippet = SimpleNamespace( - id="snippet-1", - tenant_id="tenant-1", - input_fields=None, - updated_by=None, - updated_at=None, - ) + snippet = _snippet() account = SimpleNamespace(id="account-1") - session = Mock() - session.scalars.return_value.all.return_value = [] - monkeypatch.setattr(service, "get_draft_workflow", Mock(return_value=workflow)) - service._session_maker = _session_maker(session) result = service.sync_draft_workflow( snippet=snippet, @@ -272,18 +277,26 @@ def test_sync_draft_workflow_updates_existing_draft_and_clears_variables(monkeyp assert workflow.environment_variables == [] assert workflow.conversation_variables == [] assert json.loads(snippet.input_fields) == [{"variable": "query"}] - session.commit.assert_called_once() + sqlite_session.expire_all() + assert sqlite_session.get(Workflow, workflow.id) is not None + assert sqlite_session.get(CustomizedSnippet, snippet.id) is not None -def test_update_workflow_updates_marked_fields() -> None: +def test_update_workflow_updates_marked_fields(sqlite_session: Session) -> None: service = SnippetService.__new__(SnippetService) - workflow = SimpleNamespace(marked_name="", marked_comment="", updated_by=None, updated_at=None) - session = SimpleNamespace(scalar=Mock(return_value=workflow), add=Mock()) - snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") + workflow = _create_workflow( + workflow_id="workflow-1", + version="2026-01-01 00:00:00", + graph={"nodes": []}, + features={}, + ) + snippet = _snippet() + sqlite_session.add_all([snippet, workflow]) + sqlite_session.flush() account = SimpleNamespace(id="account-1") result = service.update_workflow( - session=session, + session=sqlite_session, snippet=snippet, workflow_id="workflow-1", account=account, @@ -294,16 +307,17 @@ def test_update_workflow_updates_marked_fields() -> None: assert workflow.marked_name == "v1" assert workflow.marked_comment == "first version" assert workflow.updated_by == "account-1" - session.scalar.assert_called_once() - session.add.assert_called_once_with(workflow) + sqlite_session.flush() + stored = sqlite_session.get(Workflow, workflow.id) + assert stored is not None + assert stored.marked_name == "v1" -def test_update_workflow_returns_none_when_missing() -> None: +def test_update_workflow_returns_none_when_missing(sqlite_session: Session) -> None: service = SnippetService.__new__(SnippetService) - session = SimpleNamespace(scalar=Mock(return_value=None), add=Mock()) result = service.update_workflow( - session=session, + session=sqlite_session, snippet=SimpleNamespace(id="snippet-1", tenant_id="tenant-1"), workflow_id="missing-workflow", account=SimpleNamespace(id="account-1"), @@ -311,7 +325,6 @@ def test_update_workflow_returns_none_when_missing() -> None: ) assert result is None - session.add.assert_not_called() def test_get_default_block_configs_skips_empty_defaults(monkeypatch: pytest.MonkeyPatch) -> None: @@ -358,8 +371,10 @@ def test_get_default_block_config_returns_none_for_empty_default(monkeypatch: py def test_restore_published_snippet_workflow_to_draft_copies_source_snapshot( monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, ) -> None: - snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") + snippet = _snippet() account = SimpleNamespace(id="account-2") source_graph = {"nodes": [{"id": "llm-1", "data": {"type": "llm"}}], "edges": []} source_features = {"opening_statement": "hello"} @@ -375,10 +390,7 @@ def test_restore_published_snippet_workflow_to_draft_copies_source_snapshot( graph={"nodes": [], "edges": []}, features={}, ) - service = SnippetService.__new__(SnippetService) - session = Mock() - session.scalars.return_value.all.return_value = [] - service._session_maker = _session_maker(session) + service = SnippetService(session_maker=sqlite_session_factory) monkeypatch.setattr(service, "get_published_workflow_by_id", Mock(return_value=source_workflow)) monkeypatch.setattr(service, "get_draft_workflow", Mock(return_value=draft_workflow)) @@ -393,17 +405,19 @@ def test_restore_published_snippet_workflow_to_draft_copies_source_snapshot( assert draft_workflow.graph_dict == source_graph assert draft_workflow.features_dict == source_features assert draft_workflow.updated_by == account.id - session.add.assert_called_once_with(draft_workflow) - session.commit.assert_called_once() + sqlite_session.expire_all() + stored = sqlite_session.get(Workflow, draft_workflow.id) + assert stored is not None + assert stored.graph_dict == source_graph def test_restore_published_snippet_workflow_to_draft_raises_when_source_missing( monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], ) -> None: snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") account = SimpleNamespace(id="account-2") - service = SnippetService.__new__(SnippetService) - service._session_maker = _session_maker(SimpleNamespace(add=Mock(), commit=Mock())) + service = SnippetService(session_maker=sqlite_session_factory) monkeypatch.setattr(service, "get_published_workflow_by_id", Mock(return_value=None)) @@ -415,8 +429,12 @@ def test_restore_published_snippet_workflow_to_draft_raises_when_source_missing( ) -def test_restore_published_snippet_workflow_to_draft_adds_new_draft(monkeypatch: pytest.MonkeyPatch) -> None: - snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") +def test_restore_published_snippet_workflow_to_draft_adds_new_draft( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, +) -> None: + snippet = _snippet() account = SimpleNamespace(id="account-2") source_workflow = _create_workflow( workflow_id="published-workflow", @@ -430,10 +448,7 @@ def test_restore_published_snippet_workflow_to_draft_adds_new_draft(monkeypatch: graph={"nodes": [], "edges": []}, features={}, ) - service = SnippetService.__new__(SnippetService) - session = Mock() - session.scalars.return_value.all.return_value = [] - service._session_maker = _session_maker(session) + service = SnippetService(session_maker=sqlite_session_factory) monkeypatch.setattr(service, "get_published_workflow_by_id", Mock(return_value=source_workflow)) monkeypatch.setattr(service, "get_draft_workflow", Mock(return_value=None)) @@ -449,8 +464,8 @@ def test_restore_published_snippet_workflow_to_draft_adds_new_draft(monkeypatch: ) assert result is new_draft_workflow - session.add.assert_called_once_with(new_draft_workflow) - session.commit.assert_called_once() + sqlite_session.expire_all() + assert sqlite_session.get(Workflow, new_draft_workflow.id) is not None def test_get_published_workflow_returns_none_without_workflow_id() -> None: @@ -461,11 +476,15 @@ def test_get_published_workflow_returns_none_without_workflow_id() -> None: assert result is None -def test_get_published_workflow_by_id_raises_for_draft(monkeypatch: pytest.MonkeyPatch) -> None: - draft_workflow = SimpleNamespace(version=Workflow.VERSION_DRAFT) - session = SimpleNamespace(scalar=Mock(return_value=draft_workflow)) - service = SnippetService.__new__(SnippetService) - service._session_maker = _session_maker(session) +def test_get_published_workflow_by_id_raises_for_draft( + sqlite_session_factory: sessionmaker[Session], sqlite_session: Session +) -> None: + draft_workflow = _create_workflow( + workflow_id="workflow-1", version=Workflow.VERSION_DRAFT, graph={"nodes": []}, features={} + ) + sqlite_session.add(draft_workflow) + sqlite_session.commit() + service = SnippetService(session_maker=sqlite_session_factory) with pytest.raises(IsDraftWorkflowError): service.get_published_workflow_by_id( @@ -474,38 +493,49 @@ def test_get_published_workflow_by_id_raises_for_draft(monkeypatch: pytest.Monke ) -def test_publish_workflow_raises_when_draft_missing() -> None: +def test_publish_workflow_raises_when_draft_missing(sqlite_session: Session) -> None: service = SnippetService.__new__(SnippetService) - session = SimpleNamespace(scalar=Mock(return_value=None)) with pytest.raises(ValueError, match="No valid workflow found"): service.publish_workflow( - session=session, + session=sqlite_session, snippet=SimpleNamespace(id="snippet-1", tenant_id="tenant-1"), account=SimpleNamespace(id="account-1"), ) -def test_publish_workflow_creates_snapshot_and_updates_snippet(monkeypatch: pytest.MonkeyPatch) -> None: +def test_publish_workflow_creates_snapshot_and_updates_snippet( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: service = SnippetService.__new__(SnippetService) draft_workflow = _create_workflow( workflow_id="draft-workflow", version=Workflow.VERSION_DRAFT, - graph={"nodes": [{"id": "llm-1", "data": {"type": "llm"}}], "edges": []}, + graph={ + "nodes": [ + { + "id": "llm-1", + "data": { + "type": "llm", + "model": {"provider": "provider", "name": "model", "mode": "chat"}, + "model_selector": ["start", "MODEL_NAME"], + }, + } + ], + "edges": [], + }, features={"opening_statement": "hello"}, ) - snippet = SimpleNamespace( - id="snippet-1", - tenant_id="tenant-1", - version=1, - is_published=False, - workflow_id=None, - updated_by=None, + snippet = _snippet() + sqlite_session.add_all([draft_workflow, snippet]) + sqlite_session.flush() + monkeypatch.setattr( + "services.agent.workflow_publish_service.WorkflowAgentPublishService.copy_agent_node_bindings_to_published", + Mock(return_value=set()), ) - session = SimpleNamespace(scalar=Mock(return_value=draft_workflow), add=Mock()) - result = service.publish_workflow( - session=session, + result, retirement_candidates = service.publish_workflow( + session=sqlite_session, snippet=snippet, account=SimpleNamespace(id="account-1"), ) @@ -515,14 +545,17 @@ def test_publish_workflow_creates_snapshot_and_updates_snippet(monkeypatch: pyte assert snippet.is_published is True assert snippet.workflow_id == result.id assert snippet.updated_by == "account-1" - assert session.add.call_args_list[-1].args == (snippet,) + sqlite_session.flush() + assert sqlite_session.get(Workflow, result.id) is result + assert sqlite_session.get(CustomizedSnippet, snippet.id).workflow_id == result.id + assert retirement_candidates == set() -def test_get_all_published_workflows_returns_empty_without_current_workflow() -> None: +def test_get_all_published_workflows_returns_empty_without_current_workflow(unbound_session: Session) -> None: service = SnippetService.__new__(SnippetService) result = service.get_all_published_workflows( - session=SimpleNamespace(), + session=unbound_session, snippet=SimpleNamespace(id="snippet-1", workflow_id=None), page=1, limit=20, @@ -531,73 +564,75 @@ def test_get_all_published_workflows_returns_empty_without_current_workflow() -> assert result == ([], False) -def test_get_all_published_workflows_paginates() -> None: +def test_get_all_published_workflows_paginates(sqlite_session: Session) -> None: service = SnippetService.__new__(SnippetService) - workflows = [SimpleNamespace(id="workflow-1"), SimpleNamespace(id="workflow-2"), SimpleNamespace(id="workflow-3")] - session = SimpleNamespace(scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=workflows)))) + workflows = [ + _create_workflow( + workflow_id=f"workflow-{index}", + version=f"2026-01-0{index} 00:00:00", + graph={"nodes": []}, + features={}, + ) + for index in range(1, 4) + ] + sqlite_session.add_all(workflows) + sqlite_session.flush() result, has_more = service.get_all_published_workflows( - session=session, + session=sqlite_session, snippet=SimpleNamespace(id="snippet-1", workflow_id="workflow-current"), page=1, limit=2, ) - assert result == workflows[:2] + assert [workflow.id for workflow in result] == ["workflow-3", "workflow-2"] assert has_more is True - session.scalars.assert_called_once() -def test_delete_snippet_removes_related_records() -> None: - snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") - session = SimpleNamespace( - execute=Mock(), - scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=[]))), - delete=Mock(), +def test_delete_snippet_removes_related_records( + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session] +) -> None: + snippet = _snippet() + workflow = _create_workflow( + workflow_id="workflow-1", version=Workflow.VERSION_DRAFT, graph={"nodes": []}, features={} ) + sqlite_session.add_all([snippet, workflow]) + sqlite_session.flush() - result = SnippetService.delete_snippet(session=session, snippet=snippet) + result = SnippetService.delete_snippet(session=sqlite_session, snippet=snippet) assert result is True - executed_sql = "\n".join(str(call.args[0]) for call in session.execute.call_args_list) - assert "workflow_draft_variables" in executed_sql - assert "tool_workflow_providers" in executed_sql - assert "workflow_app_logs" in executed_sql - assert "workflow_archive_logs" in executed_sql - assert "workflow_node_executions" in executed_sql - assert "workflow_runs" in executed_sql - assert "workflows" in executed_sql - assert "kind" in executed_sql - assert "tag_bindings" in executed_sql - session.delete.assert_called_once_with(snippet) + sqlite_session.commit() + with sqlite_session_factory() as observer: + assert observer.get(CustomizedSnippet, snippet.id) is None + assert observer.get(Workflow, workflow.id) is None def test_delete_snippet_archives_owned_agents_and_schedules_backing_app_cleanup( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: - snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") - agent = SimpleNamespace( + snippet = _snippet() + agent = Agent( + id="agent-1", + tenant_id=snippet.tenant_id, + name="Snippet agent", + description="", + role="", + scope=AgentScope.WORKFLOW_ONLY, + source=AgentSource.WORKFLOW, + app_id=snippet.id, backing_app_id="backing-app-1", - status="active", - archived_by=None, - archived_at=None, + status=AgentStatus.ACTIVE, updated_by="creator-1", - updated_at=None, ) - scalar_results = [ - SimpleNamespace(all=Mock(return_value=[])), - SimpleNamespace(all=Mock(return_value=[agent])), - ] - session = SimpleNamespace( - execute=Mock(), - scalars=Mock(side_effect=scalar_results), - delete=Mock(), - ) - listen = Mock() - monkeypatch.setattr("services.snippet_service.event.listen", listen) + sqlite_session.add_all([snippet, agent]) + sqlite_session.flush() + cleanup_delay = Mock() + monkeypatch.setattr("tasks.remove_app_and_related_data_task.remove_app_and_related_data_task.delay", cleanup_delay) result = SnippetService.delete_snippet( - session=session, + session=sqlite_session, snippet=snippet, account_id="account-1", ) @@ -607,34 +642,62 @@ def test_delete_snippet_archives_owned_agents_and_schedules_backing_app_cleanup( assert agent.archived_by == "account-1" assert agent.archived_at is not None assert agent.updated_by == "account-1" - executed_sql = "\n".join(str(call.args[0]) for call in session.execute.call_args_list) - assert "DELETE FROM apps" in executed_sql - listen.assert_called_once_with(session, "after_commit", listen.call_args.args[2], once=True) + sqlite_session.commit() + assert sqlite_session.get(Agent, agent.id).status == AgentStatus.ARCHIVED + cleanup_delay.assert_called_once_with(tenant_id=snippet.tenant_id, app_id="backing-app-1") -def test_delete_draft_variable_files_removes_storage_objects(monkeypatch: pytest.MonkeyPatch) -> None: +def test_delete_draft_variable_files_removes_storage_objects( + monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], +) -> None: from extensions.ext_storage import storage - snippet = SimpleNamespace(id="snippet-1", tenant_id="tenant-1") + snippet = _snippet() storage_delete = Mock() monkeypatch.setattr(storage, "delete", storage_delete) - session = SimpleNamespace( - scalars=Mock(return_value=SimpleNamespace(all=Mock(return_value=["file-1"]))), - execute=Mock( - side_effect=[ - SimpleNamespace(all=Mock(return_value=[("file-1", "upload-1", "storage-key")])), - None, - None, - ] - ), + upload_file = UploadFile( + tenant_id=snippet.tenant_id, + storage_type=StorageType.LOCAL, + key="storage-key", + name="value.txt", + size=10, + extension=".txt", + mime_type="text/plain", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + created_at=datetime(2025, 1, 1), + used=True, ) + variable_file = WorkflowDraftVariableFile( + tenant_id=snippet.tenant_id, + app_id=snippet.id, + user_id="account-1", + upload_file_id=upload_file.id, + size=10, + length=None, + value_type=SegmentType.STRING, + ) + variable = WorkflowDraftVariable.new_node_variable( + app_id=snippet.id, + user_id="account-1", + node_id="node-1", + name="value", + value=StringSegment(value="truncated"), + node_execution_id="execution-1", + file_id=variable_file.id, + ) + sqlite_session.add_all([snippet, upload_file, variable_file, variable]) + sqlite_session.flush() - SnippetService._delete_draft_variable_files(session=session, snippet=snippet) + SnippetService._delete_draft_variable_files(session=sqlite_session, snippet=snippet) storage_delete.assert_called_once_with("storage-key") - executed_sql = "\n".join(str(call.args[0]) for call in session.execute.call_args_list) - assert "upload_files" in executed_sql - assert "workflow_draft_variable_files" in executed_sql + sqlite_session.commit() + with sqlite_session_factory() as observer: + assert observer.get(UploadFile, upload_file.id) is None + assert observer.get(WorkflowDraftVariableFile, variable_file.id) is None def test_delete_archived_workflow_run_files_removes_prefixed_objects(monkeypatch: pytest.MonkeyPatch) -> None: @@ -645,7 +708,7 @@ def test_delete_archived_workflow_run_files_removes_prefixed_objects(monkeypatch list_objects=Mock(return_value=["tenant-1/app_id=snippet-1/run.json"]), delete_object=Mock(), ) - monkeypatch.setattr(dify_config, "BILLING_ENABLED", True) + monkeypatch.setattr(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD) monkeypatch.setattr(dify_config, "ARCHIVE_STORAGE_ENABLED", True) monkeypatch.setattr("libs.archive_storage.get_archive_storage", Mock(return_value=archive_storage)) @@ -719,11 +782,14 @@ def test_workflow_run_node_executions_returns_empty_when_run_missing() -> None: service._node_execution_service_repo.get_executions_by_workflow_run.assert_not_called() -def test_increment_use_count_adds_updated_snippet() -> None: - snippet = SimpleNamespace(use_count=2) - session = SimpleNamespace(add=Mock()) +def test_increment_use_count_adds_updated_snippet(sqlite_session: Session) -> None: + snippet = _snippet() + snippet.use_count = 2 + sqlite_session.add(snippet) + sqlite_session.flush() - SnippetService.increment_use_count(session=session, snippet=snippet) + SnippetService.increment_use_count(session=sqlite_session, snippet=snippet) assert snippet.use_count == 3 - session.add.assert_called_once_with(snippet) + sqlite_session.flush() + assert sqlite_session.get(CustomizedSnippet, snippet.id).use_count == 3 diff --git a/api/tests/unit_tests/services/test_step_by_step_tour_service.py b/api/tests/unit_tests/services/test_step_by_step_tour_service.py index de08cc596fd..6cfc69dfc74 100644 --- a/api/tests/unit_tests/services/test_step_by_step_tour_service.py +++ b/api/tests/unit_tests/services/test_step_by_step_tour_service.py @@ -5,6 +5,7 @@ from datetime import UTC, datetime import pytest from sqlalchemy.exc import IntegrityError +from enums import DeploymentEdition from models.account import Account, AccountStatus from models.onboarding import AccountStepByStepTourState from services import step_by_step_tour_service as service_module @@ -104,7 +105,7 @@ def test_get_state_creates_state_and_records_first_workspace_for_eligible_accoun def test_is_eligible_does_not_depend_on_cloud_edition(monkeypatch: pytest.MonkeyPatch) -> None: _set_tour_config(monkeypatch, enabled=True, rollout_started_at=datetime(2026, 6, 1)) - monkeypatch.setattr(service_module.dify_config, "EDITION", "SELF_HOSTED") + monkeypatch.setattr(service_module.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) result = StepByStepTourService.is_eligible(_account(initialized_at=datetime(2026, 6, 28))) diff --git a/api/tests/unit_tests/services/test_summary_index_service.py b/api/tests/unit_tests/services/test_summary_index_service.py index 4af3b4cdad0..1eb1baa40a3 100644 --- a/api/tests/unit_tests/services/test_summary_index_service.py +++ b/api/tests/unit_tests/services/test_summary_index_service.py @@ -63,23 +63,24 @@ def _segment(*, has_document: bool = True) -> MagicMock: return segment -def _summary_record(*, summary_content: str = "summary", node_id: str | None = None) -> MagicMock: - record = MagicMock(spec=summary_module.DocumentSegmentSummary, name="summary_record") +def _summary_record(*, summary_content: str = "summary", node_id: str | None = None) -> DocumentSegmentSummary: + record = summary_module.DocumentSegmentSummary( + dataset_id="dataset-1", + document_id="doc-1", + chunk_id="seg-1", + summary_content=summary_content, + summary_index_node_id=node_id, + summary_index_node_hash=None, + tokens=None, + status=SummaryStatus.GENERATING, + error=None, + enabled=True, + disabled_at=None, + disabled_by=None, + ) record.id = "sum-1" - record.dataset_id = "dataset-1" - record.document_id = "doc-1" - record.chunk_id = "seg-1" - record.summary_content = summary_content - record.summary_index_node_id = node_id - record.summary_index_node_hash = None - record.tokens = None - record.status = SummaryStatus.GENERATING - record.error = None - record.enabled = True record.created_at = datetime(2024, 1, 1, tzinfo=UTC) record.updated_at = datetime(2024, 1, 1, tzinfo=UTC) - record.disabled_at = None - record.disabled_by = None return record @@ -625,9 +626,7 @@ def test_generate_and_vectorize_summary_creates_missing_record_and_logs_usage( def test_generate_summaries_for_document_skip_conditions(monkeypatch: pytest.MonkeyPatch) -> None: dataset = _dataset(indexing_technique=IndexTechniqueType.ECONOMY) - document = MagicMock(spec=summary_module.DatasetDocument) - document.id = "doc-1" - document.doc_form = IndexStructureType.PARAGRAPH_INDEX + document = summary_module.DatasetDocument(id="doc-1", doc_form=IndexStructureType.PARAGRAPH_INDEX) assert SummaryIndexService.generate_summaries_for_document(dataset, document, {"enable": True}) == [] dataset = _dataset() @@ -639,9 +638,7 @@ def test_generate_summaries_for_document_skip_conditions(monkeypatch: pytest.Mon def test_generate_summaries_for_document_runs_and_handles_errors(monkeypatch: pytest.MonkeyPatch) -> None: dataset = _dataset() - document = MagicMock(spec=summary_module.DatasetDocument) - document.id = "doc-1" - document.doc_form = IndexStructureType.PARAGRAPH_INDEX + document = summary_module.DatasetDocument(id="doc-1", doc_form=IndexStructureType.PARAGRAPH_INDEX) seg1 = _segment() seg2 = _segment() @@ -671,9 +668,7 @@ def test_generate_summaries_for_document_runs_and_handles_errors(monkeypatch: py def test_generate_summaries_for_document_no_segments_returns_empty(monkeypatch: pytest.MonkeyPatch) -> None: dataset = _dataset() - document = MagicMock(spec=summary_module.DatasetDocument) - document.id = "doc-1" - document.doc_form = IndexStructureType.PARAGRAPH_INDEX + document = summary_module.DatasetDocument(id="doc-1", doc_form=IndexStructureType.PARAGRAPH_INDEX) session = MagicMock() session.scalars.return_value.all.return_value = [] @@ -690,9 +685,7 @@ def test_generate_summaries_for_document_applies_segment_ids_and_only_parent_chu monkeypatch: pytest.MonkeyPatch, ) -> None: dataset = _dataset() - document = MagicMock(spec=summary_module.DatasetDocument) - document.id = "doc-1" - document.doc_form = IndexStructureType.PARAGRAPH_INDEX + document = summary_module.DatasetDocument(id="doc-1", doc_form=IndexStructureType.PARAGRAPH_INDEX) seg = _segment() session = MagicMock() diff --git a/api/tests/unit_tests/services/test_telemetry_service.py b/api/tests/unit_tests/services/test_telemetry_service.py new file mode 100644 index 00000000000..d931be80364 --- /dev/null +++ b/api/tests/unit_tests/services/test_telemetry_service.py @@ -0,0 +1,332 @@ +import uuid +from datetime import datetime +from unittest.mock import Mock + +import httpx +import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session + +from enums import DeploymentEdition +from models.model import DifySetup +from services import telemetry_service +from services.telemetry_service import CommunityTelemetryService + + +@pytest.fixture +def telemetry_enabled(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(telemetry_service.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + monkeypatch.setattr(telemetry_service.dify_config, "DISABLE_TELEMETRY", False) + monkeypatch.setattr(telemetry_service.dify_config, "DO_NOT_TRACK", False) + monkeypatch.setattr(telemetry_service.dify_config, "CI", False) + monkeypatch.setattr(telemetry_service.dify_config, "TELEMETRY_ENDPOINT", "https://telemetry.example.test/v1/events") + monkeypatch.setattr( + telemetry_service.dify_config, + "TELEMETRY_FALLBACK_ENDPOINT", + "https://telemetry-cn.example.test/v1/events", + ) + monkeypatch.setattr(telemetry_service.dify_config, "TELEMETRY_TIMEOUT_SECONDS", 2) + + +def test_telemetry_is_disabled_for_enterprise(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(telemetry_service.dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) + + assert CommunityTelemetryService._is_enabled() is False + + +@pytest.mark.parametrize( + ("setting", "value"), + [ + ("DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + ("DISABLE_TELEMETRY", True), + ("DO_NOT_TRACK", True), + ("CI", True), + ("TELEMETRY_ENDPOINT", ""), + ], +) +def test_telemetry_is_disabled_when_a_required_condition_is_not_met( + telemetry_enabled, monkeypatch: pytest.MonkeyPatch, setting: str, value: str | bool +): + monkeypatch.setattr(telemetry_service.dify_config, setting, value) + + assert CommunityTelemetryService._is_enabled() is False + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_reporting_without_setup_is_skipped(sqlite_session: Session, telemetry_enabled): + assert CommunityTelemetryService.report_install(session=sqlite_session) is False + assert CommunityTelemetryService.report_heartbeat(session=sqlite_session) is False + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_install_marks_reported_at(sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch): + setup = DifySetup(version="installed-version", instance_id="d246c3a1-350b-406c-92c7-6043df680758") + sqlite_session.add(setup) + sqlite_session.commit() + monkeypatch.setattr(telemetry_service.dify_config.project, "version", "running-version") + + sent_payloads: list[dict[str, str | int]] = [] + + def fake_post(url: str, json: dict[str, str | int], timeout: int): + sent_payloads.append(json) + return httpx.Response(204, request=httpx.Request("POST", url)) + + monkeypatch.setattr(telemetry_service.httpx, "post", fake_post) + + assert CommunityTelemetryService.report_install(session=sqlite_session) is True + + saved_setup = sqlite_session.scalar(select(DifySetup)) + assert saved_setup is not None + assert saved_setup.install_reported_at is not None + assert sent_payloads[0]["event"] == "install" + assert sent_payloads[0]["instance_id"] == setup.instance_id + assert sent_payloads[0]["version"] == "installed-version" + assert "installed_at" in sent_payloads[0] + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_install_generates_missing_instance_id( + sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch +): + setup = DifySetup(version="installed-version") + sqlite_session.add(setup) + sqlite_session.commit() + monkeypatch.setattr( + telemetry_service.httpx, + "post", + lambda url, json, timeout: httpx.Response(204, request=httpx.Request("POST", url)), + ) + + assert CommunityTelemetryService.report_install(session=sqlite_session) is True + + assert setup.instance_id is not None + assert str(uuid.UUID(setup.instance_id)) == setup.instance_id + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_heartbeat_generates_missing_instance_id( + sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch +): + setup = DifySetup(version="1.0.0", install_reported_at=datetime(2026, 7, 12, 8, 0, 0)) + sqlite_session.add(setup) + sqlite_session.commit() + monkeypatch.setattr( + telemetry_service.httpx, + "post", + lambda url, json, timeout: httpx.Response(204, request=httpx.Request("POST", url)), + ) + + assert ( + CommunityTelemetryService.report_heartbeat(session=sqlite_session, now=datetime(2026, 7, 13, 12, 0, 0)) is True + ) + + assert setup.instance_id is not None + assert str(uuid.UUID(setup.instance_id)) == setup.instance_id + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_install_failure_keeps_install_pending( + sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch +): + setup = DifySetup(version="1.0.0", instance_id="d246c3a1-350b-406c-92c7-6043df680758") + sqlite_session.add(setup) + sqlite_session.commit() + + def fake_post(url: str, json: dict[str, str | int], timeout: int): + raise httpx.ConnectError("offline", request=httpx.Request("POST", url)) + + monkeypatch.setattr(telemetry_service.httpx, "post", fake_post) + + assert CommunityTelemetryService.report_install(session=sqlite_session) is False + + saved_setup = sqlite_session.scalar(select(DifySetup)) + assert saved_setup is not None + assert saved_setup.install_reported_at is None + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_install_uses_fallback_endpoint_after_network_failure( + sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch +): + setup = DifySetup(version="1.0.0", instance_id="d246c3a1-350b-406c-92c7-6043df680758") + sqlite_session.add(setup) + sqlite_session.commit() + + urls: list[str] = [] + + def fake_post(url: str, json: dict[str, str | int], timeout: int): + urls.append(url) + if url == telemetry_service.dify_config.TELEMETRY_ENDPOINT: + raise httpx.ConnectError("offline", request=httpx.Request("POST", url)) + return httpx.Response(204, request=httpx.Request("POST", url)) + + monkeypatch.setattr(telemetry_service.httpx, "post", fake_post) + + assert CommunityTelemetryService.report_install(session=sqlite_session) is True + assert urls == [ + telemetry_service.dify_config.TELEMETRY_ENDPOINT, + telemetry_service.dify_config.TELEMETRY_FALLBACK_ENDPOINT, + ] + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_install_does_not_use_fallback_endpoint_after_http_error( + sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch +): + setup = DifySetup(version="1.0.0", instance_id="d246c3a1-350b-406c-92c7-6043df680758") + sqlite_session.add(setup) + sqlite_session.commit() + + post_mock = Mock( + return_value=httpx.Response( + 500, + request=httpx.Request("POST", telemetry_service.dify_config.TELEMETRY_ENDPOINT), + ) + ) + monkeypatch.setattr(telemetry_service.httpx, "post", post_mock) + + assert CommunityTelemetryService.report_install(session=sqlite_session) is False + post_mock.assert_called_once() + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_heartbeat_retries_pending_install_before_heartbeat( + sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch +): + setup = DifySetup(version="installed-version", instance_id="d246c3a1-350b-406c-92c7-6043df680758") + sqlite_session.add(setup) + sqlite_session.commit() + monkeypatch.setattr(telemetry_service.dify_config.project, "version", "running-version") + + sent_payloads: list[dict[str, str | int]] = [] + + def fake_post(url: str, json: dict[str, str | int], timeout: int): + sent_payloads.append(json) + return httpx.Response(204, request=httpx.Request("POST", url)) + + monkeypatch.setattr(telemetry_service.httpx, "post", fake_post) + now = datetime(2026, 7, 13, 0, 0, 0) + assert CommunityTelemetryService.report_heartbeat(session=sqlite_session, now=now) is True + + saved_setup = sqlite_session.scalar(select(DifySetup)) + assert saved_setup is not None + assert saved_setup.install_reported_at is not None + assert saved_setup.last_heartbeat_at == now + assert [(payload["event"], payload["version"]) for payload in sent_payloads] == [ + ("install", "installed-version"), + ("heartbeat", "running-version"), + ] + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_heartbeat_skips_when_already_sent_today( + sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch +): + setup = DifySetup( + version="1.0.0", + instance_id="d246c3a1-350b-406c-92c7-6043df680758", + install_reported_at=datetime(2026, 7, 13, 8, 0, 0), + last_heartbeat_at=datetime(2026, 7, 13, 9, 0, 0), + ) + sqlite_session.add(setup) + sqlite_session.commit() + + post_mock = Mock() + monkeypatch.setattr(telemetry_service.httpx, "post", post_mock) + + assert ( + CommunityTelemetryService.report_heartbeat(session=sqlite_session, now=datetime(2026, 7, 13, 12, 0, 0)) is False + ) + post_mock.assert_not_called() + + +@pytest.mark.parametrize("sqlite_session", [(DifySetup,)], indirect=True) +def test_report_heartbeat_failure_does_not_mark_the_day_reported( + sqlite_session: Session, telemetry_enabled, monkeypatch: pytest.MonkeyPatch +): + setup = DifySetup( + version="1.0.0", + instance_id="d246c3a1-350b-406c-92c7-6043df680758", + install_reported_at=datetime(2026, 7, 13, 8, 0, 0), + ) + sqlite_session.add(setup) + sqlite_session.commit() + + def fake_post(url: str, json: dict[str, str | int], timeout: int): + raise httpx.ConnectError("offline", request=httpx.Request("POST", url)) + + monkeypatch.setattr(telemetry_service.httpx, "post", fake_post) + + assert ( + CommunityTelemetryService.report_heartbeat(session=sqlite_session, now=datetime(2026, 7, 13, 12, 0, 0)) is False + ) + assert setup.last_heartbeat_at is None + + +def test_send_event_skips_when_telemetry_is_disabled(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(telemetry_service.dify_config, "DISABLE_TELEMETRY", True) + post_mock = Mock() + monkeypatch.setattr(telemetry_service.httpx, "post", post_mock) + + assert CommunityTelemetryService._send_event({"event": "heartbeat"}) is False + post_mock.assert_not_called() + + +def test_send_event_skips_an_empty_fallback_endpoint(telemetry_enabled, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(telemetry_service.dify_config, "TELEMETRY_FALLBACK_ENDPOINT", "") + + def fake_post(url: str, json: dict[str, str], timeout: int): + raise httpx.ConnectError("offline", request=httpx.Request("POST", url)) + + monkeypatch.setattr(telemetry_service.httpx, "post", fake_post) + + assert CommunityTelemetryService._send_event({"event": "heartbeat"}) is False + + +def test_send_event_does_not_retry_the_same_endpoint(telemetry_enabled, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr( + telemetry_service.dify_config, + "TELEMETRY_FALLBACK_ENDPOINT", + telemetry_service.dify_config.TELEMETRY_ENDPOINT, + ) + post_mock = Mock( + return_value=httpx.Response( + 204, + request=httpx.Request("POST", telemetry_service.dify_config.TELEMETRY_ENDPOINT), + ) + ) + monkeypatch.setattr(telemetry_service.httpx, "post", post_mock) + + assert CommunityTelemetryService._send_event({"event": "heartbeat"}) is True + post_mock.assert_called_once() + + +def test_heartbeat_is_not_due_without_instance_id(): + setup = DifySetup(version="1.0.0") + + assert CommunityTelemetryService._is_heartbeat_due(setup, datetime(2026, 7, 13, 12, 0, 0)) is False + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("Linux", "linux"), + ("Plan9", "unknown"), + ], +) +def test_normalize_os(value: str, expected: str): + assert CommunityTelemetryService._normalize_os(value) == expected + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("x86_64", "amd64"), + ("aarch64", "arm64"), + ("armv7l", "arm"), + ("i686", "386"), + ("riscv64", "unknown"), + ], +) +def test_normalize_arch(value: str, expected: str): + assert CommunityTelemetryService._normalize_arch(value) == expected diff --git a/api/tests/unit_tests/services/test_trigger_provider_service.py b/api/tests/unit_tests/services/test_trigger_provider_service.py index ff11bbb3035..90abc9671d9 100644 --- a/api/tests/unit_tests/services/test_trigger_provider_service.py +++ b/api/tests/unit_tests/services/test_trigger_provider_service.py @@ -1,83 +1,120 @@ +"""SQLite-backed tests for trigger provider subscription lifecycle. + +The service intentionally owns short-lived sessions for subscription and OAuth +client operations. Tests bind those session constructors to an isolated SQLite +engine and assert persisted tenant scope, commits, rollbacks, and constraints; +provider daemons, encryption, Redis locks, and caches remain external mocks. +""" + from __future__ import annotations import contextlib import json -import logging +from dataclasses import dataclass from types import SimpleNamespace -from typing import Any -from unittest.mock import MagicMock +from unittest.mock import Mock +from uuid import uuid4 import pytest -from pytest_mock import MockerFixture +from sqlalchemy import func, select +from sqlalchemy.engine import Engine +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session, sessionmaker from constants import HIDDEN_VALUE from core.plugin.entities.plugin_daemon import CredentialType +from core.trigger.entities.entities import Subscription as TriggerSubscriptionEntity +from models.base import TypeBase from models.provider_ids import TriggerProviderID +from models.trigger import ( + TriggerOAuthSystemClient, + TriggerOAuthTenantClient, + TriggerSubscription, + WorkflowPluginTrigger, +) +from services.trigger import trigger_provider_service as service_module from services.trigger.trigger_provider_service import TriggerProviderService -def _patch_redis_lock(mocker: MockerFixture) -> None: - mock_redis = mocker.patch("services.trigger.trigger_provider_service.redis_client") - mock_redis.lock.return_value = contextlib.nullcontext() +@dataclass(frozen=True) +class _DatabaseBinding: + engine: Engine -def _mock_get_trigger_provider(mocker: MockerFixture, provider: object | None) -> None: - mocker.patch( - "services.trigger.trigger_provider_service.TriggerManager.get_trigger_provider", - return_value=provider, +@dataclass(frozen=True) +class TriggerDatabase: + """Factory and identifiers for persisted subscription lifecycle state.""" + + session_maker: sessionmaker[Session] + tenant_id: str + other_tenant_id: str + user_id: str + provider_id: TriggerProviderID + + def add_subscription( + self, + *, + tenant_id: str | None = None, + subscription_id: str | None = None, + name: str = "main", + endpoint_id: str | None = None, + credential_type: CredentialType = CredentialType.API_KEY, + credentials: dict[str, str] | None = None, + properties: dict[str, object] | None = None, + parameters: dict[str, object] | None = None, + credential_expires_at: int = -1, + expires_at: int = -1, + ) -> TriggerSubscription: + subscription = TriggerSubscription( + tenant_id=tenant_id or self.tenant_id, + user_id=self.user_id, + name=name, + endpoint_id=endpoint_id or f"endpoint-{uuid4()}", + provider_id=str(self.provider_id), + parameters=parameters or {"event": "push"}, + properties=properties or {"project": "encrypted"}, + credentials=credentials or {"token": "encrypted"}, + credential_type=credential_type, + credential_expires_at=credential_expires_at, + expires_at=expires_at, + ) + if subscription_id is not None: + subscription.id = subscription_id + with self.session_maker.begin() as session: + session.add(subscription) + return subscription + + def get_subscription(self, subscription_id: str) -> TriggerSubscription | None: + with self.session_maker() as session: + return session.get(TriggerSubscription, subscription_id) + + +@pytest.fixture +def trigger_db(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> TriggerDatabase: + """Create trigger tables and bind every service-owned session to SQLite.""" + + TypeBase.metadata.create_all( + sqlite_engine, + tables=[ + TriggerSubscription.__table__, + WorkflowPluginTrigger.__table__, + TriggerOAuthTenantClient.__table__, + TriggerOAuthSystemClient.__table__, + ], + ) + monkeypatch.setattr(service_module, "db", _DatabaseBinding(engine=sqlite_engine)) + return TriggerDatabase( + session_maker=sessionmaker(bind=sqlite_engine, expire_on_commit=False), + tenant_id=str(uuid4()), + other_tenant_id=str(uuid4()), + user_id=str(uuid4()), + provider_id=TriggerProviderID("langgenius/github/github"), ) -def _encrypter_mock( - *, - decrypted: dict[str, Any] | None = None, - encrypted: dict[str, Any] | None = None, - masked: dict[str, Any] | None = None, -) -> MagicMock: - enc = MagicMock() - enc.decrypt.return_value = decrypted or {} - enc.encrypt.return_value = encrypted or {} - enc.mask_credentials.return_value = masked or {} - enc.mask_plugin_credentials.return_value = masked or {} - return enc - - @pytest.fixture -def provider_id() -> TriggerProviderID: - # Arrange - return TriggerProviderID("langgenius/github/github") - - -@pytest.fixture(autouse=True) -def mock_db_engine(mocker: MockerFixture) -> SimpleNamespace: - # Arrange - mocked_db = SimpleNamespace(engine=object()) - mocker.patch("services.trigger.trigger_provider_service.db", mocked_db) - return mocked_db - - -@pytest.fixture -def mock_session(mocker: MockerFixture) -> MagicMock: - """Mocks the database session context manager used by TriggerProviderService.""" - # Arrange - mock_session_instance = MagicMock() - mock_session_cm = MagicMock() - mock_session_cm.__enter__.return_value = mock_session_instance - mock_session_cm.__exit__.return_value = False - mocker.patch("services.trigger.trigger_provider_service.Session", return_value=mock_session_cm) - mock_begin_cm = MagicMock() - mock_begin_cm.__enter__.return_value = mock_session_instance - mock_begin_cm.__exit__.return_value = False - mock_sessionmaker_instance = MagicMock() - mock_sessionmaker_instance.begin.return_value = mock_begin_cm - mocker.patch("services.trigger.trigger_provider_service.sessionmaker", return_value=mock_sessionmaker_instance) - return mock_session_instance - - -@pytest.fixture -def provider_controller() -> MagicMock: - # Arrange - controller = MagicMock() +def provider_controller() -> Mock: + controller = Mock() controller.get_credential_schema_config.return_value = [] controller.get_properties_schema.return_value = [] controller.get_oauth_client_schema.return_value = [] @@ -85,1129 +122,534 @@ def provider_controller() -> MagicMock: return controller -def test_get_trigger_provider_should_return_api_entity_from_manager( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, -) -> None: - # Arrange - provider = MagicMock() +def _patch_provider(mocker, provider: object) -> None: + mocker.patch.object(service_module.TriggerManager, "get_trigger_provider", return_value=provider) + + +def _patch_lock(mocker) -> None: + redis = mocker.patch.object(service_module, "redis_client") + redis.lock.return_value = contextlib.nullcontext() + + +def _encrypter( + *, + decrypted: dict[str, object] | None = None, + encrypted: dict[str, object] | None = None, + masked: dict[str, object] | None = None, +) -> Mock: + result = Mock() + result.decrypt.side_effect = lambda value: decrypted if decrypted is not None else dict(value) + result.encrypt.side_effect = lambda value: encrypted if encrypted is not None else dict(value) + result.mask_credentials.side_effect = lambda value: masked if masked is not None else dict(value) + result.mask_plugin_credentials.side_effect = lambda value: masked if masked is not None else dict(value) + return result + + +def _patch_identity_encryption(mocker) -> Mock: + encrypter = _encrypter() + cache = Mock() + mocker.patch.object(service_module, "create_provider_encrypter", return_value=(encrypter, cache)) + mocker.patch.object( + service_module, + "create_trigger_provider_encrypter_for_subscription", + return_value=(encrypter, cache), + ) + mocker.patch.object( + service_module, + "create_trigger_provider_encrypter_for_properties", + return_value=(encrypter, cache), + ) + return cache + + +def test_provider_manager_entities_are_forwarded(mocker, trigger_db: TriggerDatabase) -> None: + provider = Mock() provider.to_api_entity.return_value = {"provider": "ok"} - _mock_get_trigger_provider(mocker, provider) - - # Act - result = TriggerProviderService.get_trigger_provider("tenant-1", provider_id) - - # Assert - assert result == {"provider": "ok"} - - -def test_list_trigger_providers_should_return_api_entities_from_manager(mocker: MockerFixture) -> None: - # Arrange - provider_a = MagicMock() - provider_b = MagicMock() - provider_a.to_api_entity.return_value = {"id": "a"} - provider_b.to_api_entity.return_value = {"id": "b"} - mocker.patch( - "services.trigger.trigger_provider_service.TriggerManager.list_all_trigger_providers", - return_value=[provider_a, provider_b], + _patch_provider(mocker, provider) + provider_b = Mock() + provider_b.to_api_entity.return_value = {"provider": "other"} + mocker.patch.object( + service_module.TriggerManager, "list_all_trigger_providers", return_value=[provider, provider_b] ) - # Act - result = TriggerProviderService.list_trigger_providers("tenant-1") - - # Assert - assert result == [{"id": "a"}, {"id": "b"}] + assert TriggerProviderService.get_trigger_provider(trigger_db.tenant_id, trigger_db.provider_id) == { + "provider": "ok" + } + assert TriggerProviderService.list_trigger_providers(trigger_db.tenant_id) == [ + {"provider": "ok"}, + {"provider": "other"}, + ] -def test_list_trigger_provider_subscriptions_should_return_empty_list_when_no_subscriptions( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, +def test_list_subscriptions_empty_state(trigger_db: TriggerDatabase) -> None: + assert ( + TriggerProviderService.list_trigger_provider_subscriptions(trigger_db.tenant_id, trigger_db.provider_id) == [] + ) + + +def test_list_subscriptions_masks_and_counts_distinct_apps( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - mock_session.scalars.return_value.all.return_value = [] + target = trigger_db.add_subscription(subscription_id=str(uuid4())) + trigger_db.add_subscription(tenant_id=trigger_db.other_tenant_id, name="foreign") + with trigger_db.session_maker.begin() as session: + session.add_all( + [ + WorkflowPluginTrigger( + app_id=str(uuid4()), + node_id="node-1", + tenant_id=trigger_db.tenant_id, + provider_id=str(trigger_db.provider_id), + event_name="push", + subscription_id=target.id, + ), + WorkflowPluginTrigger( + app_id=str(uuid4()), + node_id="node-2", + tenant_id=trigger_db.tenant_id, + provider_id=str(trigger_db.provider_id), + event_name="push", + subscription_id=target.id, + ), + WorkflowPluginTrigger( + app_id=str(uuid4()), + node_id="foreign", + tenant_id=trigger_db.other_tenant_id, + provider_id=str(trigger_db.provider_id), + event_name="push", + subscription_id=target.id, + ), + ] + ) + _patch_provider(mocker, provider_controller) + masked = _encrypter(masked={"secret": "****"}) + mocker.patch.object( + service_module, "create_trigger_provider_encrypter_for_subscription", return_value=(masked, Mock()) + ) + mocker.patch.object( + service_module, "create_trigger_provider_encrypter_for_properties", return_value=(masked, Mock()) + ) - # Act - result = TriggerProviderService.list_trigger_provider_subscriptions("tenant-1", provider_id) + subscriptions = TriggerProviderService.list_trigger_provider_subscriptions( + trigger_db.tenant_id, trigger_db.provider_id + ) - # Assert - assert result == [] + assert [item.id for item in subscriptions] == [target.id] + assert subscriptions[0].credentials == {"secret": "****"} + assert subscriptions[0].workflows_in_use == 2 -def test_list_trigger_provider_subscriptions_should_mask_fields_and_attach_workflow_counts( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +@pytest.mark.parametrize("credential_type", [CredentialType.API_KEY, CredentialType.UNAUTHORIZED]) +def test_add_subscription_commits_encrypted_state( + mocker, + trigger_db: TriggerDatabase, + provider_controller: Mock, + credential_type: CredentialType, ) -> None: - # Arrange - api_sub = SimpleNamespace( - id="sub-1", - credentials={"token": "enc"}, - properties={"hook": "enc"}, - parameters={"event": "push"}, - workflows_in_use=0, - ) - db_sub = SimpleNamespace(to_api_entity=lambda: api_sub) - usage_row = SimpleNamespace(subscription_id="sub-1", app_count=2) + _patch_lock(mocker) + _patch_provider(mocker, provider_controller) + encrypter = _encrypter(encrypted={"stored": "encrypted"}) + mocker.patch.object(service_module, "create_provider_encrypter", return_value=(encrypter, Mock())) + subscription_id = str(uuid4()) - mock_session.scalars.return_value.all.return_value = [db_sub] - mock_session.execute.return_value.all.return_value = [usage_row] - - _mock_get_trigger_provider(mocker, provider_controller) - cred_enc = _encrypter_mock(decrypted={"token": "plain"}, masked={"token": "****"}) - prop_enc = _encrypter_mock(decrypted={"hook": "plain"}, masked={"hook": "****"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_subscription", - return_value=(cred_enc, MagicMock()), - ) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_properties", - return_value=(prop_enc, MagicMock()), - ) - - # Act - result = TriggerProviderService.list_trigger_provider_subscriptions("tenant-1", provider_id) - - # Assert - assert len(result) == 1 - assert result[0].credentials == {"token": "****"} - assert result[0].properties == {"hook": "****"} - assert result[0].workflows_in_use == 2 - - -def test_add_trigger_subscription_should_create_subscription_successfully_for_api_key( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - _patch_redis_lock(mocker) - mock_session.scalar.side_effect = [0, None] # count=0, no existing name - - _mock_get_trigger_provider(mocker, provider_controller) - cred_enc = _encrypter_mock(encrypted={"api_key": "enc"}) - prop_enc = _encrypter_mock(encrypted={"project": "enc"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_provider_encrypter", - side_effect=[(cred_enc, MagicMock()), (prop_enc, MagicMock())], - ) - - # Act result = TriggerProviderService.add_trigger_subscription( - tenant_id="tenant-1", - user_id="user-1", + tenant_id=trigger_db.tenant_id, + user_id=trigger_db.user_id, name="main", - provider_id=provider_id, - endpoint_id="endpoint-1", - credential_type=CredentialType.API_KEY, + provider_id=trigger_db.provider_id, + endpoint_id="endpoint-main", + credential_type=credential_type, parameters={"event": "push"}, - properties={"project": "demo"}, - credentials={"api_key": "plain"}, + properties={"project": "plain"}, + credentials={"token": "plain"}, + subscription_id=subscription_id, ) - # Assert - assert result["result"] == "success" - mock_session.add.assert_called_once() + persisted = trigger_db.get_subscription(subscription_id) + assert result == {"result": "success", "id": subscription_id} + assert persisted is not None + assert persisted.properties == {"stored": "encrypted"} + expected_credentials = {} if credential_type == CredentialType.UNAUTHORIZED else {"stored": "encrypted"} + assert persisted.credentials == expected_credentials -def test_add_trigger_subscription_should_store_empty_credentials_for_unauthorized_type( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +def test_add_subscription_limit_rolls_back_without_cross_tenant_count( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - _patch_redis_lock(mocker) - mock_session.scalar.side_effect = [0, None] # count=0, no existing name + for index in range(TriggerProviderService.__MAX_TRIGGER_PROVIDER_COUNT__): + trigger_db.add_subscription(name=f"target-{index}") + for index in range(3): + trigger_db.add_subscription(tenant_id=trigger_db.other_tenant_id, name=f"foreign-{index}") + _patch_lock(mocker) + _patch_provider(mocker, provider_controller) - _mock_get_trigger_provider(mocker, provider_controller) - prop_enc = _encrypter_mock(encrypted={"p": "enc"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_provider_encrypter", - return_value=(prop_enc, MagicMock()), - ) - - # Act - result = TriggerProviderService.add_trigger_subscription( - tenant_id="tenant-1", - user_id="user-1", - name="main", - provider_id=provider_id, - endpoint_id="endpoint-1", - credential_type=CredentialType.UNAUTHORIZED, - parameters={}, - properties={"p": "v"}, - credentials={}, - subscription_id="sub-fixed", - ) - - # Assert - assert result == {"result": "success", "id": "sub-fixed"} - - -def test_add_trigger_subscription_should_raise_error_when_provider_limit_reached( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, - caplog: pytest.LogCaptureFixture, -) -> None: - # Arrange - _patch_redis_lock(mocker) - mock_session.scalar.return_value = TriggerProviderService.__MAX_TRIGGER_PROVIDER_COUNT__ - _mock_get_trigger_provider(mocker, provider_controller) - - # Act + Assert - with caplog.at_level(logging.ERROR, logger="services.trigger.trigger_provider_service"): - with pytest.raises(ValueError, match="Maximum number of providers"): - TriggerProviderService.add_trigger_subscription( - tenant_id="tenant-1", - user_id="user-1", - name="main", - provider_id=provider_id, - endpoint_id="endpoint-1", - credential_type=CredentialType.API_KEY, - parameters={}, - properties={}, - credentials={}, - ) - assert any(r.levelno >= logging.ERROR for r in caplog.records) - - -def test_add_trigger_subscription_should_raise_error_when_name_exists( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - _patch_redis_lock(mocker) - mock_session.scalar.side_effect = [0, object()] # count=0, existing name conflict - _mock_get_trigger_provider(mocker, provider_controller) - - # Act + Assert - with pytest.raises(ValueError, match="Credential name 'main' already exists"): + with pytest.raises(ValueError, match="Maximum number of providers"): TriggerProviderService.add_trigger_subscription( - tenant_id="tenant-1", - user_id="user-1", - name="main", - provider_id=provider_id, - endpoint_id="endpoint-1", - credential_type=CredentialType.API_KEY, + tenant_id=trigger_db.tenant_id, + user_id=trigger_db.user_id, + name="overflow", + provider_id=trigger_db.provider_id, + endpoint_id="overflow", + credential_type=CredentialType.UNAUTHORIZED, parameters={}, properties={}, credentials={}, ) + with trigger_db.session_maker() as session: + count = session.scalar( + select(func.count()) + .select_from(TriggerSubscription) + .where(TriggerSubscription.tenant_id == trigger_db.tenant_id) + ) + assert count == TriggerProviderService.__MAX_TRIGGER_PROVIDER_COUNT__ -def test_update_trigger_subscription_should_raise_error_when_subscription_not_found( - mocker: MockerFixture, - mock_session: MagicMock, + +def test_add_duplicate_name_rolls_back_and_database_constraint_matches_precheck( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - _patch_redis_lock(mocker) - mock_session.scalar.return_value = None + original = trigger_db.add_subscription(name="main") + _patch_lock(mocker) + _patch_provider(mocker, provider_controller) - # Act + Assert - with pytest.raises(ValueError, match="not found"): - TriggerProviderService.update_trigger_subscription("tenant-1", "sub-1") - - -def test_update_trigger_subscription_should_raise_error_when_name_conflicts( - mocker: MockerFixture, - mock_session: MagicMock, - provider_controller: MagicMock, -) -> None: - # Arrange - _patch_redis_lock(mocker) - subscription = SimpleNamespace( - id="sub-1", - name="old", - provider_id="langgenius/github/github", - credential_type=CredentialType.API_KEY, - ) - mock_session.scalar.side_effect = [subscription, object()] # found sub, name conflict - _mock_get_trigger_provider(mocker, provider_controller) - - # Act + Assert with pytest.raises(ValueError, match="already exists"): - TriggerProviderService.update_trigger_subscription("tenant-1", "sub-1", name="new-name") + TriggerProviderService.add_trigger_subscription( + tenant_id=trigger_db.tenant_id, + user_id=trigger_db.user_id, + name="main", + provider_id=trigger_db.provider_id, + endpoint_id="second-endpoint", + credential_type=CredentialType.UNAUTHORIZED, + parameters={}, + properties={}, + credentials={}, + ) + + duplicate = TriggerSubscription( + tenant_id=trigger_db.tenant_id, + user_id=trigger_db.user_id, + name="main", + endpoint_id="constraint-endpoint", + provider_id=str(trigger_db.provider_id), + parameters={}, + properties={}, + credentials={}, + credential_type=CredentialType.UNAUTHORIZED, + ) + with pytest.raises(IntegrityError): + with trigger_db.session_maker.begin() as session: + session.add(duplicate) + assert trigger_db.get_subscription(original.id) is not None -def test_update_trigger_subscription_should_update_fields_and_clear_cache( - mocker: MockerFixture, - mock_session: MagicMock, - provider_controller: MagicMock, +def test_update_subscription_persists_fields_and_preserves_hidden_property( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - _patch_redis_lock(mocker) - subscription = SimpleNamespace( - id="sub-1", - name="old", - tenant_id="tenant-1", - provider_id="langgenius/github/github", - properties={"project": "enc-old"}, - parameters={"event": "old"}, - credentials={"api_key": "enc-old"}, - credential_type=CredentialType.API_KEY, - credential_expires_at=0, - expires_at=0, + subscription = trigger_db.add_subscription(properties={"project": "old-encrypted"}) + _patch_lock(mocker) + _patch_provider(mocker, provider_controller) + properties = _encrypter(decrypted={"project": "old-value"}) + credentials = _encrypter(encrypted={"token": "new-encrypted"}) + mocker.patch.object( + service_module, + "create_provider_encrypter", + side_effect=[(properties, Mock()), (credentials, Mock())], ) - mock_session.scalar.side_effect = [subscription, None] # found sub, no name conflict + clear_cache = mocker.patch.object(service_module, "delete_cache_for_subscription") - _mock_get_trigger_provider(mocker, provider_controller) - prop_enc = _encrypter_mock(decrypted={"project": "old-value"}, encrypted={"project": "new-value"}) - cred_enc = _encrypter_mock(encrypted={"api_key": "new-key"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_provider_encrypter", - side_effect=[(prop_enc, MagicMock()), (cred_enc, MagicMock())], - ) - mock_delete_cache = mocker.patch("services.trigger.trigger_provider_service.delete_cache_for_subscription") - - # Act TriggerProviderService.update_trigger_subscription( - tenant_id="tenant-1", - subscription_id="sub-1", - name="new", + trigger_db.tenant_id, + subscription.id, + name="renamed", properties={"project": HIDDEN_VALUE, "region": "us"}, - parameters={"event": "new"}, - credentials={"api_key": "plain-key"}, + parameters={"event": "issues"}, + credentials={"token": "plain"}, credential_expires_at=100, expires_at=200, ) - # Assert - assert subscription.name == "new" - assert subscription.parameters == {"event": "new"} - assert subscription.credentials == {"api_key": "new-key"} - assert subscription.credential_expires_at == 100 - assert subscription.expires_at == 200 - - mock_delete_cache.assert_called_once() + persisted = trigger_db.get_subscription(subscription.id) + assert persisted is not None + assert persisted.name == "renamed" + assert persisted.properties == {"project": "old-value", "region": "us"} + assert persisted.credentials == {"token": "new-encrypted"} + assert persisted.expires_at == 200 + clear_cache.assert_called_once() -def test_get_subscription_by_id_should_return_none_when_missing(mocker: MockerFixture, mock_session: MagicMock) -> None: - # Arrange - mock_session.scalar.return_value = None - - # Act - result = TriggerProviderService.get_subscription_by_id("tenant-1", "sub-1") - - # Assert - assert result is None - - -def test_get_subscription_by_id_should_decrypt_credentials_and_properties( - mocker: MockerFixture, - mock_session: MagicMock, - provider_controller: MagicMock, +def test_update_missing_and_conflicting_names_leave_rows_unchanged( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - subscription = SimpleNamespace( - id="sub-1", - tenant_id="tenant-1", - provider_id="langgenius/github/github", - credentials={"token": "enc"}, - properties={"project": "enc"}, - ) - mock_session.scalar.return_value = subscription - _mock_get_trigger_provider(mocker, provider_controller) - cred_enc = _encrypter_mock(decrypted={"token": "plain"}) - prop_enc = _encrypter_mock(decrypted={"project": "plain"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_subscription", - return_value=(cred_enc, MagicMock()), - ) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_properties", - return_value=(prop_enc, MagicMock()), - ) + first = trigger_db.add_subscription(name="first") + trigger_db.add_subscription(name="second") + _patch_lock(mocker) + _patch_provider(mocker, provider_controller) - # Act - result = TriggerProviderService.get_subscription_by_id("tenant-1", "sub-1") - - # Assert - assert result is subscription - assert subscription.credentials == {"token": "plain"} - assert subscription.properties == {"project": "plain"} - - -def test_delete_trigger_provider_should_raise_error_when_subscription_missing( - mocker: MockerFixture, - mock_session: MagicMock, -) -> None: - # Arrange - mock_session.scalar.return_value = None - - # Act + Assert with pytest.raises(ValueError, match="not found"): - TriggerProviderService.delete_trigger_provider("tenant-1", "sub-1", session=mock_session) + TriggerProviderService.update_trigger_subscription(trigger_db.tenant_id, str(uuid4())) + with pytest.raises(ValueError, match="already exists"): + TriggerProviderService.update_trigger_subscription(trigger_db.tenant_id, first.id, name="second") + assert trigger_db.get_subscription(first.id).name == "first" # type: ignore[union-attr] -def test_delete_trigger_provider_should_delete_and_clear_cache_even_if_unsubscribe_fails( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +def test_get_subscription_scopes_tenant_and_decrypts( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - subscription = SimpleNamespace( - id="sub-1", - user_id="user-1", - provider_id=str(provider_id), - credential_type=CredentialType.OAUTH2, - credentials={"token": "enc"}, - to_entity=lambda: SimpleNamespace(id="sub-1"), + subscription = trigger_db.add_subscription() + _patch_provider(mocker, provider_controller) + credential = _encrypter(decrypted={"token": "plain"}) + properties = _encrypter(decrypted={"project": "plain"}) + mocker.patch.object( + service_module, "create_trigger_provider_encrypter_for_subscription", return_value=(credential, Mock()) ) - mock_session.scalar.return_value = subscription - _mock_get_trigger_provider(mocker, provider_controller) - cred_enc = _encrypter_mock(decrypted={"token": "plain"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_subscription", - return_value=(cred_enc, MagicMock()), + mocker.patch.object( + service_module, "create_trigger_provider_encrypter_for_properties", return_value=(properties, Mock()) ) - mocker.patch( - "services.trigger.trigger_provider_service.TriggerManager.unsubscribe_trigger", - side_effect=RuntimeError("remote fail"), - ) - mock_delete_cache = mocker.patch("services.trigger.trigger_provider_service.delete_cache_for_subscription") - # Act - TriggerProviderService.delete_trigger_provider("tenant-1", "sub-1", session=mock_session) - - # Assert - mock_session.delete.assert_called_once_with(subscription) - mock_delete_cache.assert_called_once() + assert TriggerProviderService.get_subscription_by_id(trigger_db.other_tenant_id, subscription.id) is None + result = TriggerProviderService.get_subscription_by_id(trigger_db.tenant_id, subscription.id) + assert result is not None + assert result.credentials == {"token": "plain"} + assert result.properties == {"project": "plain"} -def test_delete_trigger_provider_should_skip_unsubscribe_for_unauthorized( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +@pytest.mark.parametrize("credential_type", [CredentialType.API_KEY, CredentialType.UNAUTHORIZED]) +def test_delete_subscription_uses_real_caller_transaction( + mocker, + trigger_db: TriggerDatabase, + provider_controller: Mock, + credential_type: CredentialType, ) -> None: - # Arrange - subscription = SimpleNamespace( - id="sub-2", - user_id="user-1", - provider_id=str(provider_id), - credential_type=CredentialType.UNAUTHORIZED, - credentials={}, - to_entity=lambda: SimpleNamespace(id="sub-2"), - ) - mock_session.scalar.return_value = subscription - _mock_get_trigger_provider(mocker, provider_controller) - mock_unsubscribe = mocker.patch("services.trigger.trigger_provider_service.TriggerManager.unsubscribe_trigger") - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_subscription", - return_value=(_encrypter_mock(decrypted={}), MagicMock()), - ) + subscription = trigger_db.add_subscription(credential_type=credential_type) + _patch_provider(mocker, provider_controller) + _patch_identity_encryption(mocker) + unsubscribe = mocker.patch.object(service_module.TriggerManager, "unsubscribe_trigger") + mocker.patch.object(service_module, "delete_cache_for_subscription") - # Act - TriggerProviderService.delete_trigger_provider("tenant-1", "sub-2", session=mock_session) + with trigger_db.session_maker.begin() as session: + TriggerProviderService.delete_trigger_provider(trigger_db.tenant_id, subscription.id, session=session) - # Assert - mock_unsubscribe.assert_not_called() - mock_session.delete.assert_called_once_with(subscription) + assert trigger_db.get_subscription(subscription.id) is None + if credential_type == CredentialType.UNAUTHORIZED: + unsubscribe.assert_not_called() + else: + unsubscribe.assert_called_once() -def test_refresh_oauth_token_should_raise_error_when_subscription_missing( - mocker: MockerFixture, mock_session: MagicMock +def test_refresh_oauth_token_persists_credentials_after_commit( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - mock_session.scalar.return_value = None + subscription = trigger_db.add_subscription(credential_type=CredentialType.OAUTH2) + _patch_provider(mocker, provider_controller) + encrypter = _encrypter(decrypted={"refresh": "old"}, encrypted={"access": "new"}) + mocker.patch.object(service_module, "create_provider_encrypter", return_value=(encrypter, Mock())) + mocker.patch.object(TriggerProviderService, "get_oauth_client", return_value={"client": "system"}) + handler = Mock() + handler.refresh_credentials.return_value = SimpleNamespace(credentials={"access": "plain"}, expires_at=1234) + mocker.patch.object(service_module, "OAuthHandler", return_value=handler) + clear_cache = mocker.patch.object(service_module, "delete_cache_for_subscription") - # Act + Assert + result = TriggerProviderService.refresh_oauth_token(trigger_db.tenant_id, subscription.id) + + persisted = trigger_db.get_subscription(subscription.id) + assert result == {"result": "success", "expires_at": 1234} + assert persisted.credentials == {"access": "new"} # type: ignore[union-attr] + assert persisted.credential_expires_at == 1234 # type: ignore[union-attr] + clear_cache.assert_called_once() + + +def test_refresh_oauth_rejects_missing_and_non_oauth(trigger_db: TriggerDatabase) -> None: with pytest.raises(ValueError, match="not found"): - TriggerProviderService.refresh_oauth_token("tenant-1", "sub-1") + TriggerProviderService.refresh_oauth_token(trigger_db.tenant_id, str(uuid4())) + subscription = trigger_db.add_subscription(credential_type=CredentialType.API_KEY) + with pytest.raises(ValueError, match="Only OAuth"): + TriggerProviderService.refresh_oauth_token(trigger_db.tenant_id, subscription.id) -def test_refresh_oauth_token_should_raise_error_for_non_oauth_credentials( - mocker: MockerFixture, mock_session: MagicMock +def test_refresh_subscription_skips_or_persists_refreshed_properties( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - subscription = SimpleNamespace(credential_type=CredentialType.API_KEY) - mock_session.scalar.return_value = subscription - - # Act + Assert - with pytest.raises(ValueError, match="Only OAuth credentials can be refreshed"): - TriggerProviderService.refresh_oauth_token("tenant-1", "sub-1") - - -def test_refresh_oauth_token_should_refresh_and_persist_new_credentials( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace( - provider_id=str(provider_id), - user_id="user-1", - credential_type=CredentialType.OAUTH2, - credentials={"access_token": "enc"}, - credential_expires_at=0, - ) - mock_session.scalar.return_value = subscription - _mock_get_trigger_provider(mocker, provider_controller) - cache = MagicMock() - cred_enc = _encrypter_mock(decrypted={"access_token": "old"}, encrypted={"access_token": "new"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_provider_encrypter", - return_value=(cred_enc, cache), - ) - mocker.patch.object(TriggerProviderService, "get_oauth_client", return_value={"client_id": "id"}) - mock_delete_cache = mocker.patch("services.trigger.trigger_provider_service.delete_cache_for_subscription") - refreshed = SimpleNamespace(credentials={"access_token": "new"}, expires_at=12345) - oauth_handler = MagicMock() - oauth_handler.refresh_credentials.return_value = refreshed - mocker.patch("services.trigger.trigger_provider_service.OAuthHandler", return_value=oauth_handler) - - # Act - result = TriggerProviderService.refresh_oauth_token("tenant-1", "sub-1") - - # Assert - assert result == {"result": "success", "expires_at": 12345} - assert subscription.credentials == {"access_token": "new"} - assert subscription.credential_expires_at == 12345 - - cache.delete.assert_not_called() - mock_delete_cache.assert_called_once_with( - tenant_id="tenant-1", - provider_id=str(provider_id), - subscription_id="sub-1", - ) - - -def test_refresh_subscription_should_raise_error_when_subscription_missing( - mocker: MockerFixture, mock_session: MagicMock -) -> None: - # Arrange - mock_session.scalar.return_value = None - - # Act + Assert - with pytest.raises(ValueError, match="not found"): - TriggerProviderService.refresh_subscription("tenant-1", "sub-1", now=100) - - -def test_refresh_subscription_should_skip_when_not_due(mocker: MockerFixture, mock_session: MagicMock) -> None: - # Arrange - subscription = SimpleNamespace(expires_at=200) - mock_session.scalar.return_value = subscription - - # Act - result = TriggerProviderService.refresh_subscription("tenant-1", "sub-1", now=100) - - # Assert - assert result == {"result": "skipped", "expires_at": 200} - - -def test_refresh_subscription_should_refresh_and_persist_properties( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace( - id="sub-1", - tenant_id="tenant-1", - endpoint_id="endpoint-1", - expires_at=50, - provider_id=str(provider_id), + skipped = trigger_db.add_subscription(name="future", expires_at=500) + assert TriggerProviderService.refresh_subscription(trigger_db.tenant_id, skipped.id, now=100) == { + "result": "skipped", + "expires_at": 500, + } + due = trigger_db.add_subscription(name="due", expires_at=50) + _patch_provider(mocker, provider_controller) + _patch_identity_encryption(mocker) + provider_controller.refresh_trigger.return_value = TriggerSubscriptionEntity( + expires_at=900, + endpoint="https://example.test/hook", parameters={"event": "push"}, - properties={"p": "enc"}, - credentials={"c": "enc"}, - credential_type=CredentialType.API_KEY, - ) - mock_session.scalar.return_value = subscription - _mock_get_trigger_provider(mocker, provider_controller) - cred_enc = _encrypter_mock(decrypted={"c": "plain"}) - prop_cache = MagicMock() - prop_enc = _encrypter_mock(decrypted={"p": "plain"}, encrypted={"p": "new-enc"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_subscription", - return_value=(cred_enc, MagicMock()), - ) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_properties", - return_value=(prop_enc, prop_cache), - ) - mocker.patch( - "services.trigger.trigger_provider_service.generate_plugin_trigger_endpoint_url", - return_value="https://endpoint", - ) - provider_controller.refresh_trigger.return_value = SimpleNamespace(properties={"p": "new"}, expires_at=999) - - # Act - result = TriggerProviderService.refresh_subscription("tenant-1", "sub-1", now=100) - - # Assert - assert result == {"result": "success", "expires_at": 999} - assert subscription.properties == {"p": "new-enc"} - assert subscription.expires_at == 999 - - prop_cache.delete.assert_called_once() - - -def test_get_oauth_client_should_return_tenant_client_when_available( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - tenant_client = SimpleNamespace(oauth_params={"client_id": "enc"}) - mock_session.scalar.return_value = tenant_client - _mock_get_trigger_provider(mocker, provider_controller) - enc = _encrypter_mock(decrypted={"client_id": "plain"}) - mocker.patch("services.trigger.trigger_provider_service.create_provider_encrypter", return_value=(enc, MagicMock())) - - # Act - result = TriggerProviderService.get_oauth_client("tenant-1", provider_id) - - # Assert - assert result == {"client_id": "plain"} - - -def test_get_oauth_client_should_return_none_when_plugin_not_verified( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - mock_session.scalar.return_value = None # no tenant client; plugin not verified → early return - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch("services.trigger.trigger_provider_service.PluginService.is_plugin_verified", return_value=False) - - # Act - result = TriggerProviderService.get_oauth_client("tenant-1", provider_id) - - # Assert - assert result is None - - -def test_get_oauth_client_should_return_decrypted_system_client_when_verified( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - mock_session.scalar.side_effect = [None, SimpleNamespace(encrypted_oauth_params="enc")] - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch("services.trigger.trigger_provider_service.PluginService.is_plugin_verified", return_value=True) - mocker.patch( - "services.trigger.trigger_provider_service.decrypt_system_params", - return_value={"client_id": "system"}, + properties={"project": "refreshed"}, ) - # Act - result = TriggerProviderService.get_oauth_client("tenant-1", provider_id) + result = TriggerProviderService.refresh_subscription(trigger_db.tenant_id, due.id, now=100) - # Assert - assert result == {"client_id": "system"} + persisted = trigger_db.get_subscription(due.id) + assert result == {"result": "success", "expires_at": 900} + assert persisted.properties == {"project": "refreshed"} # type: ignore[union-attr] -def test_get_oauth_client_should_raise_error_when_system_decryption_fails( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +def test_oauth_client_prefers_enabled_tenant_record( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - mock_session.scalar.side_effect = [None, SimpleNamespace(encrypted_oauth_params="enc")] - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch("services.trigger.trigger_provider_service.PluginService.is_plugin_verified", return_value=True) - mocker.patch( - "services.trigger.trigger_provider_service.decrypt_system_params", - side_effect=RuntimeError("bad data"), + with trigger_db.session_maker.begin() as session: + session.add( + TriggerOAuthTenantClient( + tenant_id=trigger_db.tenant_id, + plugin_id=trigger_db.provider_id.plugin_id, + provider=trigger_db.provider_id.provider_name, + enabled=True, + encrypted_oauth_params=json.dumps({"client": "encrypted"}), + ) + ) + _patch_provider(mocker, provider_controller) + mocker.patch.object( + service_module, "create_provider_encrypter", return_value=(_encrypter(decrypted={"client": "tenant"}), Mock()) ) - # Act + Assert - with pytest.raises(ValueError, match="Error decrypting system oauth params"): - TriggerProviderService.get_oauth_client("tenant-1", provider_id) + assert TriggerProviderService.get_oauth_client(trigger_db.tenant_id, trigger_db.provider_id) == {"client": "tenant"} -def test_is_oauth_system_client_exists_should_return_false_when_unverified( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +def test_oauth_client_falls_back_to_verified_system_record( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch("services.trigger.trigger_provider_service.PluginService.is_plugin_verified", return_value=False) + with trigger_db.session_maker.begin() as session: + session.add( + TriggerOAuthSystemClient( + plugin_id=trigger_db.provider_id.plugin_id, + provider=trigger_db.provider_id.provider_name, + encrypted_oauth_params="system-encrypted", + ) + ) + _patch_provider(mocker, provider_controller) + mocker.patch.object(service_module.PluginService, "is_plugin_verified", return_value=True) + mocker.patch.object(service_module, "decrypt_system_params", return_value={"client": "system"}) - # Act - result = TriggerProviderService.is_oauth_system_client_exists("tenant-1", provider_id) - - # Assert - assert result is False + assert TriggerProviderService.get_oauth_client(trigger_db.tenant_id, trigger_db.provider_id) == {"client": "system"} + assert TriggerProviderService.is_oauth_system_client_exists(trigger_db.tenant_id, trigger_db.provider_id) -@pytest.mark.parametrize("has_client", [True, False]) -def test_is_oauth_system_client_exists_should_reflect_database_record( - has_client: bool, - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +def test_unverified_plugin_cannot_read_system_oauth( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - mock_session.scalar.return_value = object() if has_client else None - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch("services.trigger.trigger_provider_service.PluginService.is_plugin_verified", return_value=True) + _patch_provider(mocker, provider_controller) + mocker.patch.object(service_module.PluginService, "is_plugin_verified", return_value=False) - # Act - result = TriggerProviderService.is_oauth_system_client_exists("tenant-1", provider_id) - - # Assert - assert result is has_client + assert TriggerProviderService.get_oauth_client(trigger_db.tenant_id, trigger_db.provider_id) is None + assert not TriggerProviderService.is_oauth_system_client_exists(trigger_db.tenant_id, trigger_db.provider_id) -def test_save_custom_oauth_client_params_should_return_success_when_nothing_to_update( - provider_id: TriggerProviderID, +def test_custom_oauth_client_create_mask_enable_and_delete( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - # Act - result = TriggerProviderService.save_custom_oauth_client_params("tenant-1", provider_id, None, None) + _patch_provider(mocker, provider_controller) + encrypter = _encrypter( + encrypted={"client_secret": "encrypted"}, decrypted={"client_secret": "plain"}, masked={"client_secret": "****"} + ) + cache = Mock() + mocker.patch.object(service_module, "create_provider_encrypter", return_value=(encrypter, cache)) - # Assert - assert result == {"result": "success"} - - -def test_save_custom_oauth_client_params_should_create_record_and_clear_params_when_client_params_none( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - mock_session.scalar.return_value = None - _mock_get_trigger_provider(mocker, provider_controller) - fake_model = SimpleNamespace(encrypted_oauth_params="", enabled=False, oauth_params={}) - # Also mock select() so SQLAlchemy doesn't validate the patched TriggerOAuthTenantClient. - mocker.patch("services.trigger.trigger_provider_service.select", MagicMock(return_value=MagicMock())) - mocker.patch("services.trigger.trigger_provider_service.TriggerOAuthTenantClient", return_value=fake_model) - - # Act - result = TriggerProviderService.save_custom_oauth_client_params( - tenant_id="tenant-1", - provider_id=provider_id, - client_params=None, + assert TriggerProviderService.save_custom_oauth_client_params( + trigger_db.tenant_id, + trigger_db.provider_id, + client_params={"client_secret": "plain"}, enabled=True, - ) - - # Assert - assert result == {"result": "success"} - assert fake_model.encrypted_oauth_params == "{}" - assert fake_model.enabled is True - mock_session.add.assert_called_once_with(fake_model) + ) == {"result": "success"} + assert TriggerProviderService.is_oauth_custom_client_enabled(trigger_db.tenant_id, trigger_db.provider_id) + assert TriggerProviderService.get_custom_oauth_client_params(trigger_db.tenant_id, trigger_db.provider_id) == { + "client_secret": "****" + } + assert TriggerProviderService.delete_custom_oauth_client_params(trigger_db.tenant_id, trigger_db.provider_id) == { + "result": "success" + } + assert TriggerProviderService.get_custom_oauth_client_params(trigger_db.tenant_id, trigger_db.provider_id) == {} -def test_save_custom_oauth_client_params_should_merge_hidden_values_and_delete_cache( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +def test_endpoint_lookup_decrypts_persisted_subscription( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - custom_client = SimpleNamespace(oauth_params={"client_id": "enc-old"}, enabled=False) - mock_session.scalar.return_value = custom_client - _mock_get_trigger_provider(mocker, provider_controller) - cache = MagicMock() - enc = _encrypter_mock(decrypted={"client_id": "old-id"}, encrypted={"client_id": "new-id"}) - mocker.patch( - "services.trigger.trigger_provider_service.create_provider_encrypter", - return_value=(enc, cache), - ) + subscription = trigger_db.add_subscription(endpoint_id="lookup-endpoint") + _patch_provider(mocker, provider_controller) + _patch_identity_encryption(mocker) - # Act - result = TriggerProviderService.save_custom_oauth_client_params( - tenant_id="tenant-1", - provider_id=provider_id, - client_params={"client_id": HIDDEN_VALUE, "client_secret": "new"}, - enabled=None, - ) - - # Assert - assert result == {"result": "success"} - assert json.loads(custom_client.encrypted_oauth_params) == {"client_id": "new-id"} - cache.delete.assert_called_once() + assert TriggerProviderService.get_subscription_by_endpoint("missing") is None + found = TriggerProviderService.get_subscription_by_endpoint("lookup-endpoint") + assert found is not None + assert found.id == subscription.id -def test_get_custom_oauth_client_params_should_return_empty_when_record_missing( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, +@pytest.mark.parametrize("valid", [True, False]) +def test_verify_api_key_credentials_uses_persisted_subscription( + mocker, + trigger_db: TriggerDatabase, + provider_controller: Mock, + valid: bool, ) -> None: - # Arrange - mock_session.scalar.return_value = None + subscription = trigger_db.add_subscription(credentials={"token": "old"}) + _patch_provider(mocker, provider_controller) + _patch_identity_encryption(mocker) + if not valid: + provider_controller.validate_credentials.side_effect = RuntimeError("denied") - # Act - result = TriggerProviderService.get_custom_oauth_client_params("tenant-1", provider_id) - - # Assert - assert result == {} - - -def test_get_custom_oauth_client_params_should_return_masked_decrypted_values( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - custom_client = SimpleNamespace(oauth_params={"client_id": "enc"}) - mock_session.scalar.return_value = custom_client - _mock_get_trigger_provider(mocker, provider_controller) - enc = _encrypter_mock(decrypted={"client_id": "plain"}, masked={"client_id": "pl***id"}) - mocker.patch("services.trigger.trigger_provider_service.create_provider_encrypter", return_value=(enc, MagicMock())) - - # Act - result = TriggerProviderService.get_custom_oauth_client_params("tenant-1", provider_id) - - # Assert - assert result == {"client_id": "pl***id"} - - -def test_delete_custom_oauth_client_params_should_delete_record_and_commit( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, -) -> None: - # Act - result = TriggerProviderService.delete_custom_oauth_client_params("tenant-1", provider_id) - - # Assert - assert result == {"result": "success"} - - -@pytest.mark.parametrize("exists", [True, False]) -def test_is_oauth_custom_client_enabled_should_return_expected_boolean( - exists: bool, - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, -) -> None: - # Arrange - mock_session.scalar.return_value = object() if exists else None - - # Act - result = TriggerProviderService.is_oauth_custom_client_enabled("tenant-1", provider_id) - - # Assert - assert result is exists - - -def test_get_subscription_by_endpoint_should_return_none_when_not_found( - mocker: MockerFixture, mock_session: MagicMock -) -> None: - # Arrange - mock_session.scalar.return_value = None - - # Act - result = TriggerProviderService.get_subscription_by_endpoint("endpoint-1") - - # Assert - assert result is None - - -def test_get_subscription_by_endpoint_should_decrypt_credentials_and_properties( - mocker: MockerFixture, - mock_session: MagicMock, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace( - tenant_id="tenant-1", - provider_id="langgenius/github/github", - credentials={"token": "enc"}, - properties={"hook": "enc"}, - ) - mock_session.scalar.return_value = subscription - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_subscription", - return_value=(_encrypter_mock(decrypted={"token": "plain"}), MagicMock()), - ) - mocker.patch( - "services.trigger.trigger_provider_service.create_trigger_provider_encrypter_for_properties", - return_value=(_encrypter_mock(decrypted={"hook": "plain"}), MagicMock()), - ) - - # Act - result = TriggerProviderService.get_subscription_by_endpoint("endpoint-1") - - # Assert - assert result is subscription - assert subscription.credentials == {"token": "plain"} - assert subscription.properties == {"hook": "plain"} - - -def test_verify_subscription_credentials_should_raise_when_provider_not_found( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, -) -> None: - # Arrange - _mock_get_trigger_provider(mocker, None) - - # Act + Assert - with pytest.raises(ValueError, match="Provider .* not found"): - TriggerProviderService.verify_subscription_credentials( - tenant_id="tenant-1", - user_id="user-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={}, + if valid: + assert TriggerProviderService.verify_subscription_credentials( + trigger_db.tenant_id, + trigger_db.user_id, + trigger_db.provider_id, + subscription.id, + {"token": HIDDEN_VALUE}, + ) == {"verified": True} + provider_controller.validate_credentials.assert_called_once_with( + trigger_db.user_id, credentials={"token": "old"} ) + else: + with pytest.raises(ValueError, match="Invalid credentials"): + TriggerProviderService.verify_subscription_credentials( + trigger_db.tenant_id, + trigger_db.user_id, + trigger_db.provider_id, + subscription.id, + {"token": "new"}, + ) -def test_verify_subscription_credentials_should_raise_when_subscription_not_found( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, +def test_rebuild_subscription_preserves_id_endpoint_and_updates_state( + mocker, trigger_db: TriggerDatabase, provider_controller: Mock ) -> None: - # Arrange - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch.object(TriggerProviderService, "get_subscription_by_id", return_value=None) - - # Act + Assert - with pytest.raises(ValueError, match="Subscription sub-1 not found"): - TriggerProviderService.verify_subscription_credentials( - tenant_id="tenant-1", - user_id="user-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={}, - ) - - -def test_verify_subscription_credentials_should_raise_when_api_key_validation_fails( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace(credential_type=CredentialType.API_KEY, credentials={"api_key": "old"}) - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch.object(TriggerProviderService, "get_subscription_by_id", return_value=subscription) - provider_controller.validate_credentials.side_effect = RuntimeError("bad credentials") - - # Act + Assert - with pytest.raises(ValueError, match="Invalid credentials: bad credentials"): - TriggerProviderService.verify_subscription_credentials( - tenant_id="tenant-1", - user_id="user-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={"api_key": HIDDEN_VALUE}, - ) - - -def test_verify_subscription_credentials_should_return_verified_when_api_key_validation_succeeds( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace(credential_type=CredentialType.API_KEY, credentials={"api_key": "old"}) - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch.object(TriggerProviderService, "get_subscription_by_id", return_value=subscription) - - # Act - result = TriggerProviderService.verify_subscription_credentials( - tenant_id="tenant-1", - user_id="user-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={"api_key": HIDDEN_VALUE}, + subscription = trigger_db.add_subscription(endpoint_id="stable-endpoint", credentials={"token": "old"}) + _patch_provider(mocker, provider_controller) + _patch_lock(mocker) + _patch_identity_encryption(mocker) + mocker.patch.object( + service_module.TriggerManager, "unsubscribe_trigger", return_value=SimpleNamespace(success=True) ) - - # Assert - assert result == {"verified": True} - - -def test_verify_subscription_credentials_should_return_verified_for_non_api_key_credentials( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace(credential_type=CredentialType.OAUTH2, credentials={}) - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch.object(TriggerProviderService, "get_subscription_by_id", return_value=subscription) - - # Act - result = TriggerProviderService.verify_subscription_credentials( - tenant_id="tenant-1", - user_id="user-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={}, + mocker.patch.object( + service_module.TriggerManager, + "subscribe_trigger", + return_value=TriggerSubscriptionEntity( + expires_at=777, + endpoint="stable-endpoint", + parameters={"event": "issues"}, + properties={"hook": "new"}, + ), ) + mocker.patch.object(service_module, "delete_cache_for_subscription") - # Assert - assert result == {"verified": True} - - -def test_rebuild_trigger_subscription_should_raise_when_provider_not_found( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, -) -> None: - # Arrange - _mock_get_trigger_provider(mocker, None) - - # Act + Assert - with pytest.raises(ValueError, match="Provider .* not found"): - TriggerProviderService.rebuild_trigger_subscription( - tenant_id="tenant-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={}, - parameters={}, - ) - - -def test_rebuild_trigger_subscription_should_raise_when_subscription_not_found( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch.object(TriggerProviderService, "get_subscription_by_id", return_value=None) - - # Act + Assert - with pytest.raises(ValueError, match="Subscription sub-1 not found"): - TriggerProviderService.rebuild_trigger_subscription( - tenant_id="tenant-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={}, - parameters={}, - ) - - -def test_rebuild_trigger_subscription_should_raise_for_unsupported_credential_type( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace(credential_type=CredentialType.UNAUTHORIZED) - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch.object(TriggerProviderService, "get_subscription_by_id", return_value=subscription) - - # Act + Assert - with pytest.raises(ValueError, match="not supported for auto creation"): - TriggerProviderService.rebuild_trigger_subscription( - tenant_id="tenant-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={}, - parameters={}, - ) - - -def test_rebuild_trigger_subscription_should_raise_when_unsubscribe_fails( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace( - id="sub-1", - user_id="user-1", - endpoint_id="endpoint-1", - credential_type=CredentialType.API_KEY, - credentials={"api_key": "old"}, - to_entity=lambda: SimpleNamespace(id="sub-1"), - ) - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch.object(TriggerProviderService, "get_subscription_by_id", return_value=subscription) - mocker.patch( - "services.trigger.trigger_provider_service.TriggerManager.unsubscribe_trigger", - return_value=SimpleNamespace(success=False, message="remote error"), - ) - - # Act + Assert - with pytest.raises(ValueError, match="Failed to delete previous subscription"): - TriggerProviderService.rebuild_trigger_subscription( - tenant_id="tenant-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={}, - parameters={}, - ) - - -def test_rebuild_trigger_subscription_should_resubscribe_and_update_existing_subscription( - mocker: MockerFixture, - mock_session: MagicMock, - provider_id: TriggerProviderID, - provider_controller: MagicMock, -) -> None: - # Arrange - subscription = SimpleNamespace( - id="sub-1", - user_id="user-1", - endpoint_id="endpoint-1", - credential_type=CredentialType.API_KEY, - credentials={"api_key": "old-key"}, - to_entity=lambda: SimpleNamespace(id="sub-1"), - ) - new_subscription = SimpleNamespace(properties={"project": "new"}, expires_at=888) - _mock_get_trigger_provider(mocker, provider_controller) - mocker.patch.object(TriggerProviderService, "get_subscription_by_id", return_value=subscription) - mocker.patch( - "services.trigger.trigger_provider_service.TriggerManager.unsubscribe_trigger", - return_value=SimpleNamespace(success=True, message="ok"), - ) - mock_subscribe = mocker.patch( - "services.trigger.trigger_provider_service.TriggerManager.subscribe_trigger", - return_value=new_subscription, - ) - mocker.patch( - "services.trigger.trigger_provider_service.generate_plugin_trigger_endpoint_url", - return_value="https://endpoint", - ) - mock_update = mocker.patch.object(TriggerProviderService, "update_trigger_subscription") - - # Act TriggerProviderService.rebuild_trigger_subscription( - tenant_id="tenant-1", - provider_id=provider_id, - subscription_id="sub-1", - credentials={"api_key": HIDDEN_VALUE, "region": "us"}, - parameters={"event": "push"}, - name="updated", + trigger_db.tenant_id, + trigger_db.provider_id, + subscription.id, + credentials={"token": HIDDEN_VALUE}, + parameters={"event": "issues"}, + name="rebuilt", ) - # Assert - call_kwargs = mock_subscribe.call_args.kwargs - assert call_kwargs["credentials"]["api_key"] == "old-key" - assert call_kwargs["credentials"]["region"] == "us" - mock_update.assert_called_once_with( - tenant_id="tenant-1", - subscription_id="sub-1", - name="updated", - parameters={"event": "push"}, - credentials={"api_key": "old-key", "region": "us"}, - properties={"project": "new"}, - expires_at=888, - ) + persisted = trigger_db.get_subscription(subscription.id) + assert persisted is not None + assert persisted.endpoint_id == "stable-endpoint" + assert persisted.name == "rebuilt" + assert persisted.credentials == {"token": "old"} + assert persisted.properties == {"hook": "new"} + assert persisted.expires_at == 777 diff --git a/api/tests/unit_tests/services/test_trigger_subscription_builder_service.py b/api/tests/unit_tests/services/test_trigger_subscription_builder_service.py new file mode 100644 index 00000000000..6316de75792 --- /dev/null +++ b/api/tests/unit_tests/services/test_trigger_subscription_builder_service.py @@ -0,0 +1,251 @@ +from collections.abc import Callable +from contextlib import nullcontext +from unittest.mock import Mock, patch + +import pytest + +from core.plugin.entities.plugin_daemon import CredentialType +from core.trigger.entities.entities import SubscriptionBuilder, SubscriptionBuilderUpdater +from core.trigger.trigger_manager import TriggerManager +from models.provider_ids import TriggerProviderID +from services.trigger.trigger_subscription_builder_service import TriggerSubscriptionBuilderService + +PROVIDER_ID = TriggerProviderID("org/plugin/provider") + + +def subscription_builder() -> SubscriptionBuilder: + return SubscriptionBuilder( + id="builder-1", + name="Builder", + tenant_id="tenant-1", + user_id="user-1", + provider_id=str(PROVIDER_ID), + endpoint_id="builder-1", + parameters={}, + properties={}, + credentials={}, + credential_type=CredentialType.UNAUTHORIZED, + credential_expires_at=-1, + expires_at=-1, + ) + + +@pytest.mark.parametrize( + ("field", "value"), + [ + ("id", "other-builder"), + ("tenant_id", "other-tenant"), + ("user_id", "other-user"), + ("provider_id", "org/plugin/other"), + ], +) +def test_get_subscription_builder_rejects_non_owner(field: str, value: str) -> None: + builder = subscription_builder().model_copy(update={field: value}) + with patch.object( + TriggerSubscriptionBuilderService, + "_get_subscription_builder_by_endpoint_id", + return_value=builder, + ): + with pytest.raises(ValueError, match="not found"): + TriggerSubscriptionBuilderService.get_subscription_builder( + tenant_id="tenant-1", + user_id="user-1", + provider_id=PROVIDER_ID, + subscription_builder_id="builder-1", + ) + + +def test_get_subscription_builder_accepts_owner() -> None: + builder = subscription_builder() + with patch.object( + TriggerSubscriptionBuilderService, + "_get_subscription_builder_by_endpoint_id", + return_value=builder, + ): + assert ( + TriggerSubscriptionBuilderService.get_subscription_builder( + tenant_id="tenant-1", + user_id="user-1", + provider_id=PROVIDER_ID, + subscription_builder_id="builder-1", + ) + is builder + ) + + +def test_get_subscription_builder_returns_none_when_temporary_builder_is_absent() -> None: + with patch.object( + TriggerSubscriptionBuilderService, + "_get_subscription_builder_by_endpoint_id", + return_value=None, + ): + assert ( + TriggerSubscriptionBuilderService.get_subscription_builder( + tenant_id="tenant-1", + user_id="user-1", + provider_id=PROVIDER_ID, + subscription_builder_id="builder-1", + ) + is None + ) + + +def test_get_subscription_builder_rejects_mismatched_endpoint() -> None: + builder = subscription_builder().model_copy(update={"endpoint_id": "other-builder"}) + with patch( + "services.trigger.trigger_subscription_builder_service.redis_client.get", + return_value=builder.model_dump_json(), + ): + assert TriggerSubscriptionBuilderService._get_subscription_builder_by_endpoint_id("builder-1") is None + + +def test_get_subscription_builder_accepts_matching_endpoint() -> None: + builder = subscription_builder() + with patch( + "services.trigger.trigger_subscription_builder_service.redis_client.get", + return_value=builder.model_dump_json(), + ): + assert TriggerSubscriptionBuilderService._get_subscription_builder_by_endpoint_id("builder-1") == builder + + +@pytest.mark.parametrize( + ("operation", "needs_updater"), + [ + (TriggerSubscriptionBuilderService.update_trigger_subscription_builder, True), + (TriggerSubscriptionBuilderService.update_and_verify_builder, True), + (TriggerSubscriptionBuilderService.update_and_build_builder, True), + (TriggerSubscriptionBuilderService.list_logs, False), + (TriggerSubscriptionBuilderService.get_subscription_builder_by_id, False), + ], +) +def test_owner_scoped_operations_reject_missing_builder_before_side_effects( + operation: Callable[..., object], needs_updater: bool +) -> None: + kwargs: dict[str, object] = { + "tenant_id": "tenant-1", + "user_id": "user-1", + "provider_id": PROVIDER_ID, + "subscription_builder_id": "builder-1", + } + if needs_updater: + kwargs["subscription_builder_updater"] = SubscriptionBuilderUpdater(name="Updated") + + with ( + patch.object(TriggerManager, "get_trigger_provider", return_value=Mock()), + patch.object(TriggerSubscriptionBuilderService, "acquire_builder_lock", return_value=nullcontext()), + patch.object( + TriggerSubscriptionBuilderService, + "get_subscription_builder", + return_value=None, + ) as get_subscription_builder, + patch("services.trigger.trigger_subscription_builder_service.redis_client.setex") as setex, + patch("services.trigger.trigger_subscription_builder_service.redis_client.delete") as delete, + ): + with pytest.raises(ValueError, match="not found"): + operation(**kwargs) + + get_subscription_builder.assert_called_once_with( + tenant_id="tenant-1", + user_id="user-1", + provider_id=PROVIDER_ID, + subscription_builder_id="builder-1", + ) + setex.assert_not_called() + delete.assert_not_called() + + +def test_update_and_build_uses_the_owned_builder_without_refetching() -> None: + builder = subscription_builder() + cache_key = TriggerSubscriptionBuilderService.encode_cache_key(builder.id) + + with ( + patch.object(TriggerManager, "get_trigger_provider", return_value=Mock()), + patch.object(TriggerSubscriptionBuilderService, "acquire_builder_lock", return_value=nullcontext()), + patch.object( + TriggerSubscriptionBuilderService, + "get_subscription_builder", + return_value=builder, + ) as get_subscription_builder, + patch("services.trigger.trigger_subscription_builder_service.redis_client.setex") as setex, + patch( + "services.trigger.trigger_subscription_builder_service.TriggerProviderService.add_trigger_subscription" + ) as add_subscription, + patch("services.trigger.trigger_subscription_builder_service.redis_client.delete") as delete, + ): + TriggerSubscriptionBuilderService.update_and_build_builder( + tenant_id="tenant-1", + user_id="user-1", + provider_id=PROVIDER_ID, + subscription_builder_id=builder.id, + subscription_builder_updater=SubscriptionBuilderUpdater(name="Updated"), + ) + + get_subscription_builder.assert_called_once_with( + tenant_id="tenant-1", + user_id="user-1", + provider_id=PROVIDER_ID, + subscription_builder_id=builder.id, + ) + setex.assert_called_once_with(cache_key, 30 * 60, builder.model_dump_json()) + add_subscription.assert_called_once() + subscription_call = add_subscription.call_args.kwargs + assert subscription_call["subscription_id"] == builder.id + assert subscription_call["tenant_id"] == "tenant-1" + assert subscription_call["user_id"] == "user-1" + assert subscription_call["provider_id"] == PROVIDER_ID + assert subscription_call["endpoint_id"] == builder.endpoint_id + assert subscription_call["name"] == "Updated" + delete.assert_called_once_with(cache_key) + + +def test_list_logs_uses_the_owned_builder_endpoint() -> None: + builder = subscription_builder() + logs_key = f"trigger:subscription:builder:logs:{builder.endpoint_id}" + + with ( + patch.object(TriggerSubscriptionBuilderService, "get_subscription_builder", return_value=builder), + patch("services.trigger.trigger_subscription_builder_service.redis_client.get", return_value=None) as redis_get, + ): + assert ( + TriggerSubscriptionBuilderService.list_logs( + tenant_id="tenant-1", + user_id="user-1", + provider_id=PROVIDER_ID, + subscription_builder_id=builder.id, + ) + == [] + ) + + redis_get.assert_called_once_with(logs_key) + + +def test_process_validation_endpoint_uses_the_public_capability() -> None: + builder = subscription_builder() + request = Mock() + response = Mock() + controller = Mock() + controller.dispatch.return_value = Mock(response=response) + + with ( + patch.object( + TriggerSubscriptionBuilderService, + "_get_subscription_builder_by_endpoint_id", + return_value=builder, + ) as get_by_endpoint, + patch.object(TriggerSubscriptionBuilderService, "get_subscription_builder") as get_owned, + patch.object(TriggerManager, "get_trigger_provider", return_value=controller) as get_provider, + patch.object(TriggerSubscriptionBuilderService, "append_log") as append_log, + ): + assert ( + TriggerSubscriptionBuilderService.process_builder_validation_endpoint(builder.endpoint_id, request) + is response + ) + + get_by_endpoint.assert_called_once_with(builder.endpoint_id) + get_owned.assert_not_called() + get_provider.assert_called_once() + provider_call = get_provider.call_args.kwargs + assert provider_call["tenant_id"] == builder.tenant_id + assert str(provider_call["provider_id"]) == str(PROVIDER_ID) + controller.dispatch.assert_called_once() + append_log.assert_called_once() diff --git a/api/tests/unit_tests/services/test_turnstile_service.py b/api/tests/unit_tests/services/test_turnstile_service.py new file mode 100644 index 00000000000..795b914034f --- /dev/null +++ b/api/tests/unit_tests/services/test_turnstile_service.py @@ -0,0 +1,122 @@ +from unittest.mock import MagicMock + +import httpx +import pytest +from pydantic import SecretStr + +from services.turnstile_service import ( + TurnstileChallengeRejectedError, + TurnstileService, + TurnstileUpstreamError, +) + + +@pytest.fixture(autouse=True) +def configure_turnstile(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("services.turnstile_service.dify_config.TURNSTILE_SECRET_KEY", SecretStr("test-secret")) + monkeypatch.setattr("services.turnstile_service.dify_config.TURNSTILE_ALLOWED_HOSTNAMES", "dify.dev") + + +def mock_response(monkeypatch: pytest.MonkeyPatch, *, status_code: int = 200, payload: object) -> MagicMock: + response = httpx.Response( + status_code, + json=payload, + request=httpx.Request("POST", "https://challenges.cloudflare.com/turnstile/v0/siteverify"), + ) + post = MagicMock(return_value=response) + monkeypatch.setattr("services.turnstile_service._http_client.post", post) + return post + + +def test_verify_accepts_subdomain_and_forwards_remote_ip(monkeypatch: pytest.MonkeyPatch) -> None: + post = mock_response( + monkeypatch, + payload={"success": True, "action": "signin_code", "hostname": "agent.dify.dev"}, + ) + + TurnstileService.verify(token="verified-token", remote_ip="203.0.113.8") + + post.assert_called_once_with( + "https://challenges.cloudflare.com/turnstile/v0/siteverify", + data={ + "secret": "test-secret", + "response": "verified-token", + "remoteip": "203.0.113.8", + }, + ) + + +@pytest.mark.parametrize("token", [None, "", " ", "x" * 2049]) +def test_verify_rejects_missing_or_oversized_token(monkeypatch: pytest.MonkeyPatch, token: str | None) -> None: + post = MagicMock() + monkeypatch.setattr("services.turnstile_service._http_client.post", post) + + with pytest.raises(TurnstileChallengeRejectedError): + TurnstileService.verify(token=token, remote_ip=None) + + post.assert_not_called() + + +@pytest.mark.parametrize( + "payload", + [ + {"success": False, "error-codes": ["invalid-input-response"]}, + {"success": False, "error-codes": ["timeout-or-duplicate"]}, + {"success": True, "action": "different_action", "hostname": "agent.dify.dev"}, + {"success": True, "action": "signin_code", "hostname": "attacker.example"}, + ], +) +def test_verify_rejects_invalid_challenge(monkeypatch: pytest.MonkeyPatch, payload: object) -> None: + mock_response(monkeypatch, payload=payload) + + with pytest.raises(TurnstileChallengeRejectedError): + TurnstileService.verify(token="invalid-token", remote_ip=None) + + +@pytest.mark.parametrize( + "payload", + [ + {"success": False, "error-codes": ["invalid-input-secret"]}, + {"success": False, "error-codes": ["internal-error"]}, + {"unexpected": "response"}, + ], +) +def test_verify_maps_server_side_failures_to_upstream_error(monkeypatch: pytest.MonkeyPatch, payload: object) -> None: + mock_response(monkeypatch, payload=payload) + + with pytest.raises(TurnstileUpstreamError): + TurnstileService.verify(token="verified-token", remote_ip=None) + + +def test_verify_maps_http_errors_to_upstream_error(monkeypatch: pytest.MonkeyPatch) -> None: + mock_response(monkeypatch, status_code=503, payload={"error": "unavailable"}) + + with pytest.raises(TurnstileUpstreamError): + TurnstileService.verify(token="verified-token", remote_ip=None) + + +def test_verify_maps_timeout_to_upstream_error(monkeypatch: pytest.MonkeyPatch) -> None: + request = httpx.Request("POST", "https://challenges.cloudflare.com/turnstile/v0/siteverify") + monkeypatch.setattr( + "services.turnstile_service._http_client.post", + MagicMock(side_effect=httpx.ReadTimeout("timed out", request=request)), + ) + + with pytest.raises(TurnstileUpstreamError): + TurnstileService.verify(token="verified-token", remote_ip=None) + + +@pytest.mark.parametrize( + ("secret", "allowed_hostnames"), + [(None, "dify.dev"), (SecretStr("test-secret"), "")], +) +def test_verify_fails_closed_when_cloud_configuration_is_missing( + monkeypatch: pytest.MonkeyPatch, + secret: SecretStr | None, + allowed_hostnames: str, +) -> None: + monkeypatch.setattr("services.turnstile_service.dify_config.TURNSTILE_SECRET_KEY", secret) + monkeypatch.setattr("services.turnstile_service.dify_config.TURNSTILE_ALLOWED_HOSTNAMES", allowed_hostnames) + + with pytest.raises(TurnstileUpstreamError): + TurnstileService.verify(token="verified-token", remote_ip=None) diff --git a/api/tests/unit_tests/services/test_variable_truncator.py b/api/tests/unit_tests/services/test_variable_truncator.py index 931e96ef3a7..f2ba1784a28 100644 --- a/api/tests/unit_tests/services/test_variable_truncator.py +++ b/api/tests/unit_tests/services/test_variable_truncator.py @@ -673,3 +673,118 @@ def test_dummy_variable_truncator_methods(): assert isinstance(result, TruncationResult) assert result.result == segment assert result.truncated is False + + +# --------------------------------------------------------------------------- +# Regression tests for langgenius/dify#39218. +# +# Before the fix, ``_truncate_array`` had a "Dirty fix" branch that +# unconditionally appended every ``File`` element to ``truncated_value`` +# *before* the count cap and the byte-budget check, and *before* +# ``used_size`` was ever incremented. That made ``list[File]`` arrays: +# 1. uncapped by ``array_element_limit``, +# 2. uncounted against ``max_size_bytes``, and +# 3. always reported ``truncated=False``. +# The fix routes ``File`` through ``_truncate_json_primitives``'s dedicated +# ``File`` branch, which returns the file as-is with its real serialized +# size, while preserving the count cap and the byte budget. +# --------------------------------------------------------------------------- + + +class TestFileArrayTruncationRegression39218: + """``list[File]`` must respect ``array_element_limit`` and the byte budget.""" + + @pytest.fixture + def truncator(self) -> VariableTruncator: + return VariableTruncator( + array_element_limit=3, + max_size_bytes=1000, + string_length_limit=50, + ) + + @staticmethod + def _make_file(name: str = "f") -> File: + return File( + id=name, + type=FileType.DOCUMENT, + transfer_method=FileTransferMethod.REMOTE_URL, + remote_url=f"https://example.com/{name}.txt", + filename=f"{name}.txt", + extension=".txt", + mime_type="text/plain", + size=1024, + ) + + def test_file_array_respects_element_count_cap(self, truncator: VariableTruncator) -> None: + # Use a target_size larger than ``count * file_size`` so the byte + # budget never binds — only the count cap should fire. + # Each File serializes to ~237 bytes; 3 files = ~713 bytes. + files = [self._make_file(f"f{i}") for i in range(500)] + + result = truncator._truncate_array(files, target_size=10_000_000) + + # Before the fix, all 500 File entries survived (``len(value)==500``, + # ``truncated==False``). After the fix, the array is capped at + # ``array_element_limit=3`` and ``truncated`` flips to True. + assert len(result.value) == 3 + assert result.truncated is True + + def test_file_array_reports_real_used_size(self, truncator: VariableTruncator) -> None: + # Large budget so the count cap fires before the byte budget does. + files = [self._make_file(f"f{i}") for i in range(500)] + + result = truncator._truncate_array(files, target_size=10_000_000) + + # Before the fix, ``used_size`` for a File array was the empty-array + # baseline of 2 bytes (``[]``), regardless of how many File entries + # actually returned. After the fix, ``used_size`` reflects the real + # serialized size of the returned ``File`` payload. + assert result.value_size > 100 + assert result.truncated is True + + def test_file_array_respects_byte_budget(self, truncator: VariableTruncator) -> None: + # Use a small ``target_size`` so the byte budget is the binding + # constraint. Each File serializes to ~237 bytes, so even one File + # blows the 200-byte budget. + files = [self._make_file(f"f{i}") for i in range(50)] + + result = truncator._truncate_array(files, target_size=200) + + # Before the fix, all 50 File entries survived and ``used_size`` + # reported ``2`` (the empty-array baseline). After the fix, the + # loop sees the File payload: ``value_size`` reflects the real + # serialized size, and the loop stops after the first File because + # adding the next one would exceed ``target_size``. + assert len(result.value) == 1 + assert result.value_size > 100 # the File's real serialized size + assert result.value_size <= 250 # in the ballpark of the budget + + def test_mixed_array_counts_files_toward_cap(self, truncator: VariableTruncator) -> None: + mixed: list[object] = [ + self._make_file("f0"), + "a", + self._make_file("f1"), + "b", + self._make_file("f2"), + "c", + self._make_file("f3"), + "d", + ] + + result = truncator._truncate_array(mixed, target_size=10_000_000) + + # 8 items, cap of 3 → exactly 3 items. Files and primitives are + # counted together toward the cap. + assert len(result.value) == 3 + assert result.truncated is True + + def test_single_file_in_array_is_preserved(self, truncator: VariableTruncator) -> None: + result = truncator._truncate_array([self._make_file("only")], target_size=10_000_000) + + # The File itself is not truncated — the dedicated ``File`` branch + # in ``_truncate_json_primitives`` returns the file untouched. Only + # the array-shape accounting changes. + assert len(result.value) == 1 + assert isinstance(result.value[0], File) + assert result.value[0].id == "only" + assert result.truncated is False diff --git a/api/tests/unit_tests/services/test_vector_service.py b/api/tests/unit_tests/services/test_vector_service.py index eb7bd57e720..75d371c8fa9 100644 --- a/api/tests/unit_tests/services/test_vector_service.py +++ b/api/tests/unit_tests/services/test_vector_service.py @@ -4,22 +4,24 @@ from __future__ import annotations import logging from dataclasses import dataclass +from datetime import datetime from typing import Any from unittest.mock import MagicMock import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session import services.vector_service as vector_service_module from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType +from extensions.storage.storage_type import StorageType +from models import UploadFile +from models.dataset import ChildChunk, DatasetProcessRule, SegmentAttachmentBinding +from models.dataset import Document as DatasetDocument +from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, ProcessRuleMode from services.vector_service import VectorService -@dataclass(frozen=True) -class _UploadFileStub: - id: str - name: str - - @dataclass(frozen=True) class _ChildDocStub: page_content: str @@ -77,21 +79,27 @@ def _make_segment( return segment -def _mock_db_session_for_update_multimodel(*, upload_files: list[_UploadFileStub] | None) -> MagicMock: - session = MagicMock(name="session") - - # db.session.execute() is used for delete(SegmentAttachmentBinding).where(...) - session.execute = MagicMock(name="execute") - - # db.session.scalars(select(UploadFile).where(...)).all() returns upload files - session.scalars.return_value.all.return_value = upload_files or [] - - db_mock = MagicMock(name="db") - db_mock.session = session - return db_mock +def _upload_file(*, file_id: str = "file-1", name: str = "img.png") -> UploadFile: + upload_file = UploadFile( + tenant_id="tenant-1", + storage_type=StorageType.LOCAL, + key=f"uploads/{file_id}", + name=name, + size=10, + extension="png", + mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + created_at=datetime(2026, 1, 1), + used=False, + ) + upload_file.id = file_id + return upload_file -def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_segments_vector_regular_indexing_loads_documents_and_keywords( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(is_multimodal=False) segment = _make_segment() @@ -101,7 +109,7 @@ def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(mo monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) VectorService.create_segments_vector( - [["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=MagicMock() + [["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=sqlite_session ) index_processor.load.assert_called_once() @@ -113,7 +121,9 @@ def test_create_segments_vector_regular_indexing_loads_documents_and_keywords(mo assert kwargs["keywords_list"] == [["k1"]] -def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_segments_vector_regular_indexing_loads_multimodal_documents( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(is_multimodal=True) segment = _make_segment( attachments=[ @@ -127,9 +137,8 @@ def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monk factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - session = MagicMock() VectorService.create_segments_vector( - [["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=session + [["k1"]], [segment], dataset, IndexStructureType.PARAGRAPH_INDEX, session=sqlite_session ) assert index_processor.load.call_count == 2 @@ -143,43 +152,62 @@ def test_create_segments_vector_regular_indexing_loads_multimodal_documents(monk assert second_args[1] == [] assert len(second_args[2]) == 2 assert second_kwargs["with_keywords"] is False - segment.get_attachments.assert_called_once_with(session=session) + segment.get_attachments.assert_called_once_with(session=sqlite_session) -def test_create_segments_vector_with_no_segments_does_not_load(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_segments_vector_with_no_segments_does_not_load( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset() index_processor = MagicMock(name="index_processor") factory_instance = MagicMock() factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX, session=MagicMock()) + VectorService.create_segments_vector(None, [], dataset, IndexStructureType.PARAGRAPH_INDEX, session=sqlite_session) index_processor.load.assert_not_called() -def _mock_parent_child_queries( +def _persist_parent_child_rows( + session: Session, *, - dataset_document: object | None, - processing_rule: object | None, -) -> MagicMock: - session = MagicMock(name="session") - - get_dispatch: dict[object, object | None] = { - vector_service_module.DatasetDocument: dataset_document, - vector_service_module.DatasetProcessRule: processing_rule, - } - - def get_side_effect(model: object, pk: object) -> object | None: - return get_dispatch.get(model) - - session.get.side_effect = get_side_effect - db_mock = MagicMock(name="db") - db_mock.session = session - return db_mock + segment: MagicMock, + include_document: bool = True, + include_rule: bool = True, +) -> tuple[DatasetDocument | None, DatasetProcessRule | None]: + document = None + rule = None + if include_document: + document = DatasetDocument( + id=segment.document_id, + tenant_id=segment.tenant_id, + dataset_id=segment.dataset_id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + dataset_process_rule_id="rule-1", + batch="batch-1", + name="Document", + created_from=DocumentCreatedFrom.API, + created_by="user-1", + doc_language="en", + ) + session.add(document) + if include_rule: + rule = DatasetProcessRule( + dataset_id=segment.dataset_id, + mode=ProcessRuleMode.HIERARCHICAL, + rules='{"parent_mode":"full-doc"}', + created_by="user-1", + ) + rule.id = "rule-1" + session.add(rule) + session.flush() + return document, rule def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_explicit_model( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: dataset = _make_dataset( doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, @@ -188,16 +216,9 @@ def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_ex ) segment = _make_segment() - dataset_document = MagicMock(name="dataset_document") - dataset_document.id = segment.document_id - dataset_document.dataset_process_rule_id = "rule-1" - dataset_document.doc_language = "en" - dataset_document.created_by = "user-1" - - processing_rule = MagicMock(name="processing_rule") - processing_rule.to_dict.return_value = {"rules": {}} - - db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) + dataset_document, processing_rule = _persist_parent_child_rows(sqlite_session, segment=segment) + assert dataset_document is not None + assert processing_rule is not None embedding_model_instance = MagicMock(name="embedding_model_instance") model_manager_instance = MagicMock(name="model_manager_instance") @@ -215,18 +236,23 @@ def test_create_segments_vector_parent_child_calls_generate_child_chunks_with_ex monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, session=db_mock.session + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + session=sqlite_session, ) model_manager_instance.get_model_instance.assert_called_once() generate_child_chunks_mock.assert_called_once_with( - segment, dataset_document, dataset, embedding_model_instance, processing_rule, False, session=db_mock.session + segment, dataset_document, dataset, embedding_model_instance, processing_rule, False, session=sqlite_session ) index_processor.load.assert_not_called() def test_create_segments_vector_parent_child_uses_default_embedding_model_when_provider_missing( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: dataset = _make_dataset( doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, @@ -235,15 +261,7 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p ) segment = _make_segment() - dataset_document = MagicMock() - dataset_document.dataset_process_rule_id = "rule-1" - dataset_document.doc_language = "en" - dataset_document.created_by = "user-1" - - processing_rule = MagicMock() - processing_rule.to_dict.return_value = {"rules": {}} - - db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) + _persist_parent_child_rows(sqlite_session, segment=segment) embedding_model_instance = MagicMock() model_manager_instance = MagicMock() @@ -261,7 +279,11 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) VectorService.create_segments_vector( - None, [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, session=db_mock.session + None, + [segment], + dataset, + vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, + session=sqlite_session, ) model_manager_instance.get_default_model_instance.assert_called_once() @@ -271,13 +293,11 @@ def test_create_segments_vector_parent_child_uses_default_embedding_model_when_p def test_create_segments_vector_parent_child_missing_document_logs_warning_and_continues( monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, + sqlite_session: Session, ) -> None: dataset = _make_dataset(doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX) segment = _make_segment() - processing_rule = MagicMock() - db_mock = _mock_parent_child_queries(dataset_document=None, processing_rule=processing_rule) - index_processor = MagicMock() factory_instance = MagicMock() factory_instance.init_index_processor.return_value = index_processor @@ -289,19 +309,19 @@ def test_create_segments_vector_parent_child_missing_document_logs_warning_and_c [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, - session=db_mock.session, + session=sqlite_session, ) assert any(r.levelno >= logging.WARNING for r in caplog.records) index_processor.load.assert_not_called() -def test_create_segments_vector_parent_child_missing_processing_rule_raises(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_segments_vector_parent_child_missing_processing_rule_raises( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX) segment = _make_segment() - dataset_document = MagicMock() - dataset_document.dataset_process_rule_id = "rule-1" - db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=None) + _persist_parent_child_rows(sqlite_session, segment=segment, include_rule=False) with pytest.raises(ValueError, match="No processing rule found"): VectorService.create_segments_vector( @@ -309,20 +329,19 @@ def test_create_segments_vector_parent_child_missing_processing_rule_raises(monk [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, - session=db_mock.session, + session=sqlite_session, ) -def test_create_segments_vector_parent_child_non_high_quality_raises(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_segments_vector_parent_child_non_high_quality_raises( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset( doc_form=vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, indexing_technique=IndexTechniqueType.ECONOMY, ) segment = _make_segment() - dataset_document = MagicMock() - dataset_document.dataset_process_rule_id = "rule-1" - processing_rule = MagicMock() - db_mock = _mock_parent_child_queries(dataset_document=dataset_document, processing_rule=processing_rule) + _persist_parent_child_rows(sqlite_session, segment=segment) with pytest.raises(ValueError, match="not high quality"): VectorService.create_segments_vector( @@ -330,11 +349,13 @@ def test_create_segments_vector_parent_child_non_high_quality_raises(monkeypatch [segment], dataset, vector_service_module.IndexStructureType.PARENT_CHILD_INDEX, - session=db_mock.session, + session=sqlite_session, ) -def test_update_segment_vector_high_quality_uses_vector(monkeypatch: pytest.MonkeyPatch) -> None: +def test_update_segment_vector_high_quality_uses_vector( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY) segment = _make_segment() @@ -342,10 +363,9 @@ def test_update_segment_vector_high_quality_uses_vector(monkeypatch: pytest.Monk vector_cls = MagicMock(return_value=vector_instance) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - session = MagicMock() - VectorService.update_segment_vector(["k"], segment, dataset, session=session) + VectorService.update_segment_vector(["k"], segment, dataset, session=sqlite_session) - vector_cls.assert_called_once_with(dataset=dataset, session=session) + vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session) vector_instance.delete_by_ids.assert_called_once_with([segment.index_node_id]) vector_instance.add_texts.assert_called_once() add_args, add_kwargs = vector_instance.add_texts.call_args @@ -353,41 +373,45 @@ def test_update_segment_vector_high_quality_uses_vector(monkeypatch: pytest.Monk assert add_kwargs["duplicate_check"] is True -def test_update_segment_vector_economy_uses_keyword_with_keywords_list(monkeypatch: pytest.MonkeyPatch) -> None: +def test_update_segment_vector_economy_uses_keyword_with_keywords_list( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY) segment = _make_segment() keyword_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Keyword", MagicMock(return_value=keyword_instance)) - session = MagicMock() - VectorService.update_segment_vector(["a", "b"], segment, dataset, session=session) + VectorService.update_segment_vector(["a", "b"], segment, dataset, session=sqlite_session) - keyword_instance.delete_by_ids.assert_called_once_with([segment.index_node_id], session) + keyword_instance.delete_by_ids.assert_called_once_with([segment.index_node_id], sqlite_session) keyword_instance.add_texts.assert_called_once() args, kwargs = keyword_instance.add_texts.call_args assert len(args[0]) == 1 - assert args[1] is session + assert args[1] is sqlite_session assert kwargs["keywords_list"] == [["a", "b"]] -def test_update_segment_vector_economy_uses_keyword_without_keywords_list(monkeypatch: pytest.MonkeyPatch) -> None: +def test_update_segment_vector_economy_uses_keyword_without_keywords_list( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY) segment = _make_segment() keyword_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Keyword", MagicMock(return_value=keyword_instance)) - session = MagicMock() - VectorService.update_segment_vector(None, segment, dataset, session=session) + VectorService.update_segment_vector(None, segment, dataset, session=sqlite_session) keyword_instance.add_texts.assert_called_once() args, kwargs = keyword_instance.add_texts.call_args assert len(args[0]) == 1 - assert args[1] is session + assert args[1] is sqlite_session assert "keywords_list" not in kwargs -def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch: pytest.MonkeyPatch) -> None: +def test_generate_child_chunks_regenerate_cleans_then_saves_children( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(doc_form=IndexStructureType.PARAGRAPH_INDEX, tenant_id="tenant-1", dataset_id="dataset-1") segment = _make_segment(segment_id="seg-1") @@ -409,13 +433,6 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - child_chunk_ctor = MagicMock(side_effect=lambda **kwargs: kwargs) - monkeypatch.setattr(vector_service_module, "ChildChunk", child_chunk_ctor) - - db_mock = MagicMock() - db_mock.session.add = MagicMock() - db_mock.session.flush = MagicMock() - VectorService.generate_child_chunks( segment=segment, dataset_document=dataset_document, @@ -423,18 +440,20 @@ def test_generate_child_chunks_regenerate_cleans_then_saves_children(monkeypatch embedding_model_instance=MagicMock(), processing_rule=processing_rule, regenerate=True, - session=db_mock.session, + session=sqlite_session, ) index_processor.clean.assert_called_once() _, transform_kwargs = index_processor.transform.call_args assert transform_kwargs["process_rule"]["rules"]["parent_mode"] == vector_service_module.ParentMode.FULL_DOC index_processor.load.assert_called_once() - assert db_mock.session.add.call_count == 2 - db_mock.session.flush.assert_called_once() + stored = sqlite_session.query(ChildChunk).order_by(ChildChunk.position).all() + assert [chunk.content for chunk in stored] == ["c1", "c2"] -def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest.MonkeyPatch) -> None: +def test_generate_child_chunks_flushes_even_when_no_children( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(doc_form=IndexStructureType.PARAGRAPH_INDEX) segment = _make_segment() dataset_document = MagicMock() @@ -450,8 +469,6 @@ def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest factory_instance.init_index_processor.return_value = index_processor monkeypatch.setattr(vector_service_module, "IndexProcessorFactory", MagicMock(return_value=factory_instance)) - db_mock = MagicMock() - VectorService.generate_child_chunks( segment=segment, dataset_document=dataset_document, @@ -459,15 +476,16 @@ def test_generate_child_chunks_commits_even_when_no_children(monkeypatch: pytest embedding_model_instance=MagicMock(), processing_rule=processing_rule, regenerate=False, - session=db_mock.session, + session=sqlite_session, ) index_processor.load.assert_not_called() - db_mock.session.add.assert_not_called() - db_mock.session.flush.assert_called_once() + assert sqlite_session.query(ChildChunk).count() == 0 -def test_create_child_chunk_vector_high_quality_adds_texts(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_child_chunk_vector_high_quality_adds_texts( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY) child_chunk = MagicMock() child_chunk.content = "child" @@ -480,13 +498,12 @@ def test_create_child_chunk_vector_high_quality_adds_texts(monkeypatch: pytest.M vector_cls = MagicMock(return_value=vector_instance) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - session = MagicMock() - VectorService.create_child_chunk_vector(child_chunk, dataset, session=session) - vector_cls.assert_called_once_with(dataset=dataset, session=session) + VectorService.create_child_chunk_vector(child_chunk, dataset, session=sqlite_session) + vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session) vector_instance.add_texts.assert_called_once() -def test_create_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch) -> None: +def test_create_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY) vector_cls = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", vector_cls) @@ -498,11 +515,13 @@ def test_create_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch) child_chunk.document_id = "doc-1" child_chunk.dataset_id = "dataset-1" - VectorService.create_child_chunk_vector(child_chunk, dataset, session=MagicMock()) + VectorService.create_child_chunk_vector(child_chunk, dataset, session=sqlite_session) vector_cls.assert_not_called() -def test_update_child_chunk_vector_high_quality_updates_vector(monkeypatch: pytest.MonkeyPatch) -> None: +def test_update_child_chunk_vector_high_quality_updates_vector( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY) new_chunk = MagicMock() @@ -526,25 +545,24 @@ def test_update_child_chunk_vector_high_quality_updates_vector(monkeypatch: pyte vector_cls = MagicMock(return_value=vector_instance) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - session = MagicMock() - VectorService.update_child_chunk_vector([new_chunk], [upd_chunk], [del_chunk], dataset, session=session) + VectorService.update_child_chunk_vector([new_chunk], [upd_chunk], [del_chunk], dataset, session=sqlite_session) - vector_cls.assert_called_once_with(dataset=dataset, session=session) + vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session) vector_instance.delete_by_ids.assert_called_once_with(["uid", "did"]) vector_instance.add_texts.assert_called_once() docs = vector_instance.add_texts.call_args.args[0] assert len(docs) == 2 -def test_update_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch) -> None: +def test_update_child_chunk_vector_economy_noop(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY) vector_cls = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - VectorService.update_child_chunk_vector([], [], [], dataset, session=MagicMock()) + VectorService.update_child_chunk_vector([], [], [], dataset, session=sqlite_session) vector_cls.assert_not_called() -def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch) -> None: +def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: dataset = _make_dataset() child_chunk = MagicMock() child_chunk.index_node_id = "cid" @@ -553,9 +571,8 @@ def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch vector_cls = MagicMock(return_value=vector_instance) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - session = MagicMock() - VectorService.delete_child_chunk_vector(child_chunk, dataset, session=session) - vector_cls.assert_called_once_with(dataset=dataset, session=session) + VectorService.delete_child_chunk_vector(child_chunk, dataset, session=sqlite_session) + vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session) vector_instance.delete_by_ids.assert_called_once_with(["cid"]) @@ -564,156 +581,159 @@ def test_delete_child_chunk_vector_deletes_by_id(monkeypatch: pytest.MonkeyPatch # --------------------------------------------------------------------------- -def test_update_multimodel_vector_returns_when_not_high_quality(monkeypatch: pytest.MonkeyPatch) -> None: +def test_update_multimodel_vector_returns_when_not_high_quality( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.ECONOMY, is_multimodal=True) segment = _make_segment(tenant_id="t", attachments=[{"id": "a"}]) vector_cls = MagicMock() - db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) VectorService.update_multimodel_vector( - session=db_mock.session, segment=segment, attachment_ids=["a"], dataset=dataset + session=sqlite_session, segment=segment, attachment_ids=["a"], dataset=dataset ) vector_cls.assert_not_called() - db_mock.session.query.assert_not_called() + assert not sqlite_session.in_transaction() -def test_update_multimodel_vector_returns_when_no_actual_change(monkeypatch: pytest.MonkeyPatch) -> None: +def test_update_multimodel_vector_returns_when_no_actual_change( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True) segment = _make_segment(tenant_id="t", attachments=[{"id": "a"}, {"id": "b"}]) vector_cls = MagicMock() - db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) monkeypatch.setattr(vector_service_module, "Vector", vector_cls) VectorService.update_multimodel_vector( - session=db_mock.session, segment=segment, attachment_ids=["b", "a"], dataset=dataset + session=sqlite_session, segment=segment, attachment_ids=["b", "a"], dataset=dataset ) vector_cls.assert_not_called() - db_mock.session.query.assert_not_called() + assert not sqlite_session.in_transaction() def test_update_multimodel_vector_deletes_bindings_and_commits_on_empty_new_ids( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True) segment = _make_segment(tenant_id="tenant-1", attachments=[{"id": "old-1"}, {"id": "old-2"}]) vector_instance = MagicMock(name="vector_instance") vector_cls = MagicMock(return_value=vector_instance) - db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) + sqlite_session.add_all( + [ + SegmentAttachmentBinding( + tenant_id="tenant-1", + dataset_id="dataset-1", + document_id="doc-1", + segment_id="seg-1", + attachment_id=attachment_id, + ) + for attachment_id in ("old-1", "old-2") + ] + ) + sqlite_session.flush() monkeypatch.setattr(vector_service_module, "Vector", vector_cls) - VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset, session=db_mock.session) + VectorService.update_multimodel_vector(segment=segment, attachment_ids=[], dataset=dataset, session=sqlite_session) - vector_cls.assert_called_once_with(dataset=dataset, session=db_mock.session) + vector_cls.assert_called_once_with(dataset=dataset, session=sqlite_session) vector_instance.delete_by_ids.assert_called_once_with(["old-1", "old-2"]) - db_mock.session.execute.assert_called_once() - db_mock.session.flush.assert_called_once() - db_mock.session.add_all.assert_not_called() + assert sqlite_session.query(SegmentAttachmentBinding).count() == 0 vector_instance.add_texts.assert_not_called() -def test_update_multimodel_vector_commits_when_no_upload_files_found(monkeypatch: pytest.MonkeyPatch) -> None: +def test_update_multimodel_vector_flushes_when_no_upload_files_found( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True) segment = _make_segment(tenant_id="tenant-1", attachments=[{"id": "old-1"}]) vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) - db_mock = _mock_db_session_for_update_multimodel(upload_files=[]) - VectorService.update_multimodel_vector( - session=db_mock.session, segment=segment, attachment_ids=["new-1"], dataset=dataset + session=sqlite_session, segment=segment, attachment_ids=["new-1"], dataset=dataset ) - db_mock.session.flush.assert_called_once() - db_mock.session.add_all.assert_not_called() + assert sqlite_session.query(SegmentAttachmentBinding).count() == 0 vector_instance.add_texts.assert_not_called() def test_update_multimodel_vector_adds_bindings_and_vectors_and_skips_missing_upload_files( monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, + sqlite_session: Session, ) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True) segment = _make_segment(segment_id="seg-1", tenant_id="tenant-1", attachments=[{"id": "old-1"}]) vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) - db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) - - binding_ctor = MagicMock(side_effect=lambda **kwargs: kwargs) - monkeypatch.setattr(vector_service_module, "SegmentAttachmentBinding", binding_ctor) - monkeypatch.setattr(vector_service_module, "delete", MagicMock()) - monkeypatch.setattr(vector_service_module, "select", MagicMock()) + sqlite_session.add(_upload_file()) + sqlite_session.flush() with caplog.at_level(logging.WARNING, logger="services.vector_service"): VectorService.update_multimodel_vector( - session=db_mock.session, segment=segment, attachment_ids=["file-1", "missing"], dataset=dataset + session=sqlite_session, segment=segment, attachment_ids=["file-1", "missing"], dataset=dataset ) assert any(r.levelno >= logging.WARNING for r in caplog.records) - db_mock.session.add_all.assert_called_once() - bindings = db_mock.session.add_all.call_args.args[0] + bindings = sqlite_session.query(SegmentAttachmentBinding).all() assert len(bindings) == 1 - assert bindings[0]["attachment_id"] == "file-1" + assert bindings[0].attachment_id == "file-1" vector_instance.create_multimodal.assert_called_once() documents = vector_instance.create_multimodal.call_args.args[0] assert len(documents) == 1 assert documents[0].page_content == "img.png" assert documents[0].metadata["doc_id"] == "file-1" - db_mock.session.flush.assert_called_once() def test_update_multimodel_vector_updates_bindings_without_multimodal_vector_ops( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=False) segment = _make_segment(tenant_id="tenant-1", attachments=[{"id": "old-1"}]) vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) - db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) - monkeypatch.setattr( - vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs) - ) - monkeypatch.setattr(vector_service_module, "delete", MagicMock()) - monkeypatch.setattr(vector_service_module, "select", MagicMock()) + sqlite_session.add(_upload_file()) + sqlite_session.flush() VectorService.update_multimodel_vector( - session=db_mock.session, segment=segment, attachment_ids=["file-1"], dataset=dataset + session=sqlite_session, segment=segment, attachment_ids=["file-1"], dataset=dataset ) vector_instance.delete_by_ids.assert_not_called() vector_instance.add_texts.assert_not_called() - db_mock.session.add_all.assert_called_once() - db_mock.session.flush.assert_called_once() + binding = sqlite_session.query(SegmentAttachmentBinding).one() + assert binding.attachment_id == "file-1" def test_update_multimodel_vector_rolls_back_and_reraises_on_error( monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, + sqlite_session: Session, ) -> None: dataset = _make_dataset(indexing_technique=IndexTechniqueType.HIGH_QUALITY, is_multimodal=True) segment = _make_segment(segment_id="seg-1", tenant_id="tenant-1", attachments=[{"id": "old-1"}]) vector_instance = MagicMock() monkeypatch.setattr(vector_service_module, "Vector", MagicMock(return_value=vector_instance)) - db_mock = _mock_db_session_for_update_multimodel(upload_files=[_UploadFileStub(id="file-1", name="img.png")]) - db_mock.session.flush.side_effect = RuntimeError("boom") - monkeypatch.setattr( - vector_service_module, "SegmentAttachmentBinding", MagicMock(side_effect=lambda **kwargs: kwargs) - ) - monkeypatch.setattr(vector_service_module, "delete", MagicMock()) - monkeypatch.setattr(vector_service_module, "select", MagicMock()) + sqlite_session.add(_upload_file()) + sqlite_session.flush() + rollback_events: list[str] = [] + event.listen(sqlite_session, "after_rollback", lambda _session: rollback_events.append("rollback")) + monkeypatch.setattr(sqlite_session, "flush", MagicMock(side_effect=RuntimeError("boom"))) with caplog.at_level(logging.ERROR, logger="services.vector_service"): with pytest.raises(RuntimeError, match="boom"): VectorService.update_multimodel_vector( - session=db_mock.session, segment=segment, attachment_ids=["file-1"], dataset=dataset + session=sqlite_session, segment=segment, attachment_ids=["file-1"], dataset=dataset ) assert any(r.levelno >= logging.ERROR for r in caplog.records) - db_mock.session.rollback.assert_called_once() + assert rollback_events == ["rollback"] diff --git a/api/tests/unit_tests/services/test_vector_space_admission_service.py b/api/tests/unit_tests/services/test_vector_space_admission_service.py new file mode 100644 index 00000000000..4050be7ae5b --- /dev/null +++ b/api/tests/unit_tests/services/test_vector_space_admission_service.py @@ -0,0 +1,545 @@ +import json +import threading +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace, TracebackType +from typing import cast +from unittest.mock import call, patch + +import pytest +from sqlalchemy.orm import Session + +from configs import dify_config +from core.rag.datasource.vdb.vector_type import VectorType +from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType +from core.rag.models.document import AttachmentDocument, ChildDocument, Document +from enums import CloudPlan, DeploymentEdition +from models.dataset import Dataset +from services.vector_space_admission_service import ( + VECTOR_SPACE_ADMISSION_ERROR_CODE, + VectorSpaceAdmissionError, + VectorSpaceAdmissionService, + VectorStorageWorkload, + build_document_workload, + build_pipeline_workload, + estimate_tidb_storage_bytes, + format_vector_space_admission_error, + get_vector_space_admission_error_fields, + parse_vector_space_estimate_limits, +) + +_MEBIBYTE = 1024 * 1024 +_ESTIMATE_LIMITS = "sandbox:60,professional:6400,team:25600" + + +class _FakeRedisLock: + def __init__(self, lock: threading.Lock) -> None: + self._lock = lock + + def __enter__(self) -> "_FakeRedisLock": + self._lock.acquire() + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + self._lock.release() + + +class _FakeRedis: + def __init__(self) -> None: + self.values: dict[str, str] = {} + self.ttls: dict[str, int] = {} + self._locks: dict[str, threading.Lock] = {} + + def lock(self, key: str, **_kwargs: object) -> _FakeRedisLock: + return _FakeRedisLock(self._locks.setdefault(key, threading.Lock())) + + def get(self, key: str) -> str | None: + return self.values.get(key) + + def setex(self, key: str, ttl: int, value: str) -> None: + self.values[key] = value + self.ttls[key] = ttl + + +def _dataset() -> Dataset: + return cast( + Dataset, + SimpleNamespace( + id="dataset-1", + tenant_id="tenant-1", + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + embedding_model_provider="provider", + embedding_model="model", + index_struct_dict={"type": VectorType.TIDB_ON_QDRANT}, + ), + ) + + +def _workload() -> VectorStorageWorkload: + return VectorStorageWorkload(text_points=1, summary_points=0, probe_text="probe") + + +def _check_estimate( + plan: CloudPlan, + estimated_mb: float, + *, + usage_mb: float = 0, + plan_limit_mb: int = 50, + service: VectorSpaceAdmissionService | None = None, + document_id: str = "document-1", + redis: _FakeRedis | None = None, +) -> VectorSpaceAdmissionService: + service = service or VectorSpaceAdmissionService() + redis = redis or _FakeRedis() + with ( + patch.object(service, "_get_plan", return_value=plan), + patch.object(service, "_get_embedding_dimension", return_value=3072), + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + patch( + "services.vector_space_admission_service.dify_config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB", + _ESTIMATE_LIMITS, + ), + patch( + "services.vector_space_admission_service.Vector.resolve_vector_type", + return_value=VectorType.TIDB_ON_QDRANT, + ), + patch( + "services.vector_space_admission_service.estimate_tidb_storage_bytes", + return_value=estimated_mb * _MEBIBYTE, + ), + patch( + "services.vector_space_admission_service.BillingService.get_vector_space", + return_value={"size": usage_mb, "limit": plan_limit_mb}, + ), + patch("services.vector_space_admission_service.redis_client", redis), + ): + service._ensure_can_write( + dataset=_dataset(), + document_id=document_id, + workload=_workload(), + session=cast(Session, SimpleNamespace()), + ) + return service + + +def test_estimate_tidb_storage_bytes_counts_both_vector_copies_and_point_overhead() -> None: + assert estimate_tidb_storage_bytes(point_count=10, dimension=1536) == 10 * (1536 * 4 * 2 + 3584) + + +def test_parse_vector_space_estimate_limits_supports_all_plans() -> None: + assert parse_vector_space_estimate_limits("sandbox:1,professional:2,team:3") == { + CloudPlan.SANDBOX: 1, + CloudPlan.PROFESSIONAL: 2, + CloudPlan.TEAM: 3, + } + + +@pytest.mark.parametrize( + "value", + [ + "", + "sandbox", + "sandbox:60", + "unknown:60", + "sandbox:not-a-number", + "sandbox:0", + "sandbox:-1", + "sandbox:1,pro:2,team:3", + "pro:6400,professional:6401", + ], +) +def test_parse_vector_space_estimate_limits_rejects_invalid_values(value: str) -> None: + with pytest.raises(ValueError, match="Invalid vector-space estimate limit"): + parse_vector_space_estimate_limits(value) + + +def test_vector_space_admission_error_fields() -> None: + message = format_vector_space_admission_error(61, 50) + + assert get_vector_space_admission_error_fields(message) == { + "error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE, + "estimated_vector_space_mb": 61, + "vector_space_limit_mb": 50, + } + assert get_vector_space_admission_error_fields("another indexing error") == { + "error_code": None, + "estimated_vector_space_mb": None, + "vector_space_limit_mb": None, + } + + +def test_workloads_ignore_images_and_attachments() -> None: + document_workload = build_document_workload( + IndexStructureType.PARAGRAPH_INDEX, + [ + Document( + page_content="text", + attachments=[AttachmentDocument(page_content="image", metadata={"doc_id": "file-1"})], + ) + ], + include_summaries=False, + ) + pipeline_workload = build_pipeline_workload( + IndexStructureType.PARAGRAPH_INDEX, + { + "general_chunks": [ + { + "content": "text ![image](/files/file-1/file-preview)", + "files": [{"id": "file-1"}], + } + ] + }, + include_summaries=False, + ) + + assert document_workload.total_points == 1 + assert pipeline_workload.total_points == 1 + + +def test_parent_child_workload_counts_child_and_summary_vectors() -> None: + workload = build_document_workload( + IndexStructureType.PARENT_CHILD_INDEX, + [ + Document( + page_content="parent-1", + children=[ChildDocument(page_content="child-1"), ChildDocument(page_content="child-2")], + ), + Document(page_content="parent-2", children=[ChildDocument(page_content="child-3")]), + ], + include_summaries=True, + ) + + assert workload.text_points == 3 + assert workload.summary_points == 2 + assert workload.total_points == 5 + + +def test_pipeline_qa_workload_counts_question_vectors_without_summaries() -> None: + workload = build_pipeline_workload( + IndexStructureType.QA_INDEX, + { + "qa_chunks": [ + {"question": "question-1", "answer": "answer-1"}, + {"question": "question-2", "answer": "answer-2"}, + ] + }, + include_summaries=True, + ) + + assert workload.text_points == 2 + assert workload.summary_points == 0 + + +def test_admission_is_cloud_only() -> None: + service = VectorSpaceAdmissionService() + with ( + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), + patch("services.vector_space_admission_service.Vector.resolve_vector_type") as resolve_vector_type, + patch("services.vector_space_admission_service.BillingService.get_info") as get_info, + ): + service._ensure_can_write( + dataset=_dataset(), + document_id="document-1", + workload=_workload(), + session=cast(Session, SimpleNamespace()), + ) + + resolve_vector_type.assert_not_called() + get_info.assert_not_called() + + +def test_admission_skips_non_tidb_vector_backends() -> None: + service = VectorSpaceAdmissionService() + with ( + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + patch("services.vector_space_admission_service.Vector.resolve_vector_type", return_value=VectorType.QDRANT), + patch("services.vector_space_admission_service.BillingService.get_info") as get_info, + ): + service._ensure_can_write( + dataset=_dataset(), + document_id="document-1", + workload=_workload(), + session=cast(Session, SimpleNamespace()), + ) + + get_info.assert_not_called() + + +def test_sandbox_allows_60_mb_estimate() -> None: + _check_estimate(CloudPlan.SANDBOX, 60) + + +def test_sandbox_compares_current_usage_plus_document_estimate() -> None: + _check_estimate(CloudPlan.SANDBOX, 20, usage_mb=40) + + with pytest.raises(VectorSpaceAdmissionError): + _check_estimate(CloudPlan.SANDBOX, 21, usage_mb=40) + + +def test_admission_compares_fractional_usage_without_rounding_down() -> None: + _check_estimate(CloudPlan.SANDBOX, 10.5, usage_mb=49.5) + + with pytest.raises(VectorSpaceAdmissionError) as exc_info: + _check_estimate(CloudPlan.SANDBOX, 10.6, usage_mb=49.5) + + assert get_vector_space_admission_error_fields(str(exc_info.value)) == { + "error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE, + "estimated_vector_space_mb": 61, + "vector_space_limit_mb": 50, + } + + +def test_admission_uses_configured_threshold_above_nominal_limit() -> None: + _check_estimate(CloudPlan.SANDBOX, 10, usage_mb=50) + + with pytest.raises(VectorSpaceAdmissionError): + _check_estimate(CloudPlan.SANDBOX, 10.1, usage_mb=50) + + +@pytest.mark.parametrize( + ("plan", "usage_mb", "allowed_estimate_mb", "rejected_estimate_mb"), + [ + (CloudPlan.PROFESSIONAL, 5000, 1400, 1401), + (CloudPlan.TEAM, 20000, 5600, 5601), + ], +) +def test_paid_plan_projected_usage_boundaries( + plan: CloudPlan, + usage_mb: int, + allowed_estimate_mb: int, + rejected_estimate_mb: int, +) -> None: + _check_estimate(plan, allowed_estimate_mb, usage_mb=usage_mb) + + with pytest.raises(VectorSpaceAdmissionError): + _check_estimate(plan, rejected_estimate_mb, usage_mb=usage_mb) + + +def test_same_batch_accumulates_projected_usage() -> None: + service = VectorSpaceAdmissionService() + redis = _FakeRedis() + _check_estimate( + CloudPlan.SANDBOX, + 10, + usage_mb=40, + service=service, + document_id="document-1", + redis=redis, + ) + _check_estimate( + CloudPlan.SANDBOX, + 10, + usage_mb=40, + service=service, + document_id="document-2", + redis=redis, + ) + + with pytest.raises(VectorSpaceAdmissionError): + _check_estimate( + CloudPlan.SANDBOX, + 1, + usage_mb=40, + service=service, + document_id="document-3", + redis=redis, + ) + + +def test_usage_lookup_is_refreshed_for_each_document() -> None: + service = VectorSpaceAdmissionService() + redis = _FakeRedis() + with ( + patch.object(service, "_get_plan", return_value=CloudPlan.SANDBOX), + patch.object(service, "_get_embedding_dimension", return_value=3072), + patch.object(dify_config, "DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), + patch( + "services.vector_space_admission_service.dify_config.TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB", + _ESTIMATE_LIMITS, + ), + patch( + "services.vector_space_admission_service.Vector.resolve_vector_type", + return_value=VectorType.TIDB_ON_QDRANT, + ), + patch( + "services.vector_space_admission_service.estimate_tidb_storage_bytes", + side_effect=[20 * _MEBIBYTE, 1 * _MEBIBYTE], + ), + patch( + "services.vector_space_admission_service.BillingService.get_vector_space", + side_effect=[{"size": 40.0, "limit": 50}, {"size": 50.0, "limit": 50}], + ) as get_vector_space, + patch("services.vector_space_admission_service.redis_client", redis), + ): + service._ensure_can_write( + dataset=_dataset(), + document_id="document-1", + workload=_workload(), + session=cast(Session, SimpleNamespace()), + ) + with pytest.raises(VectorSpaceAdmissionError): + service._ensure_can_write( + dataset=_dataset(), + document_id="document-2", + workload=_workload(), + session=cast(Session, SimpleNamespace()), + ) + + assert get_vector_space.call_args_list == [call("tenant-1"), call("tenant-1")] + + +def test_independent_services_use_watermark_without_double_counting_fresh_usage() -> None: + redis = _FakeRedis() + _check_estimate( + CloudPlan.SANDBOX, + 10, + usage_mb=40, + service=VectorSpaceAdmissionService(), + document_id="document-1", + redis=redis, + ) + _check_estimate( + CloudPlan.SANDBOX, + 10, + usage_mb=50, + service=VectorSpaceAdmissionService(), + document_id="document-2", + redis=redis, + ) + + with pytest.raises(VectorSpaceAdmissionError): + _check_estimate( + CloudPlan.SANDBOX, + 1, + usage_mb=50, + service=VectorSpaceAdmissionService(), + document_id="document-3", + redis=redis, + ) + + state = json.loads(redis.values["tenant:tenant-1:vector_space_estimate_watermark"]) + assert state["projected_usage_bytes"] == 60 * _MEBIBYTE + assert state["document_ids"] == ["document-1", "document-2"] + assert redis.ttls["tenant:tenant-1:vector_space_estimate_watermark"] == 1800 + + +def test_fresh_usage_above_watermark_becomes_next_projection_base() -> None: + redis = _FakeRedis() + _check_estimate( + CloudPlan.SANDBOX, + 10, + usage_mb=40, + service=VectorSpaceAdmissionService(), + document_id="document-1", + redis=redis, + ) + _check_estimate( + CloudPlan.SANDBOX, + 5, + usage_mb=55, + service=VectorSpaceAdmissionService(), + document_id="document-2", + redis=redis, + ) + + state = json.loads(redis.values["tenant:tenant-1:vector_space_estimate_watermark"]) + assert state["projected_usage_bytes"] == 60 * _MEBIBYTE + + +def test_same_document_is_not_added_to_watermark_twice() -> None: + redis = _FakeRedis() + for _ in range(2): + _check_estimate( + CloudPlan.SANDBOX, + 10, + usage_mb=40, + service=VectorSpaceAdmissionService(), + document_id="document-1", + redis=redis, + ) + + _check_estimate( + CloudPlan.SANDBOX, + 10, + usage_mb=40, + service=VectorSpaceAdmissionService(), + document_id="document-2", + redis=redis, + ) + + state = json.loads(redis.values["tenant:tenant-1:vector_space_estimate_watermark"]) + assert state["projected_usage_bytes"] == 60 * _MEBIBYTE + assert state["document_ids"] == ["document-1", "document-2"] + + +def test_concurrent_services_reserve_watermark_atomically() -> None: + redis = _FakeRedis() + barrier = threading.Barrier(2) + + def reserve(document_id: str) -> bool: + barrier.wait() + _, projected_usage_bytes = VectorSpaceAdmissionService()._reserve_projected_usage( + tenant_id="tenant-1", + document_id=document_id, + current_usage_bytes=40 * _MEBIBYTE, + document_estimate_bytes=15 * _MEBIBYTE, + estimate_limit_bytes=60 * _MEBIBYTE, + ) + return projected_usage_bytes <= 60 * _MEBIBYTE + + with ( + patch("services.vector_space_admission_service.redis_client", redis), + ThreadPoolExecutor(max_workers=2) as executor, + ): + results = list(executor.map(reserve, ["document-1", "document-2"])) + + assert sorted(results) == [False, True] + state = json.loads(redis.values["tenant:tenant-1:vector_space_estimate_watermark"]) + assert state["projected_usage_bytes"] == 55 * _MEBIBYTE + assert len(state["document_ids"]) == 1 + + +@pytest.mark.parametrize( + ("plan", "estimated_mb", "plan_limit_mb"), + [ + (CloudPlan.SANDBOX, 61, 55), + (CloudPlan.PROFESSIONAL, 6401, 6000), + (CloudPlan.TEAM, 25601, 24000), + ], +) +def test_plan_threshold_rejection_reports_billing_limit( + plan: CloudPlan, + estimated_mb: int, + plan_limit_mb: int, +) -> None: + with pytest.raises(VectorSpaceAdmissionError) as exc_info: + _check_estimate(plan, estimated_mb, plan_limit_mb=plan_limit_mb) + + assert get_vector_space_admission_error_fields(str(exc_info.value)) == { + "error_code": VECTOR_SPACE_ADMISSION_ERROR_CODE, + "estimated_vector_space_mb": estimated_mb, + "vector_space_limit_mb": plan_limit_mb, + } + + +def test_2060_mb_estimate_rejects_sandbox_but_allows_pro() -> None: + with pytest.raises(VectorSpaceAdmissionError): + _check_estimate(CloudPlan.SANDBOX, 2060) + + _check_estimate(CloudPlan.PROFESSIONAL, 2060) + + +def test_billing_plan_lookup_excludes_vector_space_and_is_cached() -> None: + service = VectorSpaceAdmissionService() + with patch( + "services.vector_space_admission_service.BillingService.get_info", + return_value={"enabled": True, "subscription": {"plan": "professional"}}, + ) as get_info: + assert service._get_plan("tenant-1") == CloudPlan.PROFESSIONAL + assert service._get_plan("tenant-1") == CloudPlan.PROFESSIONAL + + get_info.assert_called_once_with("tenant-1", exclude_vector_space=True) diff --git a/api/tests/unit_tests/services/test_webhook_service.py b/api/tests/unit_tests/services/test_webhook_service.py index 338771c5a28..c4505f0b77a 100644 --- a/api/tests/unit_tests/services/test_webhook_service.py +++ b/api/tests/unit_tests/services/test_webhook_service.py @@ -15,7 +15,7 @@ from services.trigger.webhook_service import WebhookService class TestWebhookServiceUnit: - """Webhook business-logic tests with real sessions where the service owns their lifecycle.""" + """Webhook business-logic tests with isolated sessions where the service owns their lifecycle.""" def test_trigger_workflow_execution_propagates_quota_error_without_error_log( self, @@ -65,6 +65,66 @@ class TestWebhookServiceUnit: "Tenant tenant-123 quota exceeded for feature workflow, skipping webhook trigger webhook-123" ) + def test_trigger_workflow_execution_success( + self, + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + ) -> None: + monkeypatch.setattr(webhook_service_module, "db", SimpleNamespace(engine=sqlite_engine)) + webhook_trigger = MagicMock( + webhook_id="webhook-123", + tenant_id="tenant-123", + app_id="app-123", + node_id="node-123", + ) + workflow = MagicMock(id="workflow-123") + end_user = MagicMock(id="end-user-123") + quota_charge = MagicMock() + webhook_data = { + "method": "POST", + "headers": {"Authorization": "Bearer token"}, + "query_params": {"version": "1"}, + "body": {"message": "hello"}, + "files": {}, + } + + with ( + patch.object( + webhook_service_module.EndUserService, + "get_or_create_end_user_by_type", + return_value=end_user, + ), + patch.object(webhook_service_module.QuotaService, "reserve", return_value=quota_charge), + patch.object( + webhook_service_module.AsyncWorkflowService, + "trigger_workflow_async", + ) as mock_trigger, + ): + WebhookService.trigger_workflow_execution(webhook_trigger, webhook_data, workflow) + + call_session = mock_trigger.call_args.kwargs["session"] + assert call_session.get_bind() is sqlite_engine + quota_charge.commit.assert_called_once_with() + quota_charge.refund.assert_not_called() + + def test_trigger_workflow_execution_end_user_service_failure(self) -> None: + webhook_trigger = MagicMock( + webhook_id="webhook-123", + tenant_id="tenant-123", + app_id="app-123", + node_id="node-123", + ) + workflow = MagicMock(id="workflow-123") + webhook_data = {"method": "POST", "headers": {}, "query_params": {}, "body": {}, "files": {}} + + with patch.object( + webhook_service_module.EndUserService, + "get_or_create_end_user_by_type", + side_effect=ValueError("Failed to create end user"), + ): + with pytest.raises(ValueError, match="Failed to create end user"): + WebhookService.trigger_workflow_execution(webhook_trigger, webhook_data, workflow) + def test_extract_webhook_data_json(self): """Test webhook data extraction from JSON request.""" app = Flask(__name__) @@ -590,6 +650,104 @@ class TestWebhookServiceUnit: with pytest.raises(ValueError, match="HTTP method mismatch"): WebhookService.extract_and_validate_webhook_data(webhook_trigger, node_config) + def test_extract_and_validate_webhook_request_missing_required_header(self) -> None: + app = Flask(__name__) + with app.test_request_context( + "/webhook", + method="POST", + headers={"Content-Type": "application/json"}, + ): + node_config = { + "data": { + "method": "post", + "content_type": "application/json", + "headers": [{"name": "Authorization", "required": True}], + } + } + + with pytest.raises(ValueError, match="Required header missing: Authorization"): + WebhookService.extract_and_validate_webhook_data(MagicMock(), node_config) + + def test_extract_and_validate_webhook_request_case_insensitive_headers(self) -> None: + app = Flask(__name__) + with app.test_request_context( + "/webhook", + method="POST", + headers={"Content-Type": "application/json", "authorization": "Bearer token"}, + json={"message": "hello"}, + ): + node_config = { + "data": { + "method": "post", + "content_type": "application/json", + "headers": [{"name": "Authorization", "required": True}], + "body": [{"name": "message", "type": "string", "required": True}], + } + } + + result = WebhookService.extract_and_validate_webhook_data(MagicMock(), node_config) + + assert result["headers"].get("Authorization") == "Bearer token" + + def test_extract_and_validate_webhook_request_missing_required_param(self) -> None: + app = Flask(__name__) + with app.test_request_context( + "/webhook", + method="POST", + headers={"Content-Type": "application/json"}, + json={"message": "hello"}, + ): + node_config = { + "data": { + "method": "post", + "content_type": "application/json", + "params": [{"name": "version", "required": True}], + "body": [{"name": "message", "type": "string", "required": True}], + } + } + + with pytest.raises(ValueError, match="Required parameter missing: version"): + WebhookService.extract_and_validate_webhook_data(MagicMock(), node_config) + + def test_extract_and_validate_webhook_request_missing_required_body_param(self) -> None: + app = Flask(__name__) + with app.test_request_context( + "/webhook", + method="POST", + headers={"Content-Type": "application/json"}, + json={}, + ): + node_config = { + "data": { + "method": "post", + "content_type": "application/json", + "body": [{"name": "message", "type": "string", "required": True}], + } + } + + with pytest.raises(ValueError, match="Required body parameter missing: message"): + WebhookService.extract_and_validate_webhook_data(MagicMock(), node_config) + + def test_extract_and_validate_webhook_request_missing_required_file(self) -> None: + app = Flask(__name__) + with app.test_request_context( + "/webhook", + method="POST", + data={"note": "test"}, + content_type="multipart/form-data", + ): + node_config = { + "data": { + "method": "post", + "content_type": "multipart/form-data", + "body": [{"name": "file", "type": "file", "required": True}], + } + } + + result = WebhookService.extract_and_validate_webhook_data(MagicMock(), node_config) + + assert result["files"] == {} + def test_debug_mode_parameter_handling(self): """Test that the debug mode parameter is properly handled in _prepare_webhook_execution.""" from controllers.trigger.webhook import _prepare_webhook_execution diff --git a/api/tests/unit_tests/services/test_website_service.py b/api/tests/unit_tests/services/test_website_service.py index 2024aec13a3..571e8e99903 100644 --- a/api/tests/unit_tests/services/test_website_service.py +++ b/api/tests/unit_tests/services/test_website_service.py @@ -724,3 +724,23 @@ def test_scrape_with_watercrawl_calls_provider(monkeypatch: pytest.MonkeyPatch) ) assert result == {"markdown": "m"} provider_instance.scrape_url.assert_called_once_with("u") + + +def test_pooled_clients_carry_bounded_timeouts() -> None: + """Regression for #39859: the Jina and adaptive-crawl pooled clients + must carry a read/connect timeout so a stalled endpoint fails fast + instead of pinning a worker. Same shape as the WaterCrawl hardening + that landed in PR #37512. + """ + jina = website_service_module._jina_http_client + adaptive = website_service_module._adaptive_http_client + + # Read and connect bounds are set. + assert jina.timeout is not None + assert adaptive.timeout is not None + + # The values match the documented floor (read 30.0 / connect 5.0). + assert jina.timeout.read == 30.0 + assert jina.timeout.connect == 5.0 + assert adaptive.timeout.read == 30.0 + assert adaptive.timeout.connect == 5.0 diff --git a/api/tests/unit_tests/services/test_workflow_app_service_metadata.py b/api/tests/unit_tests/services/test_workflow_app_service_metadata.py new file mode 100644 index 00000000000..ded4fab1c7a --- /dev/null +++ b/api/tests/unit_tests/services/test_workflow_app_service_metadata.py @@ -0,0 +1,88 @@ +"""Unit tests for workflow app log views and trigger metadata helpers.""" + +import json +import uuid +from unittest.mock import patch + +import pytest + +from models.enums import AppTriggerType, CreatorUserRole +from models.workflow import WorkflowAppLog, WorkflowAppLogCreatedFrom +from services.workflow_app_service import LogView, WorkflowAppService + + +class TestLogView: + def test_details_and_proxy_attributes(self) -> None: + log = WorkflowAppLog( + tenant_id="tenant-1", + app_id="app-1", + workflow_id="workflow-1", + workflow_run_id="run-1", + created_from=WorkflowAppLogCreatedFrom.WEB_APP, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + ) + log.id = "log-1" + + view = LogView(log=log, details={"trigger_metadata": {"type": "plugin"}}) + + assert view.details == {"trigger_metadata": {"type": "plugin"}} + assert view.id == "log-1" + + +class TestHandleTriggerMetadata: + def test_returns_empty_dict_when_metadata_missing(self) -> None: + assert WorkflowAppService().handle_trigger_metadata("tenant-1", None) == {} + + def test_enriches_plugin_icons(self) -> None: + metadata = { + "type": AppTriggerType.TRIGGER_PLUGIN.value, + "icon_filename": "light.png", + "icon_dark_filename": "dark.png", + } + with patch( + "services.workflow_app_service.PluginService.get_plugin_icon_url", + side_effect=["https://cdn/light.png", "https://cdn/dark.png"], + ) as mock_icon: + result = WorkflowAppService().handle_trigger_metadata("tenant-1", json.dumps(metadata)) + + assert result["icon"] == "https://cdn/light.png" + assert result["icon_dark"] == "https://cdn/dark.png" + assert mock_icon.call_count == 2 + + def test_non_plugin_metadata_without_icon_lookup(self) -> None: + metadata = {"type": AppTriggerType.TRIGGER_WEBHOOK.value} + with patch("services.workflow_app_service.PluginService.get_plugin_icon_url") as mock_icon: + result = WorkflowAppService().handle_trigger_metadata("tenant-1", json.dumps(metadata)) + + assert result["type"] == AppTriggerType.TRIGGER_WEBHOOK.value + mock_icon.assert_not_called() + + +class TestSafeJsonLoads: + @pytest.mark.parametrize( + ("value", "expected"), + [ + (None, None), + ("", None), + ('{"k":"v"}', {"k": "v"}), + ("not-json", None), + ({"raw": True}, {"raw": True}), + ], + ) + def test_handles_various_inputs(self, value, expected) -> None: + assert WorkflowAppService._safe_json_loads(value) == expected + + +class TestSafeParseUuid: + def test_returns_none_for_short_or_invalid_values(self) -> None: + assert WorkflowAppService._safe_parse_uuid("short") is None + assert WorkflowAppService._safe_parse_uuid("x" * 40) is None + + def test_returns_uuid_for_valid_string(self) -> None: + raw = str(uuid.uuid4()) + + result = WorkflowAppService._safe_parse_uuid(raw) + + assert result is not None + assert str(result) == raw diff --git a/api/tests/unit_tests/services/test_workflow_comment_service.py b/api/tests/unit_tests/services/test_workflow_comment_service.py index e6db068e07c..0478351b073 100644 --- a/api/tests/unit_tests/services/test_workflow_comment_service.py +++ b/api/tests/unit_tests/services/test_workflow_comment_service.py @@ -1,35 +1,129 @@ -from unittest.mock import MagicMock, Mock, patch +"""Persistence-focused tests for :mod:`services.workflow_comment_service`. + +The service opens its own sessions for most operations, so these tests bind it to the +same disposable SQLite engine used for fixture setup and assert committed database +state. External task dispatch and the clock remain mocked at their I/O boundaries. +""" + +from datetime import datetime +from types import SimpleNamespace +from unittest.mock import Mock, patch import pytest +from sqlalchemy import event, func, select +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound +from models import App, TenantAccountJoin, WorkflowComment, WorkflowCommentMention, WorkflowCommentReply +from models.account import Account, TenantAccountRole +from models.model import AppMode from services import workflow_comment_service as service_module from services.workflow_comment_service import WorkflowCommentService +TENANT_ID = "11111111-1111-1111-1111-111111111111" +OTHER_TENANT_ID = "11111111-1111-1111-1111-111111111112" +APP_ID = "22222222-2222-2222-2222-222222222222" +OTHER_APP_ID = "22222222-2222-2222-2222-222222222223" +OWNER_ID = "33333333-3333-3333-3333-333333333333" +USER_2_ID = "33333333-3333-3333-3333-333333333334" +USER_3_ID = "33333333-3333-3333-3333-333333333335" +USER_4_ID = "33333333-3333-3333-3333-333333333336" +OUTSIDER_ID = "33333333-3333-3333-3333-333333333337" + @pytest.fixture -def mock_session(monkeypatch: pytest.MonkeyPatch) -> Mock: - session = Mock() - context_manager = MagicMock() - context_manager.__enter__.return_value = session - context_manager.__exit__.return_value = False - mock_db = MagicMock() - mock_db.engine = Mock() - empty_scalars = Mock() - empty_scalars.all.return_value = [] - session.scalars.return_value = empty_scalars - monkeypatch.setattr(service_module, "Session", Mock(return_value=context_manager)) - monkeypatch.setattr(service_module, "db", mock_db) - monkeypatch.setattr(service_module.send_workflow_comment_mention_email_task, "delay", Mock()) - return session +def delay_mock(monkeypatch: pytest.MonkeyPatch) -> Mock: + mock = Mock() + monkeypatch.setattr(service_module.send_workflow_comment_mention_email_task, "delay", mock) + return mock -def _mock_scalars(result_list: list[object]) -> Mock: - scalars = Mock() - scalars.all.return_value = result_list - return scalars +@pytest.fixture(autouse=True) +def bind_service_database( + sqlite_session: Session, + monkeypatch: pytest.MonkeyPatch, + delay_mock: Mock, +) -> None: + """Bind service-owned sessions to the engine prepared by the shared SQLite fixture.""" + monkeypatch.setattr(service_module, "db", SimpleNamespace(engine=sqlite_session.get_bind())) +def _account( + account_id: str, + *, + name: str = "Test User", + email: str = "user@example.com", + interface_language: str | None = "en-US", +) -> Account: + account = Account(name=name, email=email, interface_language=interface_language) + account.id = account_id + return account + + +def _app(*, app_id: str = APP_ID, tenant_id: str = TENANT_ID, name: str = "My App") -> App: + app = App( + tenant_id=tenant_id, + name=name, + mode=AppMode.WORKFLOW, + enable_site=False, + enable_api=False, + created_by=OWNER_ID, + ) + app.id = app_id + return app + + +def _comment( + *, + tenant_id: str = TENANT_ID, + app_id: str = APP_ID, + created_by: str = OWNER_ID, + content: str = "hello", + resolved: bool = False, + resolved_at: datetime | None = None, + resolved_by: str | None = None, +) -> WorkflowComment: + return WorkflowComment( + tenant_id=tenant_id, + app_id=app_id, + position_x=1.0, + position_y=2.0, + content=content, + created_by=created_by, + resolved=resolved, + resolved_at=resolved_at, + resolved_by=resolved_by, + ) + + +def _membership(account_id: str, *, tenant_id: str = TENANT_ID) -> TenantAccountJoin: + return TenantAccountJoin( + tenant_id=tenant_id, + account_id=account_id, + role=TenantAccountRole.NORMAL, + ) + + +def _persist(session: Session, *objects: object) -> None: + session.add_all(objects) + session.commit() + + +@pytest.mark.usefixtures("sqlite_session") +@pytest.mark.parametrize( + "sqlite_session", + [ + ( + Account, + TenantAccountJoin, + App, + WorkflowComment, + WorkflowCommentReply, + WorkflowCommentMention, + ) + ], + indirect=True, +) class TestWorkflowCommentService: def test_validate_content_rejects_empty(self) -> None: with pytest.raises(ValueError): @@ -39,42 +133,43 @@ class TestWorkflowCommentService: with pytest.raises(ValueError): WorkflowCommentService._validate_content("a" * 1001) - def test_filter_valid_mentioned_user_ids_filters_by_tenant_and_preserves_order(self, mock_session: Mock) -> None: - tenant_member_1 = "123e4567-e89b-12d3-a456-426614174000" - tenant_member_2 = "123e4567-e89b-12d3-a456-426614174002" - non_tenant_member = "123e4567-e89b-12d3-a456-426614174001" - mock_session.scalars.return_value = _mock_scalars([tenant_member_1, tenant_member_2]) + def test_filter_valid_mentioned_user_ids_filters_by_tenant_and_preserves_order( + self, sqlite_session: Session + ) -> None: + _persist( + sqlite_session, + _membership(OWNER_ID), + _membership(USER_2_ID), + _membership(USER_3_ID, tenant_id=OTHER_TENANT_ID), + ) result = WorkflowCommentService._filter_valid_mentioned_user_ids( [ - tenant_member_1, + OWNER_ID, "", 123, # type: ignore[list-item] - tenant_member_1, - non_tenant_member, - tenant_member_2, + OWNER_ID, + USER_3_ID, + USER_2_ID, ], - session=mock_session, - tenant_id="tenant-1", + session=sqlite_session, + tenant_id=TENANT_ID, ) - assert result == [ - tenant_member_1, - tenant_member_2, - ] + assert result == [OWNER_ID, USER_2_ID] def test_format_comment_excerpt_handles_short_and_long_limits(self) -> None: assert WorkflowCommentService._format_comment_excerpt(" hello ", max_length=10) == "hello" assert WorkflowCommentService._format_comment_excerpt("abcdefghijk", max_length=3) == "abc" assert WorkflowCommentService._format_comment_excerpt(" abcdefghijk ", max_length=8) == "abcde..." - def test_build_mention_email_payloads_returns_empty_for_no_candidates(self, mock_session: Mock) -> None: + def test_build_mention_email_payloads_returns_empty_for_no_candidates(self, sqlite_session: Session) -> None: assert ( WorkflowCommentService._build_mention_email_payloads( - session=mock_session, - tenant_id="tenant-1", - app_id="app-1", - mentioner_id="user-1", + session=sqlite_session, + tenant_id=TENANT_ID, + app_id=APP_ID, + mentioner_id=OWNER_ID, mentioned_user_ids=[], content="hello", ) @@ -82,11 +177,11 @@ class TestWorkflowCommentService: ) assert ( WorkflowCommentService._build_mention_email_payloads( - session=mock_session, - tenant_id="tenant-1", - app_id="app-1", - mentioner_id="user-1", - mentioned_user_ids=["user-1"], + session=sqlite_session, + tenant_id=TENANT_ID, + app_id=APP_ID, + mentioner_id=OWNER_ID, + mentioned_user_ids=[OWNER_ID], content="hello", ) == [] @@ -104,29 +199,26 @@ class TestWorkflowCommentService: assert delay_mock.call_count == 2 - def test_build_mention_email_payloads_skips_accounts_without_email(self, mock_session: Mock) -> None: - account_without_email = Mock() - account_without_email.email = None - account_without_email.name = "No Email" - account_without_email.interface_language = "en-US" - - account_with_email = Mock() - account_with_email.email = "user@example.com" - account_with_email.name = "" - account_with_email.interface_language = None - - mock_session.scalar.side_effect = ["My App", "Commenter"] - mock_session.scalars.return_value = _mock_scalars([account_without_email, account_with_email]) + def test_build_mention_email_payloads_skips_accounts_without_email(self, sqlite_session: Session) -> None: + _persist( + sqlite_session, + _app(), + _account(OWNER_ID, name="Commenter", email="commenter@example.com"), + _account(USER_2_ID, name="No Email", email=""), + _account(USER_3_ID, name="", email="user@example.com", interface_language=None), + _membership(USER_2_ID), + _membership(USER_3_ID), + ) payloads = WorkflowCommentService._build_mention_email_payloads( - session=mock_session, - tenant_id="tenant-1", - app_id="app-1", - mentioner_id="user-1", - mentioned_user_ids=["user-2"], + session=sqlite_session, + tenant_id=TENANT_ID, + app_id=APP_ID, + mentioner_id=OWNER_ID, + mentioned_user_ids=[USER_2_ID, USER_3_ID], content="hello", ) - expected_app_url = f"{service_module.dify_config.CONSOLE_WEB_URL.rstrip('/')}/app/app-1/workflow" + expected_app_url = f"{service_module.dify_config.CONSOLE_WEB_URL.rstrip('/')}/app/{APP_ID}/workflow" assert payloads == [ { @@ -140,439 +232,463 @@ class TestWorkflowCommentService: } ] - def test_create_comment_creates_mentions(self, mock_session: Mock) -> None: - comment = Mock() - comment.id = "comment-1" - comment.created_at = "ts" + def test_create_comment_creates_mentions(self, sqlite_session: Session) -> None: + _persist(sqlite_session, _membership(USER_2_ID)) - with ( - patch.object(service_module, "WorkflowComment", return_value=comment), - patch.object(service_module, "WorkflowCommentMention", return_value=Mock()), - patch.object(WorkflowCommentService, "_filter_valid_mentioned_user_ids", return_value=["user-2"]), - ): - result = WorkflowCommentService.create_comment( - tenant_id="tenant-1", - app_id="app-1", - created_by="user-1", - content="hello", - position_x=1.0, - position_y=2.0, - mentioned_user_ids=["user-2", "bad-id"], - ) + result = WorkflowCommentService.create_comment( + tenant_id=TENANT_ID, + app_id=APP_ID, + created_by=OWNER_ID, + content="hello", + position_x=1.0, + position_y=2.0, + mentioned_user_ids=[USER_2_ID, OUTSIDER_ID], + ) - assert result == {"id": "comment-1", "created_at": "ts"} - assert mock_session.add.call_args_list[0].args[0] is comment - assert mock_session.add.call_count == 2 - mock_session.commit.assert_called_once() - - def test_update_comment_raises_not_found(self, mock_session: Mock) -> None: - mock_session.scalar.return_value = None + comment = sqlite_session.get(WorkflowComment, result["id"]) + assert comment is not None + assert comment.content == "hello" + assert comment.created_at == result["created_at"] + mentions = sqlite_session.scalars( + select(WorkflowCommentMention).where(WorkflowCommentMention.comment_id == comment.id) + ).all() + assert [mention.mentioned_user_id for mention in mentions] == [USER_2_ID] + def test_update_comment_raises_not_found(self, sqlite_session: Session) -> None: with pytest.raises(NotFound): WorkflowCommentService.update_comment( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - user_id="user-1", + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id="missing-comment", + user_id=OWNER_ID, content="hello", ) - def test_update_comment_raises_forbidden(self, mock_session: Mock) -> None: - comment = Mock() - comment.created_by = "owner" - mock_session.scalar.return_value = comment + def test_update_comment_raises_forbidden(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) with pytest.raises(Forbidden): WorkflowCommentService.update_comment( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - user_id="intruder", + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + user_id=OUTSIDER_ID, content="hello", ) - def test_update_comment_replaces_mentions(self, mock_session: Mock) -> None: - comment = Mock() - comment.id = "comment-1" - comment.created_by = "owner" - mock_session.scalar.return_value = comment + def test_update_comment_replaces_mentions(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + _persist( + sqlite_session, + WorkflowCommentMention(comment_id=comment.id, mentioned_user_id=USER_3_ID), + WorkflowCommentMention(comment_id=comment.id, mentioned_user_id=USER_4_ID), + _membership(USER_2_ID), + ) - existing_mentions = [Mock(), Mock()] - mock_session.scalars.return_value = _mock_scalars(existing_mentions) + result = WorkflowCommentService.update_comment( + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + user_id=OWNER_ID, + content="updated", + mentioned_user_ids=[USER_2_ID, OUTSIDER_ID], + ) - with patch.object(WorkflowCommentService, "_filter_valid_mentioned_user_ids", return_value=["user-2"]): - result = WorkflowCommentService.update_comment( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - user_id="owner", - content="updated", - mentioned_user_ids=["user-2", "bad-id"], - ) + sqlite_session.expire_all() + persisted_comment = sqlite_session.get(WorkflowComment, comment.id) + assert persisted_comment is not None + assert persisted_comment.content == "updated" + assert result == {"id": comment.id, "updated_at": persisted_comment.updated_at} + mentions = sqlite_session.scalars( + select(WorkflowCommentMention).where(WorkflowCommentMention.comment_id == comment.id) + ).all() + assert [mention.mentioned_user_id for mention in mentions] == [USER_2_ID] - assert result == {"id": "comment-1", "updated_at": comment.updated_at} - assert mock_session.delete.call_count == 2 - assert mock_session.add.call_count == 1 - mock_session.commit.assert_called_once() - - def test_update_comment_preserves_mentions_when_mentioned_user_ids_omitted(self, mock_session: Mock) -> None: - comment = Mock() - comment.id = "comment-1" - comment.created_by = "owner" - mock_session.scalar.return_value = comment - - with ( - patch.object(WorkflowCommentService, "_filter_valid_mentioned_user_ids") as filter_mentions_mock, - patch.object(WorkflowCommentService, "_build_mention_email_payloads") as build_payloads_mock, - patch.object(WorkflowCommentService, "_dispatch_mention_emails") as dispatch_mock, - ): - result = WorkflowCommentService.update_comment( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - user_id="owner", - content="updated", - ) - - assert result == {"id": "comment-1", "updated_at": comment.updated_at} - mock_session.delete.assert_not_called() - mock_session.add.assert_not_called() - filter_mentions_mock.assert_not_called() - build_payloads_mock.assert_not_called() - dispatch_mock.assert_called_once_with([]) - mock_session.commit.assert_called_once() - - def test_update_comment_clears_mentions_when_empty_list_provided(self, mock_session: Mock) -> None: - comment = Mock() - comment.id = "comment-1" - comment.created_by = "owner" - mock_session.scalar.return_value = comment - - existing_mentions = [Mock(), Mock()] - mock_session.scalars.return_value = _mock_scalars(existing_mentions) - - with patch.object(WorkflowCommentService, "_filter_valid_mentioned_user_ids", return_value=[]): - result = WorkflowCommentService.update_comment( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - user_id="owner", - content="updated", - mentioned_user_ids=[], - ) - - assert result == {"id": "comment-1", "updated_at": comment.updated_at} - assert mock_session.delete.call_count == 2 - mock_session.add.assert_not_called() - mock_session.commit.assert_called_once() - - def test_update_comment_notifies_only_new_mentions(self, mock_session: Mock) -> None: - comment = Mock() - comment.id = "comment-1" - comment.created_by = "owner" - mock_session.scalar.return_value = comment - - existing_mention = Mock() - existing_mention.mentioned_user_id = "user-2" - mock_session.scalars.return_value = _mock_scalars([existing_mention]) - - with ( - patch.object( - WorkflowCommentService, - "_filter_valid_mentioned_user_ids", - return_value=["user-2", "user-3"], - ), - patch.object( - WorkflowCommentService, - "_build_mention_email_payloads", - return_value=[], - ) as build_payloads_mock, - patch.object(WorkflowCommentService, "_dispatch_mention_emails") as dispatch_mock, - ): - WorkflowCommentService.update_comment( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - user_id="owner", - content="updated", - mentioned_user_ids=["user-2", "user-3"], - ) - - assert build_payloads_mock.call_args.kwargs["mentioned_user_ids"] == ["user-3"] - dispatch_mock.assert_called_once_with([]) - - def test_get_comments_preloads_related_accounts(self, mock_session: Mock) -> None: - comment = Mock() - comment.created_by = "user-1" - comment.resolved_by = "user-2" - reply = Mock() - reply.created_by = "user-3" - mention = Mock() - mention.mentioned_user_id = "user-4" - comment.replies = [reply] - comment.mentions = [mention] - comment.cache_created_by_account = Mock() - comment.cache_resolved_by_account = Mock() - reply.cache_created_by_account = Mock() - mention.cache_mentioned_user_account = Mock() - - account_1 = Mock() - account_1.id = "user-1" - account_2 = Mock() - account_2.id = "user-2" - account_3 = Mock() - account_3.id = "user-3" - account_4 = Mock() - account_4.id = "user-4" - - mock_session.scalars.side_effect = [ - _mock_scalars([comment]), - _mock_scalars([account_1, account_2, account_3, account_4]), - ] - - result = WorkflowCommentService.get_comments("tenant-1", "app-1") - - assert result == [comment] - comment.cache_created_by_account.assert_called_once_with(account_1) - comment.cache_resolved_by_account.assert_called_once_with(account_2) - reply.cache_created_by_account.assert_called_once_with(account_3) - mention.cache_mentioned_user_account.assert_called_once_with(account_4) - - def test_preload_accounts_returns_early_for_empty_comments(self, mock_session: Mock) -> None: - WorkflowCommentService._preload_accounts(mock_session, []) - - mock_session.scalars.assert_not_called() - - def test_get_comment_raises_not_found_with_provided_session(self) -> None: - session = Mock() - session.scalar.return_value = None - - with pytest.raises(NotFound): - WorkflowCommentService.get_comment("tenant-1", "app-1", "comment-1", session=session) - - def test_get_comment_uses_context_manager_when_session_not_provided(self, mock_session: Mock) -> None: - comment = Mock() - comment.created_by = "user-1" - comment.resolved_by = None - comment.replies = [] - comment.mentions = [] - comment.cache_created_by_account = Mock() - comment.cache_resolved_by_account = Mock() - mock_session.scalar.return_value = comment - mock_session.scalars.return_value = _mock_scalars([]) - - result = WorkflowCommentService.get_comment("tenant-1", "app-1", "comment-1") - - assert result is comment - comment.cache_created_by_account.assert_called_once() - comment.cache_resolved_by_account.assert_called_once_with(None) - - def test_delete_comment_raises_forbidden(self, mock_session: Mock) -> None: - comment = Mock() - comment.created_by = "owner" - - with patch.object(WorkflowCommentService, "get_comment", return_value=comment): - with pytest.raises(Forbidden): - WorkflowCommentService.delete_comment("tenant-1", "app-1", "comment-1", "intruder") - - def test_delete_comment_removes_related_entities(self, mock_session: Mock) -> None: - comment = Mock() - comment.created_by = "owner" - - mentions = [Mock(), Mock()] - replies = [Mock()] - mock_session.scalars.side_effect = [_mock_scalars(mentions), _mock_scalars(replies)] - - with patch.object(WorkflowCommentService, "get_comment", return_value=comment): - WorkflowCommentService.delete_comment("tenant-1", "app-1", "comment-1", "owner") - - assert mock_session.delete.call_count == 4 - mock_session.commit.assert_called_once() - - def test_resolve_comment_sets_fields(self, mock_session: Mock) -> None: - comment = Mock() - comment.resolved = False - comment.resolved_at = None - comment.resolved_by = None - - with ( - patch.object(WorkflowCommentService, "get_comment", return_value=comment), - patch.object(service_module, "naive_utc_now", return_value="now"), - ): - result = WorkflowCommentService.resolve_comment("tenant-1", "app-1", "comment-1", "user-1") - - assert result is comment - assert comment.resolved is True - assert comment.resolved_at == "now" - assert comment.resolved_by == "user-1" - mock_session.commit.assert_called_once() - - def test_resolve_comment_noop_when_already_resolved(self, mock_session: Mock) -> None: - comment = Mock() - comment.resolved = True - - with patch.object(WorkflowCommentService, "get_comment", return_value=comment): - result = WorkflowCommentService.resolve_comment("tenant-1", "app-1", "comment-1", "user-1") - - assert result is comment - mock_session.commit.assert_not_called() - - def test_create_reply_requires_comment(self, mock_session: Mock) -> None: - mock_session.get.return_value = None - - with pytest.raises(NotFound): - WorkflowCommentService.create_reply("comment-1", "hello", "user-1") - - def test_create_reply_creates_mentions(self, mock_session: Mock) -> None: - mock_session.get.return_value = Mock() - reply = Mock() - reply.id = "reply-1" - reply.created_at = "ts" - - with ( - patch.object(service_module, "WorkflowCommentReply", return_value=reply), - patch.object(service_module, "WorkflowCommentMention", return_value=Mock()), - patch.object(WorkflowCommentService, "_filter_valid_mentioned_user_ids", return_value=["user-2"]), - ): - result = WorkflowCommentService.create_reply( - comment_id="comment-1", - content="hello", - created_by="user-1", - mentioned_user_ids=["user-2", "bad-id"], - ) - - assert result == {"id": "reply-1", "created_at": "ts"} - assert mock_session.add.call_count == 2 - mock_session.commit.assert_called_once() - - def test_update_reply_raises_not_found(self, mock_session: Mock) -> None: - mock_session.scalar.return_value = None - - with pytest.raises(NotFound): - WorkflowCommentService.update_reply( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - reply_id="reply-1", - user_id="user-1", - content="hello", - ) - - def test_update_reply_raises_forbidden(self, mock_session: Mock) -> None: - reply = Mock() - reply.created_by = "owner" - mock_session.scalar.return_value = reply - - with pytest.raises(Forbidden): - WorkflowCommentService.update_reply( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - reply_id="reply-1", - user_id="intruder", - content="hello", - ) - - def test_update_reply_replaces_mentions(self, mock_session: Mock) -> None: - reply = Mock() - reply.id = "reply-1" - reply.comment_id = "comment-1" - reply.created_by = "owner" - reply.updated_at = "updated" - mock_session.scalar.return_value = reply - mock_session.scalars.return_value = _mock_scalars([Mock()]) - comment = Mock() - comment.tenant_id = "tenant-1" - comment.app_id = "app-1" - mock_session.get.return_value = comment - - with patch.object(WorkflowCommentService, "_filter_valid_mentioned_user_ids", return_value=["user-2"]): - result = WorkflowCommentService.update_reply( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - reply_id="reply-1", - user_id="owner", - content="new", - mentioned_user_ids=["user-2", "bad-id"], - ) - - assert result == {"id": "reply-1", "updated_at": "updated"} - assert mock_session.delete.call_count == 1 - assert mock_session.add.call_count == 1 - mock_session.commit.assert_called_once() - mock_session.refresh.assert_called_once_with(reply) - - def test_update_comment_updates_position_coordinates_when_provided(self, mock_session: Mock) -> None: - comment = Mock() - comment.id = "comment-1" - comment.created_by = "owner" - comment.position_x = 1.0 - comment.position_y = 2.0 - mock_session.scalar.return_value = comment - mock_session.scalars.return_value = _mock_scalars([]) + def test_update_comment_preserves_mentions_when_mentioned_user_ids_omitted( + self, sqlite_session: Session, delay_mock: Mock + ) -> None: + comment = _comment() + _persist(sqlite_session, comment) + mention = WorkflowCommentMention(comment_id=comment.id, mentioned_user_id=USER_2_ID) + _persist(sqlite_session, mention) WorkflowCommentService.update_comment( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - user_id="owner", + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + user_id=OWNER_ID, + content="updated", + ) + + sqlite_session.expire_all() + persisted_comment = sqlite_session.get(WorkflowComment, comment.id) + assert persisted_comment is not None + assert persisted_comment.content == "updated" + assert sqlite_session.get(WorkflowCommentMention, mention.id) is not None + delay_mock.assert_not_called() + + def test_update_comment_clears_mentions_when_empty_list_provided(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + _persist( + sqlite_session, + WorkflowCommentMention(comment_id=comment.id, mentioned_user_id=USER_2_ID), + WorkflowCommentMention(comment_id=comment.id, mentioned_user_id=USER_3_ID), + ) + + WorkflowCommentService.update_comment( + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + user_id=OWNER_ID, + content="updated", + mentioned_user_ids=[], + ) + + mention_count = sqlite_session.scalar( + select(func.count()) + .select_from(WorkflowCommentMention) + .where(WorkflowCommentMention.comment_id == comment.id) + ) + assert mention_count == 0 + + def test_update_comment_notifies_only_new_mentions(self, sqlite_session: Session, delay_mock: Mock) -> None: + comment = _comment() + _persist( + sqlite_session, + _app(), + _account(OWNER_ID, name="Owner", email="owner@example.com"), + _account(USER_2_ID, name="Existing", email="existing@example.com"), + _account(USER_3_ID, name="New User", email="new@example.com"), + _membership(USER_2_ID), + _membership(USER_3_ID), + comment, + ) + _persist(sqlite_session, WorkflowCommentMention(comment_id=comment.id, mentioned_user_id=USER_2_ID)) + + WorkflowCommentService.update_comment( + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + user_id=OWNER_ID, + content="updated", + mentioned_user_ids=[USER_2_ID, USER_3_ID], + ) + + delay_mock.assert_called_once() + assert delay_mock.call_args.kwargs["to"] == "new@example.com" + mentions = sqlite_session.scalars( + select(WorkflowCommentMention).where(WorkflowCommentMention.comment_id == comment.id) + ).all() + assert {mention.mentioned_user_id for mention in mentions} == {USER_2_ID, USER_3_ID} + + def test_get_comments_preloads_related_accounts(self, sqlite_session: Session) -> None: + comment = _comment(resolved=True, resolved_by=USER_2_ID) + _persist( + sqlite_session, + _account(OWNER_ID, name="Owner"), + _account(USER_2_ID, name="Resolver"), + _account(USER_3_ID, name="Replier"), + _account(USER_4_ID, name="Mentioned"), + comment, + ) + reply = WorkflowCommentReply(comment_id=comment.id, content="reply", created_by=USER_3_ID) + mention = WorkflowCommentMention(comment_id=comment.id, mentioned_user_id=USER_4_ID) + _persist(sqlite_session, reply, mention) + + result = WorkflowCommentService.get_comments(TENANT_ID, APP_ID) + + assert len(result) == 1 + loaded_comment = result[0] + assert loaded_comment.id == comment.id + assert loaded_comment.created_by_account.id == OWNER_ID + assert loaded_comment.resolved_by_account.id == USER_2_ID + assert loaded_comment.replies[0].created_by_account.id == USER_3_ID + assert loaded_comment.mentions[0].mentioned_user_account.id == USER_4_ID + + def test_preload_accounts_returns_early_for_empty_comments(self, sqlite_session: Session) -> None: + statements: list[str] = [] + bind = sqlite_session.get_bind() + + def record_sql(*args: object) -> None: + statements.append(str(args[2])) + + event.listen(bind, "before_cursor_execute", record_sql) + try: + WorkflowCommentService._preload_accounts(sqlite_session, []) + finally: + event.remove(bind, "before_cursor_execute", record_sql) + + assert statements == [] + + def test_get_comment_raises_not_found_with_provided_session(self, sqlite_session: Session) -> None: + with pytest.raises(NotFound): + WorkflowCommentService.get_comment(TENANT_ID, APP_ID, "missing-comment", session=sqlite_session) + + def test_get_comment_uses_context_manager_when_session_not_provided(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, _account(OWNER_ID, name="Owner"), comment) + + result = WorkflowCommentService.get_comment(TENANT_ID, APP_ID, comment.id) + + assert result.id == comment.id + assert result.created_by_account.id == OWNER_ID + assert result.resolved_by_account is None + + def test_delete_comment_raises_forbidden(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + + with pytest.raises(Forbidden): + WorkflowCommentService.delete_comment(TENANT_ID, APP_ID, comment.id, OUTSIDER_ID) + + assert sqlite_session.get(WorkflowComment, comment.id) is not None + + def test_delete_comment_removes_related_entities(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + comment_id = comment.id + reply = WorkflowCommentReply(comment_id=comment.id, content="reply", created_by=USER_2_ID) + _persist(sqlite_session, reply) + _persist( + sqlite_session, + WorkflowCommentMention(comment_id=comment.id, mentioned_user_id=USER_3_ID), + WorkflowCommentMention(comment_id=comment.id, reply_id=reply.id, mentioned_user_id=USER_4_ID), + ) + + WorkflowCommentService.delete_comment(TENANT_ID, APP_ID, comment_id, OWNER_ID) + + sqlite_session.expire_all() + assert sqlite_session.get(WorkflowComment, comment_id) is None + assert ( + sqlite_session.scalar( + select(func.count()) + .select_from(WorkflowCommentReply) + .where(WorkflowCommentReply.comment_id == comment_id) + ) + == 0 + ) + assert ( + sqlite_session.scalar( + select(func.count()) + .select_from(WorkflowCommentMention) + .where(WorkflowCommentMention.comment_id == comment_id) + ) + == 0 + ) + + def test_resolve_comment_sets_fields(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + now = datetime(2026, 7, 10, 12, 0, 0) + + with patch.object(service_module, "naive_utc_now", return_value=now): + result = WorkflowCommentService.resolve_comment(TENANT_ID, APP_ID, comment.id, USER_2_ID) + + assert result.resolved is True + assert result.resolved_at == now + assert result.resolved_by == USER_2_ID + sqlite_session.expire_all() + persisted_comment = sqlite_session.get(WorkflowComment, comment.id) + assert persisted_comment is not None + assert persisted_comment.resolved is True + assert persisted_comment.resolved_at == now + assert persisted_comment.resolved_by == USER_2_ID + + def test_resolve_comment_noop_when_already_resolved(self, sqlite_session: Session) -> None: + resolved_at = datetime(2026, 7, 9, 12, 0, 0) + comment = _comment(resolved=True, resolved_at=resolved_at, resolved_by=USER_2_ID) + _persist(sqlite_session, comment) + + result = WorkflowCommentService.resolve_comment(TENANT_ID, APP_ID, comment.id, USER_3_ID) + + assert result.resolved_at == resolved_at + assert result.resolved_by == USER_2_ID + + def test_create_reply_requires_comment(self, sqlite_session: Session) -> None: + with pytest.raises(NotFound): + WorkflowCommentService.create_reply("missing-comment", "hello", OWNER_ID) + + def test_create_reply_creates_mentions(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment, _membership(USER_2_ID)) + + result = WorkflowCommentService.create_reply( + comment_id=comment.id, + content="hello", + created_by=OWNER_ID, + mentioned_user_ids=[USER_2_ID, OUTSIDER_ID], + ) + + reply = sqlite_session.get(WorkflowCommentReply, result["id"]) + assert reply is not None + assert reply.content == "hello" + assert reply.created_at == result["created_at"] + mentions = sqlite_session.scalars( + select(WorkflowCommentMention).where(WorkflowCommentMention.reply_id == reply.id) + ).all() + assert [mention.mentioned_user_id for mention in mentions] == [USER_2_ID] + + def test_update_reply_raises_not_found(self, sqlite_session: Session) -> None: + with pytest.raises(NotFound): + WorkflowCommentService.update_reply( + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id="missing-comment", + reply_id="missing-reply", + user_id=OWNER_ID, + content="hello", + ) + + def test_update_reply_raises_forbidden(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + reply = WorkflowCommentReply(comment_id=comment.id, content="reply", created_by=OWNER_ID) + _persist(sqlite_session, reply) + + with pytest.raises(Forbidden): + WorkflowCommentService.update_reply( + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + reply_id=reply.id, + user_id=OUTSIDER_ID, + content="hello", + ) + + def test_update_reply_replaces_mentions(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment, _membership(USER_2_ID)) + reply = WorkflowCommentReply(comment_id=comment.id, content="reply", created_by=OWNER_ID) + _persist(sqlite_session, reply) + _persist( + sqlite_session, + WorkflowCommentMention(comment_id=comment.id, reply_id=reply.id, mentioned_user_id=USER_3_ID), + ) + + result = WorkflowCommentService.update_reply( + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + reply_id=reply.id, + user_id=OWNER_ID, + content="new", + mentioned_user_ids=[USER_2_ID, OUTSIDER_ID], + ) + + sqlite_session.expire_all() + persisted_reply = sqlite_session.get(WorkflowCommentReply, reply.id) + assert persisted_reply is not None + assert persisted_reply.content == "new" + assert result == {"id": reply.id, "updated_at": persisted_reply.updated_at} + mentions = sqlite_session.scalars( + select(WorkflowCommentMention).where(WorkflowCommentMention.reply_id == reply.id) + ).all() + assert [mention.mentioned_user_id for mention in mentions] == [USER_2_ID] + + def test_update_comment_updates_position_coordinates_when_provided(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + + WorkflowCommentService.update_comment( + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + user_id=OWNER_ID, content="updated", position_x=10.5, position_y=20.5, mentioned_user_ids=[], ) - assert comment.position_x == 10.5 - assert comment.position_y == 20.5 + sqlite_session.expire_all() + persisted_comment = sqlite_session.get(WorkflowComment, comment.id) + assert persisted_comment is not None + assert persisted_comment.position_x == 10.5 + assert persisted_comment.position_y == 20.5 - def test_delete_reply_raises_forbidden(self, mock_session: Mock) -> None: - reply = Mock() - reply.created_by = "owner" - mock_session.scalar.return_value = reply + def test_delete_reply_raises_forbidden(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + reply = WorkflowCommentReply(comment_id=comment.id, content="reply", created_by=OWNER_ID) + _persist(sqlite_session, reply) with pytest.raises(Forbidden): WorkflowCommentService.delete_reply( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - reply_id="reply-1", - user_id="intruder", + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + reply_id=reply.id, + user_id=OUTSIDER_ID, ) - def test_delete_reply_raises_not_found(self, mock_session: Mock) -> None: - mock_session.scalar.return_value = None - + def test_delete_reply_raises_not_found(self, sqlite_session: Session) -> None: with pytest.raises(NotFound): WorkflowCommentService.delete_reply( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - reply_id="reply-1", - user_id="owner", + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id="missing-comment", + reply_id="missing-reply", + user_id=OWNER_ID, ) - def test_delete_reply_removes_mentions(self, mock_session: Mock) -> None: - reply = Mock() - reply.created_by = "owner" - mock_session.scalar.return_value = reply - mock_session.scalars.return_value = _mock_scalars([Mock(), Mock()]) - - WorkflowCommentService.delete_reply( - tenant_id="tenant-1", - app_id="app-1", - comment_id="comment-1", - reply_id="reply-1", - user_id="owner", + def test_delete_reply_removes_mentions(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + reply = WorkflowCommentReply(comment_id=comment.id, content="reply", created_by=OWNER_ID) + _persist(sqlite_session, reply) + reply_id = reply.id + _persist( + sqlite_session, + WorkflowCommentMention(comment_id=comment.id, reply_id=reply.id, mentioned_user_id=USER_2_ID), + WorkflowCommentMention(comment_id=comment.id, reply_id=reply.id, mentioned_user_id=USER_3_ID), ) - assert mock_session.delete.call_count == 3 - mock_session.commit.assert_called_once() + WorkflowCommentService.delete_reply( + tenant_id=TENANT_ID, + app_id=APP_ID, + comment_id=comment.id, + reply_id=reply_id, + user_id=OWNER_ID, + ) - def test_validate_comment_access_delegates_to_get_comment(self) -> None: - comment = Mock() - with patch.object(WorkflowCommentService, "get_comment", return_value=comment) as get_comment_mock: - result = WorkflowCommentService.validate_comment_access("comment-1", "tenant-1", "app-1") + sqlite_session.expire_all() + assert sqlite_session.get(WorkflowCommentReply, reply_id) is None + assert ( + sqlite_session.scalar( + select(func.count()) + .select_from(WorkflowCommentMention) + .where(WorkflowCommentMention.reply_id == reply_id) + ) + == 0 + ) - assert result is comment - get_comment_mock.assert_called_once_with("tenant-1", "app-1", "comment-1") + def test_validate_comment_access_delegates_to_get_comment(self, sqlite_session: Session) -> None: + comment = _comment() + _persist(sqlite_session, comment) + + result = WorkflowCommentService.validate_comment_access(comment.id, TENANT_ID, APP_ID) + + assert result.id == comment.id + + def test_reply_lookup_is_scoped_to_tenant_app_and_comment(self, sqlite_session: Session) -> None: + comment = _comment() + other_comment = _comment(app_id=OTHER_APP_ID) + _persist(sqlite_session, comment, other_comment) + reply = WorkflowCommentReply(comment_id=comment.id, content="reply", created_by=OWNER_ID) + _persist(sqlite_session, reply) + + with pytest.raises(NotFound): + WorkflowCommentService.update_reply( + tenant_id=TENANT_ID, + app_id=OTHER_APP_ID, + comment_id=other_comment.id, + reply_id=reply.id, + user_id=OWNER_ID, + content="cross-thread update", + ) + + sqlite_session.refresh(reply) + assert reply.content == "reply" diff --git a/api/tests/unit_tests/services/test_workflow_run_service.py b/api/tests/unit_tests/services/test_workflow_run_service.py index fcfa9992cd1..b5902354165 100644 --- a/api/tests/unit_tests/services/test_workflow_run_service.py +++ b/api/tests/unit_tests/services/test_workflow_run_service.py @@ -1,11 +1,16 @@ +"""Workflow-run service tests with real SQLite-bound session factories.""" + +from decimal import Decimal from types import SimpleNamespace from typing import Any, cast from unittest.mock import MagicMock import pytest -from sqlalchemy import Engine +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session, sessionmaker -from models import Account, App, EndUser, WorkflowRunTriggeredFrom +from models import Account, App, EndUser, Message, WorkflowRunTriggeredFrom +from models.enums import ConversationFromSource from services import workflow_run_service as service_module from services.workflow_run_service import WorkflowRunService @@ -22,6 +27,11 @@ def repository_factory_mocks(monkeypatch: pytest.MonkeyPatch) -> tuple[MagicMock return node_repo, workflow_run_repo, factory +@pytest.fixture +def sqlalchemy_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]: + return sessionmaker(bind=sqlite_engine, expire_on_commit=False) + + def _app_model(**kwargs: Any) -> App: return cast(App, SimpleNamespace(**kwargs)) @@ -34,13 +44,22 @@ def _end_user(**kwargs: Any) -> EndUser: return cast(EndUser, SimpleNamespace(**kwargs)) -def _fake_session_factory_returning_messages(messages: list[Any]) -> tuple[MagicMock, MagicMock]: - """Build a session factory whose session returns the given messages.""" - session = MagicMock() - session.scalars.return_value.all.return_value = messages - session_factory = MagicMock() - session_factory.return_value.__enter__.return_value = session - return session_factory, session +def _message(*, message_id: str, workflow_run_id: str, conversation_id: str) -> Message: + message = Message( + app_id="app-1", + conversation_id=conversation_id, + query="query", + message={"role": "user", "content": "query"}, + answer="answer", + message_unit_price=Decimal("0.0001"), + answer_unit_price=Decimal("0.0001"), + currency="USD", + from_source=ConversationFromSource.API, + ) + message.id = message_id + message._inputs = {} + message.workflow_run_id = workflow_run_id + return message class TestWorkflowRunServiceInitialization: @@ -48,59 +67,51 @@ class TestWorkflowRunServiceInitialization: self, monkeypatch: pytest.MonkeyPatch, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlite_engine: Engine, ) -> None: - session_factory = MagicMock(name="session_factory") - sessionmaker_mock = MagicMock(return_value=session_factory) - monkeypatch.setattr(service_module, "sessionmaker", sessionmaker_mock) - monkeypatch.setattr(service_module, "db", SimpleNamespace(engine="db-engine")) + monkeypatch.setattr(service_module, "db", SimpleNamespace(engine=sqlite_engine)) service = WorkflowRunService() - sessionmaker_mock.assert_called_once_with(bind="db-engine", expire_on_commit=False) - assert service._session_factory is session_factory + assert isinstance(service._session_factory, sessionmaker) + assert service._session_factory.kw["bind"] is sqlite_engine + assert service._session_factory.kw["expire_on_commit"] is False def test___init___should_create_sessionmaker_when_engine_is_provided( self, - monkeypatch: pytest.MonkeyPatch, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlite_engine: Engine, ) -> None: - class FakeEngine: - pass + service = WorkflowRunService(session_factory=sqlite_engine) - session_factory = MagicMock(name="session_factory") - sessionmaker_mock = MagicMock(return_value=session_factory) - monkeypatch.setattr(service_module, "Engine", FakeEngine) - monkeypatch.setattr(service_module, "sessionmaker", sessionmaker_mock) - engine = cast(Engine, FakeEngine()) - - service = WorkflowRunService(session_factory=engine) - - sessionmaker_mock.assert_called_once_with(bind=engine, expire_on_commit=False) - assert service._session_factory is session_factory + assert isinstance(service._session_factory, sessionmaker) + assert service._session_factory.kw["bind"] is sqlite_engine + assert service._session_factory.kw["expire_on_commit"] is False def test___init___should_keep_provided_sessionmaker_and_create_repositories( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: node_repo, workflow_run_repo, factory = repository_factory_mocks - session_factory = MagicMock(name="session_factory") - service = WorkflowRunService(session_factory=session_factory) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) - assert service._session_factory is session_factory + assert service._session_factory is sqlalchemy_session_factory assert service._node_execution_service_repo is node_repo assert service._workflow_run_repo is workflow_run_repo - factory.create_api_workflow_node_execution_repository.assert_called_once_with(session_factory) - factory.create_api_workflow_run_repository.assert_called_once_with(session_factory) + factory.create_api_workflow_node_execution_repository.assert_called_once_with(sqlalchemy_session_factory) + factory.create_api_workflow_run_repository.assert_called_once_with(sqlalchemy_session_factory) class TestWorkflowRunServiceQueries: def test_get_paginate_workflow_runs_should_forward_filters_and_parse_limit( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: _, workflow_run_repo, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") expected = MagicMock(name="pagination") workflow_run_repo.get_paginated_workflow_runs.return_value = expected @@ -122,20 +133,24 @@ class TestWorkflowRunServiceQueries: status="succeeded", ) + @pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) def test_get_paginate_advanced_chat_workflow_runs_should_attach_message_fields_when_message_exists( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], + sqlite_session: Session, ) -> None: - message = SimpleNamespace(id="msg-1", conversation_id="conv-1", workflow_run_id="run-1") - session_factory, session = _fake_session_factory_returning_messages([message]) - service = WorkflowRunService(session_factory=session_factory) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") run_with_message = SimpleNamespace(id="run-1", status="running") run_without_message = SimpleNamespace(id="run-2", status="succeeded") pagination = SimpleNamespace(data=[run_with_message, run_without_message]) monkeypatch.setattr(service, "get_paginate_workflow_runs", MagicMock(return_value=pagination)) + sqlite_session.add(_message(message_id="msg-1", conversation_id="conv-1", workflow_run_id="run-1")) + sqlite_session.commit() + result = service.get_paginate_advanced_chat_workflow_runs(app_model=app_model, args={"limit": "2"}) assert result is pagination @@ -145,39 +160,49 @@ class TestWorkflowRunServiceQueries: assert result.data[0].status == "running" assert not hasattr(result.data[1], "message_id") assert result.data[1].id == "run-2" - # Messages are batch-loaded in a single query, not one per run. - session_factory.assert_called_once_with() - session.scalars.assert_called_once() + @pytest.mark.parametrize("sqlite_session", [(Message,)], indirect=True) def test_get_paginate_advanced_chat_workflow_runs_batch_loads_messages_without_n_plus_one( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], + sqlite_session: Session, ) -> None: """Messages must load with a constant query count regardless of run count. Previously the deprecated WorkflowRun.message property issued one query per run (N+1); they are now batch-loaded in a single query. """ - session_factory, session = _fake_session_factory_returning_messages([]) - service = WorkflowRunService(session_factory=session_factory) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") runs = [SimpleNamespace(id=f"run-{i}", status="succeeded") for i in range(5)] pagination = SimpleNamespace(data=runs) monkeypatch.setattr(service, "get_paginate_workflow_runs", MagicMock(return_value=pagination)) - service.get_paginate_advanced_chat_workflow_runs(app_model=app_model, args={}) + message_query_count = 0 - # Exactly one message query for the whole page, independent of run count. - session_factory.assert_called_once_with() - assert session.scalars.call_count == 1 + def count_message_query(*_args: object) -> None: + nonlocal message_query_count + message_query_count += 1 + + engine = sqlite_session.get_bind() + event.listen(engine, "before_cursor_execute", count_message_query) + try: + service.get_paginate_advanced_chat_workflow_runs(app_model=app_model, args={}) + finally: + event.remove(engine, "before_cursor_execute", count_message_query) + + assert all(not hasattr(run, "message_id") for run in runs) + assert message_query_count == 1 def test_get_workflow_run_should_delegate_to_repository_by_tenant_and_app( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: _, workflow_run_repo, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") expected = MagicMock(name="workflow_run") workflow_run_repo.get_workflow_run_by_id.return_value = expected @@ -194,9 +219,10 @@ class TestWorkflowRunServiceQueries: def test_get_workflow_runs_count_should_forward_optional_filters( self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: _, workflow_run_repo, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) app_model = _app_model(tenant_id="tenant-1", id="app-1") expected = {"total": 3, "succeeded": 2} workflow_run_repo.get_workflow_runs_count.return_value = expected @@ -221,8 +247,9 @@ class TestWorkflowRunServiceQueries: self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) monkeypatch.setattr(service, "get_workflow_run", MagicMock(return_value=None)) app_model = _app_model(id="app-1") user = _account(current_tenant_id="tenant-1") @@ -235,9 +262,10 @@ class TestWorkflowRunServiceQueries: self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: node_repo, _, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) monkeypatch.setattr(service, "get_workflow_run", MagicMock(return_value=SimpleNamespace(id="run-1"))) class FakeEndUser: @@ -267,9 +295,10 @@ class TestWorkflowRunServiceQueries: self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: node_repo, _, _ = repository_factory_mocks - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) monkeypatch.setattr(service, "get_workflow_run", MagicMock(return_value=SimpleNamespace(id="run-1"))) app_model = _app_model(id="app-1") user = _account(current_tenant_id="tenant-account") @@ -293,8 +322,9 @@ class TestWorkflowRunServiceQueries: self, repository_factory_mocks: tuple[MagicMock, MagicMock, Any], monkeypatch: pytest.MonkeyPatch, + sqlalchemy_session_factory: sessionmaker[Session], ) -> None: - service = WorkflowRunService(session_factory=MagicMock(name="session_factory")) + service = WorkflowRunService(session_factory=sqlalchemy_session_factory) monkeypatch.setattr(service, "get_workflow_run", MagicMock(return_value=SimpleNamespace(id="run-1"))) app_model = _app_model(id="app-1") user = _account(current_tenant_id=None) diff --git a/api/tests/unit_tests/services/test_workflow_run_service_pause.py b/api/tests/unit_tests/services/test_workflow_run_service_pause.py index 239cc83518e..5b24cfd8a66 100644 --- a/api/tests/unit_tests/services/test_workflow_run_service_pause.py +++ b/api/tests/unit_tests/services/test_workflow_run_service_pause.py @@ -1,178 +1,57 @@ -"""Comprehensive unit tests for WorkflowRunService class. +"""Tests for the session lifecycle owned by ``WorkflowRunService``.""" -This test suite covers all pause state management operations including: -- Retrieving pause state for workflow runs -- Saving pause state with file uploads -- Marking paused workflows as resumed -- Error handling and edge cases -- Database transaction management -- Repository-based approach testing -""" - -from datetime import datetime -from unittest.mock import MagicMock, create_autospec, patch +from unittest.mock import create_autospec, patch import pytest -from sqlalchemy import Engine +from sqlalchemy import Engine, text from sqlalchemy.orm import Session, sessionmaker -from graphon.enums import WorkflowExecutionStatus -from models.workflow import WorkflowPause from repositories.api_workflow_run_repository import APIWorkflowRunRepository -from repositories.sqlalchemy_api_workflow_run_repository import _PrivateWorkflowPauseEntity -from services.workflow_run_service import ( - WorkflowRunService, -) +from services.workflow_run_service import WorkflowRunService -class TestDataFactory: - """Factory class for creating test data objects.""" - - @staticmethod - def create_workflow_run_mock( - id: str = "workflow-run-123", - tenant_id: str = "tenant-456", - app_id: str = "app-789", - workflow_id: str = "workflow-101", - status: str | WorkflowExecutionStatus = "paused", - **kwargs, - ) -> MagicMock: - """Create a mock WorkflowRun object.""" - mock_run = MagicMock() - mock_run.id = id - mock_run.tenant_id = tenant_id - mock_run.app_id = app_id - mock_run.workflow_id = workflow_id - mock_run.status = status - - for key, value in kwargs.items(): - setattr(mock_run, key, value) - - return mock_run - - @staticmethod - def create_workflow_pause_mock( - id: str = "pause-123", - tenant_id: str = "tenant-456", - app_id: str = "app-789", - workflow_id: str = "workflow-101", - workflow_execution_id: str = "workflow-execution-123", - state_file_id: str = "file-456", - resumed_at: datetime | None = None, - **kwargs, - ) -> MagicMock: - """Create a mock WorkflowPauseModel object.""" - mock_pause = MagicMock(spec=WorkflowPause) - mock_pause.id = id - mock_pause.tenant_id = tenant_id - mock_pause.app_id = app_id - mock_pause.workflow_id = workflow_id - mock_pause.workflow_execution_id = workflow_execution_id - mock_pause.state_file_id = state_file_id - mock_pause.resumed_at = resumed_at - - for key, value in kwargs.items(): - setattr(mock_pause, key, value) - - return mock_pause - - @staticmethod - def create_pause_entity_mock( - pause_model: MagicMock | None = None, - ) -> _PrivateWorkflowPauseEntity: - """Create a mock _PrivateWorkflowPauseEntity object.""" - if pause_model is None: - pause_model = TestDataFactory.create_workflow_pause_mock() - - return _PrivateWorkflowPauseEntity(pause_model=pause_model, reason_models=[], human_input_form=[]) +@pytest.fixture +def sqlite_session_factory(sqlite_engine: Engine) -> sessionmaker[Session]: + """Return a real factory whose sessions are bound to the isolated SQLite engine.""" + return sessionmaker(bind=sqlite_engine, expire_on_commit=False) -class TestWorkflowRunService: - """Comprehensive unit tests for WorkflowRunService class.""" +@pytest.fixture +def workflow_run_repository(): + """Keep the repository boundary mocked while exercising real session construction.""" + return create_autospec(APIWorkflowRunRepository) - @pytest.fixture - def mock_session_factory(self): - """Create a mock session factory with proper session management.""" - mock_session = create_autospec(Session) - # Create a mock context manager for the session - mock_session_cm = MagicMock() - mock_session_cm.__enter__ = MagicMock(return_value=mock_session) - mock_session_cm.__exit__ = MagicMock(return_value=None) +def test_init_with_session_factory( + sqlite_session_factory: sessionmaker[Session], workflow_run_repository: APIWorkflowRunRepository +) -> None: + with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as repository_factory: + repository_factory.create_api_workflow_run_repository.return_value = workflow_run_repository - # Create a mock context manager for the transaction - mock_transaction_cm = MagicMock() - mock_transaction_cm.__enter__ = MagicMock(return_value=mock_session) - mock_transaction_cm.__exit__ = MagicMock(return_value=None) + service = WorkflowRunService(sqlite_session_factory) - mock_session.begin = MagicMock(return_value=mock_transaction_cm) + assert service._session_factory is sqlite_session_factory + repository_factory.create_api_workflow_run_repository.assert_called_once_with(sqlite_session_factory) + with service._session_factory() as session: + assert session.scalar(text("SELECT 1")) == 1 - # Create mock factory that returns the context manager - mock_factory = MagicMock(spec=sessionmaker) - mock_factory.return_value = mock_session_cm - return mock_factory, mock_session +def test_init_with_engine_creates_bound_session_factory( + sqlite_engine: Engine, workflow_run_repository: APIWorkflowRunRepository +) -> None: + with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as repository_factory: + repository_factory.create_api_workflow_run_repository.return_value = workflow_run_repository - @pytest.fixture - def mock_workflow_run_repository(self): - """Create a mock APIWorkflowRunRepository.""" - mock_repo = create_autospec(APIWorkflowRunRepository) - return mock_repo + service = WorkflowRunService(sqlite_engine) - @pytest.fixture - def workflow_run_service(self, mock_session_factory, mock_workflow_run_repository): - """Create WorkflowRunService instance with mocked dependencies.""" - session_factory, _ = mock_session_factory + assert service._session_factory.kw["bind"] is sqlite_engine + assert service._session_factory.kw["expire_on_commit"] is False + repository_factory.create_api_workflow_run_repository.assert_called_once_with(service._session_factory) + with service._session_factory() as session: + assert session.scalar(text("SELECT 1")) == 1 - with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as mock_factory: - mock_factory.create_api_workflow_run_repository.return_value = mock_workflow_run_repository - service = WorkflowRunService(session_factory) - return service - @pytest.fixture - def workflow_run_service_with_engine(self, mock_session_factory, mock_workflow_run_repository): - """Create WorkflowRunService instance with Engine input.""" - mock_engine = create_autospec(Engine) - session_factory, _ = mock_session_factory +def test_init_with_default_repository_dependencies(sqlite_session_factory: sessionmaker[Session]) -> None: + service = WorkflowRunService(sqlite_session_factory) - with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as mock_factory: - mock_factory.create_api_workflow_run_repository.return_value = mock_workflow_run_repository - service = WorkflowRunService(mock_engine) - return service - - # ==================== Initialization Tests ==================== - - def test_init_with_session_factory(self, mock_session_factory, mock_workflow_run_repository): - """Test WorkflowRunService initialization with session_factory.""" - session_factory, _ = mock_session_factory - - with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as mock_factory: - mock_factory.create_api_workflow_run_repository.return_value = mock_workflow_run_repository - service = WorkflowRunService(session_factory) - - assert service._session_factory == session_factory - mock_factory.create_api_workflow_run_repository.assert_called_once_with(session_factory) - - def test_init_with_engine(self, mock_session_factory, mock_workflow_run_repository): - """Test WorkflowRunService initialization with Engine (should convert to sessionmaker).""" - mock_engine = create_autospec(Engine) - session_factory, _ = mock_session_factory - - with patch("services.workflow_run_service.DifyAPIRepositoryFactory", autospec=True) as mock_factory: - mock_factory.create_api_workflow_run_repository.return_value = mock_workflow_run_repository - with patch( - "services.workflow_run_service.sessionmaker", return_value=session_factory, autospec=True - ) as mock_sessionmaker: - service = WorkflowRunService(mock_engine) - - mock_sessionmaker.assert_called_once_with(bind=mock_engine, expire_on_commit=False) - assert service._session_factory == session_factory - mock_factory.create_api_workflow_run_repository.assert_called_once_with(session_factory) - - def test_init_with_default_dependencies(self, mock_session_factory): - """Test WorkflowRunService initialization with default dependencies.""" - session_factory, _ = mock_session_factory - - service = WorkflowRunService(session_factory) - - assert service._session_factory == session_factory + assert service._session_factory is sqlite_session_factory diff --git a/api/tests/unit_tests/services/test_workflow_service.py b/api/tests/unit_tests/services/test_workflow_service.py index b2e0e4129c9..d0ff4a085e9 100644 --- a/api/tests/unit_tests/services/test_workflow_service.py +++ b/api/tests/unit_tests/services/test_workflow_service.py @@ -11,15 +11,19 @@ This test suite covers: import json import uuid +from datetime import datetime, timedelta from types import SimpleNamespace from typing import Any, cast from unittest.mock import ANY, MagicMock, patch, sentinel import pytest -from sqlalchemy import select +from sqlalchemy import event, select +from sqlalchemy.dialects import postgresql from sqlalchemy.engine import Engine from sqlalchemy.orm import Session, sessionmaker +from core.workflow.llm_environment_variable import LLMEnvironmentVariable +from enums import DeploymentEdition from graphon.enums import ( BuiltinNodeTypes, ErrorStrategy, @@ -35,7 +39,6 @@ from graphon.variables import StringVariable from graphon.variables.input_entities import VariableEntityType from libs.datetime_utils import naive_utc_now from models.account import Account -from models.agent import WorkflowAgentNodeBinding from models.human_input import HumanInputFormRecipient, RecipientType from models.model import App, AppMode from models.tools import BuiltinToolProvider, WorkflowToolProvider @@ -199,11 +202,6 @@ class TestWorkflowAssociatedDataFactory: @pytest.mark.usefixtures("sqlite_session") -@pytest.mark.parametrize( - "sqlite_session", - [(Workflow, App, WorkflowToolProvider, HumanInputFormRecipient, WorkflowAgentNodeBinding)], - indirect=True, -) class TestWorkflowService: """ Comprehensive unit tests for WorkflowService methods. @@ -359,6 +357,31 @@ class TestWorkflowService: assert result is None + @pytest.mark.parametrize( + ("tenant_id", "app_id"), + [("other-tenant", "app-123"), ("tenant-456", "other-app")], + ) + def test_get_published_workflow_by_id_rejects_foreign_workflow( + self, + tenant_id: str, + app_id: str, + workflow_service: WorkflowService, + sqlite_session: Session, + ): + app = TestWorkflowAssociatedDataFactory.create_app() + workflow = TestWorkflowAssociatedDataFactory.create_workflow( + workflow_id="workflow-123", + tenant_id=tenant_id, + app_id=app_id, + version="v1", + ) + sqlite_session.add(workflow) + sqlite_session.commit() + + result = workflow_service.get_published_workflow_by_id(app, workflow.id, session=sqlite_session) + + assert result is None + def test_get_published_workflow_success(self, workflow_service: WorkflowService, sqlite_session: Session): """Test get_published_workflow returns published workflow.""" workflow_id = "workflow-123" @@ -449,6 +472,121 @@ class TestWorkflowService: assert workflow.features_dict == features assert workflow.updated_by == account.id + def test_sync_draft_workflow_collaborative_save_preserves_environment_variables_and_locks_row( + self, workflow_service: WorkflowService, sqlite_session: Session + ) -> None: + """A collaborative graph-only save locks the draft and keeps server environment values.""" + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + workflow = TestWorkflowAssociatedDataFactory.create_workflow() + remote_variable = StringVariable( + id="env-shared", + name="shared", + value="remote-value", + selector=["env", "shared"], + ) + stale_variable = remote_variable.model_copy(update={"value": "stale-client-value"}) + workflow.environment_variables = [remote_variable] + sqlite_session.add(workflow) + sqlite_session.commit() + unique_hash = workflow.unique_hash + statements = [] + commits = [] + + @event.listens_for(sqlite_session, "do_orm_execute") + def capture_statement(execute_state): + statements.append(execute_state.statement) + + @event.listens_for(sqlite_session, "before_commit") + def capture_commit(_session): + commits.append(True) + + result = workflow_service.sync_draft_workflow( + app_model=app, + graph=TestWorkflowAssociatedDataFactory.create_valid_workflow_graph(), + features={}, + unique_hash=unique_hash, + account=account, + environment_variables=[stale_variable], + conversation_variables=[], + session=sqlite_session, + preserve_environment_variables=True, + commit=False, + sync_agent_bindings=False, + ) + + assert "FOR UPDATE" in str(statements[0].compile(dialect=postgresql.dialect())) + assert result.environment_variables == [remote_variable] + assert commits == [] + + def test_sync_draft_workflow_merges_environment_patch_with_graph_update( + self, workflow_service: WorkflowService, sqlite_session: Session + ) -> None: + """Graph changes and a per-ID environment patch commit without replacing untouched aliases.""" + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + workflow = TestWorkflowAssociatedDataFactory.create_workflow() + remote_variable = StringVariable(id="env-a", name="a", value="remote-a", selector=["env", "a"]) + replaced_variable = StringVariable(id="env-b", name="b", value="old-b", selector=["env", "b"]) + workflow.environment_variables = [remote_variable, replaced_variable] + sqlite_session.add(workflow) + sqlite_session.commit() + unique_hash = workflow.unique_hash + next_graph = TestWorkflowAssociatedDataFactory.create_valid_workflow_graph() + next_graph["viewport"] = {"x": 10, "y": 20, "zoom": 1} + + with patch("services.workflow_service.app_draft_workflow_was_synced"): + result = workflow_service.sync_draft_workflow( + app_model=app, + graph=next_graph, + features={}, + unique_hash=unique_hash, + account=account, + environment_variables=[ + remote_variable.model_copy(update={"value": "stale-a"}), + replaced_variable, + ], + conversation_variables=[], + session=sqlite_session, + environment_variable_upserts=[replaced_variable.model_copy(update={"value": "new-b"})], + deleted_environment_variable_ids=[], + preserve_environment_variables=True, + ) + + assert result.graph_dict == next_graph + assert [(variable.id, variable.value) for variable in result.environment_variables] == [ + ("env-a", "remote-a"), + ("env-b", "new-b"), + ] + + def test_sync_draft_workflow_noncollaborative_save_replaces_environment_variables( + self, workflow_service: WorkflowService, sqlite_session: Session + ) -> None: + """Legacy full-draft saves retain their full environment replacement contract.""" + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + workflow = TestWorkflowAssociatedDataFactory.create_workflow() + workflow.environment_variables = [ + StringVariable(id="env-old", name="old", value="old", selector=["env", "old"]) + ] + sqlite_session.add(workflow) + sqlite_session.commit() + replacement = StringVariable(id="env-new", name="new", value="new", selector=["env", "new"]) + + with patch("services.workflow_service.app_draft_workflow_was_synced"): + result = workflow_service.sync_draft_workflow( + app_model=app, + graph=TestWorkflowAssociatedDataFactory.create_valid_workflow_graph(), + features={}, + unique_hash=workflow.unique_hash, + account=account, + environment_variables=[replacement], + conversation_variables=[], + session=sqlite_session, + ) + + assert result.environment_variables == [replacement] + def test_sync_draft_workflow_graph_only_preserves_independently_updated_draft_fields( self, workflow_service: WorkflowService, sqlite_session: Session ): @@ -765,6 +903,104 @@ class TestWorkflowService: session=sqlite_session, ) + def test_patch_draft_workflow_environment_variables_merges_by_id( + self, workflow_service: WorkflowService, sqlite_session: Session + ) -> None: + """A patch preserves untouched variables while applying ordered upserts and deletions.""" + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + workflow = TestWorkflowAssociatedDataFactory.create_workflow() + workflow.environment_variables = [ + StringVariable(id="env-a", name="a", value="remote-a", selector=["env", "a"]), + StringVariable(id="env-b", name="b", value="old-b", selector=["env", "b"]), + StringVariable(id="env-c", name="c", value="remote-c", selector=["env", "c"]), + ] + sqlite_session.add(workflow) + sqlite_session.commit() + + workflow_service.patch_draft_workflow_environment_variables( + app_model=app, + environment_variables=[ + StringVariable(id="env-b", name="b", value="new-b", selector=["env", "b"]), + StringVariable(id="env-d", name="d", value="new-d", selector=["env", "d"]), + ], + deleted_environment_variable_ids=["env-a"], + account=account, + session=sqlite_session, + ) + + sqlite_session.expire_all() + persisted_workflow = sqlite_session.get(Workflow, workflow.id) + assert persisted_workflow is not None + assert [(variable.id, variable.value) for variable in persisted_workflow.environment_variables] == [ + ("env-b", "new-b"), + ("env-c", "remote-c"), + ("env-d", "new-d"), + ] + assert persisted_workflow.updated_by == account.id + + def test_patch_draft_workflow_environment_variables_locks_draft_row( + self, workflow_service: WorkflowService, sqlite_session: Session + ) -> None: + """The merge reads the draft with a row lock before applying a partial update.""" + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + workflow = TestWorkflowAssociatedDataFactory.create_workflow() + sqlite_session.add(workflow) + sqlite_session.commit() + statements = [] + + @event.listens_for(sqlite_session, "do_orm_execute") + def capture_statement(execute_state): + statements.append(execute_state.statement) + + workflow_service.patch_draft_workflow_environment_variables( + app_model=app, + environment_variables=[], + deleted_environment_variable_ids=[], + account=account, + session=sqlite_session, + ) + + compiled_statement = str(statements[0].compile(dialect=postgresql.dialect())) + assert "FOR UPDATE" in compiled_statement + assert sqlite_session.get(Workflow, workflow.id) is workflow + + def test_patch_draft_workflow_environment_variables_rejects_conflicting_ids( + self, workflow_service: WorkflowService, sqlite_session: Session + ) -> None: + """A patch cannot both upsert and delete the same variable ID.""" + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + variable = StringVariable(id="env-a", name="a", value="a", selector=["env", "a"]) + sqlite_session.add(TestWorkflowAssociatedDataFactory.create_workflow()) + sqlite_session.commit() + + with pytest.raises(ValueError, match="cannot be upserted and deleted"): + workflow_service.patch_draft_workflow_environment_variables( + app_model=app, + environment_variables=[variable], + deleted_environment_variable_ids=["env-a"], + account=account, + session=sqlite_session, + ) + + def test_patch_draft_workflow_environment_variables_raises_when_missing( + self, workflow_service: WorkflowService, sqlite_session: Session + ) -> None: + """A patch fails when the app has no draft workflow to lock.""" + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + + with pytest.raises(ValueError, match="No draft workflow found."): + workflow_service.patch_draft_workflow_environment_variables( + app_model=app, + environment_variables=[], + deleted_environment_variable_ids=[], + account=account, + session=sqlite_session, + ) + def test_update_draft_workflow_conversation_variables_updates_workflow( self, workflow_service: WorkflowService, sqlite_session: Session ): @@ -830,9 +1066,9 @@ class TestWorkflowService: with ( patch("services.workflow_service.app_published_workflow_was_updated"), - patch("services.workflow_service.dify_config.BILLING_ENABLED", False), + patch("services.workflow_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY), ): - result = workflow_service.publish_workflow( + result, retirement_candidates = workflow_service.publish_workflow( session=sqlite_session, app_model=app, account=account, @@ -845,6 +1081,105 @@ class TestWorkflowService: assert result.version != Workflow.VERSION_DRAFT assert result.marked_name == "Version 1" assert result.marked_comment == "Initial release" + assert retirement_candidates == set() + + def test_publish_workflow_numbers_versions_from_one( + self, workflow_service: WorkflowService, sqlite_session: Session + ): + """ + Test publish_workflow assigns an app-scoped version number starting at #1. + + The number is what users see when a version carries no name, so it has to be + stable and monotonic per app rather than derived from list position. + """ + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + graph = TestWorkflowAssociatedDataFactory.create_valid_workflow_graph() + + draft = TestWorkflowAssociatedDataFactory.create_workflow(version=Workflow.VERSION_DRAFT, graph=graph) + sqlite_session.add(draft) + sqlite_session.commit() + + with ( + patch("services.workflow_service.app_published_workflow_was_updated"), + patch( + "services.workflow_service.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ), + ): + first, _ = workflow_service.publish_workflow(session=sqlite_session, app_model=app, account=account) + second, _ = workflow_service.publish_workflow(session=sqlite_session, app_model=app, account=account) + + assert first.version_number == 1 + assert second.version_number == 2 + + def test_publish_workflow_does_not_reuse_a_deleted_version_number( + self, workflow_service: WorkflowService, sqlite_session: Session + ): + """ + Test version numbers are never handed out twice, even after a version is deleted. + + Deployment records and audit logs refer to versions by number, so reusing one + would make two different workflows share an identity. + """ + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + graph = TestWorkflowAssociatedDataFactory.create_valid_workflow_graph() + + draft = TestWorkflowAssociatedDataFactory.create_workflow(version=Workflow.VERSION_DRAFT, graph=graph) + sqlite_session.add(draft) + sqlite_session.commit() + + with ( + patch("services.workflow_service.app_published_workflow_was_updated"), + patch( + "services.workflow_service.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ), + ): + published, _ = workflow_service.publish_workflow(session=sqlite_session, app_model=app, account=account) + sqlite_session.flush() + sqlite_session.delete(published) + sqlite_session.flush() + + republished, _ = workflow_service.publish_workflow(session=sqlite_session, app_model=app, account=account) + + assert republished.version_number == 2 + + def test_publish_workflow_numbers_each_app_independently( + self, workflow_service: WorkflowService, sqlite_session: Session + ): + """ + Test the counter is scoped per app rather than global. + + Every app starts its own sequence at #1; a busy neighbour must not advance it. + """ + account = TestWorkflowAssociatedDataFactory.create_account() + graph = TestWorkflowAssociatedDataFactory.create_valid_workflow_graph() + + published: list[Workflow] = [] + for app_id in ("app-first", "app-second"): + app = TestWorkflowAssociatedDataFactory.create_app(app_id=app_id) + draft = TestWorkflowAssociatedDataFactory.create_workflow( + workflow_id=f"draft-{app_id}", + app_id=app_id, + version=Workflow.VERSION_DRAFT, + graph=graph, + ) + sqlite_session.add(draft) + sqlite_session.commit() + + with ( + patch("services.workflow_service.app_published_workflow_was_updated"), + patch( + "services.workflow_service.dify_config.DEPLOYMENT_EDITION", + DeploymentEdition.COMMUNITY, + ), + ): + workflow, _ = workflow_service.publish_workflow(session=sqlite_session, app_model=app, account=account) + published.append(workflow) + + assert [workflow.version_number for workflow in published] == [1, 1] def test_publish_workflow_no_draft_raises_error(self, workflow_service: WorkflowService, sqlite_session: Session): """ @@ -859,6 +1194,53 @@ class TestWorkflowService: with pytest.raises(ValueError, match="No valid workflow found"): workflow_service.publish_workflow(session=sqlite_session, app_model=app, account=account) + @pytest.mark.parametrize( + ("environment_variables", "error"), + [ + ([], "was not found or is not an LLM variable"), + ( + [ + LLMEnvironmentVariable( + name="shared_model", + value={"provider": "provider", "name": "model", "mode": "completion"}, + ) + ], + "uses mode 'completion'.*uses mode 'chat'", + ), + ], + ) + def test_publish_workflow_rejects_invalid_llm_environment_reference_without_credential_validation( + self, + workflow_service: WorkflowService, + sqlite_session: Session, + environment_variables: list[LLMEnvironmentVariable], + error: str, + ) -> None: + app = TestWorkflowAssociatedDataFactory.create_app() + account = TestWorkflowAssociatedDataFactory.create_account() + graph = TestWorkflowAssociatedDataFactory.create_valid_workflow_graph() + llm_node = next(node for node in graph["nodes"] if node["data"]["type"] == BuiltinNodeTypes.LLM) + llm_node["data"]["model"] = { + "provider": "provider", + "name": "model", + "mode": "chat", + "completion_params": {}, + } + llm_node["data"]["model_selector"] = ["env", "shared_model"] + draft = TestWorkflowAssociatedDataFactory.create_workflow(graph=graph) + draft.environment_variables = environment_variables + sqlite_session.add(draft) + sqlite_session.commit() + + with ( + patch( + "services.feature_service.FeatureService.get_system_features", + return_value=SimpleNamespace(plugin_manager=SimpleNamespace(enabled=False)), + ), + pytest.raises(ValueError, match=error), + ): + workflow_service.publish_workflow(session=sqlite_session, app_model=app, account=account) + def test_publish_workflow_trigger_limit_exceeded(self, workflow_service: WorkflowService, sqlite_session: Session): """ Test publish_workflow raises error when trigger node limit exceeded in SANDBOX plan. @@ -885,7 +1267,7 @@ class TestWorkflowService: sqlite_session.commit() with ( - patch("services.workflow_service.dify_config.BILLING_ENABLED", True), + patch("services.workflow_service.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch("services.workflow_service.BillingService") as MockBillingService, ): MockBillingService.get_info.return_value = {"subscription": {"plan": "sandbox"}} @@ -922,6 +1304,44 @@ class TestWorkflowService: assert len(workflows) == 5 assert has_more is False + def test_get_all_published_workflow_lists_the_draft_first( + self, workflow_service: WorkflowService, sqlite_session: Session + ): + """ + Test the draft heads the version list no matter how old it is. + + A draft is created together with its app and its `created_at` is never refreshed, + so ordering purely by publish time would put it last — off the first page entirely + once the app has accumulated enough published versions. + """ + app = TestWorkflowAssociatedDataFactory.create_app(workflow_id="workflow-3") + app_created_at = datetime(2026, 1, 1) + + sqlite_session.add( + TestWorkflowAssociatedDataFactory.create_workflow( + workflow_id="workflow-draft", + version=Workflow.VERSION_DRAFT, + created_at=app_created_at, + ) + ) + sqlite_session.add_all( + [ + TestWorkflowAssociatedDataFactory.create_workflow( + workflow_id=f"workflow-{i}", + version=f"2026-02-0{i} 00:00:00", + created_at=app_created_at + timedelta(days=i), + ) + for i in range(1, 4) + ] + ) + sqlite_session.commit() + + workflows, _ = workflow_service.get_all_published_workflow( + session=sqlite_session, app_model=app, page=1, limit=2, user_id=None + ) + + assert [workflow.id for workflow in workflows] == ["workflow-draft", "workflow-3"] + def test_get_all_published_workflow_has_more(self, workflow_service: WorkflowService, sqlite_session: Session): """ Test get_all_published_workflow indicates has_more when results exceed limit. @@ -1420,7 +1840,6 @@ 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. @@ -1528,6 +1947,78 @@ class TestWorkflowServiceCredentialValidation: # Assert mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4") + def test_validate_workflow_credentials_should_use_llm_environment_variable_model( + self, service: WorkflowService, sqlite_session: Session + ) -> None: + workflow = self._make_workflow( + [ + { + "id": "llm-node", + "data": { + "type": "llm", + "model": { + "provider": "old-provider", + "name": "old-model", + "mode": "chat", + "completion_params": {"temperature": 0.2}, + }, + "model_selector": ["env", "shared_model"], + }, + } + ] + ) + workflow.environment_variables = [ + LLMEnvironmentVariable( + name="shared_model", + value={ + "provider": "new-provider", + "name": "new-model", + "mode": "chat", + "completion_params": {"temperature": 0.8}, + }, + ) + ] + + with ( + patch.object(service, "_validate_llm_model_config") as validate_model, + patch.object(service, "_validate_load_balancing_credentials") as validate_load_balancing, + ): + service._validate_workflow_credentials(workflow, session=sqlite_session) + + validate_model.assert_called_once_with("tenant-1", "new-provider", "new-model") + validated_node_data = validate_load_balancing.call_args.args[1] + assert validated_node_data["model"] == { + "provider": "new-provider", + "name": "new-model", + "mode": "chat", + "completion_params": {"temperature": 0.8}, + } + + def test_validate_workflow_credentials_should_reject_llm_environment_variable_mode_mismatch( + self, service: WorkflowService, sqlite_session: Session + ) -> None: + workflow = self._make_workflow( + [ + { + "id": "llm-node", + "data": { + "type": "llm", + "model": {"provider": "provider", "name": "model", "mode": "chat"}, + "model_selector": ["env", "shared_model"], + }, + } + ] + ) + workflow.environment_variables = [ + LLMEnvironmentVariable( + name="shared_model", + value={"provider": "provider", "name": "model", "mode": "completion"}, + ) + ] + + with pytest.raises(ValueError, match="uses mode 'completion'.*uses mode 'chat'"): + service._validate_workflow_credentials(workflow, session=sqlite_session) + def test_validate_workflow_credentials_should_raise_for_llm_node_missing_model( self, service: WorkflowService, sqlite_session: Session ) -> None: @@ -2642,7 +3133,6 @@ class TestWorkflowServiceHumanInputOperations: }, ) - @pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True) def test_get_human_input_form_preview_should_raise_if_workflow_not_init( self, service: WorkflowService, sqlite_session: Session ) -> None: @@ -2655,7 +3145,6 @@ class TestWorkflowServiceHumanInputOperations: session=sqlite_session, ) - @pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True) def test_get_human_input_form_preview_should_raise_if_wrong_node_type( self, service: WorkflowService, sqlite_session: Session ) -> None: @@ -2671,7 +3160,6 @@ class TestWorkflowServiceHumanInputOperations: session=sqlite_session, ) - @pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True) def test_get_human_input_form_preview_success(self, service: WorkflowService, sqlite_session: Session) -> None: app_model = TestWorkflowAssociatedDataFactory.create_app(app_id="app-1", tenant_id="tenant-1") account = TestWorkflowAssociatedDataFactory.create_account(account_id="user-1") @@ -2695,7 +3183,6 @@ class TestWorkflowServiceHumanInputOperations: mock_node.render_form_content_before_submission.assert_called_once() mock_required_cls.return_value.model_dump.assert_called_once() - @pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True) def test_submit_human_input_form_preview_success(self, service: WorkflowService, sqlite_session: Session) -> None: app_model = TestWorkflowAssociatedDataFactory.create_app(app_id="app-1", tenant_id="tenant-1") account = TestWorkflowAssociatedDataFactory.create_account(account_id="user-1") @@ -2736,7 +3223,6 @@ class TestWorkflowServiceHumanInputOperations: assert result["__rendered_content"] == "Ticket: val1" mock_saver_cls.return_value.save.assert_called_once() - @pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True) def test_test_human_input_delivery_success(self, service: WorkflowService, sqlite_session: Session) -> None: draft = self._create_human_input_workflow() service.get_draft_workflow = MagicMock(return_value=draft) @@ -2760,7 +3246,6 @@ class TestWorkflowServiceHumanInputOperations: ) mock_test_srv.return_value.send_test.assert_called_once() - @pytest.mark.parametrize("sqlite_session", [(Workflow,)], indirect=True) def test_test_human_input_delivery_failure_cases(self, service: WorkflowService, sqlite_session: Session) -> None: draft = self._create_human_input_workflow() service.get_draft_workflow = MagicMock(return_value=draft) @@ -2778,7 +3263,6 @@ class TestWorkflowServiceHumanInputOperations: session=sqlite_session, ) - @pytest.mark.parametrize("sqlite_session", [(HumanInputFormRecipient,)], indirect=True) def test_load_email_recipients_parsing_failure(self, service: WorkflowService, sqlite_session: Session) -> None: """Malformed persisted recipient payloads are skipped instead of aborting delivery tests.""" recipient = HumanInputFormRecipient( diff --git a/api/tests/unit_tests/services/test_workspace_credit_pool.py b/api/tests/unit_tests/services/test_workspace_credit_pool.py new file mode 100644 index 00000000000..08210e6d391 --- /dev/null +++ b/api/tests/unit_tests/services/test_workspace_credit_pool.py @@ -0,0 +1,76 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +from enums import CloudPlan, DeploymentEdition +from services.credit_pool_service import CreditPoolBalance +from services.workspace_service import WorkspaceService + + +@pytest.mark.parametrize( + ("quota_limit", "quota_used", "remaining_credits", "is_unlimited"), + [(500, 120, 380, False), (-1, 999, -1, True)], +) +def test_get_effective_credit_pool_prefers_available_paid_pool( + quota_limit: int, quota_used: int, remaining_credits: int, is_unlimited: bool +) -> None: + session = MagicMock() + paid_pool = CreditPoolBalance( + tenant_id="tenant-1", + pool_type="paid", + quota_limit=quota_limit, + quota_used=quota_used, + ) + billing_info = { + "enabled": True, + "subscription": {"plan": CloudPlan.TEAM}, + "next_credit_reset_date": 1775001600, + } + config = SimpleNamespace(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD) + + with ( + patch("services.workspace_service.dify_config", config), + patch("services.workspace_service.BillingService.get_info", return_value=billing_info), + patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=paid_pool) as get_pool, + ): + result = WorkspaceService.get_effective_credit_pool("tenant-1", session=session) + + get_pool.assert_called_once_with(tenant_id="tenant-1", pool_type="paid", session=session) + assert result.pool_type == "paid" + assert result.quota_limit == quota_limit + assert result.quota_used == quota_used + assert result.remaining_credits == remaining_credits + assert result.is_unlimited is is_unlimited + assert result.is_exhausted is False + assert result.next_credit_reset_date == 1775001600 + + +def test_get_effective_credit_pool_exposes_exhausted_trial_pool() -> None: + session = MagicMock() + trial_pool = CreditPoolBalance( + tenant_id="tenant-1", + pool_type="trial", + quota_limit=200, + quota_used=200, + exhausted_at=1772323200, + ) + billing_info = { + "enabled": True, + "subscription": {"plan": CloudPlan.SANDBOX}, + } + config = SimpleNamespace(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD) + + with ( + patch("services.workspace_service.dify_config", config), + patch("services.workspace_service.BillingService.get_info", return_value=billing_info), + patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=trial_pool) as get_pool, + ): + result = WorkspaceService.get_effective_credit_pool("tenant-1", session=session) + + get_pool.assert_called_once_with(tenant_id="tenant-1", pool_type="trial", session=session) + assert result.pool_type == "trial" + assert result.remaining_credits == 0 + assert result.is_unlimited is False + assert result.is_exhausted is True + assert result.exhausted_at == 1772323200 diff --git a/api/tests/unit_tests/services/test_workspace_member_query_service.py b/api/tests/unit_tests/services/test_workspace_member_query_service.py new file mode 100644 index 00000000000..cc2d3fb859d --- /dev/null +++ b/api/tests/unit_tests/services/test_workspace_member_query_service.py @@ -0,0 +1,150 @@ +from collections.abc import Mapping, Sequence +from datetime import datetime + +import pytest + +from machinery.context import RequestContext +from services.workspace_member_query_service import ( + WorkspaceMemberQueryService, + WorkspaceMemberRecord, + WorkspaceMemberRole, + WorkspaceMemberRoleSubject, + WorkspaceMemberSummary, +) + + +def make_context(*, workspace_id: str | None = "workspace-1") -> RequestContext: + return RequestContext( + request_id="request-1", + trace_id="trace-1", + account_id="actor-1", + active_workspace_id=workspace_id, + ) + + +def make_member( + member_id: str, + *, + status: str = "active", + legacy_role: str = "normal", +) -> WorkspaceMemberRecord: + created_at = datetime(2026, 1, 1) + return WorkspaceMemberRecord( + id=member_id, + name=f"Member {member_id}", + email=f"{member_id}@example.com", + avatar=None, + last_login_at=None, + last_active_at=created_at, + created_at=created_at, + status=status, + legacy_role=legacy_role, + ) + + +class RecordingMemberQuery: + def __init__(self, records: Sequence[WorkspaceMemberRecord]) -> None: + self.records = tuple(records) + self.workspace_ids: list[str] = [] + + def list_for_workspace(self, workspace_id: str) -> Sequence[WorkspaceMemberRecord]: + self.workspace_ids.append(workspace_id) + return self.records + + +class RecordingRoleResolver: + def __init__(self, roles: Mapping[str, Sequence[WorkspaceMemberRole]]) -> None: + self.roles = roles + self.calls: list[tuple[str, str, tuple[WorkspaceMemberRoleSubject, ...]]] = [] + + def resolve_many( + self, + workspace_id: str, + actor_account_id: str, + subjects: Sequence[WorkspaceMemberRoleSubject], + ) -> Mapping[str, Sequence[WorkspaceMemberRole]]: + self.calls.append((workspace_id, actor_account_id, tuple(subjects))) + return self.roles + + +class FailingRoleResolver: + def resolve_many( + self, + workspace_id: str, + actor_account_id: str, + subjects: Sequence[WorkspaceMemberRoleSubject], + ) -> Mapping[str, Sequence[WorkspaceMemberRole]]: + del workspace_id, actor_account_id, subjects + raise RoleResolutionError + + +class RoleResolutionError(Exception): + pass + + +def test_list_current_projects_members_and_merges_roles_by_account_id() -> None: + active = make_member("active", legacy_role="owner") + pending = make_member("pending", status="pending") + members = RecordingMemberQuery([active, pending]) + roles = RecordingRoleResolver( + { + active.id: [ + WorkspaceMemberRole(id="workspace.owner", name="Owner"), + WorkspaceMemberRole(id="workspace.editor", name="Editor"), + ] + } + ) + service = WorkspaceMemberQueryService(members=members, roles=roles) + + result = service.list_current(make_context()) + + by_id = {member.id: member for member in result} + assert set(by_id) == {"active", "pending"} + assert by_id["active"] == WorkspaceMemberSummary( + id=active.id, + name=active.name, + email=active.email, + avatar=active.avatar, + last_login_at=active.last_login_at, + last_active_at=active.last_active_at, + created_at=active.created_at, + role=active.legacy_role, + roles=( + WorkspaceMemberRole(id="workspace.owner", name="Owner"), + WorkspaceMemberRole(id="workspace.editor", name="Editor"), + ), + status=active.status, + ) + assert by_id["pending"].status == "pending" + assert by_id["pending"].roles == () + assert members.workspace_ids == ["workspace-1"] + assert roles.calls == [ + ( + "workspace-1", + "actor-1", + ( + WorkspaceMemberRoleSubject(account_id=active.id, legacy_role=active.legacy_role), + WorkspaceMemberRoleSubject(account_id=pending.id, legacy_role=pending.legacy_role), + ), + ) + ] + + +def test_list_current_rejects_missing_workspace_before_calling_ports() -> None: + members = RecordingMemberQuery([]) + roles = RecordingRoleResolver({}) + service = WorkspaceMemberQueryService(members=members, roles=roles) + + with pytest.raises(RuntimeError, match="Console account admission did not resolve an active workspace"): + service.list_current(make_context(workspace_id=None)) + + assert members.workspace_ids == [] + assert roles.calls == [] + + +def test_list_current_propagates_role_resolution_failure() -> None: + members = RecordingMemberQuery([make_member("member-1")]) + service = WorkspaceMemberQueryService(members=members, roles=FailingRoleResolver()) + + with pytest.raises(RoleResolutionError): + service.list_current(make_context()) diff --git a/api/tests/unit_tests/services/test_workspace_member_role_resolver.py b/api/tests/unit_tests/services/test_workspace_member_role_resolver.py new file mode 100644 index 00000000000..b40f77da45c --- /dev/null +++ b/api/tests/unit_tests/services/test_workspace_member_role_resolver.py @@ -0,0 +1,128 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from services import workspace_member_role_resolver +from services.enterprise.rbac_service import MemberRolesResponse, RBACRole +from services.workspace_member_query_service import WorkspaceMemberRole, WorkspaceMemberRoleSubject + + +def make_subject(account_id: str, *, legacy_role: str = "normal") -> WorkspaceMemberRoleSubject: + return WorkspaceMemberRoleSubject(account_id=account_id, legacy_role=legacy_role) + + +@pytest.fixture +def batch_get(monkeypatch: pytest.MonkeyPatch) -> MagicMock: + batch_get = MagicMock() + monkeypatch.setattr( + workspace_member_role_resolver.enterprise_rbac_service.RBACService.MemberRoles, + "batch_get", + batch_get, + ) + return batch_get + + +def configure_rbac(monkeypatch: pytest.MonkeyPatch, *, enabled: bool) -> None: + monkeypatch.setattr( + workspace_member_role_resolver, + "dify_config", + SimpleNamespace(RBAC_ENABLED=enabled), + ) + + +def test_legacy_mode_projects_join_roles_without_enterprise_call( + monkeypatch: pytest.MonkeyPatch, + batch_get: MagicMock, +) -> None: + configure_rbac(monkeypatch, enabled=False) + owner = make_subject("owner", legacy_role="owner") + member = make_subject("member") + + result = workspace_member_role_resolver.DeploymentWorkspaceMemberRoleResolver().resolve_many( + "workspace-1", + "actor-1", + [owner, member], + ) + + assert result == { + "owner": (WorkspaceMemberRole(id="owner", name="owner"),), + "member": (WorkspaceMemberRole(id="normal", name="normal"),), + } + batch_get.assert_not_called() + + +def test_rbac_mode_maps_batch_response_without_legacy_fallback( + monkeypatch: pytest.MonkeyPatch, + batch_get: MagicMock, +) -> None: + configure_rbac(monkeypatch, enabled=True) + owner = make_subject("owner", legacy_role="owner") + omitted = make_subject("omitted", legacy_role="admin") + batch_get.return_value = [ + MemberRolesResponse( + account_id=owner.account_id, + roles=[ + RBACRole( + id="workspace.owner", + name="Owner", + type="builtin", + ), + RBACRole( + id="workspace.editor", + name="Editor", + type="builtin", + ), + ], + ) + ] + + result = workspace_member_role_resolver.DeploymentWorkspaceMemberRoleResolver().resolve_many( + "workspace-1", + "actor-1", + [owner, omitted], + ) + + assert result == { + "owner": ( + WorkspaceMemberRole(id="workspace.owner", name="Owner"), + WorkspaceMemberRole(id="workspace.editor", name="Editor"), + ) + } + assert "omitted" not in result + batch_get.assert_called_once_with("workspace-1", "actor-1", ["owner", "omitted"]) + + +def test_rbac_failure_propagates( + monkeypatch: pytest.MonkeyPatch, + batch_get: MagicMock, +) -> None: + configure_rbac(monkeypatch, enabled=True) + batch_get.side_effect = RoleResolutionError("enterprise unavailable") + + with pytest.raises(RoleResolutionError, match="enterprise unavailable"): + workspace_member_role_resolver.DeploymentWorkspaceMemberRoleResolver().resolve_many( + "workspace-1", + "actor-1", + [make_subject("member-1")], + ) + + +def test_empty_member_list_skips_enterprise_call( + monkeypatch: pytest.MonkeyPatch, + batch_get: MagicMock, +) -> None: + configure_rbac(monkeypatch, enabled=True) + + result = workspace_member_role_resolver.DeploymentWorkspaceMemberRoleResolver().resolve_many( + "workspace-1", + "actor-1", + [], + ) + + assert result == {} + batch_get.assert_not_called() + + +class RoleResolutionError(Exception): + pass diff --git a/api/tests/unit_tests/services/test_workspace_service.py b/api/tests/unit_tests/services/test_workspace_service.py new file mode 100644 index 00000000000..9d01cb7b847 --- /dev/null +++ b/api/tests/unit_tests/services/test_workspace_service.py @@ -0,0 +1,108 @@ +from types import SimpleNamespace +from unittest.mock import MagicMock, call, patch + +from enums import CloudPlan, DeploymentEdition +from models.account import Tenant +from services.credit_pool_service import CreditPoolBalance +from services.workspace_service import WorkspaceService + + +def test_get_current_workspace_summary_sandbox_uses_trial_only() -> None: + tenant = Tenant(name="Workspace") + membership = SimpleNamespace(role="owner") + session = MagicMock() + session.scalar.return_value = membership + trial_pool = CreditPoolBalance( + tenant_id=tenant.id, + pool_type="trial", + quota_limit=200, + quota_used=20, + ) + billing_info = { + "enabled": True, + "subscription": {"plan": CloudPlan.SANDBOX}, + } + config = SimpleNamespace(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD) + + with ( + patch("services.workspace_service.dify_config", config), + patch("services.workspace_service.BillingService.get_info", return_value=billing_info) as get_info, + patch("services.credit_pool_service.CreditPoolService.get_pool", return_value=trial_pool) as get_pool, + patch("services.workspace_service.FeatureService.get_features") as get_features, + ): + result = WorkspaceService.get_current_workspace_summary(tenant, "account-1", session=session) + + assert result == { + "id": tenant.id, + "name": tenant.name, + "role": "owner", + "plan": CloudPlan.SANDBOX, + "credits": 180, + } + get_info.assert_called_once_with(tenant.id, exclude_vector_space=True) + get_pool.assert_called_once_with(tenant_id=tenant.id, pool_type="trial", session=session) + get_features.assert_not_called() + + +def test_get_current_workspace_summary_falls_back_from_exhausted_paid_pool() -> None: + tenant = Tenant(name="Workspace") + session = MagicMock() + session.scalar.return_value = SimpleNamespace(role="admin") + paid_pool = CreditPoolBalance( + tenant_id=tenant.id, + pool_type="paid", + quota_limit=500, + quota_used=500, + ) + trial_pool = CreditPoolBalance( + tenant_id=tenant.id, + pool_type="trial", + quota_limit=100, + quota_used=40, + ) + billing_info = { + "enabled": True, + "subscription": {"plan": CloudPlan.TEAM}, + } + config = SimpleNamespace(DEPLOYMENT_EDITION=DeploymentEdition.CLOUD) + + with ( + patch("services.workspace_service.dify_config", config), + patch("services.workspace_service.BillingService.get_info", return_value=billing_info), + patch( + "services.credit_pool_service.CreditPoolService.get_pool", + side_effect=[paid_pool, trial_pool], + ) as get_pool, + ): + result = WorkspaceService.get_current_workspace_summary(tenant, "account-1", session=session) + + assert result["plan"] == CloudPlan.TEAM + assert result["credits"] == 60 + assert get_pool.call_args_list == [ + call(tenant_id=tenant.id, pool_type="paid", session=session), + call(tenant_id=tenant.id, pool_type="trial", session=session), + ] + + +def test_get_current_workspace_summary_non_cloud_skips_billing_and_credits() -> None: + tenant = Tenant(name="Workspace") + session = MagicMock() + session.scalar.return_value = SimpleNamespace(role="editor") + config = SimpleNamespace(DEPLOYMENT_EDITION=DeploymentEdition.COMMUNITY) + + with ( + patch("services.workspace_service.dify_config", config), + patch("services.workspace_service.BillingService.get_info") as get_info, + patch("services.credit_pool_service.CreditPoolService.get_pool") as get_pool, + ): + result = WorkspaceService.get_current_workspace_summary(tenant, "account-1", session=session) + + assert result == { + "id": tenant.id, + "name": tenant.name, + "role": "editor", + "plan": None, + "credits": None, + } + get_info.assert_not_called() + get_pool.assert_not_called() diff --git a/api/tests/unit_tests/services/tools/test_api_tools_manage_service.py b/api/tests/unit_tests/services/tools/test_api_tools_manage_service.py index 83665beb03f..5c4825019f7 100644 --- a/api/tests/unit_tests/services/tools/test_api_tools_manage_service.py +++ b/api/tests/unit_tests/services/tools/test_api_tools_manage_service.py @@ -1,9 +1,11 @@ from unittest.mock import Mock +import pytest + from services.tools.api_tools_manage_service import ApiToolManageService -def test_get_api_tool_provider_remote_schema_uses_ssrf_proxy_get(monkeypatch) -> None: +def test_get_api_tool_provider_remote_schema_uses_ssrf_proxy_get(monkeypatch: pytest.MonkeyPatch) -> None: schema = """ { "openapi": "3.0.0", diff --git a/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py b/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py index 549f50cb370..c6926a310ed 100644 --- a/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py +++ b/api/tests/unit_tests/services/tools/test_builtin_tools_manage_service.py @@ -1,82 +1,133 @@ -from unittest.mock import MagicMock, patch +"""Unit tests for built-in tool management and its persisted credential state.""" + +import json +from types import SimpleNamespace +from unittest.mock import MagicMock import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session +from core.plugin.entities.plugin_daemon import CredentialType +from models.tools import BuiltinToolProvider, ToolOAuthSystemClient, ToolOAuthTenantClient +from services.tools import builtin_tools_manage_service as service_module from services.tools.builtin_tools_manage_service import BuiltinToolManageService -MODULE = "services.tools.builtin_tools_manage_service" + +@pytest.fixture +def repository_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session: + """Bind service-owned sessions to the shared SQLite session's engine.""" + monkeypatch.setattr(service_module, "db", SimpleNamespace(engine=sqlite_session.get_bind())) + return sqlite_session -def _mock_session(mock_session_cls): - """Helper: set up a Session context manager mock and return the inner session.""" - session = MagicMock() - mock_session_cls.return_value.__enter__ = MagicMock(return_value=session) - mock_session_cls.return_value.__exit__ = MagicMock(return_value=False) - return session +def _persist_provider( + session: Session, + *, + credential_id: str = "cred-1", + tenant_id: str = "tenant-1", + user_id: str = "user-1", + provider: str = "google", + name: str = "Google 1", + credentials: dict[str, str] | None = None, + is_default: bool = False, +) -> BuiltinToolProvider: + db_provider = BuiltinToolProvider( + tenant_id=tenant_id, + user_id=user_id, + provider=provider, + name=name, + encrypted_credentials=json.dumps(credentials or {"key": "encrypted"}), + credential_type=CredentialType.API_KEY, + is_default=is_default, + ) + db_provider.id = credential_id + session.add(db_provider) + session.commit() + return db_provider -def _mock_sessionmaker(mock_sm_cls): - """Helper: set up a sessionmaker().begin() context manager mock and return the inner session.""" - session = MagicMock() - mock_sm_cls.return_value.begin.return_value.__enter__ = MagicMock(return_value=session) - mock_sm_cls.return_value.begin.return_value.__exit__ = MagicMock(return_value=False) - return session +def _persist_tenant_oauth_client( + session: Session, + *, + tenant_id: str = "tenant-1", + plugin_id: str = "langgenius/google", + provider: str = "google", + enabled: bool = True, + encrypted_params: str = '{"encrypted": "data"}', +) -> ToolOAuthTenantClient: + client = ToolOAuthTenantClient(tenant_id=tenant_id, plugin_id=plugin_id, provider=provider) + client.enabled = enabled + client.encrypted_oauth_params = encrypted_params + session.add(client) + session.commit() + return client + + +def _persist_system_oauth_client( + session: Session, + *, + plugin_id: str = "langgenius/google", + provider: str = "google", +) -> ToolOAuthSystemClient: + client = ToolOAuthSystemClient(plugin_id=plugin_id, provider=provider, encrypted_oauth_params="enc") + session.add(client) + session.commit() + return client class TestDeleteCustomOauthClientParams: - @patch(f"{MODULE}.sessionmaker") - @patch(f"{MODULE}.db") - def test_deletes_and_returns_success(self, mock_db, mock_sm_cls): - session = _mock_sessionmaker(mock_sm_cls) + def test_deletes_matching_tenant_only(self, repository_session: Session) -> None: + _persist_tenant_oauth_client(repository_session, tenant_id="tenant-1") + _persist_tenant_oauth_client(repository_session, tenant_id="tenant-2") result = BuiltinToolManageService.delete_custom_oauth_client_params("tenant-1", "google") assert result == {"result": "success"} - session.execute.assert_called_once() + repository_session.expire_all() + clients = repository_session.scalars(select(ToolOAuthTenantClient)).all() + assert [client.tenant_id for client in clients] == ["tenant-2"] class TestListBuiltinToolProviderTools: - @patch(f"{MODULE}.ToolLabelManager") - @patch(f"{MODULE}.ToolTransformService") - @patch(f"{MODULE}.ToolManager") - def test_transforms_each_tool(self, mock_manager, mock_transform, mock_labels): - mock_controller = MagicMock() - mock_controller.get_tools.return_value = [MagicMock(), MagicMock()] - mock_manager.get_builtin_provider.return_value = mock_controller - mock_transform.convert_tool_entity_to_api_entity.return_value = MagicMock() + def test_transforms_each_tool(self, monkeypatch: pytest.MonkeyPatch) -> None: + controller = MagicMock() + controller.get_tools.return_value = [MagicMock(), MagicMock()] + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) + convert = MagicMock(return_value=MagicMock()) + monkeypatch.setattr(service_module.ToolTransformService, "convert_tool_entity_to_api_entity", convert) + monkeypatch.setattr(service_module.ToolLabelManager, "get_tool_labels", MagicMock(return_value=[])) result = BuiltinToolManageService.list_builtin_tool_provider_tools("tenant-1", "google") assert len(result) == 2 + assert convert.call_count == 2 - @patch(f"{MODULE}.ToolLabelManager") - @patch(f"{MODULE}.ToolTransformService") - @patch(f"{MODULE}.ToolManager") - def test_empty_tools(self, mock_manager, mock_transform, mock_labels): - mock_controller = MagicMock() - mock_controller.get_tools.return_value = [] - mock_manager.get_builtin_provider.return_value = mock_controller + def test_empty_tools(self, monkeypatch: pytest.MonkeyPatch) -> None: + controller = MagicMock() + controller.get_tools.return_value = list[object]() + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) assert BuiltinToolManageService.list_builtin_tool_provider_tools("t", "p") == [] class TestGetBuiltinToolProviderInfo: - @patch(f"{MODULE}.ToolTransformService") - @patch(f"{MODULE}.BuiltinToolManageService.get_builtin_provider") - @patch(f"{MODULE}.ToolManager") - def test_raises_when_not_found(self, mock_manager, mock_get, mock_transform): - mock_get.return_value = None + def test_raises_when_not_found(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(BuiltinToolManageService, "get_builtin_provider", MagicMock(return_value=None)) + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=MagicMock())) with pytest.raises(ValueError, match="you have not added provider"): BuiltinToolManageService.get_builtin_tool_provider_info("t", "no") - @patch(f"{MODULE}.ToolTransformService") - @patch(f"{MODULE}.BuiltinToolManageService.get_builtin_provider") - @patch(f"{MODULE}.ToolManager") - def test_clears_original_credentials(self, mock_manager, mock_get, mock_transform): - mock_get.return_value = MagicMock() + def test_clears_original_credentials(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(BuiltinToolManageService, "get_builtin_provider", MagicMock(return_value=MagicMock())) + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=MagicMock())) entity = MagicMock() - mock_transform.builtin_provider_to_user_provider.return_value = entity + monkeypatch.setattr( + service_module.ToolTransformService, + "builtin_provider_to_user_provider", + MagicMock(return_value=entity), + ) result = BuiltinToolManageService.get_builtin_tool_provider_info("t", "google") @@ -84,21 +135,26 @@ class TestGetBuiltinToolProviderInfo: class TestListBuiltinProviderCredentialsSchema: - @patch(f"{MODULE}.ToolManager") - def test_returns_schema(self, mock_manager): - mock_manager.get_builtin_provider.return_value.get_credentials_schema_by_type.return_value = [{"f": "k"}] + def test_returns_schema(self, monkeypatch: pytest.MonkeyPatch) -> None: + controller = MagicMock() + controller.get_credentials_schema_by_type.return_value = [{"f": "k"}] + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) - result = BuiltinToolManageService.list_builtin_provider_credentials_schema("g", "api_key", "t") + result = BuiltinToolManageService.list_builtin_provider_credentials_schema("g", CredentialType.API_KEY, "t") assert result == [{"f": "k"}] class TestGetBuiltinToolProviderIcon: - @patch(f"{MODULE}.Path") - @patch(f"{MODULE}.ToolManager") - def test_returns_bytes_and_mime(self, mock_manager, mock_path): - mock_manager.get_hardcoded_provider_icon.return_value = ("/icon.svg", "image/svg+xml") - mock_path.return_value.read_bytes.return_value = b"" + def test_returns_bytes_and_mime(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + service_module.ToolManager, + "get_hardcoded_provider_icon", + MagicMock(return_value=("/icon.svg", "image/svg+xml")), + ) + path = MagicMock() + path.return_value.read_bytes.return_value = b"" + monkeypatch.setattr(service_module, "Path", path) icon, mime = BuiltinToolManageService.get_builtin_tool_provider_icon("google") @@ -107,164 +163,140 @@ class TestGetBuiltinToolProviderIcon: class TestIsOauthSystemClientExists: - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_true_when_exists(self, mock_db, mock_session_cls): - session = _mock_session(mock_session_cls) - session.scalar.return_value = MagicMock() + def test_true_when_exists(self, repository_session: Session) -> None: + _persist_system_oauth_client(repository_session) assert BuiltinToolManageService.is_oauth_system_client_exists("google") is True - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_false_when_missing(self, mock_db, mock_session_cls): - session = _mock_session(mock_session_cls) - session.scalar.return_value = None + def test_false_when_missing(self, repository_session: Session) -> None: + _persist_system_oauth_client(repository_session, plugin_id="langgenius/slack", provider="slack") assert BuiltinToolManageService.is_oauth_system_client_exists("google") is False class TestIsOauthCustomClientEnabled: - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_true_when_enabled(self, mock_db, mock_session_cls): - session = _mock_session(mock_session_cls) - session.scalar.return_value = MagicMock(enabled=True) + def test_true_when_enabled(self, repository_session: Session) -> None: + _persist_tenant_oauth_client(repository_session) - assert BuiltinToolManageService.is_oauth_custom_client_enabled("t", "g") is True + assert BuiltinToolManageService.is_oauth_custom_client_enabled("tenant-1", "google") is True - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_false_when_none(self, mock_db, mock_session_cls): - session = _mock_session(mock_session_cls) - session.scalar.return_value = None + def test_false_when_disabled_or_other_tenant(self, repository_session: Session) -> None: + _persist_tenant_oauth_client(repository_session, tenant_id="tenant-1", enabled=False) + _persist_tenant_oauth_client(repository_session, tenant_id="tenant-2", enabled=True) - assert BuiltinToolManageService.is_oauth_custom_client_enabled("t", "g") is False + assert BuiltinToolManageService.is_oauth_custom_client_enabled("tenant-1", "google") is False class TestDeleteBuiltinToolProvider: - @patch(f"{MODULE}.BuiltinToolManageService.create_tool_encrypter") - @patch(f"{MODULE}.ToolManager") - @patch(f"{MODULE}.sessionmaker") - @patch(f"{MODULE}.db") - def test_raises_when_not_found(self, mock_db, mock_sm_cls, mock_tm, mock_enc): - session = _mock_sessionmaker(mock_sm_cls) - session.scalar.return_value = None - + def test_raises_when_not_found(self, repository_session: Session) -> None: with pytest.raises(ValueError, match="you have not added provider"): - BuiltinToolManageService.delete_builtin_tool_provider("t", "p", "id") + BuiltinToolManageService.delete_builtin_tool_provider("tenant-1", "google", "missing") - @patch(f"{MODULE}.BuiltinToolManageService.create_tool_encrypter") - @patch(f"{MODULE}.ToolManager") - @patch(f"{MODULE}.sessionmaker") - @patch(f"{MODULE}.db") - def test_deletes_provider_and_clears_cache(self, mock_db, mock_sm_cls, mock_tm, mock_enc): - session = _mock_sessionmaker(mock_sm_cls) - db_provider = MagicMock() - session.scalar.return_value = db_provider - mock_cache = MagicMock() - mock_enc.return_value = (MagicMock(), mock_cache) + def test_deletes_provider_and_clears_cache( + self, + monkeypatch: pytest.MonkeyPatch, + repository_session: Session, + ) -> None: + _persist_provider(repository_session, credential_id="cred-1") + _persist_provider(repository_session, credential_id="cred-other", tenant_id="tenant-2") + cache = MagicMock() + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=MagicMock())) + monkeypatch.setattr( + BuiltinToolManageService, + "create_tool_encrypter", + MagicMock(return_value=(MagicMock(), cache)), + ) - result = BuiltinToolManageService.delete_builtin_tool_provider("t", "p", "c") + result = BuiltinToolManageService.delete_builtin_tool_provider("tenant-1", "google", "cred-1") assert result == {"result": "success"} - session.delete.assert_called_once_with(db_provider) - mock_cache.delete.assert_called_once() + cache.delete.assert_called_once() + repository_session.expire_all() + assert repository_session.get(BuiltinToolProvider, "cred-1") is None + assert repository_session.get(BuiltinToolProvider, "cred-other") is not None class TestSetDefaultProvider: - @patch(f"{MODULE}.sessionmaker") - @patch(f"{MODULE}.db") - def test_raises_when_not_found(self, mock_db, mock_sm_cls): - session = _mock_sessionmaker(mock_sm_cls) - session.scalar.return_value = None - + def test_raises_when_not_found(self, repository_session: Session) -> None: with pytest.raises(ValueError, match="provider not found"): - BuiltinToolManageService.set_default_provider("t", "p", "id") + BuiltinToolManageService.set_default_provider("tenant-1", "google", "missing") - @patch(f"{MODULE}.sessionmaker") - @patch(f"{MODULE}.db") - def test_sets_default_and_clears_old(self, mock_db, mock_sm_cls): - session = _mock_sessionmaker(mock_sm_cls) - target = MagicMock() - session.scalar.return_value = target + def test_sets_target_and_clears_only_same_tenant_defaults(self, repository_session: Session) -> None: + _persist_provider( + repository_session, + credential_id="target", + user_id="user-2", + name="Google target", + ) + _persist_provider(repository_session, credential_id="old", name="Google old", is_default=True) + _persist_provider( + repository_session, + credential_id="other-tenant", + tenant_id="tenant-2", + name="Other tenant", + is_default=True, + ) - result = BuiltinToolManageService.set_default_provider("t", "p", "id") + result = BuiltinToolManageService.set_default_provider("tenant-1", "google", "target") assert result == {"result": "success"} + repository_session.expire_all() + target = repository_session.get(BuiltinToolProvider, "target") + old = repository_session.get(BuiltinToolProvider, "old") + other_tenant = repository_session.get(BuiltinToolProvider, "other-tenant") + assert target is not None assert target.is_default is True - - @patch(f"{MODULE}.sessionmaker") - @patch(f"{MODULE}.db") - def test_clear_default_is_tenant_scoped_not_user_scoped(self, mock_db, mock_sm_cls): - # Regression: clearing prior defaults must NOT filter by user_id, otherwise - # two workspace members can each leave their own credential as default at - # the same time (the default flag is tenant-scoped, not per-user). - session = _mock_sessionmaker(mock_sm_cls) - session.scalar.return_value = MagicMock() - - BuiltinToolManageService.set_default_provider("tenant-1", "google", "cred-id") - - session.execute.assert_called_once() - update_stmt = session.execute.call_args.args[0] - compiled = str(update_stmt.compile(compile_kwargs={"literal_binds": True})) - assert "user_id" not in compiled - assert "tenant_id" in compiled - assert "provider" in compiled + assert old is not None + assert old.is_default is False + assert other_tenant is not None + assert other_tenant.is_default is True class TestUpdateBuiltinToolProvider: - @patch(f"{MODULE}.sessionmaker") - @patch(f"{MODULE}.db") - def test_raises_when_provider_not_exists(self, mock_db, mock_sm_cls): - session = _mock_sessionmaker(mock_sm_cls) - session.scalar.return_value = None - + def test_raises_when_provider_not_exists(self, repository_session: Session) -> None: with pytest.raises(ValueError, match="you have not added provider"): - BuiltinToolManageService.update_builtin_tool_provider("u", "t", "p", "c") + BuiltinToolManageService.update_builtin_tool_provider("u", "tenant-1", "google", "missing") - @patch(f"{MODULE}.BuiltinToolManageService.create_tool_encrypter") - @patch(f"{MODULE}.CredentialType") - @patch(f"{MODULE}.ToolManager") - @patch(f"{MODULE}.sessionmaker") - @patch(f"{MODULE}.db") - def test_updates_credentials_and_commits(self, mock_db, mock_sm_cls, mock_tm, mock_cred_type, mock_enc): - session = _mock_sessionmaker(mock_sm_cls) - db_provider = MagicMock(credential_type="api_key", credentials="{}") - session.scalar.return_value = db_provider + def test_updates_persisted_credentials_and_clears_cache( + self, + monkeypatch: pytest.MonkeyPatch, + repository_session: Session, + ) -> None: + _persist_provider(repository_session, credentials={"key": "old"}) + controller = MagicMock(need_credentials=True) + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) + encrypter = MagicMock() + encrypter.decrypt.return_value = {"key": "old"} + encrypter.encrypt.return_value = {"key": "new"} + cache = MagicMock() + monkeypatch.setattr( + BuiltinToolManageService, + "create_tool_encrypter", + MagicMock(return_value=(encrypter, cache)), + ) - mock_cred_instance = MagicMock() - mock_cred_instance.is_editable.return_value = True - mock_cred_instance.is_validate_allowed.return_value = False - mock_cred_type.of.return_value = mock_cred_instance - - mock_controller = MagicMock(need_credentials=True) - mock_tm.get_builtin_provider.return_value = mock_controller - - mock_encrypter = MagicMock() - mock_encrypter.decrypt.return_value = {"key": "old"} - mock_encrypter.encrypt.return_value = {"key": "new"} - mock_cache = MagicMock() - mock_enc.return_value = (mock_encrypter, mock_cache) - - result = BuiltinToolManageService.update_builtin_tool_provider("u", "t", "p", "c", credentials={"key": "val"}) + result = BuiltinToolManageService.update_builtin_tool_provider( + "u", "tenant-1", "google", "cred-1", credentials={"key": "value"} + ) assert result == {"result": "success"} - mock_cache.delete.assert_called_once() + controller.validate_credentials.assert_called_once_with("u", {"key": "value"}) + cache.delete.assert_called_once() + repository_session.expire_all() + provider = repository_session.get(BuiltinToolProvider, "cred-1") + assert provider is not None + assert provider.credentials == {"key": "new"} class TestGetOauthClientSchema: - @patch(f"{MODULE}.BuiltinToolManageService.get_custom_oauth_client_params", return_value={}) - @patch(f"{MODULE}.BuiltinToolManageService.is_oauth_system_client_exists", return_value=False) - @patch(f"{MODULE}.BuiltinToolManageService.is_oauth_custom_client_enabled", return_value=True) - @patch(f"{MODULE}.dify_config") - @patch(f"{MODULE}.PluginService") - @patch(f"{MODULE}.ToolManager") - def test_returns_schema_dict(self, mock_tm, mock_plugin, mock_config, mock_enabled, mock_sys, mock_params): - mock_config.CONSOLE_API_URL = "https://api.example.com" - mock_controller = MagicMock() - mock_controller.get_oauth_client_schema.return_value = [] - mock_tm.get_builtin_provider.return_value = mock_controller + def test_returns_schema_dict(self, monkeypatch: pytest.MonkeyPatch) -> None: + controller = MagicMock() + controller.get_oauth_client_schema.return_value = list[object]() + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) + monkeypatch.setattr(BuiltinToolManageService, "is_oauth_custom_client_enabled", MagicMock(return_value=True)) + monkeypatch.setattr(BuiltinToolManageService, "is_oauth_system_client_exists", MagicMock(return_value=False)) + monkeypatch.setattr(BuiltinToolManageService, "get_custom_oauth_client_params", MagicMock(return_value={})) + monkeypatch.setattr(service_module.dify_config, "CONSOLE_API_URL", "https://api.example.com") result = BuiltinToolManageService.get_builtin_tool_provider_oauth_client_schema("t", "google") @@ -274,87 +306,91 @@ class TestGetOauthClientSchema: class TestGetOauthClient: - @patch(f"{MODULE}.PluginService") - @patch(f"{MODULE}.create_provider_encrypter") - @patch(f"{MODULE}.ToolManager") - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_returns_user_client_params_when_exists( - self, mock_db, mock_session_cls, mock_tm, mock_create_enc, mock_plugin - ): - session = _mock_session(mock_session_cls) - mock_controller = MagicMock() - mock_controller.get_oauth_client_schema.return_value = [] - mock_tm.get_builtin_provider.return_value = mock_controller + def test_returns_tenant_client_params_when_exists( + self, + monkeypatch: pytest.MonkeyPatch, + repository_session: Session, + ) -> None: + _persist_tenant_oauth_client(repository_session) + controller = MagicMock() + controller.get_oauth_client_schema.return_value = list[object]() + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) + encrypter = MagicMock() + encrypter.decrypt.return_value = {"client_id": "id", "client_secret": "secret"} + monkeypatch.setattr( + service_module, "create_provider_encrypter", MagicMock(return_value=(encrypter, MagicMock())) + ) - mock_encrypter = MagicMock() - mock_encrypter.decrypt.return_value = {"client_id": "id", "client_secret": "secret"} - mock_create_enc.return_value = (mock_encrypter, MagicMock()) - - user_client = MagicMock(oauth_params='{"encrypted": "data"}') - session.scalar.return_value = user_client - - result = BuiltinToolManageService.get_oauth_client("t", "google") + result = BuiltinToolManageService.get_oauth_client("tenant-1", "google") assert result == {"client_id": "id", "client_secret": "secret"} + encrypter.decrypt.assert_called_once_with({"encrypted": "data"}) - @patch(f"{MODULE}.decrypt_system_params", return_value={"sys_key": "sys_val"}) - @patch(f"{MODULE}.PluginService") - @patch(f"{MODULE}.create_provider_encrypter") - @patch(f"{MODULE}.ToolManager") - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") def test_falls_back_to_system_client( - self, mock_db, mock_session_cls, mock_tm, mock_create_enc, mock_plugin, mock_decrypt - ): - session = _mock_session(mock_session_cls) - mock_controller = MagicMock() - mock_controller.get_oauth_client_schema.return_value = [] - mock_tm.get_builtin_provider.return_value = mock_controller + self, + monkeypatch: pytest.MonkeyPatch, + repository_session: Session, + ) -> None: + _persist_system_oauth_client(repository_session) + controller = MagicMock() + controller.get_oauth_client_schema.return_value = list[object]() + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) + monkeypatch.setattr( + service_module, "create_provider_encrypter", MagicMock(return_value=(MagicMock(), MagicMock())) + ) + decrypt = MagicMock(return_value={"sys_key": "sys_val"}) + monkeypatch.setattr(service_module, "decrypt_system_params", decrypt) - mock_create_enc.return_value = (MagicMock(), MagicMock()) - - system_client = MagicMock(encrypted_oauth_params="enc") - session.scalar.side_effect = [None, system_client] - - result = BuiltinToolManageService.get_oauth_client("t", "google") + result = BuiltinToolManageService.get_oauth_client("tenant-1", "google") assert result == {"sys_key": "sys_val"} + decrypt.assert_called_once_with("enc") class TestSaveCustomOauthClientParams: - def test_returns_early_when_no_params(self): + def test_returns_early_when_no_params(self) -> None: result = BuiltinToolManageService.save_custom_oauth_client_params("t", "p") assert result == {"result": "success"} - @patch(f"{MODULE}.ToolManager") - def test_raises_when_provider_not_found(self, mock_tm): - mock_tm.get_builtin_provider.return_value = None + def test_raises_when_provider_not_found(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=None)) with pytest.raises((ValueError, Exception), match="not found|Provider"): BuiltinToolManageService.save_custom_oauth_client_params("t", "p", enable_oauth_custom_client=True) class TestGetCustomOauthClientParams: - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_returns_empty_when_none(self, mock_db, mock_session_cls): - session = _mock_session(mock_session_cls) - session.scalar.return_value = None + def test_returns_empty_when_none(self, repository_session: Session) -> None: + _persist_tenant_oauth_client(repository_session, tenant_id="other-tenant") - result = BuiltinToolManageService.get_custom_oauth_client_params("t", "p") + result = BuiltinToolManageService.get_custom_oauth_client_params("tenant-1", "google") assert result == {} class TestGetBuiltinToolProviderCredentialInfo: - @patch(f"{MODULE}.BuiltinToolManageService.is_oauth_custom_client_enabled", return_value=False) - @patch(f"{MODULE}.BuiltinToolManageService.get_builtin_tool_provider_credentials", return_value=[]) - @patch(f"{MODULE}.ToolManager") - def test_returns_credential_info(self, mock_tm, mock_creds, mock_oauth): - mock_tm.get_builtin_provider.return_value.get_supported_credential_types.return_value = ["api-key"] + def test_returns_credential_info( + self, + monkeypatch: pytest.MonkeyPatch, + repository_session: Session, + ) -> None: + controller = MagicMock() + controller.get_supported_credential_types.return_value = ["api-key"] + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) + monkeypatch.setattr( + BuiltinToolManageService, + "get_builtin_tool_provider_credentials", + MagicMock(return_value=[]), + ) + monkeypatch.setattr( + BuiltinToolManageService, + "is_oauth_custom_client_enabled", + MagicMock(return_value=False), + ) - result = BuiltinToolManageService.get_builtin_tool_provider_credential_info("t", "google", session=MagicMock()) + result = BuiltinToolManageService.get_builtin_tool_provider_credential_info( + "tenant-1", "google", session=repository_session + ) assert result.credentials == [] assert result.supported_credential_types == ["api-key"] @@ -362,113 +398,84 @@ class TestGetBuiltinToolProviderCredentialInfo: class TestGetBuiltinToolProviderCredentials: - @patch(f"{MODULE}.db") - def test_returns_empty_when_no_providers(self, mock_db): - mock_db.session.no_autoflush.__enter__ = MagicMock(return_value=None) - mock_db.session.no_autoflush.__exit__ = MagicMock(return_value=False) - mock_db.session.scalars.return_value.all.return_value = [] + def test_returns_empty_when_no_providers(self, repository_session: Session) -> None: + _persist_provider(repository_session, credential_id="other", tenant_id="other-tenant") - result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google", session=mock_db.session) + result = BuiltinToolManageService.get_builtin_tool_provider_credentials( + "tenant-1", "google", session=repository_session + ) assert result == [] - @patch(f"{MODULE}.ToolTransformService") - @patch(f"{MODULE}.BuiltinToolManageService.create_tool_encrypter") - @patch(f"{MODULE}.ToolManager") - @patch(f"{MODULE}.db") - def test_returns_credential_entities(self, mock_db, mock_tm, mock_enc, mock_transform): - mock_db.session.no_autoflush.__enter__ = MagicMock(return_value=None) - mock_db.session.no_autoflush.__exit__ = MagicMock(return_value=False) - - provider = MagicMock(provider="google", is_default=False) - mock_db.session.scalars.return_value.all.return_value = [provider] - - mock_encrypter = MagicMock() - mock_encrypter.decrypt.return_value = {"key": "decrypted"} - mock_encrypter.mask_plugin_credentials.return_value = {"key": "***"} - mock_enc.return_value = (mock_encrypter, MagicMock()) - + def test_returns_tenant_scoped_credential_entities( + self, + monkeypatch: pytest.MonkeyPatch, + repository_session: Session, + ) -> None: + _persist_provider(repository_session, is_default=False) + _persist_provider(repository_session, credential_id="other", tenant_id="other-tenant") + controller = MagicMock() + monkeypatch.setattr(service_module.ToolManager, "get_builtin_provider", MagicMock(return_value=controller)) + encrypter = MagicMock() + encrypter.decrypt.return_value = {"key": "decrypted"} + encrypter.mask_plugin_credentials.return_value = {"key": "***"} + monkeypatch.setattr( + BuiltinToolManageService, + "create_tool_encrypter", + MagicMock(return_value=(encrypter, MagicMock())), + ) credential_entity = MagicMock() - mock_transform.convert_builtin_provider_to_credential_entity.return_value = credential_entity + convert = MagicMock(return_value=credential_entity) + monkeypatch.setattr( + service_module.ToolTransformService, + "convert_builtin_provider_to_credential_entity", + convert, + ) - result = BuiltinToolManageService.get_builtin_tool_provider_credentials("t", "google", session=mock_db.session) + result = BuiltinToolManageService.get_builtin_tool_provider_credentials( + "tenant-1", "google", session=repository_session + ) - assert len(result) == 1 - assert result[0] is credential_entity - assert provider.is_default is True + assert result == [credential_entity] + converted_provider = convert.call_args.kwargs["provider"] + assert isinstance(converted_provider, BuiltinToolProvider) + assert converted_provider.tenant_id == "tenant-1" + assert converted_provider.is_default is True class TestGetBuiltinProvider: - @patch(f"{MODULE}.ToolProviderID") - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_returns_none_when_not_found(self, mock_db, mock_session_cls, mock_prov_id): - session = _mock_session(mock_session_cls) - mock_prov_id.return_value.provider_name = "google" - mock_prov_id.return_value.organization = "langgenius" - session.scalar.return_value = None + def test_returns_none_when_not_found(self, repository_session: Session) -> None: + assert BuiltinToolManageService.get_builtin_provider("google", "tenant-1") is None - result = BuiltinToolManageService.get_builtin_provider("google", "t") + def test_returns_langgenius_provider_for_matching_tenant(self, repository_session: Session) -> None: + _persist_provider(repository_session, tenant_id="tenant-1", provider="google") + _persist_provider(repository_session, credential_id="other", tenant_id="tenant-2", provider="google") - assert result is None + result = BuiltinToolManageService.get_builtin_provider("google", "tenant-1") - @patch(f"{MODULE}.ToolProviderID") - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_returns_provider_for_langgenius_org(self, mock_db, mock_session_cls, mock_prov_id): - session = _mock_session(mock_session_cls) - mock_prov_id.return_value.provider_name = "google" - mock_prov_id.return_value.organization = "langgenius" - db_provider = MagicMock(provider="google") - mock_prov_id_result = MagicMock() - mock_prov_id_result.to_string.return_value = "langgenius/google/google" + assert result is not None + assert result.id == "cred-1" + assert result.provider == "langgenius/google/google" - def prov_id_side_effect(name): - m = MagicMock() - m.provider_name = "google" - m.organization = "langgenius" - m.to_string.return_value = "langgenius/google/google" - m.plugin_id = "langgenius/google" - return m + def test_returns_non_langgenius_provider(self, repository_session: Session) -> None: + full_provider = "third-party/custom/custom-tool" + _persist_provider(repository_session, provider=full_provider) - mock_prov_id.side_effect = prov_id_side_effect - session.scalar.return_value = db_provider + result = BuiltinToolManageService.get_builtin_provider(full_provider, "tenant-1") - result = BuiltinToolManageService.get_builtin_provider("google", "t") + assert result is not None + assert result.id == "cred-1" + assert result.provider == full_provider - assert result is db_provider + def test_falls_back_on_provider_id_parse_exception( + self, + monkeypatch: pytest.MonkeyPatch, + repository_session: Session, + ) -> None: + _persist_provider(repository_session, provider="old-provider") + monkeypatch.setattr(service_module, "ToolProviderID", MagicMock(side_effect=Exception("parse error"))) - @patch(f"{MODULE}.ToolProviderID") - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_returns_provider_for_non_langgenius_org(self, mock_db, mock_session_cls, mock_prov_id): - session = _mock_session(mock_session_cls) + result = BuiltinToolManageService.get_builtin_provider("old-provider", "tenant-1") - def prov_id_side_effect(name): - m = MagicMock() - m.provider_name = "custom-tool" - m.organization = "third-party" - m.to_string.return_value = "third-party/custom/custom-tool" - m.plugin_id = "third-party/custom" - return m - - mock_prov_id.side_effect = prov_id_side_effect - db_provider = MagicMock(provider="third-party/custom/custom-tool") - session.scalar.return_value = db_provider - - result = BuiltinToolManageService.get_builtin_provider("third-party/custom/custom-tool", "t") - - assert result is db_provider - - @patch(f"{MODULE}.ToolProviderID") - @patch(f"{MODULE}.Session") - @patch(f"{MODULE}.db") - def test_falls_back_on_exception(self, mock_db, mock_session_cls, mock_prov_id): - session = _mock_session(mock_session_cls) - mock_prov_id.side_effect = Exception("parse error") - fallback = MagicMock() - session.scalar.return_value = fallback - - result = BuiltinToolManageService.get_builtin_provider("old-provider", "t") - - assert result is fallback + assert result is not None + assert result.id == "cred-1" diff --git a/api/tests/unit_tests/services/tools/test_mcp_tools_transform.py b/api/tests/unit_tests/services/tools/test_mcp_tools_transform.py index 9537d207f01..aa5f28c0535 100644 --- a/api/tests/unit_tests/services/tools/test_mcp_tools_transform.py +++ b/api/tests/unit_tests/services/tools/test_mcp_tools_transform.py @@ -3,6 +3,7 @@ from unittest.mock import Mock import pytest +from pytest_mock import MockerFixture from core.mcp.types import Tool as MCPTool from core.tools.entities.api_entities import ToolApiEntity, ToolProviderApiEntity @@ -21,33 +22,55 @@ def mock_user(): @pytest.fixture -def mock_provider(mock_user): +def mock_provider(mock_user, mocker: MockerFixture): """Provides a mock MCPToolProvider with a loaded user.""" - provider = Mock(spec=MCPToolProvider) - provider.load_user.return_value = mock_user + provider = MCPToolProvider( + name="Test Provider", + server_identifier="test-provider", + server_url="https://example.com", + server_url_hash="hash", + icon="icon", + tenant_id="tenant-id", + user_id="user-id", + ) + mocker.patch.object(provider, "load_user", return_value=mock_user) return provider @pytest.fixture -def mock_provider_no_user(): +def mock_provider_no_user(mocker: MockerFixture): """Provides a mock MCPToolProvider with no user.""" - provider = Mock(spec=MCPToolProvider) - provider.load_user.return_value = None + provider = MCPToolProvider( + name="Test Provider", + server_identifier="test-provider", + server_url="https://example.com", + server_url_hash="hash", + icon="icon", + tenant_id="tenant-id", + user_id="user-id", + ) + mocker.patch.object(provider, "load_user", return_value=None) return provider @pytest.fixture -def mock_provider_full(mock_user): +def mock_provider_full(mock_user, mocker: MockerFixture): """Provides a fully configured mock MCPToolProvider for detailed tests.""" - provider = Mock(spec=MCPToolProvider) + provider = MCPToolProvider( + name="Test MCP Provider", + server_identifier="server-identifier-456", + server_url="https://example.com", + server_url_hash="hash", + icon="icon", + tenant_id="tenant-id", + user_id="user-id", + authed=True, + timeout=30, + sse_read_timeout=300, + ) provider.id = "provider-id-123" - provider.server_identifier = "server-identifier-456" - provider.name = "Test MCP Provider" provider.provider_icon = "icon.png" - provider.authed = True provider.masked_server_url = "https://*****.com/mcp" - provider.timeout = 30 - provider.sse_read_timeout = 300 provider.masked_headers = {"Authorization": "Bearer *****"} provider.decrypted_headers = {"Authorization": "Bearer secret-token"} @@ -56,7 +79,7 @@ def mock_provider_full(mock_user): mock_updated_at.timestamp.return_value = 1234567890 provider.updated_at = mock_updated_at - provider.load_user.return_value = mock_user + mocker.patch.object(provider, "load_user", return_value=mock_user) return provider @@ -306,7 +329,7 @@ class TestMCPToolTransform: assert result[0].type == ToolParameter.ToolParameterType.STRING assert result[0].input_schema is None - def test_mcp_provider_to_user_provider_for_list(self, mock_provider_full): + def test_mcp_provider_to_user_provider_for_list(self, mock_provider_full, mocker: MockerFixture): """Test mcp_provider_to_user_provider with for_list=True.""" # Set tools data with null description mock_provider_full.tools = '[{"name": "tool1", "description": null, "inputSchema": {}}]' @@ -328,7 +351,7 @@ class TestMCPToolTransform: "label": I18nObject(en_US="Test MCP Provider", zh_Hans="Test MCP Provider"), "masked_credentials": {}, } - mock_provider_full.to_entity.return_value = mock_entity + mocker.patch.object(mock_provider_full, "to_entity", return_value=mock_entity) # Call the method with for_list=True result = ToolTransformService.mcp_provider_to_user_provider(mock_provider_full, for_list=True) @@ -343,7 +366,7 @@ class TestMCPToolTransform: assert len(result.tools) == 1 assert result.tools[0].description.en_US == "" # Should handle None description - def test_mcp_provider_to_user_provider_not_for_list(self, mock_provider_full): + def test_mcp_provider_to_user_provider_not_for_list(self, mock_provider_full, mocker: MockerFixture): """Test mcp_provider_to_user_provider with for_list=False.""" # Set tools data with description mock_provider_full.tools = '[{"name": "tool1", "description": "Tool description", "inputSchema": {}}]' @@ -367,7 +390,7 @@ class TestMCPToolTransform: "label": I18nObject(en_US="Test MCP Provider", zh_Hans="Test MCP Provider"), "masked_credentials": {}, } - mock_provider_full.to_entity.return_value = mock_entity + mocker.patch.object(mock_provider_full, "to_entity", return_value=mock_entity) # Call the method with for_list=False result = ToolTransformService.mcp_provider_to_user_provider(mock_provider_full, for_list=False) diff --git a/api/tests/unit_tests/services/workflow/test_draft_var_loader_simple.py b/api/tests/unit_tests/services/workflow/test_draft_var_loader_simple.py index fb5cf7bc6e5..2fe76f7becc 100644 --- a/api/tests/unit_tests/services/workflow/test_draft_var_loader_simple.py +++ b/api/tests/unit_tests/services/workflow/test_draft_var_loader_simple.py @@ -1,32 +1,73 @@ """Simplified unit tests for DraftVarLoader focusing on core functionality.""" import json +from datetime import datetime from unittest.mock import Mock, patch import pytest from sqlalchemy import Engine +from sqlalchemy.orm import Session from core.workflow.file_reference import build_file_reference +from extensions.storage.storage_type import StorageType from graphon.file import File, FileTransferMethod, FileType -from graphon.variables.segments import ObjectSegment, StringSegment +from graphon.variables.segments import StringSegment from graphon.variables.types import SegmentType +from models.enums import CreatorUserRole from models.model import UploadFile from models.workflow import WorkflowDraftVariable, WorkflowDraftVariableFile from services.workflow_draft_variable_service import DraftVarLoader +def _persist_offloaded_variable( + sqlite_session: Session, + *, + node_id: str, + name: str, +) -> WorkflowDraftVariable: + upload_file = UploadFile( + tenant_id="test-tenant-id", + storage_type=StorageType.LOCAL, + key=f"storage/key/{name}.txt", + name=f"{name}.txt", + size=10, + extension=".txt", + mime_type="text/plain", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="test-user-id", + created_at=datetime(2025, 1, 1), + used=True, + ) + variable_file = WorkflowDraftVariableFile( + tenant_id="test-tenant-id", + app_id="test-app-id", + user_id="test-user-id", + upload_file_id=upload_file.id, + size=10, + length=None, + value_type=SegmentType.STRING, + ) + draft_variable = WorkflowDraftVariable.new_node_variable( + app_id="test-app-id", + user_id="test-user-id", + node_id=node_id, + name=name, + value=StringSegment(value="truncated"), + node_execution_id=f"execution-{node_id}", + file_id=variable_file.id, + ) + sqlite_session.add_all([upload_file, variable_file, draft_variable]) + return draft_variable + + class TestDraftVarLoaderSimple: """Simplified unit tests for DraftVarLoader core methods.""" @pytest.fixture - def mock_engine(self) -> Engine: - return Mock(spec=Engine) - - @pytest.fixture - def draft_var_loader(self, mock_engine): + def draft_var_loader(self, sqlite_engine: Engine): """Create DraftVarLoader instance for testing.""" return DraftVarLoader( - engine=mock_engine, + engine=sqlite_engine, app_id="test-app-id", tenant_id="test-tenant-id", user_id="test-user-id", @@ -36,28 +77,45 @@ class TestDraftVarLoaderSimple: def test_load_offloaded_variable_object_type_unit(self, draft_var_loader): """Test _load_offloaded_variable with object type - isolated unit test.""" # Create mock objects - upload_file = Mock(spec=UploadFile) - upload_file.key = "storage/key/test.json" + upload_file = UploadFile( + tenant_id="tenant-id", + storage_type="opendal", + key="storage/key/test.json", + name="test.txt", + size=0, + extension="txt", + mime_type="text/plain", + created_by_role="account", + created_by="account-id", + created_at=datetime.now(), + used=False, + ) - variable_file = Mock(spec=WorkflowDraftVariableFile) - variable_file.value_type = SegmentType.OBJECT + variable_file = WorkflowDraftVariableFile( + tenant_id="tenant-id", + app_id="app-id", + user_id="user-id", + upload_file_id="upload-file-id", + size=0, + length=0, + value_type=SegmentType.OBJECT, + ) variable_file.upload_file = upload_file - draft_var = Mock(spec=WorkflowDraftVariable) - draft_var.id = "draft-var-id" - draft_var.node_id = "test-node-id" - draft_var.name = "test_object" - draft_var.description = "test description" - draft_var.get_selector.return_value = ["test-node-id", "test_object"] - draft_var.variable_file = variable_file + draft_var = WorkflowDraftVariable( + id="draft-var-id", + node_id="test-node-id", + name="test_object", + description="test description", + selector=json.dumps(["test-node-id", "test_object"]), + variable_file=variable_file, + ) test_object = {"key1": "value1", "key2": 42} test_json_content = json.dumps(test_object, ensure_ascii=False, separators=(",", ":")) with patch("services.workflow_draft_variable_service.storage") as mock_storage: mock_storage.load.return_value = test_json_content.encode() - mock_segment = ObjectSegment(value=test_object) - draft_var.build_segment_from_serialized_value.return_value = mock_segment # Execute the method selector_tuple, variable = draft_var_loader._load_offloaded_variable(draft_var) @@ -71,23 +129,32 @@ class TestDraftVarLoaderSimple: # Verify method calls mock_storage.load.assert_called_once_with("storage/key/test.json") - draft_var.build_segment_from_serialized_value.assert_called_once_with(SegmentType.OBJECT, test_object) def test_load_offloaded_variable_missing_variable_file_unit(self, draft_var_loader): """Test that assertion error is raised when variable_file is None.""" - draft_var = Mock(spec=WorkflowDraftVariable) - draft_var.variable_file = None + draft_var = WorkflowDraftVariable( + variable_file=None, + ) with pytest.raises(AssertionError): draft_var_loader._load_offloaded_variable(draft_var) def test_load_offloaded_variable_missing_upload_file_unit(self, draft_var_loader): """Test that assertion error is raised when upload_file is None.""" - variable_file = Mock(spec=WorkflowDraftVariableFile) + variable_file = WorkflowDraftVariableFile( + tenant_id="tenant-id", + app_id="app-id", + user_id="user-id", + upload_file_id="upload-file-id", + size=0, + length=0, + value_type="file", + ) variable_file.upload_file = None - draft_var = Mock(spec=WorkflowDraftVariable) - draft_var.variable_file = variable_file + draft_var = WorkflowDraftVariable( + variable_file=variable_file, + ) with pytest.raises(AssertionError): draft_var_loader._load_offloaded_variable(draft_var) @@ -106,31 +173,45 @@ class TestDraftVarLoaderSimple: def test_load_offloaded_variable_array_type_unit(self, draft_var_loader): """Test _load_offloaded_variable with array type - isolated unit test.""" # Create mock objects - upload_file = Mock(spec=UploadFile) - upload_file.key = "storage/key/test_array.json" + upload_file = UploadFile( + tenant_id="tenant-id", + storage_type="opendal", + key="storage/key/test_array.json", + name="test.txt", + size=0, + extension="txt", + mime_type="text/plain", + created_by_role="account", + created_by="account-id", + created_at=datetime.now(), + used=False, + ) - variable_file = Mock(spec=WorkflowDraftVariableFile) - variable_file.value_type = SegmentType.ARRAY_ANY + variable_file = WorkflowDraftVariableFile( + tenant_id="tenant-id", + app_id="app-id", + user_id="user-id", + upload_file_id="upload-file-id", + size=0, + length=0, + value_type=SegmentType.ARRAY_ANY, + ) variable_file.upload_file = upload_file - draft_var = Mock(spec=WorkflowDraftVariable) - draft_var.id = "draft-var-id" - draft_var.node_id = "test-node-id" - draft_var.name = "test_array" - draft_var.description = "test array description" - draft_var.get_selector.return_value = ["test-node-id", "test_array"] - draft_var.variable_file = variable_file + draft_var = WorkflowDraftVariable( + id="draft-var-id", + node_id="test-node-id", + name="test_array", + description="test array description", + selector=json.dumps(["test-node-id", "test_array"]), + variable_file=variable_file, + ) - test_array = ["item1", "item2", "item3"] + test_array = ["item1", 2, True] test_json_content = json.dumps(test_array) with patch("services.workflow_draft_variable_service.storage") as mock_storage: mock_storage.load.return_value = test_json_content.encode() - from graphon.variables.segments import ArrayAnySegment - - mock_segment = ArrayAnySegment(value=test_array) - draft_var.build_segment_from_serialized_value.return_value = mock_segment - # Execute the method selector_tuple, variable = draft_var_loader._load_offloaded_variable(draft_var) @@ -142,14 +223,31 @@ class TestDraftVarLoaderSimple: # Verify method calls mock_storage.load.assert_called_once_with("storage/key/test_array.json") - draft_var.build_segment_from_serialized_value.assert_called_once_with(SegmentType.ARRAY_ANY, test_array) def test_load_offloaded_variable_file_type_rebuilds_storage_backed_payload(self, draft_var_loader): - upload_file = Mock(spec=UploadFile) - upload_file.key = "storage/key/test_file.json" + upload_file = UploadFile( + tenant_id="tenant-id", + storage_type="opendal", + key="storage/key/test_file.json", + name="test.txt", + size=0, + extension="txt", + mime_type="text/plain", + created_by_role="account", + created_by="account-id", + created_at=datetime.now(), + used=False, + ) - variable_file = Mock(spec=WorkflowDraftVariableFile) - variable_file.value_type = SegmentType.FILE + variable_file = WorkflowDraftVariableFile( + tenant_id="tenant-id", + app_id="app-id", + user_id="user-id", + upload_file_id="upload-file-id", + size=0, + length=0, + value_type=SegmentType.FILE, + ) variable_file.upload_file = upload_file draft_var = WorkflowDraftVariable( @@ -205,131 +303,109 @@ class TestDraftVarLoaderSimple: assert variable.value == rebuilt_file rebuild_file.assert_called_once_with(file_mapping=raw_file, tenant_id="tenant-1") - def test_load_variables_with_offloaded_variables_unit(self, draft_var_loader): + @pytest.mark.parametrize( + "sqlite_session", + [(WorkflowDraftVariable, WorkflowDraftVariableFile, UploadFile)], + indirect=True, + ) + def test_load_variables_with_offloaded_variables_unit( + self, + draft_var_loader: DraftVarLoader, + sqlite_session: Session, + ): """Test load_variables method with mix of regular and offloaded variables.""" selectors = [["node1", "regular_var"], ["node2", "offloaded_var"]] - - # Mock regular variable - regular_draft_var = Mock(spec=WorkflowDraftVariable) - regular_draft_var.is_truncated.return_value = False - regular_draft_var.node_id = "node1" - regular_draft_var.name = "regular_var" - regular_draft_var.get_value.return_value = StringSegment(value="regular_value") - regular_draft_var.get_selector.return_value = ["node1", "regular_var"] - regular_draft_var.id = "regular-var-id" + regular_draft_var = WorkflowDraftVariable.new_node_variable( + app_id="test-app-id", + user_id="test-user-id", + node_id="node1", + name="regular_var", + value=StringSegment(value="regular_value"), + node_execution_id="execution-node1", + ) regular_draft_var.description = "regular description" + offloaded_draft_var = _persist_offloaded_variable( + sqlite_session, + node_id="node2", + name="offloaded_var", + ) + distractor = WorkflowDraftVariable.new_node_variable( + app_id="test-app-id", + user_id="another-user", + node_id="node1", + name="regular_var", + value=StringSegment(value="wrong user"), + node_execution_id="execution-distractor", + ) + sqlite_session.add_all([regular_draft_var, distractor]) + sqlite_session.commit() - # Mock offloaded variable - upload_file = Mock(spec=UploadFile) - upload_file.key = "storage/key/offloaded.txt" + offloaded_variable = Mock() + offloaded_variable.id = offloaded_draft_var.id + offloaded_variable.selector = ["node2", "offloaded_var"] - variable_file = Mock(spec=WorkflowDraftVariableFile) - variable_file.value_type = SegmentType.STRING - variable_file.upload_file = upload_file + with ( + patch("services.workflow_draft_variable_service.StorageKeyLoader"), + patch.object( + draft_var_loader, + "_load_offloaded_variable", + return_value=(("node2", "offloaded_var"), offloaded_variable), + ) as load_offloaded, + patch("services.workflow_draft_variable_service.ThreadPoolExecutor") as executor_cls, + ): + executor = executor_cls.return_value.__enter__.return_value + executor.map.side_effect = lambda function, values: [function(value) for value in values] - offloaded_draft_var = Mock(spec=WorkflowDraftVariable) - offloaded_draft_var.is_truncated.return_value = True - offloaded_draft_var.node_id = "node2" - offloaded_draft_var.name = "offloaded_var" - offloaded_draft_var.get_selector.return_value = ["node2", "offloaded_var"] - offloaded_draft_var.variable_file = variable_file - offloaded_draft_var.id = "offloaded-var-id" - offloaded_draft_var.description = "offloaded description" + result = draft_var_loader.load_variables(selectors) - draft_vars = [regular_draft_var, offloaded_draft_var] + assert {variable.id for variable in result} == {regular_draft_var.id, offloaded_draft_var.id} + load_offloaded.assert_called_once() + loaded_offloaded = load_offloaded.call_args.args[0] + assert isinstance(loaded_offloaded, WorkflowDraftVariable) + assert loaded_offloaded.id == offloaded_draft_var.id + assert loaded_offloaded.variable_file is not None + assert loaded_offloaded.variable_file.upload_file is not None + assert loaded_offloaded.variable_file.upload_file.key == "storage/key/offloaded_var.txt" - with patch("services.workflow_draft_variable_service.Session") as mock_session_cls: - mock_session = Mock() - mock_session_cls.return_value.__enter__.return_value = mock_session - - mock_service = Mock() - mock_service.get_draft_variables_by_selectors.return_value = draft_vars - - with patch( - "services.workflow_draft_variable_service.WorkflowDraftVariableService", return_value=mock_service - ): - with patch("services.workflow_draft_variable_service.StorageKeyLoader"): - with patch("factories.variable_factory.segment_to_variable") as mock_segment_to_variable: - # Mock regular variable creation - regular_variable = Mock() - regular_variable.selector = ["node1", "regular_var"] - - # Mock offloaded variable creation - offloaded_variable = Mock() - offloaded_variable.selector = ["node2", "offloaded_var"] - - mock_segment_to_variable.return_value = regular_variable - - with patch("services.workflow_draft_variable_service.storage") as mock_storage: - mock_storage.load.return_value = b"offloaded_content" - - with patch.object(draft_var_loader, "_load_offloaded_variable") as mock_load_offloaded: - mock_load_offloaded.return_value = (("node2", "offloaded_var"), offloaded_variable) - - with patch("concurrent.futures.ThreadPoolExecutor") as mock_executor_cls: - mock_executor = Mock() - mock_executor_cls.return_value.__enter__.return_value = mock_executor - mock_executor.map.return_value = [(("node2", "offloaded_var"), offloaded_variable)] - - # Execute the method - result = draft_var_loader.load_variables(selectors) - - # Verify results - assert len(result) == 2 - - # Verify service method was called - mock_service.get_draft_variables_by_selectors.assert_called_once_with( - draft_var_loader._app_id, - selectors, - user_id=draft_var_loader._user_id, - ) - - # Verify offloaded variable loading was called - mock_load_offloaded.assert_called_once_with(offloaded_draft_var) - - def test_load_variables_all_offloaded_variables_unit(self, draft_var_loader): + @pytest.mark.parametrize( + "sqlite_session", + [(WorkflowDraftVariable, WorkflowDraftVariableFile, UploadFile)], + indirect=True, + ) + def test_load_variables_all_offloaded_variables_unit( + self, + draft_var_loader: DraftVarLoader, + sqlite_session: Session, + ): """Test load_variables method with only offloaded variables.""" selectors = [["node1", "offloaded_var1"], ["node2", "offloaded_var2"]] + offloaded_var1 = _persist_offloaded_variable( + sqlite_session, + node_id="node1", + name="offloaded_var1", + ) + offloaded_var2 = _persist_offloaded_variable( + sqlite_session, + node_id="node2", + name="offloaded_var2", + ) + sqlite_session.commit() - # Mock first offloaded variable - offloaded_var1 = Mock(spec=WorkflowDraftVariable) - offloaded_var1.is_truncated.return_value = True - offloaded_var1.node_id = "node1" - offloaded_var1.name = "offloaded_var1" + with ( + patch("services.workflow_draft_variable_service.StorageKeyLoader"), + patch("services.workflow_draft_variable_service.ThreadPoolExecutor") as executor_cls, + ): + executor = executor_cls.return_value.__enter__.return_value + executor.map.return_value = [ + (("node1", "offloaded_var1"), Mock()), + (("node2", "offloaded_var2"), Mock()), + ] - # Mock second offloaded variable - offloaded_var2 = Mock(spec=WorkflowDraftVariable) - offloaded_var2.is_truncated.return_value = True - offloaded_var2.node_id = "node2" - offloaded_var2.name = "offloaded_var2" + result = draft_var_loader.load_variables(selectors) - draft_vars = [offloaded_var1, offloaded_var2] - - with patch("services.workflow_draft_variable_service.Session") as mock_session_cls: - mock_session = Mock() - mock_session_cls.return_value.__enter__.return_value = mock_session - - mock_service = Mock() - mock_service.get_draft_variables_by_selectors.return_value = draft_vars - - with patch( - "services.workflow_draft_variable_service.WorkflowDraftVariableService", return_value=mock_service - ): - with patch("services.workflow_draft_variable_service.StorageKeyLoader"): - with patch("services.workflow_draft_variable_service.ThreadPoolExecutor") as mock_executor_cls: - mock_executor = Mock() - mock_executor_cls.return_value.__enter__.return_value = mock_executor - mock_executor.map.return_value = [ - (("node1", "offloaded_var1"), Mock()), - (("node2", "offloaded_var2"), Mock()), - ] - - # Execute the method - result = draft_var_loader.load_variables(selectors) - - # Verify results - since we have only offloaded variables, should have 2 results - assert len(result) == 2 - - # Verify ThreadPoolExecutor was used - mock_executor_cls.assert_called_once_with(max_workers=10) - mock_executor.map.assert_called_once() + assert len(result) == 2 + executor_cls.assert_called_once_with(max_workers=10) + executor.map.assert_called_once() + loaded_draft_vars = executor.map.call_args.args[1] + assert {variable.id for variable in loaded_draft_vars} == {offloaded_var1.id, offloaded_var2.id} + assert all(variable.variable_file.upload_file is not None for variable in loaded_draft_vars) diff --git a/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py b/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py index 7f720575154..2c659b2621d 100644 --- a/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py +++ b/api/tests/unit_tests/services/workflow/test_node_output_inspector_service.py @@ -1,19 +1,20 @@ """Unit tests for NodeOutputInspectorService (Stage 4 §8). -The service reads from postgres and resolves agent v2 bindings; this suite -mocks the DB session and binding resolver so we exercise the view-construction -logic without DB / network access. +The service reads persisted workflow runs and node executions while the agent +v2 binding resolver and file URL boundaries remain isolated from network I/O. """ from __future__ import annotations import json +from collections.abc import Sequence from datetime import UTC, datetime from types import SimpleNamespace -from typing import Any +from typing import Any, Protocol from unittest.mock import MagicMock, patch import pytest +from sqlalchemy.orm import Session from core.workflow.file_reference import build_file_reference from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus @@ -22,7 +23,13 @@ from models.agent_config_entities import ( DeclaredOutputConfig, DeclaredOutputType, ) -from models.enums import WorkflowRunTriggeredFrom +from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom +from models.workflow import ( + WorkflowNodeExecutionModel, + WorkflowNodeExecutionTriggeredFrom, + WorkflowRun, + WorkflowType, +) from services.workflow.node_output_inspector_service import ( NodeOutputInspectorError, NodeOutputInspectorService, @@ -49,15 +56,28 @@ def _workflow_run( triggered_from: WorkflowRunTriggeredFrom = WorkflowRunTriggeredFrom.DEBUGGING, status: WorkflowExecutionStatus = WorkflowExecutionStatus.RUNNING, nodes: list[dict[str, Any]] | None = None, -): - return SimpleNamespace( + graph: str | None = None, +) -> WorkflowRun: + return WorkflowRun( id=run_id, workflow_id=workflow_id, tenant_id=tenant_id, app_id=app_id, + type=WorkflowType.WORKFLOW, triggered_from=triggered_from, + version="1", status=status, - graph=json.dumps({"nodes": nodes or []}), + graph=graph if graph is not None else json.dumps({"nodes": nodes or []}), + inputs="{}", + outputs="{}", + error=None, + elapsed_time=0, + total_tokens=0, + total_steps=0, + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + created_at=datetime(2024, 1, 1, tzinfo=UTC), + finished_at=None, ) @@ -72,16 +92,33 @@ def _execution( index: int = 1, created_at: datetime | None = None, finished_at: datetime | None = None, -): - return SimpleNamespace( + tenant_id: str = "tenant-1", + app_id: str = "app-1", + workflow_id: str = "workflow-1", + workflow_run_id: str = "run-1", +) -> WorkflowNodeExecutionModel: + return WorkflowNodeExecutionModel( + tenant_id=tenant_id, + app_id=app_id, + workflow_id=workflow_id, + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + workflow_run_id=workflow_run_id, + predecessor_node_id=None, + node_execution_id=None, node_id=node_id, node_type=node_type, title=title or node_id, status=status, + inputs=None, + process_data=None, outputs=json.dumps(outputs) if outputs is not None else None, + error=None, + elapsed_time=0, execution_metadata=json.dumps(execution_metadata) if execution_metadata is not None else None, index=index, created_at=created_at or datetime.now(UTC), + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", finished_at=finished_at, ) @@ -100,17 +137,32 @@ def _non_agent_node(*, node_id: str = "tool-node-1", node_type: str = "tool", ti } -def _mock_session( - *, - workflow_run: SimpleNamespace | None, - executions: list[SimpleNamespace] | None = None, -): - """Build a mock DB session with the configured rows.""" - executions = executions or [] - session = MagicMock() - session.scalar.return_value = workflow_run - session.scalars.return_value.all.return_value = executions - return session +class SessionFor(Protocol): + def __call__( + self, + *, + workflow_run: WorkflowRun | None, + executions: Sequence[WorkflowNodeExecutionModel] = (), + ) -> Session: ... + + +@pytest.fixture +def session_for(sqlite_session: Session) -> SessionFor: + """Persist one scenario in the shared SQLite session.""" + + def create_session( + *, + workflow_run: WorkflowRun | None, + executions: Sequence[WorkflowNodeExecutionModel] = (), + ) -> Session: + if workflow_run is not None: + sqlite_session.add(workflow_run) + sqlite_session.add_all(executions) + sqlite_session.commit() + sqlite_session.expunge_all() + return sqlite_session + + return create_session def _stub_binding_resolver(*, declared_outputs: list[DeclaredOutputConfig]): @@ -138,55 +190,56 @@ def _make_service(declared_outputs: list[DeclaredOutputConfig] | None = None) -> # ────────────────────────────────────────────────────────────────────────────── -def test_snapshot_404_when_workflow_run_missing(): +def test_snapshot_404_when_workflow_run_missing(session_for: SessionFor) -> None: service = _make_service() - session = _mock_session(workflow_run=None) + other_tenant_run = _workflow_run(run_id="missing", tenant_id="tenant-2") + session = session_for(workflow_run=other_tenant_run) with pytest.raises(NodeOutputInspectorError) as exc: service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="missing", session=session) assert exc.value.code == "workflow_run_not_found" -def test_snapshot_accepts_published_run_d1_lifted(): +def test_snapshot_accepts_published_run_d1_lifted(session_for: SessionFor) -> None: """D-1 was lifted 2026-05-26: any ``triggered_from`` is now accepted.""" service = _make_service() run = _workflow_run( nodes=[_agent_v2_node(node_id="agent-1")], triggered_from=WorkflowRunTriggeredFrom.APP_RUN, ) - session = _mock_session(workflow_run=run, executions=[]) + session = session_for(workflow_run=run, executions=[]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.workflow_run_id == "run-1" assert [n.node_id for n in snapshot.node_outputs] == ["agent-1"] -def test_snapshot_accepts_webhook_triggered_run(): +def test_snapshot_accepts_webhook_triggered_run(session_for: SessionFor) -> None: """Webhook / schedule / plugin triggers are also published-side.""" service = _make_service() run = _workflow_run( nodes=[_agent_v2_node(node_id="agent-1")], triggered_from=WorkflowRunTriggeredFrom.WEBHOOK, ) - session = _mock_session(workflow_run=run, executions=[]) + session = session_for(workflow_run=run, executions=[]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.workflow_run_id == "run-1" -def test_node_detail_404_when_node_id_absent_from_graph(): +def test_node_detail_404_when_node_id_absent_from_graph(session_for: SessionFor) -> None: service = _make_service() run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - session = _mock_session(workflow_run=run, executions=[]) + session = session_for(workflow_run=run, executions=[]) with pytest.raises(NodeOutputInspectorError) as exc: service.node_detail(app_model=_app_model(), workflow_run_id="run-1", node_id="ghost", session=session) assert exc.value.code == "node_not_in_workflow_run" -def test_output_preview_404_when_output_name_unknown(): +def test_output_preview_404_when_output_name_unknown(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING)], ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hello"}) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) with pytest.raises(NodeOutputInspectorError) as exc: service.output_preview( app_model=_app_model(), @@ -198,10 +251,10 @@ def test_output_preview_404_when_output_name_unknown(): assert exc.value.code == "node_output_not_declared" -def test_output_preview_404_when_node_id_absent_from_graph(): +def test_output_preview_404_when_node_id_absent_from_graph(session_for: SessionFor) -> None: service = _make_service() run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - session = _mock_session(workflow_run=run, executions=[]) + session = session_for(workflow_run=run, executions=[]) with pytest.raises(NodeOutputInspectorError) as exc: service.output_preview( app_model=_app_model(), @@ -218,12 +271,12 @@ def test_output_preview_404_when_node_id_absent_from_graph(): # ────────────────────────────────────────────────────────────────────────────── -def test_snapshot_status_pending_when_node_has_no_execution(): +def test_snapshot_status_pending_when_node_has_no_execution(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING)], ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) - session = _mock_session(workflow_run=run, executions=[]) + session = session_for(workflow_run=run, executions=[]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert len(snapshot.node_outputs) == 1 @@ -232,19 +285,19 @@ def test_snapshot_status_pending_when_node_has_no_execution(): assert node.outputs[0].status == NodeOutputStatus.PENDING -def test_snapshot_status_running(): +def test_snapshot_status_running(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING)], ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", status=WorkflowNodeExecutionStatus.RUNNING) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].node_status == NodeStatus.RUNNING assert snapshot.node_outputs[0].outputs[0].status == NodeOutputStatus.RUNNING -def test_snapshot_status_failed_node_marks_all_outputs_failed(): +def test_snapshot_status_failed_node_marks_all_outputs_failed(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig(name="a", type=DeclaredOutputType.STRING), @@ -253,26 +306,26 @@ def test_snapshot_status_failed_node_marks_all_outputs_failed(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", status=WorkflowNodeExecutionStatus.FAILED) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) statuses = {o.name: o.status for o in snapshot.node_outputs[0].outputs} assert statuses == {"a": NodeOutputStatus.FAILED, "b": NodeOutputStatus.FAILED} -def test_snapshot_status_ready_when_outputs_present_and_no_failure_metadata(): +def test_snapshot_status_ready_when_outputs_present_and_no_failure_metadata(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING)], ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hello"}) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.READY assert output.value_preview == "hello" -def test_snapshot_marks_type_check_failure(): +def test_snapshot_marks_type_check_failure(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING)], ) @@ -287,7 +340,7 @@ def test_snapshot_marks_type_check_failure(): } }, ) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.status == NodeOutputStatus.TYPE_CHECK_FAILED @@ -296,7 +349,7 @@ def test_snapshot_marks_type_check_failure(): assert output.type_check.reason == "wrong shape" -def test_snapshot_marks_output_check_failure_when_type_check_passed(): +def test_snapshot_marks_output_check_failure_when_type_check_passed(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig( @@ -317,7 +370,7 @@ def test_snapshot_marks_output_check_failure_when_type_check_passed(): }, }, ) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) with patch( "services.workflow.node_output_inspector_service.file_helpers.get_signed_file_url", return_value="https://signed.example/x", @@ -330,7 +383,7 @@ def test_snapshot_marks_output_check_failure_when_type_check_passed(): assert output.output_check.reason == "benchmark mismatch" -def test_snapshot_marks_not_produced_when_declared_output_missing_from_payload(): +def test_snapshot_marks_not_produced_when_declared_output_missing_from_payload(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING), @@ -339,7 +392,7 @@ def test_snapshot_marks_not_produced_when_declared_output_missing_from_payload() ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"text": "hi"}) # optional_meta missing - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) statuses = {o.name: o.status for o in snapshot.node_outputs[0].outputs} assert statuses == {"text": NodeOutputStatus.READY, "optional_meta": NodeOutputStatus.NOT_PRODUCED} @@ -350,7 +403,7 @@ def test_snapshot_marks_not_produced_when_declared_output_missing_from_payload() # ────────────────────────────────────────────────────────────────────────────── -def test_non_agent_node_outputs_inferred_from_payload_keys(): +def test_non_agent_node_outputs_inferred_from_payload_keys(session_for: SessionFor) -> None: service = _make_service() run = _workflow_run(nodes=[_non_agent_node(node_id="tool-1", node_type="tool")]) ex = _execution( @@ -358,7 +411,7 @@ def test_non_agent_node_outputs_inferred_from_payload_keys(): node_type="tool", outputs={"message": "sent", "thread_ts": "1234"}, ) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output_names = sorted(o.name for o in snapshot.node_outputs[0].outputs) assert output_names == ["message", "thread_ts"] @@ -372,7 +425,7 @@ def test_non_agent_node_outputs_inferred_from_payload_keys(): # ────────────────────────────────────────────────────────────────────────────── -def test_file_output_preview_includes_signed_url(): +def test_file_output_preview_includes_signed_url(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig(name="report", type=DeclaredOutputType.FILE), @@ -384,7 +437,7 @@ def test_file_output_preview_includes_signed_url(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) with patch( "services.workflow.node_output_inspector_service._resolve_preview_url", return_value="https://signed.example/x.pdf", @@ -396,7 +449,7 @@ def test_file_output_preview_includes_signed_url(): assert preview_value["reference"] == file_payload["reference"] -def test_file_output_preview_endpoint_returns_full_value_with_signed_url(): +def test_file_output_preview_endpoint_returns_full_value_with_signed_url(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig(name="report", type=DeclaredOutputType.FILE), @@ -408,7 +461,7 @@ def test_file_output_preview_endpoint_returns_full_value_with_signed_url(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) with patch( "services.workflow.node_output_inspector_service._resolve_preview_url", return_value="https://signed.example/x.pdf", @@ -450,7 +503,7 @@ def test_resolve_preview_url_uses_standard_file_factory(): resolve_file_url.assert_called_once_with(file) -def test_array_file_output_preview_includes_signed_urls_for_each_item(): +def test_array_file_output_preview_includes_signed_urls_for_each_item(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig( @@ -472,7 +525,7 @@ def test_array_file_output_preview_includes_signed_urls_for_each_item(): }, ] ex = _execution(node_id="agent-1", outputs={"files": file_payloads}) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) with patch( "services.workflow.node_output_inspector_service._resolve_preview_url", side_effect=[ @@ -506,7 +559,7 @@ def test_array_file_output_preview_includes_signed_urls_for_each_item(): ] -def test_file_output_preview_uses_none_when_signed_url_resolution_fails(): +def test_file_output_preview_uses_none_when_signed_url_resolution_fails(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig(name="report", type=DeclaredOutputType.FILE), @@ -518,7 +571,7 @@ def test_file_output_preview_uses_none_when_signed_url_resolution_fails(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"report": file_payload}) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) with patch( "services.workflow.node_output_inspector_service._resolve_preview_url", side_effect=RuntimeError("boom"), @@ -530,7 +583,7 @@ def test_file_output_preview_uses_none_when_signed_url_resolution_fails(): assert preview_value["preview_url"] is None -def test_object_output_preview_does_not_augment_canonical_file_mapping_shape(): +def test_object_output_preview_does_not_augment_canonical_file_mapping_shape(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig(name="meta", type=DeclaredOutputType.OBJECT), @@ -542,7 +595,7 @@ def test_object_output_preview_does_not_augment_canonical_file_mapping_shape(): "reference": build_file_reference(record_id="550e8400-e29b-41d4-a716-446655440000"), } ex = _execution(node_id="agent-1", outputs={"meta": raw_value}) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) with patch( "services.workflow.node_output_inspector_service._resolve_preview_url", return_value="https://signed.example/x.pdf", @@ -565,7 +618,7 @@ def test_object_output_preview_does_not_augment_canonical_file_mapping_shape(): # ────────────────────────────────────────────────────────────────────────────── -def test_retried_count_pulled_from_attempt_metadata(): +def test_retried_count_pulled_from_attempt_metadata(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[DeclaredOutputConfig(name="text", type=DeclaredOutputType.STRING)], ) @@ -575,7 +628,7 @@ def test_retried_count_pulled_from_attempt_metadata(): outputs={"text": "ok"}, execution_metadata={"attempt": 2}, ) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].outputs[0].retried == 2 @@ -585,7 +638,7 @@ def test_retried_count_pulled_from_attempt_metadata(): # ────────────────────────────────────────────────────────────────────────────── -def test_keeps_latest_execution_per_node_by_index(): +def test_keeps_latest_execution_per_node_by_index(session_for: SessionFor) -> None: """When a node has multiple executions (retries / iterations) keep the canonical one — the row with the highest ``index``.""" service = _make_service( @@ -594,7 +647,9 @@ def test_keeps_latest_execution_per_node_by_index(): run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) older = _execution(node_id="agent-1", outputs={"text": "old"}, index=1) newer = _execution(node_id="agent-1", outputs={"text": "new"}, index=5) - session = _mock_session(workflow_run=run, executions=[older, newer]) + other_tenant = _execution(node_id="agent-1", outputs={"text": "other tenant"}, index=99, tenant_id="tenant-2") + other_run = _execution(node_id="agent-1", outputs={"text": "other run"}, index=100, workflow_run_id="run-2") + session = session_for(workflow_run=run, executions=[older, newer, other_tenant, other_run]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs[0].outputs[0].value_preview == "new" @@ -604,7 +659,7 @@ def test_keeps_latest_execution_per_node_by_index(): # ────────────────────────────────────────────────────────────────────────────── -def test_array_typed_output_with_array_item_renders_correctly(): +def test_array_typed_output_with_array_item_renders_correctly(session_for: SessionFor) -> None: service = _make_service( declared_outputs=[ DeclaredOutputConfig( @@ -616,7 +671,7 @@ def test_array_typed_output_with_array_item_renders_correctly(): ) run = _workflow_run(nodes=[_agent_v2_node(node_id="agent-1")]) ex = _execution(node_id="agent-1", outputs={"files": []}) - session = _mock_session(workflow_run=run, executions=[ex]) + session = session_for(workflow_run=run, executions=[ex]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) output = snapshot.node_outputs[0].outputs[0] assert output.type == DeclaredOutputType.ARRAY @@ -627,17 +682,9 @@ def test_array_typed_output_with_array_item_renders_correctly(): # ────────────────────────────────────────────────────────────────────────────── -def test_unparseable_graph_blob_yields_empty_snapshot_not_500(): +def test_unparseable_graph_blob_yields_empty_snapshot_not_500(session_for: SessionFor) -> None: service = _make_service() - run = SimpleNamespace( - id="run-1", - workflow_id="workflow-1", - tenant_id="tenant-1", - app_id="app-1", - triggered_from=WorkflowRunTriggeredFrom.DEBUGGING, - status=WorkflowExecutionStatus.RUNNING, - graph="{not valid json", - ) - session = _mock_session(workflow_run=run, executions=[]) + run = _workflow_run(graph="{not valid json") + session = session_for(workflow_run=run, executions=[]) snapshot = service.snapshot_workflow_run(app_model=_app_model(), workflow_run_id="run-1", session=session) assert snapshot.node_outputs == [] diff --git a/api/tests/unit_tests/services/workflow/test_queue_dispatcher.py b/api/tests/unit_tests/services/workflow/test_queue_dispatcher.py index 18b46e40f52..78cce598eb6 100644 --- a/api/tests/unit_tests/services/workflow/test_queue_dispatcher.py +++ b/api/tests/unit_tests/services/workflow/test_queue_dispatcher.py @@ -1,5 +1,6 @@ from unittest.mock import patch +from enums import DeploymentEdition from services.workflow.queue_dispatcher import ( ProfessionalQueueDispatcher, QueueDispatcherManager, @@ -36,8 +37,8 @@ class TestDispatchers: class TestQueueDispatcherManager: @patch("services.workflow.queue_dispatcher.BillingService") @patch("services.workflow.queue_dispatcher.dify_config") - def test_billing_enabled_professional_plan(self, mock_config, mock_billing): - mock_config.BILLING_ENABLED = True + def test_cloud_edition_professional_plan(self, mock_config, mock_billing): + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_billing.get_info.return_value = {"subscription": {"plan": "professional"}} dispatcher = QueueDispatcherManager.get_dispatcher("tenant-1") @@ -46,8 +47,8 @@ class TestQueueDispatcherManager: @patch("services.workflow.queue_dispatcher.BillingService") @patch("services.workflow.queue_dispatcher.dify_config") - def test_billing_enabled_team_plan(self, mock_config, mock_billing): - mock_config.BILLING_ENABLED = True + def test_cloud_edition_team_plan(self, mock_config, mock_billing): + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_billing.get_info.return_value = {"subscription": {"plan": "team"}} dispatcher = QueueDispatcherManager.get_dispatcher("tenant-1") @@ -56,8 +57,8 @@ class TestQueueDispatcherManager: @patch("services.workflow.queue_dispatcher.BillingService") @patch("services.workflow.queue_dispatcher.dify_config") - def test_billing_enabled_sandbox_plan(self, mock_config, mock_billing): - mock_config.BILLING_ENABLED = True + def test_cloud_edition_sandbox_plan(self, mock_config, mock_billing): + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_billing.get_info.return_value = {"subscription": {"plan": "sandbox"}} dispatcher = QueueDispatcherManager.get_dispatcher("tenant-1") @@ -66,8 +67,8 @@ class TestQueueDispatcherManager: @patch("services.workflow.queue_dispatcher.BillingService") @patch("services.workflow.queue_dispatcher.dify_config") - def test_billing_enabled_unknown_plan_defaults_to_sandbox(self, mock_config, mock_billing): - mock_config.BILLING_ENABLED = True + def test_cloud_edition_unknown_plan_defaults_to_sandbox(self, mock_config, mock_billing): + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_billing.get_info.return_value = {"subscription": {"plan": "enterprise"}} dispatcher = QueueDispatcherManager.get_dispatcher("tenant-1") @@ -76,8 +77,8 @@ class TestQueueDispatcherManager: @patch("services.workflow.queue_dispatcher.BillingService") @patch("services.workflow.queue_dispatcher.dify_config") - def test_billing_enabled_service_failure_defaults_to_sandbox(self, mock_config, mock_billing): - mock_config.BILLING_ENABLED = True + def test_cloud_edition_billing_failure_defaults_to_sandbox(self, mock_config, mock_billing): + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_billing.get_info.side_effect = Exception("billing unavailable") dispatcher = QueueDispatcherManager.get_dispatcher("tenant-1") @@ -85,8 +86,8 @@ class TestQueueDispatcherManager: assert isinstance(dispatcher, SandboxQueueDispatcher) @patch("services.workflow.queue_dispatcher.dify_config") - def test_billing_disabled_defaults_to_team(self, mock_config): - mock_config.BILLING_ENABLED = False + def test_non_cloud_edition_defaults_to_team(self, mock_config): + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.COMMUNITY dispatcher = QueueDispatcherManager.get_dispatcher("tenant-1") @@ -95,7 +96,7 @@ class TestQueueDispatcherManager: @patch("services.workflow.queue_dispatcher.BillingService") @patch("services.workflow.queue_dispatcher.dify_config") def test_missing_subscription_key_defaults_to_sandbox(self, mock_config, mock_billing): - mock_config.BILLING_ENABLED = True + mock_config.DEPLOYMENT_EDITION = DeploymentEdition.CLOUD mock_billing.get_info.return_value = {} dispatcher = QueueDispatcherManager.get_dispatcher("tenant-1") diff --git a/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py b/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py index 60c8b5213f6..706ea415e2a 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_converter_additional.py @@ -6,6 +6,8 @@ from typing import Any, cast from unittest.mock import MagicMock, call import pytest +from sqlalchemy import event, select +from sqlalchemy.orm import Session from core.app.app_config.entities import ( AdvancedChatMessageEntity, @@ -20,7 +22,8 @@ from core.app.app_config.entities import ( from core.helper import encrypter from core.prompt.utils.prompt_template_parser import PromptTemplateParser from models.api_based_extension import APIBasedExtension, APIBasedExtensionPoint -from models.model import Account, App, AppMode, AppModelConfig +from models.model import Account, App, AppMode, AppModelConfig, IconType +from models.workflow import Workflow, WorkflowType from services.workflow import workflow_converter as converter_module from services.workflow.workflow_converter import WorkflowConverter @@ -88,7 +91,9 @@ def test__convert_to_start_node(default_variables: list[VariableEntity]) -> None assert result["data"]["variables"][0]["variable"] == "text_input" -def test__convert_to_http_request_node_for_chatbot(default_variables: list[VariableEntity]) -> None: +def test__convert_to_http_request_node_for_chatbot( + default_variables: list[VariableEntity], unbound_session: Session +) -> None: app_model = MagicMock() app_model.id = "app_id" app_model.tenant_id = "tenant_id" @@ -118,7 +123,7 @@ def test__convert_to_http_request_node_for_chatbot(default_variables: list[Varia app_model=app_model, variables=default_variables, external_data_variables=external_data_variables, - session=MagicMock(), + session=unbound_session, ) assert len(nodes) == 2 @@ -131,7 +136,9 @@ def test__convert_to_http_request_node_for_chatbot(default_variables: list[Varia assert mapping == {"external_variable": "code_1"} -def test__convert_to_http_request_node_for_workflow_app(default_variables: list[VariableEntity]) -> None: +def test__convert_to_http_request_node_for_workflow_app( + default_variables: list[VariableEntity], unbound_session: Session +) -> None: app_model = MagicMock() app_model.id = "app_id" app_model.tenant_id = "tenant_id" @@ -161,7 +168,7 @@ def test__convert_to_http_request_node_for_workflow_app(default_variables: list[ app_model=app_model, variables=default_variables, external_data_variables=external_data_variables, - session=MagicMock(), + session=unbound_session, ) body = json.loads(nodes[0]["data"]["body"]["data"]) @@ -355,9 +362,10 @@ def test__convert_to_answer_node() -> None: assert node["data"]["type"] == BuiltinNodeTypes.ANSWER -def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(converter: WorkflowConverter) -> None: +def test_convert_to_workflow_should_raise_when_app_model_config_is_missing( + converter: WorkflowConverter, unbound_session: Session +) -> None: app_model = _app_model(app_model_config_id=None) - session = MagicMock() with pytest.raises(ValueError, match="App model config is required"): converter.convert_to_workflow( @@ -367,10 +375,10 @@ def test_convert_to_workflow_should_raise_when_app_model_config_is_missing(conve icon_type="emoji", icon="robot", icon_background="#fff", - session=session, + session=unbound_session, ) - session.get.assert_not_called() + assert not unbound_session.in_transaction() @pytest.mark.parametrize( @@ -385,23 +393,26 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( monkeypatch: pytest.MonkeyPatch, source_mode: AppMode, expected_mode: AppMode, + sqlite_session: Session, ) -> None: - class FakeApp: - def __init__(self) -> None: - self.id = "new-app-id" - - workflow = SimpleNamespace(app_id=None) - monkeypatch.setattr(converter, "convert_app_model_config_to_workflow", MagicMock(return_value=workflow)) - monkeypatch.setattr(converter_module, "App", FakeApp) - - app_model_config = _app_model_config(id="config-1") - phase_events: list[str] = [] - db_session = SimpleNamespace( - add=MagicMock(), - flush=MagicMock(), - commit=MagicMock(side_effect=lambda: phase_events.append("commit")), - get=MagicMock(return_value=app_model_config), + app_model_config = AppModelConfig(app_id="source-app") + app_model_config.id = "config-1" + sqlite_session.add(app_model_config) + sqlite_session.flush() + workflow = Workflow( + tenant_id="tenant-1", + app_id="source-app", + type=WorkflowType.WORKFLOW, + version=Workflow.VERSION_DRAFT, + graph="{}", + features="{}", + created_by="account-1", + environment_variables=[], + conversation_variables=[], ) + monkeypatch.setattr(converter, "convert_app_model_config_to_workflow", MagicMock(return_value=workflow)) + phase_events: list[str] = [] + event.listen(sqlite_session, "after_commit", lambda _session: phase_events.append("commit")) send_mock = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("signal")) monkeypatch.setattr(converter_module.app_was_created, "send", send_mock) @@ -409,9 +420,10 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( account = _account(id="account-1") app_model = _app_model( tenant_id="tenant-1", + id="source-app", name="Source App", mode=source_mode, - icon_type="emoji", + icon_type=IconType.EMOJI, icon="sparkles", icon_background="#123456", enable_site=True, @@ -429,27 +441,25 @@ def test_convert_to_workflow_should_create_new_app_with_fallback_fields( icon_type="", icon="", icon_background="", - session=db_session, + session=sqlite_session, ) assert new_app.name == "Source App(workflow)" assert new_app.mode == expected_mode - assert new_app.icon_type == "emoji" + assert new_app.icon_type == IconType.EMOJI assert new_app.icon == "sparkles" assert new_app.icon_background == "#123456" assert new_app.created_by == "account-1" - assert workflow.app_id == "new-app-id" - db_session.add.assert_called_once() - db_session.flush.assert_called_once() + assert workflow.app_id == new_app.id + assert sqlite_session.get(App, new_app.id) is new_app assert phase_events == ["commit", "signal", "commit"] - assert db_session.commit.call_count == 2 - db_session.get.assert_called_once_with(AppModelConfig, "config-1") - send_mock.assert_called_once_with(new_app, account=account, session=db_session) + send_mock.assert_called_once_with(new_app, account=account, session=sqlite_session) def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_and_features( converter: WorkflowConverter, monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: app_model = _app_model(id="app-1", tenant_id="tenant-1", mode=AppMode.CHAT) app_config = SimpleNamespace( @@ -471,12 +481,6 @@ def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_a }, ) - class FakeWorkflow: - VERSION_DRAFT = "draft" - - def __init__(self, **kwargs: Any) -> None: - self.__dict__.update(kwargs) - monkeypatch.setattr(converter, "_get_new_app_mode", MagicMock(return_value=AppMode.ADVANCED_CHAT)) monkeypatch.setattr(converter, "_convert_to_app_config", MagicMock(return_value=app_config)) monkeypatch.setattr( @@ -513,15 +517,11 @@ def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_a "_convert_to_answer_node", MagicMock(return_value={"id": "answer", "position": None, "data": {"type": BuiltinNodeTypes.ANSWER}}), ) - monkeypatch.setattr(converter_module, "Workflow", FakeWorkflow) - - db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock()) - workflow = converter.convert_app_model_config_to_workflow( app_model=app_model, app_model_config=_app_model_config(id="cfg"), account_id="account-1", - session=db_session, + session=sqlite_session, ) graph = json.loads(workflow.graph) @@ -531,13 +531,13 @@ def test_convert_app_model_config_to_workflow_should_build_advanced_chat_graph_a features = json.loads(workflow.features) assert "opening_statement" in features assert "retriever_resource" in features - db_session.add.assert_called_once() - db_session.commit.assert_called_once() + assert sqlite_session.scalar(select(Workflow).where(Workflow.id == workflow.id)) is workflow def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_end_node( converter: WorkflowConverter, monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, ) -> None: app_model = _app_model(id="app-1", tenant_id="tenant-1", mode=AppMode.COMPLETION) app_config = SimpleNamespace( @@ -554,12 +554,6 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en }, ) - class FakeWorkflow: - VERSION_DRAFT = "draft" - - def __init__(self, **kwargs: Any) -> None: - self.__dict__.update(kwargs) - monkeypatch.setattr(converter, "_get_new_app_mode", MagicMock(return_value=AppMode.WORKFLOW)) monkeypatch.setattr(converter, "_convert_to_app_config", MagicMock(return_value=app_config)) monkeypatch.setattr( @@ -580,15 +574,11 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en "_convert_to_end_node", MagicMock(return_value={"id": "end", "position": None, "data": {"type": BuiltinNodeTypes.END}}), ) - monkeypatch.setattr(converter_module, "Workflow", FakeWorkflow) - - db_session = SimpleNamespace(add=MagicMock(), commit=MagicMock()) - workflow = converter.convert_app_model_config_to_workflow( app_model=app_model, app_model_config=_app_model_config(id="cfg"), account_id="account-1", - session=db_session, + session=sqlite_session, ) graph = json.loads(workflow.graph) @@ -597,11 +587,13 @@ def test_convert_app_model_config_to_workflow_should_build_workflow_mode_with_en features = json.loads(workflow.features) assert set(features.keys()) == {"text_to_speech", "file_upload", "sensitive_word_avoidance"} + assert sqlite_session.scalar(select(Workflow).where(Workflow.id == workflow.id)) is workflow def test_convert_to_app_config_should_route_to_correct_manager( converter: WorkflowConverter, monkeypatch: pytest.MonkeyPatch, + unbound_session: Session, ) -> None: agent_result = SimpleNamespace(kind="agent") chat_result = SimpleNamespace(kind="chat") @@ -614,7 +606,6 @@ def test_convert_to_app_config_should_route_to_correct_manager( monkeypatch.setattr(converter_module.ChatAppConfigManager, "get_app_config", chat_get_app_config) monkeypatch.setattr(converter_module.CompletionAppConfigManager, "get_app_config", completion_get_app_config) monkeypatch.setattr(converter_module, "load_annotation_reply_config", load_annotation_reply) - session = MagicMock() agent_mode_app = _app_model(mode=AppMode.AGENT_CHAT, is_agent_with_session=MagicMock(return_value=False)) agent_flag_app = _app_model(mode=AppMode.CHAT, is_agent_with_session=MagicMock(return_value=True)) chat_app = _app_model(mode=AppMode.CHAT, is_agent_with_session=MagicMock(return_value=False)) @@ -627,31 +618,36 @@ def test_convert_to_app_config_should_route_to_correct_manager( from_agent_mode = converter._convert_to_app_config( app_model=agent_mode_app, app_model_config=agent_mode_config, - session=session, + session=unbound_session, ) from_agent_flag = converter._convert_to_app_config( app_model=agent_flag_app, app_model_config=agent_flag_config, - session=session, + session=unbound_session, ) from_chat_mode = converter._convert_to_app_config( app_model=chat_app, app_model_config=chat_config, - session=session, + session=unbound_session, ) from_completion_mode = converter._convert_to_app_config( app_model=completion_app, app_model_config=completion_config, - session=session, + session=unbound_session, ) assert from_agent_mode is agent_result assert from_agent_flag is agent_result assert from_chat_mode is chat_result assert from_completion_mode is completion_result - agent_flag_app.is_agent_with_session.assert_called_once_with(session=session) + agent_flag_app.is_agent_with_session.assert_called_once_with(session=unbound_session) load_annotation_reply.assert_has_calls( - [call(session, "app-1"), call(session, "app-2"), call(session, "app-3"), call(session, "app-4")] + [ + call(unbound_session, "app-1"), + call(unbound_session, "app-2"), + call(unbound_session, "app-3"), + call(unbound_session, "app-4"), + ] ) assert all( manager_call.kwargs["annotation_reply"] == {"enabled": False} @@ -660,18 +656,20 @@ def test_convert_to_app_config_should_route_to_correct_manager( ) -def test_convert_to_app_config_should_raise_for_invalid_app_mode(converter: WorkflowConverter) -> None: +def test_convert_to_app_config_should_raise_for_invalid_app_mode( + converter: WorkflowConverter, unbound_session: Session +) -> None: app_model = _app_model(mode=AppMode.WORKFLOW, is_agent_with_session=MagicMock(return_value=False)) - session = MagicMock() with pytest.raises(ValueError, match="Invalid app mode"): converter._convert_to_app_config( - app_model=app_model, app_model_config=_app_model_config(id="cfg"), session=session + app_model=app_model, app_model_config=_app_model_config(id="cfg"), session=unbound_session ) def test_convert_to_http_request_node_should_skip_non_api_and_missing_extension_id( converter: WorkflowConverter, + unbound_session: Session, ) -> None: app_model = _app_model(id="app-1", tenant_id="tenant-1", mode=AppMode.CHAT) external_data_variables = [ @@ -683,7 +681,7 @@ def test_convert_to_http_request_node_should_skip_non_api_and_missing_extension_ app_model=app_model, variables=[], external_data_variables=external_data_variables, - session=MagicMock(), + session=unbound_session, ) assert nodes == [] diff --git a/api/tests/unit_tests/services/workflow/test_workflow_draft_variable_service.py b/api/tests/unit_tests/services/workflow/test_workflow_draft_variable_service.py index 9053f7c1b73..5d7b2db0013 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_draft_variable_service.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_draft_variable_service.py @@ -1,4 +1,5 @@ import dataclasses +import json import secrets import uuid from types import SimpleNamespace @@ -128,7 +129,7 @@ class TestDraftVariableSaver: assert name == c.expected_name, fail_msg def test_build_variables_from_start_mapping_rebuilds_system_files(self, sqlite_session: Session): - mock_user = MagicMock(spec=Account) + mock_user = Account(name="Test Account", email="test@example.com") mock_user.id = str(uuid.uuid4()) saver = DraftVariableSaver( session=sqlite_session, @@ -169,9 +170,8 @@ class TestDraftVariableSaver: def draft_saver(self, sqlite_session: Session): """Create DraftVariableSaver instance with user context.""" # Create a mock user - mock_user = MagicMock(spec=Account) + mock_user = Account(name="Test Account", email="test@example.com") mock_user.id = "test-user-id" - mock_user.tenant_id = "test-tenant-id" return DraftVariableSaver( session=sqlite_session, @@ -218,9 +218,8 @@ class TestDraftVariableSaver: assert draft_var.file_id == mock_draft_var_file.id def test_try_offload_large_variable_uses_resource_tenant(self, sqlite_session: Session): - mock_user = MagicMock(spec=Account) + mock_user = Account(name="Test Account", email="test@example.com") mock_user.id = "test-user-id" - mock_user.current_tenant_id = "" saver = DraftVariableSaver( session=sqlite_session, tenant_id="app-tenant-id", @@ -268,7 +267,7 @@ class TestDraftVariableSaver: self, mock_batch_upsert, sqlite_session: Session ): """Start node should persist common `sys.*` variables, not only `sys.files`.""" - mock_user = MagicMock(spec=Account) + mock_user = Account(name="Test Account", email="test@example.com") mock_user.id = "test-user-id" mock_user.tenant_id = "test-tenant-id" @@ -303,7 +302,7 @@ class TestDraftVariableSaver: @patch("services.workflow_draft_variable_service._batch_upsert_draft_variable", autospec=True) def test_start_node_save_normalizes_reserved_prefix_outputs(self, mock_batch_upsert, sqlite_session: Session): - mock_user = MagicMock(spec=Account) + mock_user = Account(name="Test Account", email="test@example.com") mock_user.id = "test-user-id" mock_user.tenant_id = "test-tenant-id" @@ -474,13 +473,11 @@ class TestWorkflowDraftVariableService: """Reset a node variable from its execution output and flush the restored value.""" service = WorkflowDraftVariableService(sqlite_session) - # Create mock execution record - mock_execution = Mock(spec=WorkflowNodeExecutionModel) - mock_execution.load_full_outputs.return_value = {"test_var": "output_value"} + execution = WorkflowNodeExecutionModel(outputs=json.dumps({"test_var": "output_value"})) # Mock the repository to return the execution record service._api_node_execution_repo = Mock() - service._api_node_execution_repo.get_execution_by_id.return_value = mock_execution + service._api_node_execution_repo.get_execution_by_id.return_value = execution test_app_id = self._get_test_app_id() workflow = self._create_test_workflow(test_app_id) @@ -549,13 +546,11 @@ class TestWorkflowDraftVariableService: sqlite_session.add(variable) sqlite_session.commit() - # Create mock execution record - mock_execution = Mock(spec=WorkflowNodeExecutionModel) - mock_execution.load_full_outputs.return_value = {"sys.files": "[]"} + execution = WorkflowNodeExecutionModel(outputs=json.dumps({"sys.files": "[]"})) # Mock the repository to return the execution record service._api_node_execution_repo = Mock() - service._api_node_execution_repo.get_execution_by_id.return_value = mock_execution + service._api_node_execution_repo.get_execution_by_id.return_value = execution with patch.object(sqlite_session, "flush", wraps=sqlite_session.flush) as flush: result = service._reset_node_var_or_sys_var(workflow, variable) @@ -584,13 +579,11 @@ class TestWorkflowDraftVariableService: sqlite_session.add(variable) sqlite_session.commit() - # Create mock execution record - mock_execution = Mock(spec=WorkflowNodeExecutionModel) - mock_execution.load_full_outputs.return_value = {"sys.query": "reset query"} + execution = WorkflowNodeExecutionModel(outputs=json.dumps({"sys.query": "reset query"})) # Mock the repository to return the execution record service._api_node_execution_repo = Mock() - service._api_node_execution_repo.get_execution_by_id.return_value = mock_execution + service._api_node_execution_repo.get_execution_by_id.return_value = execution with patch.object(sqlite_session, "flush", wraps=sqlite_session.flush) as flush: result = service._reset_node_var_or_sys_var(workflow, variable) diff --git a/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service.py b/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service.py index 58892f0ebb3..642389fe035 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service.py @@ -3,6 +3,7 @@ import queue from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import UTC, datetime +from decimal import Decimal from itertools import cycle from threading import Event from types import SimpleNamespace @@ -10,6 +11,7 @@ from typing import Any, cast, override from unittest.mock import MagicMock import pytest +from sqlalchemy import event as orm_event from sqlalchemy.orm import Session, sessionmaker from core.app.app_config.entities import WorkflowUIBasedAppConfig @@ -26,9 +28,10 @@ from core.workflow.nodes.human_input.enums import ValueSourceType from core.workflow.nodes.human_input.pause_reason import HumanInputRequired from graphon.enums import WorkflowExecutionStatus, WorkflowNodeExecutionStatus from graphon.runtime import GraphRuntimeState, VariablePool -from models.enums import CreatorUserRole -from models.human_input import RecipientType -from models.model import AppMode +from libs.datetime_utils import to_utc_timestamp +from models.enums import ConversationFromSource, CreatorUserRole +from models.human_input import HumanInputForm, HumanInputFormRecipient, RecipientType +from models.model import AppMode, Message from models.workflow import WorkflowRun from repositories.api_workflow_node_execution_repository import WorkflowNodeExecutionSnapshot from repositories.entities.workflow_pause import WorkflowPauseEntity @@ -141,8 +144,7 @@ def _build_resumption_context(task_id: str, *, select_options: list[str] | None runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0) if select_options is not None: runtime_state.variable_pool.add(("start", "options"), select_options) - runtime_state.register_paused_node("node-1") - runtime_state.outputs = {"result": "value"} + runtime_state.set_output("result", "value") wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity) return WorkflowResumptionContext( generate_entity=wrapper, @@ -247,7 +249,7 @@ def _build_resumption_context_additional(task_id: str) -> WorkflowResumptionCont workflow_execution_id="run-1", ) runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0) - runtime_state.outputs = {"answer": "ok"} + runtime_state.set_output("answer", "ok") wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity) return WorkflowResumptionContext( generate_entity=wrapper, @@ -283,23 +285,57 @@ def _build_advanced_chat_resumption_context(conversation_id: str | None) -> Work ) -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 +def _persist_message( + session_maker: sessionmaker[Session], + *, + message_id: str = "msg-1", + app_id: str = "app-1", + conversation_id: str = "conv-1", + workflow_run_id: str = "run-1", + answer: str = "answer", + created_at: datetime | None = None, +) -> None: + message = Message( + id=message_id, + app_id=app_id, + conversation_id=conversation_id, + query="question", + message={"role": "user", "content": "question"}, + answer=answer, + message_unit_price=Decimal(0), + answer_unit_price=Decimal(0), + currency="USD", + from_source=ConversationFromSource.API, + workflow_run_id=workflow_run_id, + ) + message.inputs = {} + if created_at is not None: + message.created_at = created_at + with session_maker.begin() as session: + session.add(message) -class _SessionMaker: - def __init__(self, session: Any) -> None: - self._session = session - - def __call__(self) -> _SessionContext: - return _SessionContext(self._session) +def _persist_human_input_form( + session_maker: sessionmaker[Session], + *, + recipients: Sequence[HumanInputFormRecipient] = (), +) -> datetime: + expiration_time = datetime(2024, 1, 1) + form = HumanInputForm( + id="form-1", + tenant_id="tenant-1", + app_id="app-1", + workflow_run_id="run-1", + conversation_id=None, + node_id="node-1", + form_definition='{"display_in_ui": true}', + rendered_content="content", + expiration_time=expiration_time, + ) + with session_maker.begin() as session: + session.add(form) + session.add_all(recipients) + return expiration_time class _SubscriptionContext: @@ -359,14 +395,12 @@ class _PauseEntity(WorkflowPauseEntity): return [] -def test_get_message_context_by_conversation_should_return_none_when_no_message() -> None: - # Arrange - session = SimpleNamespace(scalar=MagicMock(return_value=None)) - session_maker = _SessionMaker(session) - +def test_get_message_context_by_conversation_should_return_none_when_no_message( + sqlite_session_factory: sessionmaker[Session], +) -> None: # Act result = service_module._get_message_context_by_conversation( - cast(sessionmaker[Session], session_maker), + sqlite_session_factory, conversation_id="conv-1", workflow_run_id="run-1", ) @@ -375,69 +409,112 @@ def test_get_message_context_by_conversation_should_return_none_when_no_message( assert result is None -def test_get_message_context_by_conversation_should_scope_and_bound_message_lookup() -> None: +def test_get_message_context_by_conversation_should_scope_and_bound_message_lookup( + sqlite_session_factory: sessionmaker[Session], +) -> None: # Arrange - session = SimpleNamespace(scalar=MagicMock(return_value=None)) - session_maker = _SessionMaker(session) + _persist_message( + sqlite_session_factory, + message_id="older-match", + created_at=datetime(2024, 1, 1), + ) + _persist_message( + sqlite_session_factory, + message_id="newest-match", + app_id="other-app", + answer="newest answer", + created_at=datetime(2024, 1, 2), + ) + _persist_message( + sqlite_session_factory, + message_id="wrong-workflow", + workflow_run_id="run-2", + created_at=datetime(2024, 1, 3), + ) + _persist_message( + sqlite_session_factory, + message_id="wrong-conversation", + conversation_id="conv-2", + created_at=datetime(2024, 1, 4), + ) # Act - service_module._get_message_context_by_conversation( - cast(sessionmaker[Session], session_maker), + result = service_module._get_message_context_by_conversation( + sqlite_session_factory, conversation_id="conv-1", workflow_run_id="run-1", ) # Assert - stmt = session.scalar.call_args.args[0] - compiled = " ".join(str(stmt.compile(compile_kwargs={"literal_binds": True})).split()) - where_clause = compiled.split(" WHERE ", maxsplit=1)[1].split(" ORDER BY ", maxsplit=1)[0] - assert "messages.conversation_id = 'conv-1'" in compiled - assert "messages.workflow_run_id = 'run-1'" in compiled - assert "messages.app_id" not in where_clause - assert "ORDER BY messages.created_at DESC" in compiled - assert compiled.endswith("LIMIT 1") + assert result is not None + assert result.message_id == "newest-match" + assert result.conversation_id == "conv-1" + assert result.answer == "newest answer" -def test_get_message_context_by_app_should_scope_and_bound_compatibility_lookup() -> None: +def test_get_message_context_by_app_should_scope_and_bound_compatibility_lookup( + sqlite_session_factory: sessionmaker[Session], +) -> None: # Arrange - session = SimpleNamespace(scalar=MagicMock(return_value=None)) - session_maker = _SessionMaker(session) + _persist_message( + sqlite_session_factory, + message_id="older-match", + created_at=datetime(2024, 1, 1), + ) + _persist_message( + sqlite_session_factory, + message_id="newest-match", + conversation_id="conv-2", + answer="newest answer", + created_at=datetime(2024, 1, 2), + ) + _persist_message( + sqlite_session_factory, + message_id="wrong-app", + app_id="app-2", + created_at=datetime(2024, 1, 3), + ) + _persist_message( + sqlite_session_factory, + message_id="wrong-workflow", + workflow_run_id="run-2", + created_at=datetime(2024, 1, 4), + ) # Act - service_module._get_message_context_by_app( - cast(sessionmaker[Session], session_maker), + result = service_module._get_message_context_by_app( + sqlite_session_factory, app_id="app-1", workflow_run_id="run-1", ) # Assert - stmt = session.scalar.call_args.args[0] - compiled = " ".join(str(stmt.compile(compile_kwargs={"literal_binds": True})).split()) - where_clause = compiled.split(" WHERE ", maxsplit=1)[1].split(" ORDER BY ", maxsplit=1)[0] - assert "messages.app_id = 'app-1'" in where_clause - assert "messages.workflow_run_id = 'run-1'" in where_clause - assert "messages.conversation_id" not in where_clause - assert "ORDER BY messages.created_at DESC" in compiled - assert compiled.endswith("LIMIT 1") + assert result is not None + assert result.message_id == "newest-match" + assert result.conversation_id == "conv-2" + assert result.answer == "newest answer" -def test_get_message_context_by_conversation_should_default_created_at_to_zero_when_message_has_no_timestamp() -> None: +def test_get_message_context_by_conversation_should_default_created_at_to_zero_when_message_has_no_timestamp( + sqlite_session_factory: sessionmaker[Session], +) -> None: # Arrange - 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(sqlite_session_factory) + + def clear_created_at(message: Message, _context: Any) -> None: + # A load hook preserves coverage for legacy rows without replacing the real ORM query. + message.created_at = None # Act - result = service_module._get_message_context_by_conversation( - cast(sessionmaker[Session], session_maker), - conversation_id="conv-1", - workflow_run_id="run-1", - ) + orm_event.listen(Message, "load", clear_created_at) + try: + result = service_module._get_message_context_by_conversation( + sqlite_session_factory, + conversation_id="conv-1", + workflow_run_id="run-1", + ) + finally: + orm_event.remove(Message, "load", clear_created_at) # Assert assert result is not None @@ -643,6 +720,7 @@ def test_start_buffering_should_set_done_event_when_subscription_raises() -> Non def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_event( monkeypatch: pytest.MonkeyPatch, + unbound_session_factory: sessionmaker[Session], ) -> None: # Arrange workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED) @@ -687,7 +765,7 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even "_build_snapshot_events", MagicMock(return_value=[{"event": StreamEvent.WORKFLOW_FINISHED, "task_id": "task-1"}]), ) - session_maker = MagicMock() + session_maker = unbound_session_factory # Act events = list( @@ -740,6 +818,7 @@ def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_even ) def test_build_advanced_chat_snapshot_requires_conversation_context( monkeypatch: pytest.MonkeyPatch, + unbound_session_factory: sessionmaker[Session], resumption_context: WorkflowResumptionContext | None, expected_error: str, ) -> None: @@ -771,7 +850,7 @@ def test_build_advanced_chat_snapshot_requires_conversation_context( workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=MagicMock(), + session_maker=unbound_session_factory, ) conversation_lookup.assert_not_called() app_lookup.assert_not_called() @@ -779,6 +858,7 @@ def test_build_advanced_chat_snapshot_requires_conversation_context( def test_build_non_suspended_advanced_chat_snapshot_uses_app_scoped_fallback( monkeypatch: pytest.MonkeyPatch, + unbound_session_factory: sessionmaker[Session], ) -> None: # Arrange workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING) @@ -800,7 +880,7 @@ def test_build_non_suspended_advanced_chat_snapshot_uses_app_scoped_fallback( conversation_lookup, ) monkeypatch.setattr(service_module, "_get_message_context_by_app", app_lookup) - session_maker = MagicMock() + session_maker = unbound_session_factory # Act event_stream = build_workflow_event_stream( @@ -825,6 +905,7 @@ def test_build_non_suspended_advanced_chat_snapshot_uses_app_scoped_fallback( def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_idle_timeout( monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], ) -> None: # Arrange workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING) @@ -866,7 +947,7 @@ def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_id workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=MagicMock(), + session_maker=sqlite_session_factory, idle_timeout=20.0, ping_interval=5.0, ) @@ -879,6 +960,7 @@ def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_id def test_build_workflow_event_stream_should_exit_when_buffer_done_and_empty( monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], ) -> None: # Arrange workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.RUNNING) @@ -911,7 +993,7 @@ def test_build_workflow_event_stream_should_exit_when_buffer_done_and_empty( workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=MagicMock(), + session_maker=sqlite_session_factory, ) ) @@ -922,6 +1004,7 @@ def test_build_workflow_event_stream_should_exit_when_buffer_done_and_empty( def test_build_workflow_event_stream_should_continue_when_pause_loading_fails( monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], ) -> None: # Arrange workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED) @@ -954,7 +1037,7 @@ def test_build_workflow_event_stream_should_continue_when_pause_loading_fails( workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=MagicMock(), + session_maker=sqlite_session_factory, ) ) @@ -972,7 +1055,10 @@ def test_is_terminal_event_respects_close_on_pause_flag() -> None: assert _is_terminal_event(finish_event, close_on_pause=False) is True -def test_build_snapshot_events_preserves_public_form_token(monkeypatch: pytest.MonkeyPatch) -> None: +def test_build_snapshot_events_preserves_public_form_token( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], +) -> None: workflow_run = _build_workflow_run(WorkflowExecutionStatus.PAUSED) snapshot = _build_snapshot(WorkflowNodeExecutionStatus.PAUSED) resumption_context = _build_resumption_context("task-ctx") @@ -983,11 +1069,7 @@ def test_build_snapshot_events_preserves_public_form_token(monkeypatch: pytest.M "form-1": FormDisposition(form_token="wtok", approval_channels=[]) }, ) - session_maker = _SessionMaker( - SimpleNamespace( - execute=lambda _stmt: [("form-1", datetime(2024, 1, 1, tzinfo=UTC), '{"display_in_ui": true}')], - ) - ) + expiration_time = _persist_human_input_form(sqlite_session_factory) pause_entity = _FakePauseEntity( pause_id="pause-1", workflow_run_id="run-1", @@ -1010,34 +1092,30 @@ def test_build_snapshot_events_preserves_public_form_token(monkeypatch: pytest.M message_context=None, pause_entity=pause_entity, resumption_context=resumption_context, - session_maker=cast(sessionmaker[Session], session_maker), + session_maker=sqlite_session_factory, ) assert events[-2]["event"] == StreamEvent.HUMAN_INPUT_REQUIRED assert events[-2]["data"]["form_token"] == "wtok" - assert events[-2]["data"]["expiration_time"] == int(datetime(2024, 1, 1, tzinfo=UTC).timestamp()) + assert events[-2]["data"]["expiration_time"] == to_utc_timestamp(expiration_time) pause_data = events[-1]["data"] 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"] == to_utc_timestamp(expiration_time) -def _build_recipient_snapshot_events(recipients: Sequence[Any]) -> list[Mapping[str, Any]]: +def _build_recipient_snapshot_events( + session_maker: sessionmaker[Session], + recipients: Sequence[HumanInputFormRecipient], +) -> list[Mapping[str, Any]]: """Drive the reconnect snapshot pause path for the OPENAPI surface. - Lets the real disposition loader run against a fake session whose ``scalars`` - yields the given recipients, so the reconnect path derives the same token and - approval channels as the live path for the same recipient set. + Persisting the recipients lets the real disposition query derive the same token + and approval channels as the live path for the same recipient set. """ 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, tzinfo=UTC) - session_maker = _SessionMaker( - SimpleNamespace( - execute=lambda _stmt: [("form-1", expiration_time, '{"display_in_ui": true}')], - scalars=lambda _stmt: list(recipients), - ) - ) + expiration_time = _persist_human_input_form(session_maker, recipients=recipients) pause_entity = _FakePauseEntity( pause_id="pause-1", workflow_run_id="run-1", @@ -1059,16 +1137,31 @@ def _build_recipient_snapshot_events(recipients: Sequence[Any]) -> list[Mapping[ message_context=None, pause_entity=pause_entity, resumption_context=resumption_context, - session_maker=cast(sessionmaker[Session], session_maker), + session_maker=session_maker, human_input_surface=HumanInputSurface.OPENAPI, ) -def test_reconnect_pause_without_web_app_recipient_emits_approval_channels() -> None: +def test_reconnect_pause_without_web_app_recipient_emits_approval_channels( + sqlite_session_factory: sessionmaker[Session], +) -> None: events = _build_recipient_snapshot_events( + sqlite_session_factory, recipients=[ - SimpleNamespace(form_id="form-1", recipient_type=RecipientType.EMAIL_MEMBER, access_token="email-token"), - SimpleNamespace(form_id="form-1", recipient_type=RecipientType.BACKSTAGE, access_token="backstage-token"), + HumanInputFormRecipient( + form_id="form-1", + delivery_id="delivery-1", + recipient_type=RecipientType.EMAIL_MEMBER, + recipient_payload="{}", + access_token="email-token", + ), + HumanInputFormRecipient( + form_id="form-1", + delivery_id="delivery-2", + recipient_type=RecipientType.BACKSTAGE, + recipient_payload="{}", + access_token="backstage-token", + ), ], ) @@ -1082,15 +1175,26 @@ def test_reconnect_pause_without_web_app_recipient_emits_approval_channels() -> assert pause_data["reasons"][0]["approval_channels"] == ["console", "email"] -def test_reconnect_pause_with_web_app_recipient_sets_token_and_channels() -> None: +def test_reconnect_pause_with_web_app_recipient_sets_token_and_channels( + sqlite_session_factory: sessionmaker[Session], +) -> None: events = _build_recipient_snapshot_events( + sqlite_session_factory, recipients=[ - SimpleNamespace( + HumanInputFormRecipient( form_id="form-1", + delivery_id="delivery-1", recipient_type=RecipientType.STANDALONE_WEB_APP, + recipient_payload="{}", access_token="web-app-token", ), - SimpleNamespace(form_id="form-1", recipient_type=RecipientType.BACKSTAGE, access_token="backstage-token"), + HumanInputFormRecipient( + form_id="form-1", + delivery_id="delivery-2", + recipient_type=RecipientType.BACKSTAGE, + recipient_payload="{}", + access_token="backstage-token", + ), ], ) @@ -1104,7 +1208,10 @@ def test_reconnect_pause_with_web_app_recipient_sets_token_and_channels() -> Non assert pause_data["reasons"][0]["approval_channels"] == ["console"] -def test_build_snapshot_events_resolves_pause_reason_select_options(monkeypatch: pytest.MonkeyPatch) -> None: +def test_build_snapshot_events_resolves_pause_reason_select_options( + monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], +) -> None: workflow_run = _build_workflow_run(WorkflowExecutionStatus.PAUSED) snapshot = _build_snapshot(WorkflowNodeExecutionStatus.PAUSED) resumption_context = _build_resumption_context("task-ctx", select_options=["approve", "reject"]) @@ -1115,11 +1222,7 @@ def test_build_snapshot_events_resolves_pause_reason_select_options(monkeypatch: "form-1": FormDisposition(form_token="wtok", approval_channels=[]) }, ) - session_maker = _SessionMaker( - SimpleNamespace( - execute=lambda _stmt: [("form-1", datetime(2024, 1, 1, tzinfo=UTC), '{"display_in_ui": true}')], - ) - ) + _persist_human_input_form(sqlite_session_factory) pause_entity = _FakePauseEntity( pause_id="pause-1", workflow_run_id="run-1", @@ -1151,7 +1254,7 @@ def test_build_snapshot_events_resolves_pause_reason_select_options(monkeypatch: message_context=None, pause_entity=pause_entity, resumption_context=resumption_context, - session_maker=cast(sessionmaker[Session], session_maker), + session_maker=sqlite_session_factory, ) human_input_event = events[-2] @@ -1163,6 +1266,7 @@ def test_build_snapshot_events_resolves_pause_reason_select_options(monkeypatch: def test_build_workflow_event_stream_loads_pause_tokens_without_flask_app_context( monkeypatch: pytest.MonkeyPatch, + sqlite_session_factory: sessionmaker[Session], ) -> None: workflow_run = _build_workflow_run_additional(status=WorkflowExecutionStatus.PAUSED) topic = _Topic(_StaticSubscription()) @@ -1198,11 +1302,7 @@ def test_build_workflow_event_stream_loads_pause_tokens_without_flask_app_contex }, ) - session = SimpleNamespace( - scalar=MagicMock(return_value=None), - execute=lambda _stmt: [("form-1", datetime(2024, 1, 1, tzinfo=UTC), '{"display_in_ui": true}')], - ) - session_maker = _SessionMaker(session) + expiration_time = _persist_human_input_form(sqlite_session_factory) events = list( build_workflow_event_stream( @@ -1210,11 +1310,11 @@ def test_build_workflow_event_stream_loads_pause_tokens_without_flask_app_contex workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=cast(sessionmaker[Session], session_maker), + session_maker=sqlite_session_factory, ) ) pause_event = cast(Mapping[str, Any], events[-1]) assert pause_event["event"] == StreamEvent.WORKFLOW_PAUSED assert pause_event["data"]["reasons"][0]["form_token"] == "wtok" - assert pause_event["data"]["reasons"][0]["expiration_time"] == int(datetime(2024, 1, 1, tzinfo=UTC).timestamp()) + assert pause_event["data"]["reasons"][0]["expiration_time"] == to_utc_timestamp(expiration_time) diff --git a/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service_additional.py b/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service_additional.py index 8efd7370a73..36cdac0742e 100644 --- a/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service_additional.py +++ b/api/tests/unit_tests/services/workflow/test_workflow_event_snapshot_service_additional.py @@ -74,7 +74,7 @@ def _build_resumption_context(task_id: str) -> WorkflowResumptionContext: workflow_execution_id="run-1", ) runtime_state = GraphRuntimeState(variable_pool=VariablePool(), start_at=0.0) - runtime_state.outputs = {"answer": "ok"} + runtime_state.set_output("answer", "ok") wrapper = _WorkflowGenerateEntityWrapper(entity=generate_entity) return WorkflowResumptionContext( generate_entity=wrapper, @@ -362,6 +362,7 @@ class TestBuildWorkflowEventStream: def test_build_workflow_event_stream_should_emit_ping_and_terminal_snapshot_event( self, monkeypatch: pytest.MonkeyPatch, + message_session_maker: sessionmaker[Session], ) -> None: workflow_run = _build_workflow_run(status=WorkflowExecutionStatus.PAUSED) topic = _Topic(_StaticSubscription()) @@ -409,7 +410,7 @@ class TestBuildWorkflowEventStream: workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=MagicMock(), + session_maker=message_session_maker, ) ) @@ -424,6 +425,7 @@ class TestBuildWorkflowEventStream: def test_build_workflow_event_stream_should_emit_periodic_ping_and_stop_after_idle_timeout( self, monkeypatch: pytest.MonkeyPatch, + message_session_maker: sessionmaker[Session], ) -> None: workflow_run = _build_workflow_run(status=WorkflowExecutionStatus.RUNNING) topic = _Topic(_StaticSubscription()) @@ -463,7 +465,7 @@ class TestBuildWorkflowEventStream: workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=MagicMock(), + session_maker=message_session_maker, idle_timeout=20.0, ping_interval=5.0, ) @@ -475,6 +477,7 @@ class TestBuildWorkflowEventStream: def test_build_workflow_event_stream_should_exit_when_buffer_done_and_empty( self, monkeypatch: pytest.MonkeyPatch, + message_session_maker: sessionmaker[Session], ) -> None: workflow_run = _build_workflow_run(status=WorkflowExecutionStatus.RUNNING) topic = _Topic(_StaticSubscription()) @@ -505,7 +508,7 @@ class TestBuildWorkflowEventStream: workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=MagicMock(), + session_maker=message_session_maker, ) ) @@ -515,6 +518,7 @@ class TestBuildWorkflowEventStream: def test_build_workflow_event_stream_should_continue_when_pause_loading_fails( self, monkeypatch: pytest.MonkeyPatch, + message_session_maker: sessionmaker[Session], ) -> None: workflow_run = _build_workflow_run(status=WorkflowExecutionStatus.PAUSED) topic = _Topic(_StaticSubscription()) @@ -545,7 +549,7 @@ class TestBuildWorkflowEventStream: workflow_run=workflow_run, tenant_id="tenant-1", app_id="app-1", - session_maker=MagicMock(), + session_maker=message_session_maker, ) ) diff --git a/api/tests/unit_tests/tasks/test_agent_backend_session_cleanup_task.py b/api/tests/unit_tests/tasks/test_agent_backend_session_cleanup_task.py deleted file mode 100644 index f7bc8a4a9ac..00000000000 --- a/api/tests/unit_tests/tasks/test_agent_backend_session_cleanup_task.py +++ /dev/null @@ -1,51 +0,0 @@ -import logging - -from agenton.compositor import CompositorSessionSnapshot - -from clients.agent_backend.session_cleanup import ( - AgentBackendSessionCleanupPayload, - AgentBackendSessionCleanupResult, -) -from tasks import agent_backend_session_cleanup_task as cleanup_task_module - - -def _payload_dict() -> dict[str, object]: - return AgentBackendSessionCleanupPayload( - session_snapshot=CompositorSessionSnapshot(layers=[]), - runtime_layer_specs=[], - metadata={"tenant_id": "tenant-1", "app_id": "app-1"}, - ).model_dump(mode="json") - - -def test_run_cleanup_task_logs_info_for_skipped_result(monkeypatch, caplog): - monkeypatch.setattr(cleanup_task_module, "_create_agent_backend_client", lambda: object()) - monkeypatch.setattr( - cleanup_task_module, - "cleanup_agent_backend_session", - lambda **kwargs: AgentBackendSessionCleanupResult.skipped("missing_runtime_layer_specs"), - ) - - with caplog.at_level(logging.INFO, logger="tasks.agent_backend_session_cleanup_task"): - cleanup_task_module._run_cleanup_task(_payload_dict()) - - assert "Agent backend session cleanup skipped" in caplog.text - assert "missing_runtime_layer_specs" in caplog.text - - -def test_run_cleanup_task_logs_warning_for_failed_result(monkeypatch, caplog): - monkeypatch.setattr(cleanup_task_module, "_create_agent_backend_client", lambda: object()) - monkeypatch.setattr( - cleanup_task_module, - "cleanup_agent_backend_session", - lambda **kwargs: AgentBackendSessionCleanupResult.failed( - "backend exploded", - cleanup_run_id="cleanup-run-1", - ), - ) - - with caplog.at_level(logging.WARNING, logger="tasks.agent_backend_session_cleanup_task"): - cleanup_task_module._run_cleanup_task(_payload_dict()) - - assert "Agent backend session cleanup failed" in caplog.text - assert "backend exploded" in caplog.text - assert "cleanup-run-1" in caplog.text diff --git a/api/tests/unit_tests/tasks/test_batch_clean_document_task.py b/api/tests/unit_tests/tasks/test_batch_clean_document_task.py index 6386c72188e..59fab89ab9c 100644 --- a/api/tests/unit_tests/tasks/test_batch_clean_document_task.py +++ b/api/tests/unit_tests/tasks/test_batch_clean_document_task.py @@ -1,46 +1,73 @@ -from unittest.mock import MagicMock, patch +import uuid +from unittest.mock import patch +import pytest +from sqlalchemy.orm import Session + +import tasks.batch_clean_document_task as task_module +from models.dataset import Dataset, DocumentSegment +from models.enums import DataSourceType from tasks.batch_clean_document_task import batch_clean_document_task -def _setup_cleanup_dependencies(): - session = MagicMock() - segment = MagicMock(id="segment-1", index_node_id="node-1", content="content") - dataset = MagicMock(id="dataset-1", tenant_id="tenant-1") - session.scalars.return_value.all.return_value = [segment] - session.scalar.return_value = dataset - - context_manager = MagicMock() - context_manager.__enter__.return_value = session - context_manager.__exit__.return_value = None - return session, context_manager +@pytest.fixture +def cleanup_rows(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> tuple[str, str, str]: + tenant_id = str(uuid.uuid4()) + dataset_id = str(uuid.uuid4()) + document_id = str(uuid.uuid4()) + created_by = str(uuid.uuid4()) + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Batch cleanup dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + ) + segment = DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset_id, + document_id=document_id, + position=1, + content="content", + word_count=1, + tokens=1, + created_by=created_by, + index_node_id="node-1", + ) + sqlite_session.add_all([dataset, segment]) + sqlite_session.commit() + engine = sqlite_session.get_bind() + monkeypatch.setattr( + task_module.session_factory, + "create_session", + lambda: Session(engine, expire_on_commit=False), + ) + return dataset_id, document_id, tenant_id -def test_successful_vector_cleanup_schedules_billing_refresh(): - _, context_manager = _setup_cleanup_dependencies() +def test_successful_vector_cleanup_schedules_billing_refresh(cleanup_rows: tuple[str, str, str]): + dataset_id, document_id, tenant_id = cleanup_rows with ( - patch("tasks.batch_clean_document_task.session_factory.create_session", return_value=context_manager), patch("tasks.batch_clean_document_task.get_image_upload_file_ids", return_value=[]), patch("tasks.batch_clean_document_task.IndexProcessorFactory") as processor_factory, patch("tasks.batch_clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh, ): batch_clean_document_task( - document_ids=["document-1"], - dataset_id="dataset-1", + document_ids=[document_id], + dataset_id=dataset_id, doc_form="paragraph", file_ids=[], ) processor_factory.return_value.init_index_processor.return_value.clean.assert_called_once() - schedule_refresh.assert_called_once_with("tenant-1") + schedule_refresh.assert_called_once_with(tenant_id) -def test_failed_vector_cleanup_does_not_schedule_billing_refresh(): - _, context_manager = _setup_cleanup_dependencies() +def test_failed_vector_cleanup_does_not_schedule_billing_refresh(cleanup_rows: tuple[str, str, str]): + dataset_id, document_id, _tenant_id = cleanup_rows with ( - patch("tasks.batch_clean_document_task.session_factory.create_session", return_value=context_manager), patch("tasks.batch_clean_document_task.get_image_upload_file_ids", return_value=[]), patch("tasks.batch_clean_document_task.IndexProcessorFactory") as processor_factory, patch("tasks.batch_clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh, @@ -49,8 +76,8 @@ def test_failed_vector_cleanup_does_not_schedule_billing_refresh(): "vector cleanup failed" ) batch_clean_document_task( - document_ids=["document-1"], - dataset_id="dataset-1", + document_ids=[document_id], + dataset_id=dataset_id, doc_form="paragraph", file_ids=[], ) diff --git a/api/tests/unit_tests/tasks/test_clean_dataset_task.py b/api/tests/unit_tests/tasks/test_clean_dataset_task.py index 4a2dcc44145..3ac74ec1f4f 100644 --- a/api/tests/unit_tests/tasks/test_clean_dataset_task.py +++ b/api/tests/unit_tests/tasks/test_clean_dataset_task.py @@ -11,13 +11,33 @@ This module tests the dataset cleanup task functionality including: - Segment attachment cleanup """ +import json import uuid +from collections.abc import Iterator +from datetime import UTC, datetime from unittest.mock import MagicMock, patch import pytest +from sqlalchemy import Engine, event +from sqlalchemy.orm import Session, sessionmaker from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType -from models.enums import DataSourceType +from extensions.storage.storage_type import StorageType +from models.base import TypeBase +from models.dataset import ( + AppDatasetJoin, + DatasetMetadata, + DatasetMetadataBinding, + DatasetProcessRule, + DatasetQuery, + Document, + DocumentSegment, + Pipeline, + SegmentAttachmentBinding, +) +from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, IndexingStatus +from models.model import UploadFile +from models.workflow import Workflow, WorkflowType from tasks.clean_dataset_task import clean_dataset_task # ============================================================================ @@ -26,54 +46,54 @@ from tasks.clean_dataset_task import clean_dataset_task @pytest.fixture -def tenant_id(): +def tenant_id() -> str: """Generate a unique tenant ID for testing.""" return str(uuid.uuid4()) @pytest.fixture -def dataset_id(): +def dataset_id() -> str: """Generate a unique dataset ID for testing.""" return str(uuid.uuid4()) @pytest.fixture -def collection_binding_id(): +def collection_binding_id() -> str: """Generate a unique collection binding ID for testing.""" return str(uuid.uuid4()) @pytest.fixture -def pipeline_id(): +def pipeline_id() -> str: """Generate a unique pipeline ID for testing.""" return str(uuid.uuid4()) @pytest.fixture -def mock_db_session(): - """Mock database session via session_factory.create_session().""" - with patch("tasks.clean_dataset_task.session_factory", autospec=True) as mock_sf: - mock_session = MagicMock() - # context manager for create_session() - cm = MagicMock() - cm.__enter__.return_value = mock_session - cm.__exit__.return_value = None - mock_sf.create_session.return_value = cm - - # Setup scalars for select queries - mock_session.scalars.return_value.all.return_value = [] - - # Setup execute for JOIN queries - mock_session.execute.return_value.all.return_value = [] - - # Yield an object with a `.session` attribute to keep tests unchanged - wrapper = MagicMock() - wrapper.session = mock_session - yield wrapper +def orm_session_maker( + sqlite_engine: Engine, + sqlite_session_factory: sessionmaker[Session], +) -> sessionmaker[Session]: + """Create the cleanup tables and return the suite's real SQLite session factory.""" + models = ( + Document, + DocumentSegment, + SegmentAttachmentBinding, + UploadFile, + DatasetProcessRule, + DatasetQuery, + AppDatasetJoin, + DatasetMetadata, + DatasetMetadataBinding, + Pipeline, + Workflow, + ) + TypeBase.metadata.create_all(sqlite_engine, tables=[model.__table__ for model in models]) + return sqlite_session_factory @pytest.fixture -def mock_storage(): +def mock_storage() -> Iterator[MagicMock]: """Mock storage client.""" with patch("tasks.clean_dataset_task.storage", autospec=True) as mock_storage: mock_storage.delete.return_value = None @@ -81,7 +101,7 @@ def mock_storage(): @pytest.fixture -def mock_index_processor_factory(): +def mock_index_processor_factory() -> Iterator[dict[str, MagicMock]]: """Mock IndexProcessorFactory.""" with patch("tasks.clean_dataset_task.IndexProcessorFactory", autospec=True) as mock_factory: mock_processor = MagicMock() @@ -98,42 +118,104 @@ def mock_index_processor_factory(): @pytest.fixture -def mock_get_image_upload_file_ids(): +def mock_get_image_upload_file_ids() -> Iterator[MagicMock]: """Mock get_image_upload_file_ids function.""" with patch("tasks.clean_dataset_task.get_image_upload_file_ids", autospec=True) as mock_func: mock_func.return_value = [] yield mock_func -@pytest.fixture -def mock_document(): - """Create a mock Document object.""" - doc = MagicMock() - doc.id = str(uuid.uuid4()) - doc.tenant_id = str(uuid.uuid4()) - doc.dataset_id = str(uuid.uuid4()) - doc.data_source_type = DataSourceType.UPLOAD_FILE - doc.data_source_info = '{"upload_file_id": "test-file-id"}' - doc.data_source_info_dict = {"upload_file_id": "test-file-id"} - return doc +def _run_clean_dataset( + *, + dataset_id: str, + tenant_id: str, + collection_binding_id: str, + pipeline_id: str | None = None, +) -> None: + clean_dataset_task( + dataset_id=dataset_id, + tenant_id=tenant_id, + indexing_technique=IndexTechniqueType.HIGH_QUALITY, + index_struct='{"type": "paragraph"}', + collection_binding_id=collection_binding_id, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + pipeline_id=pipeline_id, + ) -@pytest.fixture -def mock_segment(): - """Create a mock DocumentSegment object.""" - segment = MagicMock() - segment.id = str(uuid.uuid4()) - segment.content = "Test segment content" - return segment +def _persist_document(session_maker: sessionmaker[Session], *, dataset_id: str, tenant_id: str) -> Document: + document = Document( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + dataset_id=dataset_id, + position=1, + data_source_type=DataSourceType.LOCAL_FILE, + batch="batch", + name="Document", + created_from=DocumentCreatedFrom.API, + created_by=str(uuid.uuid4()), + indexing_status=IndexingStatus.COMPLETED, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + with session_maker.begin() as session: + session.add(document) + return document -@pytest.fixture -def mock_upload_file(): - """Create a mock UploadFile object.""" - upload_file = MagicMock() - upload_file.id = str(uuid.uuid4()) - upload_file.key = f"test_files/{uuid.uuid4()}.txt" - return upload_file +def _persist_pipeline_and_workflow( + session_maker: sessionmaker[Session], + *, + pipeline_id: str, + tenant_id: str, +) -> Workflow: + pipeline = Pipeline(tenant_id=tenant_id, name="Pipeline", description="Pipeline") + pipeline.id = pipeline_id + workflow = Workflow.new( + tenant_id=tenant_id, + app_id=pipeline_id, + type=WorkflowType.RAG_PIPELINE.value, + version="v1", + graph=json.dumps({"nodes": [], "edges": []}), + features="{}", + created_by=str(uuid.uuid4()), + environment_variables=[], + conversation_variables=[], + rag_pipeline_variables=[], + ) + with session_maker.begin() as session: + session.add_all([pipeline, workflow]) + return workflow + + +def _persist_attachment( + session_maker: sessionmaker[Session], + *, + dataset_id: str, + tenant_id: str, +) -> tuple[SegmentAttachmentBinding, UploadFile]: + attachment_file = UploadFile( + tenant_id=tenant_id, + storage_type=StorageType.LOCAL, + key=f"attachments/{uuid.uuid4()}.pdf", + name="attachment.pdf", + size=10, + extension="pdf", + mime_type="application/pdf", + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid.uuid4()), + created_at=datetime.now(UTC), + used=False, + ) + binding = SegmentAttachmentBinding( + tenant_id=tenant_id, + dataset_id=dataset_id, + document_id=str(uuid.uuid4()), + segment_id=str(uuid.uuid4()), + attachment_id=attachment_file.id, + ) + with session_maker.begin() as session: + session.add_all([attachment_file, binding]) + return binding, attachment_file # ============================================================================ @@ -154,38 +236,50 @@ class TestErrorHandling: dataset_id: str, tenant_id: str, collection_binding_id: str, - mock_db_session, - mock_storage, - mock_index_processor_factory, - mock_get_image_upload_file_ids, + orm_session_maker: sessionmaker[Session], + mock_storage: MagicMock, + mock_index_processor_factory: dict[str, MagicMock], + mock_get_image_upload_file_ids: MagicMock, ): """ Test that session is closed even if rollback fails. Scenario: - Database commit fails - - Rollback also fails - - Session should still be closed + - Rollback completes, then its event hook raises + - Session cleanup should still make the factory reusable Expected behavior: - - Session.close() is called regardless of rollback failure + - The database rollback preserves the pending row + - The task session closes and a new session remains usable """ - # Arrange - mock_db_session.session.commit.side_effect = Exception("Commit failed") - mock_db_session.session.rollback.side_effect = Exception("Rollback failed") + # Arrange: persist a row whose attempted deletion must be rolled back. + document = _persist_document(orm_session_maker, dataset_id=dataset_id, tenant_id=tenant_id) + + def fail_commit(_session: Session) -> None: + raise RuntimeError("Commit failed") + + def fail_rollback(_session: Session) -> None: + raise RuntimeError("Rollback failed") + + event.listen(orm_session_maker.class_, "before_commit", fail_commit) + event.listen(orm_session_maker.class_, "after_rollback", fail_rollback) # Act - clean_dataset_task( - dataset_id=dataset_id, - tenant_id=tenant_id, - indexing_technique=IndexTechniqueType.HIGH_QUALITY, - index_struct='{"type": "paragraph"}', - collection_binding_id=collection_binding_id, - doc_form=IndexStructureType.PARAGRAPH_INDEX, - ) + try: + _run_clean_dataset( + dataset_id=dataset_id, + tenant_id=tenant_id, + collection_binding_id=collection_binding_id, + ) + finally: + event.remove(orm_session_maker.class_, "before_commit", fail_commit) + event.remove(orm_session_maker.class_, "after_rollback", fail_rollback) - # Assert - mock_db_session.session.close.assert_called_once() + # Assert: rollback happened before its hook failed, and the closed task + # session did not prevent a new real session from reading the row. + with orm_session_maker() as session: + assert session.get(Document, document.id) is not None # ============================================================================ @@ -201,11 +295,11 @@ class TestPipelineAndWorkflowDeletion: dataset_id: str, tenant_id: str, collection_binding_id: str, - pipeline_id, - mock_db_session, - mock_storage, - mock_index_processor_factory, - mock_get_image_upload_file_ids, + pipeline_id: str, + orm_session_maker: sessionmaker[Session], + mock_storage: MagicMock, + mock_index_processor_factory: dict[str, MagicMock], + mock_get_image_upload_file_ids: MagicMock, ): """ Test that pipeline and workflow are deleted when pipeline_id is provided. @@ -214,30 +308,35 @@ class TestPipelineAndWorkflowDeletion: - Pipeline record is deleted - Related workflow record is deleted """ + workflow = _persist_pipeline_and_workflow( + orm_session_maker, + pipeline_id=pipeline_id, + tenant_id=tenant_id, + ) + # Act - clean_dataset_task( + _run_clean_dataset( dataset_id=dataset_id, tenant_id=tenant_id, - indexing_technique=IndexTechniqueType.HIGH_QUALITY, - index_struct='{"type": "paragraph"}', collection_binding_id=collection_binding_id, - doc_form=IndexStructureType.PARAGRAPH_INDEX, pipeline_id=pipeline_id, ) - # Assert - verify execute was called for delete operations - # 1 attachment JOIN query + 5 base deletes + 2 pipeline/workflow deletes = 8 - assert mock_db_session.session.execute.call_count >= 8 + # Assert + with orm_session_maker() as session: + assert session.get(Pipeline, pipeline_id) is None + assert session.get(Workflow, workflow.id) is None def test_clean_dataset_task_without_pipeline_id( self, dataset_id: str, tenant_id: str, collection_binding_id: str, - mock_db_session, - mock_storage, - mock_index_processor_factory, - mock_get_image_upload_file_ids, + pipeline_id: str, + orm_session_maker: sessionmaker[Session], + mock_storage: MagicMock, + mock_index_processor_factory: dict[str, MagicMock], + mock_get_image_upload_file_ids: MagicMock, ): """ Test that pipeline/workflow deletion is skipped when pipeline_id is None. @@ -245,20 +344,24 @@ class TestPipelineAndWorkflowDeletion: Expected behavior: - Pipeline and workflow deletion queries are not executed """ + workflow = _persist_pipeline_and_workflow( + orm_session_maker, + pipeline_id=pipeline_id, + tenant_id=tenant_id, + ) + # Act - clean_dataset_task( + _run_clean_dataset( dataset_id=dataset_id, tenant_id=tenant_id, - indexing_technique=IndexTechniqueType.HIGH_QUALITY, - index_struct='{"type": "paragraph"}', collection_binding_id=collection_binding_id, - doc_form=IndexStructureType.PARAGRAPH_INDEX, pipeline_id=None, ) - # Assert - verify execute was called for delete operations - # 1 attachment JOIN query + 5 base deletes = 6 - assert mock_db_session.session.execute.call_count == 6 + # Assert + with orm_session_maker() as session: + assert session.get(Pipeline, pipeline_id) is not None + assert session.get(Workflow, workflow.id) is not None # ============================================================================ @@ -274,10 +377,10 @@ class TestSegmentAttachmentCleanup: dataset_id: str, tenant_id: str, collection_binding_id: str, - mock_db_session, - mock_storage, - mock_index_processor_factory, - mock_get_image_upload_file_ids, + orm_session_maker: sessionmaker[Session], + mock_storage: MagicMock, + mock_index_processor_factory: dict[str, MagicMock], + mock_get_image_upload_file_ids: MagicMock, ): """ Test that segment attachments are cleaned up properly. @@ -292,42 +395,34 @@ class TestSegmentAttachmentCleanup: - Binding records are deleted from database """ # Arrange - mock_binding = MagicMock() - mock_binding.attachment_id = str(uuid.uuid4()) - - mock_attachment_file = MagicMock() - mock_attachment_file.id = mock_binding.attachment_id - mock_attachment_file.key = f"attachments/{uuid.uuid4()}.pdf" - - # Setup execute to return attachment with binding - mock_db_session.session.execute.return_value.all.return_value = [(mock_binding, mock_attachment_file)] - - # Act - clean_dataset_task( + binding, attachment_file = _persist_attachment( + orm_session_maker, + dataset_id=dataset_id, + tenant_id=tenant_id, + ) + + # Act + _run_clean_dataset( dataset_id=dataset_id, tenant_id=tenant_id, - indexing_technique=IndexTechniqueType.HIGH_QUALITY, - index_struct='{"type": "paragraph"}', collection_binding_id=collection_binding_id, - doc_form=IndexStructureType.PARAGRAPH_INDEX, ) # Assert - mock_storage.delete.assert_called_with(mock_attachment_file.key) - # Attachment file and binding are deleted in batch; verify DELETEs were issued - execute_sqls = [" ".join(str(c[0][0]).split()) for c in mock_db_session.session.execute.call_args_list] - assert any("DELETE FROM upload_files" in sql for sql in execute_sqls) - assert any("DELETE FROM segment_attachment_bindings" in sql for sql in execute_sqls) + mock_storage.delete.assert_called_once_with(attachment_file.key) + with orm_session_maker() as session: + assert session.get(UploadFile, attachment_file.id) is None + assert session.get(SegmentAttachmentBinding, binding.id) is None def test_clean_dataset_task_attachment_storage_failure( self, dataset_id: str, tenant_id: str, collection_binding_id: str, - mock_db_session, - mock_storage, - mock_index_processor_factory, - mock_get_image_upload_file_ids, + orm_session_maker: sessionmaker[Session], + mock_storage: MagicMock, + mock_index_processor_factory: dict[str, MagicMock], + mock_get_image_upload_file_ids: MagicMock, ): """ Test that cleanup continues even if attachment storage deletion fails. @@ -337,32 +432,25 @@ class TestSegmentAttachmentCleanup: - Attachment file and binding are still deleted from database """ # Arrange - mock_binding = MagicMock() - mock_binding.attachment_id = str(uuid.uuid4()) - - mock_attachment_file = MagicMock() - mock_attachment_file.id = mock_binding.attachment_id - mock_attachment_file.key = f"attachments/{uuid.uuid4()}.pdf" - - mock_db_session.session.execute.return_value.all.return_value = [(mock_binding, mock_attachment_file)] + binding, attachment_file = _persist_attachment( + orm_session_maker, + dataset_id=dataset_id, + tenant_id=tenant_id, + ) mock_storage.delete.side_effect = Exception("Storage error") # Act - clean_dataset_task( + _run_clean_dataset( dataset_id=dataset_id, tenant_id=tenant_id, - indexing_technique=IndexTechniqueType.HIGH_QUALITY, - index_struct='{"type": "paragraph"}', collection_binding_id=collection_binding_id, - doc_form=IndexStructureType.PARAGRAPH_INDEX, ) # Assert - storage delete was attempted - mock_storage.delete.assert_called_once() - # Records are deleted in batch; verify DELETEs were issued - execute_sqls = [" ".join(str(c[0][0]).split()) for c in mock_db_session.session.execute.call_args_list] - assert any("DELETE FROM upload_files" in sql for sql in execute_sqls) - assert any("DELETE FROM segment_attachment_bindings" in sql for sql in execute_sqls) + mock_storage.delete.assert_called_once_with(attachment_file.key) + with orm_session_maker() as session: + assert session.get(UploadFile, attachment_file.id) is None + assert session.get(SegmentAttachmentBinding, binding.id) is None # ============================================================================ @@ -373,34 +461,35 @@ class TestSegmentAttachmentCleanup: class TestEdgeCases: """Test edge cases and boundary conditions.""" - def test_clean_dataset_task_session_always_closed( + def test_clean_dataset_task_commits_cleanup_and_factory_remains_usable( self, dataset_id: str, tenant_id: str, collection_binding_id: str, - mock_db_session, - mock_storage, - mock_index_processor_factory, - mock_get_image_upload_file_ids, + orm_session_maker: sessionmaker[Session], + mock_storage: MagicMock, + mock_index_processor_factory: dict[str, MagicMock], + mock_get_image_upload_file_ids: MagicMock, ): """ - Test that database session is always closed regardless of success or failure. + Test that cleanup commits and the task-owned session releases its resources. Expected behavior: - - Session.close() is called in finally block + - The document deletion is committed + - A subsequent real session can use the same factory """ + document = _persist_document(orm_session_maker, dataset_id=dataset_id, tenant_id=tenant_id) + # Act - clean_dataset_task( + _run_clean_dataset( dataset_id=dataset_id, tenant_id=tenant_id, - indexing_technique=IndexTechniqueType.HIGH_QUALITY, - index_struct='{"type": "paragraph"}', collection_binding_id=collection_binding_id, - doc_form=IndexStructureType.PARAGRAPH_INDEX, ) # Assert - mock_db_session.session.close.assert_called_once() + with orm_session_maker() as session: + assert session.get(Document, document.id) is None # ============================================================================ @@ -416,10 +505,10 @@ class TestIndexProcessorParameters: dataset_id: str, tenant_id: str, collection_binding_id: str, - mock_db_session, - mock_storage, - mock_index_processor_factory, - mock_get_image_upload_file_ids, + orm_session_maker: sessionmaker[Session], + mock_storage: MagicMock, + mock_index_processor_factory: dict[str, MagicMock], + mock_get_image_upload_file_ids: MagicMock, ): """ Test that correct parameters are passed to IndexProcessor.clean(). @@ -460,7 +549,9 @@ class TestIndexProcessorParameters: assert call_args[0][1] is None # Verify keyword arguments - assert call_args[1]["session"] is mock_db_session.session + cleanup_session = call_args[1]["session"] + assert isinstance(cleanup_session, Session) + assert cleanup_session.get_bind() is orm_session_maker.kw["bind"] assert call_args[1]["with_keywords"] is True assert call_args[1]["delete_child_chunks"] is True schedule_refresh.assert_called_once_with(tenant_id) @@ -470,11 +561,11 @@ class TestIndexProcessorParameters: dataset_id: str, tenant_id: str, collection_binding_id: str, - mock_db_session, - mock_storage, - mock_index_processor_factory, - mock_get_image_upload_file_ids, - ): + orm_session_maker: sessionmaker[Session], + mock_storage: MagicMock, + mock_index_processor_factory: dict[str, MagicMock], + mock_get_image_upload_file_ids: MagicMock, + ) -> None: mock_index_processor_factory["processor"].clean.side_effect = RuntimeError("vector cleanup failed") with patch("tasks.clean_dataset_task.schedule_billing_vector_space_refresh") as schedule_refresh: diff --git a/api/tests/unit_tests/tasks/test_clean_document_task.py b/api/tests/unit_tests/tasks/test_clean_document_task.py index 2f517ce1ba4..4aef0581933 100644 --- a/api/tests/unit_tests/tasks/test_clean_document_task.py +++ b/api/tests/unit_tests/tasks/test_clean_document_task.py @@ -1,62 +1,65 @@ -""" -Unit tests for clean_document_task. +"""SQLite-backed resilience tests for ``clean_document_task``. -Focuses on the resilience contract added by the billing-failure fix: -``index_processor.clean()`` is wrapped in ``try/except`` so that a transient -failure inside the vector / keyword cleanup (e.g. ``ValueError("Unable to -retrieve billing information...")`` raised by ``BillingService._send_request`` -when ``Vector(dataset)`` transitively triggers ``FeatureService.get_features``) -does not abort the entire task and leave PG with stranded ``DocumentSegment`` -/ ``ChildChunk`` / ``UploadFile`` / ``DatasetMetadataBinding`` rows. +The task must continue PostgreSQL cleanup when vector cleanup fails. Each test +starts from the production incident shape: the caller has already deleted the +``Document`` row while its segments and metadata bindings remain. """ import uuid from unittest.mock import MagicMock, patch import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session +import tasks.clean_document_task as clean_document_task_module +from models.dataset import ( + Dataset, + DatasetMetadataBinding, + Document, + DocumentSegment, + SegmentAttachmentBinding, +) +from models.enums import DataSourceType, DocumentCreatedFrom +from models.model import UploadFile from tasks.clean_document_task import clean_document_task +SQLITE_MODELS = ( + Dataset, + Document, + DocumentSegment, + SegmentAttachmentBinding, + UploadFile, + DatasetMetadataBinding, +) + +pytestmark = pytest.mark.parametrize("sqlite_session", [SQLITE_MODELS], indirect=True) + @pytest.fixture -def document_id(): +def document_id() -> str: return str(uuid.uuid4()) @pytest.fixture -def dataset_id(): +def dataset_id() -> str: return str(uuid.uuid4()) @pytest.fixture -def tenant_id(): +def tenant_id() -> str: return str(uuid.uuid4()) @pytest.fixture -def mock_session_factory(): - """Patch ``session_factory.create_session`` to return per-call mock sessions. - - Each call to ``create_session()`` yields a fresh ``MagicMock`` session so we - can assert ``execute()`` calls across the multiple short-lived transactions - used by ``clean_document_task``. - """ - with patch("tasks.clean_document_task.session_factory", autospec=True) as mock_sf: - sessions: list[MagicMock] = [] - - def _create_session(): - session = MagicMock() - session.scalars.return_value.all.return_value = [] - session.execute.return_value.all.return_value = [] - session.scalar.return_value = None - cm = MagicMock() - cm.__enter__.return_value = session - cm.__exit__.return_value = None - sessions.append(session) - return cm - - mock_sf.create_session.side_effect = _create_session - yield mock_sf, sessions +def bind_task_sessions(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session) -> None: + """Bind every short-lived task transaction to the isolated SQLite engine.""" + engine = sqlite_session.get_bind() + monkeypatch.setattr( + clean_document_task_module.session_factory, + "create_session", + lambda: Session(engine, expire_on_commit=False), + ) @pytest.fixture @@ -68,7 +71,7 @@ def mock_storage(): @pytest.fixture def mock_index_processor_factory(): - """Mock ``IndexProcessorFactory`` so we can inject behavior into ``clean``.""" + """Mock the vector/index boundary so cleanup behavior is deterministic.""" with patch("tasks.clean_document_task.IndexProcessorFactory", autospec=True) as factory_cls: processor = MagicMock() processor.clean.return_value = None @@ -83,87 +86,142 @@ def mock_index_processor_factory(): } -def _build_segment(segment_id: str, content: str = "segment content") -> MagicMock: - seg = MagicMock() - seg.id = segment_id - seg.index_node_id = f"node-{segment_id}" - seg.content = content - return seg +def _document(*, document_id: str, dataset_id: str, tenant_id: str, created_by: str) -> Document: + return Document( + id=document_id, + tenant_id=tenant_id, + dataset_id=dataset_id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name=f"{document_id}.txt", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, + ) -def _build_dataset(dataset_id: str, tenant_id: str) -> MagicMock: - ds = MagicMock() - ds.id = dataset_id - ds.tenant_id = tenant_id - return ds +def _segment(*, segment_id: str, document_id: str, dataset_id: str, tenant_id: str, created_by: str) -> DocumentSegment: + segment = DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset_id, + document_id=document_id, + position=1, + content="segment content", + word_count=2, + tokens=2, + created_by=created_by, + index_node_id=f"node-{segment_id}", + ) + segment.id = segment_id + return segment + + +def _persist_deleted_document_state( + session: Session, + *, + document_id: str, + dataset_id: str, + tenant_id: str, + target_segment_ids: list[str], +) -> tuple[str, str]: + """Persist target children after deleting their document, plus scoped control rows.""" + created_by = str(uuid.uuid4()) + other_document_id = str(uuid.uuid4()) + survivor_segment_id = str(uuid.uuid4()) + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Cleanup dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + ) + target_document = _document( + document_id=document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + created_by=created_by, + ) + other_document = _document( + document_id=other_document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + created_by=created_by, + ) + segments = [ + _segment( + segment_id=segment_id, + document_id=document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + created_by=created_by, + ) + for segment_id in target_segment_ids + ] + segments.append( + _segment( + segment_id=survivor_segment_id, + document_id=other_document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + created_by=created_by, + ) + ) + metadata_bindings = [ + DatasetMetadataBinding( + tenant_id=tenant_id, + dataset_id=dataset_id, + metadata_id=str(uuid.uuid4()), + document_id=current_document_id, + created_by=created_by, + ) + for current_document_id in (document_id, other_document_id) + ] + + session.add_all([dataset, target_document, other_document, *segments, *metadata_bindings]) + session.commit() + session.delete(target_document) + session.commit() + return other_document_id, survivor_segment_id + + +def _assert_relational_cleanup( + session: Session, + *, + document_id: str, + other_document_id: str, + survivor_segment_id: str, +) -> None: + session.expire_all() + assert session.get(Document, document_id) is None + remaining_segments = session.scalars(select(DocumentSegment)).all() + assert [(segment.id, segment.document_id) for segment in remaining_segments] == [ + (survivor_segment_id, other_document_id) + ] + remaining_binding_document_ids = set(session.scalars(select(DatasetMetadataBinding.document_id)).all()) + assert remaining_binding_document_ids == {other_document_id} class TestVectorCleanupResilience: - """Vector / keyword cleanup must not abort the task on transient failure.""" + """Vector/index failures must not abort relational cleanup.""" def test_billing_failure_during_vector_cleanup_does_not_skip_pg_cleanup( self, - document_id, - dataset_id, - tenant_id, - mock_session_factory, + document_id: str, + dataset_id: str, + tenant_id: str, + sqlite_session: Session, + bind_task_sessions: None, mock_storage, mock_index_processor_factory, - ): - """Reproduces the production incident: - - ``Vector(dataset)`` transitively calls ``FeatureService.get_features`` - which calls ``BillingService._send_request("GET", ...)``. When billing - returns non-200 it raises ``ValueError("Unable to retrieve billing - information...")``. Before the fix this propagated out of - ``clean_document_task`` and left ``DocumentSegment`` / ``ChildChunk`` / - ``UploadFile`` / ``DatasetMetadataBinding`` rows orphaned because the - already-deleted ``Document`` row had been hard-committed by the caller - (``dataset_service.delete_document``) before ``.delay()`` was invoked. - - Contract: a billing failure inside ``index_processor.clean()`` must be - caught, logged, and the rest of the task must continue so PG ends up - consistent with the deleted ``Document`` even if Qdrant retains - orphan vectors that can be reaped later. - """ - mock_sf, sessions = mock_session_factory - - # First create_session(): Step 1 (load segments + attachments). - step1_session = MagicMock() - step1_session.scalars.return_value.all.return_value = [ - _build_segment("seg-1"), - _build_segment("seg-2"), - ] - step1_session.execute.return_value.all.return_value = [] - step1_session.scalar.return_value = _build_dataset(dataset_id, tenant_id) - # Second create_session(): Step 2 (vector cleanup). Returns dataset. - step2_session = MagicMock() - step2_session.scalar.return_value = _build_dataset(dataset_id, tenant_id) - step2_session.scalars.return_value.all.return_value = [] - step2_session.execute.return_value.all.return_value = [] - # Subsequent sessions: Step 3+ (image / segment / file / metadata cleanup). - # Default fixture returns empty results which is fine for these short txns. - cm1, cm2 = MagicMock(), MagicMock() - cm1.__enter__.return_value = step1_session - cm1.__exit__.return_value = None - cm2.__enter__.return_value = step2_session - cm2.__exit__.return_value = None - - def _default_cm(): - session = MagicMock() - session.scalars.return_value.all.return_value = [] - session.execute.return_value.all.return_value = [] - session.scalar.return_value = None - cm = MagicMock() - cm.__enter__.return_value = session - cm.__exit__.return_value = None - sessions.append(session) - return cm - - mock_sf.create_session.side_effect = [cm1, cm2] + [_default_cm() for _ in range(10)] - - # Simulate the production failure: index_processor.clean() raises ValueError - # mirroring BillingService._send_request when billing returns non-200. + ) -> None: + """A transient billing failure leaves only the unrelated document's rows.""" + other_document_id, survivor_segment_id = _persist_deleted_document_state( + sqlite_session, + document_id=document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + target_segment_ids=["seg-1", "seg-2"], + ) mock_index_processor_factory["processor"].clean.side_effect = ValueError( "Unable to retrieve billing information. Please try again later or contact support." ) @@ -177,59 +235,33 @@ class TestVectorCleanupResilience: file_id=None, ) - # Assert - # 1. Vector cleanup was attempted. mock_index_processor_factory["processor"].clean.assert_called_once() - # 2. Despite the failure the task continued: at least one DocumentSegment - # delete was issued. We use the count of session.execute calls across - # later short transactions as a proxy for "Step 3+ executed". - execute_calls = sum(s.execute.call_count for s in sessions) - assert execute_calls > 0, ( - "Step 3+ DB cleanup did not run after vector cleanup failure; " - "this regression would re-introduce the orphan-segment bug." + _assert_relational_cleanup( + sqlite_session, + document_id=document_id, + other_document_id=other_document_id, + survivor_segment_id=survivor_segment_id, ) schedule_refresh.assert_not_called() def test_vector_cleanup_success_path_remains_unaffected( self, - document_id, - dataset_id, - tenant_id, - mock_session_factory, + document_id: str, + dataset_id: str, + tenant_id: str, + sqlite_session: Session, + bind_task_sessions: None, mock_storage, mock_index_processor_factory, - ): - """Backward-compat: the happy path must still call ``clean()`` exactly - once with the expected arguments and complete without errors. - """ - mock_sf, sessions = mock_session_factory - - step1_session = MagicMock() - step1_session.scalars.return_value.all.return_value = [_build_segment("seg-1")] - step1_session.execute.return_value.all.return_value = [] - step1_session.scalar.return_value = _build_dataset(dataset_id, tenant_id) - step2_session = MagicMock() - step2_session.scalar.return_value = _build_dataset(dataset_id, tenant_id) - step2_session.scalars.return_value.all.return_value = [] - step2_session.execute.return_value.all.return_value = [] - cm1, cm2 = MagicMock(), MagicMock() - cm1.__enter__.return_value = step1_session - cm1.__exit__.return_value = None - cm2.__enter__.return_value = step2_session - cm2.__exit__.return_value = None - - def _default_cm(): - session = MagicMock() - session.scalars.return_value.all.return_value = [] - session.execute.return_value.all.return_value = [] - session.scalar.return_value = None - cm = MagicMock() - cm.__enter__.return_value = session - cm.__exit__.return_value = None - sessions.append(session) - return cm - - mock_sf.create_session.side_effect = [cm1, cm2] + [_default_cm() for _ in range(10)] + ) -> None: + """The happy path calls the index boundary and completes scoped cleanup.""" + other_document_id, survivor_segment_id = _persist_deleted_document_state( + sqlite_session, + document_id=document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + target_segment_ids=["seg-1"], + ) with patch("tasks.clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh: clean_document_task( @@ -239,49 +271,42 @@ class TestVectorCleanupResilience: file_id=None, ) - assert mock_index_processor_factory["processor"].clean.call_count == 1 - # Index cleanup invoked with the expected delete_summaries / delete_child_chunks flags. + mock_index_processor_factory["processor"].clean.assert_called_once() _, kwargs = mock_index_processor_factory["processor"].clean.call_args - assert kwargs.get("with_keywords") is True - assert kwargs.get("delete_child_chunks") is True - assert kwargs.get("delete_summaries") is True + cleanup_session = kwargs.pop("session") + assert isinstance(cleanup_session, Session) + assert cleanup_session.get_bind() is sqlite_session.get_bind() + assert kwargs == { + "with_keywords": True, + "delete_child_chunks": True, + "delete_summaries": True, + } + _assert_relational_cleanup( + sqlite_session, + document_id=document_id, + other_document_id=other_document_id, + survivor_segment_id=survivor_segment_id, + ) schedule_refresh.assert_called_once_with(tenant_id) def test_no_segments_skips_vector_cleanup( self, - document_id, - dataset_id, - tenant_id, - mock_session_factory, + document_id: str, + dataset_id: str, + tenant_id: str, + sqlite_session: Session, + bind_task_sessions: None, mock_storage, mock_index_processor_factory, - ): - """When the document has no segments (e.g. indexing failed before - producing any), vector cleanup must not be attempted — and therefore - the new try/except wrapper does not change behavior here. - """ - mock_sf, sessions = mock_session_factory - - step1_session = MagicMock() - step1_session.scalars.return_value.all.return_value = [] # no segments - step1_session.execute.return_value.all.return_value = [] - step1_session.scalar.return_value = _build_dataset(dataset_id, tenant_id) - cm1 = MagicMock() - cm1.__enter__.return_value = step1_session - cm1.__exit__.return_value = None - - def _default_cm(): - session = MagicMock() - session.scalars.return_value.all.return_value = [] - session.execute.return_value.all.return_value = [] - session.scalar.return_value = None - cm = MagicMock() - cm.__enter__.return_value = session - cm.__exit__.return_value = None - sessions.append(session) - return cm - - mock_sf.create_session.side_effect = [cm1] + [_default_cm() for _ in range(10)] + ) -> None: + """A target document without segments skips the vector/index boundary.""" + other_document_id, survivor_segment_id = _persist_deleted_document_state( + sqlite_session, + document_id=document_id, + dataset_id=dataset_id, + tenant_id=tenant_id, + target_segment_ids=[], + ) with patch("tasks.clean_document_task.schedule_billing_vector_space_refresh") as schedule_refresh: clean_document_task( @@ -291,7 +316,11 @@ class TestVectorCleanupResilience: file_id=None, ) - # Vector cleanup is gated on ``index_node_ids``; when there are no - # segments the IndexProcessorFactory path is never entered. mock_index_processor_factory["factory_cls"].assert_not_called() + _assert_relational_cleanup( + sqlite_session, + document_id=document_id, + other_document_id=other_document_id, + survivor_segment_id=survivor_segment_id, + ) schedule_refresh.assert_not_called() diff --git a/api/tests/unit_tests/tasks/test_collect_agent_resources_task.py b/api/tests/unit_tests/tasks/test_collect_agent_resources_task.py new file mode 100644 index 00000000000..48fe8ff2def --- /dev/null +++ b/api/tests/unit_tests/tasks/test_collect_agent_resources_task.py @@ -0,0 +1,81 @@ +from typing import Protocol, cast +from unittest.mock import MagicMock + +import pytest + +from services.agent.home_snapshot_service import AgentHomeSnapshotService +from services.agent.workspace_service import AgentWorkspaceService +from tasks.collect_agent_resources_task import ( + collect_agent_resources, + enqueue_agent_resource_collection, +) + + +class _TaskWithQueue(Protocol): + queue: str + + +def test_collection_task_uses_retention_queue() -> None: + task = cast(_TaskWithQueue, collect_agent_resources) + assert task.queue == "retention" + + +def test_enqueue_deduplicates_ids_and_skips_empty_input(monkeypatch: pytest.MonkeyPatch) -> None: + delay = MagicMock() + monkeypatch.setattr(collect_agent_resources, "delay", delay) + + enqueue_agent_resource_collection(tenant_id="tenant-1") + enqueue_agent_resource_collection( + tenant_id="tenant-1", + binding_ids=["binding-2", "binding-1", "binding-2"], + workspace_ids=["workspace-1"], + ) + + delay.assert_called_once_with( + tenant_id="tenant-1", + binding_ids=["binding-1", "binding-2"], + workspace_ids=["workspace-1"], + home_snapshot_ids=[], + ) + + +def test_collection_continues_in_workspace_binding_snapshot_order(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[str] = [] + + def collect_workspace(**_kwargs: object) -> None: + calls.append("workspace") + raise RuntimeError("workspace failed") + + monkeypatch.setattr(AgentWorkspaceService, "collect_retired_workspace", collect_workspace) + monkeypatch.setattr( + AgentWorkspaceService, + "collect_retired_binding", + lambda **_kwargs: calls.append("binding"), + ) + monkeypatch.setattr( + AgentHomeSnapshotService, + "collect_retired_home_snapshot", + lambda **_kwargs: calls.append("home"), + ) + + collect_agent_resources.run( + tenant_id="tenant-1", + workspace_ids=["workspace-1"], + binding_ids=["binding-1"], + home_snapshot_ids=["home-1"], + ) + + assert calls == ["workspace", "binding", "home"] + + +def test_enqueue_failure_is_best_effort(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + collect_agent_resources, + "delay", + MagicMock(side_effect=RuntimeError("queue unavailable")), + ) + + enqueue_agent_resource_collection( + tenant_id="tenant-1", + binding_ids=["binding-1"], + ) diff --git a/api/tests/unit_tests/tasks/test_community_telemetry_task.py b/api/tests/unit_tests/tasks/test_community_telemetry_task.py new file mode 100644 index 00000000000..00f97a5f31b --- /dev/null +++ b/api/tests/unit_tests/tasks/test_community_telemetry_task.py @@ -0,0 +1,50 @@ +import logging +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session + +from tasks import community_telemetry_task + + +def _bind_task_to_sqlite(monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine) -> None: + """Bind the task-local sessionmaker to the isolated SQLite database.""" + monkeypatch.setattr(community_telemetry_task, "db", SimpleNamespace(engine=sqlite_engine)) + + +def test_send_community_telemetry_heartbeat_reports_with_a_database_session( + monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine +) -> None: + _bind_task_to_sqlite(monkeypatch, sqlite_engine) + received_sessions: list[Session] = [] + + def report_heartbeat(*, session: Session) -> None: + received_sessions.append(session) + assert session.get_bind() is sqlite_engine + + monkeypatch.setattr(community_telemetry_task.CommunityTelemetryService, "report_heartbeat", report_heartbeat) + + community_telemetry_task.send_community_telemetry_heartbeat.run() + + assert len(received_sessions) == 1 + assert isinstance(received_sessions[0], Session) + + +def test_send_community_telemetry_heartbeat_swallows_report_errors( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + caplog: pytest.LogCaptureFixture, +) -> None: + _bind_task_to_sqlite(monkeypatch, sqlite_engine) + monkeypatch.setattr( + community_telemetry_task.CommunityTelemetryService, + "report_heartbeat", + Mock(side_effect=RuntimeError("telemetry unavailable")), + ) + caplog.set_level(logging.DEBUG, logger=community_telemetry_task.logger.name) + + community_telemetry_task.send_community_telemetry_heartbeat.run() + + assert "Failed to process community telemetry heartbeat" in caplog.text diff --git a/api/tests/unit_tests/tasks/test_dataset_indexing_task.py b/api/tests/unit_tests/tasks/test_dataset_indexing_task.py index 39ff472ab7e..25dd16b2d58 100644 --- a/api/tests/unit_tests/tasks/test_dataset_indexing_task.py +++ b/api/tests/unit_tests/tasks/test_dataset_indexing_task.py @@ -1,29 +1,24 @@ -""" -Unit tests for dataset indexing tasks. +"""SQLite-backed tests for document indexing tasks. -This module tests the document indexing task functionality including: -- Task enqueuing to different queues (normal, priority, tenant-isolated) -- Batch processing of multiple documents -- Progress tracking through task lifecycle -- Error handling and retry mechanisms -- Task cancellation and cleanup +The indexing task deliberately uses separate transactions for validation, +status persistence, indexing, and summary dispatch. These tests persist real +ORM rows so each phase observes only committed database state. """ -import logging import uuid from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import MagicMock, Mock, patch import pytest +from sqlalchemy.orm import Session from core.indexing_runner import DocumentIsPausedError from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType -from core.rag.pipeline.queue import TenantIsolatedTaskQueue -from enums.cloud_plan import CloudPlan +from enums import CloudPlan from extensions.ext_redis import redis_client from models.dataset import Dataset, Document -from models.enums import IndexingStatus +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus from services.document_indexing_proxy.document_indexing_task_proxy import DocumentIndexingTaskProxy from tasks.document_indexing_task import ( _document_indexing, @@ -33,40 +28,26 @@ from tasks.document_indexing_task import ( priority_document_indexing_task, ) -# ============================================================================ -# Fixtures -# ============================================================================ - @pytest.fixture -def tenant_id(): - """Generate a unique tenant ID for testing.""" +def tenant_id() -> str: return str(uuid.uuid4()) @pytest.fixture -def dataset_id(): - """Generate a unique dataset ID for testing.""" +def dataset_id() -> str: return str(uuid.uuid4()) @pytest.fixture -def document_ids(): - """Generate a list of document IDs for testing.""" +def document_ids() -> list[str]: return [str(uuid.uuid4()) for _ in range(3)] @pytest.fixture -def mock_redis(): - """Mock Redis client operations.""" - # Redis is already mocked globally in conftest.py - # Reset it for each test +def mock_redis() -> MagicMock: + """Reset the external Redis boundary used by tenant-isolated queues.""" redis_client.reset_mock() - redis_client.get.reset_mock() - redis_client.setex.reset_mock() - redis_client.delete.reset_mock() - redis_client.lpush.reset_mock() - redis_client.rpop.reset_mock() redis_client.get.return_value = None redis_client.setex.return_value = True redis_client.delete.return_value = True @@ -75,1903 +56,495 @@ def mock_redis(): return redis_client -# Additional fixtures required by tests in this module - - @pytest.fixture -def mock_db_session(): - """Mock session_factory.create_session() to return a session whose queries use shared test data. - - Tests set session._shared_data = {"dataset": , "documents": [, ...]} - This fixture makes session.scalar(select(Dataset)...) return the shared dataset, - and session.scalars(select(Document)...).all() return the shared documents. - """ - with patch("tasks.document_indexing_task.session_factory") as mock_sf: - session = MagicMock() - session._shared_data = {"dataset": None, "documents": []} - - def _get_entity(stmt) -> type | None: - """Extract the mapped entity class from a SQLAlchemy select statement.""" - try: - descs = stmt.column_descriptions - if descs: - return descs[0].get("entity") - except (AttributeError, TypeError): - pass - return None - - def _extract_id_from_where(stmt) -> str | None: - """Return the value bound to the 'id' column in the WHERE clause, if present.""" - try: - where = stmt.whereclause - if where is None: - clauses = [] - else: - try: - clauses = list(where.clauses) - except AttributeError: - clauses = [where] - except Exception: - return None - - for clause in clauses: - try: - left = clause.left - right = clause.right - except AttributeError: - continue - try: - key = left.key - except AttributeError: - continue - if key == "id": - try: - return right.value - except AttributeError: - return None - return None - - def _scalar_side_effect(stmt): - entity = _get_entity(stmt) - if entity is not None: - if entity.__name__ == "Dataset": - return session._shared_data.get("dataset") - elif entity.__name__ == "Document": - docs = session._shared_data.get("documents", []) - if not docs: - return None - queried_id = _extract_id_from_where(stmt) - if queried_id: - doc_map = {d.id: d for d in docs} - return doc_map.get(queried_id, docs[0]) - return docs[0] - return None - - def _scalars_side_effect(stmt): - entity = _get_entity(stmt) - result = MagicMock() - if entity is not None: - if entity.__name__ == "Document": - result.all.return_value = list(session._shared_data.get("documents", [])) - elif entity.__name__ == "Dataset": - ds = session._shared_data.get("dataset") - result.all.return_value = [ds] if ds else [] - else: - result.all.return_value = [] - else: - result.all.return_value = [] - return result - - session.scalar.side_effect = _scalar_side_effect - session.scalars.side_effect = _scalars_side_effect - - # Implement session.begin() context manager that commits on exit - session.commit = MagicMock() - bm = MagicMock() - bm.__enter__.return_value = session - - def _bm_exit_side_effect(*args, **kwargs): - session.commit() - - bm.__exit__.side_effect = _bm_exit_side_effect - session.begin.return_value = bm - - # Context manager behavior for create_session(): ensure close() is called on exit - session.close = MagicMock() - cm = MagicMock() - cm.__enter__.return_value = session - - def _exit_side_effect(*args, **kwargs): - session.close() - - cm.__exit__.side_effect = _exit_side_effect - mock_sf.create_session.return_value = cm - - yield session +def indexing_runner(monkeypatch: pytest.MonkeyPatch) -> MagicMock: + runner = MagicMock() + runner_class = MagicMock(return_value=runner) + monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", runner_class) + runner._constructor_mock = runner_class + return runner -@pytest.fixture -def mock_dataset(dataset_id, tenant_id): - """Create a mock Dataset object.""" - dataset = Mock(spec=Dataset) - dataset.id = dataset_id - dataset.tenant_id = tenant_id - dataset.indexing_technique = IndexTechniqueType.HIGH_QUALITY - dataset.embedding_model_provider = "openai" - dataset.embedding_model = "text-embedding-ada-002" - return dataset +def _features( + *, + billing_enabled: bool = False, + plan: CloudPlan = CloudPlan.PROFESSIONAL, + vector_limit: int = 1000, + vector_size: int = 0, +) -> SimpleNamespace: + return SimpleNamespace( + billing=SimpleNamespace(enabled=billing_enabled, subscription=SimpleNamespace(plan=plan)), + vector_space=SimpleNamespace(limit=vector_limit, size=vector_size), + ) -@pytest.fixture -def mock_documents(document_ids, dataset_id): - """Create mock Document objects.""" - documents = [] - for doc_id in document_ids: - doc = Mock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.error = None - doc.stopped_at = None - doc.processing_started_at = None - # optional attribute used in some code paths - doc.doc_form = IndexStructureType.PARAGRAPH_INDEX - documents.append(doc) - return documents +def _patch_features(monkeypatch: pytest.MonkeyPatch, features: SimpleNamespace) -> MagicMock: + get_features = MagicMock(return_value=features) + monkeypatch.setattr("tasks.document_indexing_task.FeatureService.get_features", get_features) + return get_features -@pytest.fixture -def mock_indexing_runner(): - """Mock IndexingRunner for document_indexing_task module.""" - with patch("tasks.document_indexing_task.IndexingRunner") as mock_runner_class: - mock_runner = MagicMock() - mock_runner_class.return_value = mock_runner - yield mock_runner +def _persist_indexing_rows( + session: Session, + *, + tenant_id: str, + dataset_id: str, + document_ids: list[str], + indexing_technique: IndexTechniqueType = IndexTechniqueType.HIGH_QUALITY, + summary_index_setting: dict[str, bool] | None = None, + document_forms: list[IndexStructureType] | None = None, + need_summary: list[bool] | None = None, +) -> tuple[Dataset, list[Document]]: + """Persist one tenant-owned dataset and the requested document rows.""" + created_by = str(uuid.uuid4()) + dataset = Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Indexing dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + indexing_technique=indexing_technique, + embedding_model_provider="openai", + embedding_model="text-embedding-3-small", + summary_index_setting=summary_index_setting, + created_by=created_by, + ) + documents = [ + Document( + id=document_id, + tenant_id=tenant_id, + dataset_id=dataset_id, + position=position, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name=f"document-{position}.txt", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, + indexing_status=IndexingStatus.WAITING, + doc_form=(document_forms or [IndexStructureType.PARAGRAPH_INDEX] * len(document_ids))[position - 1], + need_summary=(need_summary or [False] * len(document_ids))[position - 1], + ) + for position, document_id in enumerate(document_ids, start=1) + ] + session.add_all([dataset, *documents]) + session.commit() + return dataset, documents -@pytest.fixture -def mock_feature_service(): - """Mock FeatureService for document_indexing_task module.""" - with patch("tasks.document_indexing_task.FeatureService") as mock_service: - mock_features = Mock() - mock_features.billing = Mock() - mock_features.billing.enabled = False - mock_features.vector_space = Mock() - mock_features.vector_space.size = 0 - mock_features.vector_space.limit = 1000 - mock_service.get_features.return_value = mock_features - yield mock_service - - -# ============================================================================ -# Test Task Enqueuing -# ============================================================================ +def _persisted_documents(session: Session, document_ids: list[str]) -> list[Document]: + session.expire_all() + return [document for document_id in document_ids if (document := session.get(Document, document_id)) is not None] class TestTaskEnqueuing: - """Test cases for task enqueuing to different queues.""" - - def test_enqueue_to_priority_direct_queue_for_self_hosted(self, tenant_id, dataset_id, document_ids, mock_redis): - """ - Test enqueuing to priority direct queue for self-hosted deployments. - - When billing is disabled (self-hosted), tasks should go directly to - the priority queue without tenant isolation. - """ - # Arrange - with patch.object(DocumentIndexingTaskProxy, "features") as mock_features: - mock_features.billing.enabled = False - - # Mock the class variable directly - mock_task = Mock() - with patch.object(DocumentIndexingTaskProxy, "PRIORITY_TASK_FUNC", mock_task): - proxy = DocumentIndexingTaskProxy(tenant_id, dataset_id, document_ids) - - # Act - proxy.delay() - - # Assert - mock_task.delay.assert_called_once_with( - tenant_id=tenant_id, dataset_id=dataset_id, document_ids=document_ids - ) - - def test_enqueue_to_normal_tenant_queue_for_sandbox_plan(self, tenant_id, dataset_id, document_ids, mock_redis): - """ - Test enqueuing to normal tenant queue for sandbox plan. - - Sandbox plan users should have their tasks queued with tenant isolation - in the normal priority queue. - """ - # Arrange - mock_redis.get.return_value = None # No existing task - - with patch.object(DocumentIndexingTaskProxy, "features") as mock_features: - mock_features.billing.enabled = True - mock_features.billing.subscription.plan = CloudPlan.SANDBOX - - # Mock the class variable directly - mock_task = Mock() - with patch.object(DocumentIndexingTaskProxy, "NORMAL_TASK_FUNC", mock_task): - proxy = DocumentIndexingTaskProxy(tenant_id, dataset_id, document_ids) - - # Act - proxy.delay() - - # Assert - Should set task key and call delay - assert mock_redis.setex.called - mock_task.delay.assert_called_once() - - def test_enqueue_to_priority_tenant_queue_for_paid_plan(self, tenant_id, dataset_id, document_ids, mock_redis): - """ - Test enqueuing to priority tenant queue for paid plans. - - Paid plan users should have their tasks queued with tenant isolation - in the priority queue. - """ - # Arrange - mock_redis.get.return_value = None # No existing task - - with patch.object(DocumentIndexingTaskProxy, "features") as mock_features: - mock_features.billing.enabled = True - mock_features.billing.subscription.plan = CloudPlan.PROFESSIONAL - - # Mock the class variable directly - mock_task = Mock() - with patch.object(DocumentIndexingTaskProxy, "PRIORITY_TASK_FUNC", mock_task): - proxy = DocumentIndexingTaskProxy(tenant_id, dataset_id, document_ids) - - # Act - proxy.delay() - - # Assert - assert mock_redis.setex.called - mock_task.delay.assert_called_once() - - def test_enqueue_adds_to_waiting_queue_when_task_running(self, tenant_id, dataset_id, document_ids, mock_redis): - """ - Test that new tasks are added to waiting queue when a task is already running. - - If a task is already running for the tenant (task key exists), - new tasks should be pushed to the waiting queue. - """ - # Arrange - mock_redis.get.return_value = b"1" # Task already running - - with patch.object(DocumentIndexingTaskProxy, "features") as mock_features: - mock_features.billing.enabled = True - mock_features.billing.subscription.plan = CloudPlan.PROFESSIONAL - - # Mock the class variable directly - mock_task = Mock() - with patch.object(DocumentIndexingTaskProxy, "PRIORITY_TASK_FUNC", mock_task): - proxy = DocumentIndexingTaskProxy(tenant_id, dataset_id, document_ids) - - # Act - proxy.delay() - - # Assert - Should push to queue, not call delay - assert mock_redis.lpush.called - mock_task.delay.assert_not_called() - - def test_legacy_document_indexing_task_still_works( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_documents, mock_indexing_runner - ): - """ - Test that the legacy document_indexing_task function still works. - - This ensures backward compatibility for existing code that may still - use the deprecated function. - """ - # Arrange - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - document_indexing_task(dataset_id, document_ids) - - # Assert - mock_indexing_runner.run.assert_called_once() - - -# ============================================================================ -# Test Batch Processing -# ============================================================================ - - -class TestBatchProcessing: - """Test cases for batch processing of multiple documents.""" - - def test_batch_processing_multiple_documents( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test batch processing of multiple documents. - - All documents in the batch should be processed together and their - status should be updated to 'parsing'. - """ - # Arrange - Create actual document objects that can be modified - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.error = None - doc.stopped_at = None - doc.processing_started_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - All documents should be set to 'parsing' status - for doc in mock_documents: - assert doc.indexing_status == IndexingStatus.PARSING - assert doc.processing_started_at is not None - - # IndexingRunner should be called with all documents - mock_indexing_runner.run.assert_called_once() - call_args = mock_indexing_runner.run.call_args[0][0] - assert len(call_args) == len(document_ids) - - def test_batch_processing_with_limit_check(self, dataset_id, mock_db_session, mock_dataset, mock_feature_service): - """ - Test batch processing respects upload limits. - - When the number of documents exceeds the batch upload limit, - an error should be raised and all documents should be marked as error. - """ - # Arrange - batch_limit = 10 - document_ids = [str(uuid.uuid4()) for _ in range(batch_limit + 1)] - - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.error = None - doc.stopped_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - mock_feature_service.get_features.return_value.billing.enabled = True - mock_feature_service.get_features.return_value.billing.subscription.plan = CloudPlan.PROFESSIONAL - mock_feature_service.get_features.return_value.vector_space.limit = 1000 - mock_feature_service.get_features.return_value.vector_space.size = 0 - - with patch("tasks.document_indexing_task.dify_config.BATCH_UPLOAD_LIMIT", str(batch_limit)): - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - All documents should have error status - for doc in mock_documents: - assert doc.indexing_status == "error" - assert doc.error is not None - assert "batch upload limit" in doc.error - - def test_batch_processing_sandbox_plan_single_document_only( - self, dataset_id, mock_db_session, mock_dataset, mock_feature_service - ): - """ - Test that sandbox plan only allows single document upload. - - Sandbox plan should reject batch uploads (more than 1 document). - """ - # Arrange - document_ids = [str(uuid.uuid4()) for _ in range(2)] - - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.error = None - doc.stopped_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - mock_feature_service.get_features.return_value.billing.enabled = True - mock_feature_service.get_features.return_value.billing.subscription.plan = CloudPlan.SANDBOX - mock_feature_service.get_features.return_value.vector_space.limit = 1000 - mock_feature_service.get_features.return_value.vector_space.size = 0 - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - All documents should have error status - for doc in mock_documents: - assert doc.indexing_status == "error" - assert "does not support batch upload" in doc.error - - def test_batch_processing_empty_document_list( - self, dataset_id, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test batch processing with empty document list. - - Should handle empty list gracefully without errors. - """ - # Arrange - document_ids = [] - - # Set shared mock data with empty documents list - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = [] - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - IndexingRunner should still be called with empty list - mock_indexing_runner.run.assert_called_once_with([], mock_db_session) - - -# ============================================================================ -# Test Progress Tracking -# ============================================================================ - - -class TestProgressTracking: - """Test cases for progress tracking through task lifecycle.""" - - def test_document_status_progression( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test document status progresses correctly through lifecycle. - - Documents should transition from 'waiting' -> 'parsing' -> processed. - """ - # Arrange - Create actual document objects - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.processing_started_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - Status should be 'parsing' - for doc in mock_documents: - assert doc.indexing_status == IndexingStatus.PARSING - assert doc.processing_started_at is not None - - # Verify commit was called to persist status - assert mock_db_session.commit.called - - def test_processing_started_timestamp_set( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test that processing_started_at timestamp is set correctly. - - When documents start processing, the timestamp should be recorded. - """ - # Arrange - Create actual document objects - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.processing_started_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - for doc in mock_documents: - assert doc.processing_started_at is not None - - def test_tenant_queue_processes_next_task_after_completion( - self, tenant_id, dataset_id, document_ids, mock_redis, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test that tenant queue processes next waiting task after completion. - - After a task completes, the system should check for waiting tasks - and process the next one. - """ - # Arrange - next_task_data = {"tenant_id": tenant_id, "dataset_id": dataset_id, "document_ids": ["next_doc_id"]} - - # Simulate next task in queue - from core.rag.pipeline.queue import TaskWrapper - - wrapper = TaskWrapper(data=next_task_data) - mock_redis.rpop.return_value = wrapper.serialize() - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, mock_task) - - # Assert - Next task should be enqueued - mock_task.apply_async.assert_called() - # Task key should be set for next task - assert mock_redis.setex.called - - def test_tenant_queue_clears_flag_when_no_more_tasks( - self, tenant_id, dataset_id, document_ids, mock_redis, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test that tenant queue clears flag when no more tasks are waiting. - - When there are no more tasks in the queue, the task key should be deleted. - """ - # Arrange - mock_redis.rpop.return_value = None # No more tasks - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, mock_task) - - # Assert - Task key should be deleted - assert mock_redis.delete.called - - -# ============================================================================ -# Test Error Handling and Retries -# ============================================================================ - - -class TestErrorHandling: - """Test cases for error handling and retry mechanisms.""" - - def test_error_handling_sets_document_error_status( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_feature_service - ): - """ - Test that errors during validation set document error status. - - When validation fails (e.g., limit exceeded), documents should be - marked with error status and error message. - """ - # Arrange - Create actual document objects - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.error = None - doc.stopped_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - # Set up to trigger vector space limit error - mock_feature_service.get_features.return_value.billing.enabled = True - mock_feature_service.get_features.return_value.billing.subscription.plan = CloudPlan.PROFESSIONAL - mock_feature_service.get_features.return_value.vector_space.limit = 100 - mock_feature_service.get_features.return_value.vector_space.size = 100 # At limit - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - for doc in mock_documents: - assert doc.indexing_status == "error" - assert doc.error is not None - assert "over the limit" in doc.error - assert doc.stopped_at is not None - - def test_error_handling_during_indexing_runner( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_documents, mock_indexing_runner - ): - """ - Test error handling when IndexingRunner raises an exception. - - Errors during indexing should be caught and logged, but not crash the task. - """ - # Arrange - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - # Make IndexingRunner raise an exception - mock_indexing_runner.run.side_effect = Exception("Indexing failed") - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - Should not raise exception - _document_indexing(dataset_id, document_ids) - - # Assert - Session should be closed even after error - assert mock_db_session.close.called - - def test_document_paused_error_handling( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_documents, mock_indexing_runner - ): - """ - Test handling of DocumentIsPausedError. - - When a document is paused, the error should be caught and logged - but not treated as a failure. - """ - # Arrange - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - # Make IndexingRunner raise DocumentIsPausedError - mock_indexing_runner.run.side_effect = DocumentIsPausedError("Document is paused") - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - Should not raise exception - _document_indexing(dataset_id, document_ids) - - # Assert - Session should be closed - assert mock_db_session.close.called - - def test_dataset_not_found_error_handling(self, dataset_id, document_ids, mock_db_session): - """ - Test handling when dataset is not found. - - If the dataset doesn't exist, the task should exit gracefully. - """ - # Arrange - dataset is not in _shared_data (None by default), so scalar() returns None - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - Session should be closed - assert mock_db_session.close.called - - def test_tenant_queue_error_handling_still_processes_next_task( - self, - tenant_id, - dataset_id, - document_ids, - mock_redis, - mock_db_session, - mock_dataset, - mock_indexing_runner, - caplog: pytest.LogCaptureFixture, - ): - """ - Test that errors don't prevent processing next task in tenant queue. - - Even if the current task fails, the next task should still be processed. - """ - # Arrange - next_task_data = {"tenant_id": tenant_id, "dataset_id": dataset_id, "document_ids": ["next_doc_id"]} - - from core.rag.pipeline.queue import TaskWrapper - - wrapper = TaskWrapper(data=next_task_data) - # Set up rpop to return task once for concurrency check - mock_redis.rpop.side_effect = [wrapper.serialize(), None] - - # Make _document_indexing raise an error - with patch("tasks.document_indexing_task._document_indexing") as mock_indexing: - mock_indexing.side_effect = Exception("Processing failed") - - with caplog.at_level(logging.ERROR, logger="tasks.document_indexing_task"): - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, mock_task) - - # Assert - Next task should still be enqueued despite error - mock_task.apply_async.assert_called() - assert ( - f"Error processing document indexing {dataset_id} for tenant {tenant_id}: {document_ids}" - in caplog.messages - ) - - def test_concurrent_task_limit_respected( - self, tenant_id, dataset_id, document_ids, mock_redis, mock_db_session, mock_dataset - ): - """ - Test that tenant isolated task concurrency limit is respected. - - Should pull only TENANT_ISOLATED_TASK_CONCURRENCY tasks at a time. - """ - # Arrange - concurrency_limit = 2 - - # Create multiple tasks in queue - tasks = [] - for i in range(5): - task_data = {"tenant_id": tenant_id, "dataset_id": dataset_id, "document_ids": [f"doc_{i}"]} - from core.rag.pipeline.queue import TaskWrapper - - wrapper = TaskWrapper(data=task_data) - tasks.append(wrapper.serialize()) - - # Mock rpop to return tasks one by one - mock_redis.rpop.side_effect = tasks[:concurrency_limit] + [None] - - mock_db_session._shared_data["dataset"] = mock_dataset - - with patch("tasks.document_indexing_task.dify_config.TENANT_ISOLATED_TASK_CONCURRENCY", concurrency_limit): - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, mock_task) - - # Assert - Should enqueue exactly concurrency_limit tasks - assert mock_task.apply_async.call_count == concurrency_limit - - -# ============================================================================ -# Test Task Cancellation -# ============================================================================ - - -class TestTaskCancellation: - """Test cases for task cancellation and cleanup.""" - - def test_task_isolation_between_tenants(self, mock_redis): - """ - Test that tasks are properly isolated between different tenants. - - Each tenant should have their own queue and task key. - """ - # Arrange - tenant_1 = str(uuid.uuid4()) - tenant_2 = str(uuid.uuid4()) - dataset_id = str(uuid.uuid4()) - document_ids = [str(uuid.uuid4())] - - # Act - queue_1 = TenantIsolatedTaskQueue(tenant_1, "document_indexing") - queue_2 = TenantIsolatedTaskQueue(tenant_2, "document_indexing") - - # Assert - Different tenants should have different queue keys - assert queue_1._queue != queue_2._queue - assert queue_1._task_key != queue_2._task_key - assert tenant_1 in queue_1._queue - assert tenant_2 in queue_2._queue - - -# ============================================================================ -# Integration Tests -# ============================================================================ - - -class TestAdvancedScenarios: - """Advanced test scenarios for edge cases and complex workflows.""" - - def test_multiple_documents_with_mixed_success_and_failure( - self, dataset_id, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test handling of mixed success and failure scenarios in batch processing. - - When processing multiple documents, some may succeed while others fail. - This tests that the system handles partial failures gracefully. - - Scenario: - - Process 3 documents in a batch - - First document succeeds - - Second document is not found (skipped) - - Third document succeeds - - Expected behavior: - - Only found documents are processed - - Missing documents are skipped without crashing - - IndexingRunner receives only valid documents - """ - # Arrange - Create document IDs with one missing - document_ids = [str(uuid.uuid4()) for _ in range(3)] - - # Create only 2 documents (simulate one missing) - # The new code uses .all() which will only return existing documents - mock_documents = [] - for i, doc_id in enumerate([document_ids[0], document_ids[2]]): # Skip middle one - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.processing_started_at = None - mock_documents.append(doc) - - # Set shared mock data - .all() will only return existing documents - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - Only 2 documents should be processed (missing one skipped) - mock_indexing_runner.run.assert_called_once() - call_args = mock_indexing_runner.run.call_args[0][0] - assert len(call_args) == 2 # Only found documents - - def test_tenant_queue_with_multiple_concurrent_tasks( - self, tenant_id, dataset_id, mock_redis, mock_db_session, mock_dataset - ): - """ - Test concurrent task processing with tenant isolation. - - This tests the scenario where multiple tasks are queued for the same tenant - and need to be processed respecting the concurrency limit. - - Scenario: - - 5 tasks are waiting in the queue - - Concurrency limit is 2 - - After current task completes, pull and enqueue next 2 tasks - - Expected behavior: - - Exactly 2 tasks are pulled from queue (respecting concurrency) - - Each task is enqueued with correct parameters - - Task waiting time is set for each new task - """ - # Arrange - concurrency_limit = 2 - document_ids = [str(uuid.uuid4())] - - # Create multiple waiting tasks - waiting_tasks = [] - for i in range(5): - task_data = {"tenant_id": tenant_id, "dataset_id": dataset_id, "document_ids": [f"doc_{i}"]} - from core.rag.pipeline.queue import TaskWrapper - - wrapper = TaskWrapper(data=task_data) - waiting_tasks.append(wrapper.serialize()) - - # Mock rpop to return tasks up to concurrency limit - mock_redis.rpop.side_effect = waiting_tasks[:concurrency_limit] + [None] - mock_db_session._shared_data["dataset"] = mock_dataset - - with patch("tasks.document_indexing_task.dify_config.TENANT_ISOLATED_TASK_CONCURRENCY", concurrency_limit): - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, mock_task) - - # Assert - # Should enqueue exactly concurrency_limit tasks - assert mock_task.apply_async.call_count == concurrency_limit - - # Verify task waiting time was set for each task - assert mock_redis.setex.call_count >= concurrency_limit - - def test_vector_space_limit_edge_case_at_exact_limit( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_feature_service - ): - """ - Test vector space limit validation at exact boundary. - - Edge case: When vector space is exactly at the limit (not over), - the upload should still be rejected. - - Scenario: - - Vector space limit: 100 - - Current size: 100 (exactly at limit) - - Try to upload 3 documents - - Expected behavior: - - Upload is rejected with appropriate error message - - All documents are marked with error status - """ - # Arrange - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.error = None - doc.stopped_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - # Set vector space exactly at limit - mock_feature_service.get_features.return_value.billing.enabled = True - mock_feature_service.get_features.return_value.billing.subscription.plan = CloudPlan.PROFESSIONAL - mock_feature_service.get_features.return_value.vector_space.limit = 100 - mock_feature_service.get_features.return_value.vector_space.size = 100 # Exactly at limit - - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - All documents should have error status - for doc in mock_documents: - assert doc.indexing_status == "error" - assert "over the limit" in doc.error - - def test_task_queue_fifo_ordering(self, tenant_id, dataset_id, mock_redis, mock_db_session, mock_dataset): - """ - Test that tasks are processed in FIFO (First-In-First-Out) order. - - The tenant isolated queue should maintain task order, ensuring - that tasks are processed in the sequence they were added. - - Scenario: - - Task A added first - - Task B added second - - Task C added third - - When pulling tasks, should get A, then B, then C - - Expected behavior: - - Tasks are retrieved in the order they were added - - FIFO ordering is maintained throughout processing - """ - # Arrange - document_ids = [str(uuid.uuid4())] - - # Create tasks with identifiable document IDs to track order - task_order = ["task_A", "task_B", "task_C"] - tasks = [] - for task_name in task_order: - task_data = {"tenant_id": tenant_id, "dataset_id": dataset_id, "document_ids": [task_name]} - from core.rag.pipeline.queue import TaskWrapper - - wrapper = TaskWrapper(data=task_data) - tasks.append(wrapper.serialize()) - - # Mock rpop to return tasks in FIFO order - mock_redis.rpop.side_effect = tasks + [None] - mock_db_session._shared_data["dataset"] = mock_dataset - - with patch("tasks.document_indexing_task.dify_config.TENANT_ISOLATED_TASK_CONCURRENCY", 3): - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, mock_task) - - # Assert - Verify tasks were enqueued in correct order - assert mock_task.apply_async.call_count == 3 - - # Check that document_ids in calls match expected order - for i, call_obj in enumerate(mock_task.apply_async.call_args_list): - called_doc_ids = call_obj[1]["kwargs"]["document_ids"] - assert called_doc_ids == [task_order[i]] - - def test_empty_queue_after_task_completion_cleans_up( - self, tenant_id, dataset_id, document_ids, mock_redis, mock_db_session, mock_dataset - ): - """ - Test cleanup behavior when queue becomes empty after task completion. - - After processing the last task in the queue, the system should: - 1. Detect that no more tasks are waiting - 2. Delete the task key to indicate tenant is idle - 3. Allow new tasks to start fresh processing - - Scenario: - - Process a task - - Check queue for next tasks - - Queue is empty - - Task key should be deleted - - Expected behavior: - - Task key is deleted when queue is empty - - Tenant is marked as idle (no active tasks) - """ - # Arrange - mock_redis.rpop.return_value = None # Empty queue - mock_db_session._shared_data["dataset"] = mock_dataset - - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, mock_task) - - # Assert - expected_task_key = f"tenant_document_indexing_task:{tenant_id}" - - # Verify the task key for this tenant was deleted (do not assert call count; fixtures may be shared). - mock_redis.delete.assert_any_call(expected_task_key) - - deleted_keys = [delete_call.args[0] for delete_call in mock_redis.delete.call_args_list if delete_call.args] - assert expected_task_key in deleted_keys - - deleted_task_key = next(key for key in deleted_keys if key == expected_task_key) - assert tenant_id in deleted_task_key - assert "document_indexing" in deleted_task_key - - def test_billing_disabled_skips_limit_checks( - self, dataset_id, document_ids, mock_db_session, mock_dataset, mock_indexing_runner, mock_feature_service - ): - """ - Test that billing limit checks are skipped when billing is disabled. - - For self-hosted or enterprise deployments where billing is disabled, - the system should not enforce vector space or batch upload limits. - - Scenario: - - Billing is disabled - - Upload 100 documents (would normally exceed limits) - - No limit checks should be performed - - Expected behavior: - - Documents are processed without limit validation - - No errors related to limits - - All documents proceed to indexing - """ - # Arrange - Create many documents - large_batch_ids = [str(uuid.uuid4()) for _ in range(100)] - - mock_documents = [] - for doc_id in large_batch_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.processing_started_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - # Billing disabled - limits should not be checked - mock_feature_service.get_features.return_value.billing.enabled = False - - # Act - _document_indexing(dataset_id, large_batch_ids) - - # Assert - # All documents should be set to parsing (no limit errors) - for doc in mock_documents: - assert doc.indexing_status == IndexingStatus.PARSING - - # IndexingRunner should be called with all documents - mock_indexing_runner.run.assert_called_once() - call_args = mock_indexing_runner.run.call_args[0][0] - assert len(call_args) == 100 - - -class TestIntegration: - """Integration tests for complete task workflows.""" - - def test_complete_workflow_normal_task( - self, tenant_id, dataset_id, document_ids, mock_redis, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test complete workflow for normal document indexing task. - - This tests the full flow from task receipt to completion. - """ - # Arrange - Create actual document objects - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.processing_started_at = None - mock_documents.append(doc) - - # Set up rpop to return None for concurrency check (no more tasks) - mock_redis.rpop.side_effect = [None] - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - normal_document_indexing_task(tenant_id, dataset_id, document_ids) - - # Assert - # Documents should be processed - mock_indexing_runner.run.assert_called_once() - # Session should be closed - assert mock_db_session.close.called - # Task key should be deleted (no more tasks) - assert mock_redis.delete.called - - def test_complete_workflow_priority_task( - self, tenant_id, dataset_id, document_ids, mock_redis, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test complete workflow for priority document indexing task. - - Priority tasks should follow the same flow as normal tasks. - """ - # Arrange - Create actual document objects - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.processing_started_at = None - mock_documents.append(doc) - - # Set up rpop to return None for concurrency check (no more tasks) - mock_redis.rpop.side_effect = [None] - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - # Act - priority_document_indexing_task(tenant_id, dataset_id, document_ids) - - # Assert - mock_indexing_runner.run.assert_called_once() - assert mock_db_session.close.called - assert mock_redis.delete.called - - def test_queue_chain_processing( - self, tenant_id, dataset_id, mock_redis, mock_db_session, mock_dataset, mock_indexing_runner - ): - """ - Test that multiple tasks in queue are processed in sequence. - - When tasks are queued, they should be processed one after another. - """ - # Arrange - task_1_docs = [str(uuid.uuid4())] - task_2_docs = [str(uuid.uuid4())] - - task_2_data = {"tenant_id": tenant_id, "dataset_id": dataset_id, "document_ids": task_2_docs} - - from core.rag.pipeline.queue import TaskWrapper - - wrapper = TaskWrapper(data=task_2_data) - - # First call returns task 2, second call returns None - mock_redis.rpop.side_effect = [wrapper.serialize(), None] - - mock_db_session._shared_data["dataset"] = mock_dataset - - with patch("tasks.document_indexing_task.FeatureService.get_features") as mock_features: - mock_features.return_value.billing.enabled = False - - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - Process first task - _document_indexing_with_tenant_queue(tenant_id, dataset_id, task_1_docs, mock_task) - - # Assert - Second task should be enqueued - assert mock_task.apply_async.called - call_args = mock_task.apply_async.call_args - assert call_args[1]["kwargs"]["document_ids"] == task_2_docs - - -# ============================================================================ -# Additional Edge Case Tests -# ============================================================================ - - -class TestEdgeCases: - """Test edge cases and boundary conditions.""" - - def test_rapid_successive_task_enqueuing(self, tenant_id, dataset_id, mock_redis): - """ - Test rapid successive task enqueuing to the same tenant queue. - - When multiple tasks are enqueued rapidly for the same tenant, - the system should queue them properly without race conditions. - - Scenario: - - First task starts processing (task key exists) - - Multiple tasks enqueued rapidly while first is running - - All should be added to waiting queue - - Expected behavior: - - All tasks are queued (not executed immediately) - - No tasks are lost - - Queue maintains all tasks - """ - # Arrange - document_ids_list = [[str(uuid.uuid4())] for _ in range(5)] - - # Simulate task already running - mock_redis.get.return_value = b"1" - - with patch.object(DocumentIndexingTaskProxy, "features") as mock_features: - mock_features.billing.enabled = True - mock_features.billing.subscription.plan = CloudPlan.PROFESSIONAL - - # Mock the class variable directly - mock_task = Mock() - with patch.object(DocumentIndexingTaskProxy, "PRIORITY_TASK_FUNC", mock_task): - # Act - Enqueue multiple tasks rapidly - for doc_ids in document_ids_list: - proxy = DocumentIndexingTaskProxy(tenant_id, dataset_id, doc_ids) - proxy.delay() - - # Assert - All tasks should be pushed to queue, none executed - assert mock_redis.lpush.call_count == 5 - mock_task.delay.assert_not_called() - - -class TestPerformanceScenarios: - """Test performance-related scenarios and optimizations.""" - - def test_large_document_batch_processing( - self, dataset_id, mock_db_session, mock_dataset, mock_indexing_runner, mock_feature_service - ): - """ - Test processing a large batch of documents at batch limit. - - When processing the maximum allowed batch size, the system - should handle it efficiently without errors. - - Scenario: - - Process exactly batch_upload_limit documents (e.g., 50) - - All documents are valid - - Billing is enabled - - Expected behavior: - - All documents are processed successfully - - No timeout or memory issues - - Batch limit is not exceeded - """ - # Arrange - batch_limit = 50 - document_ids = [str(uuid.uuid4()) for _ in range(batch_limit)] - - mock_documents = [] - for doc_id in document_ids: - doc = MagicMock(spec=Document) - doc.id = doc_id - doc.dataset_id = dataset_id - doc.indexing_status = "waiting" - doc.processing_started_at = None - mock_documents.append(doc) - - # Set shared mock data so all sessions can access it - mock_db_session._shared_data["dataset"] = mock_dataset - mock_db_session._shared_data["documents"] = mock_documents - - # Configure billing with sufficient limits - mock_feature_service.get_features.return_value.billing.enabled = True - mock_feature_service.get_features.return_value.billing.subscription.plan = CloudPlan.PROFESSIONAL - mock_feature_service.get_features.return_value.vector_space.limit = 10000 - mock_feature_service.get_features.return_value.vector_space.size = 0 - - with patch("tasks.document_indexing_task.dify_config.BATCH_UPLOAD_LIMIT", str(batch_limit)): - # Act - _document_indexing(dataset_id, document_ids) - - # Assert - for doc in mock_documents: - assert doc.indexing_status == IndexingStatus.PARSING - - mock_indexing_runner.run.assert_called_once() - call_args = mock_indexing_runner.run.call_args[0][0] - assert len(call_args) == batch_limit - - def test_tenant_queue_handles_burst_traffic(self, tenant_id, dataset_id, mock_redis, mock_db_session, mock_dataset): - """ - Test tenant queue handling burst traffic scenarios. - - When many tasks arrive in a burst for the same tenant, - the queue should handle them efficiently without dropping tasks. - - Scenario: - - 20 tasks arrive rapidly - - Concurrency limit is 3 - - Tasks should be queued and processed in batches - - Expected behavior: - - First 3 tasks are processed immediately - - Remaining tasks wait in queue - - No tasks are lost - """ - # Arrange - num_tasks = 20 - concurrency_limit = 3 - document_ids = [str(uuid.uuid4())] - - # Create waiting tasks - waiting_tasks = [] - for i in range(num_tasks): - task_data = { - "tenant_id": tenant_id, - "dataset_id": dataset_id, - "document_ids": [f"doc_{i}"], - } - from core.rag.pipeline.queue import TaskWrapper - - wrapper = TaskWrapper(data=task_data) - waiting_tasks.append(wrapper.serialize()) - - # Mock rpop to return tasks up to concurrency limit - mock_redis.rpop.side_effect = waiting_tasks[:concurrency_limit] + [None] - mock_db_session._shared_data["dataset"] = mock_dataset - - with patch("tasks.document_indexing_task.dify_config.TENANT_ISOLATED_TASK_CONCURRENCY", concurrency_limit): - with patch("tasks.document_indexing_task.normal_document_indexing_task") as mock_task: - # Act - _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, mock_task) - - # Assert - Should process exactly concurrency_limit tasks - assert mock_task.apply_async.call_count == concurrency_limit - - def test_multiple_tenants_isolated_processing(self, mock_redis): - """ - Test that multiple tenants process tasks in isolation. - - When multiple tenants have tasks running simultaneously, - they should not interfere with each other. - - Scenario: - - Tenant A has tasks in queue - - Tenant B has tasks in queue - - Both process independently - - Expected behavior: - - Each tenant has separate queue - - Each tenant has separate task key - - No cross-tenant interference - """ - # Arrange - tenant_a = str(uuid.uuid4()) - tenant_b = str(uuid.uuid4()) - dataset_id = str(uuid.uuid4()) - document_ids = [str(uuid.uuid4())] - - # Create queues for both tenants - queue_a = TenantIsolatedTaskQueue(tenant_a, "document_indexing") - queue_b = TenantIsolatedTaskQueue(tenant_b, "document_indexing") - - # Act - Set task keys for both tenants - queue_a.set_task_waiting_time() - queue_b.set_task_waiting_time() - - # Assert - Each tenant has independent queue and key - assert queue_a._queue != queue_b._queue - assert queue_a._task_key != queue_b._task_key - assert tenant_a in queue_a._queue - assert tenant_b in queue_b._queue - assert tenant_a in queue_a._task_key - assert tenant_b in queue_b._task_key - - -class TestRobustness: - """Test system robustness and resilience.""" - - def test_task_proxy_handles_feature_service_failure(self, tenant_id, dataset_id, document_ids, mock_redis): - """ - Test that task proxy handles FeatureService failures gracefully. - - If FeatureService fails to retrieve features, the system should - have a fallback or handle the error appropriately. - - Scenario: - - FeatureService.get_features() raises an exception during dispatch - - Task enqueuing should handle the error - - Expected behavior: - - Exception is raised when trying to dispatch - - System doesn't crash unexpectedly - - Error is propagated appropriately - """ - # Arrange - with patch("services.document_indexing_proxy.base.FeatureService.get_features") as mock_get_features: - # Simulate FeatureService failure - mock_get_features.side_effect = Exception("Feature service unavailable") - - # Create proxy instance - proxy = DocumentIndexingTaskProxy(tenant_id, dataset_id, document_ids) - - # Act & Assert - Should raise exception when trying to delay (which accesses features) - with pytest.raises(Exception) as exc_info: - proxy.delay() - - # Verify the exception message - assert "Feature service" in str(exc_info.value) or isinstance(exc_info.value, Exception) - - -class _SessionContext: - def __init__(self, session: MagicMock) -> None: - self._session = session - - def __enter__(self) -> MagicMock: - return self._session - - def __exit__(self, exc_type, exc, tb) -> None: # type: ignore[override] - return None - - -class TestDocumentIndexingTaskSummaryFlow: - """Additional coverage for summary and tenant queue branches.""" - - def test_should_return_when_dataset_missing(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test early return when dataset does not exist.""" - # Arrange - session = MagicMock() - session = MagicMock() - session.scalar.return_value = None # dataset not found - - create_session_mock = MagicMock(return_value=_SessionContext(session)) - monkeypatch.setattr("tasks.document_indexing_task.session_factory.create_session", create_session_mock) - features_mock = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task.FeatureService.get_features", features_mock) - - # Act - _document_indexing("dataset-1", ["doc-1"]) - - # Assert - features_mock.assert_not_called() - - def test_should_mark_documents_error_when_batch_upload_limit_exceeded( - self, monkeypatch: pytest.MonkeyPatch + def test_self_hosted_dispatches_directly_to_priority_task( + self, tenant_id: str, dataset_id: str, document_ids: list[str], mock_redis: MagicMock ) -> None: - """Test batch upload limit triggers error handling.""" - # Arrange - dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") - document = SimpleNamespace(id="doc-1", indexing_status=None, error=None, stopped_at=None) + with ( + patch.object(DocumentIndexingTaskProxy, "features") as features, + patch.object(DocumentIndexingTaskProxy, "PRIORITY_TASK_FUNC", Mock()) as task, + ): + features.billing.enabled = False + DocumentIndexingTaskProxy(tenant_id, dataset_id, document_ids).delay() - session = MagicMock() - - def _scalar_se(stmt): - entity = stmt.column_descriptions[0].get("entity") - if entity is Dataset: - return dataset - return document - - session.scalar.side_effect = _scalar_se - - monkeypatch.setattr( - "tasks.document_indexing_task.session_factory.create_session", - MagicMock(return_value=_SessionContext(session)), + task.delay.assert_called_once_with( + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=document_ids, ) - features = SimpleNamespace( - billing=SimpleNamespace( - enabled=True, - subscription=SimpleNamespace(plan=CloudPlan.PROFESSIONAL), - ), - vector_space=SimpleNamespace(limit=0, size=0), + @pytest.mark.parametrize( + ("plan", "task_attribute"), + [ + (CloudPlan.SANDBOX, "NORMAL_TASK_FUNC"), + (CloudPlan.PROFESSIONAL, "PRIORITY_TASK_FUNC"), + ], + ) + def test_cloud_dispatches_first_task_through_tenant_queue( + self, + tenant_id: str, + dataset_id: str, + document_ids: list[str], + mock_redis: MagicMock, + plan: CloudPlan, + task_attribute: str, + ) -> None: + with ( + patch.object(DocumentIndexingTaskProxy, "features") as features, + patch.object(DocumentIndexingTaskProxy, task_attribute, Mock()) as task, + ): + features.billing.enabled = True + features.billing.subscription.plan = plan + DocumentIndexingTaskProxy(tenant_id, dataset_id, document_ids).delay() + + mock_redis.setex.assert_called() + task.delay.assert_called_once() + + def test_running_tenant_task_queues_followup_work( + self, tenant_id: str, dataset_id: str, document_ids: list[str], mock_redis: MagicMock + ) -> None: + mock_redis.get.return_value = b"1" + with ( + patch.object(DocumentIndexingTaskProxy, "features") as features, + patch.object(DocumentIndexingTaskProxy, "PRIORITY_TASK_FUNC", Mock()) as task, + ): + features.billing.enabled = True + features.billing.subscription.plan = CloudPlan.PROFESSIONAL + DocumentIndexingTaskProxy(tenant_id, dataset_id, document_ids).delay() + + mock_redis.lpush.assert_called_once() + task.delay.assert_not_called() + + +class TestDocumentIndexing: + def test_legacy_task_persists_parsing_before_running( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + document_ids: list[str], + indexing_runner: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=document_ids, ) - monkeypatch.setattr( - "tasks.document_indexing_task.FeatureService.get_features", MagicMock(return_value=features) + _patch_features(monkeypatch, _features()) + + def assert_committed_parsing(documents: list[Document], session: Session) -> None: + assert all(document.indexing_status == IndexingStatus.PARSING for document in documents) + assert all(document.processing_started_at is not None for document in documents) + assert all(session.get(Document, document.id) is document for document in documents) + + indexing_runner.run.side_effect = assert_committed_parsing + document_indexing_task.run(dataset_id, document_ids) + + persisted = _persisted_documents(sqlite_session, document_ids) + assert [document.indexing_status for document in persisted] == [IndexingStatus.PARSING] * 3 + indexing_runner._constructor_mock.assert_called_once_with(enforce_vector_space_admission=True) + indexing_runner.run.assert_called_once() + assert isinstance(indexing_runner.run.call_args.args[1], Session) + + def test_only_existing_documents_are_processed( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + document_ids: list[str], + indexing_runner: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + existing_ids = [document_ids[0], document_ids[2]] + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=existing_ids, ) - monkeypatch.setattr("tasks.document_indexing_task.dify_config.BATCH_UPLOAD_LIMIT", "1") + _patch_features(monkeypatch, _features()) - # Act - _document_indexing("dataset-1", ["doc-1", "doc-2"]) + _document_indexing(dataset_id, document_ids) - # Assert - assert document.indexing_status == "error" - assert "batch upload limit" in document.error - session.commit.assert_called_once() + processed = indexing_runner.run.call_args.args[0] + assert {document.id for document in processed} == set(existing_ids) + assert sqlite_session.get(Document, document_ids[1]) is None - def test_should_queue_summary_generation_for_completed_documents(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test summary generation is queued for eligible documents.""" - # Arrange - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - indexing_technique="high_quality", + def test_empty_batch_still_reaches_runner( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + indexing_runner: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=[], + ) + _patch_features(monkeypatch, _features()) + + _document_indexing(dataset_id, []) + + assert indexing_runner.run.call_args.args[0] == [] + assert isinstance(indexing_runner.run.call_args.args[1], Session) + + def test_missing_dataset_returns_before_feature_lookup( + self, dataset_id: str, document_ids: list[str], monkeypatch: pytest.MonkeyPatch + ) -> None: + get_features = _patch_features(monkeypatch, _features()) + runner_class = MagicMock() + monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", runner_class) + + _document_indexing(dataset_id, document_ids) + + get_features.assert_not_called() + runner_class.assert_not_called() + + @pytest.mark.parametrize( + ("features", "batch_limit", "message"), + [ + (_features(billing_enabled=True), 1, "batch upload limit"), + (_features(billing_enabled=True, plan=CloudPlan.SANDBOX), 100, "does not support batch upload"), + (_features(billing_enabled=True, vector_limit=100, vector_size=100), 100, "over the limit"), + ], + ) + def test_validation_failure_marks_every_scoped_document_error( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + document_ids: list[str], + monkeypatch: pytest.MonkeyPatch, + features: SimpleNamespace, + batch_limit: int, + message: str, + ) -> None: + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=document_ids, + ) + control_dataset_id = str(uuid.uuid4()) + control_document_id = str(uuid.uuid4()) + _persist_indexing_rows( + sqlite_session, + tenant_id=str(uuid.uuid4()), + dataset_id=control_dataset_id, + document_ids=[control_document_id], + ) + _patch_features(monkeypatch, features) + monkeypatch.setattr("tasks.document_indexing_task.dify_config.BATCH_UPLOAD_LIMIT", str(batch_limit)) + + _document_indexing(dataset_id, document_ids) + + persisted = _persisted_documents(sqlite_session, document_ids) + assert all(document.indexing_status == IndexingStatus.ERROR for document in persisted) + assert all(document.error and message in document.error for document in persisted) + assert all(document.stopped_at is not None for document in persisted) + control = sqlite_session.get(Document, control_document_id) + assert control is not None + assert control.indexing_status == IndexingStatus.WAITING + + @pytest.mark.parametrize("error", [DocumentIsPausedError("paused"), RuntimeError("boom")]) + def test_runner_failure_stops_before_summary_dispatch( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + document_ids: list[str], + indexing_runner: MagicMock, + monkeypatch: pytest.MonkeyPatch, + error: Exception, + ) -> None: + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=document_ids, summary_index_setting={"enable": True}, + need_summary=[True] * len(document_ids), ) + _patch_features(monkeypatch, _features()) + indexing_runner.run.side_effect = error + summary_delay = MagicMock() + monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", summary_delay) - doc_eligible = SimpleNamespace( - id="doc-1", - indexing_status="completed", - doc_form="text", - need_summary=True, - ) - doc_skip_form = SimpleNamespace( - id="doc-2", - indexing_status="completed", - doc_form="qa_model", - need_summary=True, - ) - doc_skip_status = SimpleNamespace( - id="doc-3", - indexing_status="processing", - doc_form="text", - need_summary=True, - ) + _document_indexing(dataset_id, document_ids) - phase1_docs = [SimpleNamespace(id="doc-1"), SimpleNamespace(id="doc-2"), SimpleNamespace(id="doc-3")] + summary_delay.assert_not_called() + persisted = _persisted_documents(sqlite_session, document_ids) + assert all(document.indexing_status == IndexingStatus.PARSING for document in persisted) - session1 = MagicMock() - session2 = MagicMock() - session2.begin.return_value = nullcontext() - session3 = MagicMock() - session4 = MagicMock() - session1.scalar.return_value = dataset - session2.scalars.return_value = MagicMock(all=MagicMock(return_value=phase1_docs)) - session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=phase1_docs)) - session4.scalar.return_value = dataset - session4.scalars.return_value = MagicMock( - all=MagicMock(return_value=[doc_eligible, doc_skip_form, doc_skip_status]) - ) - - create_session_mock = MagicMock( - side_effect=[ - _SessionContext(session1), - _SessionContext(session2), - _SessionContext(session3), - _SessionContext(session4), - ] - ) - monkeypatch.setattr("tasks.document_indexing_task.session_factory.create_session", create_session_mock) - - features = SimpleNamespace( - billing=SimpleNamespace(enabled=False), - vector_space=SimpleNamespace(limit=0, size=0), - ) - monkeypatch.setattr( - "tasks.document_indexing_task.FeatureService.get_features", MagicMock(return_value=features) - ) - - indexing_runner = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", MagicMock(return_value=indexing_runner)) - delay_mock = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", delay_mock) - - # Act - _document_indexing("dataset-1", ["doc-1", "doc-2", "doc-3"]) - - # Assert - delay_mock.assert_called_once_with("dataset-1", "doc-1", None) - - def test_should_continue_when_summary_queue_fails(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test summary queueing errors are swallowed.""" - # Arrange - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - indexing_technique="high_quality", +class TestSummaryDispatch: + def test_only_eligible_completed_documents_queue_summaries( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + document_ids: list[str], + indexing_runner: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=document_ids, summary_index_setting={"enable": True}, + document_forms=[ + IndexStructureType.PARAGRAPH_INDEX, + IndexStructureType.QA_INDEX, + IndexStructureType.PARAGRAPH_INDEX, + ], + need_summary=[True, True, True], ) + _patch_features(monkeypatch, _features()) - doc_eligible = SimpleNamespace( - id="doc-1", - indexing_status="completed", - doc_form="text", - need_summary=True, - ) + def finish_documents(documents: list[Document], _session: Session) -> None: + documents[0].indexing_status = IndexingStatus.COMPLETED + documents[1].indexing_status = IndexingStatus.COMPLETED + documents[2].indexing_status = IndexingStatus.INDEXING - session1 = MagicMock() - session2 = MagicMock() - session2.begin.return_value = nullcontext() - session3 = MagicMock() - session4 = MagicMock() + indexing_runner.run.side_effect = finish_documents + summary_delay = MagicMock() + monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", summary_delay) - session1.scalar.return_value = dataset - session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session4.scalar.return_value = dataset - session4.scalars.return_value = MagicMock(all=MagicMock(return_value=[doc_eligible])) + _document_indexing(dataset_id, document_ids) - monkeypatch.setattr( - "tasks.document_indexing_task.session_factory.create_session", - MagicMock( - side_effect=[ - _SessionContext(session1), - _SessionContext(session2), - _SessionContext(session3), - _SessionContext(session4), - ] - ), - ) + summary_delay.assert_called_once_with(dataset_id, document_ids[0], None) - features = SimpleNamespace( - billing=SimpleNamespace(enabled=False), - vector_space=SimpleNamespace(limit=0, size=0), - ) - monkeypatch.setattr( - "tasks.document_indexing_task.FeatureService.get_features", MagicMock(return_value=features) - ) - - indexing_runner = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", MagicMock(return_value=indexing_runner)) - delay_mock = MagicMock(side_effect=Exception("boom")) - monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", delay_mock) - - # Act - _document_indexing("dataset-1", ["doc-1"]) - - # Assert - delay_mock.assert_called_once_with("dataset-1", "doc-1", None) - - def test_should_return_when_dataset_missing_after_indexing(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test early return when dataset is missing after indexing.""" - # Arrange - dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") - - session1 = MagicMock() - session2 = MagicMock() - session2.begin.return_value = nullcontext() - session3 = MagicMock() - session4 = MagicMock() - session1.scalar.return_value = dataset - session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session4.scalar.return_value = None # dataset not found after indexing - - monkeypatch.setattr( - "tasks.document_indexing_task.session_factory.create_session", - MagicMock( - side_effect=[ - _SessionContext(session1), - _SessionContext(session2), - _SessionContext(session3), - _SessionContext(session4), - ] - ), - ) - - features = SimpleNamespace( - billing=SimpleNamespace(enabled=False), - vector_space=SimpleNamespace(limit=0, size=0), - ) - monkeypatch.setattr( - "tasks.document_indexing_task.FeatureService.get_features", MagicMock(return_value=features) - ) - monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", MagicMock(return_value=MagicMock())) - - # Act - _document_indexing("dataset-1", ["doc-1"]) - - # Assert - session4.scalar.assert_called() - - def test_should_skip_summary_when_not_high_quality(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test summary generation skipped when indexing_technique is not high_quality.""" - # Arrange - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - indexing_technique="economy", + def test_summary_queue_failure_does_not_fail_indexing( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + indexing_runner: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + document_id = str(uuid.uuid4()) + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=[document_id], summary_index_setting={"enable": True}, + need_summary=[True], ) - session1 = MagicMock() - session2 = MagicMock() - session2.begin.return_value = nullcontext() - session3 = MagicMock() - session4 = MagicMock() - - session1.scalar.return_value = dataset - session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session4.scalar.return_value = dataset - - monkeypatch.setattr( - "tasks.document_indexing_task.session_factory.create_session", - MagicMock( - side_effect=[ - _SessionContext(session1), - _SessionContext(session2), - _SessionContext(session3), - _SessionContext(session4), - ] - ), + _patch_features(monkeypatch, _features()) + indexing_runner.run.side_effect = lambda documents, _session: setattr( + documents[0], "indexing_status", IndexingStatus.COMPLETED ) + summary_delay = MagicMock(side_effect=RuntimeError("queue unavailable")) + monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", summary_delay) - features = SimpleNamespace( - billing=SimpleNamespace(enabled=False), - vector_space=SimpleNamespace(limit=0, size=0), - ) - monkeypatch.setattr( - "tasks.document_indexing_task.FeatureService.get_features", MagicMock(return_value=features) - ) - monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", MagicMock(return_value=MagicMock())) + _document_indexing(dataset_id, [document_id]) - delay_mock = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", delay_mock) + summary_delay.assert_called_once_with(dataset_id, document_id, None) + persisted = _persisted_documents(sqlite_session, [document_id])[0] + assert persisted.indexing_status == IndexingStatus.COMPLETED - # Act - _document_indexing("dataset-1", ["doc-1"]) - - # Assert - delay_mock.assert_not_called() - - def test_should_skip_summary_generation_when_indexing_paused(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test summary generation is skipped when indexing is paused.""" - # Arrange - dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") - - session1 = MagicMock() - session2 = MagicMock() - session2.begin.return_value = nullcontext() - session1.scalar.return_value = dataset - session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session3 = MagicMock() - session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - - create_session_mock = MagicMock( - side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)] - ) - monkeypatch.setattr("tasks.document_indexing_task.session_factory.create_session", create_session_mock) - - features = SimpleNamespace( - billing=SimpleNamespace(enabled=False), - vector_space=SimpleNamespace(limit=0, size=0), - ) - monkeypatch.setattr( - "tasks.document_indexing_task.FeatureService.get_features", MagicMock(return_value=features) - ) - - runner = MagicMock() - runner.run.side_effect = DocumentIsPausedError("paused") - monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", MagicMock(return_value=runner)) - delay_mock = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", delay_mock) - - # Act - _document_indexing("dataset-1", ["doc-1"]) - - # Assert - delay_mock.assert_not_called() - - def test_should_handle_indexing_runner_exception(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test generic indexing runner exception is handled.""" - # Arrange - dataset = SimpleNamespace(id="dataset-1", tenant_id="tenant-1") - - session1 = MagicMock() - session2 = MagicMock() - session2.begin.return_value = nullcontext() - session1.scalar.return_value = dataset - session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session3 = MagicMock() - session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - - monkeypatch.setattr( - "tasks.document_indexing_task.session_factory.create_session", - MagicMock(side_effect=[_SessionContext(session1), _SessionContext(session2), _SessionContext(session3)]), - ) - - features = SimpleNamespace( - billing=SimpleNamespace(enabled=False), - vector_space=SimpleNamespace(limit=0, size=0), - ) - monkeypatch.setattr( - "tasks.document_indexing_task.FeatureService.get_features", MagicMock(return_value=features) - ) - - runner = MagicMock() - runner.run.side_effect = RuntimeError("boom") - monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", MagicMock(return_value=runner)) - - delay_mock = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", delay_mock) - - # Act - _document_indexing("dataset-1", ["doc-1"]) - - # Assert - delay_mock.assert_not_called() - - def test_should_log_missing_document_entry_in_summary_list(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test falsey document entries are handled in summary iteration.""" - - # Arrange - class _FalseyDocument: - def __init__(self, doc_id: str) -> None: - self.id = doc_id - - def __bool__(self) -> bool: - return False - - dataset = SimpleNamespace( - id="dataset-1", - tenant_id="tenant-1", - indexing_technique="high_quality", + def test_economy_indexing_skips_summary_generation( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + indexing_runner: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + document_id = str(uuid.uuid4()) + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=[document_id], + indexing_technique=IndexTechniqueType.ECONOMY, summary_index_setting={"enable": True}, + need_summary=[True], ) - session1 = MagicMock() - session2 = MagicMock() - session2.begin.return_value = nullcontext() - session3 = MagicMock() - session4 = MagicMock() + _patch_features(monkeypatch, _features()) + indexing_runner.run.side_effect = lambda documents, _session: setattr( + documents[0], "indexing_status", IndexingStatus.COMPLETED + ) + summary_delay = MagicMock() + monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", summary_delay) - session1.scalar.return_value = dataset - session2.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session3.scalar.return_value = dataset - session3.scalars.return_value = MagicMock(all=MagicMock(return_value=[SimpleNamespace(id="doc-1")])) - session4.scalar.return_value = dataset - session4.scalars.return_value = MagicMock(all=MagicMock(return_value=[_FalseyDocument("missing-doc")])) + _document_indexing(dataset_id, [document_id]) + summary_delay.assert_not_called() + + def test_dataset_removed_by_runner_is_absent_from_summary_phase( + self, + sqlite_session: Session, + tenant_id: str, + dataset_id: str, + indexing_runner: MagicMock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + document_id = str(uuid.uuid4()) + _persist_indexing_rows( + sqlite_session, + tenant_id=tenant_id, + dataset_id=dataset_id, + document_ids=[document_id], + summary_index_setting={"enable": True}, + need_summary=[True], + ) + _patch_features(monkeypatch, _features()) + + def remove_dataset(_documents: list[Document], session: Session) -> None: + dataset = session.get(Dataset, dataset_id) + assert dataset is not None + session.delete(dataset) + + indexing_runner.run.side_effect = remove_dataset + summary_delay = MagicMock() + monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", summary_delay) + + _document_indexing(dataset_id, [document_id]) + + sqlite_session.expire_all() + assert sqlite_session.get(Dataset, dataset_id) is None + summary_delay.assert_not_called() + + +class TestTenantQueue: + def test_followup_tasks_are_dispatched_with_one_shared_producer( + self, tenant_id: str, dataset_id: str, document_ids: list[str], monkeypatch: pytest.MonkeyPatch + ) -> None: + next_documents = [str(uuid.uuid4())] + queue = MagicMock() + queue.pull_tasks.return_value = [ + {"tenant_id": tenant_id, "dataset_id": dataset_id, "document_ids": next_documents} + ] + monkeypatch.setattr("tasks.document_indexing_task.TenantIsolatedTaskQueue", MagicMock(return_value=queue)) + monkeypatch.setattr("tasks.document_indexing_task._document_indexing", MagicMock()) + producer = object() monkeypatch.setattr( - "tasks.document_indexing_task.session_factory.create_session", - MagicMock( - side_effect=[ - _SessionContext(session1), - _SessionContext(session2), - _SessionContext(session3), - _SessionContext(session4), - ] - ), + "tasks.document_indexing_task.current_app.producer_or_acquire", + MagicMock(return_value=nullcontext(producer)), ) + task = MagicMock() - features = SimpleNamespace( - billing=SimpleNamespace(enabled=False), - vector_space=SimpleNamespace(limit=0, size=0), + _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, task) + + task.apply_async.assert_called_once_with( + kwargs={"tenant_id": tenant_id, "dataset_id": dataset_id, "document_ids": next_documents}, + producer=producer, ) - monkeypatch.setattr( - "tasks.document_indexing_task.FeatureService.get_features", MagicMock(return_value=features) - ) - monkeypatch.setattr("tasks.document_indexing_task.IndexingRunner", MagicMock(return_value=MagicMock())) + queue.set_task_waiting_time.assert_called_once() + queue.delete_task_key.assert_not_called() - delay_mock = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task.generate_summary_index_task.delay", delay_mock) + def test_queue_cleanup_runs_when_indexing_fails( + self, tenant_id: str, dataset_id: str, document_ids: list[str], monkeypatch: pytest.MonkeyPatch + ) -> None: + queue = MagicMock() + queue.pull_tasks.return_value = [] + monkeypatch.setattr("tasks.document_indexing_task.TenantIsolatedTaskQueue", MagicMock(return_value=queue)) + indexing = MagicMock(side_effect=RuntimeError("indexing failed")) + monkeypatch.setattr("tasks.document_indexing_task._document_indexing", indexing) - # Act - _document_indexing("dataset-1", ["doc-1"]) + _document_indexing_with_tenant_queue(tenant_id, dataset_id, document_ids, MagicMock()) - # Assert - delay_mock.assert_not_called() + queue.delete_task_key.assert_called_once() - def test_normal_document_indexing_task_should_delegate(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test normal indexing task delegates to tenant queue handler.""" - # Arrange - handler = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task._document_indexing_with_tenant_queue", handler) + @pytest.mark.parametrize("task", [normal_document_indexing_task, priority_document_indexing_task]) + def test_celery_entrypoints_delegate_to_tenant_queue( + self, + task: object, + tenant_id: str, + dataset_id: str, + document_ids: list[str], + monkeypatch: pytest.MonkeyPatch, + ) -> None: + delegate = MagicMock() + monkeypatch.setattr("tasks.document_indexing_task._document_indexing_with_tenant_queue", delegate) - # Act - normal_document_indexing_task("tenant-1", "dataset-1", ["doc-1"]) + task.run(tenant_id, dataset_id, document_ids) # type: ignore[attr-defined] - # Assert - handler.assert_called_once_with("tenant-1", "dataset-1", ["doc-1"], normal_document_indexing_task) - - def test_priority_document_indexing_task_should_delegate(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Test priority indexing task delegates to tenant queue handler.""" - # Arrange - handler = MagicMock() - monkeypatch.setattr("tasks.document_indexing_task._document_indexing_with_tenant_queue", handler) - - # Act - priority_document_indexing_task("tenant-1", "dataset-1", ["doc-1"]) - - # Assert - handler.assert_called_once_with("tenant-1", "dataset-1", ["doc-1"], priority_document_indexing_task) + delegate.assert_called_once() diff --git a/api/tests/unit_tests/tasks/test_document_indexing_update_task.py b/api/tests/unit_tests/tasks/test_document_indexing_update_task.py index 59eaf497960..62a39e9ebb0 100644 --- a/api/tests/unit_tests/tasks/test_document_indexing_update_task.py +++ b/api/tests/unit_tests/tasks/test_document_indexing_update_task.py @@ -1,520 +1,276 @@ -""" -Unit tests for document_indexing_update_task summary generation. +"""SQLite-backed tests for document update indexing and summary generation.""" -After updating a document via the API, the summary index should be -regenerated under the same conditions as during initial creation: -- indexing_technique is HIGH_QUALITY -- summary_index_setting has enable=True -- document.indexing_status is COMPLETED -- document.doc_form is not QA_INDEX -- document.need_summary is True -""" +from __future__ import annotations -from contextlib import nullcontext -from types import SimpleNamespace +import uuid +from collections.abc import Callable from unittest.mock import MagicMock import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session +import tasks.document_indexing_update_task as task_module from core.indexing_runner import DocumentIsPausedError +from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus from tasks.document_indexing_update_task import document_indexing_update_task -class _SessionContext: - """Minimal context manager that yields a mock session.""" - - def __init__(self, session: MagicMock) -> None: - self._session = session - - def __enter__(self) -> MagicMock: - return self._session - - def __exit__(self, exc_type, exc, tb) -> None: # type: ignore[override] - return None +@pytest.fixture +def task_harness( + sqlite_session: Session, + monkeypatch: pytest.MonkeyPatch, +) -> tuple[MagicMock, MagicMock]: + """Bind task-owned sessions to SQLite and keep only external boundaries mocked.""" + engine = sqlite_session.get_bind() + monkeypatch.setattr( + task_module.session_factory, + "create_session", + lambda: Session(engine, expire_on_commit=False), + ) + runner = MagicMock() + processor = MagicMock() + monkeypatch.setattr(task_module, "IndexingRunner", MagicMock(return_value=runner)) + monkeypatch.setattr( + task_module, + "IndexProcessorFactory", + MagicMock(return_value=MagicMock(init_index_processor=MagicMock(return_value=processor))), + ) + return runner, processor -def _make_dataset_and_documents( +def _persist_rows( + session: Session, *, - dataset_id: str = "ds-1", - document_id: str = "doc-1", - indexing_technique: str = "high_quality", + indexing_technique: IndexTechniqueType = IndexTechniqueType.HIGH_QUALITY, summary_index_setting: dict | None = None, - doc_form: str = "text_model", + doc_form: IndexStructureType = IndexStructureType.PARAGRAPH_INDEX, need_summary: bool = True, -): - """Create mock dataset and document objects. - - Returns (dataset, doc_for_session1, doc_for_session3). - - session1 doc: before IndexingRunner runs (status irrelevant for summary). - session3 doc: re-queried after IndexingRunner completes — normally COMPLETED. - """ - dataset = SimpleNamespace( + with_segment: bool = False, +) -> tuple[Dataset, Document]: + tenant_id = str(uuid.uuid4()) + dataset_id = str(uuid.uuid4()) + document_id = str(uuid.uuid4()) + created_by = str(uuid.uuid4()) + dataset = Dataset( id=dataset_id, + tenant_id=tenant_id, + name="Update dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, indexing_technique=indexing_technique, summary_index_setting=summary_index_setting, ) - doc_s1 = SimpleNamespace( + document = Document( id=document_id, + tenant_id=tenant_id, dataset_id=dataset_id, - indexing_status="waiting", + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="document.txt", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, + indexing_status=IndexingStatus.WAITING, doc_form=doc_form, need_summary=need_summary, ) - # After IndexingRunner.run the document status is COMPLETED in the DB - doc_s3 = SimpleNamespace( - id=document_id, - dataset_id=dataset_id, - indexing_status="completed", - doc_form=doc_form, - need_summary=need_summary, + rows: list[object] = [dataset, document] + if with_segment: + rows.append( + DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset_id, + document_id=document_id, + position=1, + content="segment", + word_count=1, + tokens=1, + created_by=created_by, + index_node_id="node-1", + ) + ) + session.add_all(rows) + session.commit() + return dataset, document + + +def _complete_indexing(documents: list[Document], _session: Session) -> None: + for document in documents: + document.indexing_status = IndexingStatus.COMPLETED + + +def test_queues_summary_when_all_persisted_conditions_match( + sqlite_session: Session, + task_harness: tuple[MagicMock, MagicMock], + monkeypatch: pytest.MonkeyPatch, +) -> None: + runner, _processor = task_harness + dataset, document = _persist_rows(sqlite_session, summary_index_setting={"enable": True}) + runner.run.side_effect = _complete_indexing + delay = MagicMock() + monkeypatch.setattr(task_module.generate_summary_index_task, "delay", delay) + + document_indexing_update_task(dataset.id, document.id) + + delay.assert_called_once_with(dataset.id, document.id, None) + sqlite_session.expire_all() + assert sqlite_session.get(Document, document.id).indexing_status == IndexingStatus.COMPLETED # type: ignore[union-attr] + + +@pytest.mark.parametrize( + ("dataset_changes", "document_changes"), + [ + ({"indexing_technique": IndexTechniqueType.ECONOMY}, {}), + ({"summary_index_setting": None}, {}), + ({"summary_index_setting": {"enable": False}}, {}), + ({"summary_index_setting": {"enable": True}}, {"need_summary": False}), + ( + {"summary_index_setting": {"enable": True}}, + {"doc_form": IndexStructureType.QA_INDEX}, + ), + ], +) +def test_skips_summary_when_persisted_eligibility_does_not_match( + sqlite_session: Session, + task_harness: tuple[MagicMock, MagicMock], + monkeypatch: pytest.MonkeyPatch, + dataset_changes: dict, + document_changes: dict, +) -> None: + runner, _processor = task_harness + dataset, document = _persist_rows( + sqlite_session, + summary_index_setting={"enable": True}, ) - return dataset, doc_s1, doc_s3 + for key, value in dataset_changes.items(): + setattr(dataset, key, value) + for key, value in document_changes.items(): + setattr(document, key, value) + sqlite_session.commit() + runner.run.side_effect = _complete_indexing + delay = MagicMock() + monkeypatch.setattr(task_module.generate_summary_index_task, "delay", delay) + + document_indexing_update_task(dataset.id, document.id) + + delay.assert_not_called() -def _patch_all(monkeypatch: pytest.MonkeyPatch, *, sessions, runner, processor): - """Wire up all mocks for document_indexing_update_task.""" +@pytest.mark.parametrize( + "error_factory", + [ + lambda _document_id: RuntimeError("indexing failed"), + lambda document_id: DocumentIsPausedError(f"{document_id} is paused"), + ], +) +def test_skips_summary_when_indexing_fails_or_is_paused( + sqlite_session: Session, + task_harness: tuple[MagicMock, MagicMock], + monkeypatch: pytest.MonkeyPatch, + error_factory: Callable[[str], Exception], +) -> None: + runner, _processor = task_harness + dataset, document = _persist_rows(sqlite_session, summary_index_setting={"enable": True}) + runner.run.side_effect = error_factory(document.id) + delay = MagicMock() + monkeypatch.setattr(task_module.generate_summary_index_task, "delay", delay) + + document_indexing_update_task(dataset.id, document.id) + + delay.assert_not_called() + + +def test_returns_without_opening_external_boundaries_when_document_is_missing( + task_harness: tuple[MagicMock, MagicMock], +) -> None: + runner, processor = task_harness + + document_indexing_update_task(str(uuid.uuid4()), str(uuid.uuid4())) + + runner.run.assert_not_called() + processor.clean.assert_not_called() + + +def test_skips_summary_when_dataset_is_removed_after_indexing( + sqlite_session: Session, + task_harness: tuple[MagicMock, MagicMock], + monkeypatch: pytest.MonkeyPatch, +) -> None: + runner, _processor = task_harness + dataset, document = _persist_rows(sqlite_session, summary_index_setting={"enable": True}) + + def complete_and_remove(documents: list[Document], session: Session) -> None: + _complete_indexing(documents, session) + persisted_dataset = session.get(Dataset, dataset.id) + assert persisted_dataset is not None + session.delete(persisted_dataset) + + runner.run.side_effect = complete_and_remove + delay = MagicMock() + monkeypatch.setattr(task_module.generate_summary_index_task, "delay", delay) + + document_indexing_update_task(dataset.id, document.id) + + delay.assert_not_called() + + +def test_skips_summary_when_runner_leaves_document_incomplete( + sqlite_session: Session, + task_harness: tuple[MagicMock, MagicMock], + monkeypatch: pytest.MonkeyPatch, +) -> None: + runner, _processor = task_harness + dataset, document = _persist_rows(sqlite_session, summary_index_setting={"enable": True}) + runner.run.return_value = None + delay = MagicMock() + monkeypatch.setattr(task_module.generate_summary_index_task, "delay", delay) + + document_indexing_update_task(dataset.id, document.id) + + delay.assert_not_called() + + +def test_queue_failure_is_swallowed_after_successful_indexing( + sqlite_session: Session, + task_harness: tuple[MagicMock, MagicMock], + monkeypatch: pytest.MonkeyPatch, +) -> None: + runner, _processor = task_harness + dataset, document = _persist_rows(sqlite_session, summary_index_setting={"enable": True}) + runner.run.side_effect = _complete_indexing monkeypatch.setattr( - "tasks.document_indexing_update_task.session_factory.create_session", - MagicMock(side_effect=sessions), - ) - monkeypatch.setattr( - "tasks.document_indexing_update_task.IndexProcessorFactory", - MagicMock(return_value=MagicMock(init_index_processor=MagicMock(return_value=processor))), - ) - monkeypatch.setattr( - "tasks.document_indexing_update_task.IndexingRunner", - MagicMock(return_value=runner), + task_module.generate_summary_index_task, + "delay", + MagicMock(side_effect=RuntimeError("queue unavailable")), ) - -def _session_with_begin(): - """Create a mock session with a begin() context manager.""" - s = MagicMock() - s.begin.return_value = nullcontext() - return s - - -class TestUpdateTaskSummaryGeneration: - """Tests for summary index generation in the document update task. - - The update task creates sessions in this order: - 1. session1: fetch document + dataset + segments (uses begin()) - 2. session2: delete segments — only if segments exist (uses begin()) - 3. session3: summary check — only if indexing succeeded (no begin()) - - With empty segments (default), only sessions 1 and 3 are created. - """ - - def test_should_queue_summary_when_conditions_met(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary task is queued when all conditions are met.""" - dataset, doc_s1, doc_s3 = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - session3 = MagicMock() - session3.scalar.side_effect = [doc_s3, dataset] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_called_once_with("ds-1", "doc-1", None) - - def test_should_not_queue_when_not_high_quality(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when indexing_technique is not high_quality.""" - dataset, doc_s1, _ = _make_dataset_and_documents( - indexing_technique="economy", - summary_index_setting={"enable": True}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - session3 = MagicMock() - session3.scalar.side_effect = [doc_s1, dataset] # dataset.indexing_technique == "economy" - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_not_queue_when_summary_setting_disabled(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when summary_index_setting has enable=False.""" - dataset, doc_s1, _ = _make_dataset_and_documents( - summary_index_setting={"enable": False}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - session3 = MagicMock() - session3.scalar.side_effect = [doc_s1, dataset] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_not_queue_when_summary_setting_none(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when summary_index_setting is None.""" - dataset, doc_s1, _ = _make_dataset_and_documents( - summary_index_setting=None, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - session3 = MagicMock() - session3.scalar.side_effect = [doc_s1, dataset] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_not_queue_when_need_summary_false(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when document.need_summary is False.""" - dataset, doc_s1, doc_s3 = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - need_summary=False, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - session3 = MagicMock() - session3.scalar.side_effect = [doc_s3, dataset] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_not_queue_when_qa_index_form(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when doc_form is QA_INDEX.""" - dataset, doc_s1, doc_s3 = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - doc_form="qa_model", - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - session3 = MagicMock() - session3.scalar.side_effect = [doc_s3, dataset] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_not_queue_when_indexing_fails(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when IndexingRunner.run raises.""" - dataset, doc_s1, _ = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - runner = MagicMock() - runner.run.side_effect = Exception("indexing failed") - processor = MagicMock() - - # Only session1 needed — task returns early after indexing failure - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_not_queue_when_document_is_paused(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when IndexingRunner raises DocumentIsPausedError.""" - - dataset, doc_s1, _ = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - runner = MagicMock() - runner.run.side_effect = DocumentIsPausedError("doc-1 is paused") - processor = MagicMock() - - # Only session1 needed — task returns early after paused error - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_not_queue_when_dataset_not_found_after_indexing(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when the dataset disappears after indexing.""" - dataset, doc_s1, _ = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - # Session 3: dataset is None - session3 = MagicMock() - session3.scalar.side_effect = [doc_s1, None] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_not_queue_when_document_not_completed_after_indexing(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Summary is skipped when document indexing_status is not COMPLETED after indexing.""" - dataset, doc_s1, _ = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - # Document still in error status after indexing - doc_s3_error = SimpleNamespace( - id="doc-1", - dataset_id="ds-1", - indexing_status="error", - doc_form="text_model", - need_summary=True, - ) - session3 = MagicMock() - session3.scalar.side_effect = [doc_s3_error, dataset] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_not_called() - - def test_should_swallow_summary_queue_error(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Task should not raise when generate_summary_index_task.delay raises.""" - dataset, doc_s1, doc_s3 = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[])) - - session3 = MagicMock() - session3.scalar.side_effect = [doc_s3, dataset] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[_SessionContext(session1), _SessionContext(session3)], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock(side_effect=Exception("queue full")) - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - # Should not raise - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_called_once_with("ds-1", "doc-1", None) - - def test_should_queue_summary_with_segments_and_session2(self, monkeypatch: pytest.MonkeyPatch) -> None: - """When segments exist, session2 is also created for deletion. - Verify summary generation still works correctly.""" - dataset, doc_s1, doc_s3 = _make_dataset_and_documents( - summary_index_setting={"enable": True}, - ) - - session1 = _session_with_begin() - session1.scalar.side_effect = [doc_s1, dataset] - seg = SimpleNamespace(index_node_id="node-1") - session1.scalars.return_value = MagicMock(all=MagicMock(return_value=[seg])) - - session3 = MagicMock() - session3.scalar.side_effect = [doc_s3, dataset] - - runner = MagicMock() - processor = MagicMock() - - _patch_all( - monkeypatch, - sessions=[ - _SessionContext(session1), - _SessionContext(session3), - ], - runner=runner, - processor=processor, - ) - - delay_mock = MagicMock() - monkeypatch.setattr( - "tasks.document_indexing_update_task.generate_summary_index_task.delay", - delay_mock, - ) - - document_indexing_update_task("ds-1", "doc-1") - - delay_mock.assert_called_once_with("ds-1", "doc-1", None) + document_indexing_update_task(dataset.id, document.id) + + sqlite_session.expire_all() + assert sqlite_session.get(Document, document.id).indexing_status == IndexingStatus.COMPLETED # type: ignore[union-attr] + + +def test_cleans_and_deletes_persisted_segments_with_real_session( + sqlite_session: Session, + task_harness: tuple[MagicMock, MagicMock], + monkeypatch: pytest.MonkeyPatch, +) -> None: + runner, processor = task_harness + dataset, document = _persist_rows( + sqlite_session, + summary_index_setting={"enable": True}, + with_segment=True, + ) + runner.run.side_effect = _complete_indexing + delay = MagicMock() + monkeypatch.setattr(task_module.generate_summary_index_task, "delay", delay) + + document_indexing_update_task(dataset.id, document.id) + + processor.clean.assert_called_once() + assert isinstance(processor.clean.call_args.kwargs["session"], Session) + assert sqlite_session.scalars(select(DocumentSegment).where(DocumentSegment.document_id == document.id)).all() == [] + delay.assert_called_once_with(dataset.id, document.id, None) diff --git a/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py b/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py index af7ecb8c570..fb64f5b5ba4 100644 --- a/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py +++ b/api/tests/unit_tests/tasks/test_enable_segment_index_tasks.py @@ -1,43 +1,97 @@ -from contextlib import nullcontext -from types import SimpleNamespace +import uuid +from collections.abc import Iterator +from contextlib import contextmanager from unittest.mock import MagicMock, patch +import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker + from core.rag.index_processor.constant.index_type import IndexStructureType -from models.enums import IndexingStatus, SegmentStatus +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus, SegmentStatus from tasks.enable_segment_to_index_task import enable_segment_to_index_task from tasks.enable_segments_to_index_task import enable_segments_to_index_task -def test_enable_segment_commits_index_rows_after_loading() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, +@pytest.fixture +def indexed_segment(sqlite_session: Session) -> tuple[Dataset, Document, DocumentSegment]: + """Persist the complete owner chain consumed by segment indexing tasks.""" + tenant_id = str(uuid.uuid4()) + created_by = str(uuid.uuid4()) + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Indexing dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + is_multimodal=False, + ) + document = Document( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + dataset_id=dataset.id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="document.txt", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, indexing_status=IndexingStatus.COMPLETED, doc_form=IndexStructureType.PARAGRAPH_INDEX, ) - segment = SimpleNamespace( - id="segment-1", - status=SegmentStatus.COMPLETED, + segment = DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset.id, + document_id=document.id, + position=1, content="content", + word_count=1, + tokens=1, + created_by=created_by, index_node_id="node-1", index_node_hash="hash-1", - document_id=document.id, - dataset_id=dataset.id, - get_dataset=MagicMock(return_value=dataset), - get_document=MagicMock(return_value=document), + status=SegmentStatus.COMPLETED, ) - session = MagicMock() - session.scalar.return_value = segment + sqlite_session.add_all([dataset, document, segment]) + sqlite_session.commit() + return dataset, document, segment + + +@contextmanager +def _record_transaction_events( + sqlite_session_factory: sessionmaker[Session], phase_events: list[str] +) -> Iterator[None]: + """Record real transaction boundaries from task-owned SQLite sessions.""" + session_type = sqlite_session_factory.class_ + + def after_commit(_session: Session) -> None: + phase_events.append("commit") + + def after_rollback(_session: Session) -> None: + phase_events.append("rollback") + + event.listen(session_type, "after_commit", after_commit) + event.listen(session_type, "after_rollback", after_rollback) + try: + yield + finally: + event.remove(session_type, "after_commit", after_commit) + event.remove(session_type, "after_rollback", after_rollback) + + +def test_enable_segment_commits_index_rows_after_loading( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, _document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") index_processor = MagicMock() index_processor.load.side_effect = lambda *_args, **_kwargs: phase_events.append("load") enable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) with ( - patch("tasks.enable_segment_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.enable_segment_to_index_task.IndexProcessorFactory") as processor_factory, patch( "services.summary_index_service.SummaryIndexService.enable_summaries_for_segments", @@ -49,53 +103,28 @@ def test_enable_segment_commits_index_rows_after_loading() -> None: enable_segment_to_index_task.run(segment.id) assert phase_events == ["load", "commit", "summary"] + enable_summaries.assert_called_once() + assert enable_summaries.call_args.kwargs["dataset"].id == dataset.id + assert enable_summaries.call_args.kwargs["segment_ids"] == [segment.id] -def test_enable_segment_rolls_back_before_error_compensation() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status=IndexingStatus.COMPLETED, - doc_form=IndexStructureType.PARAGRAPH_INDEX, - ) - segment = SimpleNamespace( - id="segment-1", - status=SegmentStatus.COMPLETED, - content="content", - index_node_id="node-1", - index_node_hash="hash-1", - document_id=document.id, - dataset_id=dataset.id, - enabled=True, - disabled_at=None, - error=None, - get_dataset=MagicMock(return_value=dataset), - get_document=MagicMock(return_value=document), - ) +def test_enable_segment_rolls_back_before_error_compensation( + sqlite_session: Session, + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + _dataset, _document, segment = indexed_segment phase_events: list[str] = [] - session = MagicMock() - session.scalar.return_value = segment - session.rollback.side_effect = lambda: phase_events.append("rollback") - - def commit() -> None: - assert segment.enabled is False - assert segment.status == SegmentStatus.ERROR - assert segment.error == "load failed" - phase_events.append("commit") - - session.commit.side_effect = commit index_processor = MagicMock() - def fail_load(*_args, **_kwargs) -> None: + def fail_load(*_args: object, **_kwargs: object) -> None: phase_events.append("load") raise RuntimeError("load failed") index_processor.load.side_effect = fail_load with ( - patch("tasks.enable_segment_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.enable_segment_to_index_task.IndexProcessorFactory") as processor_factory, patch("services.summary_index_service.SummaryIndexService.enable_summaries_for_segments") as enable_summaries, patch("tasks.enable_segment_to_index_task.redis_client.delete"), @@ -103,38 +132,29 @@ def test_enable_segment_rolls_back_before_error_compensation() -> None: processor_factory.return_value.init_index_processor.return_value = index_processor enable_segment_to_index_task.run(segment.id) + sqlite_session.expire_all() + persisted_segment = sqlite_session.get(DocumentSegment, segment.id) + assert persisted_segment is not None + assert persisted_segment.enabled is False + assert persisted_segment.status == SegmentStatus.ERROR + assert persisted_segment.error == "load failed" + assert persisted_segment.disabled_at is not None assert phase_events == ["load", "rollback", "commit"] enable_summaries.assert_not_called() -def test_enable_segments_commits_index_rows_after_loading() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status="completed", - doc_form=IndexStructureType.PARAGRAPH_INDEX, - ) - segment = SimpleNamespace( - id="segment-1", - content="content", - index_node_id="node-1", - index_node_hash="hash-1", - document_id=document.id, - dataset_id=dataset.id, - ) - session = MagicMock() - session.scalar.side_effect = [dataset, document] - session.scalars.return_value.all.return_value = [segment] +def test_enable_segments_commits_index_rows_after_loading( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") index_processor = MagicMock() index_processor.load.side_effect = lambda *_args, **_kwargs: phase_events.append("load") enable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) with ( - patch("tasks.enable_segments_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.enable_segments_to_index_task.IndexProcessorFactory") as processor_factory, patch( "services.summary_index_service.SummaryIndexService.enable_summaries_for_segments", @@ -146,42 +166,28 @@ def test_enable_segments_commits_index_rows_after_loading() -> None: enable_segments_to_index_task.run([segment.id], dataset.id, document.id) assert phase_events == ["load", "commit", "summary"] + enable_summaries.assert_called_once() + assert enable_summaries.call_args.kwargs["dataset"].id == dataset.id + assert enable_summaries.call_args.kwargs["segment_ids"] == [segment.id] -def test_enable_segments_rolls_back_before_error_compensation() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status="completed", - doc_form=IndexStructureType.PARAGRAPH_INDEX, - ) - segment = SimpleNamespace( - id="segment-1", - content="content", - index_node_id="node-1", - index_node_hash="hash-1", - document_id=document.id, - dataset_id=dataset.id, - ) +def test_enable_segments_rolls_back_before_error_compensation( + sqlite_session: Session, + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, document, segment = indexed_segment phase_events: list[str] = [] - session = MagicMock() - session.scalar.side_effect = [dataset, document] - session.scalars.return_value.all.return_value = [segment] - session.rollback.side_effect = lambda: phase_events.append("rollback") - session.execute.side_effect = lambda *_args, **_kwargs: phase_events.append("compensate") - session.commit.side_effect = lambda: phase_events.append("commit") index_processor = MagicMock() - def fail_load(*_args, **_kwargs) -> None: + def fail_load(*_args: object, **_kwargs: object) -> None: phase_events.append("load") raise RuntimeError("load failed") index_processor.load.side_effect = fail_load with ( - patch("tasks.enable_segments_to_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.enable_segments_to_index_task.IndexProcessorFactory") as processor_factory, patch("services.summary_index_service.SummaryIndexService.enable_summaries_for_segments") as enable_summaries, patch("tasks.enable_segments_to_index_task.redis_client.delete"), @@ -189,5 +195,12 @@ def test_enable_segments_rolls_back_before_error_compensation() -> None: processor_factory.return_value.init_index_processor.return_value = index_processor enable_segments_to_index_task.run([segment.id], dataset.id, document.id) - assert phase_events == ["load", "rollback", "compensate", "commit"] + sqlite_session.expire_all() + persisted_segment = sqlite_session.get(DocumentSegment, segment.id) + assert persisted_segment is not None + assert persisted_segment.enabled is False + assert persisted_segment.status == SegmentStatus.ERROR + assert persisted_segment.error == "load failed" + assert persisted_segment.disabled_at is not None + assert phase_events == ["load", "rollback", "commit"] enable_summaries.assert_not_called() diff --git a/api/tests/unit_tests/tasks/test_human_input_timeout_tasks.py b/api/tests/unit_tests/tasks/test_human_input_timeout_tasks.py index be837acd2f0..b8cca3a1171 100644 --- a/api/tests/unit_tests/tasks/test_human_input_timeout_tasks.py +++ b/api/tests/unit_tests/tasks/test_human_input_timeout_tasks.py @@ -2,74 +2,22 @@ from __future__ import annotations from datetime import datetime, timedelta from types import SimpleNamespace -from typing import Any +from unittest.mock import MagicMock import pytest +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session, sessionmaker +import core.db.session_factory as session_factory_module +from core.repositories.human_input_repository import HumanInputFormSubmissionRepository +from core.workflow.nodes.human_input.entities import FormDefinition from core.workflow.nodes.human_input.enums import HumanInputFormKind, HumanInputFormStatus +from models.human_input import HumanInputForm from tasks import human_input_timeout_tasks as task_module -class _FakeScalarResult: - def __init__(self, items: list[Any]): - self._items = items - - def all(self) -> list[Any]: - return self._items - - -class _FakeSession: - def __init__(self, items: list[Any], capture: dict[str, Any]): - self._items = items - self._capture = capture - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb): - return False - - def scalars(self, stmt): - self._capture["stmt"] = stmt - return _FakeScalarResult(self._items) - - -class _FakeSessionFactory: - def __init__(self, items: list[Any], capture: dict[str, Any]): - self._items = items - self._capture = capture - self._capture["session_factory"] = self - - def __call__(self): - session = _FakeSession(self._items, self._capture) - self._capture["session"] = session - return session - - -class _FakeFormRepo: - def __init__(self, form_map: dict[str, Any] | None = None): - self.calls: list[dict[str, Any]] = [] - self._form_map = form_map or {} - - def mark_timeout(self, *, form_id: str, timeout_status: HumanInputFormStatus, reason: str | None = None): - self.calls.append( - { - "form_id": form_id, - "timeout_status": timeout_status, - "reason": reason, - } - ) - form = self._form_map.get(form_id) - return SimpleNamespace( - form_id=form_id, - workflow_run_id=getattr(form, "workflow_run_id", None), - conversation_id=getattr(form, "conversation_id", None), - node_id=getattr(form, "node_id", None), - ) - - class _FakeService: - def __init__(self, _session_factory, form_repository=None): + def __init__(self): self.enqueued: list[str] = [] self.agent_app_resumed: list[tuple[str, str]] = [] @@ -90,22 +38,49 @@ def _build_form( workflow_run_id: str | None, node_id: str, conversation_id: str | None = None, -) -> SimpleNamespace: - return SimpleNamespace( +) -> HumanInputForm: + form_definition = FormDefinition( + form_content="", + rendered_content="", + expiration_time=expiration_time, + ) + return HumanInputForm( id=form_id, + tenant_id="tenant-1", + app_id="app-1", form_kind=form_kind, created_at=created_at, expiration_time=expiration_time, workflow_run_id=workflow_run_id, conversation_id=conversation_id, node_id=node_id, + form_definition=form_definition.model_dump_json(), + rendered_content="", status=HumanInputFormStatus.WAITING, ) +@pytest.fixture +def sqlite_task_database( + sqlite_engine: Engine, + sqlite_session: Session, + monkeypatch: pytest.MonkeyPatch, +) -> None: + repository_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) + monkeypatch.setattr(session_factory_module, "_session_maker", repository_session_maker) + monkeypatch.setattr(task_module, "db", SimpleNamespace(engine=sqlite_engine)) + + def test_is_global_timeout_uses_created_at(): now = datetime(2025, 1, 1, 12, 0, 0) - form = SimpleNamespace(created_at=now - timedelta(seconds=61), workflow_run_id="run-1") + form = _build_form( + form_id="form-1", + form_kind=HumanInputFormKind.RUNTIME, + created_at=now - timedelta(seconds=61), + expiration_time=now + timedelta(hours=1), + workflow_run_id="run-1", + node_id="node-1", + ) assert task_module._is_global_timeout(form, 60, now=now) is True @@ -119,11 +94,16 @@ def test_is_global_timeout_uses_created_at(): assert task_module._is_global_timeout(form, 0, now=now) is False -def test_check_and_handle_human_input_timeouts_marks_and_routes(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) +def test_check_and_handle_human_input_timeouts_marks_and_routes( + monkeypatch: pytest.MonkeyPatch, + sqlite_task_database: None, + sqlite_engine: Engine, + sqlite_session: Session, +): now = datetime(2025, 1, 1, 12, 0, 0) monkeypatch.setattr(task_module, "naive_utc_now", lambda: now) monkeypatch.setattr(task_module.dify_config, "HUMAN_INPUT_GLOBAL_TIMEOUT_SECONDS", 3600) - monkeypatch.setattr(task_module, "db", SimpleNamespace(engine=object())) forms = [ _build_form( @@ -151,74 +131,131 @@ def test_check_and_handle_human_input_timeouts_marks_and_routes(monkeypatch: pyt node_id="node-delivery", ), ] + sqlite_session.add_all(forms) + sqlite_session.commit() - capture: dict[str, Any] = {} - monkeypatch.setattr(task_module, "sessionmaker", lambda *args, **kwargs: _FakeSessionFactory(forms, capture)) + repo = HumanInputFormSubmissionRepository() + mark_timeout_spy = MagicMock(wraps=repo.mark_timeout) + monkeypatch.setattr(repo, "mark_timeout", mark_timeout_spy) + service = _FakeService() + service_factory = MagicMock(return_value=service) + global_timeout_handler = MagicMock() - form_map = {form.id: form for form in forms} - repo = _FakeFormRepo(form_map=form_map) - - def _repo_factory(): - return repo - - service = _FakeService(None) - - def _service_factory(_session_factory, form_repository=None): - return service - - global_calls: list[dict[str, Any]] = [] - - monkeypatch.setattr(task_module, "HumanInputFormSubmissionRepository", _repo_factory) - monkeypatch.setattr(task_module, "HumanInputService", _service_factory) - monkeypatch.setattr(task_module, "_handle_global_timeout", lambda **kwargs: global_calls.append(kwargs)) + monkeypatch.setattr(task_module, "HumanInputFormSubmissionRepository", lambda: repo) + monkeypatch.setattr(task_module, "HumanInputService", service_factory) + monkeypatch.setattr(task_module, "_handle_global_timeout", global_timeout_handler) task_module.check_and_handle_human_input_timeouts(limit=100) - assert {(call["form_id"], call["timeout_status"], call["reason"]) for call in repo.calls} == { + assert { + (call.kwargs["form_id"], call.kwargs["timeout_status"], call.kwargs["reason"]) + for call in mark_timeout_spy.call_args_list + } == { ("form-global", HumanInputFormStatus.EXPIRED, "global_timeout"), ("form-node", HumanInputFormStatus.TIMEOUT, "node_timeout"), ("form-delivery", HumanInputFormStatus.TIMEOUT, "delivery_test_timeout"), } assert service.enqueued == ["run-node"] - assert global_calls == [ - { - "form_id": "form-global", - "workflow_run_id": "run-global", - "node_id": "node-global", - "session_factory": capture.get("session_factory"), - } - ] + global_timeout_handler.assert_called_once() + global_timeout_call = global_timeout_handler.call_args.kwargs + assert global_timeout_call["form_id"] == "form-global" + assert global_timeout_call["workflow_run_id"] == "run-global" + assert global_timeout_call["node_id"] == "node-global" + task_session_maker = global_timeout_call["session_factory"] + assert isinstance(task_session_maker, sessionmaker) + assert task_session_maker.kw["bind"] is sqlite_engine + service_factory.assert_called_once_with(task_session_maker, form_repository=repo) - stmt = capture.get("stmt") - assert stmt is not None - stmt_text = str(stmt) - assert "created_at <=" in stmt_text - assert "expiration_time <=" in stmt_text - assert "ORDER BY human_input_forms.id" in stmt_text + sqlite_session.expire_all() + assert sqlite_session.get(HumanInputForm, "form-global").status == HumanInputFormStatus.EXPIRED + assert sqlite_session.get(HumanInputForm, "form-node").status == HumanInputFormStatus.TIMEOUT + assert sqlite_session.get(HumanInputForm, "form-delivery").status == HumanInputFormStatus.TIMEOUT -def test_check_and_handle_human_input_timeouts_omits_global_filter_when_disabled(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) +def test_check_and_handle_human_input_timeouts_orders_by_id_before_limit( + monkeypatch: pytest.MonkeyPatch, + sqlite_task_database: None, + sqlite_session: Session, +): now = datetime(2025, 1, 1, 12, 0, 0) monkeypatch.setattr(task_module, "naive_utc_now", lambda: now) monkeypatch.setattr(task_module.dify_config, "HUMAN_INPUT_GLOBAL_TIMEOUT_SECONDS", 0) - monkeypatch.setattr(task_module, "db", SimpleNamespace(engine=object())) - capture: dict[str, Any] = {} - monkeypatch.setattr(task_module, "sessionmaker", lambda *args, **kwargs: _FakeSessionFactory([], capture)) - monkeypatch.setattr(task_module, "HumanInputFormSubmissionRepository", _FakeFormRepo) - monkeypatch.setattr(task_module, "HumanInputService", _FakeService) - monkeypatch.setattr(task_module, "_handle_global_timeout", lambda **_kwargs: None) + forms = [ + _build_form( + form_id=form_id, + form_kind=HumanInputFormKind.DELIVERY_TEST, + created_at=now - timedelta(minutes=1), + expiration_time=now - timedelta(seconds=1), + workflow_run_id=None, + node_id=f"node-{form_id}", + ) + for form_id in ("form-b", "form-a") + ] + sqlite_session.add_all(forms) + sqlite_session.commit() + + repo = HumanInputFormSubmissionRepository() + mark_timeout_spy = MagicMock(wraps=repo.mark_timeout) + monkeypatch.setattr(repo, "mark_timeout", mark_timeout_spy) + monkeypatch.setattr(task_module, "HumanInputFormSubmissionRepository", lambda: repo) + monkeypatch.setattr(task_module, "HumanInputService", MagicMock(return_value=_FakeService())) task_module.check_and_handle_human_input_timeouts(limit=1) - stmt = capture.get("stmt") - assert stmt is not None - stmt_text = str(stmt) - assert "created_at <=" not in stmt_text + mark_timeout_spy.assert_called_once_with( + form_id="form-a", + timeout_status=HumanInputFormStatus.TIMEOUT, + reason="delivery_test_timeout", + ) + sqlite_session.expire_all() + assert sqlite_session.get(HumanInputForm, "form-a").status == HumanInputFormStatus.TIMEOUT + assert sqlite_session.get(HumanInputForm, "form-b").status == HumanInputFormStatus.WAITING +@pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) +def test_check_and_handle_human_input_timeouts_omits_global_filter_when_disabled( + monkeypatch: pytest.MonkeyPatch, + sqlite_task_database: None, + sqlite_session: Session, +): + now = datetime(2025, 1, 1, 12, 0, 0) + monkeypatch.setattr(task_module, "naive_utc_now", lambda: now) + monkeypatch.setattr(task_module.dify_config, "HUMAN_INPUT_GLOBAL_TIMEOUT_SECONDS", 0) + + old_unexpired_form = _build_form( + form_id="form-old", + form_kind=HumanInputFormKind.RUNTIME, + created_at=now - timedelta(hours=2), + expiration_time=now + timedelta(hours=1), + workflow_run_id="run-old", + node_id="node-old", + ) + sqlite_session.add(old_unexpired_form) + sqlite_session.commit() + + repo = HumanInputFormSubmissionRepository() + mark_timeout_spy = MagicMock(wraps=repo.mark_timeout) + monkeypatch.setattr(repo, "mark_timeout", mark_timeout_spy) + monkeypatch.setattr(task_module, "HumanInputFormSubmissionRepository", lambda: repo) + monkeypatch.setattr(task_module, "HumanInputService", MagicMock(return_value=_FakeService())) + global_timeout_handler = MagicMock() + monkeypatch.setattr(task_module, "_handle_global_timeout", global_timeout_handler) + + task_module.check_and_handle_human_input_timeouts(limit=1) + + mark_timeout_spy.assert_not_called() + global_timeout_handler.assert_not_called() + sqlite_session.refresh(old_unexpired_form) + assert old_unexpired_form.status == HumanInputFormStatus.WAITING + + +@pytest.mark.parametrize("sqlite_session", [(HumanInputForm,)], indirect=True) def test_check_and_handle_human_input_timeouts_routes_conversation_owned_form_to_agent_app_resume( monkeypatch: pytest.MonkeyPatch, + sqlite_task_database: None, + sqlite_session: Session, ): # ENG-635 (review): a conversation-owned Agent v2 chat ask_human form has no # workflow_run_id. On timeout it must enqueue the Agent App resume (so the @@ -227,24 +264,23 @@ def test_check_and_handle_human_input_timeouts_routes_conversation_owned_form_to now = datetime(2025, 1, 1, 12, 0, 0) monkeypatch.setattr(task_module, "naive_utc_now", lambda: now) monkeypatch.setattr(task_module.dify_config, "HUMAN_INPUT_GLOBAL_TIMEOUT_SECONDS", 3600) - monkeypatch.setattr(task_module, "db", SimpleNamespace(engine=object())) - forms = [ - _build_form( - form_id="form-chat", - form_kind=HumanInputFormKind.RUNTIME, - created_at=now - timedelta(minutes=5), - expiration_time=now - timedelta(seconds=1), - workflow_run_id=None, - conversation_id="conv-1", - node_id="agent", - ), - ] - capture: dict[str, Any] = {} - monkeypatch.setattr(task_module, "sessionmaker", lambda *args, **kwargs: _FakeSessionFactory(forms, capture)) + form = _build_form( + form_id="form-chat", + form_kind=HumanInputFormKind.RUNTIME, + created_at=now - timedelta(minutes=5), + expiration_time=now - timedelta(seconds=1), + workflow_run_id=None, + conversation_id="conv-1", + node_id="agent", + ) + sqlite_session.add(form) + sqlite_session.commit() - repo = _FakeFormRepo(form_map={form.id: form for form in forms}) - service = _FakeService(None) + repo = HumanInputFormSubmissionRepository() + mark_timeout_spy = MagicMock(wraps=repo.mark_timeout) + monkeypatch.setattr(repo, "mark_timeout", mark_timeout_spy) + service = _FakeService() monkeypatch.setattr(task_module, "HumanInputFormSubmissionRepository", lambda: repo) monkeypatch.setattr(task_module, "HumanInputService", lambda *_args, **_kwargs: service) monkeypatch.setattr(task_module, "_handle_global_timeout", lambda **_kwargs: None) @@ -252,8 +288,10 @@ def test_check_and_handle_human_input_timeouts_routes_conversation_owned_form_to task_module.check_and_handle_human_input_timeouts(limit=100) # Node timeout (conversation forms are never "global"), routed to Agent App resume. - assert repo.calls == [ - {"form_id": "form-chat", "timeout_status": HumanInputFormStatus.TIMEOUT, "reason": "node_timeout"} - ] + mark_timeout_spy.assert_called_once_with( + form_id="form-chat", timeout_status=HumanInputFormStatus.TIMEOUT, reason="node_timeout" + ) assert service.agent_app_resumed == [("conv-1", "form-chat")] assert service.enqueued == [] + sqlite_session.refresh(form) + assert form.status == HumanInputFormStatus.TIMEOUT diff --git a/api/tests/unit_tests/tasks/test_initialize_created_app_rbac_access_task.py b/api/tests/unit_tests/tasks/test_initialize_created_app_rbac_access_task.py index 2bc5b8a69aa..bc04cd5a800 100644 --- a/api/tests/unit_tests/tasks/test_initialize_created_app_rbac_access_task.py +++ b/api/tests/unit_tests/tasks/test_initialize_created_app_rbac_access_task.py @@ -11,7 +11,7 @@ def test_initialize_created_app_rbac_access_task_uses_rbac_queue(): assert initialize_created_app_rbac_access_task.queue == APP_RBAC_QUEUE -def test_initialize_created_app_rbac_access_task_batches_workspace_members(monkeypatch): +def test_initialize_created_app_rbac_access_task_batches_workspace_members(monkeypatch: pytest.MonkeyPatch): import tasks.initialize_created_app_rbac_access_task as task_module from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task @@ -44,7 +44,7 @@ def test_initialize_created_app_rbac_access_task_batches_workspace_members(monke assert call.kwargs["payload"].access_policy_ids == [task_module.APP_RBAC_DEFAULT_ACCESS_POLICY_ID] -def test_initialize_created_app_rbac_access_task_retries_on_failure(monkeypatch): +def test_initialize_created_app_rbac_access_task_retries_on_failure(monkeypatch: pytest.MonkeyPatch): import tasks.initialize_created_app_rbac_access_task as task_module from tasks.initialize_created_app_rbac_access_task import initialize_created_app_rbac_access_task diff --git a/api/tests/unit_tests/tasks/test_mail_send_task.py b/api/tests/unit_tests/tasks/test_mail_send_task.py index a7af192d3a7..6862daaeb90 100644 --- a/api/tests/unit_tests/tasks/test_mail_send_task.py +++ b/api/tests/unit_tests/tasks/test_mail_send_task.py @@ -13,10 +13,15 @@ import smtplib from unittest.mock import ANY, MagicMock, patch import pytest +from python_http_client.exceptions import ForbiddenError, UnauthorizedError from configs import dify_config from configs.feature import TemplateMode -from libs.email_i18n import EmailType +from extensions.ext_mail import Mail +from libs.email_i18n import EmailI18nConfig, EmailI18nService, EmailLanguage, EmailTemplate, EmailType +from libs.sendgrid import SendGridClient +from libs.smtp import SMTPClient +from services.entities.feature_entities import BrandingModel from tasks.mail_inner_task import _render_template_with_strategy, send_inner_email_task from tasks.mail_register_task import ( send_email_register_mail_task, @@ -131,8 +136,6 @@ class TestSMTPIntegration: def test_smtp_send_with_tls_ssl(self, mock_smtp_ssl): """Test SMTP send with TLS using SMTP_SSL.""" # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp_ssl.return_value = mock_server @@ -161,8 +164,6 @@ class TestSMTPIntegration: def test_smtp_send_with_opportunistic_tls(self, mock_smtp): """Test SMTP send with opportunistic TLS (STARTTLS).""" # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp.return_value = mock_server @@ -193,8 +194,6 @@ class TestSMTPIntegration: def test_smtp_send_without_tls(self, mock_smtp): """Test SMTP send without TLS encryption.""" # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp.return_value = mock_server @@ -223,8 +222,6 @@ class TestSMTPIntegration: def test_smtp_send_without_authentication(self, mock_smtp): """Test SMTP send without authentication (empty credentials).""" # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp.return_value = mock_server @@ -252,8 +249,6 @@ class TestSMTPIntegration: def test_smtp_send_authentication_failure(self, mock_smtp_ssl): """Test SMTP send handles authentication failure.""" # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp_ssl.return_value = mock_server mock_server.login.side_effect = smtplib.SMTPAuthenticationError(535, b"Authentication failed") @@ -280,8 +275,6 @@ class TestSMTPIntegration: def test_smtp_send_timeout_error(self, mock_smtp_ssl): """Test SMTP send handles timeout errors.""" # Arrange - from libs.smtp import SMTPClient - mock_smtp_ssl.side_effect = TimeoutError("Connection timeout") client = SMTPClient( @@ -304,8 +297,6 @@ class TestSMTPIntegration: def test_smtp_send_connection_refused(self, mock_smtp_ssl): """Test SMTP send handles connection refused errors.""" # Arrange - from libs.smtp import SMTPClient - mock_smtp_ssl.side_effect = ConnectionRefusedError("Connection refused") client = SMTPClient( @@ -328,8 +319,6 @@ class TestSMTPIntegration: def test_smtp_send_ensures_cleanup_on_error(self, mock_smtp_ssl): """Test SMTP send ensures cleanup even when errors occur.""" # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp_ssl.return_value = mock_server mock_server.sendmail.side_effect = smtplib.SMTPException("Send failed") @@ -610,8 +599,6 @@ class TestSendGridIntegration: def test_sendgrid_send_success(self, mock_sg_client): """Test SendGrid client sends email successfully.""" # Arrange - from libs.sendgrid import SendGridClient - mock_client_instance = MagicMock() mock_sg_client.return_value = mock_client_instance mock_response = MagicMock() @@ -633,8 +620,6 @@ class TestSendGridIntegration: def test_sendgrid_send_missing_recipient(self, mock_sg_client): """Test SendGrid client raises error when recipient is missing.""" # Arrange - from libs.sendgrid import SendGridClient - client = SendGridClient(sendgrid_api_key="test_api_key", _from="noreply@example.com") mail_data = {"to": "", "subject": "Test Subject", "html": "

Test Content

"} @@ -647,10 +632,6 @@ class TestSendGridIntegration: def test_sendgrid_send_unauthorized_error(self, mock_sg_client): """Test SendGrid client handles unauthorized errors.""" # Arrange - from python_http_client.exceptions import UnauthorizedError - - from libs.sendgrid import SendGridClient - mock_client_instance = MagicMock() mock_sg_client.return_value = mock_client_instance mock_client_instance.client.mail.send.post.side_effect = UnauthorizedError( @@ -669,10 +650,6 @@ class TestSendGridIntegration: def test_sendgrid_send_forbidden_error(self, mock_sg_client): """Test SendGrid client handles forbidden errors.""" # Arrange - from python_http_client.exceptions import ForbiddenError - - from libs.sendgrid import SendGridClient - mock_client_instance = MagicMock() mock_sg_client.return_value = mock_client_instance mock_client_instance.client.mail.send.post.side_effect = ForbiddenError(MagicMock(status_code=403), "Forbidden") @@ -689,8 +666,6 @@ class TestSendGridIntegration: def test_sendgrid_send_timeout_error(self, mock_sg_client): """Test SendGrid client handles timeout errors.""" # Arrange - from libs.sendgrid import SendGridClient - mock_client_instance = MagicMock() mock_sg_client.return_value = mock_client_instance mock_client_instance.client.mail.send.post.side_effect = TimeoutError("Request timeout") @@ -711,8 +686,6 @@ class TestMailExtension: def test_mail_init_smtp_configuration(self, mock_config): """Test mail extension initializes SMTP client correctly.""" # Arrange - from extensions.ext_mail import Mail - mock_config.MAIL_TYPE = "smtp" mock_config.SMTP_SERVER = "smtp.example.com" mock_config.SMTP_PORT = 465 @@ -736,8 +709,6 @@ class TestMailExtension: def test_mail_init_without_mail_type(self, mock_config): """Test mail extension skips initialization when MAIL_TYPE is not set.""" # Arrange - from extensions.ext_mail import Mail - mock_config.MAIL_TYPE = None mail = Mail() @@ -753,8 +724,6 @@ class TestMailExtension: def test_mail_send_validates_parameters(self, mock_config): """Test mail send validates required parameters.""" # Arrange - from extensions.ext_mail import Mail - mail = Mail() mail._client = MagicMock() mail._default_send_from = "noreply@example.com" @@ -775,8 +744,6 @@ class TestMailExtension: def test_mail_send_uses_default_from(self, mock_config): """Test mail send uses default from address when not provided.""" # Arrange - from extensions.ext_mail import Mail - mail = Mail() mock_client = MagicMock() mail._client = mock_client @@ -800,9 +767,6 @@ class TestEmailI18nService: def test_email_service_sends_with_branding(self, mock_renderer_class, mock_branding_class, mock_sender_class): """Test email service sends email with branding support.""" # Arrange - from libs.email_i18n import EmailI18nConfig, EmailI18nService, EmailLanguage, EmailTemplate, EmailType - from services.feature_service import BrandingModel - mock_renderer = MagicMock() mock_renderer.render_template.return_value = "Rendered content" mock_renderer_class.return_value = mock_renderer @@ -848,8 +812,6 @@ class TestEmailI18nService: def test_email_service_send_raw_email_single_recipient(self, mock_sender_class): """Test email service sends raw email to single recipient.""" # Arrange - from libs.email_i18n import EmailI18nConfig, EmailI18nService - mock_sender = MagicMock() mock_sender_class.return_value = mock_sender @@ -872,8 +834,6 @@ class TestEmailI18nService: def test_email_service_send_raw_email_multiple_recipients(self, mock_sender_class): """Test email service sends raw email to multiple recipients.""" # Arrange - from libs.email_i18n import EmailI18nConfig, EmailI18nService - mock_sender = MagicMock() mock_sender_class.return_value = mock_sender @@ -941,8 +901,6 @@ class TestEdgeCasesAndErrorHandling: configuration parameters are not provided. """ # Arrange - from extensions.ext_mail import Mail - mock_config.MAIL_TYPE = "smtp" mock_config.SMTP_SERVER = None # Missing required parameter mock_config.SMTP_PORT = 465 @@ -963,8 +921,6 @@ class TestEdgeCasesAndErrorHandling: This test ensures the configuration is validated properly. """ # Arrange - from extensions.ext_mail import Mail - mock_config.MAIL_TYPE = "smtp" mock_config.SMTP_SERVER = "smtp.example.com" mock_config.SMTP_PORT = 587 @@ -987,8 +943,6 @@ class TestEdgeCasesAndErrorHandling: are accepted and invalid types are rejected. """ # Arrange - from extensions.ext_mail import Mail - mock_config.MAIL_TYPE = "unsupported_provider" mail = Mail() @@ -1007,8 +961,6 @@ class TestEdgeCasesAndErrorHandling: emails with empty subjects without crashing. """ # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp_ssl.return_value = mock_server @@ -1040,8 +992,6 @@ class TestEdgeCasesAndErrorHandling: subject lines and email bodies. """ # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp_ssl.return_value = mock_server @@ -1141,8 +1091,6 @@ class TestResendIntegration: and the client is initialized. """ # Arrange - from extensions.ext_mail import Mail - mock_config.MAIL_TYPE = "resend" mock_config.RESEND_API_KEY = "re_test_api_key" mock_config.RESEND_API_URL = None @@ -1183,8 +1131,6 @@ class TestResendIntegration: This test ensures custom URLs are properly configured. """ # Arrange - from extensions.ext_mail import Mail - mock_config.MAIL_TYPE = "resend" mock_config.RESEND_API_KEY = "re_test_api_key" mock_config.RESEND_API_URL = "https://custom-resend.example.com" @@ -1224,8 +1170,6 @@ class TestResendIntegration: proper validation of required configuration. """ # Arrange - from extensions.ext_mail import Mail - mock_config.MAIL_TYPE = "resend" mock_config.RESEND_API_KEY = None # Missing API key @@ -1333,8 +1277,6 @@ class TestEmailValidation: this test documents the current behavior. """ # Arrange - from extensions.ext_mail import Mail - mail = Mail() mock_client = MagicMock() mail._client = mock_client @@ -1364,8 +1306,6 @@ class TestSMTPEdgeCases: or extensive formatting. This test ensures they're handled. """ # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp_ssl.return_value = mock_server @@ -1401,8 +1341,6 @@ class TestSMTPEdgeCases: recipient per call. This test documents that behavior. """ # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp_ssl.return_value = mock_server @@ -1434,8 +1372,6 @@ class TestSMTPEdgeCases: whitespace to avoid authentication with blank credentials. """ # Arrange - from libs.smtp import SMTPClient - mock_server = MagicMock() mock_smtp.return_value = mock_server diff --git a/api/tests/unit_tests/tasks/test_refresh_billing_vector_space_task.py b/api/tests/unit_tests/tasks/test_refresh_billing_vector_space_task.py index 8f4820be660..63bf6fcedfd 100644 --- a/api/tests/unit_tests/tasks/test_refresh_billing_vector_space_task.py +++ b/api/tests/unit_tests/tasks/test_refresh_billing_vector_space_task.py @@ -2,6 +2,7 @@ from unittest.mock import patch import pytest +from enums import DeploymentEdition from tasks.refresh_billing_vector_space_task import ( refresh_billing_vector_space_task, schedule_billing_vector_space_refresh, @@ -10,7 +11,7 @@ from tasks.refresh_billing_vector_space_task import ( def test_refresh_invalidates_vector_space_cache(): with ( - patch("tasks.refresh_billing_vector_space_task.dify_config.BILLING_ENABLED", True), + patch("tasks.refresh_billing_vector_space_task.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch( "tasks.refresh_billing_vector_space_task.BillingService.invalidate_vector_space_cache" ) as invalidate_cache, @@ -24,7 +25,7 @@ def test_refresh_failure_schedules_retry(): error = RuntimeError("billing unavailable") with ( - patch("tasks.refresh_billing_vector_space_task.dify_config.BILLING_ENABLED", True), + patch("tasks.refresh_billing_vector_space_task.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch( "tasks.refresh_billing_vector_space_task.BillingService.invalidate_vector_space_cache", side_effect=error, @@ -39,7 +40,7 @@ def test_refresh_failure_schedules_retry(): def test_dispatch_failure_does_not_propagate(): with ( - patch("tasks.refresh_billing_vector_space_task.dify_config.BILLING_ENABLED", True), + patch("tasks.refresh_billing_vector_space_task.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.CLOUD), patch.object(refresh_billing_vector_space_task, "delay", side_effect=RuntimeError("broker unavailable")), ): schedule_billing_vector_space_refresh("tenant-1") diff --git a/api/tests/unit_tests/tasks/test_remove_app_and_related_data_task.py b/api/tests/unit_tests/tasks/test_remove_app_and_related_data_task.py index 7d2705447db..3f0ce2475ef 100644 --- a/api/tests/unit_tests/tasks/test_remove_app_and_related_data_task.py +++ b/api/tests/unit_tests/tasks/test_remove_app_and_related_data_task.py @@ -1,22 +1,17 @@ import logging -from collections.abc import Generator from datetime import UTC, datetime from unittest.mock import MagicMock, call, patch from uuid import uuid4 import pytest -from agenton.compositor import CompositorSessionSnapshot -from sqlalchemy import delete, select from sqlalchemy.orm import Session -from core.db.session_factory import session_factory from graphon.enums import WorkflowExecutionStatus from libs.archive_storage import ArchiveStorageNotConfiguredError -from models import AgentRuntimeSession, AgentRuntimeSessionOwnerType, AgentRuntimeSessionStatus, AppStar +from models import AppStar from models.enums import CreatorUserRole, WorkflowRunTriggeredFrom from models.workflow import WorkflowArchiveLog from tasks.remove_app_and_related_data_task import ( - _cleanup_active_agent_runtime_sessions_for_app, _delete_app_stars, _delete_app_workflow_archive_logs, _delete_archived_workflow_run_files, @@ -26,29 +21,6 @@ from tasks.remove_app_and_related_data_task import ( ) -@pytest.fixture(autouse=True) -def _create_agent_runtime_sessions_table() -> Generator[None, None, None]: - engine = session_factory.get_session_maker().kw["bind"] - AgentRuntimeSession.__table__.create(bind=engine, checkfirst=True) - yield - with session_factory.create_session() as session: - session.execute(delete(AgentRuntimeSession)) - session.commit() - AgentRuntimeSession.__table__.drop(bind=engine, checkfirst=True) - - -def _runtime_session_specs_json() -> str: - return ( - '[{"name":"execution_context","type":"dify.execution_context","deps":{},"metadata":{},' - '"config":{"tenant_id":"tenant-1"}},{"name":"history","type":"pydantic_ai.history","deps":{},' - '"metadata":{},"config":null}]' - ) - - -def _snapshot_json() -> str: - return CompositorSessionSnapshot(layers=[]).model_dump_json() - - class TestDeleteDraftVariablesBatch: def test_delete_draft_variables_batch_invalid_batch_size(self): """Test that invalid batch size raises ValueError.""" @@ -89,18 +61,12 @@ class TestDeleteDraftVariableOffloadData: """Test handling of database operation failures.""" mock_conn = MagicMock() file_ids = ["file-1"] - - # Make execute raise an exception mock_conn.execute.side_effect = Exception("Database error") - # Execute function - should not raise, but log error with caplog.at_level(logging.ERROR): result = _delete_draft_variable_offload_data(mock_conn, file_ids) - # Should return 0 when error occurs assert result == 0 - - # Verify error was logged assert "Error deleting draft variable offload data:" in caplog.text @@ -217,215 +183,3 @@ class TestDeleteArchivedWorkflowRunFiles: storage.list_objects.assert_called_once_with("tenant-1/app_id=app-1/") storage.delete_object.assert_has_calls([call("key-1"), call("key-2")], any_order=False) assert "Deleted 2 archive objects for app app-1" in caplog.text - - -class TestCleanupActiveAgentRuntimeSessionsForApp: - @patch("tasks.remove_app_and_related_data_task.cleanup_workflow_agent_runtime_session") - @patch("tasks.remove_app_and_related_data_task.cleanup_conversation_agent_runtime_session") - def test_enqueues_cleanup_for_active_rows_and_marks_rows_cleaned( - self, - mock_conversation_cleanup, - mock_workflow_cleanup, - ): - with session_factory.create_session() as session: - session.add( - AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id="agent-1", - agent_config_snapshot_id="snap-1", - conversation_id="conv-1", - backend_run_id="run-conv", - session_snapshot=_snapshot_json(), - composition_layer_specs=_runtime_session_specs_json(), - status=AgentRuntimeSessionStatus.ACTIVE, - ) - ) - session.add( - AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.WORKFLOW_RUN, - agent_id="agent-2", - workflow_id="wf-1", - workflow_run_id="wf-run-1", - node_id="node-1", - backend_run_id="run-wf", - session_snapshot=_snapshot_json(), - composition_layer_specs=_runtime_session_specs_json(), - status=AgentRuntimeSessionStatus.ACTIVE, - ) - ) - session.add( - AgentRuntimeSession( - tenant_id="tenant-1", - app_id="other-app", - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id="agent-3", - conversation_id="conv-2", - backend_run_id="run-other", - session_snapshot=_snapshot_json(), - composition_layer_specs=_runtime_session_specs_json(), - status=AgentRuntimeSessionStatus.ACTIVE, - ) - ) - session.commit() - - _cleanup_active_agent_runtime_sessions_for_app("tenant-1", "app-1", batch_size=1) - - assert mock_conversation_cleanup.delay.call_count == 1 - assert mock_workflow_cleanup.delay.call_count == 1 - conversation_payload = mock_conversation_cleanup.delay.call_args.args[0] - workflow_payload = mock_workflow_cleanup.delay.call_args.args[0] - assert conversation_payload["metadata"]["conversation_id"] == "conv-1" - assert workflow_payload["metadata"]["workflow_run_id"] == "wf-run-1" - assert conversation_payload["idempotency_key"].startswith("tenant-1:app-1:conv-1:agent-1:app-delete-cleanup:") - assert workflow_payload["idempotency_key"].startswith( - "tenant-1:app-1:wf-run-1:node-1:agent-2:app-delete-cleanup:" - ) - with session_factory.create_session() as session: - app_rows = session.scalars( - select(AgentRuntimeSession).where( - AgentRuntimeSession.tenant_id == "tenant-1", - AgentRuntimeSession.app_id == "app-1", - ) - ).all() - assert {row.status for row in app_rows} == {AgentRuntimeSessionStatus.CLEANED} - other_row = session.scalar( - select(AgentRuntimeSession).where( - AgentRuntimeSession.tenant_id == "tenant-1", - AgentRuntimeSession.app_id == "other-app", - ) - ) - assert other_row is not None - assert other_row.status == AgentRuntimeSessionStatus.ACTIVE - - @patch("tasks.remove_app_and_related_data_task.cleanup_workflow_agent_runtime_session") - @patch("tasks.remove_app_and_related_data_task.cleanup_conversation_agent_runtime_session") - def test_marks_rows_cleaned_even_when_enqueue_fails( - self, - mock_conversation_cleanup, - mock_workflow_cleanup, - ): - mock_conversation_cleanup.delay.side_effect = RuntimeError("queue down") - with session_factory.create_session() as session: - session.add( - AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id="agent-1", - conversation_id="conv-1", - backend_run_id="run-conv", - session_snapshot=_snapshot_json(), - composition_layer_specs=_runtime_session_specs_json(), - status=AgentRuntimeSessionStatus.ACTIVE, - ) - ) - session.add( - AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.WORKFLOW_RUN, - agent_id="agent-2", - workflow_id="wf-1", - workflow_run_id="wf-run-1", - node_id="node-1", - backend_run_id="run-wf", - session_snapshot=_snapshot_json(), - composition_layer_specs=_runtime_session_specs_json(), - status=AgentRuntimeSessionStatus.ACTIVE, - ) - ) - session.commit() - - _cleanup_active_agent_runtime_sessions_for_app("tenant-1", "app-1") - - mock_conversation_cleanup.delay.assert_called_once() - mock_workflow_cleanup.delay.assert_called_once() - with session_factory.create_session() as session: - rows = session.scalars( - select(AgentRuntimeSession).where( - AgentRuntimeSession.tenant_id == "tenant-1", - AgentRuntimeSession.app_id == "app-1", - ) - ).all() - assert rows - assert {row.status for row in rows} == {AgentRuntimeSessionStatus.CLEANED} - - @patch("tasks.remove_app_and_related_data_task.cleanup_workflow_agent_runtime_session") - @patch("tasks.remove_app_and_related_data_task.cleanup_conversation_agent_runtime_session") - def test_uses_row_identity_to_keep_distinct_app_delete_cleanup_jobs_distinct( - self, - mock_conversation_cleanup, - mock_workflow_cleanup, - ): - del mock_workflow_cleanup - with session_factory.create_session() as session: - first = AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id="agent-1", - agent_config_snapshot_id="snap-1", - conversation_id="conv-1", - backend_run_id="run-conv-1", - session_snapshot=_snapshot_json(), - composition_layer_specs=_runtime_session_specs_json(), - status=AgentRuntimeSessionStatus.ACTIVE, - ) - second = AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id="agent-1", - agent_config_snapshot_id="snap-2", - conversation_id="conv-1", - backend_run_id="run-conv-2", - session_snapshot=_snapshot_json(), - composition_layer_specs=_runtime_session_specs_json(), - status=AgentRuntimeSessionStatus.ACTIVE, - ) - session.add(first) - session.add(second) - session.commit() - - _cleanup_active_agent_runtime_sessions_for_app("tenant-1", "app-1") - - assert mock_conversation_cleanup.delay.call_count == 2 - payloads = [queued_call.args[0] for queued_call in mock_conversation_cleanup.delay.call_args_list] - assert payloads[0]["idempotency_key"] != payloads[1]["idempotency_key"] - assert payloads[0]["idempotency_key"].endswith(first.id) - assert payloads[1]["idempotency_key"].endswith(second.id) - - @patch("tasks.remove_app_and_related_data_task.cleanup_workflow_agent_runtime_session") - @patch("tasks.remove_app_and_related_data_task.cleanup_conversation_agent_runtime_session") - def test_marks_empty_runtime_layer_specs_rows_clean_without_enqueue( - self, - mock_conversation_cleanup, - mock_workflow_cleanup, - ): - del mock_workflow_cleanup - with session_factory.create_session() as session: - row = AgentRuntimeSession( - tenant_id="tenant-1", - app_id="app-1", - owner_type=AgentRuntimeSessionOwnerType.CONVERSATION, - agent_id="agent-1", - conversation_id="conv-1", - backend_run_id="run-no-specs", - session_snapshot=_snapshot_json(), - composition_layer_specs="[]", - status=AgentRuntimeSessionStatus.ACTIVE, - ) - session.add(row) - session.commit() - - _cleanup_active_agent_runtime_sessions_for_app("tenant-1", "app-1") - - mock_conversation_cleanup.delay.assert_not_called() - with session_factory.create_session() as session: - stored_row = session.scalar(select(AgentRuntimeSession).where(AgentRuntimeSession.id == row.id)) - assert stored_row is not None - assert stored_row.status == AgentRuntimeSessionStatus.CLEANED diff --git a/api/tests/unit_tests/tasks/test_resume_agent_app_task.py b/api/tests/unit_tests/tasks/test_resume_agent_app_task.py index 3b1c887e8ca..684ee06edf3 100644 --- a/api/tests/unit_tests/tasks/test_resume_agent_app_task.py +++ b/api/tests/unit_tests/tasks/test_resume_agent_app_task.py @@ -58,6 +58,7 @@ def test_resume_happy_path_account_user_sets_tenant_and_runs(mocker: MockerFixtu gen.return_value.resume_after_form_submission.assert_called_once() kwargs = gen.return_value.resume_after_form_submission.call_args.kwargs assert kwargs["conversation_id"] == "conv-1" + assert kwargs["form_id"] == "form-1" assert kwargs["user"] is account assert kwargs["app_model"] is app assert kwargs["invoke_from"] == InvokeFrom.WEB_APP diff --git a/api/tests/unit_tests/tasks/test_retry_document_indexing_task.py b/api/tests/unit_tests/tasks/test_retry_document_indexing_task.py new file mode 100644 index 00000000000..e5e1961112d --- /dev/null +++ b/api/tests/unit_tests/tasks/test_retry_document_indexing_task.py @@ -0,0 +1,65 @@ +from unittest.mock import MagicMock, patch +from uuid import uuid4 + +from sqlalchemy.orm import Session + +from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType +from models import Account, Tenant, TenantAccountJoin +from models.account import TenantAccountRole +from models.dataset import Dataset, Document +from models.enums import DatasetRuntimeMode, DataSourceType, DocumentCreatedFrom, IndexingStatus +from tasks.retry_document_indexing_task import retry_document_indexing_task + + +def test_retry_enforces_vector_space_admission(sqlite_session: Session) -> None: + tenant = Tenant(name="Retry tenant") + user = Account(name="Retry user", email=f"retry-{uuid4()}@example.com") + membership = TenantAccountJoin( + tenant_id=tenant.id, + account_id=user.id, + current=True, + role=TenantAccountRole.OWNER, + ) + dataset = Dataset( + id=str(uuid4()), + tenant_id=tenant.id, + name="Retry dataset", + created_by=user.id, + data_source_type=DataSourceType.UPLOAD_FILE, + indexing_technique=IndexTechniqueType.ECONOMY, + chunk_structure=IndexStructureType.PARAGRAPH_INDEX, + runtime_mode=DatasetRuntimeMode.GENERAL, + ) + document = Document( + id=str(uuid4()), + tenant_id=tenant.id, + dataset_id=dataset.id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + data_source_info="{}", + batch="retry-batch", + name="Retry document", + created_from=DocumentCreatedFrom.WEB, + created_by=user.id, + indexing_status=IndexingStatus.COMPLETED, + enabled=True, + archived=False, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + sqlite_session.add_all([tenant, user, membership, dataset, document]) + sqlite_session.commit() + features = MagicMock() + features.billing.enabled = False + + with ( + patch("tasks.retry_document_indexing_task.FeatureService.get_features", return_value=features), + patch("tasks.retry_document_indexing_task.IndexProcessorFactory"), + patch("tasks.retry_document_indexing_task.IndexingRunner") as indexing_runner, + patch("tasks.retry_document_indexing_task.redis_client"), + ): + retry_document_indexing_task.run(dataset.id, [document.id], user.id) + + indexing_runner.assert_called_once_with(enforce_vector_space_admission=True) + run_documents, run_session = indexing_runner.return_value.run.call_args.args + assert [item.id for item in run_documents] == [document.id] + assert isinstance(run_session, Session) diff --git a/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py b/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py index 3454035fe5f..8af7ac6a3b5 100644 --- a/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py +++ b/api/tests/unit_tests/tasks/test_segment_index_cleanup_tasks.py @@ -1,36 +1,94 @@ -from contextlib import nullcontext -from types import SimpleNamespace +import uuid +from collections.abc import Generator +from contextlib import contextmanager from unittest.mock import MagicMock, patch -from models.enums import SegmentStatus +import pytest +from sqlalchemy import event +from sqlalchemy.orm import Session, sessionmaker + +from core.rag.index_processor.constant.index_type import IndexStructureType +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus, SegmentStatus from tasks.delete_segment_from_index_task import delete_segment_from_index_task from tasks.disable_segment_from_index_task import disable_segment_from_index_task from tasks.disable_segments_from_index_task import disable_segments_from_index_task -def test_disable_segment_commits_index_cleanup() -> None: - dataset = SimpleNamespace(id="dataset-1") - document = SimpleNamespace(enabled=True, archived=False, indexing_status="completed", doc_form="text_model") - segment = SimpleNamespace( - id="segment-1", - status=SegmentStatus.COMPLETED, - index_node_id="node-1", - disabled_by="user-1", - get_dataset=MagicMock(return_value=dataset), - get_document=MagicMock(return_value=document), +@pytest.fixture +def indexed_segment(sqlite_session: Session) -> tuple[Dataset, Document, DocumentSegment]: + """Persist the complete owner chain consumed by segment cleanup tasks.""" + tenant_id = str(uuid.uuid4()) + created_by = str(uuid.uuid4()) + dataset = Dataset( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Cleanup dataset", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + is_multimodal=False, ) - session = MagicMock() - session.scalar.return_value = segment + document = Document( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + dataset_id=dataset.id, + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + batch="batch-1", + name="document.txt", + created_from=DocumentCreatedFrom.WEB, + created_by=created_by, + indexing_status=IndexingStatus.COMPLETED, + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + segment = DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset.id, + document_id=document.id, + position=1, + content="content", + word_count=1, + tokens=1, + created_by=created_by, + index_node_id="node-1", + index_node_hash="hash-1", + disabled_by=created_by, + status=SegmentStatus.COMPLETED, + ) + sqlite_session.add_all([dataset, document, segment]) + sqlite_session.commit() + return dataset, document, segment + + +@contextmanager +def _record_transaction_events( + sqlite_session_factory: sessionmaker[Session], phase_events: list[str] +) -> Generator[None]: + """Record real commits made by the task-owned SQLite session.""" + session_type = sqlite_session_factory.class_ + + def after_commit(_session: Session) -> None: + phase_events.append("commit") + + event.listen(session_type, "after_commit", after_commit) + try: + yield + finally: + event.remove(session_type, "after_commit", after_commit) + + +def test_disable_segment_commits_index_cleanup( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, _document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") processor = MagicMock() processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") disable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) with ( - patch( - "tasks.disable_segment_from_index_task.session_factory.create_session", return_value=nullcontext(session) - ), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.disable_segment_from_index_task.IndexProcessorFactory") as processor_factory, patch( "services.summary_index_service.SummaryIndexService.disable_summaries_for_segments", @@ -42,31 +100,24 @@ def test_disable_segment_commits_index_cleanup() -> None: disable_segment_from_index_task.run(segment.id) assert phase_events == ["clean", "commit", "summary"] + disable_summaries.assert_called_once() + assert disable_summaries.call_args.kwargs["dataset"].id == dataset.id + assert disable_summaries.call_args.kwargs["segment_ids"] == [segment.id] + assert disable_summaries.call_args.kwargs["disabled_by"] == segment.disabled_by -def test_disable_segments_commits_index_cleanup() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status="completed", - doc_form="text_model", - ) - segment = SimpleNamespace(id="segment-1", index_node_id="node-1", disabled_by="user-1") - session = MagicMock() - session.scalar.side_effect = [dataset, document] - session.scalars.return_value.all.return_value = [segment] +def test_disable_segments_commits_index_cleanup( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") processor = MagicMock() processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") disable_summaries = MagicMock(side_effect=lambda *_args, **_kwargs: phase_events.append("summary")) with ( - patch( - "tasks.disable_segments_from_index_task.session_factory.create_session", return_value=nullcontext(session) - ), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.disable_segments_from_index_task.IndexProcessorFactory") as processor_factory, patch( "services.summary_index_service.SummaryIndexService.disable_summaries_for_segments", @@ -78,29 +129,26 @@ def test_disable_segments_commits_index_cleanup() -> None: disable_segments_from_index_task.run([segment.id], dataset.id, document.id) assert phase_events == ["clean", "commit", "summary"] + disable_summaries.assert_called_once() + assert disable_summaries.call_args.kwargs["dataset"].id == dataset.id + assert disable_summaries.call_args.kwargs["segment_ids"] == [segment.id] + assert disable_summaries.call_args.kwargs["disabled_by"] == segment.disabled_by -def test_delete_segment_commits_index_cleanup_without_attachments() -> None: - dataset = SimpleNamespace(id="dataset-1", is_multimodal=False) - document = SimpleNamespace( - id="document-1", - enabled=True, - archived=False, - indexing_status="completed", - doc_form="text_model", - ) - session = MagicMock() - session.scalar.side_effect = [dataset, document] +def test_delete_segment_commits_index_cleanup_without_attachments( + indexed_segment: tuple[Dataset, Document, DocumentSegment], + sqlite_session_factory: sessionmaker[Session], +) -> None: + dataset, document, segment = indexed_segment phase_events: list[str] = [] - session.commit.side_effect = lambda: phase_events.append("commit") processor = MagicMock() processor.clean.side_effect = lambda *_args, **_kwargs: phase_events.append("clean") with ( - patch("tasks.delete_segment_from_index_task.session_factory.create_session", return_value=nullcontext(session)), + _record_transaction_events(sqlite_session_factory, phase_events), patch("tasks.delete_segment_from_index_task.IndexProcessorFactory") as processor_factory, ): processor_factory.return_value.init_index_processor.return_value = processor - delete_segment_from_index_task.run(["node-1"], dataset.id, document.id, ["segment-1"]) + delete_segment_from_index_task.run(["node-1"], dataset.id, document.id, [segment.id]) assert phase_events == ["clean", "commit"] diff --git a/api/tests/unit_tests/tasks/test_sync_website_document_indexing_task.py b/api/tests/unit_tests/tasks/test_sync_website_document_indexing_task.py new file mode 100644 index 00000000000..b5f88f1a706 --- /dev/null +++ b/api/tests/unit_tests/tasks/test_sync_website_document_indexing_task.py @@ -0,0 +1,98 @@ +import uuid +from unittest.mock import MagicMock, patch + +from sqlalchemy import select +from sqlalchemy.orm import Session + +from core.rag.index_processor.constant.index_type import IndexStructureType +from models.dataset import Dataset, Document, DocumentSegment +from models.enums import DataSourceType, DocumentCreatedFrom +from tasks.sync_website_document_indexing_task import sync_website_document_indexing_task + + +def _dataset(tenant_id: str) -> Dataset: + return Dataset( + id=str(uuid.uuid4()), + tenant_id=tenant_id, + name="Website dataset", + data_source_type=DataSourceType.WEBSITE_CRAWL, + created_by=str(uuid.uuid4()), + ) + + +def _document(dataset: Dataset) -> Document: + return Document( + id=str(uuid.uuid4()), + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + position=1, + data_source_type=DataSourceType.WEBSITE_CRAWL, + batch="batch-1", + name="Website document", + created_from=DocumentCreatedFrom.WEB, + created_by=str(uuid.uuid4()), + doc_form=IndexStructureType.PARAGRAPH_INDEX, + ) + + +def _segment(*, tenant_id: str, dataset_id: str, document_id: str) -> DocumentSegment: + return DocumentSegment( + tenant_id=tenant_id, + dataset_id=dataset_id, + document_id=document_id, + position=1, + content="content", + word_count=1, + tokens=1, + created_by=str(uuid.uuid4()), + ) + + +def test_rejects_document_outside_dataset_before_side_effects(sqlite_session: Session) -> None: + tenant_id = str(uuid.uuid4()) + requested_dataset = _dataset(tenant_id) + foreign_dataset = _dataset(tenant_id) + foreign_document = _document(foreign_dataset) + sqlite_session.add_all([requested_dataset, foreign_dataset, foreign_document]) + sqlite_session.commit() + + with ( + patch("tasks.sync_website_document_indexing_task.FeatureService") as feature_service, + patch("tasks.sync_website_document_indexing_task.IndexProcessorFactory") as processor_factory, + ): + sync_website_document_indexing_task(requested_dataset.id, foreign_document.id) + + feature_service.get_features.assert_not_called() + processor_factory.assert_not_called() + + +def test_cleanup_is_owner_scoped_and_skips_empty_vector_ids(sqlite_session: Session) -> None: + tenant_id = str(uuid.uuid4()) + dataset = _dataset(tenant_id) + document = _document(dataset) + owned_segment = _segment(tenant_id=tenant_id, dataset_id=dataset.id, document_id=document.id) + other_dataset = _dataset(tenant_id) + decoy_segments = [ + _segment(tenant_id=tenant_id, dataset_id=dataset.id, document_id=str(uuid.uuid4())), + _segment(tenant_id=tenant_id, dataset_id=other_dataset.id, document_id=document.id), + _segment(tenant_id=str(uuid.uuid4()), dataset_id=dataset.id, document_id=document.id), + ] + for index, segment in enumerate(decoy_segments): + segment.index_node_id = f"decoy-node-{index}" + sqlite_session.add_all([dataset, document, owned_segment, other_dataset, *decoy_segments]) + sqlite_session.commit() + + features = MagicMock() + features.billing.enabled = False + with ( + patch("tasks.sync_website_document_indexing_task.FeatureService.get_features", return_value=features), + patch("tasks.sync_website_document_indexing_task.IndexProcessorFactory") as processor_factory, + patch("tasks.sync_website_document_indexing_task.IndexingRunner") as indexing_runner, + patch("tasks.sync_website_document_indexing_task.redis_client"), + ): + sync_website_document_indexing_task(dataset.id, document.id) + + processor_factory.return_value.init_index_processor.return_value.clean.assert_not_called() + indexing_runner.return_value.run.assert_called_once() + sqlite_session.expire_all() + assert set(sqlite_session.scalars(select(DocumentSegment.id))) == {segment.id for segment in decoy_segments} diff --git a/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py b/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py index cd5df4466e3..84f95d649c4 100644 --- a/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py +++ b/api/tests/unit_tests/tasks/test_trigger_processing_tasks.py @@ -55,10 +55,6 @@ class TestDispatchTriggeredWorkflow: (``get_workflows``, ``reserve``, ``create_end_user_batch``, ...) to drive the path it targets. """ - session_cm = MagicMock() - session_cm.__enter__.return_value = MagicMock() - session_cm.__exit__.return_value = False - invoke_response = MagicMock() invoke_response.cancelled = False invoke_response.variables = {} @@ -105,11 +101,6 @@ class TestDispatchTriggeredWorkflow: "create_end_user_batch", return_value={}, ) as create_end_user_batch, - patch.object( - trigger_processing_tasks_module.session_factory, - "create_session", - return_value=session_cm, - ), patch.object( trigger_processing_tasks_module.QuotaService, "reserve", diff --git a/api/tests/unit_tests/tasks/test_workflow_execute_task.py b/api/tests/unit_tests/tasks/test_workflow_execute_task.py index 3b9cad30018..7c7f8f34b08 100644 --- a/api/tests/unit_tests/tasks/test_workflow_execute_task.py +++ b/api/tests/unit_tests/tasks/test_workflow_execute_task.py @@ -3,6 +3,7 @@ from __future__ import annotations import json import logging import uuid +from collections.abc import Generator, Mapping from contextlib import nullcontext from datetime import datetime from decimal import Decimal @@ -17,6 +18,7 @@ from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import AdvancedChatAppGenerateEntity, InvokeFrom, WorkflowAppGenerateEntity from graphon.entities import WorkflowStartReason from graphon.enums import WorkflowExecutionStatus +from models.account import Account from models.base import TypeBase from models.enums import ConversationFromSource, CreatorUserRole, WorkflowRunTriggeredFrom from models.model import App, AppMode, Conversation, Message @@ -36,6 +38,7 @@ from tasks.app_generate.workflow_execute_task import ( class _StreamEventModel(BaseModel): event: object | None = None task_id: object | None = None + message: object | None = None def _build_advanced_chat_generate_entity(conversation_id: str | None) -> AdvancedChatAppGenerateEntity: @@ -248,6 +251,21 @@ def test_get_task_id(event: object, expected: str | None): assert workflow_execute_task_module._get_task_id(event) == expected +@pytest.mark.parametrize( + ("event", "expected"), + [ + ({"message": "workflow error"}, "workflow error"), + (_StreamEventModel(message="workflow error"), "workflow error"), + ({"message": ""}, None), + ({"message": 123}, None), + ({}, None), + ("workflow error", None), + ], +) +def test_get_error_message(event: str | Mapping[str, object] | BaseModel, expected: str | None): + assert workflow_execute_task_module._get_error_message(event) == expected + + @pytest.fixture def mock_topic(monkeypatch: pytest.MonkeyPatch) -> MagicMock: topic = MagicMock() @@ -486,6 +504,38 @@ def test_publish_streaming_response_publishes_failed_terminal_on_exhaustion_with assert "ended without a terminal event" in caplog.text +def test_publish_streaming_response_uses_error_message_for_failed_terminal(mock_topic: MagicMock): + def response_stream() -> Generator[str | Mapping[str, object] | BaseModel, None, None]: + yield { + "event": "error", + "workflow_run_id": "workflow-run-id", + "code": "invalid_param", + "message": "LLM provider and model are required.", + "status": 400, + } + + _publish_streaming_response( + response_stream(), + "workflow-run-id", + app_mode=AppMode.WORKFLOW, + workflow_id="workflow-id", + inputs={}, + started_reason=WorkflowStartReason.INITIAL, + ) + + payloads = _published_payloads(mock_topic) + error_payload = payloads[0] + finished_payload = payloads[-1] + assert isinstance(error_payload, dict) + assert isinstance(finished_payload, dict) + assert error_payload["status"] == 400 + assert error_payload["message"] == "LLM provider and model are required." + finished_data = finished_payload["data"] + assert isinstance(finished_data, dict) + assert finished_data["status"] == WorkflowExecutionStatus.FAILED + assert finished_data["error"] == "LLM provider and model are required." + + def test_publish_streaming_response_does_not_publish_synthetic_failure_after_terminal_event(mock_topic: MagicMock): response_stream = iter( [ @@ -583,44 +633,50 @@ def test_app_runner_streaming_failure_publishes_started_then_failed_workflow_fin assert finished_payload["data"]["files"] == [] -def test_app_runner_resolves_account_without_switching_tenant(monkeypatch: pytest.MonkeyPatch): +def test_app_runner_resolves_account_without_switching_tenant( + sqlite_session_factory: sessionmaker[Session], +): + account_id = str(uuid.uuid4()) exec_params = AppExecutionParams( app_id="app-id", workflow_id="workflow-id", tenant_id="resource-tenant-id", app_mode=AppMode.WORKFLOW, - user={"TYPE": "account", "user_id": "user-id"}, + user={"TYPE": "account", "user_id": account_id}, args={"inputs": {}}, invoke_from=InvokeFrom.EXPLORE, streaming=True, workflow_run_id="workflow-run-id", ) - runner = _AppRunner(session_factory=MagicMock(), exec_params=exec_params) - account = MagicMock() - session = MagicMock() - session.get.return_value = account - monkeypatch.setattr(runner, "_session", lambda: nullcontext(session)) + with sqlite_session_factory() as session: + account = Account(name="Runner Account", email="runner@example.com") + account.id = account_id + session.add(account) + session.commit() + runner = _AppRunner(session_factory=sqlite_session_factory, exec_params=exec_params) resolved_user = runner._resolve_user() - assert resolved_user is account - account.set_tenant_id_with_session.assert_not_called() + assert resolved_user.id == account_id + assert resolved_user.current_tenant is None -def test_resolve_account_for_run_without_switching_tenant(): - account = MagicMock() - session = MagicMock() - session.get.return_value = account +def test_resolve_account_for_run_without_switching_tenant(sqlite_session: Session): + account_id = str(uuid.uuid4()) + account = Account(name="Run Account", email="run@example.com") + account.id = account_id + sqlite_session.add(account) + sqlite_session.commit() workflow_run = MagicMock( created_by_role=CreatorUserRole.ACCOUNT, - created_by="user-id", + created_by=account_id, tenant_id="resource-tenant-id", ) - resolved_user = workflow_execute_task_module._resolve_user_for_run(session, workflow_run) + resolved_user = workflow_execute_task_module._resolve_user_for_run(sqlite_session, workflow_run) assert resolved_user is account - account.set_tenant_id_with_session.assert_not_called() + assert account.current_tenant is None def test_app_runner_streaming_failure_keeps_existing_pre_runtime_helper_behavior( @@ -798,6 +854,7 @@ def test_resume_app_execution_returns_early_when_advanced_chat_missing_conversat def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs( monkeypatch: pytest.MonkeyPatch, + sqlite_session: Session, sqlite_session_factory: sessionmaker[Session], ): generate_entity = _build_advanced_chat_generate_entity(conversation_id="conversation-id") @@ -824,8 +881,6 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs( "tasks.app_generate.workflow_execute_task.DifyCoreRepositoryFactory.create_workflow_node_execution_repository", lambda **kwargs: MagicMock(), ) - session = MagicMock() - _resume_advanced_chat( app_model=SimpleNamespace(id="app-id", tenant_id="resource-tenant-id"), workflow=workflow, @@ -839,12 +894,12 @@ def test_resume_advanced_chat_publishes_events_for_originally_blocking_runs( pause_state_config=MagicMock(), workflow_run_id="workflow-run-id", workflow_run=SimpleNamespace(triggered_from="app_run"), - session=session, + session=sqlite_session, ) resumed_entity = generator_instance.resume.call_args.kwargs["application_generate_entity"] assert resumed_entity.stream is True - assert generator_instance.resume.call_args.kwargs["session"] is session + assert generator_instance.resume.call_args.kwargs["session"] is sqlite_session publish_streaming_response.assert_called_once_with( response_stream, "workflow-run-id", diff --git a/api/tests/unit_tests/tasks/test_workflow_node_execution_identity.py b/api/tests/unit_tests/tasks/test_workflow_node_execution_identity.py new file mode 100644 index 00000000000..e1e8f8bc0cf --- /dev/null +++ b/api/tests/unit_tests/tasks/test_workflow_node_execution_identity.py @@ -0,0 +1,28 @@ +from datetime import UTC, datetime + +from graphon.entities import WorkflowNodeExecution +from models.workflow import WorkflowNodeExecutionModel +from tasks.workflow_node_execution_tasks import _update_node_execution_from_domain + + +def test_celery_update_preserves_workflow_agent_binding_identity() -> None: + stored = WorkflowNodeExecutionModel( + process_data='{"workflow_agent_binding_id": "workflow-binding-1"}', + ) + incoming = WorkflowNodeExecution( + id="execution-1", + workflow_id="workflow-1", + node_id="node-1", + node_type="agent", + title="Agent", + index=1, + created_at=datetime.now(UTC), + process_data={"retry": True}, + ) + + _update_node_execution_from_domain(stored, incoming) + + assert stored.process_data_dict == { + "retry": True, + "workflow_agent_binding_id": "workflow-binding-1", + } diff --git a/api/tests/unit_tests/test_app_factory.py b/api/tests/unit_tests/test_app_factory.py new file mode 100644 index 00000000000..acdeecc07c0 --- /dev/null +++ b/api/tests/unit_tests/test_app_factory.py @@ -0,0 +1,319 @@ +"""Enterprise license gating performed by the global ``before_request`` hook.""" + +from unittest.mock import patch + +import pytest +from flask import Blueprint, Flask +from flask_restx import Resource + +from app_factory import create_flask_app_with_configs +from enums import DeploymentEdition +from libs.external_api import ExternalApi +from services.entities.feature_entities import LicenseStatus + +INVALID_STATUSES = [LicenseStatus.INACTIVE, LicenseStatus.EXPIRED, LicenseStatus.LOST] +VALID_STATUSES = [LicenseStatus.ACTIVE, LicenseStatus.EXPIRING] + + +def _license(status: LicenseStatus | None): + return patch("app_factory.EnterpriseService.get_cached_license_status", return_value=status) + + +def _enterprise(): + return patch("app_factory.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.ENTERPRISE) + + +def _community(): + return patch("app_factory.dify_config.DEPLOYMENT_EDITION", DeploymentEdition.COMMUNITY) + + +@pytest.fixture +def gated_app() -> Flask: + app = create_flask_app_with_configs() + + @app.route("/v1/chat-messages", methods=["POST"]) + def service_api_route(): + return {"surface": "service_api"} + + @app.route("/v1/") + def service_api_index_route(): + return {"surface": "service_api_index"} + + @app.route("/mcp/server//mcp", methods=["POST"]) + def mcp_route(server_code: str): + return {"surface": "mcp"} + + @app.route("/triggers/webhook/", methods=["POST"]) + def trigger_route(webhook_id: str): + return {"surface": "triggers"} + + @app.route("/console/api/apps") + def console_route(): + return {"surface": "console"} + + @app.route("/console/api/login", methods=["POST"]) + def console_bootstrap_route(): + return {"surface": "console_bootstrap"} + + @app.route("/api/messages") + def webapp_route(): + return {"surface": "webapp"} + + @app.route("/api/system-features") + def webapp_bootstrap_route(): + return {"surface": "webapp_bootstrap"} + + @app.route("/health") + def health_route(): + return {"surface": "health"} + + @app.route("/inner/api/rbac/check-access", methods=["POST"]) + def inner_api_route(): + return {"surface": "inner_api"} + + @app.route("/files/upload/for-plugin", methods=["POST"]) + def files_route(): + return {"surface": "files"} + + return app + + +class TestServiceApiLicenseGate: + """/v1 is a bearer-token surface, so it is gated with an opaque 403.""" + + @pytest.mark.parametrize("status", INVALID_STATUSES) + def test_blocks_when_license_invalid(self, gated_app: Flask, status: LicenseStatus): + with _enterprise(), _license(status): + response = gated_app.test_client().post("/v1/chat-messages") + + assert response.status_code == 403 + + def test_block_response_carries_machine_readable_marker(self, gated_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().post("/v1/chat-messages") + + assert b"license_required" in response.data + + def test_block_response_does_not_leak_license_status(self, gated_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().post("/v1/chat-messages") + + assert b"expired" not in response.data.lower() + + def test_blocks_when_license_status_unavailable(self, gated_app: Flask): + with _enterprise(), _license(None): + response = gated_app.test_client().post("/v1/chat-messages") + + assert response.status_code == 403 + + def test_blocks_when_license_lookup_raises(self, gated_app: Flask): + lookup_failed = patch( + "app_factory.EnterpriseService.get_cached_license_status", + side_effect=RuntimeError("enterprise api unreachable"), + ) + with _enterprise(), lookup_failed: + response = gated_app.test_client().post("/v1/chat-messages") + + assert response.status_code == 403 + + def test_blocks_index_route(self, gated_app: Flask): + """/v1 has no sign-in page to bootstrap, so nothing on it is exempt.""" + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().get("/v1/") + + assert response.status_code == 403 + + @pytest.mark.parametrize("status", VALID_STATUSES) + def test_allows_when_license_valid(self, gated_app: Flask, status: LicenseStatus): + with _enterprise(), _license(status): + response = gated_app.test_client().post("/v1/chat-messages") + + assert response.status_code == 200 + + def test_allows_unclassified_status(self, gated_app: Flask): + """LicenseStatus.NONE is not in the blocked set — parity with console/webapp.""" + with _enterprise(), _license(LicenseStatus.NONE): + response = gated_app.test_client().post("/v1/chat-messages") + + assert response.status_code == 200 + + @pytest.mark.parametrize("status", INVALID_STATUSES) + def test_does_not_gate_community_edition(self, gated_app: Flask, status: LicenseStatus): + with _community(), _license(status): + response = gated_app.test_client().post("/v1/chat-messages") + + assert response.status_code == 200 + + +class TestMcpLicenseGate: + """/mcp invokes apps for external MCP clients, so it is gated like the Service API.""" + + @pytest.mark.parametrize("status", INVALID_STATUSES) + def test_blocks_when_license_invalid(self, gated_app: Flask, status: LicenseStatus): + with _enterprise(), _license(status): + response = gated_app.test_client().post("/mcp/server/srv-code/mcp") + + assert response.status_code == 403 + + def test_block_response_carries_machine_readable_marker(self, gated_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().post("/mcp/server/srv-code/mcp") + + assert b"license_required" in response.data + + @pytest.mark.parametrize("status", VALID_STATUSES) + def test_allows_when_license_valid(self, gated_app: Flask, status: LicenseStatus): + with _enterprise(), _license(status): + response = gated_app.test_client().post("/mcp/server/srv-code/mcp") + + assert response.status_code == 200 + + @pytest.mark.parametrize("status", INVALID_STATUSES) + def test_does_not_gate_community_edition(self, gated_app: Flask, status: LicenseStatus): + with _community(), _license(status): + response = gated_app.test_client().post("/mcp/server/srv-code/mcp") + + assert response.status_code == 200 + + +class TestTriggerLicenseGate: + """Inbound webhooks are refused so senders retry, rather than dropping events.""" + + @pytest.mark.parametrize("status", INVALID_STATUSES) + def test_blocks_when_license_invalid(self, gated_app: Flask, status: LicenseStatus): + with _enterprise(), _license(status): + response = gated_app.test_client().post("/triggers/webhook/hook-id") + + assert response.status_code == 503 + + def test_block_response_carries_machine_readable_marker(self, gated_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().post("/triggers/webhook/hook-id") + + assert b"license_required" in response.data + + @pytest.mark.parametrize("status", VALID_STATUSES) + def test_allows_when_license_valid(self, gated_app: Flask, status: LicenseStatus): + with _enterprise(), _license(status): + response = gated_app.test_client().post("/triggers/webhook/hook-id") + + assert response.status_code == 200 + + @pytest.mark.parametrize("status", INVALID_STATUSES) + def test_does_not_gate_community_edition(self, gated_app: Flask, status: LicenseStatus): + with _community(), _license(status): + response = gated_app.test_client().post("/triggers/webhook/hook-id") + + assert response.status_code == 200 + + +class TestGateThroughRealErrorHandlers: + """Gate errors must survive each blueprint's error handling: flask-restx vs plain Flask.""" + + @pytest.fixture + def wired_app(self) -> Flask: + app = create_flask_app_with_configs() + + service_api_bp = Blueprint("service_api_test", __name__, url_prefix="/v1") + api = ExternalApi(service_api_bp) + + @api.route("/chat-messages") + class ChatMessages(Resource): + def post(self): + return {"surface": "service_api"} + + app.register_blueprint(service_api_bp) + + trigger_bp = Blueprint("trigger_test", __name__, url_prefix="/triggers") + + @trigger_bp.route("/webhook/", methods=["POST"]) + def webhook_route(webhook_id: str): + return {"surface": "triggers"} + + app.register_blueprint(trigger_bp) + return app + + def test_service_api_block_is_json_with_license_marker(self, wired_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = wired_app.test_client().post("/v1/chat-messages") + + assert response.status_code == 403 + body = response.get_json() + assert body["message"] == "license_required" + assert body["status"] == 403 + + def test_service_api_block_does_not_clear_cookies(self, wired_app: Flask): + """Force-logout cookie clearing belongs to the cookie-authed surfaces only.""" + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = wired_app.test_client().post("/v1/chat-messages") + + assert response.headers.getlist("Set-Cookie") == [] + + def test_trigger_block_survives_plain_blueprint_handling(self, wired_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = wired_app.test_client().post("/triggers/webhook/hook-id") + + assert response.status_code == 503 + assert b"license_required" in response.data + + def test_surfaces_are_reachable_when_license_valid(self, wired_app: Flask): + with _enterprise(), _license(LicenseStatus.ACTIVE): + service_api = wired_app.test_client().post("/v1/chat-messages") + triggers = wired_app.test_client().post("/triggers/webhook/hook-id") + + assert service_api.status_code == 200 + assert triggers.status_code == 200 + + +class TestUngatedSurfaces: + """Surfaces that must stay reachable while the license is invalid.""" + + def test_inner_api_is_not_gated(self, gated_app: Flask): + """dify-enterprise control plane — gating it could block license recovery itself.""" + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().post("/inner/api/rbac/check-access") + + assert response.status_code == 200 + + def test_files_data_plane_is_not_gated(self, gated_app: Flask): + """Signed file URLs are fetched by the plugin daemon and by LLM vendors.""" + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().post("/files/upload/for-plugin") + + assert response.status_code == 200 + + +class TestSessionSurfaceLicenseGate: + """Console and webapp are cookie-authed, so they keep force-logout 401 semantics.""" + + @pytest.mark.parametrize("status", INVALID_STATUSES) + def test_blocks_console_with_force_logout(self, gated_app: Flask, status: LicenseStatus): + with _enterprise(), _license(status): + response = gated_app.test_client().get("/console/api/apps") + + assert response.status_code == 401 + + @pytest.mark.parametrize("status", INVALID_STATUSES) + def test_blocks_webapp_with_force_logout(self, gated_app: Flask, status: LicenseStatus): + with _enterprise(), _license(status): + response = gated_app.test_client().get("/api/messages") + + assert response.status_code == 401 + + def test_console_bootstrap_route_stays_reachable(self, gated_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().post("/console/api/login") + + assert response.status_code == 200 + + def test_webapp_bootstrap_route_stays_reachable(self, gated_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().get("/api/system-features") + + assert response.status_code == 200 + + def test_health_route_is_never_gated(self, gated_app: Flask): + with _enterprise(), _license(LicenseStatus.EXPIRED): + response = gated_app.test_client().get("/health") + + assert response.status_code == 200 diff --git a/api/tests/unit_tests/test_sqlite_fixtures.py b/api/tests/unit_tests/test_sqlite_fixtures.py new file mode 100644 index 00000000000..11a60b407c7 --- /dev/null +++ b/api/tests/unit_tests/test_sqlite_fixtures.py @@ -0,0 +1,103 @@ +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from threading import Barrier + +import pytest +from sqlalchemy import create_engine, inspect, text +from sqlalchemy.engine import URL, Engine +from sqlalchemy.exc import UnboundExecutionError +from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.pool import QueuePool + +import core.db.session_factory as session_factory_module +from models.account import Account +from models.base import TypeBase +from models.model import ExporleBanner + + +def test_sqlite_session_contains_the_full_registered_schema(sqlite_session: Session) -> None: + table_names = set(inspect(sqlite_session.get_bind()).get_table_names()) + + assert table_names == set(TypeBase.metadata.tables) + + +@pytest.mark.parametrize("sqlite_session", [(Account,)], indirect=True) +def test_sqlite_session_accepts_deferred_legacy_indirect_parameters(sqlite_session: Session) -> None: + """Prove legacy model parameters no longer limit the copied schema.""" + + assert inspect(sqlite_session.get_bind()).has_table(ExporleBanner.__tablename__) + + +def test_sqlite_engine_is_a_pristine_file_copy( + sqlite_engine: Engine, + request: pytest.FixtureRequest, +) -> None: + sqlite_database_template: Path = request.getfixturevalue("_sqlite_database_template") + assert isinstance(sqlite_engine.pool, QueuePool) + assert sqlite_engine.url.database != str(sqlite_database_template) + + with sqlite_engine.begin() as connection: + connection.execute(text("CREATE TABLE per_test_mutation (value INTEGER NOT NULL)")) + + template_engine = create_engine(URL.create("sqlite", database=str(sqlite_database_template))) + try: + assert not inspect(template_engine).has_table("per_test_mutation") + finally: + template_engine.dispose() + + +def test_core_session_factory_uses_the_shared_sqlite_session_factory( + sqlite_session_factory: sessionmaker[Session], +) -> None: + assert session_factory_module.session_factory.get_session_maker() is sqlite_session_factory + + with sqlite_session_factory.begin() as session: + session.execute(text("CREATE TABLE global_factory_probe (value INTEGER NOT NULL)")) + session.execute(text("INSERT INTO global_factory_probe (value) VALUES (42)")) + + with session_factory_module.session_factory.create_session() as session: + assert session.scalar(text("SELECT value FROM global_factory_probe")) == 42 + + +def test_unbound_session_factory_disables_explicit_and_global_database_access( + unbound_session_factory: sessionmaker[Session], +) -> None: + assert session_factory_module.session_factory.get_session_maker() is unbound_session_factory + + with unbound_session_factory() as session: + with pytest.raises(UnboundExecutionError): + session.get_bind() + + with session_factory_module.session_factory.create_session() as session: + with pytest.raises(UnboundExecutionError): + session.execute(text("SELECT 1")) + + +def test_unbound_session_rejects_database_access(unbound_session: Session) -> None: + with pytest.raises(UnboundExecutionError): + unbound_session.scalar(text("SELECT 1")) + + +def test_sqlite_session_factory_shares_one_database_across_worker_sessions( + sqlite_session_factory: sessionmaker[Session], +) -> None: + with sqlite_session_factory.begin() as session: + session.execute(text("CREATE TABLE thread_probe (value INTEGER NOT NULL)")) + session.execute(text("INSERT INTO thread_probe (value) VALUES (42)")) + + worker_barrier = Barrier(2) + + def read_value() -> tuple[int, int]: + with sqlite_session_factory() as session: + connection = session.connection() + worker_barrier.wait(timeout=1) + value = session.scalar(text("SELECT value FROM thread_probe")) + connection_id = id(connection.connection.dbapi_connection) + return connection_id, value + + with ThreadPoolExecutor(max_workers=2) as executor: + futures = [executor.submit(read_value) for _ in range(2)] + results = [future.result() for future in futures] + + assert {value for _, value in results} == {42} + assert len({connection_id for connection_id, _ in results}) == 2 diff --git a/api/tests/unit_tests/tools/test_mcp_tool.py b/api/tests/unit_tests/tools/test_mcp_tool.py index eb7bb34dbe4..00a84f71f7f 100644 --- a/api/tests/unit_tests/tools/test_mcp_tool.py +++ b/api/tests/unit_tests/tools/test_mcp_tool.py @@ -135,6 +135,19 @@ class TestMCPToolInvoke: values = {m.message.variable_name: m.message.variable_value for m in var_msgs} assert values == {"a": 1, "b": "x"} + def test_invoke_yields_json_when_structured_content_has_no_output_schema(self, orm_session: Session) -> None: + tool = _make_mcp_tool() + result = CallToolResult(content=[], structuredContent={"a": 1, "b": "x"}) + + with patch.object(tool, "invoke_remote_mcp_tool", return_value=result): + messages = list(tool._invoke(session=orm_session, user_id="test_user", tool_parameters={})) + + assert len(messages) == 1 + msg = messages[0] + assert msg.type == ToolInvokeMessage.MessageType.JSON + assert isinstance(msg.message, ToolInvokeMessage.JsonMessage) + assert msg.message.json_object == {"a": 1, "b": "x"} + class TestMCPToolUsageExtraction: """Test usage metadata extraction from MCP tool results.""" diff --git a/api/uv.lock b/api/uv.lock index 125cfca452b..2c2a689f541 100644 --- a/api/uv.lock +++ b/api/uv.lock @@ -51,7 +51,7 @@ members = [ "dify-vdb-weaviate", ] overrides = [ - { name = "cryptography", specifier = ">=49.0.0,<50.0.0" }, + { name = "cryptography", specifier = ">=50.0.0,<51.0.0" }, { name = "litellm", specifier = ">=1.83.10,<2.0.0" }, { name = "pyarrow", specifier = ">=23.0.1,<24.0.0" }, { name = "setuptools", specifier = ">=80.10.2,<81" }, @@ -89,7 +89,7 @@ wheels = [ [[package]] name = "aiohttp" -version = "3.14.1" +version = "3.14.3" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "aiohappyeyeballs" }, @@ -101,26 +101,26 @@ dependencies = [ { name = "typing-extensions" }, { name = "yarl" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/82/78/8ea7308cac6934de8c74a14f3d5f65d1c89287426688be79538d0e5c013d/aiohttp-3.14.1.tar.gz", hash = "sha256:307f2cff90a764d329e77040603fa032db89c5c24fdad50c4c15334cba744035", size = 7955794, upload-time = "2026-06-07T21:09:35.529Z" } +sdist = { url = "https://files.pythonhosted.org/packages/58/d9/22ce5786ac0c1653ae8b6c23bded02c1686d11f0dbb45b31ce128e0df985/aiohttp-3.14.3.tar.gz", hash = "sha256:9491196535a88924a60afd5b5f434b5b203b6cc616250878dbdb223a8f7844bc", size = 7971213, upload-time = "2026-07-23T01:57:27.037Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/1d/21/151624b51cd92553d95424daf4bf19f19ce9be9002d19253e7e7ce67197b/aiohttp-3.14.1-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:d35143e27778b4bb0fb189562d7f275bff79c62ab8e98459717c0ea617ff2480", size = 757402, upload-time = "2026-06-07T21:06:40.311Z" }, - { url = "https://files.pythonhosted.org/packages/c2/82/280619e0bd7bf2454987e19282616e84762255dd9c8468f62382e8c191f1/aiohttp-3.14.1-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:bcfb80a2cc36fba2534e5e5b5264dc7ae6fcd9bf15256da3e53d2f499e6fa29d", size = 512310, upload-time = "2026-06-07T21:06:42.207Z" }, - { url = "https://files.pythonhosted.org/packages/55/b2/2aac325583aaa1353045f96dffa586d8a34e8322e14a7ba49cffeb103ab4/aiohttp-3.14.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:27fd7c91e51729b4f7e1577865fa6d34c9adccbc39aabe9000285b48af9f0ec2", size = 512448, upload-time = "2026-06-07T21:06:43.813Z" }, - { url = "https://files.pythonhosted.org/packages/8a/72/a60607cb849faa8af8a356c9329ea2eb6f395d49e82cc82ccba1fd8deb8f/aiohttp-3.14.1-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:64c567bf9eaf664280116a8688f63016e6b32db2505908e2bdaca1b6438142f2", size = 1766854, upload-time = "2026-06-07T21:06:45.391Z" }, - { url = "https://files.pythonhosted.org/packages/b5/d3/d9fe1c9ec7557ab4d0d82bebaa728c6418f0b93295ec2f4ab015f7710cc7/aiohttp-3.14.1-cp312-cp312-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:f5e6ff2bdbb8f4cd3fbe41f99e25bbcd58e3bf9f13d3dd31a11e7917251cc77a", size = 1740884, upload-time = "2026-06-07T21:06:47.413Z" }, - { url = "https://files.pythonhosted.org/packages/c1/dc/f2cecfaf9337ba3e63f181500814ff502aa3d00d9c7ec93a9d23d10a27b2/aiohttp-3.14.1-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:2f73e01dc37122325caf079982621262f96d74823c179038a82fddfc50359264", size = 1810034, upload-time = "2026-06-07T21:06:50.165Z" }, - { url = "https://files.pythonhosted.org/packages/66/d7/2ff65c5e65c0d7476daf7e15c032e0805e36811185b9623e3238ad6c763e/aiohttp-3.14.1-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:bb2c0c80d431c0d03f2c7dbf125150fedd4f0de17366a7ca33f7ccb822391842", size = 1904054, upload-time = "2026-06-07T21:06:52.035Z" }, - { url = "https://files.pythonhosted.org/packages/20/9c/d445818389df371f56d141d881153ba23183c4735a03f7356ffb43f7757d/aiohttp-3.14.1-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:3e6fc1a85fa7194a1a7d19f44e8609180f4a8eb5fa4c7ed8b4355f080fad235c", size = 1790278, upload-time = "2026-06-07T21:06:54.049Z" }, - { url = "https://files.pythonhosted.org/packages/4d/aa/bf04cb4d865fc6101c2229a294ad744973b72e513fdc5a6b791e6983d72a/aiohttp-3.14.1-cp312-cp312-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:686b6c0d3911ec387b444ddf5dc62fb7f7c0a7d5186a7861626496a5ab4aff95", size = 1591795, upload-time = "2026-06-07T21:06:55.911Z" }, - { url = "https://files.pythonhosted.org/packages/dc/b4/4dac0038960427ba832f6609dfb4ea5437d7fd80c72001b9e48f834f428b/aiohttp-3.14.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:c6fa4dc7ad6f8109c70bb1499e589f76b0b792baf39f9b017eb92c8a81d0a199", size = 1728397, upload-time = "2026-06-07T21:06:57.777Z" }, - { url = "https://files.pythonhosted.org/packages/2b/f9/7cd4e8ad7aa3b75f17d56bb5498dd604a93d4e6eece822ba0568c413fff0/aiohttp-3.14.1-cp312-cp312-musllinux_1_2_armv7l.whl", hash = "sha256:87a5eea1b2a5e21e1ebdbb33ad4165359189327e63fc4e4894693e7f821ac817", size = 1766504, upload-time = "2026-06-07T21:07:00.009Z" }, - { url = "https://files.pythonhosted.org/packages/f9/df/fc01d9fcad0f73fed3f3d361f1f94f975947b50dff82919f6dc2bf4316cc/aiohttp-3.14.1-cp312-cp312-musllinux_1_2_ppc64le.whl", hash = "sha256:1c1421eb01d4fd608d88cc8290211d177a58532b55ad94076fb349c5bf467f0a", size = 1777806, upload-time = "2026-06-07T21:07:02.064Z" }, - { url = "https://files.pythonhosted.org/packages/41/09/47e2d090bddcc8fb4ccb4c314aadc32d7c5d9bb55f50f6ad1c92fc15d501/aiohttp-3.14.1-cp312-cp312-musllinux_1_2_riscv64.whl", hash = "sha256:34b257ec41345c1e8f2df68fa908a7952f5de932723871eb633ecbbff396c9a4", size = 1580707, upload-time = "2026-06-07T21:07:03.942Z" }, - { url = "https://files.pythonhosted.org/packages/3d/36/f1a4ce904ae0b6930cfe9afc96d0896f7ec1a620c400405d63783bb95a9c/aiohttp-3.14.1-cp312-cp312-musllinux_1_2_s390x.whl", hash = "sha256:de538791a80e5d862addbc183f70f0158ac9b9bb872bb147f1fd2a683691e087", size = 1798121, upload-time = "2026-06-07T21:07:05.987Z" }, - { url = "https://files.pythonhosted.org/packages/70/0a/e0075ce9ca0279ee1d4f0c0b85f54fea02ebc83c3007651a72bece658fec/aiohttp-3.14.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:6f71173be42d3241d428f760122febb748de0623f44308a6f120d0dd9ec572e3", size = 1767580, upload-time = "2026-06-07T21:07:07.873Z" }, - { url = "https://files.pythonhosted.org/packages/3e/61/a0c0a8f327a9c52095cdd8e312391b00d3ed64ab6c72bb5c33d8ec251cf7/aiohttp-3.14.1-cp312-cp312-win32.whl", hash = "sha256:ec8dc383ee57ea3e883477dcca3f11b65d58199f1080acaf4cd6ad9a99698be4", size = 452771, upload-time = "2026-06-07T21:07:09.669Z" }, - { url = "https://files.pythonhosted.org/packages/df/d9/ea367c75f16ac9c6cdc8febb25e8318fa21a2b1bc8d6514d4b2d890bface/aiohttp-3.14.1-cp312-cp312-win_amd64.whl", hash = "sha256:2aa92c87868cd13674989f9ee83e5f9f7ea4237589b728048e1f0c8f6caa3271", size = 479873, upload-time = "2026-06-07T21:07:11.538Z" }, - { url = "https://files.pythonhosted.org/packages/03/64/8d96784a7851156db8a4c6c3f6f91042fdf39fb15a4cc38c8b3c14833c45/aiohttp-3.14.1-cp312-cp312-win_arm64.whl", hash = "sha256:2c840c90759922cb5e6dda94596e079a30fb5a5ba548e7e0dc00574703940847", size = 448073, upload-time = "2026-06-07T21:07:13.637Z" }, + { url = "https://files.pythonhosted.org/packages/18/d4/eb96299230e20acf2efae207cb8d69051f1f68e357e5ea5e479bf6fb097a/aiohttp-3.14.3-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:39aded8c7f3b935b54aab1d8d73c70ec0ee2d3ec3b943e0e86611bc150ba47f5", size = 754690, upload-time = "2026-07-23T01:53:47.332Z" }, + { url = "https://files.pythonhosted.org/packages/88/11/e7a70a209eb9a067c0d3212b518a0134e3484f5178c7533878b6b514d469/aiohttp-3.14.3-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:5bcb6ff3fdab1258a192679ff1a05d44f59626430aa05cd1a9d2447423599228", size = 509484, upload-time = "2026-07-23T01:53:51.159Z" }, + { url = "https://files.pythonhosted.org/packages/30/07/4bbc222cc8dbe31d4c3e8a5baad2286e4d42026ac0c570027b89afce6344/aiohttp-3.14.3-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:617105e2c3018ee38d0c8ce5ee3c84f621a6d8b9f723202aacaff28449ca91ee", size = 511949, upload-time = "2026-07-23T01:53:55.083Z" }, + { url = "https://files.pythonhosted.org/packages/54/b9/42e74c46b7b7c794b995bbc1f573fb48950c38b19d8600c62a6804ee2d67/aiohttp-3.14.3-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:f631fe87a6f30df5fbe6d79640b25e4cffb38c31c7fb6f10871517b84b0f8c1a", size = 1765282, upload-time = "2026-07-23T01:53:59.662Z" }, + { url = "https://files.pythonhosted.org/packages/6b/ed/62bc4d74363ad346d518e0720363a949f63e2e23439a79eb5813d4d29bb3/aiohttp-3.14.3-cp312-cp312-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:a94dbaae5ae27bd849c93570669bff91e0510f33a80805738e3de72a7be0447b", size = 1741511, upload-time = "2026-07-23T01:54:04.063Z" }, + { url = "https://files.pythonhosted.org/packages/d0/9f/181e8a8bc79e47d13c7fc4540bd7a3b729d9505609c61f392a8dd2fbfe55/aiohttp-3.14.3-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:8f2f1c4c032c7cedd7d8da6f54c97b70266c6570c3108d3fdffee7188bb70529", size = 1810680, upload-time = "2026-07-23T01:54:09.882Z" }, + { url = "https://files.pythonhosted.org/packages/5c/9a/dec94d6ad694552fe3424e3f1928d7a606a5d9d9433a04e7ecdd9d38ae7f/aiohttp-3.14.3-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:ea05e1f97ceea523942d9b2a7d7c0359d781d683d6b043f5943a602b14da4787", size = 1905646, upload-time = "2026-07-23T01:54:13.475Z" }, + { url = "https://files.pythonhosted.org/packages/52/b7/7cd31f29d6055bd711ae6e669367fba6f5ae9de463910a793e30556a8db7/aiohttp-3.14.3-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:543906c127fb1d929b95076db19b83fa2d46751006ff1e23b093aa5ac4d8db42", size = 1792122, upload-time = "2026-07-23T01:54:15.752Z" }, + { url = "https://files.pythonhosted.org/packages/66/73/10b1ef93afa61f4963c746257b70ced619cf31a4798671de5fdb2608501d/aiohttp-3.14.3-cp312-cp312-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:0a5ff2dfbb9ce645fa5b8ef3e02c6c0b9cc3f6030ff863d0c51fffc50cb5541b", size = 1591127, upload-time = "2026-07-23T01:54:19.489Z" }, + { url = "https://files.pythonhosted.org/packages/49/ed/3b203fa6de1b338c14acdc06bf6ca9b043b7944f005966958c2ced932cde/aiohttp-3.14.3-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:041badb8f84396357c4d3ad26de6afd7a32b112f43d3c63045c0c8278cfd2043", size = 1725210, upload-time = "2026-07-23T01:54:24.129Z" }, + { url = "https://files.pythonhosted.org/packages/28/b7/1c2aab8c706436dcc28598452488ac9cd7c409da815237c28c27d58993e6/aiohttp-3.14.3-cp312-cp312-musllinux_1_2_armv7l.whl", hash = "sha256:530125ee1163c4219af35dc3aa1206e541e7b31b6efc1a3f93b70a136f65d427", size = 1764848, upload-time = "2026-07-23T01:54:27.973Z" }, + { url = "https://files.pythonhosted.org/packages/54/50/94c28f08b131c4bf10984ea2c7a536c9920608bb2d6e7f95642c30cc87b7/aiohttp-3.14.3-cp312-cp312-musllinux_1_2_ppc64le.whl", hash = "sha256:c8653fd547c93a61aadc612007790f5555cdd18946fa48cf45e26d8ea4ea473d", size = 1777102, upload-time = "2026-07-23T01:54:31.775Z" }, + { url = "https://files.pythonhosted.org/packages/13/d4/e7d09ba7d345fb2d74440fd2fa033c5e079fac05552927705986f41a364f/aiohttp-3.14.3-cp312-cp312-musllinux_1_2_riscv64.whl", hash = "sha256:89176250f686cb9853c0fb7ead90e639e915b84a6f43eedc2a4e7ec21f1037f0", size = 1580205, upload-time = "2026-07-23T01:54:34.518Z" }, + { url = "https://files.pythonhosted.org/packages/a3/84/072a91d68e1e1eb587985b54baab94221277f877e8ef274fc213a0ceae28/aiohttp-3.14.3-cp312-cp312-musllinux_1_2_s390x.whl", hash = "sha256:3a26434dafe408229ff3403458ca58de24fb51936504decac49ce6755f77e59d", size = 1797219, upload-time = "2026-07-23T01:54:36.995Z" }, + { url = "https://files.pythonhosted.org/packages/e0/eb/aad34e897e668424d6e995da5dff8a4a09af93363d3392488772957a63aa/aiohttp-3.14.3-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:d1558173930a5a8d3069cee5c92fc91c87c4dbcb099debbb3622053717145a19", size = 1768629, upload-time = "2026-07-23T01:54:40.103Z" }, + { url = "https://files.pythonhosted.org/packages/b6/2b/6bb88ddba0fecd9122aa3ebcad25996cf6c083a4a7040dbb3a4f97972af6/aiohttp-3.14.3-cp312-cp312-win32.whl", hash = "sha256:16100ad3ab8d649fdfbee87602d9d2dcdca9df0b9eda8a1b5fdc0d41f96da559", size = 451481, upload-time = "2026-07-23T01:54:42.547Z" }, + { url = "https://files.pythonhosted.org/packages/76/9b/f2f8f108da17ecef2cc3efc424e8b7ad3782b1a8360f7b8eae8ced84f6ea/aiohttp-3.14.3-cp312-cp312-win_amd64.whl", hash = "sha256:33a2d7c28d33797a2e99923dffa63f83d908a19b6bf26cfe80fa790aa5e1a75a", size = 476845, upload-time = "2026-07-23T01:54:44.853Z" }, + { url = "https://files.pythonhosted.org/packages/3e/44/28dac80a8941b604f4da10ce21097614ca1bf905ce93dca28d8d7de9c1e7/aiohttp-3.14.3-cp312-cp312-win_arm64.whl", hash = "sha256:362a3fd481769cac1a824514bcd86fda51c65e8fe6e051099e008fddde6db17c", size = 448050, upload-time = "2026-07-23T01:54:47.087Z" }, ] [[package]] @@ -194,16 +194,16 @@ sdist = { url = "https://files.pythonhosted.org/packages/ab/98/d7111245f17935bf7 [[package]] name = "alibabacloud-gpdb20160503" -version = "5.2.0" +version = "5.9.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "alibabacloud-credentials" }, { name = "alibabacloud-tea-openapi" }, { name = "darabonba-core" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/95/ba/606601479707f90138be38493b7b4d8457da10bbc58e84cd000108468a44/alibabacloud_gpdb20160503-5.2.0.tar.gz", hash = "sha256:d8f41bfcdc189f9d0283a87df2c3fa26a27617bc2d604652c7763bf9dd3ba22d", size = 299202, upload-time = "2026-04-02T19:27:25.639Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e2/f2/9a3a5fb89db032bd591be026684aa039ccad8f95af2b23da5f3bf49104c9/alibabacloud_gpdb20160503-5.9.0.tar.gz", hash = "sha256:79ac492c27b7353102de00c258c672e9ad5e626869b4adbbf9c315fe1b397462", size = 340677, upload-time = "2026-07-14T17:33:09.317Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/8f/a3/eee56773d22b8ee4039f2a4754bcf957631302d2e59e5b110cdd768e25ac/alibabacloud_gpdb20160503-5.2.0-py3-none-any.whl", hash = "sha256:b2bad9d2f7e0247985120c25f6cd42e75447fb9157dff817f64eae1734abcbd7", size = 857108, upload-time = "2026-04-02T19:27:24.446Z" }, + { url = "https://files.pythonhosted.org/packages/cc/ae/7d52d8a550f1bf8a4eb3958233f0ceb1bf69aba07b235feac1aba4ca1579/alibabacloud_gpdb20160503-5.9.0-py3-none-any.whl", hash = "sha256:f4318bd5b512e1edea5dfb7de2f985270535d0b2b94f94684680dd99d617b0f5", size = 974562, upload-time = "2026-07-14T17:33:07.942Z" }, ] [[package]] @@ -424,6 +424,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/49/9a/417b3a533e01953a7c618884df2cb05a71e7b68bdbce4fbdb62349d2a2e8/azure_identity-1.25.3-py3-none-any.whl", hash = "sha256:f4d0b956a8146f30333e071374171f3cfa7bdb8073adb8c3814b65567aa7447c", size = 192138, upload-time = "2026-03-13T01:12:22.951Z" }, ] +[[package]] +name = "azure-keyvault-keys" +version = "4.11.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "azure-core" }, + { name = "cryptography" }, + { name = "isodate" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/f4/03/5ce6db28b545427d4ab572f6a4ef2a727b6b4e7bf6941cedddf98822535b/azure_keyvault_keys-4.11.1.tar.gz", hash = "sha256:90caa3a7b2c8f6b53c247ec115cf1c1dad7f107cc3aa9f35aff4838bbce7e562", size = 260915, upload-time = "2026-05-19T20:01:08.041Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/0c/3d/7bed91ae9268cf48124cf6990d8cd2c3daff7a2bb1e91a439d26ee90d705/azure_keyvault_keys-4.11.1-py3-none-any.whl", hash = "sha256:f46cdf6ee7a9baf27f70e6838327032886c8a087041dd56397773b1639da8fe2", size = 200651, upload-time = "2026-05-19T20:01:09.809Z" }, +] + [[package]] name = "azure-storage-blob" version = "12.30.0" @@ -475,7 +490,7 @@ wheels = [ [[package]] name = "bce-python-sdk" -version = "0.9.72" +version = "0.9.76" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "crc32c" }, @@ -483,9 +498,9 @@ dependencies = [ { name = "pycryptodome" }, { name = "six" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/32/bb/1ccb8b28bfa0802356f8588e479adb61cfd2268b37fabc6c8a805d645cd5/bce_python_sdk-0.9.72.tar.gz", hash = "sha256:d9db568698792d74db4245252d98776eae8c9c5225fc0ba86548dfc52d478fcc", size = 302207, upload-time = "2026-06-08T12:10:32.326Z" } +sdist = { url = "https://files.pythonhosted.org/packages/0a/ed/b41906366c0e1e3a1a883e0f1bcfc927403d9697d74b1eb7e89502870218/bce_python_sdk-0.9.76.tar.gz", hash = "sha256:01c630ce8dcbf8be0563d65f18b5201eaac6bd954b3dc776040aec268f81b6ff", size = 317399, upload-time = "2026-07-24T05:12:57.005Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/2d/3a/f84b025ff6c8ec8fe222430cf9201b513dd11ada744417649a64ad295b6d/bce_python_sdk-0.9.72-py3-none-any.whl", hash = "sha256:54a0c121134d6f183f6013d9b33dbf5b6678b0815b6d09616bbceed0202dd797", size = 417800, upload-time = "2026-06-08T12:10:30.556Z" }, + { url = "https://files.pythonhosted.org/packages/b9/06/4d4a26c3dfcdb29599f563329b87eefe55b02b3c2b9d8f8c3aeaa38a33a6/bce_python_sdk-0.9.76-py3-none-any.whl", hash = "sha256:e629181d060f4ed8f29749f47139bf08f1f9b0d8b15daedc956900785587e8d8", size = 435179, upload-time = "2026-07-24T05:12:55.064Z" }, ] [[package]] @@ -598,16 +613,16 @@ wheels = [ [[package]] name = "boto3" -version = "1.43.46" +version = "1.43.56" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "botocore" }, { name = "jmespath" }, { name = "s3transfer" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/f2/e7/976bf3dfe0aa5d7f31bec2f2cf57c79641620c910a39bc843a237aa9592d/boto3-1.43.46.tar.gz", hash = "sha256:66c0d943b049a46a492ec4ec2ebe73c930b1842c7137bee83aad6d93e95d4d96", size = 112654, upload-time = "2026-07-10T19:32:12.498Z" } +sdist = { url = "https://files.pythonhosted.org/packages/53/05/23e1aa8c9e4b0399a61e7fd65c4f9cc0625121f24760e37471f776404abb/boto3-1.43.56.tar.gz", hash = "sha256:57c90df9fb026f2e6ae22530861198130203733c5c9ec4e5cca3a4037f5a8db4", size = 112673, upload-time = "2026-07-24T19:31:48.606Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ef/1d/c52e66ff32ba7911664e6c4c2ac62e1c6d2d1e7550c7ac185d3f4b70a8a4/boto3-1.43.46-py3-none-any.whl", hash = "sha256:69453e2c1bcb9fd9806527ab99950cacfc2826cb0dce9a3a0414d19270c06c3c", size = 140031, upload-time = "2026-07-10T19:32:11.129Z" }, + { url = "https://files.pythonhosted.org/packages/b8/57/3a960c9f581c00f2a591901b46e035ff79ab3956d16607f12306b3b8d483/boto3-1.43.56-py3-none-any.whl", hash = "sha256:feb699d4ab241ef5c1b80bb58277be2aaad365cd4b672d7817e0bc59ee45131b", size = 140026, upload-time = "2026-07-24T19:31:47.155Z" }, ] [[package]] @@ -630,16 +645,16 @@ bedrock-runtime = [ [[package]] name = "botocore" -version = "1.43.46" +version = "1.43.56" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "jmespath" }, { name = "python-dateutil" }, { name = "urllib3" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/7d/f1/1917891851ac5ac09bb9f4862b8fc9252a009d7c24e8688bb67e4383d9e7/botocore-1.43.46.tar.gz", hash = "sha256:59f2e1ac3cdc66d191cae91c0804bc41847ce817dc8147cf43eaada8f76a5533", size = 15694635, upload-time = "2026-07-10T19:32:00.437Z" } +sdist = { url = "https://files.pythonhosted.org/packages/b9/cc/7f84a5d3071fe878380e9f610ab36ca87b8cbbc4aa81ba2727f90e1f3ea3/botocore-1.43.56.tar.gz", hash = "sha256:6c01f85f0ff9863076f4c761e74ee3aa96c5ccc1ad09fc1efd62ef8f2d22bf57", size = 15733117, upload-time = "2026-07-24T19:31:38.125Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/0e/f2/4bd8f2f419088feb3ce55f0ca91040ff902f402edfd197450b20a2e1d533/botocore-1.43.46-py3-none-any.whl", hash = "sha256:cb673891e623ae6e6a1bf24d94ef169504f3eb02584adb5d5bee2f6aae819b60", size = 15380350, upload-time = "2026-07-10T19:31:57.616Z" }, + { url = "https://files.pythonhosted.org/packages/5c/cd/86fe9e659e9699f62f8dd5ecd8c6725474334b23cab8aa71d82b5f56f1a4/botocore-1.43.56-py3-none-any.whl", hash = "sha256:aafc741f1b10f6fd63253eaf6ea029680c1ff436d87e1b8969d62aefa0c76976", size = 15418773, upload-time = "2026-07-24T19:31:34.758Z" }, ] [[package]] @@ -1068,19 +1083,19 @@ wheels = [ [[package]] name = "couchbase" -version = "4.6.0" +version = "4.6.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/8d/be/1e6974158348dfa634ebbc32b76448f84945e15494852e0cea85607825b5/couchbase-4.6.0.tar.gz", hash = "sha256:61229d6112597f35f6aca687c255e12f495bde9051cd36063b4fddd532ab8f7f", size = 6697937, upload-time = "2026-03-31T23:29:50.602Z" } +sdist = { url = "https://files.pythonhosted.org/packages/93/84/b5228cb6fe74c1c2dcd1bf7fddf451a4d608d577d9e2eee485e718191235/couchbase-4.6.2.tar.gz", hash = "sha256:c660549f8c87354837364aa48b84268a4778fd4a0c0a9db41cbf5eab1b11d70c", size = 6710448, upload-time = "2026-06-18T02:50:30.275Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/84/dc/bea38235bfabd4fcf3d11e05955e38311869f173328475c369199a6b076b/couchbase-4.6.0-cp312-cp312-macosx_10_15_x86_64.whl", hash = "sha256:8d1244fd0581cc23aaf2fa3148e9c2d8cfba1d5489c123ee6bf975624d861f7a", size = 5521692, upload-time = "2026-03-31T23:29:07.933Z" }, - { url = "https://files.pythonhosted.org/packages/d1/18/cd1c751005cb67d3e2b090cd11626b8922b9d6a882516e57c1a3aedeed18/couchbase-4.6.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:8efa57a86e35ceb7ae249cfa192e3f2c32a4a5b37098830196d3936994d55a67", size = 4667116, upload-time = "2026-03-31T23:29:10.706Z" }, - { url = "https://files.pythonhosted.org/packages/64/e9/1212bd59347e1cecdb02c6735704650e25f9195b634bf8df73d3382ffa14/couchbase-4.6.0-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:7106e334acdacab64ae3530a181b8fabf0a1b91e7a1a1e41e259f995bdc78330", size = 5511873, upload-time = "2026-03-31T23:29:13.414Z" }, - { url = "https://files.pythonhosted.org/packages/86/a3/f676ee10f8ea2370700c1c4d03cbe8c3064a3e0cf887941a39333f3bdd97/couchbase-4.6.0-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:c84e625f3e2ac895fafd2053fa50af2fbb63ab3cdd812eff2bc4171d9f934bde", size = 5782875, upload-time = "2026-03-31T23:29:16.258Z" }, - { url = "https://files.pythonhosted.org/packages/c5/34/45d167bc18d5d91b9ff95dcd4e24df60d424567611d48191a29bf19fdbc8/couchbase-4.6.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:a2619c966b308948900e51f1e4e1488e09ad50b119b1d5c31b697870aa82a6ce", size = 7234591, upload-time = "2026-03-31T23:29:19.148Z" }, - { url = "https://files.pythonhosted.org/packages/41/1f/cc4d1503463cf243959532424a30e79f34aadafde5bcb21754b19b2b9dde/couchbase-4.6.0-cp312-cp312-win_amd64.whl", hash = "sha256:f64a017416958f10a07312a6d39c9b362827854de173fdef9bffdac71c8f3345", size = 4517477, upload-time = "2026-03-31T23:29:21.955Z" }, + { url = "https://files.pythonhosted.org/packages/ff/c8/886d80b17a1f265930dd56bd80e9499ccab28798286bf74e02ce272318c3/couchbase-4.6.2-cp312-cp312-macosx_10_15_x86_64.whl", hash = "sha256:1dfad6adb0e0985a4c7206f9b8ae46c037c059f1b9df2c2bcbdce6a436847aad", size = 5642182, upload-time = "2026-06-18T02:49:32.66Z" }, + { url = "https://files.pythonhosted.org/packages/15/22/40c3da8a7efb404113b6a9bb80fc3359639081409e6538dd5e12f538785f/couchbase-4.6.2-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:51803b25b84f2b0ac4b4676d8c5180c9ea783e09be7ea1beff8a2cdd384693b6", size = 4776221, upload-time = "2026-06-18T02:49:35.459Z" }, + { url = "https://files.pythonhosted.org/packages/1e/89/3eb91bdffb273bdc74998671a53c83955fe65ec412d99374b2af9c7a4813/couchbase-4.6.2-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:390707f14c7ed8541758bff2addef90e6a173e32676e9b54d9b4b122f349deef", size = 5609885, upload-time = "2026-06-18T02:49:38.459Z" }, + { url = "https://files.pythonhosted.org/packages/39/7a/cad8c0ed349ae83b711ee879e0ee270c2a043c4a678b1a497ca0af340bc0/couchbase-4.6.2-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:84e725336b4de8f2d9b653a0fd650e1270f316747be272f420712c8545dcd80d", size = 5883188, upload-time = "2026-06-18T02:49:41.68Z" }, + { url = "https://files.pythonhosted.org/packages/7a/e9/bea010f955fd6eb5a81af3c17ce2eb637cbf32bb331ccad2c75d1242b241/couchbase-4.6.2-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:443d9404103d3bbf6c6997438641abde36c3f6c1e4fad5589f240cecb53876aa", size = 7377393, upload-time = "2026-06-18T02:49:45.29Z" }, + { url = "https://files.pythonhosted.org/packages/64/6e/9fe84d86846770f4a4e787a0f3ced32351a9c79fdf5493bc4be30cf52446/couchbase-4.6.2-cp312-cp312-win_amd64.whl", hash = "sha256:26e2defa8540997466f8cd3fab649646451a45e068c91561e7c111d8fe7d6f7c", size = 4547006, upload-time = "2026-06-18T02:49:48.881Z" }, ] [[package]] @@ -1144,39 +1159,39 @@ wheels = [ [[package]] name = "cryptography" -version = "49.0.0" +version = "50.0.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "cffi", marker = "platform_python_implementation != 'PyPy'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1f/99/d1c90d6041656cc6ee229dc99cd67fd0cd5aec3c5f7d72fffc27cc750054/cryptography-49.0.0.tar.gz", hash = "sha256:f89660a348f4f78a92366240a61404e337586ef7f5909a2fef59ca88ef505493", size = 854345, upload-time = "2026-06-12T20:02:30.512Z" } +sdist = { url = "https://files.pythonhosted.org/packages/de/41/6cbdcf9142d00fe82836fbb51e503e58088575cf7a0fe1dbff6695bf0840/cryptography-50.0.0.tar.gz", hash = "sha256:eeac2acb5a20ed25e0ad6d1df9891a520b78b404266b6d11778f25d5d691a6c9", size = 880201, upload-time = "2026-07-31T14:25:10.11Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/9b/22/adf66990e63584a68dfb50c24f48a125c07b1699899381c8151e63ed458c/cryptography-49.0.0-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:966fe0e9c67490071f14c0d2b1cb2dfb3023c5ce39457343931415f08382f2db", size = 4032100, upload-time = "2026-06-12T20:02:32.143Z" }, - { url = "https://files.pythonhosted.org/packages/09/41/3797cfaf69cae04a13ee78ebd83f0678d9c02b4779d21ce24445326f1a69/cryptography-49.0.0-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:36d1709f992593689b45bda411498d62c6e365f2ca00b84657d4dadd24de16db", size = 4692978, upload-time = "2026-06-12T20:01:21.305Z" }, - { url = "https://files.pythonhosted.org/packages/e6/8b/43011f7ebe515a8aa20d61f290a326cd890c2e738e16e59eaff8d9c3a412/cryptography-49.0.0-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:0e959b578856a3924bc0cbb710fc12c387b9412a951389f3ca61704a9e25f325", size = 4716422, upload-time = "2026-06-12T20:01:48.566Z" }, - { url = "https://files.pythonhosted.org/packages/4a/91/01ce7303a4579e6d3a6abef01bd322848e9ea7a219adcabc5048b9033571/cryptography-49.0.0-cp311-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:53ecee2e23f7169b6117e99fc8a944e5e50f79e69758a83b52a00cb98ab2b2d2", size = 4700503, upload-time = "2026-06-12T20:02:47.091Z" }, - { url = "https://files.pythonhosted.org/packages/62/99/a2c95cf8293f07491e9e27c20cc4dcd18176d944e674679adeb1d0173fd6/cryptography-49.0.0-cp311-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:2eda353d8a27bcbcaa4cbed18994a74ab4d19a2ca897db188ea269ab9b71419b", size = 5309779, upload-time = "2026-06-12T20:02:08.987Z" }, - { url = "https://files.pythonhosted.org/packages/20/2c/0622f20ff02b2ef32558733443805dc82fd4c275be01b2d19d14676f3a1b/cryptography-49.0.0-cp311-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:2afe9051da7ae7bd5905da5a949280c7d2bb75682e188f650a9d0f2756b834c6", size = 4749683, upload-time = "2026-06-12T20:02:03.335Z" }, - { url = "https://files.pythonhosted.org/packages/a3/5b/c5246635d5fd3b64e0d45ae10e99fd32fe9676a79915ccfe5a61ba9af1a5/cryptography-49.0.0-cp311-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:0b82e28ee398a386f0807bba7884d30f25218855690f45115831bcce5d90822c", size = 4337874, upload-time = "2026-06-12T20:02:54.323Z" }, - { url = "https://files.pythonhosted.org/packages/6d/88/05563c7fe2e914e87d1a536d06fe83e66b4e1d95cb593e05aea375531da8/cryptography-49.0.0-cp311-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:ccac2bfebc306b862133e3bb71f3f6ee8bb525240089b2d952e4144b3a6d5da7", size = 4700283, upload-time = "2026-06-12T20:01:34.822Z" }, - { url = "https://files.pythonhosted.org/packages/c4/b6/d7696e4e890d6ae1469935164c9e5215c557671cb78d6e3f458ccceaa632/cryptography-49.0.0-cp311-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:d0527ce944105f257f605a827d6ebead966c752038b6e8656abb9c5edee6fc68", size = 5265844, upload-time = "2026-06-12T20:01:24.09Z" }, - { url = "https://files.pythonhosted.org/packages/a9/3c/f3ad17eecc1a57b0ba236dc01f90e783c51f4a2f35f64777cc4f47a184b2/cryptography-49.0.0-cp311-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:cbc77da8c523d5abd028635ba850a6966fcee2c82e2bf65a41d1d8afe0f98be9", size = 4749290, upload-time = "2026-06-12T20:01:30.848Z" }, - { url = "https://files.pythonhosted.org/packages/4f/01/339573cf1023163a400b0b5d16f6d507de413b9f60be6fd1b77feeaf6737/cryptography-49.0.0-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:b87e65d263b3e5d3bb92a57e2a6638e2f31110fa7aa890c7b2dbba42248d0a3f", size = 4834612, upload-time = "2026-06-12T20:01:29.246Z" }, - { url = "https://files.pythonhosted.org/packages/71/fd/577302e213a1be9468f92d1afef66fcf1ef83d516819d9992ca547f592bd/cryptography-49.0.0-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:66ec79c3904820572d7e987abdf304281f141d37ad9a489b8e97066e7b9b6459", size = 4980804, upload-time = "2026-06-12T20:01:42.853Z" }, - { url = "https://files.pythonhosted.org/packages/1f/09/f42b1d190c5ba75f72062a387f8030d1d75f6ab035788f1d9c4b01de6525/cryptography-49.0.0-cp311-abi3-win_amd64.whl", hash = "sha256:e5dfc1e64de5677cec922ffa8da89c546d0415bf6efdf081842e5d44c84e1f0e", size = 3810026, upload-time = "2026-06-12T20:02:39.262Z" }, - { url = "https://files.pythonhosted.org/packages/19/2a/5bb823f5bedcf80718cea7fbc95ec5515cca3769633c4b01a32be7f30e7c/cryptography-49.0.0-cp39-abi3-macosx_11_0_arm64.whl", hash = "sha256:ec5e529fb80935c94fe7b729f9972b50e351a0e6b50aa294fd5cabb109fcc29a", size = 4025947, upload-time = "2026-06-12T20:01:25.745Z" }, - { url = "https://files.pythonhosted.org/packages/3d/df/40577043ca124e17012f408ddddaeb213b856336ac82ddb3bc915f39e29f/cryptography-49.0.0-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:f78ff2c9ed8dc2d036b0f4d640e22522213d047c1b14e61205a7e55c80a494d4", size = 4692429, upload-time = "2026-06-12T20:01:53.628Z" }, - { url = "https://files.pythonhosted.org/packages/2c/99/2d13299eb3dd27b02dcfaafcc91d6b5cb3329f7cbd6d8f51921acd566c1a/cryptography-49.0.0-cp39-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:35b151772baff2c74cba7fa290ceaff4c3b11c0c881eb93eb5dbc05a7cfbba18", size = 4700968, upload-time = "2026-06-12T20:02:45.383Z" }, - { url = "https://files.pythonhosted.org/packages/a5/4d/9c0cd02f95e2602dd5e563da149ee0830abef3537be8b34dc56281ebe27a/cryptography-49.0.0-cp39-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:0f21641cf4b30fca7aee061ced0ec7ad7b073518088b7c9969a297c0ae796c69", size = 4697758, upload-time = "2026-06-12T20:01:41.13Z" }, - { url = "https://files.pythonhosted.org/packages/24/01/186c825898477d77e2324d5360fefe622ff1d8d1963ec0554e2cada8ec77/cryptography-49.0.0-cp39-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:9e82dcc8e56052715fb18b2429e3bca4823b1629136a2084fc45a9a5cecb9b64", size = 5298863, upload-time = "2026-06-12T20:02:24.579Z" }, - { url = "https://files.pythonhosted.org/packages/b8/7b/62cbbab75d0659865bf0273790031544a0b16c8072d258f9428dcd8190dc/cryptography-49.0.0-cp39-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:6f2debedf9ca60cf1d5bd466475638af5130f89965605cd818484d19987d3a21", size = 4735983, upload-time = "2026-06-12T20:01:50.14Z" }, - { url = "https://files.pythonhosted.org/packages/6c/72/3e798c064bc39e471008075d0f9bc9daf77a80879c092e4a8e170c585ed4/cryptography-49.0.0-cp39-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:8c25ceb16df5b9435f3f6a9829204985b0e0cbee3b48aacd432c7d2c850b44d9", size = 4334173, upload-time = "2026-06-12T20:01:44.743Z" }, - { url = "https://files.pythonhosted.org/packages/f0/ee/6fca21d1ac73e06f8bef71940abfd4d2f6472b4bca284d770f32bd4086f6/cryptography-49.0.0-cp39-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:28d8b15e6275f12c8a207dc309dfa957903c927d08d0cc937ee3f63f200693cc", size = 4697298, upload-time = "2026-06-12T20:02:20.918Z" }, - { url = "https://files.pythonhosted.org/packages/67/d0/a5fcd3515f0bae49a7b6d0413cc1bdccdcc1fc0047037a0d480642cdc5d6/cryptography-49.0.0-cp39-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:6fc361c34fb6aac015ce19435876635e5c6d21db31998b0920f675f131e043b8", size = 5254338, upload-time = "2026-06-12T20:02:22.737Z" }, - { url = "https://files.pythonhosted.org/packages/a0/84/84fe36f19caf857d61cb7fc9c63035a47ffabd84ea12d1d393148efa3615/cryptography-49.0.0-cp39-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:2400ef9c9e2299a25614eb1dea3db54a69b1349efd043bfac9c67630d136df36", size = 4735650, upload-time = "2026-06-12T20:02:41.389Z" }, - { url = "https://files.pythonhosted.org/packages/6c/a0/db537264e234f7273a73ec020873d6d6b39dfd8a53db78b550ca8320440e/cryptography-49.0.0-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:67e1d20ad9ef3a563c59ef22e7a8a0b8210bd26604369ea4a30a7c66aefe504e", size = 4834820, upload-time = "2026-06-12T20:01:51.847Z" }, - { url = "https://files.pythonhosted.org/packages/93/77/8df9eb486495979bccecd1062e2eaf435250e84437040295b57d09048b0b/cryptography-49.0.0-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:42b0684e0e40cf26122427802486f6d93aea593612603a94fbf260c7eb1e9c1b", size = 4967968, upload-time = "2026-06-12T20:02:12.524Z" }, - { url = "https://files.pythonhosted.org/packages/c2/e6/f60198ea8d9dfa15fff9ed4ca02ce362f6eadd9ba757dcc50634c4257b63/cryptography-49.0.0-cp39-abi3-win_amd64.whl", hash = "sha256:026ac7423e6fa66872d3bf889be5974507da3944f866f704fa200eadacd00001", size = 3785547, upload-time = "2026-06-12T20:02:26.847Z" }, + { url = "https://files.pythonhosted.org/packages/c5/5c/59086b4aac5e879d38ddbcf74e4be7ade89cebc3eb199a55da998c3bb46a/cryptography-50.0.0-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:031e2d5dd4bb9caa3ca9c82e5a197fd8ae680232cee62603d1a813f3f07e3d03", size = 4001252, upload-time = "2026-07-31T14:23:33.331Z" }, + { url = "https://files.pythonhosted.org/packages/57/ef/8f2df13c7216bcad3e1c74e07f6e193d93e998e114f524a53877c9af27ad/cryptography-50.0.0-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:fd9192b7b70c573d7f214eb1ae35e00d359f6f5e4b27c7e21e30de1fc6204645", size = 4719554, upload-time = "2026-07-31T14:23:35.611Z" }, + { url = "https://files.pythonhosted.org/packages/d9/41/029086c34d91052fc3b88bcc8056f709a7c915c7a23b235a54eb800b1c97/cryptography-50.0.0-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:06a32a980526a6ab9a4b9bf8f7385800791e2bb960903cb6b530e4817509a3b7", size = 4702130, upload-time = "2026-07-31T14:23:37.635Z" }, + { url = "https://files.pythonhosted.org/packages/7d/ff/b6ce0954962e7f7b969f850a883744197bb3910bdfd7b6da162eab7d9f68/cryptography-50.0.0-cp311-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:a1b30560f2acc95aa8b2e06e716a13dbfc97314747b80d9707e307f77b40d6b3", size = 4725244, upload-time = "2026-07-31T14:23:39.471Z" }, + { url = "https://files.pythonhosted.org/packages/06/1e/63a1027cb7fec360a182208e1b7767d5aa1fe57be3d6aa856e69a321edc0/cryptography-50.0.0-cp311-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:8d89f3976b10b4ce31118de72329025f70d2c6ead14a8217c5514dd2c6d5a78f", size = 5342265, upload-time = "2026-07-31T14:23:41.286Z" }, + { url = "https://files.pythonhosted.org/packages/6b/72/a1116d683a6d7ece94590013882515de087edf9ef0e6292aae615a44df73/cryptography-50.0.0-cp311-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:b42a28c1844fd9de8f3f7d540e36b66f3a9c83fceac7170ebc7a6a19edd9dcae", size = 4734609, upload-time = "2026-07-31T14:23:43.139Z" }, + { url = "https://files.pythonhosted.org/packages/15/37/36a9c479bbe49acea2636c7fd3360d20f7b7e079c300352011c44850b181/cryptography-50.0.0-cp311-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:900131fafd8aead39ac7dd3a7e833be754c17a95cfd91221636949fe4eb0aa8a", size = 4356517, upload-time = "2026-07-31T14:23:44.939Z" }, + { url = "https://files.pythonhosted.org/packages/32/98/8a151d64367204cbc63ec65d37502f1d9c53cf4bfc6ec3c532614dbec60d/cryptography-50.0.0-cp311-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:07949c449a1abcf60d1ee6e88956d89404c7df3c8258f46589e912988e551987", size = 4724529, upload-time = "2026-07-31T14:23:46.93Z" }, + { url = "https://files.pythonhosted.org/packages/22/f6/ec13b470172126464a86bf54d2294a46d29837fc51ba3e45d4047946fb5e/cryptography-50.0.0-cp311-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:f89831ef99dd7dd169ab06d63a831adb9e20a87aac6d380266bbda5823349169", size = 5299852, upload-time = "2026-07-31T14:23:48.851Z" }, + { url = "https://files.pythonhosted.org/packages/da/3a/f05e32c99d440c9bb891ea0e36c9091891e36be5a9a87ab2ee6ea20729f6/cryptography-50.0.0-cp311-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:82148ec5bddac30b51a5b3c1945075f896fa022cb93f8e4a01e9f6ee95292c5f", size = 4734462, upload-time = "2026-07-31T14:23:50.861Z" }, + { url = "https://files.pythonhosted.org/packages/ca/dc/bd72b26be8953f80625f63151efd38eee71c76ca6cf591c08ff34615a79e/cryptography-50.0.0-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:1489e263a8048bb8b6a8bac662eb2d402ea5d2b7b4699b72f385f1e2772db105", size = 4852708, upload-time = "2026-07-31T14:23:52.715Z" }, + { url = "https://files.pythonhosted.org/packages/27/20/c930314a2ab476d15dec966ec87e2e9637bb02b06106b12c0396c57bb603/cryptography-50.0.0-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:7cec5b856506da6defb290f30c9ee687d5f5e8cb0bd3f6459dde43b0b4fa40ef", size = 5004179, upload-time = "2026-07-31T14:23:54.887Z" }, + { url = "https://files.pythonhosted.org/packages/32/2e/c9db68a0c4bfa28e310707527c0ee3a2bd254104d2e02e68f368e197aa4c/cryptography-50.0.0-cp311-abi3-win_amd64.whl", hash = "sha256:bd1c592e4d5974f0d08d4888e432157adba757c66da0246918e43677fafa2d30", size = 3840395, upload-time = "2026-07-31T14:23:56.677Z" }, + { url = "https://files.pythonhosted.org/packages/03/37/73d005be173aff344af30e9fd2a576575cb2391a7101d9cd3842e1fa8cce/cryptography-50.0.0-cp39-abi3-macosx_11_0_arm64.whl", hash = "sha256:ccdc4a71a4dabae05de219404f9f4abc38e3b58422177ff93d0da05967dafa07", size = 4036009, upload-time = "2026-07-31T14:24:24.122Z" }, + { url = "https://files.pythonhosted.org/packages/ff/c6/7a6202a534e32103a285b7834a120869557fe198d51d7cfe59754c8bda9c/cryptography-50.0.0-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:910e1d2668e7de9648f2bcee30e180db2a6b15c30f887d7c4c93ddf96e3992e3", size = 4745252, upload-time = "2026-07-31T14:24:26.118Z" }, + { url = "https://files.pythonhosted.org/packages/85/4f/0fa8c2f4428198f15d9ff8d63400e27afbf94ce833f6108da1eb3753f945/cryptography-50.0.0-cp39-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:a91296cb61e8df6f86d0c19cc4068228da256bf59bf86049fbd821084565327f", size = 4728939, upload-time = "2026-07-31T14:24:27.994Z" }, + { url = "https://files.pythonhosted.org/packages/d1/63/54dd723490ba2dc09b299682c10b38db38f159728bcaae8c591b8af2f22d/cryptography-50.0.0-cp39-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:e722f16708d854fe924790e051061f6704a472c3bac347b6fd88033ea8dd0dc5", size = 4748483, upload-time = "2026-07-31T14:24:30.254Z" }, + { url = "https://files.pythonhosted.org/packages/1d/dd/7c77d26285cc7f6991efce64a0f5b4f9383bfa5dd8c5033003eaf7db4cdb/cryptography-50.0.0-cp39-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:d764dcf130c428ef66786f866dd750f53182bc608813489915e9fc106bb0c82f", size = 5367599, upload-time = "2026-07-31T14:24:32.457Z" }, + { url = "https://files.pythonhosted.org/packages/46/c9/f60aed34c013f317f92817b6c171c2d22a78270fa41109bd4b08af26b194/cryptography-50.0.0-cp39-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:105110f43a471dbd0060b9c9516cb8a6a79233631a04cc2ba16f28323ac6e025", size = 4762647, upload-time = "2026-07-31T14:24:34.599Z" }, + { url = "https://files.pythonhosted.org/packages/be/f3/f9a0173b139372c3a48ed98154b45cc6b9de17c789d5ab552e621c293609/cryptography-50.0.0-cp39-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:828743d939e9629bc267b8e2d08d8bb67cd4319c771a33d4b18b22dd8fb7440a", size = 4385197, upload-time = "2026-07-31T14:24:36.647Z" }, + { url = "https://files.pythonhosted.org/packages/d8/36/83bb81f6e569bc38e1e4a7bc80f29b46bb9601920bc455fc8e888f5d5742/cryptography-50.0.0-cp39-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:2a8183b489dc1f7f80f135780fadc1108f14b31b8a40411c7a5b17425f65f28b", size = 4748095, upload-time = "2026-07-31T14:24:39.493Z" }, + { url = "https://files.pythonhosted.org/packages/6b/16/d3008eff98c764979865834c3d386d4fd041b5f52e7f34fc29ac1a5eb515/cryptography-50.0.0-cp39-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:6e7d61120573a7f2cd94cc095f9e81f6967c61ccdf194285aa143ecec8e0b708", size = 5325948, upload-time = "2026-07-31T14:24:41.556Z" }, + { url = "https://files.pythonhosted.org/packages/9c/f8/d97f9603efda3888187bfdb893f26c41be4735c10631d05d284ee6b047c4/cryptography-50.0.0-cp39-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:37fdb0d0111f1e2ff07139dfb79f1b49531f8e213c46f1163dd7642979b58c47", size = 4762400, upload-time = "2026-07-31T14:24:43.636Z" }, + { url = "https://files.pythonhosted.org/packages/64/a2/4615c8f7d81a00b1d6e6afe19f694e1543582349fb5f4076f6cb5dc36485/cryptography-50.0.0-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:c87f62a3d3b9888ed0fdde100ec06aa61ca9cd44bad9057d1dff9a516b5f5bb9", size = 4878208, upload-time = "2026-07-31T14:24:45.522Z" }, + { url = "https://files.pythonhosted.org/packages/d2/1a/efcfb02f91407149a0dacffffab791f7e19bf6385f63b3666dc8b5e5c9c8/cryptography-50.0.0-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:65c2c3add92b45fd0709db8594536aea39c2a67af0e27ffcf049c498501140b7", size = 5037050, upload-time = "2026-07-31T14:24:47.697Z" }, + { url = "https://files.pythonhosted.org/packages/57/30/4a22984d4f1bdfb8c054f07a92bc176b97a3134cc1d6c4b3bffb1f3688b4/cryptography-50.0.0-cp39-abi3-win_amd64.whl", hash = "sha256:d24fead1d4d076e1bfb006dcec392074a3cd8d7b4fc8a595aa64073b2b7a96ba", size = 3874135, upload-time = "2026-07-31T14:24:50.085Z" }, ] [[package]] @@ -1281,7 +1296,7 @@ wheels = [ [[package]] name = "dify-agent" -version = "1.16.0" +version = "1.16.1" source = { editable = "../dify-agent" } dependencies = [ { name = "httpx" }, @@ -1293,15 +1308,14 @@ dependencies = [ [package.metadata] requires-dist = [ + { name = "e2b", marker = "extra == 'server'", specifier = ">=2.34.0,<3.0.0" }, { name = "fastapi", marker = "extra == 'server'", specifier = "==0.136.0" }, { name = "graphon", marker = "extra == 'server'", specifier = "==0.5.2" }, - { name = "grpclib", extras = ["protobuf"], marker = "extra == 'grpc'", specifier = ">=0.4.9,<0.5.0" }, { name = "httpx", specifier = "==0.28.1" }, { name = "httpx2", specifier = ">=2.5.0,<3.0.0" }, { name = "jsonschema", marker = "extra == 'server'", specifier = ">=4.23.0,<5.0.0" }, { name = "jwcrypto", marker = "extra == 'server'", specifier = ">=1.5.6,<2" }, { name = "logfire", extras = ["fastapi", "httpx", "redis"], marker = "extra == 'server'", specifier = ">=4.37.0,<5.0.0" }, - { name = "protobuf", marker = "extra == 'grpc'", specifier = ">=6.33.5,<7.0.0" }, { name = "pydantic", specifier = ">=2.12.5,<2.13" }, { name = "pydantic-ai-slim", specifier = ">=1.102.0,<2.0.0" }, { name = "pydantic-ai-slim", extras = ["anthropic", "google", "openai"], marker = "extra == 'server'", specifier = ">=1.85.1,<2.0.0" }, @@ -1310,13 +1324,12 @@ requires-dist = [ { name = "typing-extensions", specifier = ">=4.12.2,<5.0.0" }, { name = "uvicorn", extras = ["standard"], marker = "extra == 'server'", specifier = "==0.46.0" }, ] -provides-extras = ["grpc", "server"] +provides-extras = ["server"] [package.metadata.requires-dev] dev = [ { name = "basedpyright", specifier = ">=1.39.3" }, { name = "coverage", extras = ["toml"], specifier = ">=7.10.7" }, - { name = "grpcio-tools", specifier = ">=1.81.0,<2.0.0" }, { name = "pytest", specifier = ">=9.0.3" }, { name = "pytest-examples", specifier = ">=0.0.18" }, { name = "pytest-mock", specifier = ">=3.14.0" }, @@ -1331,7 +1344,7 @@ docs = [ [[package]] name = "dify-api" -version = "1.16.0" +version = "1.16.1" source = { virtual = "." } dependencies = [ { name = "aliyun-log-python-sdk" }, @@ -1402,9 +1415,11 @@ dev = [ { name = "testcontainers" }, { name = "types-aiofiles" }, { name = "types-beautifulsoup4" }, + { name = "types-bleach" }, { name = "types-cachetools" }, { name = "types-cffi" }, { name = "types-colorama" }, + { name = "types-croniter" }, { name = "types-defusedxml" }, { name = "types-deprecated" }, { name = "types-docutils" }, @@ -1428,6 +1443,7 @@ dev = [ { name = "types-pyopenssl" }, { name = "types-python-dateutil" }, { name = "types-python-http-client" }, + { name = "types-pytz" }, { name = "types-pywin32" }, { name = "types-pyyaml" }, { name = "types-redis" }, @@ -1441,6 +1457,9 @@ dev = [ { name = "types-ujson" }, { name = "xinference-client" }, ] +kms = [ + { name = "azure-keyvault-keys" }, +] storage = [ { name = "azure-storage-blob" }, { name = "bce-python-sdk" }, @@ -1621,7 +1640,7 @@ requires-dist = [ { name = "aliyun-log-python-sdk", specifier = "==0.9.44" }, { name = "azure-identity", specifier = ">=1.25.3,<2.0.0" }, { name = "bleach", specifier = ">=6.4.0,<7.0.0" }, - { name = "boto3", specifier = ">=1.43.46,<2.0.0" }, + { name = "boto3", specifier = ">=1.43.56,<2.0.0" }, { name = "celery", specifier = ">=5.6.3,<6.0.0" }, { name = "croniter", specifier = ">=6.2.2,<7.0.0" }, { name = "dify-agent", editable = "../dify-agent" }, @@ -1638,17 +1657,17 @@ requires-dist = [ { name = "gmpy2", specifier = ">=2.3.0,<3.0.0" }, { name = "google-api-python-client", specifier = ">=2.198.0,<3.0.0" }, { name = "google-cloud-aiplatform", specifier = ">=1.160.0,<2.0.0" }, - { name = "graphon", specifier = "==0.6.0" }, + { name = "graphon", specifier = "==0.7.0" }, { name = "gunicorn", specifier = ">=26.0.0,<27.0.0" }, { name = "httpx", extras = ["socks"], specifier = "==0.28.1" }, { name = "httpx-sse", specifier = "==0.4.3" }, { name = "json-repair", specifier = "==0.60.1" }, - { name = "opentelemetry-distro", specifier = "==0.62b1" }, - { name = "opentelemetry-instrumentation-celery", specifier = "==0.62b1" }, - { name = "opentelemetry-instrumentation-flask", specifier = "==0.62b1" }, - { name = "opentelemetry-instrumentation-httpx", specifier = "==0.62b1" }, - { name = "opentelemetry-instrumentation-redis", specifier = "==0.62b1" }, - { name = "opentelemetry-instrumentation-sqlalchemy", specifier = "==0.62b1" }, + { name = "opentelemetry-distro", specifier = "==0.65b0" }, + { name = "opentelemetry-instrumentation-celery", specifier = "==0.65b0" }, + { name = "opentelemetry-instrumentation-flask", specifier = "==0.65b0" }, + { name = "opentelemetry-instrumentation-httpx", specifier = "==0.65b0" }, + { name = "opentelemetry-instrumentation-redis", specifier = "==0.65b0" }, + { name = "opentelemetry-instrumentation-sqlalchemy", specifier = "==0.65b0" }, { name = "opentelemetry-propagator-b3", specifier = ">=1.41.1,<2.0.0" }, { name = "psycogreen", specifier = ">=1.0.2,<2.0.0" }, { name = "psycopg2-binary", specifier = ">=2.9.12,<3.0.0" }, @@ -1673,7 +1692,7 @@ dev = [ { name = "lxml-stubs", specifier = ">=0.5.1" }, { name = "mypy", specifier = ">=1.20.2" }, { name = "pandas-stubs", specifier = ">=3.0.0" }, - { name = "pyrefly", specifier = ">=1.0.0" }, + { name = "pyrefly", specifier = ">=1.2.0" }, { name = "pytest", specifier = ">=9.0.3" }, { name = "pytest-benchmark", specifier = ">=5.2.3" }, { name = "pytest-cov", specifier = ">=7.1.0" }, @@ -1686,9 +1705,11 @@ dev = [ { name = "testcontainers", specifier = ">=4.14.2" }, { name = "types-aiofiles", specifier = ">=25.1.0" }, { name = "types-beautifulsoup4", specifier = ">=4.12.0" }, + { name = "types-bleach", specifier = ">=6.4.0.20260728" }, { name = "types-cachetools", specifier = ">=7.0.0.20260503" }, { name = "types-cffi", specifier = ">=2.0.0.20260429" }, { name = "types-colorama", specifier = ">=0.4.15" }, + { name = "types-croniter", specifier = ">=6.2.4.20260711" }, { name = "types-defusedxml", specifier = ">=0.7.0" }, { name = "types-deprecated", specifier = ">=1.3.1" }, { name = "types-docutils", specifier = ">=0.22.3" }, @@ -1712,6 +1733,7 @@ dev = [ { name = "types-pyopenssl", specifier = ">=24.1.0" }, { name = "types-python-dateutil", specifier = ">=2.9.0" }, { name = "types-python-http-client", specifier = ">=3.3.7.20260408" }, + { name = "types-pytz", specifier = ">=2026.3.1.20260727" }, { name = "types-pywin32", specifier = ">=311.0.0" }, { name = "types-pyyaml", specifier = ">=6.0.12" }, { name = "types-redis", specifier = ">=4.6.0.20241004" }, @@ -1725,12 +1747,13 @@ dev = [ { name = "types-ujson", specifier = ">=5.10.0" }, { name = "xinference-client", specifier = ">=2.7.0" }, ] +kms = [{ name = "azure-keyvault-keys", specifier = ">=4.10.0,<5.0.0" }] storage = [ { name = "azure-storage-blob", specifier = ">=12.30.0,<13.0.0" }, - { name = "bce-python-sdk", specifier = "==0.9.72" }, + { name = "bce-python-sdk", specifier = "==0.9.76" }, { name = "cos-python-sdk-v5", specifier = ">=1.9.44,<2.0.0" }, { name = "esdk-obs-python", specifier = ">=3.26.6,<4.0.0" }, - { name = "google-cloud-storage", specifier = ">=3.12.1,<4.0.0" }, + { name = "google-cloud-storage", specifier = ">=3.13.0,<4.0.0" }, { name = "opendal", specifier = "==0.46.0" }, { name = "oss2", specifier = ">=2.19.1,<3.0.0" }, { name = "supabase", specifier = ">=2.31.0,<3.0.0" }, @@ -1738,7 +1761,7 @@ storage = [ ] tools = [ { name = "cloudscraper", specifier = ">=1.2.71,<2.0.0" }, - { name = "nltk", specifier = ">=3.9.1,<4.0.0" }, + { name = "nltk", specifier = ">=3.10.0,<4.0.0" }, ] trace-aliyun = [{ name = "dify-trace-aliyun", editable = "providers/trace/trace-aliyun" }] trace-all = [ @@ -1935,7 +1958,7 @@ dependencies = [ ] [package.metadata] -requires-dist = [{ name = "mysql-connector-python", specifier = ">=9.3.0,<10.0.0" }] +requires-dist = [{ name = "mysql-connector-python", specifier = ">=26.7.0,<27.0.0" }] [[package]] name = "dify-vdb-analyticdb" @@ -1949,7 +1972,7 @@ dependencies = [ [package.metadata] requires-dist = [ - { name = "alibabacloud-gpdb20160503", specifier = "~=5.2.0" }, + { name = "alibabacloud-gpdb20160503", specifier = "~=5.9.0" }, { name = "alibabacloud-tea-openapi", specifier = "==0.4.4" }, { name = "clickhouse-connect", specifier = "==0.15.1" }, ] @@ -1963,7 +1986,7 @@ dependencies = [ ] [package.metadata] -requires-dist = [{ name = "pymochow", specifier = "==2.4.0" }] +requires-dist = [{ name = "pymochow", specifier = "==2.4.1" }] [[package]] name = "dify-vdb-chroma" @@ -1996,7 +2019,7 @@ dependencies = [ ] [package.metadata] -requires-dist = [{ name = "couchbase", specifier = "~=4.6.0" }] +requires-dist = [{ name = "couchbase", specifier = "~=4.6.2" }] [[package]] name = "dify-vdb-elasticsearch" @@ -2040,7 +2063,7 @@ dependencies = [ ] [package.metadata] -requires-dist = [{ name = "intersystems-irispython", specifier = ">=5.1.0,<6.0.0" }] +requires-dist = [{ name = "intersystems-irispython", specifier = ">=5.4.0,<6.0.0" }] [[package]] name = "dify-vdb-lindorm" @@ -2101,7 +2124,7 @@ dependencies = [ [package.metadata] requires-dist = [ - { name = "mysql-connector-python", specifier = ">=9.3.0,<10.0.0" }, + { name = "mysql-connector-python", specifier = ">=26.7.0,<27.0.0" }, { name = "pyobvector", specifier = "==0.2.25" }, ] @@ -2130,7 +2153,7 @@ dependencies = [ ] [package.metadata] -requires-dist = [{ name = "oracledb", specifier = "==3.4.2" }] +requires-dist = [{ name = "oracledb", specifier = "==4.0.2" }] [[package]] name = "dify-vdb-pgvecto-rs" @@ -2152,7 +2175,7 @@ dependencies = [ ] [package.metadata] -requires-dist = [{ name = "pgvector", specifier = "==0.4.2" }] +requires-dist = [{ name = "pgvector", specifier = "==0.5.0" }] [[package]] name = "dify-vdb-qdrant" @@ -2179,7 +2202,7 @@ dependencies = [ ] [package.metadata] -requires-dist = [{ name = "tablestore", specifier = "==6.4.4" }] +requires-dist = [{ name = "tablestore", specifier = "==6.4.8" }] [[package]] name = "dify-vdb-tencent" @@ -2256,7 +2279,7 @@ dependencies = [ ] [package.metadata] -requires-dist = [{ name = "weaviate-client", specifier = "==4.20.5" }] +requires-dist = [{ name = "weaviate-client", specifier = "==4.22.0" }] [[package]] name = "diskcache-weave" @@ -2710,14 +2733,14 @@ wheels = [ [[package]] name = "gitpython" -version = "3.1.52" +version = "3.1.58" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "gitdb" }, ] -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" } +sdist = { url = "https://files.pythonhosted.org/packages/26/d6/5f358ff283325580c2003a6d953aea18cfe10ae87b46f5ebc80fa3a386dc/gitpython-3.1.58.tar.gz", hash = "sha256:621416df10ef3fd0e19fabf9172ddeed0fa704d353d04f194eec56a625a95b22", size = 228498, upload-time = "2026-08-04T15:05:49.47Z" } wheels = [ - { 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" }, + { url = "https://files.pythonhosted.org/packages/ec/0c/9d8752098bc442f0726e64aa6135940b3a96809915d1aa4206c1bb97881d/gitpython-3.1.58-py3-none-any.whl", hash = "sha256:d331e722577f0fd7fc1f857419b3ecc07af66282b933d2a4d95f84a042fdd50f", size = 220183, upload-time = "2026-08-04T15:05:48.025Z" }, ] [[package]] @@ -2891,7 +2914,7 @@ wheels = [ [[package]] name = "google-cloud-storage" -version = "3.12.1" +version = "3.13.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "google-api-core" }, @@ -2901,9 +2924,9 @@ dependencies = [ { name = "google-resumable-media" }, { name = "requests" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/da/ac/60b4cb0a6c8c6bb7cedb8971ba5e34a94096acf76e2cc242bcf1e6fc5c49/google_cloud_storage-3.12.1.tar.gz", hash = "sha256:1d81491c7663bc26c5056d00b834356f2253b910ef467f9cf9928a87fca1e04b", size = 17339353, upload-time = "2026-07-08T17:03:59.142Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e3/25/355ed97c1723c787dfaa888808d55db18371f82c38ff862357b1e902cd19/google_cloud_storage-3.13.0.tar.gz", hash = "sha256:d11d8706ea1520fba0f21043bcb7897caf7015d76ce1ad9a4f60237e4d7a9f6c", size = 17340960, upload-time = "2026-07-13T19:10:07.524Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/80/6e/ca176e95bafac0fe7befeee7e0420e686de147571cd2908e308c5fe71bda/google_cloud_storage-3.12.1-py3-none-any.whl", hash = "sha256:9297ae0c2ce3f5400b1f2bb3a3e6d2cd256614366e03cd30600871df8e903afb", size = 340845, upload-time = "2026-07-08T17:03:31.418Z" }, + { url = "https://files.pythonhosted.org/packages/81/e8/b3678a0931ee7d4b3fdaf0813e6206d66e0922b2c26d912f308728b5b95a/google_cloud_storage-3.13.0-py3-none-any.whl", hash = "sha256:648af3ef8a6acc674e1359d3c920c67eb89a7a5ab66b336bd3ac43fed6b5ab84", size = 341428, upload-time = "2026-07-13T19:09:52.39Z" }, ] [[package]] @@ -2991,7 +3014,7 @@ httpx = [ [[package]] name = "graphon" -version = "0.6.0" +version = "0.7.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "charset-normalizer" }, @@ -3013,9 +3036,9 @@ dependencies = [ { name = "unstructured", extra = ["docx", "epub", "md", "ppt", "pptx"] }, { name = "webvtt-py" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/ee/6c/9ea051ed30dc3306e9e77c4486b5a2e5462af45e35daf230d9ec886eb07e/graphon-0.6.0.tar.gz", hash = "sha256:2d3a386899dc7ab8e9767ab96c694ff7e6eb454c045a1e801505cba9c615160d", size = 264404, upload-time = "2026-06-29T15:26:27.437Z" } +sdist = { url = "https://files.pythonhosted.org/packages/6a/c6/6f16398e28bceb11304f8dd743d5f4f2c1e1aedf20eef6531764edb40a7e/graphon-0.7.0.tar.gz", hash = "sha256:e3f19284432d6e4a947fb285ef36bc79eb48c41af36867b03c0d70f062fe8563", size = 266715, upload-time = "2026-07-29T10:00:51.303Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/32/36/5b0ece2d61fa091f7d74aea5a48f9cf927202099e8d22b856e0319218e9d/graphon-0.6.0-py3-none-any.whl", hash = "sha256:f1445ccef40c0d0eb50a60af85c1028b26d2d782f357315047acee816acbcb33", size = 376038, upload-time = "2026-06-29T15:26:26.186Z" }, + { url = "https://files.pythonhosted.org/packages/9c/9b/3b02d944a20b3693bb5bb77cbb16b3afed44a965096970c47625560c0235/graphon-0.7.0-py3-none-any.whl", hash = "sha256:6f7608aaaa65607c4935735137ed4b7029b1e6ca80202cbe829af3370cf4404e", size = 380138, upload-time = "2026-07-29T10:00:49.654Z" }, ] [[package]] @@ -3184,15 +3207,15 @@ wheels = [ [[package]] name = "h2" -version = "4.3.0" +version = "4.4.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "hpack" }, { name = "hyperframe" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1d/17/afa56379f94ad0fe8defd37d6eb3f89a25404ffc71d4d848893d270325fc/h2-4.3.0.tar.gz", hash = "sha256:6c59efe4323fa18b47a632221a1888bd7fde6249819beda254aeca909f221bf1", size = 2152026, upload-time = "2025-08-23T18:12:19.778Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e7/85/7c366e69d84c17bb778fe41419e1fbcce3033d5b7ce29bbffff0a98b859f/h2-4.4.1.tar.gz", hash = "sha256:4e866ffb1a869ae14dd9b5e6beb5c24a13da0495ad72b65925ded182521c1516", size = 2157281, upload-time = "2026-08-03T11:45:09.509Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/69/b2/119f6e6dcbd96f9069ce9a2665e0146588dc9f88f29549711853645e736a/h2-4.3.0-py3-none-any.whl", hash = "sha256:c438f029a25f7945c69e0ccf0fb951dc3f73a5f6412981daee861431b70e2bdd", size = 61779, upload-time = "2025-08-23T18:12:17.779Z" }, + { url = "https://files.pythonhosted.org/packages/7e/22/e85faf23bd72a92d1921e37d674ca56eb298a3c8be31fdecef0ff2b3aaac/h2-4.4.1-py3-none-any.whl", hash = "sha256:0e25f1462b23c9cb82d9eb02e28bc706dac2a68cb457c6a0d74d63c8a2a5d0e6", size = 62636, upload-time = "2026-08-03T11:44:59.164Z" }, ] [[package]] @@ -3248,11 +3271,11 @@ wheels = [ [[package]] name = "hpack" -version = "4.1.0" +version = "4.2.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/2c/48/71de9ed269fdae9c8057e5a4c0aa7402e8bb16f2c6e90b3aa53327b113f8/hpack-4.1.0.tar.gz", hash = "sha256:ec5eca154f7056aa06f196a557655c5b009b382873ac8d1e66e79e87535f1dca", size = 51276, upload-time = "2025-01-22T21:44:58.347Z" } +sdist = { url = "https://files.pythonhosted.org/packages/26/5b/fcabf6028144a8723726318b07a32c2f3314acdff6265743cf08a344b18e/hpack-4.2.0.tar.gz", hash = "sha256:0895cfa3b5531fc65fe439c05eb65144f123bf7a394fcaa56aa423548d8e45c0", size = 51300, upload-time = "2026-06-23T18:34:46.667Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/07/c6/80c95b1b2b94682a72cbdbfb85b81ae2daffa4291fbfa1b1464502ede10d/hpack-4.1.0-py3-none-any.whl", hash = "sha256:157ac792668d995c657d93111f46b4535ed114f0c9c8d672271bbec7eae1b496", size = 34357, upload-time = "2025-01-22T21:44:56.92Z" }, + { url = "https://files.pythonhosted.org/packages/71/b4/4a9fcfb2aef6ba44d9073ecd301443aa00b3dac95de5619f2a7de7ec8a91/hpack-4.2.0-py3-none-any.whl", hash = "sha256:858ac0b02280fa582b5080d68db0899c62a80375e0e5413a74970c5e518b6986", size = 34246, upload-time = "2026-06-23T18:34:45.472Z" }, ] [[package]] @@ -3487,14 +3510,14 @@ wheels = [ [[package]] name = "intersystems-irispython" -version = "5.3.2" +version = "5.4.0" source = { registry = "https://pypi.org/simple" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d2/23/0a7bc92e68480d523015eb454aa0ec73a33320975d10d5500ba54ccd124e/intersystems_irispython-5.3.2-cp38.cp39.cp310.cp311.cp312.cp313.cp314-cp38.cp39.cp310.cp311.cp312.cp313.cp314-macosx_10_9_universal2.whl", hash = "sha256:8af5e31273ad97c391141111630e8303d510272360b609990a8c85e56a7850ac", size = 7121915, upload-time = "2026-03-31T18:53:12.205Z" }, - { url = "https://files.pythonhosted.org/packages/22/cc/2f066a0dc82fae884b655d2f862bd51dd21a4322d4b9f898117f74c010b4/intersystems_irispython-5.3.2-cp38.cp39.cp310.cp311.cp312.cp313.cp314-cp38.cp39.cp310.cp311.cp312.cp313.cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:25663d3cce7b414451a781ffaeb785e8f8439d0275920ffd4f05add2c056abfd", size = 16247974, upload-time = "2026-03-31T18:53:13.798Z" }, - { url = "https://files.pythonhosted.org/packages/27/cd/cef09a8310541d99fdbe89b2eccc21a6d776384325a9a6e740ad01e8461f/intersystems_irispython-5.3.2-cp38.cp39.cp310.cp311.cp312.cp313.cp314-cp38.cp39.cp310.cp311.cp312.cp313.cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:d5cb6efc3e2b9651f1c37539a3f69a823e80c32210d11d745cffad1eca4c7995", size = 15900577, upload-time = "2026-03-31T18:53:15.958Z" }, - { url = "https://files.pythonhosted.org/packages/37/91/0e08555834de10f59810ef6c615af72c3f234920c70cc0421d455ba9c359/intersystems_irispython-5.3.2-cp38.cp39.cp310.cp311.cp312.cp313.cp314-cp38.cp39.cp310.cp311.cp312.cp313.cp314-win32.whl", hash = "sha256:a250b21067c9e8275232ca798dcfe0719a970cd6ec9f2023923c810fffa46f41", size = 3046761, upload-time = "2026-03-31T18:53:09.151Z" }, - { url = "https://files.pythonhosted.org/packages/21/28/00b6b03b648005cb9c14dc75943e7cccce83eb5fd8fdba502028c25c7fc4/intersystems_irispython-5.3.2-cp38.cp39.cp310.cp311.cp312.cp313.cp314-cp38.cp39.cp310.cp311.cp312.cp313.cp314-win_amd64.whl", hash = "sha256:43feb7e23bc9f77db7bb140d1b55c22090b0c46691b570b1faaf6875baa6452d", size = 3742519, upload-time = "2026-03-31T18:53:10.597Z" }, + { url = "https://files.pythonhosted.org/packages/7c/d8/00b2e06068975d93910f18111f51754fb921243b9148d0fe50010edf550c/intersystems_irispython-5.4.0-cp39.cp310.cp311.cp312.cp313.cp314-cp39.cp310.cp311.cp312.cp313.cp314-macosx_10_9_universal2.whl", hash = "sha256:ae3d147698305161e33ba9e408fde020af917c22f5e738cb3d7e2fb3d113179f", size = 7074943, upload-time = "2026-07-30T19:07:03.024Z" }, + { url = "https://files.pythonhosted.org/packages/fb/f8/b8a044a0c5afc8e81bf3fa75c09698b32a85550d4caa1b68c709b57b5107/intersystems_irispython-5.4.0-cp39.cp310.cp311.cp312.cp313.cp314-cp39.cp310.cp311.cp312.cp313.cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:a8f060b0c25338d3de7bdebd9975d1b206ac07ea31325440a97bfeaede0473f2", size = 15494467, upload-time = "2026-07-30T19:07:14.941Z" }, + { url = "https://files.pythonhosted.org/packages/69/13/c2758a0d4d727813961cf4bde66b618a5eef2580e5d9578f27205555bc3d/intersystems_irispython-5.4.0-cp39.cp310.cp311.cp312.cp313.cp314-cp39.cp310.cp311.cp312.cp313.cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:737d6aa96202e216cf567a1ea04048923dc32ed070aca69cd3e3c1d1471b98a6", size = 15157284, upload-time = "2026-07-30T19:07:06.101Z" }, + { url = "https://files.pythonhosted.org/packages/6d/d4/27cb8aff27cdb9223e4dbff2b85e15ddf3fc87d477ea21bfb87c8eb51dfc/intersystems_irispython-5.4.0-cp39.cp310.cp311.cp312.cp313.cp314-cp39.cp310.cp311.cp312.cp313.cp314-win32.whl", hash = "sha256:8df3580f0d82b7d9d7c0b09d40ea04fbebbe3f78ab13251bd6d83d15ebf96e88", size = 3054941, upload-time = "2026-07-30T19:07:08.703Z" }, + { url = "https://files.pythonhosted.org/packages/96/b3/97b7da76d63e0b7e119872db6f85ae19371a39fe57f037c9fd08fd4e07af/intersystems_irispython-5.4.0-cp39.cp310.cp311.cp312.cp313.cp314-cp39.cp310.cp311.cp312.cp313.cp314-win_amd64.whl", hash = "sha256:97f3568c351de3f5b6cff088f104990b728b3243faab5947ed11357cd646d671", size = 3704377, upload-time = "2026-07-30T19:07:00.885Z" }, ] [[package]] @@ -4090,16 +4113,15 @@ wheels = [ [[package]] name = "mysql-connector-python" -version = "9.6.0" +version = "26.7.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/6f/6e/c89babc7de3df01467d159854414659c885152579903a8220c8db02a3835/mysql_connector_python-9.6.0.tar.gz", hash = "sha256:c453bb55347174d87504b534246fb10c589daf5d057515bf615627198a3c7ef1", size = 12254999, upload-time = "2026-02-10T12:04:52.63Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f2/ce/a53b169388f8c6a595cfa9a653138381f3afef1d2af60f5c1972d015f52f/mysql_connector_python-26.7.0.tar.gz", hash = "sha256:d8ff5ee236ea46661ee639336323e124ed868e37f3ea991bdc5de5a146f39fd5", size = 12256009, upload-time = "2026-07-29T10:24:13.314Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/8f/d9/2a4b4d90b52f4241f0f71618cd4bd8779dd6d18db8058b0a4dd83ec0541c/mysql_connector_python-9.6.0-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:9664e217c72dd6fb700f4c8512af90261f72d2f5d7c00c4e13e4c1e09bfa3d5e", size = 17585672, upload-time = "2026-02-10T12:03:52.955Z" }, - { url = "https://files.pythonhosted.org/packages/33/91/2495835733a054e716a17dc28404748b33f2dc1da1ae4396fb45574adf40/mysql_connector_python-9.6.0-cp312-cp312-macosx_14_0_x86_64.whl", hash = "sha256:1ed4b5c4761e5333035293e746683890e4ef2e818e515d14023fd80293bc31fa", size = 18452624, upload-time = "2026-02-10T12:03:56.153Z" }, - { url = "https://files.pythonhosted.org/packages/7a/69/e83abbbbf7f8eed855b5a5ff7285bc0afb1199418ac036c7691edf41e154/mysql_connector_python-9.6.0-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:5095758dcb89a6bce2379f349da336c268c407129002b595c5dba82ce387e2a5", size = 34169154, upload-time = "2026-02-10T12:03:58.831Z" }, - { url = "https://files.pythonhosted.org/packages/82/44/67bb61c71f398fbc739d07e8dcadad94e2f655874cb32ae851454066bea0/mysql_connector_python-9.6.0-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:4ae4e7780fad950a4f267dea5851048d160f5b71314a342cdbf30b154f1c74f7", size = 34542947, upload-time = "2026-02-10T12:04:02.408Z" }, - { url = "https://files.pythonhosted.org/packages/ba/39/994c4f7e9c59d3ca534a831d18442ac4c529865db20aeaa4fd94e2af5efd/mysql_connector_python-9.6.0-cp312-cp312-win_amd64.whl", hash = "sha256:c180e0b4100d7402e03993bfac5c97d18e01d7ca9d198d742fffc245077f8ffe", size = 16515709, upload-time = "2026-02-10T12:04:04.924Z" }, - { url = "https://files.pythonhosted.org/packages/15/dd/b3250826c29cee7816de4409a2fe5e469a68b9a89f6bfaa5eed74f05532c/mysql_connector_python-9.6.0-py2.py3-none-any.whl", hash = "sha256:44b0fb57207ebc6ae05b5b21b7968a9ed33b29187fe87b38951bad2a334d75d5", size = 480527, upload-time = "2026-02-10T12:04:36.176Z" }, + { url = "https://files.pythonhosted.org/packages/e3/24/8c3516ee9edc3d740059253410bcb30a0dd7686bfd876eb905af6b92b4e6/mysql_connector_python-26.7.0-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:9b63672cc381f097966faecd6940beb2a91f9efdf674a5807dc45f4efdb42e62", size = 20279974, upload-time = "2026-07-29T10:18:12.334Z" }, + { url = "https://files.pythonhosted.org/packages/ad/9c/d2f705d2caaf2c91d05bdea0d67dde7cc243a62245ee91fba4e201e24ea1/mysql_connector_python-26.7.0-cp312-cp312-macosx_14_0_x86_64.whl", hash = "sha256:375f9a042398e515a670003f26ab07e9b2ecaecba6eaf678f66fb5c2d36bd305", size = 19840279, upload-time = "2026-07-29T10:18:15.057Z" }, + { url = "https://files.pythonhosted.org/packages/4f/1c/73b9f540fbbcf78949e00d673db576f61ab93e7824e31dfb568295e58f2c/mysql_connector_python-26.7.0-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:227063e73c3d7b0326bc6ce9af47b8aa081e53a4e97e8683bba658569cb57271", size = 21954262, upload-time = "2026-07-29T10:18:18.148Z" }, + { url = "https://files.pythonhosted.org/packages/2d/0c/7dc1c1367a6a7881f20b0ed8bf92133d4710ea0dce22657972f56bf2a365/mysql_connector_python-26.7.0-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:e118030bdd6a9882f5aa80a0c0c88d684d0e82d78747e5baf3c3e54a5a2d1348", size = 21746245, upload-time = "2026-07-29T10:18:20.65Z" }, + { url = "https://files.pythonhosted.org/packages/ff/01/07efeb3e3fc730f945af40ddf1f4c879c13945b193a94eba8e642c0d4176/mysql_connector_python-26.7.0-py2.py3-none-any.whl", hash = "sha256:43ee453b53556f690e4f4b795d3f17a1c6d4b45f729657bb0d999c8dfc9261d5", size = 480733, upload-time = "2026-07-29T10:18:49.777Z" }, ] [[package]] @@ -4113,17 +4135,18 @@ wheels = [ [[package]] name = "nltk" -version = "3.9.4" +version = "3.10.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "click" }, + { name = "defusedxml" }, { name = "joblib" }, { name = "regex" }, { name = "tqdm" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/74/a1/b3b4adf15585a5bc4c357adde150c01ebeeb642173ded4d871e89468767c/nltk-3.9.4.tar.gz", hash = "sha256:ed03bc098a40481310320808b2db712d95d13ca65b27372f8a403949c8b523d0", size = 2946864, upload-time = "2026-03-24T06:13:40.641Z" } +sdist = { url = "https://files.pythonhosted.org/packages/96/02/df4f105b28a7c16b0e41423bc09cf0f1b8a305df4ef0b10ca74a2e4c648c/nltk-3.10.0.tar.gz", hash = "sha256:4fbac1d98203cbcd1b5d94a2877fb822300072d80604a5e7fae49d2c5f84e8c1", size = 3089244, upload-time = "2026-07-08T02:39:13.562Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/9d/91/04e965f8e717ba0ab4bdca5c112deeab11c9e750d94c4d4602f050295d39/nltk-3.9.4-py3-none-any.whl", hash = "sha256:f2fa301c3a12718ce4a0e9305c5675299da5ad9e26068218b69d692fda84828f", size = 1552087, upload-time = "2026-03-24T06:13:38.47Z" }, + { url = "https://files.pythonhosted.org/packages/6e/89/a0b0f35e2820d6a99d75ea1c11977ee6d5c9e6658eceb45b0c7620881faa/nltk-3.10.0-py3-none-any.whl", hash = "sha256:54ff84d4916d3ef127e8953bee0023f6a6b320b75d634a19e06ef056d3d244bf", size = 1716144, upload-time = "2026-07-08T02:39:09.753Z" }, ] [[package]] @@ -4322,59 +4345,58 @@ wheels = [ [[package]] name = "opentelemetry-api" -version = "1.41.1" +version = "1.44.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "importlib-metadata" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/fa/fc/b7564cbef36601aef0d6c9bc01f7badb64be8e862c2e1c3c5c3b43b53e4f/opentelemetry_api-1.41.1.tar.gz", hash = "sha256:0ad1814d73b875f84494387dae86ce0b12c68556331ce6ce8fe789197c949621", size = 71416, upload-time = "2026-04-24T13:15:38.262Z" } +sdist = { url = "https://files.pythonhosted.org/packages/ee/8b/aa9e2d8b8dfa7c946f7dec5d1f8f6ba8eca062f43509a06bdb5ce93d26c0/opentelemetry_api-1.44.0.tar.gz", hash = "sha256:67647e5e9566edcf421166fdf022b3537f818635daa852b289e34604dc6fb33a", size = 72406, upload-time = "2026-07-16T15:25:32.678Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/29/59/3e7118ed140f76b0982ba4321bdaed1997a0473f9720de2d10788a577033/opentelemetry_api-1.41.1-py3-none-any.whl", hash = "sha256:a22df900e75c76dc08440710e51f52f1aa6b451b429298896023e60db5b3139f", size = 69007, upload-time = "2026-04-24T13:15:15.662Z" }, + { url = "https://files.pythonhosted.org/packages/ca/6f/a04e900f465ff3221ccc395522503e2d10e79fa21f2723c8e177aae1e0d1/opentelemetry_api-1.44.0-py3-none-any.whl", hash = "sha256:94b98c893a91b88657eaac1e3ba89618cdb85be6918196705354f34728b2cdef", size = 60018, upload-time = "2026-07-16T15:25:11.657Z" }, ] [[package]] name = "opentelemetry-distro" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, { name = "opentelemetry-instrumentation" }, { name = "opentelemetry-sdk" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/45/f1/314e5015e353a001948e03f48a6935ca7ef00e99107b8e3e63871426b0f6/opentelemetry_distro-0.62b1.tar.gz", hash = "sha256:0169b128b9d6d5cab809ae4c4fb3d576bfc5d3f30b32d8a43b770b587f04f253", size = 2606, upload-time = "2026-04-24T13:22:29.403Z" } +sdist = { url = "https://files.pythonhosted.org/packages/5f/61/29e94c78299b173486e6e3b10d63e89cbe340745743c6ed6ae2055433743/opentelemetry_distro-0.65b0.tar.gz", hash = "sha256:e2fef26fdf72978ab172530c13338aff367edae41ae8d90b7f6133efedde10ab", size = 2334, upload-time = "2026-07-16T15:25:47.917Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b9/19/c58c119a299298f03d0797fcb780f221880e8d725959c71bcfb4ae034738/opentelemetry_distro-0.62b1-py3-none-any.whl", hash = "sha256:fd938de6ca1d047ffd15a65fa09d89f4b4ca7dd97ef25601a12d6d10efd693a0", size = 3348, upload-time = "2026-04-24T13:21:27.389Z" }, + { url = "https://files.pythonhosted.org/packages/25/d5/603a29962542152331c6e37a121129a26e24917daa4298ca00bea97571bb/opentelemetry_distro-0.65b0-py3-none-any.whl", hash = "sha256:0e15bb1a54c638e6c361bc66bc43e9de9ead442e1ae06f8099a2c1d27b6ef845", size = 2775, upload-time = "2026-07-16T15:24:47.562Z" }, ] [[package]] name = "opentelemetry-exporter-otlp" -version = "1.41.1" +version = "1.44.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-exporter-otlp-proto-grpc" }, { name = "opentelemetry-exporter-otlp-proto-http" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/42/84/d55baf8e1a222f40282956083e67de9fa92d5fa451108df4839505fa2a24/opentelemetry_exporter_otlp-1.41.1.tar.gz", hash = "sha256:299a2f0541ca175df186f5ac58fd5db177ba1e9b72b0826049062f750d55b47f", size = 6152, upload-time = "2026-04-24T13:15:40.006Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f2/45/7af37fe54e5d3e66e7dcd7ba8b8aeee73f202bfac909cc94b8c4e428f9ac/opentelemetry_exporter_otlp-1.44.0.tar.gz", hash = "sha256:af1cde7c33ea8ed624bf04ac49a885730fe44c1f1ad698656e592c38f70ce106", size = 6090, upload-time = "2026-07-16T15:25:34.585Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/6d/d5/ea4aa7dfc458fd537bd9519ea0e7226eef2a6212dfe952694984167daaba/opentelemetry_exporter_otlp-1.41.1-py3-none-any.whl", hash = "sha256:db276c5a80c02b063994e80950d00ca1bfddcf6520f608335b7dc2db0c0eb9c6", size = 7025, upload-time = "2026-04-24T13:15:17.839Z" }, + { url = "https://files.pythonhosted.org/packages/18/c3/7b466a9463944e70b37b744072a0c1b88a425dade3fff0631adec66c9bcc/opentelemetry_exporter_otlp-1.44.0-py3-none-any.whl", hash = "sha256:4a498fa8d8fd8be9e8e2d175fe5524a3fe581ccffadd8509db86526a5fb97051", size = 6727, upload-time = "2026-07-16T15:25:14.445Z" }, ] [[package]] name = "opentelemetry-exporter-otlp-proto-common" -version = "1.41.1" +version = "1.44.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-proto" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/ae/fa/f9e3bd3c4d692b3ce9a2880a167d1f79681a1bea11f00d5bf76adc03e6ea/opentelemetry_exporter_otlp_proto_common-1.41.1.tar.gz", hash = "sha256:0e253156ea9c36b0bd3d2440c5c9ba7dd1f3fb64ba7a08fc85fbac536b56e1fb", size = 20409, upload-time = "2026-04-24T13:15:40.924Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/09/4d717852c1cf3f854b76c7110a5d00883bc3c99288b9b0dbcbeb9e306eb6/opentelemetry_exporter_otlp_proto_common-1.44.0.tar.gz", hash = "sha256:dc87a5a5bc58f149a56d1547e4691588fa12994cdc3bc039a694ccb3375862ac", size = 20202, upload-time = "2026-07-16T15:25:37.658Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/29/48/bce76d3ea772b609757e9bc844e02ab408a6446609bf74fb562062ba6b71/opentelemetry_exporter_otlp_proto_common-1.41.1-py3-none-any.whl", hash = "sha256:10da74dad6a49344b9b7b21b6182e3060373a235fde1528616d5f01f92e66aa9", size = 18366, upload-time = "2026-04-24T13:15:18.917Z" }, + { url = "https://files.pythonhosted.org/packages/5e/71/65fd9d54c10b860f87c045ccee1264cab7011268895d3528818a29c1172a/opentelemetry_exporter_otlp_proto_common-1.44.0-py3-none-any.whl", hash = "sha256:9a9fe61bba73d802904bc989f1d6b4a7b1ee40f06c40e98d6f85af65aaebb694", size = 17045, upload-time = "2026-07-16T15:25:18.201Z" }, ] [[package]] name = "opentelemetry-exporter-otlp-proto-grpc" -version = "1.41.1" +version = "1.44.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "googleapis-common-protos" }, @@ -4385,14 +4407,14 @@ dependencies = [ { name = "opentelemetry-sdk" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1e/9b/e4503060b8695579dbaad187dc8cef4554188de68748c88060599b77489e/opentelemetry_exporter_otlp_proto_grpc-1.41.1.tar.gz", hash = "sha256:b05df8fa1333dc9a3fda36b676b96b5095ab6016d3f0c3296d430d629ba1443b", size = 25755, upload-time = "2026-04-24T13:15:41.93Z" } +sdist = { url = "https://files.pythonhosted.org/packages/1f/47/80d9e9d468dc5de3af5096f5ccdb065fa4dd1470f74495cc53e59e397f47/opentelemetry_exporter_otlp_proto_grpc-1.44.0.tar.gz", hash = "sha256:40d1ae9e03fcc36de3cbac610cc99f35894938bff9cfd90fc4ec68bd85448463", size = 27225, upload-time = "2026-07-16T15:25:38.308Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ac/f2/c54f33c92443d087703e57e52e55f22f111373a5c4c4aa349ea60efe512e/opentelemetry_exporter_otlp_proto_grpc-1.41.1-py3-none-any.whl", hash = "sha256:537926dcef951136992479af1d9cd88f25e33d56c530e9f020ed57774dca2f94", size = 20297, upload-time = "2026-04-24T13:15:20.212Z" }, + { url = "https://files.pythonhosted.org/packages/54/29/6ae42ba32b153ae0a44ae125f0caff2188bbe62d99c82d1768da30864e72/opentelemetry_exporter_otlp_proto_grpc-1.44.0-py3-none-any.whl", hash = "sha256:6a1a645ea182a2f59440c51fa8301d309f3324a8f9d65f8395584b064b67ee4e", size = 19624, upload-time = "2026-07-16T15:25:19.096Z" }, ] [[package]] name = "opentelemetry-exporter-otlp-proto-http" -version = "1.41.1" +version = "1.44.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "googleapis-common-protos" }, @@ -4403,14 +4425,14 @@ dependencies = [ { name = "requests" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/33/5b/9d3c7f70cca10136ba82a81e738dee626c8e7fc61c6887ea9a58bf34c606/opentelemetry_exporter_otlp_proto_http-1.41.1.tar.gz", hash = "sha256:4747a9604c8550ab38c6fd6180e2fcb80de3267060bef2c306bad3cb443302bc", size = 24139, upload-time = "2026-04-24T13:15:42.977Z" } +sdist = { url = "https://files.pythonhosted.org/packages/1a/87/95e2a5aaa795b4e2260d74e16df2d5541deb2ea9de010bcd615f4dee2654/opentelemetry_exporter_otlp_proto_http-1.44.0.tar.gz", hash = "sha256:c633d7270ad6b57cd4cfbe8b0007a9e2e7c0cb50bd6c50fe2a7b245f721a09d8", size = 25806, upload-time = "2026-07-16T15:25:39.162Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ba/4d/ef07ff2fc630849f2080ae0ae73a61f67257905b7ac79066640bfa0c5739/opentelemetry_exporter_otlp_proto_http-1.41.1-py3-none-any.whl", hash = "sha256:1a21e8f49c7a946d935551e90947d6c3eb39236723c6624401da0f33d68edcb4", size = 22673, upload-time = "2026-04-24T13:15:21.313Z" }, + { url = "https://files.pythonhosted.org/packages/cd/d0/fdeb1a98d8d3a6205f5f297c51b4a9bfe65126ab60339669bbe3dd54c2e2/opentelemetry_exporter_otlp_proto_http-1.44.0-py3-none-any.whl", hash = "sha256:838592fce774c1c8bb7b9a0a7facbfa82e17be5a8a4e94cef10cb84ae026bae3", size = 21850, upload-time = "2026-07-16T15:25:20.006Z" }, ] [[package]] name = "opentelemetry-instrumentation" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -4418,14 +4440,14 @@ dependencies = [ { name = "packaging" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/52/cb/0523b92c112a6cc70be43724343dc45225d3af134419844d7879a07755d4/opentelemetry_instrumentation-0.62b1.tar.gz", hash = "sha256:90e92a905ba4f84db06ac3aec96701df6c079b2d66e9379f8739f0a1bdcc7f45", size = 34043, upload-time = "2026-04-24T13:22:31.997Z" } +sdist = { url = "https://files.pythonhosted.org/packages/13/91/3c58961cb0360cd60509064734f0be4275383c8681d73c580a40ca83ddce/opentelemetry_instrumentation-0.65b0.tar.gz", hash = "sha256:071d9d9eced9bd6460444ec3b0c77229870ed05a881c22c84fdede58e4eed09b", size = 42689, upload-time = "2026-07-16T15:25:50.275Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/4d/0f/45adbaea1f81b847cffdcee4f4b5f89297e42facf7fac78c7aaac4c38e75/opentelemetry_instrumentation-0.62b1-py3-none-any.whl", hash = "sha256:976fc6e640f2006599e97429c949e622c108d0c17c2059347d1e6c93c707f257", size = 34163, upload-time = "2026-04-24T13:21:31.722Z" }, + { url = "https://files.pythonhosted.org/packages/40/7b/85eab1215f72adf0e68d3dc4a679b9bff993fa679ff34cd8dd378e2659fd/opentelemetry_instrumentation-0.65b0-py3-none-any.whl", hash = "sha256:ea967a72b9939b5fcfdad572753b4306c59dcb99e3f382d95dae04286805e137", size = 36717, upload-time = "2026-07-16T15:24:51.424Z" }, ] [[package]] name = "opentelemetry-instrumentation-asgi" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "asgiref" }, @@ -4434,28 +4456,28 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-util-http" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/54/43/b2f0703ff46718ff7b17d7fbf8e9d7f20e26a23c7c325092dd762d09cf9d/opentelemetry_instrumentation_asgi-0.62b1.tar.gz", hash = "sha256:7cf5f5d5c493bbb1edd2bd6d51fa879d964e94048904017258a32ffa47329310", size = 26781, upload-time = "2026-04-24T13:22:37.158Z" } +sdist = { url = "https://files.pythonhosted.org/packages/17/83/8e8e83b7ac285281687c7be2fd305213ccccbb8c0a2dd4fb45a8ccaf12c7/opentelemetry_instrumentation_asgi-0.65b0.tar.gz", hash = "sha256:892bca67c56522ffa85a8a83cf934d7b50b3be2132e45cbee705825f0a5ba426", size = 26140, upload-time = "2026-07-16T15:25:54.544Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/d0/41/968c1fe12fb90abffca6620e65d4af91451c02ecca8f74a17a62cac490de/opentelemetry_instrumentation_asgi-0.62b1-py3-none-any.whl", hash = "sha256:b7f89be48528512619bd54fa2459f72afb1695ba71d7024d382ad96d467e7fa8", size = 17011, upload-time = "2026-04-24T13:21:38.006Z" }, + { url = "https://files.pythonhosted.org/packages/0b/9c/376962840b619d2d55fe8ee2285f8c70971c090e5fff614516fc654a6f3a/opentelemetry_instrumentation_asgi-0.65b0-py3-none-any.whl", hash = "sha256:3a845a8ebd1c4ef0d8263401e6545f5b219b2feee612090d50f578a87e71fd65", size = 15903, upload-time = "2026-07-16T15:24:57.198Z" }, ] [[package]] name = "opentelemetry-instrumentation-celery" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, { name = "opentelemetry-instrumentation" }, { name = "opentelemetry-semantic-conventions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/35/86/9e78c174b2f6ea92af3f99aa7488807b74290a5cd44a8e05bfbfd7b109be/opentelemetry_instrumentation_celery-0.62b1.tar.gz", hash = "sha256:f0035abd464a2989414a9c5ecdd79a25c87bd8c43f96c7f39e07000c6f25dfef", size = 14809, upload-time = "2026-04-24T13:22:45.656Z" } +sdist = { url = "https://files.pythonhosted.org/packages/b1/ba/c4ef088927a6bf1b75916b3bfca93552526f0e3e85d19557666f858fc5b6/opentelemetry_instrumentation_celery-0.65b0.tar.gz", hash = "sha256:ff4d192b6371cfe06969005862dc558cc8b79b03b0bbdfbf65b9a59e4dcd7fee", size = 16071, upload-time = "2026-07-16T15:26:00.867Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/24/51/f38a31ac8f8e3bd365f301f697661679addaf548d52a05cfdde4448a5493/opentelemetry_instrumentation_celery-0.62b1-py3-none-any.whl", hash = "sha256:50567a47b7adc4ea552d09709de4d73fea7b4ff24ab0e9d38739d03fcd3f95ef", size = 13864, upload-time = "2026-04-24T13:21:46.557Z" }, + { url = "https://files.pythonhosted.org/packages/1d/3a/3a53db4c82bd42f67a16c846aa972aef75043557e26e9d6a4ca8ffac9699/opentelemetry_instrumentation_celery-0.65b0-py3-none-any.whl", hash = "sha256:8510357a6a9e2cb1a6485fe98642a39f06c2970eb58656de06260b922ffe7888", size = 13556, upload-time = "2026-07-16T15:25:05.427Z" }, ] [[package]] name = "opentelemetry-instrumentation-fastapi" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -4464,14 +4486,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-util-http" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/77/38/91780475a25370b6d483afbaed3e1e170459d6351c5f7c08d66b65e2172e/opentelemetry_instrumentation_fastapi-0.62b1.tar.gz", hash = "sha256:b377d4ba32868fb1ff0f64da3fcdd3aa154d698fc83d65f5d380ea21bf31ee19", size = 25054, upload-time = "2026-04-24T13:22:50.222Z" } +sdist = { url = "https://files.pythonhosted.org/packages/30/23/b057f8196d06efdc1b50e3ff11fbc499a7d96b35c87f217eb7885542f4ea/opentelemetry_instrumentation_fastapi-0.65b0.tar.gz", hash = "sha256:10a3a95486036230413a58fe4fdf4a83fa6bba46918407e527476994bd92bd97", size = 26236, upload-time = "2026-07-16T15:26:05.954Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/8c/6f/602e4081d3fe82731aff7e3e9c2f1662d85701841d6dc25f16a1874e11cd/opentelemetry_instrumentation_fastapi-0.62b1-py3-none-any.whl", hash = "sha256:93fa9cc4f315819aee5f4fceb6196c1e5b0fbd789c5520c631de228bd3e5285b", size = 13484, upload-time = "2026-04-24T13:21:54.538Z" }, + { url = "https://files.pythonhosted.org/packages/fa/b0/c9b0300d33349ecc3dfd2362516eaffc44877e90970e6a52178ff953fec3/opentelemetry_instrumentation_fastapi-0.65b0-py3-none-any.whl", hash = "sha256:cda2610a0ec1b22d19886f33e4d861e9f5dbb886aeaa3a1263b47aff82c36943", size = 13261, upload-time = "2026-07-16T15:25:12.429Z" }, ] [[package]] name = "opentelemetry-instrumentation-flask" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -4481,14 +4503,14 @@ dependencies = [ { name = "opentelemetry-util-http" }, { name = "packaging" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/d3/08/e52e6eab550db1736c5657a7e38484c22a101009e77fc67eb00b272a96c1/opentelemetry_instrumentation_flask-0.62b1.tar.gz", hash = "sha256:37662ad159570dab1e3017a2a415193c014a5798fc32d33f3bdd254469e8c69a", size = 24100, upload-time = "2026-04-24T13:22:50.845Z" } +sdist = { url = "https://files.pythonhosted.org/packages/5f/45/7ddc536b91d133ad9e40fb0adc9f48d619400e47d8b3b936cb9d41443e04/opentelemetry_instrumentation_flask-0.65b0.tar.gz", hash = "sha256:887de3a97c09953da09ae713fbb777172900f33b2924d85dad314a033156ef66", size = 24149, upload-time = "2026-07-16T15:26:06.676Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/2b/58/d0e5e82d225365987bd192576095b1125f6b172decc4db79963373c92b74/opentelemetry_instrumentation_flask-0.62b1-py3-none-any.whl", hash = "sha256:6df32684a7dd5dab5feb499c0748a4628b3fd139bffd8171326fb479aa525367", size = 16007, upload-time = "2026-04-24T13:21:55.462Z" }, + { url = "https://files.pythonhosted.org/packages/66/4f/0fc0c3f78323b0c3cdb8d3e2bde68c2a8a98e4cabdd9c41076c6e02d0825/opentelemetry_instrumentation_flask-0.65b0-py3-none-any.whl", hash = "sha256:d5337dac3b2af7f658fbc11c879667c9978910e38744b9706508f0b9908f7841", size = 15081, upload-time = "2026-07-16T15:25:13.494Z" }, ] [[package]] name = "opentelemetry-instrumentation-httpx" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -4497,14 +4519,14 @@ dependencies = [ { name = "opentelemetry-util-http" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/33/cb/7a418e69c7dad281803529cb4f6de1b747d802cca44c38032668690b4836/opentelemetry_instrumentation_httpx-0.62b1.tar.gz", hash = "sha256:a1fac9bcc3a6ef5996a7990563f1af0798468b2c146de535fd598369383fba7e", size = 24181, upload-time = "2026-04-24T13:22:52.124Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/03/a529140241addd4d0acc73bafbd6f74691651b92fc0ae9b4513cf80f07fa/opentelemetry_instrumentation_httpx-0.65b0.tar.gz", hash = "sha256:4627aa9c6bb99bf4462c8b565b0ef6aeb9ffad95c6c92868be1ef7895de112ee", size = 26309, upload-time = "2026-07-16T15:26:07.973Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/c7/e0/eca824e9492ccec00e055bdd243aeda8eb7c5eda746d98af4d7a2d97ecf3/opentelemetry_instrumentation_httpx-0.62b1-py3-none-any.whl", hash = "sha256:88614015df451d61bc7e73f22524e6f223611f80b6caad2f6bdcbe05fa0df653", size = 17201, upload-time = "2026-04-24T13:21:58.072Z" }, + { url = "https://files.pythonhosted.org/packages/9d/0f/c6144096b4914bbf44b43ba21c962e8f333ff045770b50a3e79ed8bd455f/opentelemetry_instrumentation_httpx-0.65b0-py3-none-any.whl", hash = "sha256:400f1b78afa4ee2332b5debe58e1ed1b317913d58812c952576be76660aeadb1", size = 17436, upload-time = "2026-07-16T15:25:15.772Z" }, ] [[package]] name = "opentelemetry-instrumentation-redis" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -4512,14 +4534,14 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/f5/ff/35414ad80409bd9e472c7959832524c5f2c8f63965af08c41c2b42d3a6a6/opentelemetry_instrumentation_redis-0.62b1.tar.gz", hash = "sha256:2d3c421d95e05ade075bee5becbe34e743b1cdf5bdee2085cb524f88c4f13dcb", size = 14796, upload-time = "2026-04-24T13:23:01.138Z" } +sdist = { url = "https://files.pythonhosted.org/packages/21/e5/b369897a9ff902eb028f1f999c9f9a4602f0b30fe1e97298d2c15fe5b8b0/opentelemetry_instrumentation_redis-0.65b0.tar.gz", hash = "sha256:f16409f189092984ff922f26939e7be79509365f5c9d202308c13fddf1368147", size = 17218, upload-time = "2026-07-16T15:26:17.243Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/31/37/bc2271f3472e3041eeade8b8da1cfd3b06badae76fe5d0ff135b6285e70c/opentelemetry_instrumentation_redis-0.62b1-py3-none-any.whl", hash = "sha256:9aedd02c1acf631251d1d676634db47da9da04e0a626cd0c7d83fe0eb791d165", size = 15501, upload-time = "2026-04-24T13:22:11.705Z" }, + { url = "https://files.pythonhosted.org/packages/9d/2a/5b874bf8824e7cb995b22e1e61e1b5bc42810023d81612c44edbfd2c48a3/opentelemetry_instrumentation_redis-0.65b0-py3-none-any.whl", hash = "sha256:7d135f61db9d72416e1f382012fbc13de6f5e726491bebd0344a9c42a60e6769", size = 14658, upload-time = "2026-07-16T15:25:29.146Z" }, ] [[package]] name = "opentelemetry-instrumentation-sqlalchemy" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -4528,14 +4550,14 @@ dependencies = [ { name = "packaging" }, { name = "wrapt" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1a/53/fa511ab998dd66b4eb66a36d8c262d0604cc5bad7a9c82e923be038dda97/opentelemetry_instrumentation_sqlalchemy-0.62b1.tar.gz", hash = "sha256:bdeac015351a1de057e8ea39f1fe26c9e60ea6bedbf1d5ad6a8262a516b3dc7d", size = 18539, upload-time = "2026-04-24T13:23:03.169Z" } +sdist = { url = "https://files.pythonhosted.org/packages/ae/9d/72af21527ce6f286ac78dc2c75ce44b76cc45343a28f0bcbb13f213421fc/opentelemetry_instrumentation_sqlalchemy-0.65b0.tar.gz", hash = "sha256:8ec2e79f1e00808c5dc639ab5b2cfcdf9dcdef55efd6442c6019707fe2894028", size = 18007, upload-time = "2026-07-16T15:26:19.197Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/2d/c5/aa2abcf8752a435536901636c5d540ba7a2c0ba2c4e98c7d119482e04262/opentelemetry_instrumentation_sqlalchemy-0.62b1-py3-none-any.whl", hash = "sha256:613542ecd52aabeec83d8813b5c287a3fb6c9ac3cd660694c94c0571f066e972", size = 15536, upload-time = "2026-04-24T13:22:14.767Z" }, + { url = "https://files.pythonhosted.org/packages/71/94/3eb3195a60b894cc1852a5657a2b64764a998085c50d7fb3452148af8f18/opentelemetry_instrumentation_sqlalchemy-0.65b0-py3-none-any.whl", hash = "sha256:4d8a2e5afc7b505a48d05cb1fb6db5f8b31681814c1d7767bd6d4b5bdd3a3047", size = 14410, upload-time = "2026-07-16T15:25:32.04Z" }, ] [[package]] name = "opentelemetry-instrumentation-wsgi" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, @@ -4543,9 +4565,9 @@ dependencies = [ { name = "opentelemetry-semantic-conventions" }, { name = "opentelemetry-util-http" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/36/db/19f1d66cead56e52291fccaa235b07ad45a5c24be1c740301a840c68235a/opentelemetry_instrumentation_wsgi-0.62b1.tar.gz", hash = "sha256:02a364fd9c940a46b19c825c5bfe386b007d5292ef91573894164836953fe831", size = 19919, upload-time = "2026-04-24T13:23:09.796Z" } +sdist = { url = "https://files.pythonhosted.org/packages/17/f7/bfe74ba3baea8c61290c3ab1d692f841b695ccb6e40faf5ad55fd5412cf2/opentelemetry_instrumentation_wsgi-0.65b0.tar.gz", hash = "sha256:d4a62ae98667ddfe04fe538c3c54abad538feb8c9c7a407ba19f016e1ce4a89a", size = 19666, upload-time = "2026-07-16T15:26:25.485Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f7/0e/60fec0780e16929c821df7c55c4f0bea45d6ef562e662c5f27f47d0ff195/opentelemetry_instrumentation_wsgi-0.62b1-py3-none-any.whl", hash = "sha256:a2df11de0113f504043e2b0fa0288238a93ee49ff607bd5100cb2d3a75bc771f", size = 14629, upload-time = "2026-04-24T13:22:23.951Z" }, + { url = "https://files.pythonhosted.org/packages/05/af/9dfe500d3b816e4bd65f467c6275688744ed8186894c73223f7e3d6d9745/opentelemetry_instrumentation_wsgi-0.65b0-py3-none-any.whl", hash = "sha256:af23e6686c7cd2abcd7d14ac03fb7e3b438273eb2d54a8f8dc401dc71bc52a9f", size = 13787, upload-time = "2026-07-16T15:25:42.876Z" }, ] [[package]] @@ -4563,50 +4585,50 @@ wheels = [ [[package]] name = "opentelemetry-proto" -version = "1.41.1" +version = "1.44.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "protobuf" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/99/e8/633c6d8a9c8840338b105907e55c32d3da1983abab5e52f899f72a82c3d1/opentelemetry_proto-1.41.1.tar.gz", hash = "sha256:4b9d2eb631237ea43b80e16c073af438554e32bc7e9e3f8ca4a9582f900020e5", size = 45670, upload-time = "2026-04-24T13:15:49.768Z" } +sdist = { url = "https://files.pythonhosted.org/packages/64/01/40ac4ae9a149263cc52c2cee200ddd80cb6d8db1a4610abf8eabce0fe771/opentelemetry_proto-1.44.0.tar.gz", hash = "sha256:c547a79c2f8c0c515d31509154682e5921c7cfd5ca67b70e1f9266e2c3e103f3", size = 46488, upload-time = "2026-07-16T15:25:45.34Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/e4/1e/5cd77035e3e82070e2265a63a760f715aacd3cb16dddc7efee913f297fcc/opentelemetry_proto-1.41.1-py3-none-any.whl", hash = "sha256:0496713b804d127a4147e32849fbaf5683fac8ee98550e8e7679cd706c289720", size = 72076, upload-time = "2026-04-24T13:15:32.542Z" }, + { url = "https://files.pythonhosted.org/packages/d1/7c/8be563d68e93bbefa5c8affb82ddcff91b3ad858ce49957ba7b16fd3e0ab/opentelemetry_proto-1.44.0-py3-none-any.whl", hash = "sha256:898b155a0e1557afd867478fb6158e8122a46329ca0bb8dc53cc55e98f017f56", size = 72483, upload-time = "2026-07-16T15:25:28.429Z" }, ] [[package]] name = "opentelemetry-sdk" -version = "1.41.1" +version = "1.44.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, { name = "opentelemetry-semantic-conventions" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/58/d0/54ee30dab82fb0acda23d144502771ff76ef8728459c83c3e89ef9fb1825/opentelemetry_sdk-1.41.1.tar.gz", hash = "sha256:724b615e1215b5aeacda0abb8a6a8922c9a1853068948bd0bd225a56d0c792e6", size = 230180, upload-time = "2026-04-24T13:15:50.991Z" } +sdist = { url = "https://files.pythonhosted.org/packages/5d/77/a6592cbc7c8d9bcc9d6757a9df45e04a7c585e3e6e7a13456da522b21109/opentelemetry_sdk-1.44.0.tar.gz", hash = "sha256:cebe7f65dc12f26ead75c6064de12fd2a9052e5060c0272d402cfa203aae123b", size = 208624, upload-time = "2026-07-16T15:25:46.078Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/b4/e7/a1420b698aad018e1cf60fdbaaccbe49021fb415e2a0d81c242f4c518f54/opentelemetry_sdk-1.41.1-py3-none-any.whl", hash = "sha256:edee379c126c1bce952b0c812b48fe8ff35b30df0eecf17e98afa4d598b7d85d", size = 180213, upload-time = "2026-04-24T13:15:33.767Z" }, + { url = "https://files.pythonhosted.org/packages/e7/23/ff077e61886ee020a17ce9c8b6fa11c601c8d8345b09ea24f605445df62a/opentelemetry_sdk-1.44.0-py3-none-any.whl", hash = "sha256:df081c4c6bcfdb1211e3e86140376792643128a25f8d72d1d27675936e7e96ad", size = 137221, upload-time = "2026-07-16T15:25:29.534Z" }, ] [[package]] name = "opentelemetry-semantic-conventions" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "opentelemetry-api" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/9e/de/911ac9e309052aca1b20b2d5549d3db45d1011e1a610e552c6ccdd1b64f8/opentelemetry_semantic_conventions-0.62b1.tar.gz", hash = "sha256:c5cc6e04a7f8c7cdd30be2ed81499fa4e75bfbd52c9cb70d40af1f9cd3619802", size = 145750, upload-time = "2026-04-24T13:15:52.236Z" } +sdist = { url = "https://files.pythonhosted.org/packages/8f/73/0cbdebcb4cf545fdd328da14f5137e37d0770c3f26185e478b0d15d94f50/opentelemetry_semantic_conventions-0.65b0.tar.gz", hash = "sha256:f9b2b81e9d5b64f11bc952075e7e9c7fb0aab075c7fd1c46d597f1b919852d60", size = 148774, upload-time = "2026-07-16T15:25:46.902Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/eb/a6/83dc2ab6fa397ee66fba04fe2e74bdf7be3b3870005359ceb7689103c058/opentelemetry_semantic_conventions-0.62b1-py3-none-any.whl", hash = "sha256:cf506938103d331fbb78eded0d9788095f7fd59016f2bda813c3324e5a74a93c", size = 231620, upload-time = "2026-04-24T13:15:35.454Z" }, + { url = "https://files.pythonhosted.org/packages/a6/0e/49df70d9b81fb5cbae4bbf2a49d865b09bcbcbc4eb53f5851b1027738d78/opentelemetry_semantic_conventions-0.65b0-py3-none-any.whl", hash = "sha256:1cacde7b0ad306f84c5ef08c3dbe1bbaf20165bba6f8bff43b670e555a086bcb", size = 204645, upload-time = "2026-07-16T15:25:30.688Z" }, ] [[package]] name = "opentelemetry-util-http" -version = "0.62b1" +version = "0.65b0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/24/1b/aa71b63e18d30a8384036b9937f40f7618f8030a7aa213155fb54f6f2b47/opentelemetry_util_http-0.62b1.tar.gz", hash = "sha256:adf6facbb89aef8f8bc566e2f04624942ba08a7b678b3479a91051a8f4dc70a3", size = 11393, upload-time = "2026-04-24T13:23:12.994Z" } +sdist = { url = "https://files.pythonhosted.org/packages/32/a9/d7525a59fdd240e69b5af4a6338e78fafa1b4203394122cbd6701fb5f84a/opentelemetry_util_http-0.65b0.tar.gz", hash = "sha256:84f82d826978bba416ab453460ff6a7391cdc3534c93a786595e4068680016b7", size = 11243, upload-time = "2026-07-16T15:26:27.898Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5d/85/a9d9d32161c1ced61346267db4c9702da54f81ec5dc88214bc65c23f4e9d/opentelemetry_util_http-0.62b1-py3-none-any.whl", hash = "sha256:c57e8a6c19fc422c288e6074e882f506f85030b69b7376182f74f9257b9261f0", size = 9295, upload-time = "2026-04-24T13:22:28.078Z" }, + { url = "https://files.pythonhosted.org/packages/23/3f/ab8d29df207ce5f470a07fa96ebb48af4e95b7fab7e7635311b9a32f2fab/opentelemetry_util_http-0.65b0-py3-none-any.whl", hash = "sha256:7553b606f963097cb190536dc30556cce85090692e471a422fff30ca29b04348", size = 8245, upload-time = "2026-07-16T15:25:46.482Z" }, ] [[package]] @@ -4656,19 +4678,22 @@ numpy = [ [[package]] name = "oracledb" -version = "3.4.2" +version = "4.0.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "cryptography" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/f7/02/70a872d1a4a739b4f7371ab8d3d5ed8c6e57e142e2503531aafcb220893c/oracledb-3.4.2.tar.gz", hash = "sha256:46e0f2278ff1fe83fbc33a3b93c72d429323ec7eed47bc9484e217776cd437e5", size = 855467, upload-time = "2026-01-28T17:25:39.91Z" } +sdist = { url = "https://files.pythonhosted.org/packages/87/ae/4576e5df7b8eadec51bb7d981a2bcc0d8387d7d6998a51d146c0886a523e/oracledb-4.0.2.tar.gz", hash = "sha256:0a380ab72853487ea2764c5df772f35026b4219868fc3eba68e193c9aea230ac", size = 881658, upload-time = "2026-07-14T17:21:28.876Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/8f/81/2e6154f34b71cd93b4946c73ea13b69d54b8d45a5f6bbffe271793240d21/oracledb-3.4.2-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:a7396664e592881225ba66385ee83ce339d864f39003d6e4ca31a894a7e7c552", size = 4220806, upload-time = "2026-01-28T17:26:04.322Z" }, - { url = "https://files.pythonhosted.org/packages/ab/a9/a1d59aaac77d8f727156ec6a3b03399917c90b7da4f02d057f92e5601f56/oracledb-3.4.2-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:0f04a2d62073407672f114d02529921de0677c6883ed7c64d8d1a3c04caa3238", size = 2233795, upload-time = "2026-01-28T17:26:05.877Z" }, - { url = "https://files.pythonhosted.org/packages/94/ec/8c4a38020cd251572bd406ddcbde98ca052ec94b5684f9aa9ef1ddfcc68c/oracledb-3.4.2-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:d8d75e4f879b908be66cce05ba6c05791a5dbb4a15e39abc01aa25c8a2492bd9", size = 2424756, upload-time = "2026-01-28T17:26:07.35Z" }, - { url = "https://files.pythonhosted.org/packages/fa/7d/c251c2a8567151ccfcfbe3467ea9a60fb5480dc4719342e2e6b7a9679e5d/oracledb-3.4.2-cp312-cp312-win32.whl", hash = "sha256:31b7ee83c23d0439778303de8a675717f805f7e8edb5556d48c4d8343bcf14f5", size = 1453486, upload-time = "2026-01-28T17:26:08.869Z" }, - { url = "https://files.pythonhosted.org/packages/4c/78/c939f3c16fb39400c4734d5a3340db5659ba4e9dce23032d7b33ccfd3fe5/oracledb-3.4.2-cp312-cp312-win_amd64.whl", hash = "sha256:ac25a0448fc830fb7029ad50cd136cdbfcd06975d53967e269772cc5cb8c203a", size = 1794445, upload-time = "2026-01-28T17:26:10.66Z" }, + { url = "https://files.pythonhosted.org/packages/a5/89/2568c2d32afb3c0a66cef2e74dac7769f69e117a14402301b106970475dd/oracledb-4.0.2-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:cbe3a75b463a334da9b3d76026062bcfb24ec563b7df291f700c9f7eefcbeefe", size = 4366002, upload-time = "2026-07-14T17:21:58.951Z" }, + { url = "https://files.pythonhosted.org/packages/9c/c8/1825d240aa68b255eb78867c62964d2a69c6ae6197ab3543fe220539bf69/oracledb-4.0.2-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9d11723018b6aeae035f4e1feae4150b7cca1161023c2d68d957d7310fe0dd56", size = 2286174, upload-time = "2026-07-14T17:22:00.649Z" }, + { url = "https://files.pythonhosted.org/packages/10/31/5b9d6fa28942ff38e2d76dcef3184e7569d4b0006d0425f92a3d4a764dc4/oracledb-4.0.2-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:579f2c568433523a990cde5bea73c980d144754dc54d3ab2cd37efd670dc31d6", size = 2490198, upload-time = "2026-07-14T17:22:02.322Z" }, + { url = "https://files.pythonhosted.org/packages/28/4e/77b21ec50c786270a7471abd6f48ec3c2e629ce304ec792eeeb6f03ba4d0/oracledb-4.0.2-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:b82d3d83bdc6246e53f8ad8b598d2508e9f5d983e352ce2e8a71a020200ba2d6", size = 2330058, upload-time = "2026-07-14T17:23:04.213Z" }, + { url = "https://files.pythonhosted.org/packages/24/83/834e07805b8b3aaab7e87b818ef44ab3cb5394a73149a580ecf73cd92d4d/oracledb-4.0.2-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:b4f2b51837248d4acfb323a64c48d89b62467b0f73d1211c99db9b80df0cdfbc", size = 2510815, upload-time = "2026-07-14T17:23:05.838Z" }, + { url = "https://files.pythonhosted.org/packages/49/10/f05a3348a50ecab07b86317424751a4f63891e31c228a6d95f92b9538afd/oracledb-4.0.2-cp312-cp312-win32.whl", hash = "sha256:b5d203426f0d191842b4cfd87c223163ccc57f7e7875ec7ed073f163de2229af", size = 1488879, upload-time = "2026-07-14T17:23:07.53Z" }, + { url = "https://files.pythonhosted.org/packages/85/e2/99fe3fa29466533df10fdfed90faee133aa1e7147b9edfc13b839e06f738/oracledb-4.0.2-cp312-cp312-win_amd64.whl", hash = "sha256:5fe6e07ed29f84a6656e0580b4e662109d3bb90fc13d8f98ae3a56a67c75b80e", size = 1865771, upload-time = "2026-07-14T17:23:09.374Z" }, + { url = "https://files.pythonhosted.org/packages/87/cb/009980df826442900419a7bc317e35dc85a031bc2d3806256d469de7d670/oracledb-4.0.2-cp312-cp312-win_arm64.whl", hash = "sha256:a1f46b01a089e0e0dd44c4ac4c5c79903760976346c74a698955c1fb76d68bf5", size = 1519969, upload-time = "2026-07-14T17:23:11.531Z" }, ] [[package]] @@ -4811,14 +4836,11 @@ sqlalchemy = [ [[package]] name = "pgvector" -version = "0.4.2" +version = "0.5.0" source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "numpy" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/25/6c/6d8b4b03b958c02fa8687ec6063c49d952a189f8c91ebbe51e877dfab8f7/pgvector-0.4.2.tar.gz", hash = "sha256:322cac0c1dc5d41c9ecf782bd9991b7966685dee3a00bc873631391ed949513a", size = 31354, upload-time = "2025-12-05T01:07:17.87Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a7/ec/6eb80aebc728200f95229219882994c1b0585b956ca47da5edb9d062627a/pgvector-0.5.0.tar.gz", hash = "sha256:07a9dcf735696879406983afc6eba9a787cef7c0cf6c367ca1a5779f036dee74", size = 35170, upload-time = "2026-07-06T18:27:27.767Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5a/26/6cee8a1ce8c43625ec561aff19df07f9776b7525d9002c86bceb3e0ac970/pgvector-0.4.2-py3-none-any.whl", hash = "sha256:549d45f7a18593783d5eec609ea1684a724ba8405c4cb182a0b2b08aeff04e08", size = 27441, upload-time = "2025-12-05T01:07:16.536Z" }, + { url = "https://files.pythonhosted.org/packages/5c/e4/a5573f2c579ca9ad133293bfb624148ba0893674ca4a6eeec85ced9a6a09/pgvector-0.5.0-py3-none-any.whl", hash = "sha256:fedc9800894e6da2be51358d7b7c574bf34f247ca741a5a09513622135f5964f", size = 30958, upload-time = "2026-07-06T18:27:26.797Z" }, ] [[package]] @@ -5314,16 +5336,16 @@ wheels = [ [[package]] name = "pymochow" -version = "2.4.0" +version = "2.4.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "future" }, { name = "orjson" }, { name = "requests" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1d/06/ba1b9ad8939a7289196df73934eb805bdd3e38473ccf2edcc06018f156c5/pymochow-2.4.0.tar.gz", hash = "sha256:63d9f9abc44d3643b4384fd233005978a0079b45bbb35700a81ccb99c1442cfd", size = 51300, upload-time = "2026-04-02T10:24:11.883Z" } +sdist = { url = "https://files.pythonhosted.org/packages/1c/17/16e8db9d6299cbb85a73bee162601ad2d4bd619567f0b71fef8f710b51af/pymochow-2.4.1.tar.gz", hash = "sha256:e53a700337c9e538bdd62597e13c05532136451a7722459b90474ce38c7c62b0", size = 51571, upload-time = "2026-05-09T10:24:59.328Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f3/f8/d3c23f0e1d15c66ce3e431cf1866309c375c0685ff0ed6e4ae21f72161b2/pymochow-2.4.0-py3-none-any.whl", hash = "sha256:52d128aa9bea643f51aded91fed99af4d6421922e7696dfe9a1877684469d172", size = 79149, upload-time = "2026-04-02T10:24:10.029Z" }, + { url = "https://files.pythonhosted.org/packages/e8/ec/c8ec3d989c7391e58403ad670d996611b630e4969b74298cc977442425ac/pymochow-2.4.1-py3-none-any.whl", hash = "sha256:a8c7b269e9b5997347b15806c953b62b1305c6a770f4f83af3f96199498252c3", size = 79391, upload-time = "2026-05-09T10:24:56.259Z" }, ] [[package]] @@ -5386,11 +5408,11 @@ wheels = [ [[package]] name = "pypdf" -version = "6.14.2" +version = "6.15.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/03/72/7dfd5ff1c9c37de97a731701f51af091325f123d9d4270361c9c69e4431f/pypdf-6.14.2.tar.gz", hash = "sha256:7873f502fe4385e79539b21d872392dc0c4e3714327c15881cbc7fbfd1f95b25", size = 6491182, upload-time = "2026-06-23T14:18:30.859Z" } +sdist = { url = "https://files.pythonhosted.org/packages/17/17/ee75a92718ec7212de831e71454d702225aa5e474a805cce169806044453/pypdf-6.15.0.tar.gz", hash = "sha256:d39c4d955a76409284a905e2d65b40076d77ab76129e0faaeeb6612403ecfc79", size = 6993794, upload-time = "2026-08-06T13:06:49.929Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/49/e6/136aa8993a2ae7214e0b0ef2edaa0d2e08d1d4e4982635b08a835ff31ec8/pypdf-6.14.2-py3-none-any.whl", hash = "sha256:3f07891af76dc002657e04993ab9b4de81de29f9013b9761d0b7968bff12e946", size = 349514, upload-time = "2026-06-23T14:18:28.867Z" }, + { url = "https://files.pythonhosted.org/packages/af/72/ce3067ac31e214a66388159f8462ddb8c13dd00170f24d555a1f1ae8ee91/pypdf-6.15.0-py3-none-any.whl", hash = "sha256:14e001d6504822cb1ca9c7ed9a69bccb320f59b320730f55af804361abe4d5ee", size = 378123, upload-time = "2026-08-06T13:06:47.709Z" }, ] [[package]] @@ -5448,19 +5470,21 @@ wheels = [ [[package]] name = "pyrefly" -version = "1.0.0" +version = "1.2.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/9f/3a/9045b0097ac58979c7c30a4fa0e673db942d4adbc7b6d439bd54ae58c441/pyrefly-1.0.0.tar.gz", hash = "sha256:5c2b810ffcebd84be71de5df1223651edee951653a66935c6f091e957c452455", size = 5677995, upload-time = "2026-05-12T20:12:46.812Z" } +sdist = { url = "https://files.pythonhosted.org/packages/89/01/a86e9f24722b095c3f88e3616132b75a21b0df53804bdc6a45314dd4d93c/pyrefly-1.2.0.tar.gz", hash = "sha256:5485f960fc2481617068c918335c39ab1507ef90b6b5bd35bf57726e60e73185", size = 6243654, upload-time = "2026-08-01T02:56:27.592Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f4/c6/90788819bac9c61dd7bacba53b79f3c12d47ccbe5e51b3d6d89f2387e1d2/pyrefly-1.0.0-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:e355a0908555348ed4b9585ef25c76ff566673e345c866c325f1633f44d890b6", size = 13122950, upload-time = "2026-05-12T20:12:20.711Z" }, - { url = "https://files.pythonhosted.org/packages/82/91/a3cf2a1e87d336eaa804a1e6fc93266faf6dc2a97eecdbc7eae289628022/pyrefly-1.0.0-py3-none-macosx_11_0_arm64.whl", hash = "sha256:a7038efc3a40f8294edee339895633cf22db268c0d434cdbcbefc34f78a9ecc3", size = 12599494, upload-time = "2026-05-12T20:12:23.495Z" }, - { url = "https://files.pythonhosted.org/packages/cd/ab/74d1e11e737e99b1c003ecc5d7d2e846c4ea1f328966bfdbbd0ac63fad0a/pyrefly-1.0.0-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:da331ca515ed1c08791da2b5f664cf9c1294c48fd802133262e7d5d51e0f4416", size = 12995507, upload-time = "2026-05-12T20:12:25.951Z" }, - { url = "https://files.pythonhosted.org/packages/7c/ac/2df0899f8464c97e5d995f994c97c5cb5b0f58610432aa90d26d924e1db5/pyrefly-1.0.0-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:c74219d8f3e63cdaa5501a0b21d1c9d37011820f9606728d0ed06f09ae86a878", size = 13947693, upload-time = "2026-05-12T20:12:29.188Z" }, - { url = "https://files.pythonhosted.org/packages/6b/3e/b247c24321e36f04b7d51f9ccf3df93e5009e4b29939524b36ec2e17dc2a/pyrefly-1.0.0-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:c0d05543b1bb6ee6d64149eb5d6b2fb15aa72d3962d6a97abca0afaca8b0c131", size = 13925803, upload-time = "2026-05-12T20:12:31.904Z" }, - { url = "https://files.pythonhosted.org/packages/61/16/cfa2d61a4aa1e1f7bca48bb37acd01c6a09db4864b16a54f9587092765ff/pyrefly-1.0.0-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:1382d5b1fcdb49a4de9f34d112d2bddf290a78ff93ee8149492ad5f1077ddffc", size = 13470398, upload-time = "2026-05-12T20:12:35.302Z" }, - { url = "https://files.pythonhosted.org/packages/cb/2b/6372c7dddb326223e24a46b17efd0d4bd7b4fe22c821e523157577eed2d2/pyrefly-1.0.0-py3-none-win32.whl", hash = "sha256:aa8b5d0e47080e3202a2547b39f7a5a61d2c781c712b3b67884f745ca2c759d2", size = 12222643, upload-time = "2026-05-12T20:12:38.618Z" }, - { url = "https://files.pythonhosted.org/packages/be/ad/1d23be700b6b2ddaeb362360c7145917a8edbbf7240ae428d40541772fce/pyrefly-1.0.0-py3-none-win_amd64.whl", hash = "sha256:c8abcb0f2082e83c890375128f9cff4aa4d3f210b85eea7b3046c1ae764e77f5", size = 13146369, upload-time = "2026-05-12T20:12:41.423Z" }, - { url = "https://files.pythonhosted.org/packages/8c/38/16589134f3012fd097a10dcc85771555f1a5fb76e04b682597180743af30/pyrefly-1.0.0-py3-none-win_arm64.whl", hash = "sha256:d150fa9e40e8392832be81c3bcfc0497c146674ce4d0f8e04e1ec29e775ffb8c", size = 12538326, upload-time = "2026-05-12T20:12:43.996Z" }, + { url = "https://files.pythonhosted.org/packages/7d/9d/3c0ef1d4843987b22f996ed381ec9cf5a3b1273e29804db276252e4c95eb/pyrefly-1.2.0-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:7f46d983ac49ddd2b043694960a01dc6a19a5cfd8eec609d6bd9c42866f91b4e", size = 14026305, upload-time = "2026-08-01T02:56:02.611Z" }, + { url = "https://files.pythonhosted.org/packages/0a/06/03bbb78fbea54cdc65b626619f3597d5611aca4fdef11e72a4e8360e7e63/pyrefly-1.2.0-py3-none-macosx_11_0_arm64.whl", hash = "sha256:756f669b5555090f5c1a4fef30db1785fabe657764f7e4e6dc88994dfb8ca82d", size = 13463880, upload-time = "2026-08-01T02:56:04.93Z" }, + { url = "https://files.pythonhosted.org/packages/13/5a/7d8bc00a38e93bbc9c3e7bd14d305f7948717e667c9bcddeab9dd42fd255/pyrefly-1.2.0-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:e3465812ce5ef4781fb592edbf2724547296f0a3124be115d73c7e8b2401862d", size = 13907329, upload-time = "2026-08-01T02:56:07.104Z" }, + { url = "https://files.pythonhosted.org/packages/be/94/9e08b4bf799d0b8f36b55a2783c7ba5f51730cf0632a85a67b5b5ed876cd/pyrefly-1.2.0-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:5de7b2ad2bba5c8055181681a84b74143eac2234a48ba5d1b7ed7e7a722b02bd", size = 15039020, upload-time = "2026-08-01T02:56:09.208Z" }, + { url = "https://files.pythonhosted.org/packages/5b/bd/bca5fd0c80f4daf8ee6903a29df9f3de1feb05ff0946b8f35ec8c5096b13/pyrefly-1.2.0-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:25822ea9505f589ea8a725e4268b475132fb89e038fbf092e446510443ac142a", size = 14986199, upload-time = "2026-08-01T02:56:11.924Z" }, + { url = "https://files.pythonhosted.org/packages/97/f7/f07087f3d185ad2eced0c56cef89ca5474dfb4ff25f146cd50a861c97553/pyrefly-1.2.0-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:90efe75e17491ef5d636e10469e9278d7d0256b3b4c5e1f4750069bf3ae0f5d1", size = 14393715, upload-time = "2026-08-01T02:56:14.143Z" }, + { url = "https://files.pythonhosted.org/packages/d3/70/0d142c320e284b9e3ce35e9b1e58b8ce2ee1f578f2a7234bc30e5022b94f/pyrefly-1.2.0-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:368aaf7eee4f511ddc0f8e564cf14e01ab2f10b0db9105c6d5b153bf498d07bf", size = 13933008, upload-time = "2026-08-01T02:56:16.525Z" }, + { url = "https://files.pythonhosted.org/packages/5d/e8/e84f11b6e1f63fd453ad3654213b9a0f6f4de8cef6b58038eef2d0d5955d/pyrefly-1.2.0-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:d52d5da7bc65fb7675fbaa80eda879d4f8787c494f04cac21603330d3abbdbbe", size = 14431827, upload-time = "2026-08-01T02:56:18.645Z" }, + { url = "https://files.pythonhosted.org/packages/0f/06/810d31380f66c75e1c0779a408d3b16117b1b368b57894f6aa66bef21686/pyrefly-1.2.0-py3-none-win32.whl", hash = "sha256:8c90751de8506d938e8f802659c74cf35bd7a0036510ee6c634a38eebb280bfa", size = 13229447, upload-time = "2026-08-01T02:56:20.921Z" }, + { url = "https://files.pythonhosted.org/packages/ed/98/4dafa3c7a1caed2dc8cc708dde09ba27963c7736508f55b626fff3024113/pyrefly-1.2.0-py3-none-win_amd64.whl", hash = "sha256:8a8964c224ccc4882730130955815de21ff443c1ac3f0b90685b19bf63848170", size = 14087387, upload-time = "2026-08-01T02:56:23.188Z" }, + { url = "https://files.pythonhosted.org/packages/1b/1c/df3cb0a2e5591660ded7a1836cd2f29dc48c91adb1c0a3a700a96f6d09e1/pyrefly-1.2.0-py3-none-win_arm64.whl", hash = "sha256:3a90bb8df39dfbac74b1f3b2e9d7c526b8f80568884c3944d955023a73ebf61e", size = 13430873, upload-time = "2026-08-01T02:56:25.425Z" }, ] [[package]] @@ -6366,7 +6390,7 @@ wheels = [ [[package]] name = "tablestore" -version = "6.4.4" +version = "6.4.8" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "aiohttp" }, @@ -6379,9 +6403,9 @@ dependencies = [ { name = "six" }, { name = "urllib3" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/2f/bc/84d7592188b060950f4fe5713eb3b03068d42b2e43ad37decdb5242c1879/tablestore-6.4.4.tar.gz", hash = "sha256:0f40834030aff0c67e568b09deaab97144229b569710d66557edf7a06a5dcb19", size = 5076731, upload-time = "2026-04-09T09:40:20.399Z" } +sdist = { url = "https://files.pythonhosted.org/packages/b8/e5/e0babae5e9b591080b1a229a0d25a13323025e53b3d708aac7b3be301758/tablestore-6.4.8.tar.gz", hash = "sha256:772b8505f324e6bad81b27165ce3c0d387ecf843c5c9c71cb7424d8af8041d04", size = 5114056, upload-time = "2026-07-10T09:26:47.36Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/45/3f/48af1e72e59d60481724b326317bd311615bdedc31f8f81f9508fb84cda6/tablestore-6.4.4-py3-none-any.whl", hash = "sha256:984f086fa7acabaa3558da93205ad6df562b266b85fd249bc5891f2dd1d65814", size = 5118758, upload-time = "2026-04-09T09:40:17.209Z" }, + { url = "https://files.pythonhosted.org/packages/28/a0/e2b34ebc075ead0081506952319ed667b1cb23fa9542eb88ed7dc445db5f/tablestore-6.4.8-py3-none-any.whl", hash = "sha256:6f87d7410569c16cd011acd5c0501f75d1f66c6ba8005218649221bbf147838e", size = 5153545, upload-time = "2026-07-10T09:26:45.307Z" }, ] [[package]] @@ -6650,6 +6674,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/7c/79/d84de200a80085b32f12c5820d4fd0addcbe7ba6dce8c1c9d8605e833c8e/types_beautifulsoup4-4.12.0.20250516-py3-none-any.whl", hash = "sha256:5923399d4a1ba9cc8f0096fe334cc732e130269541d66261bb42ab039c0376ee", size = 16879, upload-time = "2025-05-16T03:09:09.051Z" }, ] +[[package]] +name = "types-bleach" +version = "6.4.0.20260728" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "types-html5lib" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/c0/58/b9f2678dd538de3aa8613825986dab68608447a25bd32aa336e65a6e2cdc/types_bleach-6.4.0.20260728.tar.gz", hash = "sha256:526a4fec997514bf3fe989ef59b86ec694bbb2ed70db22c59b622eb5a9b80fae", size = 11785, upload-time = "2026-07-28T04:52:17.359Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/3c/3d/ac2809b7afe8ac7c907faacdd60bd3d194de0802d809b82327e969bc05da/types_bleach-6.4.0.20260728-py3-none-any.whl", hash = "sha256:ea9feed53196b24a6d9f0cb471c1cfbb0788b3194e8335dc8da1b951b4965d7f", size = 12012, upload-time = "2026-07-28T04:52:16.344Z" }, +] + [[package]] name = "types-cachetools" version = "7.0.0.20260503" @@ -6680,6 +6716,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/b9/65/d03948be8ae9362ad26f36443eab051fe5524295fe008126cd65792f9833/types_colorama-0.4.15.20260408-py3-none-any.whl", hash = "sha256:7327a51c760d94f7df2e8c72c275a4468c03c3abb606d23995cb37e3d24d9132", size = 10763, upload-time = "2026-04-08T04:28:30.688Z" }, ] +[[package]] +name = "types-croniter" +version = "6.2.4.20260711" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/c0/e1/ef8900515a103efa41b29ce81e59efb18e72bc9642139c5a3a17f6022641/types_croniter-6.2.4.20260711.tar.gz", hash = "sha256:807793e109f4e0d51ec18f6277a1a5f1d133e425ac6bee3a7c3254ab1793e0ea", size = 12164, upload-time = "2026-07-11T04:51:18.624Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/c8/3f/896a98180a55f855bd068b97e8c99fd4ce1885a4b43f68fd7b29641084f6/types_croniter-6.2.4.20260711-py3-none-any.whl", hash = "sha256:601b68139a47fe1b2177f433132cc81fd91155f3059a49c5cb63404c389dff01", size = 9746, upload-time = "2026-07-11T04:51:17.662Z" }, +] + [[package]] name = "types-defusedxml" version = "0.7.0.20260408" @@ -6908,6 +6953,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f1/d8/45c6c04924086e8856e7f9a33a38ee713992d9ae9cd6d449de97badcba3c/types_python_http_client-3.3.7.20260408-py3-none-any.whl", hash = "sha256:3f310282e0fe2a18c5291f935538e1f97b9f80d2c5571aad155e66806719017c", size = 8851, upload-time = "2026-04-08T04:27:08.877Z" }, ] +[[package]] +name = "types-pytz" +version = "2026.3.1.20260727" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/b1/cf/eae96172a036b942e0ad0ec49512108b4ef4b97cd6c2aed5677540a597e7/types_pytz-2026.3.1.20260727.tar.gz", hash = "sha256:4364075b6867dd15b210bb8c1d29727d609917129b45600defe1d4b3eda5ecb9", size = 10914, upload-time = "2026-07-27T05:36:32.389Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9c/24/a654176875944981c75516857cd8750600eeb9ff7fc07a2b82573619a57c/types_pytz-2026.3.1.20260727-py3-none-any.whl", hash = "sha256:42ac44e83645bfceb46597342c8fb8ec028a1407e2feae9c5c539986eab6b57a", size = 10132, upload-time = "2026-07-27T05:36:31.52Z" }, +] + [[package]] name = "types-pywin32" version = "311.0.0.20260408" @@ -7327,11 +7381,10 @@ wheels = [ [[package]] name = "wandb" -version = "0.28.0" +version = "0.28.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "click" }, - { name = "gitpython" }, { name = "packaging" }, { name = "platformdirs" }, { name = "protobuf" }, @@ -7341,17 +7394,17 @@ dependencies = [ { name = "sentry-sdk" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/5f/a7/683bfbd6cbade3012bc90d3e9c4cfc72dd62566195bf4c30321946d64b77/wandb-0.28.0.tar.gz", hash = "sha256:b20e5af0fe80e2e2a466b0466a1d60cedcc578dce0f036eca04f4a0adcad95b6", size = 40558332, upload-time = "2026-06-23T00:38:50.115Z" } +sdist = { url = "https://files.pythonhosted.org/packages/92/fb/8d3f96a8b143060d6fa145462d0785981373e04694e4152555ccb5d23939/wandb-0.28.1.tar.gz", hash = "sha256:870ccb1a01238b0ac07c6fd96a0810a1f79090aba04ea29f4ee012ac8327705d", size = 40578119, upload-time = "2026-07-16T18:47:05.413Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/c0/47/1723605f76c5d6446b6d0db65b83eda1599721bc8c1e65bd76cc1682b1a7/wandb-0.28.0-py3-none-macosx_12_0_arm64.whl", hash = "sha256:c3dab1205a5aca4abbad1eca08902cdba86add0edfa83d8d61b4429d0e79fa87", size = 24335272, upload-time = "2026-06-23T00:38:26.002Z" }, - { url = "https://files.pythonhosted.org/packages/81/ff/42b539bc75bc48fc86981dccde89327ba9b71504b805b9ba42cba7c26de9/wandb-0.28.0-py3-none-macosx_12_0_x86_64.whl", hash = "sha256:ae255da18726ee8e731ef82cbc85035b901a28ae14cf91604c361b44b8d44ce0", size = 25557959, upload-time = "2026-06-23T00:38:28.993Z" }, - { url = "https://files.pythonhosted.org/packages/15/55/c3db03d04aeab3726066a418b2ef6a1f8119774ee510f4fbe992f52b7472/wandb-0.28.0-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:6dbcba12ab168aa37561f2f32dcdef8713495fc25fa7d30fdc9bfb37989694dd", size = 24878557, upload-time = "2026-06-23T00:38:31.417Z" }, - { url = "https://files.pythonhosted.org/packages/d8/5d/1385ce3c219cb5bd30d4027687e3f8d25969c7dfd09adad1cbd5080e1a72/wandb-0.28.0-py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:325b2d0bd88be6eda5db10542499bad3710927f2569c81a84dc5eeaffc76825c", size = 26764727, upload-time = "2026-06-23T00:38:33.775Z" }, - { url = "https://files.pythonhosted.org/packages/00/58/23b6c17a6d3d5422b007707961c4496b2f6f892624d2910c9f7742fcc202/wandb-0.28.0-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:8954bc1c62ae43914dce2bebfd1d9957f72350f8fbb78e5cdfe2ca9b6be8a7b8", size = 25051656, upload-time = "2026-06-23T00:38:36.281Z" }, - { url = "https://files.pythonhosted.org/packages/89/67/9be00fb2db2281063af24a148636d2dd363d337317642ab5d8e93572c794/wandb-0.28.0-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:9fec6c908554c2dad33110c1312bc3028cc2e430f0679f16b84f82c8ea801e3b", size = 27074113, upload-time = "2026-06-23T00:38:38.737Z" }, - { url = "https://files.pythonhosted.org/packages/59/b1/f7a96c09cab0c5131b1e6466659b093b401e1653cbe6bb77b462fc1c361d/wandb-0.28.0-py3-none-win32.whl", hash = "sha256:8834ef3a7c8c43b701654162783caa7ad37af48a0ff06fc35d0d65a411f76ccd", size = 24525206, upload-time = "2026-06-23T00:38:42.041Z" }, - { url = "https://files.pythonhosted.org/packages/c6/c4/c7bed5e981679c74e9fbb22c03ff31c42e95f266199d03d8d325f4d0e6df/wandb-0.28.0-py3-none-win_amd64.whl", hash = "sha256:ac1f82292e2da4f98297b78c3a46726b3a6c5734ecb75fc39b8db2c8a4989159", size = 24525214, upload-time = "2026-06-23T00:38:44.549Z" }, - { url = "https://files.pythonhosted.org/packages/f0/77/b5ce9696c8cb955521a7941fbc443e78b2f504894c6ae1a2d0b1de6e12ae/wandb-0.28.0-py3-none-win_arm64.whl", hash = "sha256:c5b0faf1b84cf79ebabed77538c1940a4c6053e815f767a4004e877a1354bed1", size = 22378208, upload-time = "2026-06-23T00:38:47.148Z" }, + { url = "https://files.pythonhosted.org/packages/06/21/8df50164d07623cfcefec19bbf9327d9be84b637a827cea1f0c06db005fd/wandb-0.28.1-py3-none-macosx_12_0_arm64.whl", hash = "sha256:da909a76e65c64c0d93acc485d2a19f66e336f1e3f725f1c98a070883e084943", size = 24277925, upload-time = "2026-07-16T18:46:42.383Z" }, + { url = "https://files.pythonhosted.org/packages/8f/18/6c3da7e6cb215ad363324db8dc4d83b93626f5e339822b05b1c38a6097fd/wandb-0.28.1-py3-none-macosx_12_0_x86_64.whl", hash = "sha256:3da3db219c54bfd1082c00e9061c8ea894ba43e42733b5af00bb10c09d7158fe", size = 25480852, upload-time = "2026-07-16T18:46:45.102Z" }, + { url = "https://files.pythonhosted.org/packages/e2/1a/d15bcfb4417fa69edcaa33db8ea012db733da1057e193b047e3f69fdd671/wandb-0.28.1-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:ae9ae6fb29e2e2b1d097ed8b75c0c0240c778c2a8cad1d996dee870a1e401c2c", size = 24832138, upload-time = "2026-07-16T18:46:47.433Z" }, + { url = "https://files.pythonhosted.org/packages/b3/da/49924c7df2952dfd82c86c3779c339c0c3d6f6439387c03d97d0470c3658/wandb-0.28.1-py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:8cfb898b6a6c884d9c9294b02764e88bce65049f027a124d6bee53fe722469b6", size = 26486533, upload-time = "2026-07-16T18:46:49.839Z" }, + { url = "https://files.pythonhosted.org/packages/11/c0/06b23518e29690784f1b3081e39c7679ca076cb0af094cb9b4bb309150f5/wandb-0.28.1-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:cf2b1533945395e4fdbe6182b272bb0ca8a02c10b3086a395e2d57686ae3ed0d", size = 25022635, upload-time = "2026-07-16T18:46:52.376Z" }, + { url = "https://files.pythonhosted.org/packages/23/30/6de2f7995a8a6eecbd03d24c79a139a734c0168f5520cf4c7ccb43c1dbbc/wandb-0.28.1-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:7233061080507a4b4098bed1ccb381ce6f890c60397cd4153d060285bcb267bd", size = 27008895, upload-time = "2026-07-16T18:46:55.025Z" }, + { url = "https://files.pythonhosted.org/packages/b2/83/49deab9447687625371435ca21b6da82f223f1c7d014d77386b9cb91833c/wandb-0.28.1-py3-none-win32.whl", hash = "sha256:4bc461cda3ce23a19d8df5e42981a664d95fa3231efb10fd1e85d9d4824c7d29", size = 24418398, upload-time = "2026-07-16T18:46:57.43Z" }, + { url = "https://files.pythonhosted.org/packages/bf/6f/ed6616b11ea15b8ceabedcaa567286c1c9ec65fa50563230a90bfb627cc5/wandb-0.28.1-py3-none-win_amd64.whl", hash = "sha256:d98a10370162b1e970850237114c56e9c4c58f3cb701e4b8cb38f36f6749fd52", size = 24418404, upload-time = "2026-07-16T18:47:00.427Z" }, + { url = "https://files.pythonhosted.org/packages/07/78/75b6827a6665337a715c5347c5edbd84eca660f7a0f48d8d6d24d1f66bee/wandb-0.28.1-py3-none-win_arm64.whl", hash = "sha256:4aa07f13dd3bcac2c0524c8d0f49f76e83ab5c1054fd09f3b1a436cfcde146a6", size = 22299006, upload-time = "2026-07-16T18:47:02.71Z" }, ] [[package]] @@ -7444,19 +7497,20 @@ wheels = [ [[package]] name = "weaviate-client" -version = "4.20.5" +version = "4.22.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "authlib" }, { name = "grpcio" }, { name = "httpx" }, + { name = "packaging" }, { name = "protobuf" }, { name = "pydantic" }, { name = "validators" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/81/c8/aa47cfa0a2b1e260846eaf04ce4cc2ab1bb03f29d793e7b009bc3e3babc7/weaviate_client-4.20.5.tar.gz", hash = "sha256:c07c688f0e6b78723dfecbcfeebf897cefa75f1a89c63ebd84aab88c662e4394", size = 811866, upload-time = "2026-04-09T20:08:45.268Z" } +sdist = { url = "https://files.pythonhosted.org/packages/34/2a/73cf7d6c7c6aa638738dfcb0d318e0404aab3c0673f82a1a4d89455b21a5/weaviate_client-4.22.0.tar.gz", hash = "sha256:0c50fbef546a522262a87d1138cde0509c7a8a48e702e967be33472e9f7fbae3", size = 860126, upload-time = "2026-06-18T06:08:30.202Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/58/ba/d55f1a665802f736436d09198afc0d00806a405aadb9977193a2f009cfcb/weaviate_client-4.20.5-py3-none-any.whl", hash = "sha256:3f508e3dc08257f85230f9d2ea0562443ed0715e7e89156f22b7e950d6c08cdb", size = 620766, upload-time = "2026-04-09T20:08:43.215Z" }, + { url = "https://files.pythonhosted.org/packages/e1/a3/27353ea3fbf9e7d4f03375176a838c632d009f5167d801adcbb5f6d3bd07/weaviate_client-4.22.0-py3-none-any.whl", hash = "sha256:ff2dbc8d1fc25739942402c22d1aba2350a16ba4a6b6ed0ba140689c70adf1d9", size = 652691, upload-time = "2026-06-18T06:08:28.622Z" }, ] [[package]] diff --git a/cli/AGENTS.md b/cli/AGENTS.md index 56a59128c03..446c85abdcc 100644 --- a/cli/AGENTS.md +++ b/cli/AGENTS.md @@ -1,99 +1,26 @@ # AGENTS.md — difyctl (TypeScript CLI) -TypeScript port of difyctl. Stack: custom CLI framework (`src/framework/`), Node 22+, ESM, ky for HTTP, Vitest, and Vite+ formatting and linting. +This package is the Node 22+, ESM TypeScript implementation of `difyctl`. Development also requires the Bun version pinned in `.bun-version`; command-tree generation and the `dev`, `test`, and `build` pre-scripts invoke it. Read [`ARD.md`] before adding a command or changing shared CLI infrastructure. Read `src/commands/AGENTS.md` for command-folder and registry rules. -> Architecture patterns, scaffolding recipe, printer chain, strategy pattern, testing conventions, anti-patterns: see **[`ARD.md`]**. +## Architecture Boundaries -## Code rules +- Every leaf command extends `DifyCommand`; command classes own framework parsing and delegate behavior to domain modules. +- Each command folder keeps its framework shell in `index.ts`. Extract behavior into sibling modules such as `run.ts` and `handlers.ts` when it needs an independently testable owner; those modules receive typed dependencies and do not import `src/framework/`. +- `src/http/` owns ky middleware and client construction; `src/api/` owns resource clients; `src/sys/io/` owns process streams and progress UI; `src/types/` remains a pure data and schema leaf. +- Preserve flags, output, and exit codes during refactors. Do not add dependencies or compatibility shims unless the task explicitly requires them. +- `ARD.md` owns CLI code structure. Keep wire behavior aligned with typed API clients and the real mock-server behavior tests. -- **Spaces, not tabs.** -- **Minimum comments.** Code speak for self. Comment only non-obvious WHY — hidden constraints, subtle invariants, bug-workaround notes. Never restate code. Never reference tasks, PRs, current callers. -- **No magic strings or numbers.** Enums or named constants for bounded value sets. -- **No long positional arg lists.** Use options objects. -- **No long if/switch ladders on discriminator.** Polymorphism, dispatch tables, or strategy pattern. Name concept, let implementations plug in. -- **No `any`. No `unknown` outside genuine wire boundaries** (HTTP body parse, env vars). Narrow types everywhere else. -- **Avoid `!` non-null assertions.** Narrow instead. -- **`readonly` on inputs not mutated.** -- **Discriminated unions** for variant data (SSE events, run outputs, error shapes), not optional-field bags. -- **No backwards-compat shims.** No re-exports of old names, no `// removed:` markers, no deprecation notes. Delete, update callers. -- **No new dependencies without explicit approval.** -- **No CLI behavior changes in refactor commit.** Same flags, same output, same exit codes. -- **Every leaf command extends `DifyCommand`.** Add `static agentGuide` string when command benefits from agent workflow docs — see `src/commands/AGENTS.md`. +## Commands -## Layering +Run package scripts from `cli/`: -| Layer | Path | Role | -| --------- | -------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------- | -| commands | `src/commands/` | Command class shells (extend `DifyCommand`). Only place framework imports run. | -| domain | `src/run/`, `src/get/`, etc. | Plain TS modules. Take typed deps via options. Testable without the framework. | -| api | `src/api/` | One typed client per resource. Each takes `KyInstance`. | -| http | `src/http/` | `createClient` + middleware (auth, retry, logging, error mapping). Only place ky runs. | -| io | `src/io/` | Streams + spinner. Fence between data-out and progress UI. | -| printers | `src/printers/` | `CompositePrintFlags` + `-o {json,yaml,name,wide,text}` matrix. | -| errors | `src/errors/` | `BaseError`, `ErrorCode` enum, `ExitCode` enum, dispatch table, `formatErrorForCli`. | -| guide | `src/commands/**//guide.ts` | Per-command agent guide string. Export `agentGuide`, assign `static agentGuide = agentGuide` in command class. Surfaced via `--help`. | -| cache | `src/cache/` | On-disk caches (app-info, etc.). | -| auth | `src/auth/` | Hosts file, token store, login flow. | -| config | `src/config/` | XDG dir resolution, config.yml load/save. | -| workspace | `src/workspace/` | Resolver: flag → env → bundle. | -| types | `src/types/` | Pure data + zod schemas for server contracts. No runtime imports outward. | +- Source CLI: `pnpm dev [args...]` +- Tests: `pnpm test` +- Build: `pnpm build` +- Regenerate and verify the registry: `pnpm tree:gen` and `pnpm tree:check` -## Command Structure +Run the scoped static check from the repository root with `vp check cli`. -Scaffold recipe + checklist: see `ARD.md §New command scaffold`. Full folder convention (subcommands, guide.ts): see `src/commands/AGENTS.md`. - -Layer rules: - -- Commands thin shells. Use `this.authedCtx(opts)` for bearer context; delegate to domain function. -- Domain receives deps via options; never imports `src/framework/`. -- Only `src/http/client.ts` and `src/api/*` import ky at runtime; elsewhere use `import type { KyInstance }`. -- `process.*` lives in `src/io/`, `src/store/dir.ts`, `src/util/browser.ts`. Nowhere else. -- No circular imports. `types/` pure leaf. - -## Dev commands - -```sh -pnpm install # one-time -pnpm dev [args...] # run CLI from source (no -- separator) -pnpm test # vitest -pnpm test:coverage # with coverage -pnpm -w check # repository-wide static check -pnpm -w check:fix # repository-wide static fixes -pnpm build # production bundle (vp pack) -pnpm tree:gen # regenerate src/commands/tree.ts (registry) -pnpm tree:check # verify tree.ts is up-to-date with the fs -``` - -Release binaries (5 platform targets, Bun-compiled) are produced by `pnpm build:bin` (called from `.github/workflows/cli-release.yml`). - -## Tests - -- Behavior tests run against real Hono mock at `test/fixtures/dify-mock/`. No `nock`, `msw`, or `fetchMock` — every test exercises real HTTP. -- Test files co-located: `foo.test.ts` next to `foo.ts`. -- The repository-wide static check and full test suite must be green before any commit. - -## Spec docs (`docs/specs/`) - -Behavior contracts. Living tree — amended in place, no version subfolders. - -**Keep:** HTTP wire shape (req/resp JSON, headers, status codes), SQL DDL, Redis keys + TTL, state transitions, audit event names + payload, error/exit codes, rate-limit values, JWS/cookie envelope claims. - -**Cut:** language type decls, internal helper sigs, decorator snippets, file-path tables, pseudocode mirroring code, "Open items"/"Handler walk"/"CI guard"/"Migration" sections, rationale (`Rejected:`/`Why X not Y`/`Historical note:`/product comparisons), release-pipeline lines, version-pinning (`in v1.0`, `post-v1.0`, milestone codes), frontmatter `date`/`status`/`author`. - -**Test:** "rewrite in Rust tomorrow, does spec hold?" HTTP/SQL/Redis stays; type defs go. - -**Rules:** behavior, not rationale. One topic per file; cross-refs = `auth.md §Storage`. Tables beat prose. Code wins on drift — update spec. - -## Out of scope for unrelated work - -Do not modify in passing: - -- `test/fixtures/dify-mock/` public surface (endpoints, JSON shapes, status codes, scenario names) — that's the dify-api contract. -- `bin/`, `scripts/`, `Makefile`, `lint.config.ts`, `tsconfig*.json`, `package.json` (unless the change is required by the task). - -## Commits - -- One concern per commit. Style: `(): ` lowercase. Body explains why if non-obvious. -- Never push, amend, force-push, or skip hooks (`--no-verify`) without explicit user approval. +Behavior tests use the real Hono server under `test/fixtures/dify-mock/`; do not replace it with `nock`, `msw`, or `fetchMock`. Keep tests colocated with their source files. [`ARD.md`]: ARD.md diff --git a/cli/ARD.md b/cli/ARD.md index a87b01dd85f..e485e74be60 100644 --- a/cli/ARD.md +++ b/cli/ARD.md @@ -2,8 +2,6 @@ Onboarding ref for `dify/cli/` contributors. Cover canonical patterns, layer contracts, scaffolding recipe, dev workflow, anti-patterns. Read before adding command or touching shared infra. -Spec authority: [`docs/specs/`]. Specs own HTTP wire shape + server behavior; this file owns CLI code structure. - --- ## Project layout @@ -17,7 +15,7 @@ src/ config/ config.yml read/write errors/ BaseError, ErrorCode, exit codes http/ ky client factory + middleware - io/ IOStreams, spinner, printer chain + sys/io/ IOStreams, prompts, spinner, output rendering limit/ --limit flag parsing types/ shared TypeScript types util/ small pure helpers @@ -38,17 +36,17 @@ src/commands/// Examples: `get/app/`, `auth/devices/revoke/`, `describe/app/`. -**2. Mandatory files** +**2. Mandatory file** -| File | Responsibility | -| ---------- | --------------------------------------------------------------------------------------- | -| `index.ts` | `DifyCommand` subclass. Flag/arg declaration + `run()` wiring only. No business logic. | -| `run.ts` | Pure async function. Typed options + deps. Returns string. No `src/framework/` imports. | +| File | Responsibility | +| ---------- | ------------------------------------------------------------------------------------ | +| `index.ts` | `DifyCommand` subclass. Owns flag/arg parsing, framework output, and command wiring. | **3. Optional files — add as needed** | File | Purpose | | ------------------ | ------------------------------------------------------------------ | +| `run.ts` | Typed behavior owner when logic merits independent tests or reuse | | `handlers.ts` | Output types implementing `FormattedPrintable` or `TablePrintable` | | `payload-shape.ts` | Response type narrowing/transformation | | `run.test.ts` | Behavior tests against `run.ts` | @@ -58,12 +56,12 @@ Examples: `get/app/`, `auth/devices/revoke/`, `describe/app/`. - [ ] `index.ts` extends `DifyCommand` - [ ] Authed command calls `this.authedCtx()`; non-authed skips -- [ ] No try/catch in `run()` — `DifyCommand.catch()` handles `BaseError` -- [ ] `run.ts` returns string; no direct stdout write -- [ ] `run.ts` no `src/framework/` imports +- [ ] Let the command boundary handle `BaseError`; catch only when the command owns recovery +- [ ] Keep framework parsing and output construction in `index.ts` +- [ ] When present, `run.ts` returns typed behavior data or owns explicit streaming/interactive I/O and does not import `src/framework/` - [ ] HTTP client via factory dep, not direct -- [ ] `run.test.ts` written before impl (test-first) -- [ ] `pnpm tree:gen` run after adding command (updates `src/commands/tree.ts`) +- [ ] Add focused behavior tests when the command changes an observable contract +- [ ] `pnpm tree:gen` run after adding command (updates `src/commands/tree.generated.ts`) - [ ] README command table updated by hand --- @@ -74,27 +72,22 @@ All commands extend `DifyCommand`, not `Command`. ```typescript export default class MyCommand extends DifyCommand { - async run(): Promise { + async run(argv: string[]) { const { args, flags } = this.parse(MyCommand, argv) - // Authed: authedCtx() sets outputFormat + builds context - const ctx = await this.authedCtx({ format: flags.output }) - - process.stdout.write( - await runMyThing( - { - // args - }, - { bundle: ctx.bundle, http: ctx.http, io: ctx.io }, - ), + const ctx = await this.authedCtx({ retryFlag: undefined, format: flags.output }) + const result = await runMyThing( + { id: args.id }, + { active: ctx.active, http: ctx.http, io: ctx.io }, ) + return formatted({ format: flags.output, data: result.data }) } } ``` -**`authedCtx(opts)`** — wraps `buildAuthedContext`. Sets `this.outputFormat` as side effect. Required for any command needing bearer token. +**`authedCtx(opts)`** — wraps `buildAuthedContext` and returns the authenticated registry, account, HTTP, I/O, and optional cache dependencies. Pass the selected output format so authentication failures use the same serialization contract. Required for commands that need a bearer token. -**`catch(err)` override** — auto-handles `BaseError` with format-aware serialization. Never wrap `run()` in try/catch. Throw `BaseError`; base class catches. +The framework runner in `src/framework/run.ts` catches command errors, normalizes unknown failures, and serializes `BaseError` according to the selected output format. Catch inside a command only when that command owns a real recovery path. --- @@ -103,8 +96,8 @@ export default class MyCommand extends DifyCommand { Throw `BaseError`. Never throw raw `Error` for domain failures. ```typescript -import { BaseError } from '../../errors/base.js' -import { ErrorCode } from '../../errors/codes.js' +import { BaseError } from '@/errors/base' +import { ErrorCode } from '@/errors/codes' throw new BaseError({ code: ErrorCode.UsageMissingArg, @@ -113,7 +106,7 @@ throw new BaseError({ }) ``` -`ErrorCode` exhaustive const object — never use raw strings. `exitFor(code)` maps to exit codes auto. `DifyCommand.catch()` calls `formatErrorForCli` with `outputFormat` so JSON/YAML consumers get machine-readable error output. +`ErrorCode` is the exhaustive error-code object; do not scatter raw code strings. `exitFor(code)` maps it to a process exit code, and the framework runner calls `formatErrorForCli` so JSON/YAML consumers receive machine-readable errors. | Exit | Meaning | | ---- | ----------------------------------------- | @@ -122,6 +115,7 @@ throw new BaseError({ | 2 | Usage error (bad flag, missing arg) | | 4 | Auth error (not logged in, token expired) | | 6 | Version/compat error | +| 7 | Rate limited | New error code: add to `ErrorCode` + map to `ExitCode` in `codes.ts`. Never scatter exit codes inline. @@ -172,7 +166,7 @@ Output rendering separated from data fetching via protocol objects. - Data classes implement `TablePrintable` or `FormattedPrintable` from `src/framework/output`. - Streaming commands implement `StreamPrinter` from `src/framework/stream`. -- `index.ts` wraps the result with `table({format, data})` or `formatted({format, data})` and returns it; the base class calls `stringifyOutput()`. +- `index.ts` wraps the result with `table({format, data})` or `formatted({format, data})` and returns it; `src/framework/run.ts` calls `stringifyOutput()`. - Commands that write incrementally (streaming) write directly from the strategy via `deps.io.out.write(stringifyOutput(...))`. ```typescript @@ -220,51 +214,15 @@ New mode = new class + one line in picker. Singletons avoid per-call allocation. ## HTTP clients -One file per resource under `src/api/`. Each exports class wrapping `KyInstance`. +Keep resource clients under `src/api/`. They receive the shared `HttpClient` and call generated oRPC operations through `createOpenApiClient(...)` when the OpenAPI contract covers the endpoint. Reuse generated request and response types instead of duplicating wire shapes. -```typescript -export class AppsClient { - private readonly http: KyInstance - constructor(http: KyInstance) { - this.http = http - } - - async list(params: ListParams): Promise { - /* ... */ throw new Error('elided') - } - async describe(id: string, workspaceId: string, fields: string[]): Promise { - /* ... */ throw new Error('elided') - } -} -``` - -Inject via factory dep in `run.ts` for testability: - -```typescript -type GetAppDeps = { - appsFactory?: (http: KyInstance) => AppsClient -} -// default: (h) => new AppsClient(h) -``` - -Never instantiate clients in `index.ts`. +Pass `HttpClient` into behavior owners. Add a client or factory dependency only when it owns a real substitution or lifecycle boundary; behavior tests normally exercise the real client stack against `test/fixtures/dify-mock/`. Keep client construction out of `index.ts` so the command remains a framework and output boundary. --- ## Testing -**Test-first.** Write failing test, run to confirm fail, then implement. - -Tests live in `run.test.ts` alongside command. Test `run.ts` direct — never the `DifyCommand` class. - -```typescript -const io = bufferStreams() -const result = await runGetApp( - { format: 'json', appId: 'app-1' }, - { bundle, http: mockHttp, io, appsFactory: () => fakeClient }, -) -expect(JSON.parse(result).data).toHaveLength(1) -``` +Keep tests beside the owner as `*.test.ts`. When a command has a behavior module, test that public function directly for domain and protocol behavior. Test the command class or framework boundary when argument parsing, flags, help, output construction, or command wiring is the observable contract. Establish a failing case first when practical for behavior changes and bug fixes. ### dify-mock fixture server @@ -306,14 +264,14 @@ expect(JSON.parse(out).workspaces).toHaveLength(2) | `pnpm dev [args]` | Run CLI from source during dev | | `pnpm test` | Full vitest suite — run before every commit | | `pnpm test:coverage` | Coverage report | -| `pnpm -w check` | Repository-wide static check | -| `pnpm -w check:fix` | Repository-wide static fixes | +| `vp check cli` | Scoped static check from the repository root | +| `vp check --fix cli` | Scoped static fixes from the repository root | | `pnpm build` | Production bundle (`vp pack`) | -| `pnpm tree:gen` | Regenerate `src/commands/tree.ts` (registry) | -| `pnpm tree:check` | Verify `tree.ts` matches the filesystem | +| `pnpm tree:gen` | Regenerate `src/commands/tree.generated.ts` | +| `pnpm tree:check` | Verify the generated tree matches the commands | | `pnpm build:bin` | Cross-compile standalone binaries via Bun (CI) | -**`pnpm tree:gen` rule:** run after adding, removing, renaming any command. The generated `tree.ts` is the runtime command registry — stale tree causes commands to be invisible at runtime. (Runs implicitly via `prebuild`/`predev`/`pretest`.) +**`pnpm tree:gen` rule:** run after adding, removing, or renaming any command. The generated `tree.generated.ts` is the runtime command registry; a stale tree makes commands invisible at runtime. It also runs through `prebuild`, `predev`, and `pretest`. **README hand-maintained.** When adding a command, update the command table in `README.md` manually. @@ -331,7 +289,7 @@ The repository runs Vite+ Oxlint as the primary code-quality linter, an explicit | `unicorn/no-new-array` | Use `Array.from({ length: n })` not `new Array(n)` | | `noUncheckedIndexedAccess` (tsc) | `arr[i]` is `T \| undefined`; guard before use | -Run `pnpm -w check:fix` for Oxlint, ESLint, TypeScript, and Oxfmt fixes and diagnostics. +Run `vp check --fix cli` from the repository root for scoped formatting, lint, and TypeScript fixes and diagnostics. --- @@ -350,14 +308,12 @@ Run `pnpm -w check:fix` for Oxlint, ESLint, TypeScript, and Oxfmt fixes and diag | Pattern | Do instead | | -------------------------------------------------------------------- | -------------------------------------------------------------------------- | | `if (format === 'json') { ... }` in `run.ts` | Printer handler per format | -| `try { ... } catch (e) { if (isBaseError(e)) ... }` in every command | Throw `BaseError`; `DifyCommand.catch()` handles | +| `try { ... } catch (e) { if (isBaseError(e)) ... }` in every command | Throw `BaseError`; `src/framework/run.ts` normalizes and formats it | | Raw string error codes `'not_logged_in'` | `ErrorCode.NotLoggedIn` | | `enabled: !isHuman` in `runWithSpinner` | Set `outputFormat` on `IOStreams`; spinner auto-detects | | Long positional arg lists | Options struct | | `Record` dispatch map | Named singletons + picker function | | `src/framework/` import in `run.ts`, `api/`, or `auth/` | Framework imports belong in `index.ts`, `handlers.ts`, and strategies only | | `buildAuthedContext(this, opts)` in command body | `this.authedCtx(opts)` | -| `console.log` in `src/` | Return string from `run.ts`; write in `index.ts` | +| `console.log` in `src/` | Return `CommandOutput` from the command or use owned I/O for streaming | | New dependency without approval | Check first | - -[`docs/specs/`]: docs/specs/ diff --git a/cli/package.json b/cli/package.json index f1e0be53a78..8b8f8d027cc 100644 --- a/cli/package.json +++ b/cli/package.json @@ -71,7 +71,7 @@ "channel": "alpha", "compat": { "minDify": "1.16.0", - "maxDify": "1.16.0" + "maxDify": "1.16.1" }, "release": { "tagPrefix": "difyctl-v", diff --git a/cli/scripts/release-naming.mjs b/cli/scripts/release-naming.mjs index 2af172f1333..d9fcc75c441 100644 --- a/cli/scripts/release-naming.mjs +++ b/cli/scripts/release-naming.mjs @@ -114,9 +114,15 @@ function die(msg) { process.exit(1) } +// Tests point this at a fixture manifest so their assertions stay fixed while +// the real version and compat window move with every release. The name is +// mirrored in test/fixtures/pkg-manifest.ts rather than imported from here, +// because this file's shebang breaks the Windows test runner. +const PKG_PATH_ENV = 'DIFYCTL_PKG_PATH' + function loadPkg() { - const pkgUrl = new URL('../package.json', import.meta.url) - const pkg = JSON.parse(readFileSync(pkgUrl, 'utf8')) + const pkgPath = process.env[PKG_PATH_ENV] || new URL('../package.json', import.meta.url) + const pkg = JSON.parse(readFileSync(pkgPath, 'utf8')) if (!pkg.difyctl?.release) die('cli/package.json missing difyctl.release') return { version: pkg.version, diff --git a/cli/scripts/release-naming.test.ts b/cli/scripts/release-naming.test.ts index 97552ed0823..e6f36391801 100644 --- a/cli/scripts/release-naming.test.ts +++ b/cli/scripts/release-naming.test.ts @@ -1,12 +1,19 @@ import { execFileSync } from 'node:child_process' import { fileURLToPath } from 'node:url' import { describe, expect, it } from 'vitest' +import { FIXTURE_COMPAT, pkgManifestEnv } from '../test/fixtures/pkg-manifest' const SCRIPT = fileURLToPath(new URL('./release-naming.mjs', import.meta.url)) -function run(args: string[]): { code: number; stdout: string; stderr: string } { +function run( + args: string[], + env: Record = {}, +): { code: number; stdout: string; stderr: string } { try { - const stdout = execFileSync('node', [SCRIPT, ...args], { encoding: 'utf8' }) + const stdout = execFileSync('node', [SCRIPT, ...args], { + encoding: 'utf8', + env: { ...process.env, ...env }, + }) return { code: 0, stdout, stderr: '' } } catch (e) { const err = e as { status?: number; stdout?: string; stderr?: string } @@ -14,45 +21,50 @@ function run(args: string[]): { code: number; stdout: string; stderr: string } { } } -describe('release-naming compat-check (compat 1.16.0..1.16.0)', () => { +describe('release-naming compat-check', () => { + const { minDify, maxDify } = FIXTURE_COMPAT // 2.0.0 .. 2.5.0 + const pkgEnv = pkgManifestEnv() + const compatCheck = (difyVersion?: string) => + run(difyVersion === undefined ? ['compat-check'] : ['compat-check', difyVersion], pkgEnv).code + it('accepts a version inside the window', () => { - expect(run(['compat-check', '1.16.0']).code).toBe(0) + expect(compatCheck('2.3.0')).toBe(0) }) it('accepts the inclusive lower bound', () => { - expect(run(['compat-check', '1.16.0']).code).toBe(0) + expect(compatCheck(minDify)).toBe(0) }) it('accepts the inclusive upper bound', () => { - expect(run(['compat-check', '1.16.0']).code).toBe(0) + expect(compatCheck(maxDify)).toBe(0) }) it('accepts a v-prefixed tag', () => { - expect(run(['compat-check', 'v1.16.0']).code).toBe(0) + expect(compatCheck('v2.3.0')).toBe(0) }) it('rejects a version below the lower bound', () => { - expect(run(['compat-check', '1.15.9']).code).not.toBe(0) + expect(compatCheck('1.9.9')).not.toBe(0) }) it('rejects a version above the upper bound', () => { - expect(run(['compat-check', '1.16.1']).code).not.toBe(0) + expect(compatCheck('2.5.1')).not.toBe(0) }) - it('treats a prerelease of the bound as below it (1.16.0-rc1 < 1.16.0)', () => { - expect(run(['compat-check', '1.16.0-rc1']).code).not.toBe(0) + it('treats a prerelease of the lower bound as below it', () => { + expect(compatCheck(`${minDify}-rc1`)).not.toBe(0) }) - it('ignores build metadata on the bound (1.16.0+build == 1.16.0)', () => { - expect(run(['compat-check', '1.16.0+build123']).code).toBe(0) + it('ignores build metadata on the bound', () => { + expect(compatCheck(`${maxDify}+build123`)).toBe(0) }) - it('ignores build metadata when out of range (1.16.1+build still rejected)', () => { - expect(run(['compat-check', '1.16.1+build123']).code).not.toBe(0) + it('ignores build metadata when out of range', () => { + expect(compatCheck('2.5.1+build123')).not.toBe(0) }) it('requires a version argument', () => { - expect(run(['compat-check']).code).not.toBe(0) + expect(compatCheck()).not.toBe(0) }) }) @@ -67,6 +79,14 @@ describe('release-naming github-env', () => { for (const key of ['version', 'channel', 'prerelease', 'minDify', 'maxDify', 'tagPrefix']) expect(stdout).toMatch(new RegExp(`^${key}=`, 'm')) }) + + // The only assertion against the live manifest: the window must exist and be + // well-formed, whatever release it currently points at. + it('emits a well-formed compat window from the real cli/package.json', () => { + const { stdout } = run(['github-env']) + expect(stdout).toMatch(/^minDify=\d+\.\d+\.\d+$/m) + expect(stdout).toMatch(/^maxDify=\d+\.\d+\.\d+$/m) + }) }) describe('release-naming edge channel', () => { diff --git a/cli/scripts/release-r2-edge.test.ts b/cli/scripts/release-r2-edge.test.ts index 7c93133fddf..3054a8793b6 100644 --- a/cli/scripts/release-r2-edge.test.ts +++ b/cli/scripts/release-r2-edge.test.ts @@ -4,14 +4,20 @@ import { tmpdir } from 'node:os' import { join } from 'node:path' import { fileURLToPath } from 'node:url' import { describe, expect, it } from 'vitest' +import { FIXTURE_COMPAT, pkgManifestEnv } from '../test/fixtures/pkg-manifest' const SCRIPT = fileURLToPath(new URL('./release-r2-edge.mjs', import.meta.url)) +const PKG_ENV = pkgManifestEnv() + function run(args: string[]): { code: number; stdout: string; stderr: string } { try { return { code: 0, - stdout: execFileSync('node', [SCRIPT, ...args], { encoding: 'utf8' }), + stdout: execFileSync('node', [SCRIPT, ...args], { + encoding: 'utf8', + env: { ...process.env, ...PKG_ENV }, + }), stderr: '', } } catch (e) { @@ -108,7 +114,7 @@ describe('release-r2-edge manifest', () => { it('carries the compat window from package.json', () => { const { json } = buildManifest() - expect(json.compat).toEqual({ minDify: '1.16.0', maxDify: '1.16.0' }) + expect(json.compat).toEqual(FIXTURE_COMPAT) }) it('lists all 5 targets with asset name + sha256 from the checksums file', () => { diff --git a/cli/src/api/meta.test.ts b/cli/src/api/meta.test.ts index c3e081a39f9..a73ca6ecf2c 100644 --- a/cli/src/api/meta.test.ts +++ b/cli/src/api/meta.test.ts @@ -38,7 +38,7 @@ describe('MetaClient', () => { const info = await client.serverVersion() expect(info.version).toBe('') - expect(info.edition).toBe('SELF_HOSTED') + expect(info.edition).toBe('COMMUNITY') }) it('throws when the host has no Dify on it', async () => { diff --git a/cli/src/commands/AGENTS.md b/cli/src/commands/AGENTS.md index 5df83849779..6d675ef70a9 100644 --- a/cli/src/commands/AGENTS.md +++ b/cli/src/commands/AGENTS.md @@ -13,7 +13,7 @@ src/commands/ / / index.ts ← command class (extends DifyCommand; the ONLY file the registry discovers) - run.ts ← business logic (not a command, invisible to the registry) + run.ts ← optional behavior owner (not a command, invisible to the registry) handlers.ts ← helpers guide.ts ← agent guide string (optional) *.test.ts ← tests @@ -23,7 +23,7 @@ src/commands/ .ts ``` -The registry generator (`pnpm tree:gen` → `src/commands/tree.ts`) discovers +The registry generator (`pnpm tree:gen` → `src/commands/tree.generated.ts`) discovers commands only via `**/index.+(js|cjs|mjs|ts)`. All other files in command folders are invisible to the registry — add freely without glob exclusions. Folders prefixed with `_` (e.g. `_shared/`, `_strategies/`) are excluded from @@ -32,7 +32,7 @@ registry discovery and from coverage checks. ## Adding a new command 1. Create `src/commands///index.ts` extending `DifyCommand`. -1. Add business logic in sibling files (e.g. `run.ts`, `handlers.ts`). +1. Keep small owner-local behavior in `index.ts`; extract sibling modules such as `run.ts` or `handlers.ts` when logic needs independent tests, reuse, or a clearer owner. 1. Run `pnpm tree:gen` to regenerate the command tree (also runs implicitly via `prebuild`/`predev`/`pretest`). 1. Run `pnpm test` to verify coverage. @@ -54,7 +54,9 @@ registry discovery and from coverage checks. import { agentGuide } from './guide.js' export default class MyCmd extends DifyCommand { - static agentGuide = agentGuide + override agentGuide(): string { + return agentGuide + } } ``` 1. The guide appears at the bottom of `difyctl --help` automatically. diff --git a/cli/src/version/enforce.test.ts b/cli/src/version/enforce.test.ts index 92c9d51022d..cca7b33b6e9 100644 --- a/cli/src/version/enforce.test.ts +++ b/cli/src/version/enforce.test.ts @@ -19,7 +19,7 @@ function fakeStore(fresh = false): CompatStore & { readonly marked: string[] } { } } -const server = (version: string): ServerVersionResponse => ({ version, edition: 'SELF_HOSTED' }) +const server = (version: string): ServerVersionResponse => ({ version, edition: 'COMMUNITY' }) describe('enforceDifyVersion', () => { it('throws version_skew (exit 6) when the server is too old, and never caches it', async () => { diff --git a/cli/src/version/nudge.test.ts b/cli/src/version/nudge.test.ts index 53038a28b46..04816689b0f 100644 --- a/cli/src/version/nudge.test.ts +++ b/cli/src/version/nudge.test.ts @@ -15,7 +15,7 @@ const fixedNow = () => NOW type Probe = (host: string) => Promise -const UNSUPPORTED: ServerVersionResponse = { version: '99.0.0', edition: 'SELF_HOSTED' } +const UNSUPPORTED: ServerVersionResponse = { version: '99.0.0', edition: 'COMMUNITY' } const COMPATIBLE: ServerVersionResponse = { version: '1.6.4', edition: 'CLOUD' } function emitterSpy() { @@ -122,7 +122,7 @@ describe('maybeNudgeCompat', () => { it('does not warn when server version yields unknown verdict', async () => { const probe = vi.fn( - async () => ({ version: '', edition: 'SELF_HOSTED' }) as ServerVersionResponse, + async () => ({ version: '', edition: 'COMMUNITY' }) as ServerVersionResponse, ) const { emit, lines } = emitterSpy() diff --git a/cli/src/version/probe.test.ts b/cli/src/version/probe.test.ts index 232be7a5fd1..688c1d1a70c 100644 --- a/cli/src/version/probe.test.ts +++ b/cli/src/version/probe.test.ts @@ -111,7 +111,7 @@ describe('runVersionProbe', () => { const report = await runVersionProbe({ skipServer: false, loadActive: async () => active(), - probe: async () => ({ version: '99.0.0', edition: 'SELF_HOSTED' }), + probe: async () => ({ version: '99.0.0', edition: 'COMMUNITY' }), }) expect(report.server.reachable).toBe(true) @@ -122,7 +122,7 @@ describe('runVersionProbe', () => { const report = await runVersionProbe({ skipServer: false, loadActive: async () => active(), - probe: async (): Promise => ({ version: '', edition: 'SELF_HOSTED' }), + probe: async (): Promise => ({ version: '', edition: 'COMMUNITY' }), }) expect(report.server.reachable).toBe(true) @@ -149,7 +149,7 @@ describe('runVersionProbe', () => { const report = await runVersionProbe({ skipServer: false, loadActive: async () => active({ host: 'localhost:5001', scheme: 'http' }), - probe: async () => ({ version: '1.6.4', edition: 'SELF_HOSTED' }), + probe: async () => ({ version: '1.6.4', edition: 'COMMUNITY' }), }) expect(report.server.endpoint).toBe('http://localhost:5001') diff --git a/cli/src/version/render.test.ts b/cli/src/version/render.test.ts index 5ae33ded085..46033c0358c 100644 --- a/cli/src/version/render.test.ts +++ b/cli/src/version/render.test.ts @@ -136,7 +136,7 @@ describe('renderVersionText', () => { endpoint: 'https://cloud.dify.ai', reachable: true, version: '99.0.0', - edition: 'SELF_HOSTED', + edition: 'COMMUNITY', }, compat: { minDify: '1.6.0', @@ -185,7 +185,7 @@ describe('renderVersionText', () => { endpoint: 'https://cloud.dify.ai', reachable: true, version: '99.0.0', - edition: 'SELF_HOSTED', + edition: 'COMMUNITY', }, compat: { minDify: '1.6.0', diff --git a/cli/test/fixtures/dify-mock/server.ts b/cli/test/fixtures/dify-mock/server.ts index ee9bfe3a00a..42f4d561f6b 100644 --- a/cli/test/fixtures/dify-mock/server.ts +++ b/cli/test/fixtures/dify-mock/server.ts @@ -150,9 +150,9 @@ export function buildApp(getScenario: () => Scenario, state?: MockState): Hono { app.get('/openapi/v1/_version', (c) => { const scenario = getScenario() - if (scenario === 'server-version-empty') return c.json({ version: '', edition: 'SELF_HOSTED' }) + if (scenario === 'server-version-empty') return c.json({ version: '', edition: 'COMMUNITY' }) if (scenario === 'server-version-unsupported') - return c.json({ version: '99.0.0', edition: 'SELF_HOSTED' }) + return c.json({ version: '99.0.0', edition: 'COMMUNITY' }) return c.json({ version: '1.6.4', edition: 'CLOUD' }) }) diff --git a/cli/test/fixtures/pkg-manifest.ts b/cli/test/fixtures/pkg-manifest.ts new file mode 100644 index 00000000000..92f2299ad35 --- /dev/null +++ b/cli/test/fixtures/pkg-manifest.ts @@ -0,0 +1,57 @@ +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join } from 'node:path' + +// Mirrors PKG_PATH_ENV in scripts/release-naming.mjs, which cannot be imported +// here: its shebang breaks the Windows test runner. Divergence is self- +// reporting, not silent — the script would fall back to the real +// cli/package.json and every fixture-window assertion would fail. +const PKG_PATH_ENV = 'DIFYCTL_PKG_PATH' + +// release-naming.mjs and release-r2-edge.mjs read their data from +// cli/package.json. Tests spawn them against this fixture instead, so +// assertions can name exact versions without tracking the live release. + +// Deliberately far from any real Dify version, and min != max so "inside the +// window" is a case distinct from either bound. +export const FIXTURE_COMPAT = { minDify: '2.0.0', maxDify: '2.5.0' } + +export const FIXTURE_TARGET_IDS = [ + 'linux-x64', + 'linux-arm64', + 'darwin-x64', + 'darwin-arm64', + 'windows-x64', +] as const + +const FIXTURE_RELEASE = { + tagPrefix: 'difyctl-v', + binName: 'difyctl', + checksumsSuffix: '-checksums.txt', + targets: FIXTURE_TARGET_IDS.map((id) => ({ + id, + bunTarget: `bun-${id}`, + exe: id.startsWith('windows'), + })), +} + +export type PkgManifestOverrides = { + version?: string + channel?: string + compat?: { minDify: string; maxDify: string } +} + +// Returns the env additions that point a spawned script at the fixture. +export function pkgManifestEnv(overrides: PkgManifestOverrides = {}): Record { + const manifest = { + version: overrides.version ?? '0.2.0-alpha', + difyctl: { + channel: overrides.channel ?? 'alpha', + compat: overrides.compat ?? FIXTURE_COMPAT, + release: FIXTURE_RELEASE, + }, + } + const path = join(mkdtempSync(join(tmpdir(), 'difyctl-pkg-')), 'package.json') + writeFileSync(path, JSON.stringify(manifest)) + return { [PKG_PATH_ENV]: path } +} diff --git a/dev/pyrefly-check-local b/dev/pyrefly-check-local index 4b975792701..b14dfba996f 100755 --- a/dev/pyrefly-check-local +++ b/dev/pyrefly-check-local @@ -7,6 +7,7 @@ REPO_ROOT="$SCRIPT_DIR/.." cd "$REPO_ROOT" EXCLUDES_FILE="api/pyrefly-local-excludes.txt" +UNIT_TESTS_CONFIG="tests/unit_tests/pyrefly.toml" TEST_CONTAINERS_DIR="tests/test_containers_integration_tests" TEST_CONTAINERS_CONFIG="$TEST_CONTAINERS_DIR/pyrefly.toml" @@ -74,6 +75,19 @@ fi run_pyrefly "${pyrefly_command[@]}" || status=$? if (( ${#target_paths[@]} == 0 )); then + unit_tests_args=( + "--summary=none" + "--use-ignore-files=false" + "--config=$UNIT_TESTS_CONFIG" + ) + if [[ "${PYREFLY_OUTPUT_FORMAT:-}" == "github" ]]; then + unit_tests_args+=("--output-format=github") + fi + run_pyrefly \ + uv run --directory api --dev pyrefly check \ + "${unit_tests_args[@]}" \ + || status=$? + test_containers_args=( "--summary=none" "--use-ignore-files=false" diff --git a/dev/start-worker b/dev/start-worker index 8baa36f1ed4..9d1be839667 100755 --- a/dev/start-worker +++ b/dev/start-worker @@ -99,17 +99,14 @@ if [[ -n "${ENV_FILE}" ]]; then set +a fi -# If no queues specified, use edition-based defaults +# If no queues are specified, use product-edition defaults if [[ -z "${QUEUES}" ]]; then - # Get EDITION from environment, default to SELF_HOSTED (community edition) - EDITION=${EDITION:-"SELF_HOSTED"} - - # Configure queues based on edition - if [[ "${EDITION}" == "CLOUD" ]]; then + # Configure queues based on product edition + if [[ "${DEPLOYMENT_EDITION:-COMMUNITY}" == "CLOUD" ]]; then # Cloud edition: separate queues for dataset and trigger tasks QUEUES="dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow_professional,workflow_team,workflow_sandbox,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_executor,retention,workflow_based_app_execution" else - # Community edition (SELF_HOSTED): dataset and workflow have separate queues + # Self-hosted editions: dataset and workflow have separate queues QUEUES="dataset,dataset_summary,priority_dataset,priority_pipeline,pipeline,mail,ops_trace,app_deletion,plugin,workflow_storage,conversation,workflow,schedule_poller,schedule_executor,triggered_workflow_dispatcher,trigger_refresh_executor,retention,workflow_based_app_execution" fi diff --git a/dify-agent-runtime/Makefile b/dify-agent-runtime/Makefile index fffa071cc3b..76e853a03f9 100644 --- a/dify-agent-runtime/Makefile +++ b/dify-agent-runtime/Makefile @@ -1,7 +1,5 @@ -.PHONY: build clean test lint proto gen-cli-help integration integration-up integration-test integration-down +.PHONY: build clean test lint gen-cli-help integration integration-up integration-test integration-down BIN_DIR := bin -PROTO_SRC := ../dify-agent/proto/dify/agent/stub/v1/agent_stub.proto -PROTO_OUT := gen AGENT_CLI_HELP_JSON := ../dify-agent/src/dify_agent/layers/_agent_cli_help.json build: $(BIN_DIR)/shellctl $(BIN_DIR)/shellctl-sanitize-pty $(BIN_DIR)/shellctl-runner-exit $(BIN_DIR)/shellctl-runner $(BIN_DIR)/dify-agent gen-cli-help @@ -19,7 +17,7 @@ $(BIN_DIR)/shellctl-runner: $(shell find cmd/runner internal/landlock internal/c go build -o $@ ./cmd/runner -$(BIN_DIR)/dify-agent: $(shell find cmd/dify-agent-cli internal/agentcli internal/stubclient -name '*.go') $(shell find gen -name '*.go' 2>/dev/null) +$(BIN_DIR)/dify-agent: $(shell find cmd/dify-agent-cli internal/agentcli -name '*.go') go build -o $@ ./cmd/dify-agent-cli # Regenerate the JSON help snapshot from the Go command tree. Runs as part of @@ -33,14 +31,6 @@ test: lint lint: golangci-lint run ./... -proto: - @mkdir -p $(PROTO_OUT) - protoc \ - --proto_path=../dify-agent/proto \ - --go_out=$(PROTO_OUT) --go_opt=paths=source_relative \ - --go-grpc_out=$(PROTO_OUT) --go-grpc_opt=paths=source_relative \ - $(PROTO_SRC) - clean: rm -rf $(BIN_DIR) diff --git a/dify-agent-runtime/README.md b/dify-agent-runtime/README.md index 49ecc5890bf..59b1d451167 100644 --- a/dify-agent-runtime/README.md +++ b/dify-agent-runtime/README.md @@ -2,21 +2,16 @@ Go implementation of the shellctl server and runtime utilities. -This is a rewrite of the Python `shellctl` package (`dify-agent/src/shellctl/` and -`dify-agent/src/shellctl_runtime/`). The original Python code is kept as reference. - ## Architecture ``` cmd/ - shellctl/ — main server binary (shellctl serve) - sanitize-pty/ — tmux pipe-pane PTY sanitizer (stdin→stdout filter) - runner-exit/ — post-drain SQLite exit recorder - -internal/ - sanitize/ — PTY ANSI stripping + CR normalization - runner_exit/ — SQLite CAS update for job exit - server/ — HTTP API, job service, tmux controller, output reader + shellctl/ - main server binary (shellctl serve) + sanitize-pty/ - tmux pipe-pane PTY sanitizer (stdin→stdout filter) + runner-exit/ - post-drain SQLite exit recorder + dify-agent-cli/ - cli tool talking to agent backend + runner/ - process runner to bootstrap agent commands +internal/ - internal implementations ``` ## Building @@ -27,11 +22,6 @@ make build Produces binaries in `bin/`: -- `shellctl` — the main server (`shellctl serve --listen 0.0.0.0:5004`) -- `shellctl-sanitize-pty` — PTY sanitizer for tmux pipe-pane -- `shellctl-runner-exit` — exit state writer -- `shellctl-runner` — job runner with integrated Landlock isolation - ### Building docker image ``` @@ -41,7 +31,7 @@ docker build -f dify-agent-runtime/docker/Dockerfile \ dify-agent-runtime/ ``` -### Runing docker container +### Running docker container ``` docker run -d --name dify-agent-runtime \ @@ -78,27 +68,12 @@ The runner automatically creates `$CWD/.tmp` and sets `TMPDIR`, `TMP`, `TEMP` to ### Environment Variables -| Variable | Default | Description | -| -------------------------------- | --------------- | ------------------------------------------------ | -| `SHELLCTL_ENABLE_PATH_ISOLATION` | `true` | Set to `false` to disable Landlock entirely | -| `SHELLCTL_LANDLOCK_RW_PATHS` | _(empty)_ | Comma-separated RW directories (besides `$HOME`) | -| `SHELLCTL_LANDLOCK_RO_PATHS` | `/usr,/bin,...` | Comma-separated RO+exec directories | -| `SHELLCTL_LANDLOCK_RW_DEV_PATHS` | `/dev/null,...` | Comma-separated device files with RW access | +See [here](./internal/envvar/envvar.go) Requires Linux ≥ 5.13. On unsupported kernels, a warning is printed to stderr. ## Dependencies -- Go 1.23+ +- Go 1.26 - `modernc.org/sqlite` (pure-Go SQLite driver, no CGO required) - tmux (runtime dependency, not a build dependency) - -## Migration from Python - -The Go binaries are drop-in replacements for the Python console scripts: - -- `shellctl-sanitize-pty` replaces the Python `shellctl-sanitize-pty` entrypoint -- `shellctl-runner-exit` replaces the Python `shellctl-runner-exit` entrypoint -- `shellctl serve` replaces the Python `shellctl serve` (FastAPI/uvicorn) - -The HTTP API contract, SQLite schema, and filesystem artifact layout are identical. diff --git a/dify-agent-runtime/cmd/dify-agent-cli/main.go b/dify-agent-runtime/cmd/dify-agent-cli/main.go index bcee7dc1fe4..2111d6cd4f1 100644 --- a/dify-agent-runtime/cmd/dify-agent-cli/main.go +++ b/dify-agent-runtime/cmd/dify-agent-cli/main.go @@ -1,5 +1,5 @@ // dify-agent-cli is the Go replacement for the Python dify-agent CLI. -// It communicates with the Agent Stub server via HTTP or gRPC to provide +// It communicates with the Agent Stub server via HTTP to provide // connect, file, drive, and config operations inside the sandbox container. package main @@ -103,14 +103,21 @@ func newFileCommand() *cobra.Command { Short: "Upload or download workflow files through the Agent Stub.", } + var noDownloadLink bool upload := &cobra.Command{ Use: "upload PATH", Short: "Upload one sandbox-local file as a ToolFile output reference.", Args: cobra.ExactArgs(1), RunE: withEnv(func(env *agentcli.Environment, args []string, _ *cobra.Command) error { - return agentcli.RunFileUpload(env, args[0]) + return agentcli.RunFileUpload(env, args[0], noDownloadLink) }), } + upload.Flags().BoolVar( + &noDownloadLink, + "no-download-link", + false, + "Skip creating a public download link after upload.", + ) var downloadTo string download := &cobra.Command{ diff --git a/dify-agent-runtime/cmd/dify-agent-cli/main_test.go b/dify-agent-runtime/cmd/dify-agent-cli/main_test.go index 89b2677e96d..8e024122c6a 100644 --- a/dify-agent-runtime/cmd/dify-agent-cli/main_test.go +++ b/dify-agent-runtime/cmd/dify-agent-cli/main_test.go @@ -53,7 +53,12 @@ func TestCommandHelp(t *testing.T) { { name: "file upload", args: []string{"file", "upload", "--help"}, - want: []string{"dify-agent file upload", "Upload one sandbox-local file"}, + want: []string{ + "dify-agent file upload", + "Upload one sandbox-local file", + "--no-download-link", + "Skip creating a public download link after upload.", + }, }, { name: "file download", diff --git a/dify-agent-runtime/docker/Dockerfile b/dify-agent-runtime/docker/Dockerfile index 798e36a8d41..ca058727e45 100644 --- a/dify-agent-runtime/docker/Dockerfile +++ b/dify-agent-runtime/docker/Dockerfile @@ -40,6 +40,7 @@ RUN apt-get update \ openssh-client \ procps \ ripgrep \ + tini \ tmux \ unzip \ xz-utils \ @@ -79,4 +80,5 @@ WORKDIR /home/dify EXPOSE 5004 +ENTRYPOINT ["/usr/bin/tini", "-g", "--"] CMD ["shellctl", "serve", "--listen", "0.0.0.0:5004"] diff --git a/dify-agent-runtime/docker/sync-e2b-template.sh b/dify-agent-runtime/docker/sync-e2b-template.sh new file mode 100755 index 00000000000..09a9558f98a --- /dev/null +++ b/dify-agent-runtime/docker/sync-e2b-template.sh @@ -0,0 +1,253 @@ +#!/usr/bin/env bash + +set -euo pipefail + +readonly DEFAULT_TEMPLATE_NAME='dify-agent-local-sandbox' +readonly DEFAULT_IMAGE_REPOSITORY='langgenius/dify-agent-local-sandbox' +readonly DEFAULT_PLATFORM='linux/amd64' +readonly DEFAULT_E2B_CLI_VERSION='2.13.3' +readonly DEFAULT_WAIT_SECONDS='600' +readonly DEFAULT_POLL_SECONDS='10' +readonly DEFAULT_CPU_COUNT='2' +readonly DEFAULT_MEMORY_MB='1024' +readonly START_COMMAND='/usr/bin/env SHELLCTL_ENABLE_PATH_ISOLATION=true /usr/local/bin/shellctl serve --listen 0.0.0.0:5004' +readonly READY_COMMAND='curl -fsS http://localhost:5004/healthz' + +SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" + +usage() { + cat <&2 + exit 1 +} + +require_command() { + command -v "$1" >/dev/null 2>&1 || die "required command not found: $1" +} + +require_non_negative_integer() { + local name="$1" + local value="$2" + + [[ "${value}" =~ ^[0-9]+$ ]] || die "${name} must be a non-negative integer: ${value}" +} + +require_positive_integer() { + local name="$1" + local value="$2" + + require_non_negative_integer "${name}" "${value}" + (( value > 0 )) || die "${name} must be greater than zero" +} + +print_command() { + printf 'Command:' + printf ' %q' "$@" + printf '\n' +} + +source_ref='HEAD' +image_repository="${E2B_BASE_IMAGE_REPOSITORY:-${DEFAULT_IMAGE_REPOSITORY}}" +image_tag='' +template_name="${E2B_TEMPLATE_NAME:-${DEFAULT_TEMPLATE_NAME}}" +platform="${DEFAULT_PLATFORM}" +e2b_cli_version="${E2B_CLI_VERSION:-${DEFAULT_E2B_CLI_VERSION}}" +wait_seconds="${DEFAULT_WAIT_SECONDS}" +poll_seconds="${DEFAULT_POLL_SECONDS}" +cpu_count="${DEFAULT_CPU_COUNT}" +memory_mb="${DEFAULT_MEMORY_MB}" +dry_run=false +no_cache=false + +while (( $# > 0 )); do + case "$1" in + --source-ref) + (( $# >= 2 )) || die '--source-ref requires a value' + source_ref="$2" + shift 2 + ;; + --image-repository) + (( $# >= 2 )) || die '--image-repository requires a value' + image_repository="$2" + shift 2 + ;; + --image-tag) + (( $# >= 2 )) || die '--image-tag requires a value' + image_tag="$2" + shift 2 + ;; + --template-name) + (( $# >= 2 )) || die '--template-name requires a value' + template_name="$2" + shift 2 + ;; + --platform) + (( $# >= 2 )) || die '--platform requires a value' + platform="$2" + shift 2 + ;; + --wait-seconds) + (( $# >= 2 )) || die '--wait-seconds requires a value' + wait_seconds="$2" + shift 2 + ;; + --poll-seconds) + (( $# >= 2 )) || die '--poll-seconds requires a value' + poll_seconds="$2" + shift 2 + ;; + --cpu-count) + (( $# >= 2 )) || die '--cpu-count requires a value' + cpu_count="$2" + shift 2 + ;; + --memory-mb) + (( $# >= 2 )) || die '--memory-mb requires a value' + memory_mb="$2" + shift 2 + ;; + --no-cache) + no_cache=true + shift + ;; + --dry-run) + dry_run=true + shift + ;; + -h|--help) + usage + exit 0 + ;; + *) + die "unknown option: $1" + ;; + esac +done + +[[ -n "${image_repository}" ]] || die 'image repository cannot be empty' +[[ -n "${template_name}" ]] || die 'template name cannot be empty' +[[ "${platform}" == */* ]] || die "platform must use OS/ARCH form: ${platform}" +require_non_negative_integer 'wait-seconds' "${wait_seconds}" +require_positive_integer 'poll-seconds' "${poll_seconds}" +require_positive_integer 'cpu-count' "${cpu_count}" +require_positive_integer 'memory-mb' "${memory_mb}" + +require_command git +require_command docker +require_command jq +require_command mktemp + +source_sha="$(git -C "${SCRIPT_DIR}" rev-parse --verify "${source_ref}^{commit}" 2>/dev/null)" \ + || die "cannot resolve source ref: ${source_ref}" +[[ -n "${image_tag}" ]] || image_tag="${source_sha}" + +image_ref="${image_repository}:${image_tag}" +platform_os="${platform%%/*}" +platform_arch="${platform#*/}" +deadline=$((SECONDS + wait_seconds)) +manifest_json='' +inspect_error='' + +while true; do + if manifest_json="$(docker buildx imagetools inspect --format '{{json .Manifest}}' "${image_ref}" 2>&1)"; then + break + fi + inspect_error="${manifest_json}" + manifest_json='' + + (( SECONDS < deadline )) || die "image is not available: ${image_ref}${inspect_error:+ (${inspect_error})}" + echo "Waiting for published image ${image_ref}..." >&2 + sleep "${poll_seconds}" +done + +platform_digest="$( + jq -r \ + --arg os "${platform_os}" \ + --arg arch "${platform_arch}" \ + '[.manifests[]? | select(.platform.os == $os and .platform.architecture == $arch) | .digest][0] // empty' \ + <<<"${manifest_json}" +)" +[[ "${platform_digest}" == sha256:* ]] \ + || die "image ${image_ref} does not contain platform ${platform}" + +resolved_image="${image_repository}@${platform_digest}" +image_json="$(docker buildx imagetools inspect --format '{{json .Image}}' "${resolved_image}")" \ + || die "cannot inspect resolved image: ${resolved_image}" + +actual_os="$(jq -r '.os // empty' <<<"${image_json}")" +actual_arch="$(jq -r '.architecture // empty' <<<"${image_json}")" +actual_revision="$(jq -r '.config.Labels["org.opencontainers.image.revision"] // empty' <<<"${image_json}")" + +[[ "${actual_os}/${actual_arch}" == "${platform}" ]] \ + || die "resolved image platform ${actual_os}/${actual_arch} does not match ${platform}" +[[ "${actual_revision}" == "${source_sha}" ]] \ + || die "image revision ${actual_revision:-} does not match source commit ${source_sha}" + +tmp_dir="$(mktemp -d)" +cleanup() { + rm -rf "${tmp_dir}" +} +trap cleanup EXIT + +dockerfile="${tmp_dir}/e2b.Dockerfile" +printf 'FROM %s\nUSER dify\nWORKDIR /home/dify\n' "${resolved_image}" > "${dockerfile}" + +e2b_command=( + npx --yes "@e2b/cli@${e2b_cli_version}" + template create "${template_name}" + --path "${tmp_dir}" + --dockerfile e2b.Dockerfile + --cmd "${START_COMMAND}" + --ready-cmd "${READY_COMMAND}" + --cpu-count "${cpu_count}" + --memory-mb "${memory_mb}" +) +if [[ "${no_cache}" == true ]]; then + e2b_command+=(--no-cache) +fi + +echo "Source commit: ${source_sha}" +echo "Published image: ${image_ref}" +echo "Resolved image: ${resolved_image} (${platform})" +echo "E2B Template: ${template_name}" +print_command "${e2b_command[@]}" + +if [[ "${dry_run}" == true ]]; then + echo 'Dry run; E2B Template was not changed.' + exit 0 +fi + +require_command npx +[[ -n "${E2B_API_KEY:-}" ]] || die 'E2B_API_KEY is required' + +"${e2b_command[@]}" +echo "E2B Template synchronized: ${template_name} <- ${resolved_image}" diff --git a/dify-agent-runtime/gen/dify/agent/stub/v1/agent_stub.pb.go b/dify-agent-runtime/gen/dify/agent/stub/v1/agent_stub.pb.go deleted file mode 100644 index 4cfc11903d8..00000000000 --- a/dify-agent-runtime/gen/dify/agent/stub/v1/agent_stub.pb.go +++ /dev/null @@ -1,515 +0,0 @@ -// Code generated by protoc-gen-go. DO NOT EDIT. -// versions: -// protoc-gen-go v1.36.11 -// protoc v5.28.3 -// source: dify/agent/stub/v1/agent_stub.proto - -package stubv1 - -import ( - protoreflect "google.golang.org/protobuf/reflect/protoreflect" - protoimpl "google.golang.org/protobuf/runtime/protoimpl" - reflect "reflect" - sync "sync" - unsafe "unsafe" -) - -const ( - // Verify that this generated code is sufficiently up-to-date. - _ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion) - // Verify that runtime/protoimpl is sufficiently up-to-date. - _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) -) - -type ConnectRequest struct { - state protoimpl.MessageState `protogen:"open.v1"` - ProtocolVersion int32 `protobuf:"varint,1,opt,name=protocol_version,json=protocolVersion,proto3" json:"protocol_version,omitempty"` - Argv []string `protobuf:"bytes,2,rep,name=argv,proto3" json:"argv,omitempty"` - MetadataJson string `protobuf:"bytes,3,opt,name=metadata_json,json=metadataJson,proto3" json:"metadata_json,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache -} - -func (x *ConnectRequest) Reset() { - *x = ConnectRequest{} - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[0] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) -} - -func (x *ConnectRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*ConnectRequest) ProtoMessage() {} - -func (x *ConnectRequest) ProtoReflect() protoreflect.Message { - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[0] - if x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use ConnectRequest.ProtoReflect.Descriptor instead. -func (*ConnectRequest) Descriptor() ([]byte, []int) { - return file_dify_agent_stub_v1_agent_stub_proto_rawDescGZIP(), []int{0} -} - -func (x *ConnectRequest) GetProtocolVersion() int32 { - if x != nil { - return x.ProtocolVersion - } - return 0 -} - -func (x *ConnectRequest) GetArgv() []string { - if x != nil { - return x.Argv - } - return nil -} - -func (x *ConnectRequest) GetMetadataJson() string { - if x != nil { - return x.MetadataJson - } - return "" -} - -type ConnectResponse struct { - state protoimpl.MessageState `protogen:"open.v1"` - ConnectionId string `protobuf:"bytes,1,opt,name=connection_id,json=connectionId,proto3" json:"connection_id,omitempty"` - Status string `protobuf:"bytes,2,opt,name=status,proto3" json:"status,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache -} - -func (x *ConnectResponse) Reset() { - *x = ConnectResponse{} - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[1] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) -} - -func (x *ConnectResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*ConnectResponse) ProtoMessage() {} - -func (x *ConnectResponse) ProtoReflect() protoreflect.Message { - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[1] - if x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use ConnectResponse.ProtoReflect.Descriptor instead. -func (*ConnectResponse) Descriptor() ([]byte, []int) { - return file_dify_agent_stub_v1_agent_stub_proto_rawDescGZIP(), []int{1} -} - -func (x *ConnectResponse) GetConnectionId() string { - if x != nil { - return x.ConnectionId - } - return "" -} - -func (x *ConnectResponse) GetStatus() string { - if x != nil { - return x.Status - } - return "" -} - -type FileUploadRequest struct { - state protoimpl.MessageState `protogen:"open.v1"` - Filename string `protobuf:"bytes,1,opt,name=filename,proto3" json:"filename,omitempty"` - Mimetype string `protobuf:"bytes,2,opt,name=mimetype,proto3" json:"mimetype,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache -} - -func (x *FileUploadRequest) Reset() { - *x = FileUploadRequest{} - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[2] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) -} - -func (x *FileUploadRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*FileUploadRequest) ProtoMessage() {} - -func (x *FileUploadRequest) ProtoReflect() protoreflect.Message { - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[2] - if x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use FileUploadRequest.ProtoReflect.Descriptor instead. -func (*FileUploadRequest) Descriptor() ([]byte, []int) { - return file_dify_agent_stub_v1_agent_stub_proto_rawDescGZIP(), []int{2} -} - -func (x *FileUploadRequest) GetFilename() string { - if x != nil { - return x.Filename - } - return "" -} - -func (x *FileUploadRequest) GetMimetype() string { - if x != nil { - return x.Mimetype - } - return "" -} - -type FileUploadResponse struct { - state protoimpl.MessageState `protogen:"open.v1"` - UploadUrl string `protobuf:"bytes,1,opt,name=upload_url,json=uploadUrl,proto3" json:"upload_url,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache -} - -func (x *FileUploadResponse) Reset() { - *x = FileUploadResponse{} - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[3] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) -} - -func (x *FileUploadResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*FileUploadResponse) ProtoMessage() {} - -func (x *FileUploadResponse) ProtoReflect() protoreflect.Message { - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[3] - if x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use FileUploadResponse.ProtoReflect.Descriptor instead. -func (*FileUploadResponse) Descriptor() ([]byte, []int) { - return file_dify_agent_stub_v1_agent_stub_proto_rawDescGZIP(), []int{3} -} - -func (x *FileUploadResponse) GetUploadUrl() string { - if x != nil { - return x.UploadUrl - } - return "" -} - -type FileMapping struct { - state protoimpl.MessageState `protogen:"open.v1"` - TransferMethod string `protobuf:"bytes,1,opt,name=transfer_method,json=transferMethod,proto3" json:"transfer_method,omitempty"` - Reference *string `protobuf:"bytes,2,opt,name=reference,proto3,oneof" json:"reference,omitempty"` - Url *string `protobuf:"bytes,3,opt,name=url,proto3,oneof" json:"url,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache -} - -func (x *FileMapping) Reset() { - *x = FileMapping{} - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[4] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) -} - -func (x *FileMapping) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*FileMapping) ProtoMessage() {} - -func (x *FileMapping) ProtoReflect() protoreflect.Message { - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[4] - if x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use FileMapping.ProtoReflect.Descriptor instead. -func (*FileMapping) Descriptor() ([]byte, []int) { - return file_dify_agent_stub_v1_agent_stub_proto_rawDescGZIP(), []int{4} -} - -func (x *FileMapping) GetTransferMethod() string { - if x != nil { - return x.TransferMethod - } - return "" -} - -func (x *FileMapping) GetReference() string { - if x != nil && x.Reference != nil { - return *x.Reference - } - return "" -} - -func (x *FileMapping) GetUrl() string { - if x != nil && x.Url != nil { - return *x.Url - } - return "" -} - -type FileDownloadRequest struct { - state protoimpl.MessageState `protogen:"open.v1"` - File *FileMapping `protobuf:"bytes,1,opt,name=file,proto3" json:"file,omitempty"` - ForExternal *bool `protobuf:"varint,2,opt,name=for_external,json=forExternal,proto3,oneof" json:"for_external,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache -} - -func (x *FileDownloadRequest) Reset() { - *x = FileDownloadRequest{} - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[5] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) -} - -func (x *FileDownloadRequest) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*FileDownloadRequest) ProtoMessage() {} - -func (x *FileDownloadRequest) ProtoReflect() protoreflect.Message { - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[5] - if x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use FileDownloadRequest.ProtoReflect.Descriptor instead. -func (*FileDownloadRequest) Descriptor() ([]byte, []int) { - return file_dify_agent_stub_v1_agent_stub_proto_rawDescGZIP(), []int{5} -} - -func (x *FileDownloadRequest) GetFile() *FileMapping { - if x != nil { - return x.File - } - return nil -} - -func (x *FileDownloadRequest) GetForExternal() bool { - if x != nil && x.ForExternal != nil { - return *x.ForExternal - } - return false -} - -type FileDownloadResponse struct { - state protoimpl.MessageState `protogen:"open.v1"` - Filename string `protobuf:"bytes,1,opt,name=filename,proto3" json:"filename,omitempty"` - MimeType *string `protobuf:"bytes,2,opt,name=mime_type,json=mimeType,proto3,oneof" json:"mime_type,omitempty"` - Size int64 `protobuf:"varint,3,opt,name=size,proto3" json:"size,omitempty"` - DownloadUrl string `protobuf:"bytes,4,opt,name=download_url,json=downloadUrl,proto3" json:"download_url,omitempty"` - unknownFields protoimpl.UnknownFields - sizeCache protoimpl.SizeCache -} - -func (x *FileDownloadResponse) Reset() { - *x = FileDownloadResponse{} - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[6] - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - ms.StoreMessageInfo(mi) -} - -func (x *FileDownloadResponse) String() string { - return protoimpl.X.MessageStringOf(x) -} - -func (*FileDownloadResponse) ProtoMessage() {} - -func (x *FileDownloadResponse) ProtoReflect() protoreflect.Message { - mi := &file_dify_agent_stub_v1_agent_stub_proto_msgTypes[6] - if x != nil { - ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) - if ms.LoadMessageInfo() == nil { - ms.StoreMessageInfo(mi) - } - return ms - } - return mi.MessageOf(x) -} - -// Deprecated: Use FileDownloadResponse.ProtoReflect.Descriptor instead. -func (*FileDownloadResponse) Descriptor() ([]byte, []int) { - return file_dify_agent_stub_v1_agent_stub_proto_rawDescGZIP(), []int{6} -} - -func (x *FileDownloadResponse) GetFilename() string { - if x != nil { - return x.Filename - } - return "" -} - -func (x *FileDownloadResponse) GetMimeType() string { - if x != nil && x.MimeType != nil { - return *x.MimeType - } - return "" -} - -func (x *FileDownloadResponse) GetSize() int64 { - if x != nil { - return x.Size - } - return 0 -} - -func (x *FileDownloadResponse) GetDownloadUrl() string { - if x != nil { - return x.DownloadUrl - } - return "" -} - -var File_dify_agent_stub_v1_agent_stub_proto protoreflect.FileDescriptor - -const file_dify_agent_stub_v1_agent_stub_proto_rawDesc = "" + - "\n" + - "#dify/agent/stub/v1/agent_stub.proto\x12\x12dify.agent.stub.v1\"t\n" + - "\x0eConnectRequest\x12)\n" + - "\x10protocol_version\x18\x01 \x01(\x05R\x0fprotocolVersion\x12\x12\n" + - "\x04argv\x18\x02 \x03(\tR\x04argv\x12#\n" + - "\rmetadata_json\x18\x03 \x01(\tR\fmetadataJson\"N\n" + - "\x0fConnectResponse\x12#\n" + - "\rconnection_id\x18\x01 \x01(\tR\fconnectionId\x12\x16\n" + - "\x06status\x18\x02 \x01(\tR\x06status\"K\n" + - "\x11FileUploadRequest\x12\x1a\n" + - "\bfilename\x18\x01 \x01(\tR\bfilename\x12\x1a\n" + - "\bmimetype\x18\x02 \x01(\tR\bmimetype\"3\n" + - "\x12FileUploadResponse\x12\x1d\n" + - "\n" + - "upload_url\x18\x01 \x01(\tR\tuploadUrl\"\x86\x01\n" + - "\vFileMapping\x12'\n" + - "\x0ftransfer_method\x18\x01 \x01(\tR\x0etransferMethod\x12!\n" + - "\treference\x18\x02 \x01(\tH\x00R\treference\x88\x01\x01\x12\x15\n" + - "\x03url\x18\x03 \x01(\tH\x01R\x03url\x88\x01\x01B\f\n" + - "\n" + - "_referenceB\x06\n" + - "\x04_url\"\x83\x01\n" + - "\x13FileDownloadRequest\x123\n" + - "\x04file\x18\x01 \x01(\v2\x1f.dify.agent.stub.v1.FileMappingR\x04file\x12&\n" + - "\ffor_external\x18\x02 \x01(\bH\x00R\vforExternal\x88\x01\x01B\x0f\n" + - "\r_for_external\"\x99\x01\n" + - "\x14FileDownloadResponse\x12\x1a\n" + - "\bfilename\x18\x01 \x01(\tR\bfilename\x12 \n" + - "\tmime_type\x18\x02 \x01(\tH\x00R\bmimeType\x88\x01\x01\x12\x12\n" + - "\x04size\x18\x03 \x01(\x03R\x04size\x12!\n" + - "\fdownload_url\x18\x04 \x01(\tR\vdownloadUrlB\f\n" + - "\n" + - "_mime_type2\xc0\x02\n" + - "\x10AgentStubService\x12R\n" + - "\aConnect\x12\".dify.agent.stub.v1.ConnectRequest\x1a#.dify.agent.stub.v1.ConnectResponse\x12h\n" + - "\x17CreateFileUploadRequest\x12%.dify.agent.stub.v1.FileUploadRequest\x1a&.dify.agent.stub.v1.FileUploadResponse\x12n\n" + - "\x19CreateFileDownloadRequest\x12'.dify.agent.stub.v1.FileDownloadRequest\x1a(.dify.agent.stub.v1.FileDownloadResponseB\x1bZ\x19dify/agent/stub/v1;stubv1b\x06proto3" - -var ( - file_dify_agent_stub_v1_agent_stub_proto_rawDescOnce sync.Once - file_dify_agent_stub_v1_agent_stub_proto_rawDescData []byte -) - -func file_dify_agent_stub_v1_agent_stub_proto_rawDescGZIP() []byte { - file_dify_agent_stub_v1_agent_stub_proto_rawDescOnce.Do(func() { - file_dify_agent_stub_v1_agent_stub_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_dify_agent_stub_v1_agent_stub_proto_rawDesc), len(file_dify_agent_stub_v1_agent_stub_proto_rawDesc))) - }) - return file_dify_agent_stub_v1_agent_stub_proto_rawDescData -} - -var file_dify_agent_stub_v1_agent_stub_proto_msgTypes = make([]protoimpl.MessageInfo, 7) -var file_dify_agent_stub_v1_agent_stub_proto_goTypes = []any{ - (*ConnectRequest)(nil), // 0: dify.agent.stub.v1.ConnectRequest - (*ConnectResponse)(nil), // 1: dify.agent.stub.v1.ConnectResponse - (*FileUploadRequest)(nil), // 2: dify.agent.stub.v1.FileUploadRequest - (*FileUploadResponse)(nil), // 3: dify.agent.stub.v1.FileUploadResponse - (*FileMapping)(nil), // 4: dify.agent.stub.v1.FileMapping - (*FileDownloadRequest)(nil), // 5: dify.agent.stub.v1.FileDownloadRequest - (*FileDownloadResponse)(nil), // 6: dify.agent.stub.v1.FileDownloadResponse -} -var file_dify_agent_stub_v1_agent_stub_proto_depIdxs = []int32{ - 4, // 0: dify.agent.stub.v1.FileDownloadRequest.file:type_name -> dify.agent.stub.v1.FileMapping - 0, // 1: dify.agent.stub.v1.AgentStubService.Connect:input_type -> dify.agent.stub.v1.ConnectRequest - 2, // 2: dify.agent.stub.v1.AgentStubService.CreateFileUploadRequest:input_type -> dify.agent.stub.v1.FileUploadRequest - 5, // 3: dify.agent.stub.v1.AgentStubService.CreateFileDownloadRequest:input_type -> dify.agent.stub.v1.FileDownloadRequest - 1, // 4: dify.agent.stub.v1.AgentStubService.Connect:output_type -> dify.agent.stub.v1.ConnectResponse - 3, // 5: dify.agent.stub.v1.AgentStubService.CreateFileUploadRequest:output_type -> dify.agent.stub.v1.FileUploadResponse - 6, // 6: dify.agent.stub.v1.AgentStubService.CreateFileDownloadRequest:output_type -> dify.agent.stub.v1.FileDownloadResponse - 4, // [4:7] is the sub-list for method output_type - 1, // [1:4] is the sub-list for method input_type - 1, // [1:1] is the sub-list for extension type_name - 1, // [1:1] is the sub-list for extension extendee - 0, // [0:1] is the sub-list for field type_name -} - -func init() { file_dify_agent_stub_v1_agent_stub_proto_init() } -func file_dify_agent_stub_v1_agent_stub_proto_init() { - if File_dify_agent_stub_v1_agent_stub_proto != nil { - return - } - file_dify_agent_stub_v1_agent_stub_proto_msgTypes[4].OneofWrappers = []any{} - file_dify_agent_stub_v1_agent_stub_proto_msgTypes[5].OneofWrappers = []any{} - file_dify_agent_stub_v1_agent_stub_proto_msgTypes[6].OneofWrappers = []any{} - type x struct{} - out := protoimpl.TypeBuilder{ - File: protoimpl.DescBuilder{ - GoPackagePath: reflect.TypeOf(x{}).PkgPath(), - RawDescriptor: unsafe.Slice(unsafe.StringData(file_dify_agent_stub_v1_agent_stub_proto_rawDesc), len(file_dify_agent_stub_v1_agent_stub_proto_rawDesc)), - NumEnums: 0, - NumMessages: 7, - NumExtensions: 0, - NumServices: 1, - }, - GoTypes: file_dify_agent_stub_v1_agent_stub_proto_goTypes, - DependencyIndexes: file_dify_agent_stub_v1_agent_stub_proto_depIdxs, - MessageInfos: file_dify_agent_stub_v1_agent_stub_proto_msgTypes, - }.Build() - File_dify_agent_stub_v1_agent_stub_proto = out.File - file_dify_agent_stub_v1_agent_stub_proto_goTypes = nil - file_dify_agent_stub_v1_agent_stub_proto_depIdxs = nil -} diff --git a/dify-agent-runtime/gen/dify/agent/stub/v1/agent_stub_grpc.pb.go b/dify-agent-runtime/gen/dify/agent/stub/v1/agent_stub_grpc.pb.go deleted file mode 100644 index f7b649df503..00000000000 --- a/dify-agent-runtime/gen/dify/agent/stub/v1/agent_stub_grpc.pb.go +++ /dev/null @@ -1,197 +0,0 @@ -// Code generated by protoc-gen-go-grpc. DO NOT EDIT. -// versions: -// - protoc-gen-go-grpc v1.6.2 -// - protoc v5.28.3 -// source: dify/agent/stub/v1/agent_stub.proto - -package stubv1 - -import ( - context "context" - grpc "google.golang.org/grpc" - codes "google.golang.org/grpc/codes" - status "google.golang.org/grpc/status" -) - -// This is a compile-time assertion to ensure that this generated file -// is compatible with the grpc package it is being compiled against. -// Requires gRPC-Go v1.64.0 or later. -const _ = grpc.SupportPackageIsVersion9 - -const ( - AgentStubService_Connect_FullMethodName = "/dify.agent.stub.v1.AgentStubService/Connect" - AgentStubService_CreateFileUploadRequest_FullMethodName = "/dify.agent.stub.v1.AgentStubService/CreateFileUploadRequest" - AgentStubService_CreateFileDownloadRequest_FullMethodName = "/dify.agent.stub.v1.AgentStubService/CreateFileDownloadRequest" -) - -// AgentStubServiceClient is the client API for AgentStubService service. -// -// For semantics around ctx use and closing/ending streaming RPCs, please refer to https://pkg.go.dev/google.golang.org/grpc/?tab=doc#ClientConn.NewStream. -type AgentStubServiceClient interface { - Connect(ctx context.Context, in *ConnectRequest, opts ...grpc.CallOption) (*ConnectResponse, error) - CreateFileUploadRequest(ctx context.Context, in *FileUploadRequest, opts ...grpc.CallOption) (*FileUploadResponse, error) - CreateFileDownloadRequest(ctx context.Context, in *FileDownloadRequest, opts ...grpc.CallOption) (*FileDownloadResponse, error) -} - -type agentStubServiceClient struct { - cc grpc.ClientConnInterface -} - -func NewAgentStubServiceClient(cc grpc.ClientConnInterface) AgentStubServiceClient { - return &agentStubServiceClient{cc} -} - -func (c *agentStubServiceClient) Connect(ctx context.Context, in *ConnectRequest, opts ...grpc.CallOption) (*ConnectResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(ConnectResponse) - err := c.cc.Invoke(ctx, AgentStubService_Connect_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *agentStubServiceClient) CreateFileUploadRequest(ctx context.Context, in *FileUploadRequest, opts ...grpc.CallOption) (*FileUploadResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(FileUploadResponse) - err := c.cc.Invoke(ctx, AgentStubService_CreateFileUploadRequest_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -func (c *agentStubServiceClient) CreateFileDownloadRequest(ctx context.Context, in *FileDownloadRequest, opts ...grpc.CallOption) (*FileDownloadResponse, error) { - cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...) - out := new(FileDownloadResponse) - err := c.cc.Invoke(ctx, AgentStubService_CreateFileDownloadRequest_FullMethodName, in, out, cOpts...) - if err != nil { - return nil, err - } - return out, nil -} - -// AgentStubServiceServer is the server API for AgentStubService service. -// All implementations must embed UnimplementedAgentStubServiceServer -// for forward compatibility. -type AgentStubServiceServer interface { - Connect(context.Context, *ConnectRequest) (*ConnectResponse, error) - CreateFileUploadRequest(context.Context, *FileUploadRequest) (*FileUploadResponse, error) - CreateFileDownloadRequest(context.Context, *FileDownloadRequest) (*FileDownloadResponse, error) - mustEmbedUnimplementedAgentStubServiceServer() -} - -// UnimplementedAgentStubServiceServer must be embedded to have -// forward compatible implementations. -// -// NOTE: this should be embedded by value instead of pointer to avoid a nil -// pointer dereference when methods are called. -type UnimplementedAgentStubServiceServer struct{} - -func (UnimplementedAgentStubServiceServer) Connect(context.Context, *ConnectRequest) (*ConnectResponse, error) { - return nil, status.Error(codes.Unimplemented, "method Connect not implemented") -} -func (UnimplementedAgentStubServiceServer) CreateFileUploadRequest(context.Context, *FileUploadRequest) (*FileUploadResponse, error) { - return nil, status.Error(codes.Unimplemented, "method CreateFileUploadRequest not implemented") -} -func (UnimplementedAgentStubServiceServer) CreateFileDownloadRequest(context.Context, *FileDownloadRequest) (*FileDownloadResponse, error) { - return nil, status.Error(codes.Unimplemented, "method CreateFileDownloadRequest not implemented") -} -func (UnimplementedAgentStubServiceServer) mustEmbedUnimplementedAgentStubServiceServer() {} -func (UnimplementedAgentStubServiceServer) testEmbeddedByValue() {} - -// UnsafeAgentStubServiceServer may be embedded to opt out of forward compatibility for this service. -// Use of this interface is not recommended, as added methods to AgentStubServiceServer will -// result in compilation errors. -type UnsafeAgentStubServiceServer interface { - mustEmbedUnimplementedAgentStubServiceServer() -} - -func RegisterAgentStubServiceServer(s grpc.ServiceRegistrar, srv AgentStubServiceServer) { - // If the following call panics, it indicates UnimplementedAgentStubServiceServer was - // embedded by pointer and is nil. This will cause panics if an - // unimplemented method is ever invoked, so we test this at initialization - // time to prevent it from happening at runtime later due to I/O. - if t, ok := srv.(interface{ testEmbeddedByValue() }); ok { - t.testEmbeddedByValue() - } - s.RegisterService(&AgentStubService_ServiceDesc, srv) -} - -func _AgentStubService_Connect_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(ConnectRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AgentStubServiceServer).Connect(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AgentStubService_Connect_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AgentStubServiceServer).Connect(ctx, req.(*ConnectRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AgentStubService_CreateFileUploadRequest_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(FileUploadRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AgentStubServiceServer).CreateFileUploadRequest(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AgentStubService_CreateFileUploadRequest_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AgentStubServiceServer).CreateFileUploadRequest(ctx, req.(*FileUploadRequest)) - } - return interceptor(ctx, in, info, handler) -} - -func _AgentStubService_CreateFileDownloadRequest_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) { - in := new(FileDownloadRequest) - if err := dec(in); err != nil { - return nil, err - } - if interceptor == nil { - return srv.(AgentStubServiceServer).CreateFileDownloadRequest(ctx, in) - } - info := &grpc.UnaryServerInfo{ - Server: srv, - FullMethod: AgentStubService_CreateFileDownloadRequest_FullMethodName, - } - handler := func(ctx context.Context, req interface{}) (interface{}, error) { - return srv.(AgentStubServiceServer).CreateFileDownloadRequest(ctx, req.(*FileDownloadRequest)) - } - return interceptor(ctx, in, info, handler) -} - -// AgentStubService_ServiceDesc is the grpc.ServiceDesc for AgentStubService service. -// It's only intended for direct use with grpc.RegisterService, -// and not to be introspected or modified (even as a copy) -var AgentStubService_ServiceDesc = grpc.ServiceDesc{ - ServiceName: "dify.agent.stub.v1.AgentStubService", - HandlerType: (*AgentStubServiceServer)(nil), - Methods: []grpc.MethodDesc{ - { - MethodName: "Connect", - Handler: _AgentStubService_Connect_Handler, - }, - { - MethodName: "CreateFileUploadRequest", - Handler: _AgentStubService_CreateFileUploadRequest_Handler, - }, - { - MethodName: "CreateFileDownloadRequest", - Handler: _AgentStubService_CreateFileDownloadRequest_Handler, - }, - }, - Streams: []grpc.StreamDesc{}, - Metadata: "dify/agent/stub/v1/agent_stub.proto", -} diff --git a/dify-agent-runtime/go.mod b/dify-agent-runtime/go.mod index 297b144dcd1..9fa6c29ac6b 100644 --- a/dify-agent-runtime/go.mod +++ b/dify-agent-runtime/go.mod @@ -5,8 +5,6 @@ go 1.26 require ( github.com/landlock-lsm/go-landlock v0.9.0 github.com/spf13/cobra v1.10.2 - google.golang.org/grpc v1.82.1 - google.golang.org/protobuf v1.36.11 modernc.org/sqlite v1.37.1 ) @@ -19,10 +17,9 @@ require ( github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect github.com/spf13/pflag v1.0.9 // indirect golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect - golang.org/x/net v0.53.0 // indirect - golang.org/x/sys v0.43.0 // indirect - golang.org/x/text v0.36.0 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 // indirect + golang.org/x/mod v0.35.0 // indirect + golang.org/x/sys v0.45.0 // indirect + golang.org/x/tools v0.44.0 // indirect kernel.org/pub/linux/libs/security/libcap/psx v1.2.77 // indirect modernc.org/libc v1.65.7 // indirect modernc.org/mathutil v1.7.1 // indirect diff --git a/dify-agent-runtime/go.sum b/dify-agent-runtime/go.sum index b4fbc8b3322..c5698692187 100644 --- a/dify-agent-runtime/go.sum +++ b/dify-agent-runtime/go.sum @@ -1,16 +1,6 @@ -github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= -github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= -github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= -github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= -github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= -github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= -github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= -github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= -github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= -github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= @@ -30,42 +20,18 @@ github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU= github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4= github.com/spf13/pflag v1.0.9 h1:9exaQaMOCwffKiiiYk6/BndUBv+iRViNW+4lEMi0PvY= github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= -go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= -go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= -go.opentelemetry.io/otel v1.43.0 h1:mYIM03dnh5zfN7HautFE4ieIig9amkNANT+xcVxAj9I= -go.opentelemetry.io/otel v1.43.0/go.mod h1:JuG+u74mvjvcm8vj8pI5XiHy1zDeoCS2LB1spIq7Ay0= -go.opentelemetry.io/otel/metric v1.43.0 h1:d7638QeInOnuwOONPp4JAOGfbCEpYb+K6DVWvdxGzgM= -go.opentelemetry.io/otel/metric v1.43.0/go.mod h1:RDnPtIxvqlgO8GRW18W6Z/4P462ldprJtfxHxyKd2PY= -go.opentelemetry.io/otel/sdk v1.43.0 h1:pi5mE86i5rTeLXqoF/hhiBtUNcrAGHLKQdhg4h4V9Dg= -go.opentelemetry.io/otel/sdk v1.43.0/go.mod h1:P+IkVU3iWukmiit/Yf9AWvpyRDlUeBaRg6Y+C58QHzg= -go.opentelemetry.io/otel/sdk/metric v1.43.0 h1:S88dyqXjJkuBNLeMcVPRFXpRw2fuwdvfCGLEo89fDkw= -go.opentelemetry.io/otel/sdk/metric v1.43.0/go.mod h1:C/RJtwSEJ5hzTiUz5pXF1kILHStzb9zFlIEe85bhj6A= -go.opentelemetry.io/otel/trace v1.43.0 h1:BkNrHpup+4k4w+ZZ86CZoHHEkohws8AY+WTX09nk+3A= -go.opentelemetry.io/otel/trace v1.43.0/go.mod h1:/QJhyVBUUswCphDVxq+8mld+AvhXZLhe+8WVFxiFff0= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 h1:R84qjqJb5nVJMxqWYb3np9L5ZsaDtB+a39EqjV0JSUM= golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0/go.mod h1:S9Xr4PYopiDyqSyp5NjCrhFrqg6A5zA2E/iPHPhqnS8= -golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI= -golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY= -golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA= -golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs= +golang.org/x/mod v0.35.0 h1:Ww1D637e6Pg+Zb2KrWfHQUnH2dQRLBQyAtpr/haaJeM= +golang.org/x/mod v0.35.0/go.mod h1:+GwiRhIInF8wPm+4AoT6L0FA1QWAad3OMdTRx4tFYlU= golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI= -golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -golang.org/x/text v0.36.0 h1:JfKh3XmcRPqZPKevfXVpI1wXPTqbkE5f7JA92a55Yxg= -golang.org/x/text v0.36.0/go.mod h1:NIdBknypM8iqVmPiuco0Dh6P5Jcdk8lJL0CUebqK164= -golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s= -golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0= -gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= -gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 h1:RmoJA1ujG+/lRGNfUnOMfhCy5EipVMyvUE+KNbPbTlw= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= -google.golang.org/grpc v1.82.1 h1:NnAxzGRA0677vCa4BUkOAnO5+FfQqVl9iUXeD0IqcGE= -google.golang.org/grpc v1.82.1/go.mod h1:yzTZ1TB1Z3SG+LIYaI+WiE8D5+PZ3ArnrSp8zF3+/ZA= -google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= -google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= +golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/tools v0.44.0 h1:UP4ajHPIcuMjT1GqzDWRlalUEoY+uzoZKnhOjbIPD2c= +golang.org/x/tools v0.44.0/go.mod h1:KA0AfVErSdxRZIsOVipbv3rQhVXTnlU6UhKxHd1seDI= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= kernel.org/pub/linux/libs/security/libcap/psx v1.2.77 h1:Z06sMOzc0GNCwp6efaVrIrz4ywGJ1v+DP0pjVkOfDuA= kernel.org/pub/linux/libs/security/libcap/psx v1.2.77/go.mod h1:+l6Ee2F59XiJ2I6WR5ObpC1utCQJZ/VLsEbQCD8RG24= diff --git a/dify-agent-runtime/internal/agentcli/client.go b/dify-agent-runtime/internal/agentcli/client.go index 328386cdac1..07c80515954 100644 --- a/dify-agent-runtime/internal/agentcli/client.go +++ b/dify-agent-runtime/internal/agentcli/client.go @@ -2,13 +2,12 @@ package agentcli import "context" -// StubClient abstracts Agent Stub control-plane and data-plane operations. -// Business logic depends only on this interface, never on HTTP/gRPC details. +// StubClient abstracts Agent Stub HTTP control-plane and file data-plane operations. type StubClient interface { - // Control-plane: available via gRPC or HTTP + // HTTP control-plane Connect(ctx context.Context, argv []string, metadataJSON string) (*ConnectResponse, error) CreateFileUploadURL(ctx context.Context, filename, mimetype string) (string, error) - CreateFileDownloadURL(ctx context.Context, transferMethod string, reference, url *string, forExternal bool) (*FileDownloadResponse, error) + CreateFileDownloadURL(ctx context.Context, transferMethod string, reference, url *string, forFrontend bool) (*FileDownloadResponse, error) // Drive operations (HTTP-only control-plane) GetDriveManifest(ctx context.Context, prefix string, includeDownloadURL bool) (*DriveManifestResponse, error) @@ -16,8 +15,7 @@ type StubClient interface { // Config operations (HTTP-only control-plane) GetConfigManifest(ctx context.Context) ([]byte, error) - PullConfigSkill(ctx context.Context, name string) ([]byte, error) - PullConfigFile(ctx context.Context, name string) ([]byte, error) + CreateConfigDownloadURL(ctx context.Context, kind, name string) (*FileDownloadResponse, error) PushConfig(ctx context.Context, payload any) ([]byte, error) PatchConfigEnv(ctx context.Context, envText string) ([]byte, error) PutConfigNote(ctx context.Context, note string) ([]byte, error) @@ -29,17 +27,13 @@ type StubClient interface { Close() error } -// NewStubClient creates the appropriate client based on the endpoint scheme. +// NewStubClient creates the HTTP Agent Stub client. func NewStubClient(env *Environment) (StubClient, error) { endpoint, err := ParseEndpoint(env.URL) if err != nil { return nil, err } - - httpClient := newHTTPStubClient(env) - - if endpoint.IsGRPC { - return newGRPCStubClient(endpoint, httpClient) - } - return httpClient, nil + normalizedEnv := *env + normalizedEnv.URL = endpoint.URL + return newHTTPStubClient(&normalizedEnv), nil } diff --git a/dify-agent-runtime/internal/agentcli/client_grpc.go b/dify-agent-runtime/internal/agentcli/client_grpc.go deleted file mode 100644 index 6a5f14b715d..00000000000 --- a/dify-agent-runtime/internal/agentcli/client_grpc.go +++ /dev/null @@ -1,69 +0,0 @@ -package agentcli - -import ( - "context" - "fmt" - - "github.com/langgenius/dify/dify-agent-runtime/internal/stubclient" -) - -// grpcStubClient implements StubClient with gRPC for Connect/FileUpload/FileDownload -// and delegates all other operations to the embedded HTTP client. -type grpcStubClient struct { - *httpStubClient - grpc *stubclient.Client -} - -func newGRPCStubClient(endpoint *Endpoint, httpClient *httpStubClient) (*grpcStubClient, error) { - target := endpoint.Host + ":" + endpoint.Port - client, err := stubclient.Dial(target) - if err != nil { - return nil, fmt.Errorf("gRPC dial %s: %w", target, err) - } - return &grpcStubClient{ - httpStubClient: httpClient, - grpc: client, - }, nil -} - -func (c *grpcStubClient) Close() error { - return c.grpc.Close() -} - -func (c *grpcStubClient) Connect(ctx context.Context, argv []string, metadataJSON string) (*ConnectResponse, error) { - if metadataJSON == "" { - metadataJSON = "{}" - } - result, err := c.grpc.Connect(ctx, agentStubProtocolVersion, argv, metadataJSON) - if err != nil { - return nil, err - } - return &ConnectResponse{ - ConnectionID: result.ConnectionID, - Status: result.Status, - }, nil -} - -func (c *grpcStubClient) CreateFileUploadURL(ctx context.Context, filename, mimetype string) (string, error) { - result, err := c.grpc.CreateFileUpload(ctx, filename, mimetype) - if err != nil { - return "", err - } - if result.UploadURL == "" { - return "", fmt.Errorf("signed file upload response is missing upload_url") - } - return result.UploadURL, nil -} - -func (c *grpcStubClient) CreateFileDownloadURL(ctx context.Context, transferMethod string, reference, url *string, forExternal bool) (*FileDownloadResponse, error) { - result, err := c.grpc.CreateFileDownload(ctx, transferMethod, reference, url, forExternal) - if err != nil { - return nil, err - } - return &FileDownloadResponse{ - Filename: result.Filename, - MimeType: result.MimeType, - Size: result.Size, - DownloadURL: result.DownloadURL, - }, nil -} diff --git a/dify-agent-runtime/internal/agentcli/client_http.go b/dify-agent-runtime/internal/agentcli/client_http.go index 2a6031a54fd..fa01c818711 100644 --- a/dify-agent-runtime/internal/agentcli/client_http.go +++ b/dify-agent-runtime/internal/agentcli/client_http.go @@ -72,7 +72,7 @@ func (c *httpStubClient) CreateFileUploadURL(_ context.Context, filename, mimety return resp.UploadURL, nil } -func (c *httpStubClient) CreateFileDownloadURL(_ context.Context, transferMethod string, reference, url *string, forExternal bool) (*FileDownloadResponse, error) { +func (c *httpStubClient) CreateFileDownloadURL(_ context.Context, transferMethod string, reference, url *string, forFrontend bool) (*FileDownloadResponse, error) { fileMapping := map[string]any{ "transfer_method": transferMethod, } @@ -85,7 +85,7 @@ func (c *httpStubClient) CreateFileDownloadURL(_ context.Context, transferMethod payload := map[string]any{ "file": fileMapping, - "for_external": forExternal, + "for_frontend": forFrontend, } body, statusCode, err := c.http.postJSON("/files/download-request", payload) if err != nil { @@ -152,26 +152,33 @@ func (c *httpStubClient) GetConfigManifest(_ context.Context) ([]byte, error) { return body, nil } -func (c *httpStubClient) PullConfigSkill(_ context.Context, name string) ([]byte, error) { - body, statusCode, err := c.http.getRaw(fmt.Sprintf("/config/skills/%s/pull", name), nil) +func (c *httpStubClient) CreateConfigDownloadURL( + _ context.Context, + kind, name string, +) (*FileDownloadResponse, error) { + payload := map[string]any{ + "config": map[string]string{ + "kind": kind, + "name": name, + }, + "for_frontend": false, + } + body, statusCode, err := c.http.postJSON("/files/download-request", payload) if err != nil { return nil, err } - if err := checkHTTPError(body, statusCode, "config skill pull"); err != nil { + if err := checkHTTPError(body, statusCode, "config download request"); err != nil { return nil, err } - return body, nil -} -func (c *httpStubClient) PullConfigFile(_ context.Context, name string) ([]byte, error) { - body, statusCode, err := c.http.getRaw(fmt.Sprintf("/config/files/%s/pull", name), nil) - if err != nil { - return nil, err + var resp FileDownloadResponse + if err := json.Unmarshal(body, &resp); err != nil { + return nil, fmt.Errorf("parse config download response: %w", err) } - if err := checkHTTPError(body, statusCode, "config file pull"); err != nil { - return nil, err + if resp.DownloadURL == "" { + return nil, fmt.Errorf("signed config download response is missing download_url") } - return body, nil + return &resp, nil } func (c *httpStubClient) PushConfig(_ context.Context, payload any) ([]byte, error) { diff --git a/dify-agent-runtime/internal/agentcli/config.go b/dify-agent-runtime/internal/agentcli/config.go index 6b00d0147b3..d627625c542 100644 --- a/dify-agent-runtime/internal/agentcli/config.go +++ b/dify-agent-runtime/internal/agentcli/config.go @@ -74,10 +74,14 @@ func RunConfigSkillsPull(env *Environment, names []string, localDir string, json var items []pullItem for _, name := range names { - archiveBytes, err := client.PullConfigSkill(ctx, name) + download, err := client.CreateConfigDownloadURL(ctx, "skill", name) if err != nil { return err } + archiveBytes, err := client.DownloadFromURL(download.DownloadURL) + if err != nil { + return fmt.Errorf("download config skill %q: %w", name, err) + } archivePath := filepath.Join(targetDir, name+".zip") skillDir := filepath.Join(targetDir, name) @@ -168,10 +172,14 @@ func RunConfigFilesPull(env *Environment, names []string, localDir string, jsonO var items []fileItem for _, name := range names { - payload, err := client.PullConfigFile(ctx, name) + download, err := client.CreateConfigDownloadURL(ctx, "file", name) if err != nil { return err } + payload, err := client.DownloadFromURL(download.DownloadURL) + if err != nil { + return fmt.Errorf("download config file %q: %w", name, err) + } targetPath := filepath.Join(targetDir, name) if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil { diff --git a/dify-agent-runtime/internal/agentcli/config_test.go b/dify-agent-runtime/internal/agentcli/config_test.go new file mode 100644 index 00000000000..8f67d5c0863 --- /dev/null +++ b/dify-agent-runtime/internal/agentcli/config_test.go @@ -0,0 +1,177 @@ +package agentcli + +import ( + "archive/zip" + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestConfigPullRequestsURLThenDownloadsFromDataPlane(t *testing.T) { + skillArchive := zipFixture(t, map[string]string{"SKILL.md": "# Alpha\n", "reference.md": "guide"}) + tests := []struct { + name string + kind string + assetName string + payload []byte + run func(*Environment, string) error + assertFiles func(*testing.T, string) + }{ + { + name: "file", + kind: "file", + assetName: "guide.txt", + payload: []byte("guide"), + run: func(env *Environment, targetDir string) error { + return RunConfigFilesPull(env, []string{"guide.txt"}, targetDir, true) + }, + assertFiles: func(t *testing.T, targetDir string) { + data, err := os.ReadFile(filepath.Join(targetDir, "guide.txt")) + if err != nil { + t.Fatalf("read pulled config file: %v", err) + } + if string(data) != "guide" { + t.Fatalf("pulled config file = %q", data) + } + }, + }, + { + name: "skill", + kind: "skill", + assetName: "alpha", + payload: skillArchive, + run: func(env *Environment, targetDir string) error { + return RunConfigSkillsPull(env, []string{"alpha"}, targetDir, true) + }, + assertFiles: func(t *testing.T, targetDir string) { + data, err := os.ReadFile(filepath.Join(targetDir, "alpha", "SKILL.md")) + if err != nil { + t.Fatalf("read pulled skill: %v", err) + } + if string(data) != "# Alpha\n" { + t.Fatalf("pulled SKILL.md = %q", data) + } + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + var controlPayload map[string]any + dataPlaneCalls := 0 + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/agent-stub/files/download-request": + if err := json.NewDecoder(r.Body).Decode(&controlPayload); err != nil { + http.Error(w, "bad request", http.StatusBadRequest) + return + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "filename": test.assetName, + "mime_type": "application/octet-stream", + "size": len(test.payload), + "download_url": server.URL + "/files/config-asset", + }) + case "/files/config-asset": + dataPlaneCalls++ + _, _ = w.Write(test.payload) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + targetDir := t.TempDir() + err := test.run(&Environment{URL: server.URL + "/agent-stub", AuthJWE: "token"}, targetDir) + if err != nil { + t.Fatalf("pull config %s: %v", test.kind, err) + } + + config, ok := controlPayload["config"].(map[string]any) + if !ok { + t.Fatalf("control payload config = %#v", controlPayload["config"]) + } + if config["kind"] != test.kind || config["name"] != test.assetName { + t.Fatalf("control payload config = %#v", config) + } + if controlPayload["for_frontend"] != false { + t.Fatalf("for_frontend = %#v", controlPayload["for_frontend"]) + } + if dataPlaneCalls != 1 { + t.Fatalf("data-plane calls = %d, want 1", dataPlaneCalls) + } + test.assertFiles(t, targetDir) + }) + } +} + +func TestConfigPullReportsControlAndDataPlaneFailures(t *testing.T) { + tests := []struct { + name string + controlStatus int + dataStatus int + want string + }{ + {name: "control", controlStatus: http.StatusNotFound, dataStatus: http.StatusOK, want: "config download request failed"}, + {name: "data plane", controlStatus: http.StatusOK, dataStatus: http.StatusBadGateway, want: "download config file"}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/agent-stub/files/download-request": + if test.controlStatus != http.StatusOK { + w.WriteHeader(test.controlStatus) + _, _ = w.Write([]byte(`{"detail":"missing"}`)) + return + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "filename": "guide.txt", "size": 5, "download_url": server.URL + "/files/config-asset", + }) + case "/files/config-asset": + w.WriteHeader(test.dataStatus) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + err := RunConfigFilesPull( + &Environment{URL: server.URL + "/agent-stub", AuthJWE: "token"}, + []string{"guide.txt"}, + t.TempDir(), + true, + ) + if err == nil || !strings.Contains(err.Error(), test.want) { + t.Fatalf("error = %v, want substring %q", err, test.want) + } + }) + } +} + +func zipFixture(t *testing.T, files map[string]string) []byte { + t.Helper() + var buffer bytes.Buffer + archive := zip.NewWriter(&buffer) + for name, content := range files { + writer, err := archive.Create(name) + if err != nil { + t.Fatalf("create zip member: %v", err) + } + if _, err := writer.Write([]byte(content)); err != nil { + t.Fatalf("write zip member: %v", err) + } + } + if err := archive.Close(); err != nil { + t.Fatalf("close zip fixture: %v", err) + } + return buffer.Bytes() +} diff --git a/dify-agent-runtime/internal/agentcli/connect.go b/dify-agent-runtime/internal/agentcli/connect.go index 6d17cc45e6f..3a81bc1155c 100644 --- a/dify-agent-runtime/internal/agentcli/connect.go +++ b/dify-agent-runtime/internal/agentcli/connect.go @@ -6,8 +6,6 @@ import ( "fmt" ) -const agentStubProtocolVersion = 1 - // ConnectResponse is the JSON output for `dify-agent connect --json`. type ConnectResponse struct { ConnectionID string `json:"connection_id"` diff --git a/dify-agent-runtime/internal/agentcli/env.go b/dify-agent-runtime/internal/agentcli/env.go index 1db58b057d2..04d6f5cf16d 100644 --- a/dify-agent-runtime/internal/agentcli/env.go +++ b/dify-agent-runtime/internal/agentcli/env.go @@ -1,6 +1,5 @@ // Package agentcli implements the dify-agent CLI that runs inside the sandbox -// container. It communicates with the Agent Stub server on the host via HTTP -// or gRPC to provide connect, file, drive, and config operations. +// container. It communicates with the Agent Stub server on the host via HTTP. package agentcli import ( @@ -27,13 +26,9 @@ type Environment struct { AuthJWE string } -// Endpoint represents a parsed Agent Stub endpoint with transport info. +// Endpoint represents a normalized HTTP Agent Stub endpoint. type Endpoint struct { - URL string - Scheme string // "http", "https", or "grpc" - Host string - Port string - IsGRPC bool + URL string } var ErrMissingEnvironment = errors.New("missing required Agent Stub environment variables") @@ -90,14 +85,10 @@ func ParseEndpoint(rawURL string) (*Endpoint, error) { return nil, fmt.Errorf("invalid URL: %w", err) } - switch parsed.Scheme { - case "http", "https": - return parseHTTPEndpoint(parsed) - case "grpc": - return parseGRPCEndpoint(parsed) - default: - return nil, errors.New("agent stub URL must use http, https, or grpc") + if parsed.Scheme != "http" && parsed.Scheme != "https" { + return nil, errors.New("agent stub URL must use http or https") } + return parseHTTPEndpoint(parsed) } func parseHTTPEndpoint(parsed *url.URL) (*Endpoint, error) { @@ -120,33 +111,6 @@ func parseHTTPEndpoint(parsed *url.URL) (*Endpoint, error) { normalizedURL := fmt.Sprintf("%s://%s%s", parsed.Scheme, parsed.Host, path) return &Endpoint{ - URL: normalizedURL, - Scheme: parsed.Scheme, - Host: parsed.Hostname(), - Port: parsed.Port(), - IsGRPC: false, - }, nil -} - -func parseGRPCEndpoint(parsed *url.URL) (*Endpoint, error) { - if parsed.Host == "" { - return nil, errors.New("gRPC agent stub URL must include a host") - } - path := strings.TrimRight(parsed.Path, "/") - if path != "" && path != "/" { - return nil, errors.New("gRPC agent stub URL must not include a path") - } - port := parsed.Port() - if port == "" { - return nil, errors.New("gRPC agent stub URL must include an explicit port") - } - - normalizedURL := fmt.Sprintf("grpc://%s:%s", parsed.Hostname(), port) - return &Endpoint{ - URL: normalizedURL, - Scheme: "grpc", - Host: parsed.Hostname(), - Port: port, - IsGRPC: true, + URL: normalizedURL, }, nil } diff --git a/dify-agent-runtime/internal/agentcli/env_test.go b/dify-agent-runtime/internal/agentcli/env_test.go index 6d02f8014a8..d5b333c82e3 100644 --- a/dify-agent-runtime/internal/agentcli/env_test.go +++ b/dify-agent-runtime/internal/agentcli/env_test.go @@ -1,6 +1,10 @@ package agentcli import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" "testing" ) @@ -58,55 +62,6 @@ func TestParseEndpoint_HTTP(t *testing.T) { if ep.URL != tc.wantURL { t.Errorf("URL = %q, want %q", ep.URL, tc.wantURL) } - if ep.IsGRPC { - t.Error("expected IsGRPC=false") - } - }) - } -} - -func TestParseEndpoint_GRPC(t *testing.T) { - tests := []struct { - name string - input string - wantURL string - wantErr bool - }{ - { - name: "valid grpc endpoint", - input: "grpc://localhost:50051", - wantURL: "grpc://localhost:50051", - }, - { - name: "grpc without port rejects", - input: "grpc://localhost", - wantErr: true, - }, - { - name: "grpc with path rejects", - input: "grpc://localhost:50051/some-path", - wantErr: true, - }, - } - - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - ep, err := ParseEndpoint(tc.input) - if tc.wantErr { - if err == nil { - t.Fatal("expected error") - } - return - } - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if ep.URL != tc.wantURL { - t.Errorf("URL = %q, want %q", ep.URL, tc.wantURL) - } - if !ep.IsGRPC { - t.Error("expected IsGRPC=true") - } }) } } @@ -132,6 +87,34 @@ func TestParseEndpoint_Invalid(t *testing.T) { } } +func TestNewStubClient_NormalizesServiceRootWithoutMutatingEnvironment(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/agent-stub/connections" { + t.Errorf("request path = %q, want %q", r.URL.Path, "/agent-stub/connections") + } + _ = json.NewEncoder(w).Encode(ConnectResponse{ConnectionID: "connection-1", Status: "connected"}) + })) + defer server.Close() + + env := &Environment{URL: server.URL, AuthJWE: "test-token"} + client, err := NewStubClient(env) + if err != nil { + t.Fatalf("NewStubClient() error = %v", err) + } + defer func() { _ = client.Close() }() + + response, err := client.Connect(context.Background(), nil, "") + if err != nil { + t.Fatalf("Connect() error = %v", err) + } + if response.ConnectionID != "connection-1" { + t.Errorf("connection ID = %q, want %q", response.ConnectionID, "connection-1") + } + if env.URL != server.URL { + t.Errorf("caller environment URL = %q, want unchanged %q", env.URL, server.URL) + } +} + func TestReadEnvironment_Missing(t *testing.T) { t.Setenv(EnvAPIBaseURL, "") t.Setenv(EnvAuthJWE, "") diff --git a/dify-agent-runtime/internal/agentcli/file.go b/dify-agent-runtime/internal/agentcli/file.go index d0f8ab75f1e..b1f5179bb60 100644 --- a/dify-agent-runtime/internal/agentcli/file.go +++ b/dify-agent-runtime/internal/agentcli/file.go @@ -2,8 +2,10 @@ package agentcli import ( "context" + "encoding/base64" "encoding/json" "fmt" + "io" "mime" "os" "path/filepath" @@ -12,9 +14,9 @@ import ( // FileUploadResponse is the JSON output for `dify-agent file upload`. type FileUploadResponse struct { - TransferMethod string `json:"transfer_method"` - Reference string `json:"reference"` - DownloadURL string `json:"download_url"` + TransferMethod string `json:"transfer_method"` + Reference string `json:"reference"` + PublicDownloadURL string `json:"public_download_url,omitempty"` } // FileDownloadResponse is the response from a file download request. @@ -26,7 +28,28 @@ type FileDownloadResponse struct { } // RunFileUpload executes the `file upload` command. -func RunFileUpload(env *Environment, path string) error { +func RunFileUpload(env *Environment, path string, noDownloadLink bool) error { + client, err := NewStubClient(env) + if err != nil { + return err + } + defer func() { _ = client.Close() }() + + return runFileUpload(client, path, noDownloadLink, os.Stdout) +} + +type fileUploadClient interface { + CreateFileUploadURL(ctx context.Context, filename, mimetype string) (string, error) + UploadFileToURL(uploadURL, filePath, filename, mimetype string) ([]byte, error) + CreateFileDownloadURL( + ctx context.Context, + transferMethod string, + reference, url *string, + forFrontend bool, + ) (*FileDownloadResponse, error) +} + +func runFileUpload(client fileUploadClient, path string, noDownloadLink bool, output io.Writer) error { absPath, err := filepath.Abs(path) if err != nil { return fmt.Errorf("resolve path: %w", err) @@ -40,12 +63,6 @@ func RunFileUpload(env *Environment, path string) error { mimetype := guessMIMEType(filename) ctx := context.Background() - client, err := NewStubClient(env) - if err != nil { - return err - } - defer func() { _ = client.Close() }() - // Step 1: Request a signed upload URL uploadURL, err := client.CreateFileUploadURL(ctx, filename, mimetype) if err != nil { @@ -67,24 +84,47 @@ func RunFileUpload(env *Environment, path string) error { if reference == "" { return fmt.Errorf("signed file upload response is missing reference") } - - // Step 3: Request download URL for the uploaded file - ref := reference - dlResp, err := client.CreateFileDownloadURL(ctx, "tool_file", &ref, nil, false) - if err != nil { - return err + if noDownloadLink && !isCanonicalDifyFileReference(reference) { + return fmt.Errorf("signed file upload response has invalid reference") } result := FileUploadResponse{ TransferMethod: "tool_file", Reference: reference, - DownloadURL: dlResp.DownloadURL, + } + if !noDownloadLink { + // Step 3: Request a browser-visible URL unless the caller only needs the + // canonical ToolFile reference. + ref := reference + dlResp, err := client.CreateFileDownloadURL(ctx, "tool_file", &ref, nil, true) + if err != nil { + return err + } + result.PublicDownloadURL = dlResp.DownloadURL } out, _ := json.Marshal(result) - fmt.Println(string(out)) + _, _ = fmt.Fprintln(output, string(out)) return nil } +func isCanonicalDifyFileReference(reference string) bool { + encodedPayload, found := strings.CutPrefix(reference, "dify-file-ref:") + if !found || encodedPayload == "" { + return false + } + payloadJSON, err := base64.URLEncoding.DecodeString(encodedPayload) + if err != nil { + return false + } + var payload struct { + RecordID string `json:"record_id"` + } + if err := json.Unmarshal(payloadJSON, &payload); err != nil { + return false + } + return payload.RecordID != "" +} + // RunFileDownload executes the `file download` command. func RunFileDownload(env *Environment, transferMethod string, referenceOrURL string, localDir string) error { var reference *string diff --git a/dify-agent-runtime/internal/agentcli/file_test.go b/dify-agent-runtime/internal/agentcli/file_test.go new file mode 100644 index 00000000000..a610457e6c2 --- /dev/null +++ b/dify-agent-runtime/internal/agentcli/file_test.go @@ -0,0 +1,250 @@ +package agentcli + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" +) + +type fakeFileUploadClient struct { + forFrontend bool + downloadRequestCall int + uploadResponse []byte + calls []string + filename string + mimetype string + uploadURL string + uploadedBytes []byte + downloadReference string +} + +func (f *fakeFileUploadClient) CreateFileUploadURL(_ context.Context, filename, mimetype string) (string, error) { + f.calls = append(f.calls, "upload-request") + f.filename = filename + f.mimetype = mimetype + return "https://sandbox-files.example.com/files/upload/for-plugin?sign=1", nil +} + +func (f *fakeFileUploadClient) UploadFileToURL(uploadURL, filePath, filename, mimetype string) ([]byte, error) { + f.calls = append(f.calls, "multipart-upload") + f.uploadURL = uploadURL + f.filename = filename + f.mimetype = mimetype + f.uploadedBytes, _ = os.ReadFile(filePath) + if f.uploadResponse != nil { + return f.uploadResponse, nil + } + return []byte(`{"reference":"dify-file-ref:eyJyZWNvcmRfaWQiOiJ0b29sLTEifQ=="}`), nil +} + +func (f *fakeFileUploadClient) CreateFileDownloadURL( + _ context.Context, + _ string, + reference, _ *string, + forFrontend bool, +) (*FileDownloadResponse, error) { + f.calls = append(f.calls, "download-request") + f.forFrontend = forFrontend + f.downloadRequestCall++ + if reference != nil { + f.downloadReference = *reference + } + return &FileDownloadResponse{ + Filename: "report.pdf", + MimeType: "application/pdf", + Size: 123, + DownloadURL: "/files/tools/report.pdf?sign=2", + }, nil +} + +func TestRunFileUploadReturnsFrontendDisplayURL(t *testing.T) { + filePath := t.TempDir() + "/report.pdf" + if err := os.WriteFile(filePath, []byte("report"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + client := &fakeFileUploadClient{} + var output bytes.Buffer + if err := runFileUpload(client, filePath, false, &output); err != nil { + t.Fatalf("run file upload: %v", err) + } + + if !client.forFrontend { + t.Fatal("download request did not select frontend display URL") + } + if got, want := strings.Join(client.calls, ","), "upload-request,multipart-upload,download-request"; got != want { + t.Fatalf("call order = %s, want %s", got, want) + } + if client.filename != "report.pdf" || client.mimetype != "application/pdf" { + t.Fatalf("upload metadata = (%q, %q), want report.pdf/application/pdf", client.filename, client.mimetype) + } + if client.uploadURL != "https://sandbox-files.example.com/files/upload/for-plugin?sign=1" { + t.Fatalf("upload URL = %q", client.uploadURL) + } + if string(client.uploadedBytes) != "report" { + t.Fatalf("uploaded bytes = %q, want report", client.uploadedBytes) + } + if client.downloadReference != "dify-file-ref:eyJyZWNvcmRfaWQiOiJ0b29sLTEifQ==" { + t.Fatalf("download reference = %q", client.downloadReference) + } + got := strings.TrimSpace(output.String()) + want := `{"transfer_method":"tool_file","reference":"dify-file-ref:eyJyZWNvcmRfaWQiOiJ0b29sLTEifQ==","public_download_url":"/files/tools/report.pdf?sign=2"}` + if got != want { + t.Fatalf("output = %s, want %s", got, want) + } +} + +func TestRunFileUploadWithoutDownloadLinkReturnsOnlyCanonicalMapping(t *testing.T) { + filePath := filepath.Join(t.TempDir(), "report.pdf") + if err := os.WriteFile(filePath, []byte("report"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + client := &fakeFileUploadClient{} + var output bytes.Buffer + if err := runFileUpload(client, filePath, true, &output); err != nil { + t.Fatalf("run file upload: %v", err) + } + + if client.downloadRequestCall != 0 { + t.Fatalf("download request calls = %d, want 0", client.downloadRequestCall) + } + if got, want := strings.Join(client.calls, ","), "upload-request,multipart-upload"; got != want { + t.Fatalf("call order = %s, want %s", got, want) + } + got := strings.TrimSpace(output.String()) + want := `{"transfer_method":"tool_file","reference":"dify-file-ref:eyJyZWNvcmRfaWQiOiJ0b29sLTEifQ=="}` + if got != want { + t.Fatalf("output = %s, want %s", got, want) + } +} + +func TestRunFileUploadDefaultAcceptsLegacyNonemptyReference(t *testing.T) { + filePath := filepath.Join(t.TempDir(), "report.pdf") + if err := os.WriteFile(filePath, []byte("report"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + client := &fakeFileUploadClient{uploadResponse: []byte(`{"reference":"raw-id"}`)} + var output bytes.Buffer + err := runFileUpload(client, filePath, false, &output) + if err != nil { + t.Fatalf("run file upload: %v", err) + } + if client.downloadRequestCall != 1 || client.downloadReference != "raw-id" { + t.Fatalf("download request = (%d, %q), want legacy reference", client.downloadRequestCall, client.downloadReference) + } + if !strings.Contains(output.String(), `"reference":"raw-id"`) { + t.Fatalf("output = %q, want legacy reference", output.String()) + } +} + +func TestRunFileUploadWithoutDownloadLinkRejectsNonCanonicalReference(t *testing.T) { + filePath := filepath.Join(t.TempDir(), "report.pdf") + if err := os.WriteFile(filePath, []byte("report"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + client := &fakeFileUploadClient{uploadResponse: []byte(`{"reference":"raw-id"}`)} + var output bytes.Buffer + err := runFileUpload(client, filePath, true, &output) + if err == nil || !strings.Contains(err.Error(), "invalid reference") { + t.Fatalf("error = %v, want invalid reference", err) + } + if client.downloadRequestCall != 0 { + t.Fatalf("download request calls = %d, want 0", client.downloadRequestCall) + } + if output.Len() != 0 { + t.Fatalf("output = %q, want empty", output.String()) + } +} + +func TestRunFileUploadRejectsMissingReferenceInBothModes(t *testing.T) { + filePath := filepath.Join(t.TempDir(), "report.pdf") + if err := os.WriteFile(filePath, []byte("report"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + for _, noDownloadLink := range []bool{false, true} { + client := &fakeFileUploadClient{uploadResponse: []byte(`{"reference":""}`)} + var output bytes.Buffer + err := runFileUpload(client, filePath, noDownloadLink, &output) + if err == nil || !strings.Contains(err.Error(), "missing reference") { + t.Fatalf("noDownloadLink=%t error = %v, want missing reference", noDownloadLink, err) + } + if client.downloadRequestCall != 0 { + t.Fatalf("noDownloadLink=%t download request calls = %d, want 0", noDownloadLink, client.downloadRequestCall) + } + } +} + +func TestRunFileUploadRejectsInvalidUploadResponseBeforeDownloadRequest(t *testing.T) { + filePath := filepath.Join(t.TempDir(), "report.pdf") + if err := os.WriteFile(filePath, []byte("report"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + client := &fakeFileUploadClient{uploadResponse: []byte("not-json")} + var output bytes.Buffer + err := runFileUpload(client, filePath, true, &output) + if err == nil || !strings.Contains(err.Error(), "parse upload result") { + t.Fatalf("error = %v, want parse upload result failure", err) + } + if client.downloadRequestCall != 0 { + t.Fatalf("download request calls = %d, want 0", client.downloadRequestCall) + } +} + +func TestRunFileDownloadRequestsSandboxURLAndWritesFile(t *testing.T) { + var requestPayload map[string]json.RawMessage + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/agent-stub/files/download-request": + if err := json.NewDecoder(r.Body).Decode(&requestPayload); err != nil { + t.Errorf("decode download request: %v", err) + http.Error(w, "bad request", http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"filename":"report.pdf","mime_type":"application/pdf","size":6,"download_url":"` + server.URL + `/files/report.pdf"}`)) + case "/files/report.pdf": + _, _ = w.Write([]byte("report")) + default: + http.NotFound(w, r) + } + })) + defer server.Close() + + targetDir := t.TempDir() + err := RunFileDownload( + &Environment{URL: server.URL + "/agent-stub", AuthJWE: "test-token"}, + "tool_file", + "dify-file-ref:canonical", + targetDir, + ) + if err != nil { + t.Fatalf("run file download: %v", err) + } + + var forFrontend bool + if err := json.Unmarshal(requestPayload["for_frontend"], &forFrontend); err != nil { + t.Fatalf("decode for_frontend: %v", err) + } + if forFrontend { + t.Fatal("download request selected a frontend URL") + } + data, err := os.ReadFile(filepath.Join(targetDir, "report.pdf")) + if err != nil { + t.Fatalf("read downloaded file: %v", err) + } + if string(data) != "report" { + t.Fatalf("downloaded file = %q, want report", data) + } +} diff --git a/dify-agent-runtime/internal/agentcli/httpclient.go b/dify-agent-runtime/internal/agentcli/httpclient.go index 95c3614bf1b..0c2c156c16b 100644 --- a/dify-agent-runtime/internal/agentcli/httpclient.go +++ b/dify-agent-runtime/internal/agentcli/httpclient.go @@ -2,7 +2,9 @@ package agentcli import ( "bytes" + "context" "encoding/json" + "errors" "fmt" "io" "mime/multipart" @@ -14,29 +16,45 @@ import ( // HTTPClient wraps HTTP interactions with the Agent Stub server. type HTTPClient struct { - baseURL string - authJWE string - client *http.Client + baseURL string + authJWE string + client *http.Client + openUploadFile func(string) (io.ReadCloser, error) + doUploadRequest func(*http.Request) (*http.Response, error) } +var errUploadRequestAborted = errors.New("upload request aborted") + // NewHTTPClient creates a new HTTP client for the Agent Stub API. func NewHTTPClient(env *Environment) *HTTPClient { return &HTTPClient{ - baseURL: env.URL, - authJWE: env.AuthJWE, - client: &http.Client{Timeout: 30 * time.Second}, + baseURL: env.URL, + authJWE: env.AuthJWE, + client: &http.Client{Timeout: 30 * time.Second}, + openUploadFile: openUploadSource, + doUploadRequest: doUploadRequest, } } // NewHTTPClientWithTimeout creates a client with a custom timeout. func NewHTTPClientWithTimeout(env *Environment, timeout time.Duration) *HTTPClient { return &HTTPClient{ - baseURL: env.URL, - authJWE: env.AuthJWE, - client: &http.Client{Timeout: timeout}, + baseURL: env.URL, + authJWE: env.AuthJWE, + client: &http.Client{Timeout: timeout}, + openUploadFile: openUploadSource, + doUploadRequest: doUploadRequest, } } +func openUploadSource(path string) (io.ReadCloser, error) { + return os.Open(path) +} + +func doUploadRequest(req *http.Request) (*http.Response, error) { + return (&http.Client{Timeout: 120 * time.Second}).Do(req) +} + // postJSON sends a POST request with JSON body and returns the response body. func (c *HTTPClient) postJSON(path string, payload any) ([]byte, int, error) { body, err := json.Marshal(payload) @@ -93,11 +111,6 @@ func (c *HTTPClient) getJSON(path string, params map[string]string) ([]byte, int return respBody, resp.StatusCode, nil } -// getRaw sends a GET request and returns raw bytes (for binary downloads). -func (c *HTTPClient) getRaw(path string, params map[string]string) ([]byte, int, error) { - return c.getJSON(path, params) -} - // patchJSON sends a PATCH request with JSON body. func (c *HTTPClient) patchJSON(path string, payload any) ([]byte, int, error) { body, err := json.Marshal(payload) @@ -154,49 +167,105 @@ func (c *HTTPClient) putJSON(path string, payload any) ([]byte, int, error) { // uploadFile uploads a file to a signed URL using multipart form. func (c *HTTPClient) uploadFile(uploadURL string, filePath string, filename string, mimetype string) ([]byte, error) { - file, err := os.Open(filePath) + file, err := c.openUploadFile(filePath) if err != nil { return nil, fmt.Errorf("open file: %w", err) } - defer func() { _ = file.Close() }() + pipeReader, pipeWriter := io.Pipe() + + req, err := http.NewRequest("POST", uploadURL, pipeReader) + if err != nil { + _ = pipeReader.Close() + _ = pipeWriter.Close() + _ = file.Close() + return nil, errors.New("create upload request: invalid signed upload URL") + } + multipartWriter := multipart.NewWriter(pipeWriter) + req.Header.Set("Content-Type", multipartWriter.FormDataContentType()) + + writerDone := make(chan error, 1) + go func() { + writeErr := writeMultipartFile(multipartWriter, file, filename, mimetype) + if closeErr := file.Close(); writeErr == nil && closeErr != nil { + writeErr = fmt.Errorf("close file: %w", closeErr) + } + if writeErr != nil { + _ = pipeWriter.CloseWithError(writeErr) + } else { + writeErr = pipeWriter.Close() + } + writerDone <- writeErr + }() + + resp, requestErr := c.doUploadRequest(req) + _ = pipeReader.CloseWithError(errUploadRequestAborted) + + const maxUploadResponseBytes = 1024 * 1024 + var respBody []byte + var responseErr error + if resp != nil { + defer func() { _ = resp.Body.Close() }() + respBody, responseErr = io.ReadAll(io.LimitReader(resp.Body, maxUploadResponseBytes+1)) + } + + writerErr := <-writerDone + if resp != nil && (resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices) { + if len(respBody) > maxUploadResponseBytes { + respBody = respBody[:maxUploadResponseBytes] + } + return nil, fmt.Errorf("upload failed with status %d: %s", resp.StatusCode, string(respBody)) + } + if requestErr != nil { + if writerErr != nil && !isUploadWriterAbort(writerErr) { + return nil, writerErr + } + if errors.Is(requestErr, context.DeadlineExceeded) || os.IsTimeout(requestErr) { + return nil, errors.New("upload request timed out") + } + return nil, errors.New("upload request failed") + } + if writerErr != nil && !isUploadWriterAbort(writerErr) { + return nil, writerErr + } + if isUploadWriterAbort(writerErr) { + return nil, errors.New("upload request completed before multipart body was fully written") + } + if responseErr != nil { + return nil, fmt.Errorf("read upload response: %w", responseErr) + } + if len(respBody) > maxUploadResponseBytes { + return nil, fmt.Errorf("upload response exceeds %d bytes", maxUploadResponseBytes) + } + return respBody, nil +} + +func isUploadWriterAbort(err error) bool { + return errors.Is(err, errUploadRequestAborted) || errors.Is(err, io.ErrClosedPipe) +} + +func writeMultipartFile( + writer *multipart.Writer, + file io.Reader, + filename string, + mimetype string, +) (resultErr error) { + defer func() { + if closeErr := writer.Close(); resultErr == nil && closeErr != nil { + resultErr = fmt.Errorf("close multipart writer: %w", closeErr) + } + }() - var buf bytes.Buffer - writer := multipart.NewWriter(&buf) h := make(textproto.MIMEHeader) h.Set("Content-Disposition", fmt.Sprintf(`form-data; name="file"; filename="%s"`, filename)) h.Set("Content-Type", mimetype) part, err := writer.CreatePart(h) if err != nil { - return nil, fmt.Errorf("create form file: %w", err) + return fmt.Errorf("create form file: %w", err) } if _, err := io.Copy(part, file); err != nil { - return nil, fmt.Errorf("copy file content: %w", err) + return fmt.Errorf("copy file content: %w", err) } - if err := writer.Close(); err != nil { - return nil, fmt.Errorf("close multipart writer: %w", err) - } - - uploadClient := &http.Client{Timeout: 120 * time.Second} - req, err := http.NewRequest("POST", uploadURL, &buf) - if err != nil { - return nil, fmt.Errorf("create upload request: %w", err) - } - req.Header.Set("Content-Type", writer.FormDataContentType()) - - resp, err := uploadClient.Do(req) - if err != nil { - return nil, fmt.Errorf("upload request failed: %w", err) - } - defer func() { _ = resp.Body.Close() }() - - respBody, err := io.ReadAll(resp.Body) - if err != nil { - return nil, fmt.Errorf("read upload response: %w", err) - } - if resp.StatusCode >= 400 { - return nil, fmt.Errorf("upload failed with status %d: %s", resp.StatusCode, string(respBody)) - } - return respBody, nil + return nil } // downloadFromURL downloads bytes from a signed URL. diff --git a/dify-agent-runtime/internal/agentcli/httpclient_test.go b/dify-agent-runtime/internal/agentcli/httpclient_test.go new file mode 100644 index 00000000000..67ac9dbd11e --- /dev/null +++ b/dify-agent-runtime/internal/agentcli/httpclient_test.go @@ -0,0 +1,678 @@ +package agentcli + +import ( + "bytes" + "errors" + "fmt" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "testing" + "time" +) + +const fifoTestDeadline = 3 * time.Second + +func receiveWithin[T any](ch <-chan T, timeout time.Duration) (T, bool) { + timer := time.NewTimer(timeout) + defer timer.Stop() + select { + case value := <-ch: + return value, true + case <-timer.C: + var zero T + return zero, false + } +} + +type multipartWriteRecorder struct { + headerErr error + cleanupErr error + source *terminalRecordingReader + headerAttempted bool + headerFailed bool + cleanupAttempted bool +} + +func (w *multipartWriteRecorder) Write(p []byte) (int, error) { + if w.headerFailed || (w.source != nil && w.source.finished) { + w.cleanupAttempted = true + if w.cleanupErr != nil { + return 0, w.cleanupErr + } + return len(p), nil + } + if !w.headerAttempted { + w.headerAttempted = true + if w.headerErr != nil { + w.headerFailed = true + return 0, w.headerErr + } + } + return len(p), nil +} + +type terminalRecordingReader struct { + reader io.Reader + terminalErr error + finished bool +} + +func (r *terminalRecordingReader) Read(p []byte) (int, error) { + n, err := r.reader.Read(p) + if errors.Is(err, io.EOF) { + r.finished = true + if r.terminalErr != nil { + return n, r.terminalErr + } + } + return n, err +} + +type dataThenErrorReader struct { + data []byte + sourceErr error + delivered bool +} + +type closeRecordingBody struct { + reader io.Reader + closed bool +} + +type gatedUploadSource struct { + payload []byte + started chan struct{} + release chan struct{} + completed chan struct{} + releaseOnce sync.Once + delivered bool + closed bool +} + +func newGatedUploadSource(payload []byte) *gatedUploadSource { + return &gatedUploadSource{ + payload: payload, + started: make(chan struct{}), + release: make(chan struct{}), + completed: make(chan struct{}), + } +} + +func (s *gatedUploadSource) Read(p []byte) (int, error) { + if s.delivered { + return 0, io.EOF + } + s.delivered = true + close(s.started) + <-s.release + n := copy(p, s.payload) + close(s.completed) + return n, nil +} + +func (s *gatedUploadSource) Close() error { + s.unblock() + s.closed = true + return nil +} + +func (s *gatedUploadSource) unblock() { + s.releaseOnce.Do(func() { close(s.release) }) +} + +type releaseOnReadBody struct { + reader io.Reader + release func() + closed bool +} + +func (b *releaseOnReadBody) Read(p []byte) (int, error) { + b.release() + return b.reader.Read(p) +} + +func (b *releaseOnReadBody) Close() error { + b.release() + b.closed = true + return nil +} + +func (b *closeRecordingBody) Read(p []byte) (int, error) { + return b.reader.Read(p) +} + +func (b *closeRecordingBody) Close() error { + b.closed = true + return nil +} + +func (r *dataThenErrorReader) Read(p []byte) (int, error) { + if r.delivered { + return 0, r.sourceErr + } + r.delivered = true + return copy(p, r.data), nil +} + +func TestUploadFileStreamsMultipartBody(t *testing.T) { + payload := bytes.Repeat([]byte("streamed-payload-"), 128*1024) + filePath := filepath.Join(t.TempDir(), "payload.bin") + if err := os.WriteFile(filePath, payload, 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.ContentLength != -1 { + t.Errorf("content length = %d, want -1 for streamed request", r.ContentLength) + } + reader, err := r.MultipartReader() + if err != nil { + http.Error(w, fmt.Sprintf("create multipart reader: %v", err), http.StatusBadRequest) + return + } + part, err := reader.NextPart() + if err != nil { + http.Error(w, fmt.Sprintf("read multipart part: %v", err), http.StatusBadRequest) + return + } + defer func() { _ = part.Close() }() + if part.FormName() != "file" || part.FileName() != "payload.bin" { + http.Error(w, "unexpected multipart metadata", http.StatusBadRequest) + return + } + if got := part.Header.Get("Content-Type"); got != "application/octet-stream" { + http.Error(w, "unexpected multipart content type: "+got, http.StatusBadRequest) + return + } + got, err := io.ReadAll(part) + if err != nil { + http.Error(w, fmt.Sprintf("read multipart content: %v", err), http.StatusBadRequest) + return + } + if !bytes.Equal(got, payload) { + http.Error(w, "multipart content mismatch", http.StatusBadRequest) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"reference":"dify-file-ref:canonical"}`)) + })) + defer server.Close() + + client := NewHTTPClient(&Environment{}) + body, err := client.uploadFile(server.URL, filePath, "payload.bin", "application/octet-stream") + if err != nil { + t.Fatalf("upload file: %v", err) + } + if string(body) != `{"reference":"dify-file-ref:canonical"}` { + t.Fatalf("response body = %s", body) + } +} + +func TestUploadFileStartsRequestBeforeSourceEOF(t *testing.T) { + if testing.Short() { + t.Skip("uses a local HTTP server and FIFO coordination") + } + mkfifo, err := exec.LookPath("mkfifo") + if err != nil { + t.Skip("mkfifo is unavailable") + } + fifoPath := filepath.Join(t.TempDir(), "stream.bin") + if err := exec.Command(mkfifo, "-m", "600", fifoPath).Run(); err != nil { + t.Fatalf("create FIFO: %v", err) + } + + requestStarted := make(chan struct{}) + continueSource := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + close(requestStarted) + reader, err := r.MultipartReader() + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + part, err := reader.NextPart() + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + defer func() { _ = part.Close() }() + body, err := io.ReadAll(part) + if err != nil || string(body) != "prefix-suffix" { + http.Error(w, "unexpected streamed body", http.StatusBadRequest) + return + } + _, _ = w.Write([]byte(`{"reference":"dify-file-ref:eyJyZWNvcmRfaWQiOiJ0b29sLTEifQ=="}`)) + })) + defer server.Close() + + writerDone := make(chan error, 1) + go func() { + file, err := os.OpenFile(fifoPath, os.O_WRONLY, 0) + if err != nil { + writerDone <- err + return + } + if _, err = file.WriteString("prefix-"); err == nil { + <-continueSource + _, err = file.WriteString("suffix") + } + if closeErr := file.Close(); err == nil { + err = closeErr + } + writerDone <- err + }() + + uploadDone := make(chan error, 1) + go func() { + client := NewHTTPClient(&Environment{}) + _, err := client.uploadFile(server.URL, fifoPath, "stream.bin", "application/octet-stream") + uploadDone <- err + }() + + releaseSource := sync.OnceFunc(func() { close(continueSource) }) + writerJoined := false + uploadJoined := false + defer func() { + releaseSource() + server.CloseClientConnections() + if !writerJoined { + if err, ok := receiveWithin(writerDone, fifoTestDeadline); !ok { + t.Errorf("cleanup: FIFO writer did not finish within %s", fifoTestDeadline) + } else if err != nil { + t.Errorf("cleanup: FIFO writer failed: %v", err) + } + } + if !uploadJoined { + if err, ok := receiveWithin(uploadDone, fifoTestDeadline); !ok { + t.Errorf("cleanup: upload did not finish within %s", fifoTestDeadline) + } else if err != nil { + t.Errorf("cleanup: upload failed: %v", err) + } + } + }() + + if _, ok := receiveWithin(requestStarted, fifoTestDeadline); !ok { + t.Fatalf("HTTP request did not start within %s before the source reached EOF", fifoTestDeadline) + } + releaseSource() + + writerErr, ok := receiveWithin(writerDone, fifoTestDeadline) + if !ok { + t.Fatalf("FIFO writer did not finish within %s after source release", fifoTestDeadline) + } + writerJoined = true + if writerErr != nil { + t.Fatalf("write FIFO: %v", writerErr) + } + uploadErr, ok := receiveWithin(uploadDone, fifoTestDeadline) + if !ok { + t.Fatalf("upload did not finish within %s after source EOF", fifoTestDeadline) + } + uploadJoined = true + if uploadErr != nil { + t.Fatalf("upload FIFO: %v", uploadErr) + } +} + +func TestUploadFileTransportFailureDoesNotBlockWriter(t *testing.T) { + const querySecret = "signed-upload-credential" + uploadURL := "https://upload.example/path?X-Amz-Credential=" + querySecret + source := &closeRecordingBody{reader: strings.NewReader("payload")} + client := NewHTTPClient(&Environment{}) + client.openUploadFile = func(string) (io.ReadCloser, error) { + return source, nil + } + requestBodyReady := make(chan io.ReadCloser, 1) + client.doUploadRequest = func(req *http.Request) (*http.Response, error) { + requestBodyReady <- req.Body + return nil, errors.New("deterministic transport failure") + } + + uploadDone := make(chan error, 1) + go func() { + _, err := client.uploadFile(uploadURL, "source-path", "payload.bin", "application/octet-stream") + uploadDone <- err + }() + + var requestBody io.ReadCloser + uploadJoined := false + defer func() { + _ = source.Close() + if requestBody == nil { + requestBody, _ = receiveWithin(requestBodyReady, fifoTestDeadline) + } + if requestBody != nil { + _ = requestBody.Close() + } + if !uploadJoined { + if _, ok := receiveWithin(uploadDone, fifoTestDeadline); !ok { + t.Errorf("cleanup: transport-failure upload did not finish within %s", fifoTestDeadline) + } + } + }() + + var ok bool + requestBody, ok = receiveWithin(requestBodyReady, fifoTestDeadline) + if !ok { + t.Fatalf("upload request did not start within %s", fifoTestDeadline) + } + + err, ok := receiveWithin(uploadDone, fifoTestDeadline) + if !ok { + t.Fatalf("transport-failure upload did not finish within %s", fifoTestDeadline) + } + uploadJoined = true + + if err == nil || err.Error() != "upload request failed" { + t.Fatalf("error = %v, want upload request failure", err) + } + if !source.closed { + t.Fatal("source file was not closed after transport failure") + } + if strings.Contains(err.Error(), uploadURL) || strings.Contains(err.Error(), querySecret) { + t.Fatalf("error leaked signed upload URL credentials: %v", err) + } +} + +func TestUploadFileInvalidSignedURLDoesNotLeakCredentials(t *testing.T) { + filePath := filepath.Join(t.TempDir(), "payload.txt") + if err := os.WriteFile(filePath, []byte("payload"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + const querySecret = "signed-upload-credential" + uploadURL := "http://example.test/upload?X-Amz-Credential=" + querySecret + "\n" + + client := NewHTTPClient(&Environment{}) + _, err := client.uploadFile(uploadURL, filePath, "payload.txt", "text/plain") + if err == nil || err.Error() != "create upload request: invalid signed upload URL" { + t.Fatalf("error = %v, want invalid signed upload URL failure", err) + } + if strings.Contains(err.Error(), uploadURL) || strings.Contains(err.Error(), querySecret) { + t.Fatalf("error leaked signed upload URL credentials: %v", err) + } +} + +func TestUploadFileRejectsOversizedResponse(t *testing.T) { + filePath := filepath.Join(t.TempDir(), "payload.txt") + if err := os.WriteFile(filePath, []byte("payload"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = io.Copy(io.Discard, r.Body) + _, _ = w.Write(bytes.Repeat([]byte("x"), 1024*1024+1)) + })) + defer server.Close() + + client := NewHTTPClient(&Environment{}) + _, err := client.uploadFile(server.URL, filePath, "payload.txt", "text/plain") + if err == nil || !strings.Contains(err.Error(), "upload response exceeds") { + t.Fatalf("error = %v, want bounded-response failure", err) + } +} + +func TestUploadFileReturnsNonSuccessStatus(t *testing.T) { + filePath := filepath.Join(t.TempDir(), "payload.txt") + if err := os.WriteFile(filePath, []byte("payload"), 0o600); err != nil { + t.Fatalf("write fixture: %v", err) + } + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = io.Copy(io.Discard, r.Body) + http.Error(w, "too large", http.StatusRequestEntityTooLarge) + })) + defer server.Close() + + client := NewHTTPClient(&Environment{}) + _, err := client.uploadFile(server.URL, filePath, "payload.txt", "text/plain") + if err == nil || !strings.Contains(err.Error(), "upload failed with status 413") { + t.Fatalf("error = %v, want non-success status", err) + } +} + +func uploadFileWithGatedEarlyResponse(t *testing.T, statusCode int, body string) error { + t.Helper() + source := newGatedUploadSource([]byte("payload")) + responseBody := &releaseOnReadBody{reader: strings.NewReader(body), release: source.unblock} + requestDone := make(chan error, 1) + var orderingErr error + client := NewHTTPClient(&Environment{}) + client.openUploadFile = func(string) (io.ReadCloser, error) { + return source, nil + } + client.doUploadRequest = func(req *http.Request) (*http.Response, error) { + go func() { + _, err := io.Copy(io.Discard, req.Body) + requestDone <- err + }() + if _, ok := receiveWithin(source.started, fifoTestDeadline); !ok { + orderingErr = errors.New("upload source did not start before the early response") + source.unblock() + return nil, orderingErr + } + select { + case <-source.completed: + orderingErr = errors.New("upload source completed before the early response") + return nil, orderingErr + default: + } + return &http.Response{StatusCode: statusCode, Body: responseBody}, nil + } + + uploadDone := make(chan error, 1) + go func() { + _, err := client.uploadFile("https://upload.example/path", "source-path", "payload.bin", "application/octet-stream") + uploadDone <- err + }() + uploadJoined := false + requestJoined := false + defer func() { + source.unblock() + if !uploadJoined { + if _, ok := receiveWithin(uploadDone, fifoTestDeadline); !ok { + t.Errorf("cleanup: early-response upload did not finish within %s", fifoTestDeadline) + } + } + if !requestJoined { + if _, ok := receiveWithin(requestDone, fifoTestDeadline); !ok { + t.Errorf("cleanup: early-response request drain did not finish within %s", fifoTestDeadline) + } + } + }() + + uploadErr, ok := receiveWithin(uploadDone, fifoTestDeadline) + if !ok { + t.Fatalf("early-response upload did not finish within %s", fifoTestDeadline) + } + uploadJoined = true + if _, ok := receiveWithin(requestDone, fifoTestDeadline); !ok { + t.Fatalf("early-response request drain did not finish within %s", fifoTestDeadline) + } + requestJoined = true + if orderingErr != nil { + t.Fatalf("invalid early-response ordering: %v", orderingErr) + } + if _, ok := receiveWithin(source.completed, fifoTestDeadline); !ok { + t.Fatalf("upload source did not finish within %s after response processing", fifoTestDeadline) + } + if !source.closed { + t.Fatal("upload source was not closed before uploadFile returned") + } + if !responseBody.closed { + t.Fatal("early response body was not closed before uploadFile returned") + } + return uploadErr +} + +func TestUploadFileReturnsEarlyNonSuccessStatusInsteadOfWriterAbort(t *testing.T) { + err := uploadFileWithGatedEarlyResponse( + t, + http.StatusRequestEntityTooLarge, + "too large without reading body\n", + ) + if err == nil || !strings.Contains(err.Error(), "upload failed with status 413: too large without reading body") { + t.Fatalf("error = %v, want early HTTP 413 response", err) + } + if strings.Contains(err.Error(), "upload request aborted") || strings.Contains(err.Error(), "closed pipe") { + t.Fatalf("error exposed internal multipart abort: %v", err) + } +} + +func TestUploadFileRejectsEarlySuccessBeforeMultipartCompletes(t *testing.T) { + err := uploadFileWithGatedEarlyResponse(t, http.StatusOK, `{"reference":"incomplete"}`) + if err == nil || err.Error() != "upload request completed before multipart body was fully written" { + t.Fatalf("error = %v, want incomplete multipart failure", err) + } +} + +func TestUploadFilePropagatesSourceReadFailure(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = io.Copy(io.Discard, r.Body) + _, _ = w.Write([]byte(`{"reference":"unused"}`)) + })) + defer server.Close() + + client := NewHTTPClient(&Environment{}) + _, err := client.uploadFile(server.URL, t.TempDir(), "directory", "application/octet-stream") + if err == nil || !strings.Contains(err.Error(), "copy file content") { + t.Fatalf("error = %v, want source read failure", err) + } +} + +func TestUploadFileClosesResourcesAndJoinsWriterOnResponseOutcomes(t *testing.T) { + responseReadErr := errors.New("response read failed") + tests := []struct { + name string + statusCode int + responseReader io.Reader + wantError string + }{ + { + name: "success", + statusCode: http.StatusOK, + responseReader: strings.NewReader(`{"reference":"dify-file-ref:canonical"}`), + }, + { + name: "non-2xx", + statusCode: http.StatusRequestEntityTooLarge, + responseReader: strings.NewReader("too large"), + wantError: "upload failed with status 413", + }, + { + name: "response body read error", + statusCode: http.StatusOK, + responseReader: &dataThenErrorReader{data: []byte("partial"), sourceErr: responseReadErr}, + wantError: "read upload response: response read failed", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + source := &closeRecordingBody{reader: strings.NewReader("payload")} + var responseBody *closeRecordingBody + client := NewHTTPClient(&Environment{}) + client.openUploadFile = func(path string) (io.ReadCloser, error) { + if path != "source-path" { + t.Fatalf("source path = %q", path) + } + return source, nil + } + client.doUploadRequest = func(req *http.Request) (*http.Response, error) { + if _, err := io.Copy(io.Discard, req.Body); err != nil { + t.Fatalf("drain request body: %v", err) + } + responseBody = &closeRecordingBody{reader: tt.responseReader} + return &http.Response{StatusCode: tt.statusCode, Body: responseBody}, nil + } + + _, err := client.uploadFile("https://upload.example/path", "source-path", "payload.txt", "text/plain") + + if tt.wantError == "" { + if err != nil { + t.Fatalf("upload file: %v", err) + } + } else if err == nil || !strings.Contains(err.Error(), tt.wantError) { + t.Fatalf("error = %v, want containing %q", err, tt.wantError) + } + if !source.closed { + t.Fatal("source file was not closed before uploadFile returned") + } + if responseBody == nil || !responseBody.closed { + t.Fatal("response body was not closed before uploadFile returned") + } + }) + } +} + +func TestWriteMultipartFileAlwaysAttemptsCloseAndPreservesPrimaryError(t *testing.T) { + t.Run("create part failure", func(t *testing.T) { + createErr := errors.New("create part failed") + closeErr := errors.New("close failed") + destination := &multipartWriteRecorder{headerErr: createErr, cleanupErr: closeErr} + writer := multipart.NewWriter(destination) + + err := writeMultipartFile( + writer, + strings.NewReader("payload"), + "payload.txt", + "text/plain", + ) + + if !errors.Is(err, createErr) { + t.Fatalf("error = %v, want CreatePart failure", err) + } + if !destination.headerAttempted || !destination.cleanupAttempted { + t.Fatal("multipart header and cleanup were not both attempted") + } + }) + + t.Run("copy failure", func(t *testing.T) { + sourceErr := errors.New("source read failed") + closeErr := errors.New("close failed") + source := &terminalRecordingReader{reader: strings.NewReader("payload"), terminalErr: sourceErr} + destination := &multipartWriteRecorder{cleanupErr: closeErr, source: source} + writer := multipart.NewWriter(destination) + + err := writeMultipartFile( + writer, + source, + "payload.txt", + "text/plain", + ) + + if !errors.Is(err, sourceErr) { + t.Fatalf("error = %v, want source copy failure", err) + } + if !destination.cleanupAttempted { + t.Fatal("multipart cleanup was not attempted after the source read failure") + } + }) + + t.Run("close failure", func(t *testing.T) { + closeErr := errors.New("close failed") + source := &terminalRecordingReader{reader: strings.NewReader("payload")} + destination := &multipartWriteRecorder{cleanupErr: closeErr, source: source} + writer := multipart.NewWriter(destination) + + err := writeMultipartFile( + writer, + source, + "payload.txt", + "text/plain", + ) + + if !errors.Is(err, closeErr) || !strings.Contains(err.Error(), "close multipart writer") { + t.Fatalf("error = %v, want multipart Close failure", err) + } + if !destination.cleanupAttempted { + t.Fatal("multipart cleanup was not attempted after source EOF") + } + }) +} diff --git a/dify-agent-runtime/internal/envvar/envvar.go b/dify-agent-runtime/internal/envvar/envvar.go index e93af4e828c..40f9f9a97ad 100644 --- a/dify-agent-runtime/internal/envvar/envvar.go +++ b/dify-agent-runtime/internal/envvar/envvar.go @@ -23,7 +23,7 @@ const ( // --- Agent Stub --- const ( - // EnvAgentStubAPIBaseURL is the Agent Stub HTTP/gRPC endpoint. + // EnvAgentStubAPIBaseURL is the Agent Stub HTTP endpoint. EnvAgentStubAPIBaseURL = "DIFY_AGENT_STUB_API_BASE_URL" // EnvAgentStubAuthJWE is the per-request JWE token for Agent Stub auth. diff --git a/dify-agent-runtime/internal/server/tmux.go b/dify-agent-runtime/internal/server/tmux.go index 42a05c7de5f..dc5cfac9565 100644 --- a/dify-agent-runtime/internal/server/tmux.go +++ b/dify-agent-runtime/internal/server/tmux.go @@ -21,7 +21,11 @@ func NewTmuxController(config *Config) *TmuxController { // StartServer ensures the tmux server is running. func (t *TmuxController) StartServer() error { - _, err := t.runTmux("start-server") + _, err := t.runTmux( + "start-server", + ";", + "set-option", "-g", "exit-empty", "off", + ) return err } diff --git a/dify-agent-runtime/internal/server/tmux_test.go b/dify-agent-runtime/internal/server/tmux_test.go index 186b57dbfa7..68a8003a2c6 100644 --- a/dify-agent-runtime/internal/server/tmux_test.go +++ b/dify-agent-runtime/internal/server/tmux_test.go @@ -1,9 +1,91 @@ package server import ( + "os/exec" + "strings" "testing" ) +func newTestTmuxController(t *testing.T) *TmuxController { + t.Helper() + + if _, err := exec.LookPath("tmux"); err != nil { + t.Skip("tmux is not installed") + } + + config := DefaultConfig() + config.RuntimeDir = t.TempDir() + controller := NewTmuxController(config) + t.Cleanup(func() { + _, _ = controller.runTmuxNoCheck("kill-server") + }) + + return controller +} + +func assertTmuxServerRunning(t *testing.T, controller *TmuxController) { + t.Helper() + + exitEmpty, err := controller.runTmux("show-options", "-gv", "exit-empty") + if err != nil { + t.Fatalf("tmux server is not running: %v", err) + } + if got := strings.TrimSpace(exitEmpty); got != "off" { + t.Fatalf("exit-empty = %q, want %q", got, "off") + } +} + +func TestStartServerKeepsTmuxServerRunningWithoutSessions(t *testing.T) { + controller := newTestTmuxController(t) + + if err := controller.StartServer(); err != nil { + t.Fatalf("StartServer() error = %v", err) + } + + assertTmuxServerRunning(t, controller) + + sessions, err := controller.ListSessions() + if err != nil { + t.Fatalf("ListSessions() error = %v", err) + } + if len(sessions) != 0 { + t.Fatalf("ListSessions() = %v, want no sessions", sessions) + } +} + +func TestTmuxServerRemainsRunningAfterLastSessionIsDeleted(t *testing.T) { + controller := newTestTmuxController(t) + + if err := controller.StartServer(); err != nil { + t.Fatalf("StartServer() error = %v", err) + } + + const jobID = "last-session" + sessionName := JobSessionName(jobID) + if _, err := controller.runTmux("new-session", "-d", "-s", sessionName); err != nil { + t.Fatalf("create tmux session: %v", err) + } + + sessions, err := controller.ListSessions() + if err != nil { + t.Fatalf("ListSessions() before cleanup error = %v", err) + } + if !sessions[sessionName] { + t.Fatalf("ListSessions() = %v, want session %q", sessions, sessionName) + } + + controller.CleanupSession(jobID) + assertTmuxServerRunning(t, controller) + + sessions, err = controller.ListSessions() + if err != nil { + t.Fatalf("ListSessions() after cleanup error = %v", err) + } + if len(sessions) != 0 { + t.Fatalf("ListSessions() after cleanup = %v, want no sessions", sessions) + } +} + func TestShellQuote(t *testing.T) { tests := []struct { input string diff --git a/dify-agent-runtime/internal/stubclient/client.go b/dify-agent-runtime/internal/stubclient/client.go deleted file mode 100644 index d9b6bdddc8a..00000000000 --- a/dify-agent-runtime/internal/stubclient/client.go +++ /dev/null @@ -1,143 +0,0 @@ -// Package stubclient provides a gRPC client for the Dify Agent Stub service. -// -// The Agent Stub service runs on the host (dify-agent backend) and exposes -// Connect, FileUpload, and FileDownload RPCs to sandbox-resident processes. -// This client is used by the dify-agent CLI binary inside the sandbox container. -package stubclient - -import ( - "context" - "fmt" - "time" - - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" - - stubv1 "github.com/langgenius/dify/dify-agent-runtime/gen/dify/agent/stub/v1" -) - -// Client wraps the gRPC AgentStubService client with connection lifecycle management. -type Client struct { - conn *grpc.ClientConn - stub stubv1.AgentStubServiceClient - timeout time.Duration -} - -// Option configures a Client. -type Option func(*Client) - -// WithTimeout sets the default per-call timeout. Default is 30s. -func WithTimeout(d time.Duration) Option { - return func(c *Client) { - c.timeout = d - } -} - -// Dial creates a new Client connected to the given target address (host:port). -func Dial(target string, opts ...Option) (*Client, error) { - c := &Client{timeout: 30 * time.Second} - for _, o := range opts { - o(c) - } - - conn, err := grpc.NewClient(target, grpc.WithTransportCredentials(insecure.NewCredentials())) - if err != nil { - return nil, fmt.Errorf("stubclient: dial %s: %w", target, err) - } - c.conn = conn - c.stub = stubv1.NewAgentStubServiceClient(conn) - return c, nil -} - -// Close closes the underlying gRPC connection. -func (c *Client) Close() error { - if c.conn != nil { - return c.conn.Close() - } - return nil -} - -// ConnectResult holds the response from a Connect RPC. -type ConnectResult struct { - ConnectionID string - Status string -} - -// Connect establishes a logical connection with the Agent Stub server. -func (c *Client) Connect(ctx context.Context, protocolVersion int32, argv []string, metadataJSON string) (*ConnectResult, error) { - ctx, cancel := context.WithTimeout(ctx, c.timeout) - defer cancel() - - resp, err := c.stub.Connect(ctx, &stubv1.ConnectRequest{ - ProtocolVersion: protocolVersion, - Argv: argv, - MetadataJson: metadataJSON, - }) - if err != nil { - return nil, fmt.Errorf("stubclient: connect: %w", err) - } - return &ConnectResult{ - ConnectionID: resp.GetConnectionId(), - Status: resp.GetStatus(), - }, nil -} - -// FileUploadResult holds the response from a CreateFileUploadRequest RPC. -type FileUploadResult struct { - UploadURL string -} - -// CreateFileUpload requests an upload URL from the Agent Stub server. -func (c *Client) CreateFileUpload(ctx context.Context, filename, mimetype string) (*FileUploadResult, error) { - ctx, cancel := context.WithTimeout(ctx, c.timeout) - defer cancel() - - resp, err := c.stub.CreateFileUploadRequest(ctx, &stubv1.FileUploadRequest{ - Filename: filename, - Mimetype: mimetype, - }) - if err != nil { - return nil, fmt.Errorf("stubclient: file upload: %w", err) - } - return &FileUploadResult{ - UploadURL: resp.GetUploadUrl(), - }, nil -} - -// FileDownloadResult holds the response from a CreateFileDownloadRequest RPC. -type FileDownloadResult struct { - Filename string - MimeType string - Size int64 - DownloadURL string -} - -// CreateFileDownload requests a download URL from the Agent Stub server. -func (c *Client) CreateFileDownload(ctx context.Context, transferMethod string, reference, url *string, forExternal bool) (*FileDownloadResult, error) { - ctx, cancel := context.WithTimeout(ctx, c.timeout) - defer cancel() - - fileMapping := &stubv1.FileMapping{ - TransferMethod: transferMethod, - } - if reference != nil { - fileMapping.Reference = reference - } - if url != nil { - fileMapping.Url = url - } - - resp, err := c.stub.CreateFileDownloadRequest(ctx, &stubv1.FileDownloadRequest{ - File: fileMapping, - ForExternal: &forExternal, - }) - if err != nil { - return nil, fmt.Errorf("stubclient: file download: %w", err) - } - return &FileDownloadResult{ - Filename: resp.GetFilename(), - MimeType: resp.GetMimeType(), - Size: resp.GetSize(), - DownloadURL: resp.GetDownloadUrl(), - }, nil -} diff --git a/dify-agent-runtime/internal/stubclient/client_test.go b/dify-agent-runtime/internal/stubclient/client_test.go deleted file mode 100644 index 13499b67499..00000000000 --- a/dify-agent-runtime/internal/stubclient/client_test.go +++ /dev/null @@ -1,114 +0,0 @@ -package stubclient - -import ( - "context" - "net" - "testing" - - "google.golang.org/grpc" - - stubv1 "github.com/langgenius/dify/dify-agent-runtime/gen/dify/agent/stub/v1" -) - -// fakeServer implements the AgentStubService for testing. -type fakeServer struct { - stubv1.UnimplementedAgentStubServiceServer -} - -func (f *fakeServer) Connect(_ context.Context, req *stubv1.ConnectRequest) (*stubv1.ConnectResponse, error) { - return &stubv1.ConnectResponse{ - ConnectionId: "conn-123", - Status: "connected", - }, nil -} - -func (f *fakeServer) CreateFileUploadRequest(_ context.Context, req *stubv1.FileUploadRequest) (*stubv1.FileUploadResponse, error) { - return &stubv1.FileUploadResponse{ - UploadUrl: "https://upload.example.com/" + req.GetFilename(), - }, nil -} - -func (f *fakeServer) CreateFileDownloadRequest(_ context.Context, req *stubv1.FileDownloadRequest) (*stubv1.FileDownloadResponse, error) { - return &stubv1.FileDownloadResponse{ - Filename: "downloaded.txt", - MimeType: strPtr("text/plain"), - Size: 1024, - DownloadUrl: "https://download.example.com/file", - }, nil -} - -func strPtr(s string) *string { return &s } - -func startFakeServer(t *testing.T) string { - t.Helper() - lis, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatalf("failed to listen: %v", err) - } - srv := grpc.NewServer() - stubv1.RegisterAgentStubServiceServer(srv, &fakeServer{}) - go func() { _ = srv.Serve(lis) }() - t.Cleanup(srv.Stop) - return lis.Addr().String() -} - -func TestConnect(t *testing.T) { - addr := startFakeServer(t) - c, err := Dial(addr) - if err != nil { - t.Fatalf("Dial failed: %v", err) - } - defer func() { _ = c.Close() }() - - result, err := c.Connect(context.Background(), 1, []string{"hello"}, `{"key":"val"}`) - if err != nil { - t.Fatalf("Connect failed: %v", err) - } - if result.ConnectionID != "conn-123" { - t.Errorf("expected conn-123, got %s", result.ConnectionID) - } - if result.Status != "connected" { - t.Errorf("expected connected, got %s", result.Status) - } -} - -func TestCreateFileUpload(t *testing.T) { - addr := startFakeServer(t) - c, err := Dial(addr) - if err != nil { - t.Fatalf("Dial failed: %v", err) - } - defer func() { _ = c.Close() }() - - result, err := c.CreateFileUpload(context.Background(), "test.txt", "text/plain") - if err != nil { - t.Fatalf("CreateFileUpload failed: %v", err) - } - if result.UploadURL != "https://upload.example.com/test.txt" { - t.Errorf("unexpected upload URL: %s", result.UploadURL) - } -} - -func TestCreateFileDownload(t *testing.T) { - addr := startFakeServer(t) - c, err := Dial(addr) - if err != nil { - t.Fatalf("Dial failed: %v", err) - } - defer func() { _ = c.Close() }() - - ref := "dify-file-ref:abc123" - result, err := c.CreateFileDownload(context.Background(), "local_file", &ref, nil, false) - if err != nil { - t.Fatalf("CreateFileDownload failed: %v", err) - } - if result.Filename != "downloaded.txt" { - t.Errorf("unexpected filename: %s", result.Filename) - } - if result.Size != 1024 { - t.Errorf("unexpected size: %d", result.Size) - } - if result.DownloadURL != "https://download.example.com/file" { - t.Errorf("unexpected download URL: %s", result.DownloadURL) - } -} diff --git a/dify-agent/.example.env b/dify-agent/.example.env index 26f25293efd..31220ebb52c 100644 --- a/dify-agent/.example.env +++ b/dify-agent/.example.env @@ -4,7 +4,7 @@ # Redis # Redis connection URL for run records and per-run event streams. -DIFY_AGENT_REDIS_URL=redis://localhost:6379/0 +DIFY_AGENT_REDIS_URL=redis://:difyai123456localhost:6379/0 # Prefix for Redis run-record and event-stream keys. DIFY_AGENT_REDIS_PREFIX=dify-agent @@ -18,7 +18,7 @@ DIFY_AGENT_RUN_RETENTION_SECONDS=259200 # Base URL for the Dify plugin daemon used by local runs. DIFY_AGENT_PLUGIN_DAEMON_URL=http://localhost:5002 # API key sent to the Dify plugin daemon. -DIFY_AGENT_PLUGIN_DAEMON_API_KEY= +DIFY_AGENT_PLUGIN_DAEMON_API_KEY=lYkiYYT6owG+71oLerGzA7GXCgOT++6ovaezWAjpCjf+Sjc3ZtU+qUEi # Dify API inner endpoints # Base URL for Dify API inner endpoints used by Agent Stub config/file/drive requests. @@ -26,48 +26,59 @@ DIFY_AGENT_INNER_API_URL=http://localhost:5001 # Must match API/worker INNER_API_KEY_FOR_PLUGIN, not the generic INNER_API_KEY. DIFY_AGENT_INNER_API_KEY= -# Shell layer -# Shell backend used by the dify.shell layer: "shellctl" (default) or "enterprise". -DIFY_AGENT_SHELL_PROVIDER=shellctl -# shellctl provider: base URL for the shellctl server. Leave empty to disable shell layer use. -DIFY_AGENT_SHELLCTL_ENTRYPOINT= -# Optional bearer token sent to the shellctl server. -DIFY_AGENT_SHELLCTL_AUTH_TOKEN= -# Root directory for per-Agent shell HOME directories. Use a writable path such -# as /tmp/dify-agent-home for local macOS development. -DIFY_AGENT_SHELL_HOME_ROOT=/home -# enterprise provider: sandbox gateway endpoint. Leave empty to disable shell layer use. +# Runtime resources +# Select one coherent Home Snapshot + Execution Binding backend: local, enterprise, or e2b. +DIFY_AGENT_RUNTIME_BACKEND=local +# Local backend: shellctl data-plane URL and optional bearer token. +# Leave the endpoint empty when this server will not provide dify.runtime or resource endpoints. +DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT= +DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN= +# Enterprise resource operations currently fail fast with NotImplementedError. +# These names are retained for the configured Enterprise Gateway boundary. DIFY_AGENT_ENTERPRISE_SANDBOX_GATEWAY_ENDPOINT= -# Optional X-Inner-Api-Key sent to the enterprise sandbox gateway. DIFY_AGENT_ENTERPRISE_SANDBOX_GATEWAY_AUTH_TOKEN= DIFY_AGENT_ENTERPRISE_SANDBOX_GATEWAY_TIMEOUT=30 DIFY_AGENT_ENTERPRISE_SANDBOX_PROXY_TIMEOUT=60 +# E2B backend: API key, prepared shellctl template, and active Binding policy. +DIFY_AGENT_E2B_API_KEY= +DIFY_AGENT_E2B_TEMPLATE=difys-default-team/dify-agent-local-sandbox +# Maximum continuous active time for the RuntimeLease that spans one complete Agent run. +# Binding resources pause; temporary Home initialization resources are killed. +# This is not a retention TTL for paused resources or immutable snapshots. +DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS=3600 +DIFY_AGENT_E2B_SHELLCTL_AUTH_TOKEN= +DIFY_AGENT_E2B_SHELLCTL_PORT=5004 +# JSON array of regex patterns to redact from shell output shown to the agent. +DIFY_AGENT_SHELL_REDACT_PATTERNS= # Agent Stub # Public Agent Stub URL reachable from shellctl-managed remote machines. -# Use http(s)://.../agent-stub for HTTP or grpc://host:port for gRPC. +# Use an HTTP(S) service root or an explicit /agent-stub API root. # Leave empty to avoid injecting DIFY_AGENT_STUB_* into shell.run jobs. DIFY_AGENT_STUB_API_BASE_URL=http://localhost:5050/agent-stub -# Optional bind override used only when DIFY_AGENT_STUB_API_BASE_URL uses grpc://. -DIFY_AGENT_STUB_GRPC_BIND_ADDRESS= +# Dify API base URL reachable from the Sandbox for the signed /files/* data plane, +# including Config file and skill pulls. +DIFY_AGENT_SANDBOX_FILES_BASE_URL=http://localhost:5001 # Server-wide root secret used to derive Agent Stub JWE keys. # This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens. # Replace this development default in production. # Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))' DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY -# Shared plugin-daemon HTTP client timeouts and limits. -# Plugin-daemon HTTP connect timeout in seconds. -DIFY_AGENT_PLUGIN_DAEMON_CONNECT_TIMEOUT=10 -# Plugin-daemon HTTP read timeout in seconds. -DIFY_AGENT_PLUGIN_DAEMON_READ_TIMEOUT=600 -# Plugin-daemon HTTP write timeout in seconds. -DIFY_AGENT_PLUGIN_DAEMON_WRITE_TIMEOUT=30 -# Plugin-daemon HTTP connection-pool wait timeout in seconds. -DIFY_AGENT_PLUGIN_DAEMON_POOL_TIMEOUT=10 -# Maximum total plugin-daemon HTTP connections. -DIFY_AGENT_PLUGIN_DAEMON_MAX_CONNECTIONS=100 -# Maximum idle keep-alive plugin-daemon HTTP connections. -DIFY_AGENT_PLUGIN_DAEMON_MAX_KEEPALIVE_CONNECTIONS=20 -# Keep-alive expiry in seconds for idle plugin-daemon HTTP connections. -DIFY_AGENT_PLUGIN_DAEMON_KEEPALIVE_EXPIRY=30 +# Inbound Bearer token for /runs API authentication. +# Must match AGENT_BACKEND_API_TOKEN on the Dify API side. +# Replace this development default in production. +# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))' +DIFY_AGENT_API_TOKEN=dify-agent-run-token-for-dev-only + +# Shared outbound HTTP client timeouts and limits. +DIFY_AGENT_OUTBOUND_HTTP_CONNECT_TIMEOUT=10 +DIFY_AGENT_OUTBOUND_HTTP_READ_TIMEOUT=600 +DIFY_AGENT_OUTBOUND_HTTP_WRITE_TIMEOUT=30 +DIFY_AGENT_OUTBOUND_HTTP_POOL_TIMEOUT=10 +DIFY_AGENT_OUTBOUND_HTTP_MAX_CONNECTIONS=100 +DIFY_AGENT_OUTBOUND_HTTP_MAX_KEEPALIVE_CONNECTIONS=20 +DIFY_AGENT_OUTBOUND_HTTP_KEEPALIVE_EXPIRY=30 +DIFY_AGENT_LOCAL_SANDBOX_MATERIALIZED_HOME_ROOT= +DIFY_AGENT_LOCAL_SANDBOX_WORKSPACE_ROOT= +DIFY_AGENT_LOCAL_SANDBOX_HOME_SNAPSHOT_ROOT= diff --git a/dify-agent/.gitignore b/dify-agent/.gitignore index f644f2e5451..5f6b9ad4495 100644 --- a/dify-agent/.gitignore +++ b/dify-agent/.gitignore @@ -1 +1,2 @@ dify-aio +site/ diff --git a/dify-agent/AGENTS.md b/dify-agent/AGENTS.md index 43c68448f20..ee33943a3aa 100644 --- a/dify-agent/AGENTS.md +++ b/dify-agent/AGENTS.md @@ -1,184 +1,21 @@ # Agent Guide -## Notes for Agent (must-check) +Read surrounding docstrings and non-obvious comments before changing behavior. They are local contracts; update them only when their owned behavior changes, and keep them aligned with the current code. Read `docs/dify-agent/index.md` when changing the public runtime contract. -Before changing any source code under this folder, you MUST read the surrounding docstrings and comments. These notes contain required context (invariants, edge cases, trade-offs) and are treated as part of the spec. +## Commands -Look for: +Run package commands from `dify-agent/`: -- The module (file) docstring at the top of a source code file -- Docstrings on classes and functions/methods -- Paragraph/block comments for non-obvious logic +- Lint: `make check` +- Format and fix lint: `make fix` +- Type check: `make typecheck` +- Tests: `make test` -### What to write where +Use the package's `uv` environment and Pydantic v2 APIs. Inspect current dependency source or official documentation before integrating, implementing, or mocking an API whose runtime contract is not established locally. -- Keep notes scoped: module notes cover module-wide context, class notes cover class-wide context, function/method notes cover behavioural contracts, and paragraph/block comments cover local “why”. Avoid duplicating the same content across scopes unless repetition prevents misuse. -- **Module (file) docstring**: purpose, boundaries, key invariants, and “gotchas” that a new reader must know before editing. - - Include cross-links to the key collaborators (modules/services) when discovery is otherwise hard. - - Prefer stable facts (invariants, contracts) over ephemeral “today we…” notes. -- **Class docstring**: responsibility, lifecycle, invariants, and how it should be used (or not used). - - If the class is intentionally stateful, note what state exists and what methods mutate it. - - If concurrency/async assumptions matter, state them explicitly. -- **Function/method docstring**: behavioural contract. - - Document arguments, return shape, side effects (DB writes, external I/O, task dispatch), and raised domain exceptions. - - Add examples only when they prevent misuse. -- **Paragraph/block comments**: explain *why* (trade-offs, historical constraints, surprising edge cases), not what the code already states. - - Keep comments adjacent to the logic they justify; delete or rewrite comments that no longer match reality. +## Tests And Boundaries -### Rules (must follow) - -In this section, “notes” means module/class/function docstrings plus any relevant paragraph/block comments. - -- **Before working** - - Read the notes in the area you’ll touch; treat them as part of the spec. - - If a docstring or comment conflicts with the current code, treat the **code as the single source of truth** and update the docstring or comment to match reality. - - If important intent/invariants/edge cases are missing, add them in the closest docstring or comment (module for overall scope, function for behaviour). -- **During working** - - Keep the notes in sync as you discover constraints, make decisions, or change approach. - - If you move/rename responsibilities across modules/classes, update the affected docstrings and comments so readers can still find the “why” and the invariants. - - Record non-obvious edge cases, trade-offs, and the test/verification plan in the nearest docstring or comment that will stay correct. - - Keep the notes **coherent**: integrate new findings into the relevant docstrings and comments; avoid append-only “recent fix” / changelog-style additions. -- **When finishing** - - Update the notes to reflect what changed, why, and any new edge cases/tests. - - Remove or rewrite any comments that could be mistaken as current guidance but no longer apply. - - Keep docstrings and comments concise and accurate; they are meant to prevent repeated rediscovery. - -## Coding Style - -This is the default standard for backend code in this repo. Follow it for new code and use it as the checklist when reviewing changes. - -### Linting & Formatting - -- Use Ruff for formatting and linting (follow `.ruff.toml`). -- Keep each line under 120 characters (including spaces). - -### Naming Conventions - -- Use `snake_case` for variables and functions. -- Use `PascalCase` for classes. -- Use `UPPER_CASE` for constants. - -### Typing & Class Layout - -- Code should usually include type annotations that match the repo’s current Python version (avoid untyped public APIs and “mystery” values). -- Prefer modern typing forms (e.g. `list[str]`, `dict[str, int]`) and avoid `Any` unless there’s a strong reason. -- For dictionary-like data with known keys and value types, prefer `TypedDict` over `dict[...]` or `Mapping[...]`. -- For optional keys in typed payloads, use `NotRequired[...]` (or `total=False` when most fields are optional). -- Keep `dict[...]` / `Mapping[...]` for truly dynamic key spaces where the key set is unknown. - -```python -from datetime import datetime -from typing import NotRequired, TypedDict - - -class UserProfile(TypedDict): - user_id: str - email: str - created_at: datetime - nickname: NotRequired[str] -``` - -- For classes, declare all member variables explicitly with types at the top of the class body (before `__init__`), even when the class is not a dataclass or Pydantic model, so the class shape is obvious at a glance: - -```python -from datetime import datetime - - -class Example: - user_id: str - created_at: datetime - - def __init__(self, user_id: str, created_at: datetime) -> None: - self.user_id = user_id - self.created_at = created_at -``` - -- For dataclasses, prefer `field(default_factory=...)` over `field(init=False)` when a default can be provided declaratively. -- Prefer dataclasses with `slots=True` when defining lightweight data containers: - -```python -from dataclasses import dataclass -from datetime import datetime - - -@dataclass(slots=True) -class Example: - user_id: str - created_at: datetime -``` - -### General Rules - -- Use Pydantic v2 conventions. -- Use `uv` for Python package management in this repo (usually with `--project dify-agent`). -- Use `make typecheck` to run `basedpyright` against `dify-agent/src` and `dify-agent/tests`. -- Keep type checking passing after every edit you make. -- Use `pytest` for all tests in this package. -- When integrating with, implementing, or mocking a dependency, inspect the dependency's source code to confirm its API shape and runtime behavior instead of guessing from names alone. -- Prefer simple functions over small “utility classes” for lightweight helpers. -- Avoid implementing dunder methods unless it’s clearly needed and matches existing patterns. -- Keep code readable and explicit—avoid clever hacks. - -### Testing - -- Work in TDD style: write or update a failing test first when changing behavior, then make the implementation pass, then refactor while keeping tests and typecheck green. -- Use `make test` to run the agent pytest suite. -- Keep local tests under `dify-agent/tests/local/`. -- Mirror the `dify-agent/src/` package structure inside `dify-agent/tests/local/` so test locations stay predictable. - -#### Local Tests - -- Write local tests for stable, externally observable behavior that can run quickly without real external services. -- In this repo, code, comments, docs, and tests are expected to change together. Because of that, a local test is only useful if it would still be correct after an internal refactor that does not change the intended contract. -- Local tests should verify: - - what callers and downstream code can observe and rely on - - how the unit is expected to use its dependencies at the boundary - - how the unit handles dependency success, failure, empty responses, malformed responses, and documented error cases - - documented invariants, error mapping, and output/input shape guarantees -- When asserting dependency interactions, assert only the parts of the request or response that are part of the real boundary contract. Do not over-specify incidental details that callers or dependencies do not rely on. -- It is acceptable to mock dependencies in local tests, but only when the mock represents a real contract, schema, documented behavior, or known regression. -- Tests may use line-scoped type-ignore comments when intentionally exercising runtime validation paths that static typing would normally reject. Keep the ignore on the exact invalid call. -- Do not use local tests to prove real integration, network wiring, serialization, framework configuration, or third-party runtime behavior; cover those in higher-level tests. -- Meaningless local tests include: - - tests that only mirror the current implementation or must be updated whenever internal code changes even though the contract did not change - - tests of private helpers, local variables, temporary state, internal branching, or exact internal call order unless those details are part of the published contract - - tests with mocked dependency behavior that is invented only to make the current implementation pass - - tests that add no value beyond static type checking or linting - -### Logging & Errors - -- Never use `print`; use a module-level logger: - - `logger = logging.getLogger(__name__)` -- Include tenant/app/workflow identifiers in log context when relevant. -- Raise domain-specific exceptions and translate them into HTTP responses in controllers. -- Log retryable events at `warning`, terminal failures at `error`. - -### Pydantic Usage - -- Define DTOs with Pydantic v2 models and forbid extras by default. -- Use `@field_validator` / `@model_validator` for domain rules. - -Example: - -```python -from pydantic import BaseModel, ConfigDict, HttpUrl, field_validator - - -class TriggerConfig(BaseModel): - endpoint: HttpUrl - secret: str - - model_config = ConfigDict(extra="forbid") - - @field_validator("secret") - def ensure_secret_prefix(cls, value: str) -> str: - if not value.startswith("dify_"): - raise ValueError("secret must start with dify_") - return value -``` - -### Generics & Protocols - -- Use `typing.Protocol` to define behavioural contracts (e.g., cache interfaces). -- Apply generics (`TypeVar`, `Generic`) for reusable utilities like caches or providers. -- Validate dynamic inputs at runtime when generics cannot enforce safety alone. +- Keep local tests under `tests/local/` and mirror the `src/` package structure. +- Test stable behavior and real dependency boundaries. Do not use local mocks to claim real network, framework wiring, serialization, or third-party runtime coverage. +- Keep tests, public docs, and local contracts aligned with behavior changes. +- Preserve the existing runtime and layer owners; do not add generic utilities or compatibility boundaries to bypass them. diff --git a/dify-agent/docs/dify-agent/concepts/runtime-resources/index.md b/dify-agent/docs/dify-agent/concepts/runtime-resources/index.md new file mode 100644 index 00000000000..0745f1f9fce --- /dev/null +++ b/dify-agent/docs/dify-agent/concepts/runtime-resources/index.md @@ -0,0 +1,201 @@ +# Runtime resources + +Dify separates persistent product resources from request-time execution: + +- a **Home Snapshot** is immutable Agent-owned Home content; +- a **Workspace** is mutable working data owned by a product scope such as a + conversation, Build Draft, or Workflow run; +- an **Execution Binding** is one materialized Agent participant, including its + private Home and resumable session, attached to a Workspace; +- a **RuntimeLease** is operation-scoped access to the physical Binding. + +`AgentWorkspaceBinding.id` is the participant, materialized Home, and persisted +Agenton session identity. `agent_id` identifies the source Agent. The same +Agent can therefore have multiple active Bindings in one Workspace: +each has an independent Home and session, while all may share Workspace files. + +Home and Workspace are logically independent. A backend may still couple their +physical representation. For example, current E2B maps one Binding and its +Workspace to one E2B resource, while Local can attach multiple materialized +Homes to one shared Workspace. + +## Runtime layer graph + +Agent requests do not expose separate Home, Workspace, or Sandbox layers. Dify +API resolves the Binding selected by the product flow and sends its opaque +backend ref to the `dify.runtime` layer: + +```mermaid +flowchart LR + EC["dify.execution_context
request identity"] + RT["dify.runtime
opaque backend_binding_ref"] + SH["dify.shell
commands and jobs"] + + EC --> SH + RT --> SH +``` + +`DifyRuntimeLayer` calls the selected `ExecutionBindingBackend.acquire()` when +its resource context opens and `release()` when the operation ends. It exposes +the resulting `RuntimeLease` only while that context is active. The layer does +not create, retire, or destroy persistent resources, and it stores no backend +SDK object in an Agenton session snapshot. + +The Shell layer consumes `RuntimeLease.commands` and `RuntimeLease.layout`. It +tracks only request-local shell job ids and offsets. +Closing a run clears that job state; it does not retire the Binding. + +## State ownership + +Dify API is the lifecycle ledger. It stores three resource records: + +| Record | Meaning | Backend field | +| --- | --- | --- | +| `agent_home_snapshots` | One immutable Home version owned by an Agent. | `snapshot_ref` | +| `agent_workspaces` | One mutable Workspace owned by a product scope. | `backend_workspace_ref` | +| `agent_workspace_bindings` | One materialized participant, private Home, and resumable session attached to a Workspace. | `backend_binding_ref` | + +Backend refs are opaque strings interpreted only by the selected backend +adapter. Dify API stores the latest Agenton session snapshot on the Binding, but +it does not serialize `RuntimeLease`, SDK clients, credentials, or temporary +access tokens. + +Dify Agent does not connect to the Dify product database and has no persistent +resource registry. Its private control-plane endpoints create or destroy +backend resources from requests made by Dify API. Redis run records and event +streams are observability state, not the Home/Workspace/Binding ledger. + +## Creation and execution flow + +Agent creation does not create a Home Snapshot. A config with no logical Home +Snapshot asks the selected backend to materialize its deployment-default Home +when the Binding is created. This default Home is mutable and private to the +Binding; it does not produce an `agent_home_snapshots` row or an implicit +snapshot ref. + +Build Draft Apply uses `POST /home-snapshots/from-binding`: Dify Agent acquires +the exact source Binding, snapshots its materialized Home through the +backend-native operation, releases the lease, and returns a new opaque snapshot +ref. Dify API then stores a new immutable `agent_home_snapshots` row and records +its logical id on the resulting config version. There is no replay or fallback +when the source Binding is unavailable. + +Before an Agent request, Dify API loads the specific product context. If it has +no associated Binding, Dify API materializes one and saves the Binding id in the +same database transaction. Otherwise it resolves only that Binding and validates +its owner and config/Home generation. Missing, retired, or mismatched Bindings +fail fast; Dify API does not search by Agent, Workspace, candidate count, or +recency, and it does not create a replacement implicitly. + +`POST /execution-bindings` accepts either an exact `home_snapshot_ref` or +`null`. An exact ref must be materialized without fallback; `null` selects the +backend's deployment-default Home. It returns opaque Binding and Workspace +refs. Every create request represents a new participant, even when the Agent, +Snapshot, config generation, and Workspace match another Binding. The request +composition contains: + +```json +{ + "name": "runtime", + "type": "dify.runtime", + "config": {"backend_binding_ref": "opaque-backend-binding-ref"} +} +``` + +Each Agent request acquires that ref for the duration of the run and releases it +afterward. Local release closes the operation's shellctl connection. E2B release +also pauses the underlying E2B resource with memory preserved. A later request +or Binding file operation acquires a new lease for the same Binding ref. If a +backend confirms the resource is gone, acquisition fails; it does not create an +empty replacement Workspace. + +## Retirement and collection + +Retirement is a database transition from `ACTIVE` to `RETIRED`. It prevents new +product use without performing network I/O inside the caller's transaction. +Product lifecycle paths commit this transition synchronously. After the +transaction commits, one Celery task asks Dify Agent to destroy the physical +resources. A successful collector deletes the corresponding ledger row; a +failed collector logs the failure and leaves the RETIRED row available for a +future retry or reconciler. + +The unified `collect_agent_resources` task is registered on normal Celery +workers and explicitly uses the existing `retention` queue. Standard workers +already consume that queue, so no dedicated Agent resource worker or new queue +is required. At a Workflow terminal event, the graph layer synchronously retires +and commits the run's Workspaces before enqueueing collection. When a Workflow +change may orphan Workflow-only Agents, the main product transaction commits +first; a fresh session then rechecks effective ownership and retires only Agents +that remain unowned. + +Retiring a final Binding also retires its Workspace. Workspace collection +destroys the physical Workspace through one Binding and then collects remaining +materialized Homes. Home Snapshots are retired when their owning Agent is +retired and are collected only after no draft or config snapshot references +them. Celery performs physical collection only; it does not decide or perform +the initial retirement. Dify Agent itself remains stateless. + +There is currently no age-based TTL, periodic GC, or global orphan reconciler. +Backend destroy operations are idempotent where supported. Dify API does not +perform cross-system compensation after a backend create returns success. Any +later API failure, including Python, flush, or commit failure, may leave a +physical orphan for a future global reconciler. + +Backends still clean up partial resources when a create operation fails before +returning success. For example, E2B kills a Sandbox when its initialization +fails, and Local removes paths created by an incomplete operation. This +backend-local cleanup does not cross the database commit boundary. + +## Binding file boundary + +Dify API's public file APIs accept a product locator, not a Binding id or +backend ref: a Conversation, a debug Build Draft, or a Workflow Node Execution. +Dify API authorizes that object and resolves its associated active Binding. It +does not select the latest Binding or fall back to another product context. + +The resolved request reaches Dify Agent through its private +`POST /execution-bindings/files/list`, `POST /execution-bindings/files/read`, +and `POST /execution-bindings/files/download` endpoints. Each operation receives a +`backend_binding_ref`, acquires a fresh RuntimeLease, performs the file action, +and releases the lease. + +`BindingFileService` resolves relative paths from `workspace_dir`, `~` and +`~/...` from `home_dir`, and leaves absolute paths in the Binding filesystem +namespace. It does not enforce Workspace containment or reject `..`; the +selected backend's isolation policy remains authoritative. List and preview +run bounded inspection scripts through `RuntimeLease.commands`. Download runs +`dify-agent file upload --no-download-link` inside the Binding so bytes stream +directly from the runtime to Dify's existing ToolFile endpoint. Dify Agent +returns only the canonical ToolFile reference and releases the lease before +Dify API signs a browser URL. + +`RuntimeLayout.home_dir` and `RuntimeLayout.workspace_dir` are canonical paths +inside the backend execution namespace. They are not host paths, product ids, +or request configuration. Shell commands start in `workspace_dir`, and `HOME` +is forced to `home_dir`. On Local, sibling materialized Homes may exist in the +same shellctl namespace, while path isolation restricts the active lease to its +own Home plus the shared Workspace. + +## Backend support + +| Backend | Home Snapshot operations | Binding operations | Physical relationship | +| --- | --- | --- | --- | +| Local | Supported | Supported, including default empty Homes and attaching multiple Bindings to one Workspace | Snapshot directory, per-Binding materialized Home, and Workspace directory are separate. | +| E2B | Supported | Supported with template-backed default Homes, without shared-Workspace attachment | Binding and Workspace refs map to the same E2B resource; checkpoints use E2B snapshots. | +| Enterprise | Not implemented | Default-Home Binding creation, acquire, and coupled destroy are supported | Binding and Workspace refs map to one Gateway sandbox. Explicit Home Snapshot materialization fails fast. | + +Local creates a new Home for every Binding id. Destroying one Binding without +the Workspace leaves sibling Homes and the shared Workspace intact. Current E2B +rejects `existing_workspace_ref` with `shared_workspace_unsupported`, because +its Binding and Workspace are one Sandbox. It also rejects binding-only destroy. +Neither path creates a fallback Workspace or switches backends. + +`DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS` limits continuous active time for an E2B +resource to one hour. The limit covers the complete Agent run held by one +RuntimeLease rather than an individual tool call. Runtime resources pause on +timeout. It is not a retention TTL and does not delete paused resources or +immutable snapshots. + +See the [Shell layer](../../user-manual/shell-layer/index.md) for request +composition and the [Operations Guide](../../guide/index.md) for Local and E2B +validation. diff --git a/dify-agent/docs/dify-agent/get-started/index.md b/dify-agent/docs/dify-agent/get-started/index.md index 9c0056d772c..71990baad5a 100644 --- a/dify-agent/docs/dify-agent/get-started/index.md +++ b/dify-agent/docs/dify-agent/get-started/index.md @@ -73,11 +73,40 @@ The minimum settings are: See `.example.env` for the full server settings template. -If you plan to run `dify.shell`, also configure `DIFY_AGENT_SHELLCTL_ENTRYPOINT` -and, when shell jobs need to call back with the `dify-agent` command, set -`DIFY_AGENT_STUB_API_BASE_URL`. The supplied default configs include a +If you plan to run `dify.shell`, select a coherent Home Snapshot and Execution +Binding backend. A standalone Local server normally uses: + +```env +DIFY_AGENT_RUNTIME_BACKEND=local +DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT=http://127.0.0.1:5004 +DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN= +# Optional when shellctl runs directly on a host without /home/dify: +# DIFY_AGENT_LOCAL_SANDBOX_MATERIALIZED_HOME_ROOT=/tmp/dify-agent/materialized-homes +# DIFY_AGENT_LOCAL_SANDBOX_WORKSPACE_ROOT=/tmp/dify-agent/workspaces +# DIFY_AGENT_LOCAL_SANDBOX_HOME_SNAPSHOT_ROOT=/tmp/dify-agent/home-snapshots +``` + +E2B requires `DIFY_AGENT_E2B_API_KEY` and defaults to the prepared +`difys-default-team/dify-agent-local-sandbox` template. The E2B active timeout +pauses the physical resource behind a Binding; it is not a retention TTL. +Enterprise supports Bindings created from its deployment-default Home. +Immutable Home Snapshot creation and materialization remain unsupported there. + +A shell-enabled request includes Execution Context, `dify.runtime`, and +`dify.shell`. Dify API creates or resolves the specific persistent Binding for +the request's product context; +`DifyRuntimeLayerConfig.backend_binding_ref` carries only that opaque ref and +opens a new operation-scoped `RuntimeLease` for the run. When shell jobs need to +call back with the `dify-agent` command, also set +`DIFY_AGENT_STUB_API_BASE_URL` and the Sandbox-reachable Dify API base +`DIFY_AGENT_SANDBOX_FILES_BASE_URL`. The supplied default configs include a development `DIFY_AGENT_SERVER_SECRET_KEY`, but production deployments should -override it with a unique 32-byte base64url value as documented in `.example.env`. +override it with a unique 32-byte base64url value as documented in +`.example.env`. + +See [Runtime resources](../concepts/runtime-resources/index.md) for the layer +graph and state ownership, and the [Operations Guide](../guide/index.md) for +backend-specific configuration and validation commands. ## Start the Dify Agent server diff --git a/dify-agent/docs/dify-agent/guide/index.md b/dify-agent/docs/dify-agent/guide/index.md index d588e5b7960..75f5a4f2ba6 100644 --- a/dify-agent/docs/dify-agent/guide/index.md +++ b/dify-agent/docs/dify-agent/guide/index.md @@ -15,7 +15,8 @@ uv run --project dify-agent uvicorn dify_agent.server.app:app --reload By default, the FastAPI lifespan creates: - one Redis-backed run store used by HTTP routes -- one shared plugin-daemon `httpx.AsyncClient` used by local run tasks +- shared plugin-daemon and Dify API inner `httpx.AsyncClient` instances +- one deployment-selected, stateless runtime backend profile when configured - one process-local scheduler that starts background `asyncio` run tasks This means local development needs one uvicorn process plus Redis, and @@ -38,19 +39,32 @@ also reads `.env` and `dify-agent/.env` when present. | `DIFY_AGENT_PLUGIN_DAEMON_API_KEY` | empty | API key sent to the Dify plugin daemon. | | `DIFY_AGENT_INNER_API_URL` | `http://localhost:5001` | Dify API service root used when dify-agent calls `/inner/api/...` endpoints. | | `DIFY_AGENT_INNER_API_KEY` | empty | API key sent to Dify API inner plugin endpoints. Set this to Dify API `INNER_API_KEY_FOR_PLUGIN` (Docker: `PLUGIN_DIFY_INNER_API_KEY`). | -| `DIFY_AGENT_SHELLCTL_ENTRYPOINT` | empty | Base URL for the shellctl server used by `dify.shell`; required when runs include the shell layer. | -| `DIFY_AGENT_SHELLCTL_AUTH_TOKEN` | empty | Optional bearer token sent to the shellctl server. | -| `DIFY_AGENT_SHELL_HOME_ROOT` | `/home` | Root for per-Agent shell HOME directories. Set this to a writable path such as `/tmp/dify-agent-home` for local macOS development. | -| `DIFY_AGENT_STUB_API_BASE_URL` | empty | Public Agent Stub API base URL reachable from shellctl-managed remote machines. HTTP may be the service root or `/agent-stub`; gRPC must be `grpc://host:port`. Enables `DIFY_AGENT_STUB_*` env injection for user `shell.run` jobs. | -| `DIFY_AGENT_STUB_GRPC_BIND_ADDRESS` | empty | Optional `host:port` bind override used only when `DIFY_AGENT_STUB_API_BASE_URL` uses `grpc://`. | +| `DIFY_AGENT_RUNTIME_BACKEND` | `local` | Selects one coherent `local`, `enterprise`, or `e2b` Home Snapshot + Execution Binding backend profile. | +| `DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT` | empty | Local shellctl data-plane URL. With the default Local selection, leaving it empty disables `dify.runtime` and resource endpoints. | +| `DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN` | empty | Optional bearer token sent to Local shellctl. | +| `DIFY_AGENT_LOCAL_SANDBOX_MATERIALIZED_HOME_ROOT` | `/home/dify/.dify-agent-materialized-homes` | Root directory, on the Local shellctl filesystem, for per-Binding materialized Homes. | +| `DIFY_AGENT_LOCAL_SANDBOX_WORKSPACE_ROOT` | `/home/dify/.dify-agent-workspaces` | Root directory, on the Local shellctl filesystem, for mutable Workspaces. | +| `DIFY_AGENT_LOCAL_SANDBOX_HOME_SNAPSHOT_ROOT` | `/home/dify/.dify-agent-home-snapshots` | Root directory, on the Local shellctl filesystem, for immutable Home Snapshots. | +| `DIFY_AGENT_ENTERPRISE_SANDBOX_GATEWAY_ENDPOINT` | empty | Enterprise Gateway endpoint required by configuration. Default-Home Bindings are supported; immutable Home Snapshot operations remain unsupported. | +| `DIFY_AGENT_ENTERPRISE_SANDBOX_GATEWAY_AUTH_TOKEN` | empty | Optional `X-Inner-Api-Key` sent to the Enterprise Gateway. | +| `DIFY_AGENT_ENTERPRISE_SANDBOX_GATEWAY_TIMEOUT` | `30` | Enterprise control-plane timeout in seconds. | +| `DIFY_AGENT_ENTERPRISE_SANDBOX_PROXY_TIMEOUT` | `60` | Enterprise shellctl-proxy timeout in seconds. | +| `DIFY_AGENT_E2B_API_KEY` | empty | E2B API key; required for E2B. | +| `DIFY_AGENT_E2B_TEMPLATE` | `difys-default-team/dify-agent-local-sandbox` | Prepared E2B template containing shellctl and the deployment-default Home environment. | +| `DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS` | `3600` | Maximum continuous active time for the RuntimeLease spanning one complete Agent run. Binding resources pause on timeout. This is not a retention TTL. | +| `DIFY_AGENT_E2B_SHELLCTL_AUTH_TOKEN` | empty | Optional bearer token expected by shellctl inside the E2B template. | +| `DIFY_AGENT_E2B_SHELLCTL_PORT` | `5004` | shellctl port exposed by the E2B template. | +| `DIFY_AGENT_SHELL_REDACT_PATTERNS` | empty | JSON array of additional regex patterns redacted from Shell output. | +| `DIFY_AGENT_STUB_API_BASE_URL` | empty | HTTP(S) Agent Stub API base URL reachable from the Sandbox. It may be the service root or `/agent-stub`. Enables `DIFY_AGENT_STUB_*` env injection for user `shell.run` jobs. | +| `DIFY_AGENT_SANDBOX_FILES_BASE_URL` | empty | Dify API base URL reachable from the Sandbox for signed `/files/*` upload/download bytes, including Config file and skill pulls. Required when Agent Stub file operations are enabled. May include an ingress path prefix, but not a query or fragment. | | `DIFY_AGENT_SERVER_SECRET_KEY` | empty | Security-sensitive server-wide root secret used to derive the JWE encryption key for Agent Stub bearer tokens; required when `DIFY_AGENT_STUB_API_BASE_URL` is set. The supplied default config uses a development value; set a unique unpadded base64url 32-byte secret in production. | -| `DIFY_AGENT_PLUGIN_DAEMON_CONNECT_TIMEOUT` | `10` | Plugin-daemon HTTP connect timeout in seconds. | -| `DIFY_AGENT_PLUGIN_DAEMON_READ_TIMEOUT` | `600` | Plugin-daemon HTTP read timeout in seconds. | -| `DIFY_AGENT_PLUGIN_DAEMON_WRITE_TIMEOUT` | `30` | Plugin-daemon HTTP write timeout in seconds. | -| `DIFY_AGENT_PLUGIN_DAEMON_POOL_TIMEOUT` | `10` | Plugin-daemon HTTP connection-pool wait timeout in seconds. | -| `DIFY_AGENT_PLUGIN_DAEMON_MAX_CONNECTIONS` | `100` | Maximum total plugin-daemon HTTP connections. | -| `DIFY_AGENT_PLUGIN_DAEMON_MAX_KEEPALIVE_CONNECTIONS` | `20` | Maximum idle keep-alive plugin-daemon HTTP connections. | -| `DIFY_AGENT_PLUGIN_DAEMON_KEEPALIVE_EXPIRY` | `30` | Keep-alive expiry in seconds for idle plugin-daemon HTTP connections. | +| `DIFY_AGENT_OUTBOUND_HTTP_CONNECT_TIMEOUT` | `10` | Shared outbound HTTP connect timeout in seconds. | +| `DIFY_AGENT_OUTBOUND_HTTP_READ_TIMEOUT` | `600` | Shared outbound HTTP read timeout in seconds. | +| `DIFY_AGENT_OUTBOUND_HTTP_WRITE_TIMEOUT` | `30` | Shared outbound HTTP write timeout in seconds. | +| `DIFY_AGENT_OUTBOUND_HTTP_POOL_TIMEOUT` | `10` | Shared outbound connection-pool wait timeout in seconds. | +| `DIFY_AGENT_OUTBOUND_HTTP_MAX_CONNECTIONS` | `100` | Maximum total shared outbound HTTP connections. | +| `DIFY_AGENT_OUTBOUND_HTTP_MAX_KEEPALIVE_CONNECTIONS` | `20` | Maximum idle shared outbound HTTP connections. | +| `DIFY_AGENT_OUTBOUND_HTTP_KEEPALIVE_EXPIRY` | `30` | Idle keep-alive expiry in seconds. | Example `.env`: @@ -63,20 +77,176 @@ DIFY_AGENT_PLUGIN_DAEMON_URL=http://localhost:5002 DIFY_AGENT_PLUGIN_DAEMON_API_KEY=replace-with-daemon-key DIFY_AGENT_INNER_API_URL=http://localhost:5001 DIFY_AGENT_INNER_API_KEY=replace-with-dify-inner-api-key-for-plugin -DIFY_AGENT_SHELLCTL_ENTRYPOINT=http://127.0.0.1:5004 -DIFY_AGENT_SHELLCTL_AUTH_TOKEN=replace-with-shellctl-token -DIFY_AGENT_SHELL_HOME_ROOT=/tmp/dify-agent-home +DIFY_AGENT_RUNTIME_BACKEND=local +DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT=http://127.0.0.1:5004 +DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN=replace-with-shellctl-token +# Set these when shellctl runs directly on a host that does not have /home/dify. +DIFY_AGENT_LOCAL_SANDBOX_MATERIALIZED_HOME_ROOT=/tmp/dify-agent/materialized-homes +DIFY_AGENT_LOCAL_SANDBOX_WORKSPACE_ROOT=/tmp/dify-agent/workspaces +DIFY_AGENT_LOCAL_SANDBOX_HOME_SNAPSHOT_ROOT=/tmp/dify-agent/home-snapshots DIFY_AGENT_STUB_API_BASE_URL=https://agent.example.com/agent-stub +DIFY_AGENT_SANDBOX_FILES_BASE_URL=https://dify.example.com # This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens. # Replace this development default in production. # Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))' DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY ``` +The two Sandbox-facing base URLs have different owners. Agent Stub control +requests use `DIFY_AGENT_STUB_API_BASE_URL`; signed file bytes use +`DIFY_AGENT_SANDBOX_FILES_BASE_URL`. `DIFY_AGENT_INNER_API_URL` remains a +trusted service-to-service URL and is never returned to the Sandbox. + +Config file and skill pulls use the same split: Agent Stub authorizes the +Config target and returns a short-lived URL, then the Sandbox fetches the bytes +directly from the Dify API `/files/*` data plane. + +Removing Agent Stub gRPC is a breaking transport migration: replace every +`grpc://` Agent Stub URL with HTTP(S), remove +`DIFY_AGENT_STUB_GRPC_BIND_ADDRESS`, and deploy without a gRPC fallback. + +For a remote Sandbox, expose only `/agent-stub/*` from Agent Backend and the +existing `/files/*` Dify API data plane. The `/files/*` ingress must preserve +the complete signed query string, allow the configured upload body size, and +use response streaming and timeouts suitable for large downloads. Do not expose +Agent Backend `/runs`, Workspace, or Binding management routes through the +Sandbox ingress. + +Browser presentation URLs are independent. Configure Dify API `FILES_URL` to a +browser-reachable public origin, or leave it empty so responses use same-origin +relative `/files/...` URIs. Never set `FILES_URL` to a Docker-only service name +such as `http://api:5001`. + +`DIFY_AGENT_SHELLCTL_ENTRYPOINT` and `DIFY_AGENT_SHELLCTL_AUTH_TOKEN` remain +accepted only as legacy aliases for the two Local settings. New deployments +must use `DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT` and +`DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN`. There is no compatibility setting for +the removed shell-provider selector. + +The backend selection is deployment-private. Shell-enabled run requests use an +Execution Context, `dify.runtime`, and `dify.shell` graph. Runtime config carries +only the opaque `backend_binding_ref` resolved by Dify API. See +[Runtime resources](../concepts/runtime-resources/index.md) for the ownership +and lifecycle contract. + Run records and event streams use the same retention. Status writes refresh the record TTL, and event writes refresh both the stream TTL and the corresponding record TTL so active runs that keep producing events remain observable. +## Validate the E2B Compose deployment + +The E2B overlay requires a prepared template that starts shellctl on port 5004. +The default is `difys-default-team/dify-agent-local-sandbox`. It also requires +an E2B API key in `DIFY_AGENT_E2B_API_KEY`; the Compose interpolation accepts +`E2B_API_KEY` and `E2B_API_TOKEN` only as deployment-level fallbacks. + +From the repository root, keep the normal `docker/.env` unchanged and export the +secret for this shell: + +```bash +export DIFY_AGENT_E2B_API_KEY="${DIFY_AGENT_E2B_API_KEY:-${E2B_API_KEY:-${E2B_API_TOKEN:-}}}" +test -n "$DIFY_AGENT_E2B_API_KEY" +export DIFY_AGENT_E2B_TEMPLATE=difys-default-team/dify-agent-local-sandbox + +docker compose \ + --env-file docker/.env \ + -f docker/docker-compose.yaml \ + -f docker/docker-compose.e2b.yaml \ + up -d --build +``` + +The overlay builds `dify-api:e2b-local` and +`dify-agent-backend:e2b-local` from the current checkout. It disables the +normal Local sandbox service, switches Dify Agent to E2B, and mounts PostgreSQL +on the separate `dify_e2b_postgres_data` Compose volume. That database is empty +when the volume is first created and is isolated from the normal stack; later +starts reuse it until an operator explicitly removes the volume. + +Verify the merged deployment and the branch-built Dify Agent API: + +```bash +docker compose \ + --env-file docker/.env \ + -f docker/docker-compose.yaml \ + -f docker/docker-compose.e2b.yaml \ + ps + +agent_backend_port="$( + docker compose \ + --env-file docker/.env \ + -f docker/docker-compose.yaml \ + -f docker/docker-compose.e2b.yaml \ + port agent_backend 5050 | awk -F: 'NR == 1 { print $NF }' +)" +test -n "$agent_backend_port" +curl --fail --silent --show-error \ + --connect-timeout 2 --max-time 5 \ + --retry 12 --retry-delay 1 --retry-connrefused --retry-max-time 60 \ + "http://127.0.0.1:${agent_backend_port}/openapi.json" \ + >/dev/null + +docker compose \ + --env-file docker/.env \ + -f docker/docker-compose.yaml \ + -f docker/docker-compose.e2b.yaml \ + logs --tail=100 agent_backend api worker +``` + +Stop the validation stack without deleting its isolated database volume: + +```bash +docker compose \ + --env-file docker/.env \ + -f docker/docker-compose.yaml \ + -f docker/docker-compose.e2b.yaml \ + down +``` + +`DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS` controls continuous active E2B time. +The physical resource behind a Binding pauses when that timeout fires, preserving +the current Workspace. The setting is not a resource-age TTL and does not delete +paused resources or immutable snapshots. + +## Run runtime-backend integration contracts + +Run the disposable Local contract from the `dify-agent` directory. The script +starts one local-sandbox container on an unused port and removes that exact +container on exit: + +```bash +cd dify-agent +DIFY_AGENT_TEST_LOCAL_SANDBOX_IMAGE=langgenius/dify-agent-local-sandbox:1.16.0 \ + tests/integration/dify_agent/runtime_backend/run_local_integration.sh +``` + +To use an already managed Local shellctl endpoint instead: + +```bash +cd dify-agent +DIFY_AGENT_TEST_LOCAL_SHELLCTL_ENDPOINT=http://127.0.0.1:5004 \ +DIFY_AGENT_TEST_LOCAL_SHELLCTL_AUTH_TOKEN=replace-with-shellctl-token \ + pdm run pytest --import-mode=importlib \ + tests/integration/dify_agent/runtime_backend/test_working_environment.py \ + -k local -q -rs +``` + +Run the real E2B contract with an explicit test credential and template: + +```bash +cd dify-agent +DIFY_AGENT_TEST_E2B_API_KEY="$E2B_API_TOKEN" \ +DIFY_AGENT_TEST_E2B_TEMPLATE=difys-default-team/dify-agent-local-sandbox \ + pdm run pytest --import-mode=importlib \ + tests/integration/dify_agent/runtime_backend/test_working_environment.py \ + -k e2b -q -rs +``` + +The Local auth token is optional when shellctl has authentication disabled. +The E2B contract uses the one-hour `E2B_MAX_ACTIVE_TIMEOUT_SECONDS` RuntimeLease +limit. This is continuous active test time, not a post-test retention TTL. Both +contracts create unique resources and perform explicit cleanup in `finally` +blocks. + ## Scheduling and shutdown semantics `POST /runs` persists a `running` run record and starts an `asyncio` task in the @@ -85,15 +255,39 @@ automatic retry layer. Request-shaped runtime failures such as bad composition, prompt, output, or snapshot inputs are reported later as failed runs rather than rejected synchronously once the request DTO itself is accepted. +Each run explicitly limits Pydantic AI to 100 model-request steps. Tool calls do +not have a separate count limit, but every model request used to continue the +tool loop consumes one of those steps. + During FastAPI shutdown the scheduler rejects new runs, waits up to `DIFY_AGENT_SHUTDOWN_GRACE_SECONDS` for active tasks, then cancels remaining tasks -and best-effort appends a `run_failed` event plus failed status. A hard process -crash can still leave active runs stuck as `running`; there is no in-service -recovery or worker handoff. +and attempts to finalize them as failed. Success, failure, cancellation, and this +shutdown path all use one atomic Redis transition: only the first transition from +`running` appends a terminal event and updates the run record. A later terminal +attempt leaves both the record and event stream unchanged. A hard process crash +can still leave active runs stuck as `running`; there is no in-service recovery +or worker handoff. Horizontal scaling is possible by running multiple API processes against the same Redis prefix, but each process executes only the runs it accepted. Redis provides -shared status/event visibility, not load balancing or queued-job recovery. +shared status/event visibility, not load balancing or queued-job recovery. The +cancel endpoint can atomically accept a running run on any process. The process +that owns the runner observes the shared `run_cancelled` event, then cancels and +cleans up its local task. The HTTP response confirms that logical cancellation is +durable; local runner cleanup may still be in progress. Retrying a cancellation +after the run is already `cancelled` is idempotent. + +Atomic terminal finalization currently assumes the configured Redis URL targets +one Redis deployment that can execute both run keys in a Lua script. The existing +record and event key names are unchanged and do not contain a shared Redis +Cluster hash tag, so Redis Cluster is not supported for this transition. During +a rolling upgrade, older processes can still use the former split event/status +writes; treat the single-terminal invariant as active only after those processes +have exited. Deploy atomic terminal finalization everywhere first, then ensure +every process that can own a runner has the cancellation observer before relying +on route-independent cancellation. Operators should then alert on more than one +terminal event per run and on disagreement between the run record status and +terminal event type. ## Run inputs and session snapshots @@ -133,21 +327,42 @@ Use the HTTP status endpoint for coarse state and the event endpoints for detail progress: - `POST /runs` creates a running run and schedules it locally. -- `GET /runs/{run_id}` returns `running`, `succeeded`, or `failed`. +- `GET /runs/{run_id}` returns `running`, `succeeded`, `failed`, or `cancelled`. + Failed records can also expose a stable machine-readable `error_type` alongside + the diagnostic `error` text. +- `POST /runs/{run_id}/cancel` atomically accepts cancellation on any API process + and emits `run_cancelled`; it returns `409` only when a success/failure terminal + already won. Runner cleanup continues asynchronously on the owner process. - `GET /runs/{run_id}/events` polls the Redis Stream event log with `after` and `next_cursor` cursors. - `GET /runs/{run_id}/events/sse` replays and streams events over SSE. The SSE `id` is the event Redis Stream ID. `after` query cursors take precedence over - `Last-Event-ID` headers. + `Last-Event-ID` headers. The server closes the SSE response normally after + delivering a terminal event. Clients must stop reconnecting after consuming + that event. Both cursor forms remain exclusive resume cursors, so the server + does not resend a terminal event that the supplied cursor already excludes. Successful runs emit `run_started`, zero or more `pydantic_ai_event`, and -`run_succeeded`. Failed runs end with `run_failed`. Event envelopes retain `id`, -`run_id`, `type`, `data`, and `created_at`; `data` is typed per event type, +`run_succeeded`. Failed runs end with `run_failed`, and accepted cancellations +end with `run_cancelled`. Each run can append at most one of these terminal +events. Event envelopes retain `id`, `run_id`, `type`, `data`, and `created_at`; +`data` is typed per event type, including Pydantic AI's `AgentStreamEvent` payload for `pydantic_ai_event` and a terminal `run_succeeded.data` object containing a `CompositorSessionSnapshot` for resumption. A successful run has exactly one active result branch: JSON-safe `output` for final answers, or `deferred_tool_call` when a layer such as `dify.ask_human` ends the current agent run with an external deferred tool call. +Failed event payloads contain the diagnostic `error`, optional source-specific +`reason`, and optional stable `error_type`. Pydantic AI request/step budget +exhaustion enforced by Dify Agent is reported as +`error_type: "agent_run_limit_exceeded"`; consumers should branch on that value +rather than parsing the error text. This type does not classify wall-clock run +timeouts, whose classification is not implemented in this release, or provider +and connection timeouts. The matching failed run record and terminal event are +committed atomically with the same error type. For independently deployed Agent +backend and API services, deploy consumers that accept the optional field before +producers begin emitting it because the public protocol models reject unknown +fields. ## Examples diff --git a/dify-agent/docs/dify-agent/user-manual/shell-layer/index.md b/dify-agent/docs/dify-agent/user-manual/shell-layer/index.md index f3f6a77c4de..0169af0369d 100644 --- a/dify-agent/docs/dify-agent/user-manual/shell-layer/index.md +++ b/dify-agent/docs/dify-agent/user-manual/shell-layer/index.md @@ -1,126 +1,158 @@ # Shell layer -The shell layer lets a Dify Agent run expose a `shellctl`-backed workspace to the -model. This page is for Dify Agent clients that build `CreateRunRequest` -payloads. It explains how to add the layer to a run composition and how the -server-side runtime must be wired. +The `dify.shell` layer exposes shellctl-backed commands to an Agent. +It does not select a backend or own persistent Home, Workspace, or Binding +resources. It consumes the operation-scoped `RuntimeLease` opened by a sibling +`dify.runtime` layer. -The layer type id is `dify.shell`. Its public config is intentionally empty: +## Public configuration ```python -from dify_agent.layers.shell import DIFY_SHELL_LAYER_TYPE_ID, DifyShellLayerConfig +from dify_agent.layers.shell import ( + DIFY_SHELL_LAYER_TYPE_ID, + DifyShellEnvVarConfig, + DifyShellLayerConfig, +) from dify_agent.protocol import RunLayerSpec RunLayerSpec( name="shell", type=DIFY_SHELL_LAYER_TYPE_ID, - config=DifyShellLayerConfig(), -) -``` - -Server-only settings, such as the `shellctl` HTTP entrypoint and auth token, are -injected by the Dify Agent runtime provider. They are not part of -`DifyShellLayerConfig` and should not be submitted by clients in the public run -request. - -## Runtime requirements - -When a run includes `dify.shell`, the Dify Agent server must construct its layer -providers with a non-empty shellctl entrypoint: - -```python -from dify_agent.adapters.shell.shellctl import ShellctlProvider -from dify_agent.runtime.compositor_factory import create_default_layer_providers - -layer_providers = create_default_layer_providers( - plugin_daemon_url="http://localhost:5002", - plugin_daemon_api_key="replace-with-plugin-daemon-key", - shell_provider=ShellctlProvider( - entrypoint="http://127.0.0.1:5004", - token="replace-with-shellctl-token", # optional; defaults to empty string + deps={"execution_context": "execution_context", "runtime": "runtime"}, + config=DifyShellLayerConfig( + env=[DifyShellEnvVarConfig(name="REPORT_FORMAT", value="markdown")], + redact_patterns=["private-[A-Za-z0-9]+"], ), ) ``` -In the FastAPI server, these values are read from environment-backed -`ServerSettings` fields: +| Config field | Meaning | +| --- | --- | +| `agent_stub_drive_ref` | Optional Drive ref used by shell-visible Agent Stub commands. | +| `cli_tools` | CLI bootstrap declarations with install commands and scoped environment metadata. | +| `env` | Normal environment variables exported to Shell commands. | +| `secret_refs` | Names of secret environment variables supplied by the backend environment. | +| `redact_patterns` | Request-level regex patterns removed from Shell output shown to the model. | -```env -DIFY_AGENT_SHELLCTL_ENTRYPOINT=http://127.0.0.1:5004 -DIFY_AGENT_SHELLCTL_AUTH_TOKEN=replace-with-shellctl-token -DIFY_AGENT_SHELL_HOME_ROOT=/tmp/dify-agent-home +Endpoints, credentials, Home/Workspace paths, resource refs, timeouts, and +network policy are not Shell config. Backend selection is server-private, and +the opaque Binding ref belongs to `DifyRuntimeLayerConfig`. + +## Runtime requirements + +The server constructs one coherent runtime backend profile. Local and E2B +implement Home Snapshot and Execution Binding operations. Enterprise implements +default-Home Binding creation, acquisition, and coupled destruction, while +immutable Home Snapshot operations fail fast; there is no compatibility +fallback to the retired Sandbox protocol. + +```python +from dify_agent.runtime.compositor_factory import create_default_layer_providers +from dify_agent.runtime_backend.profile import RuntimeBackendSettings, create_runtime_backend_profile + +runtime_backend_profile = create_runtime_backend_profile( + RuntimeBackendSettings( + runtime_backend="local", + local_sandbox_endpoint="http://127.0.0.1:5004", + local_sandbox_auth_token="replace-with-shellctl-token", + ) +) + +layer_providers = create_default_layer_providers( + plugin_daemon_url="http://localhost:5002", + plugin_daemon_api_key="replace-with-plugin-daemon-key", + runtime_backend_profile=runtime_backend_profile, +) ``` -`DIFY_AGENT_SHELLCTL_AUTH_TOKEN` defaults to `None`/empty, which keeps the shell -client on the no-token path. Set it only when the shellctl server is started with -bearer authentication. +Equivalent standalone environment settings are: -`DIFY_AGENT_SHELL_HOME_ROOT` defaults to `/home`, matching Linux and container -deployments. For local macOS development, set it to a writable directory such as -`/tmp/dify-agent-home`; the shell layer creates per-Agent home directories under -that root. +```env +DIFY_AGENT_RUNTIME_BACKEND=local +DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT=http://127.0.0.1:5004 +DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN=replace-with-shellctl-token +# Optional when shellctl runs directly on a host without /home/dify: +# DIFY_AGENT_LOCAL_SANDBOX_MATERIALIZED_HOME_ROOT=/tmp/dify-agent/materialized-homes +# DIFY_AGENT_LOCAL_SANDBOX_WORKSPACE_ROOT=/tmp/dify-agent/workspaces +# DIFY_AGENT_LOCAL_SANDBOX_HOME_SNAPSHOT_ROOT=/tmp/dify-agent/home-snapshots +``` -To let commands inside user-visible shell jobs call back to the Dify Agent server -with `dify-agent ...`, also enable the Agent Stub: +The auth token may be empty when shellctl authentication is disabled. E2B uses +`DIFY_AGENT_E2B_API_KEY`, the prepared template, and its shellctl settings. + +To let shell jobs call the Agent Stub with `dify-agent ...`, configure a +Sandbox-reachable Agent Stub URL and a unique production secret. Remote +deployments normally use a public Agent ingress. Local Compose uses +`http://agent_backend:5050/agent-stub`, reached through the existing +`agent_ssrf_proxy`; this configuration does not change the Compose network +topology. ```env DIFY_AGENT_STUB_API_BASE_URL=https://agent.example.com/agent-stub -# This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens. -# Replace this development default in production. -# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))' -DIFY_AGENT_SERVER_SECRET_KEY=MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY +DIFY_AGENT_SANDBOX_FILES_BASE_URL=https://dify.example.com +DIFY_AGENT_SERVER_SECRET_KEY=replace-with-unpadded-base64url-for-32-random-bytes ``` -HTTP `DIFY_AGENT_STUB_API_BASE_URL` may be either the service root or the -explicit `/agent-stub` API root; the server normalizes the service root to -`/agent-stub`. Other HTTP paths are rejected at startup. +HTTP URLs may be either the service root or the explicit `/agent-stub` root. +The server normalizes a service root and rejects unrelated paths. The separate +Sandbox file base must point to the Dify API ingress serving `/files/*`; it is +used for CLI upload/download bytes, including Config file and skill pulls. -The supplied Docker and `.example.env` configs use a development -`DIFY_AGENT_SERVER_SECRET_KEY`. Override it in production with unpadded base64url -text for exactly 32 decoded bytes. One way to generate it is: +After `dify-agent file upload ` succeeds, the CLI prints JSON such as: -```bash -python -c 'import secrets; print(secrets.token_urlsafe(32))' +```json +{ + "transfer_method": "tool_file", + "reference": "dify-file-ref:...", + "public_download_url": "https://dify.example.com/files/tools/..." +} ``` -## Client request shape +`reference` is the persistent canonical file identity and should be stored in +structured output. `public_download_url` is a short-lived frontend presentation +address: it is an absolute URL when Dify API `FILES_URL` has a public origin, +or a same-origin `/files/...` relative URI when `FILES_URL` is empty. The CLI +does not access this field. The former ambiguous `download_url` upload-output +key is now named `public_download_url`. -A client adds the shell layer as an ordinary composition layer. Basic shell jobs -do not need dependencies. To inject `DIFY_AGENT_STUB_API_BASE_URL`, -`DIFY_AGENT_STUB_AUTH_JWE`, and `DIFY_AGENT_STUB_DRIVE_BASE` into user-visible -`shell.run` jobs, declare the execution-context layer as the shell layer's -`execution_context` dependency. When the run also includes `dify.drive`, declare -it as the shell layer's `drive` dependency; the injected drive base is then -computed from the fixed Agent Stub drive mount and the drive reference, for -example `/mnt/drive/agent-123`. Without a drive dependency, the CLI keeps the -historical `/mnt/drive` fallback. A typical run still also includes: +Server-side Binding downloads use +`dify-agent file upload --no-download-link `. This additive mode performs +the same streaming ToolFile upload but skips the download-request step and +prints only `transfer_method` plus the canonical `reference`. The regular +`file upload` command keeps the link-producing behavior shown above. -- a prompt layer that supplies the task; -- an execution-context layer carrying tenant/user context; -- an LLM layer named `llm`. +## Request graph -When clients want the shell workspace and shellctl job records to be removed at -the end of the run, set `on_exit.default` to `delete`. +A shell-enabled run contains Execution Context, Runtime, and Shell layers: -## Example: CSV analysis run +```mermaid +flowchart LR + EC["execution_context"] --> SH["shell"] + RT["runtime
backend_binding_ref"] --> SH +``` -The following example mirrors a verified run with a real Gemini model and a -temporary shellctl server. The client gives the model a small CSV-shaped dataset -and asks for computed metrics without prescribing the exact shell commands. +`DifyRuntimeLayer` acquires the Binding when the run's resource context opens +and releases it when that operation exits. `DifyShellLayer` uses the active +lease's commands, Home path, and Workspace path. It performs only +best-effort cleanup of shell jobs; the persistent Binding lifecycle remains in +Dify API. -### Request +## Example request + +The Binding must already have been resolved or created by Dify API. Its backend +ref is opaque to the request builder: ```python {test="skip" lint="skip"} -from agenton.layers import ExitIntent from agenton_collections.layers.plain import PromptLayerConfig from dify_agent.layers.dify_plugin.configs import DifyPluginLLMLayerConfig from dify_agent.layers.execution_context import ( DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, DifyExecutionContextLayerConfig, ) +from dify_agent.layers.runtime import DIFY_RUNTIME_LAYER_TYPE_ID, DifyRuntimeLayerConfig from dify_agent.layers.shell import DIFY_SHELL_LAYER_TYPE_ID, DifyShellLayerConfig from dify_agent.protocol import DIFY_AGENT_MODEL_LAYER_ID -from dify_agent.protocol.schemas import CreateRunRequest, LayerExitSignals, RunComposition, RunLayerSpec +from dify_agent.protocol.schemas import CreateRunRequest, RunComposition, RunLayerSpec request = CreateRunRequest( @@ -130,35 +162,33 @@ request = CreateRunRequest( name="prompt", type="plain.prompt", config=PromptLayerConfig( - prefix="You are a practical data analyst. Give a concise final answer.", - user="""Analyze this small sales dataset with pandas. Use any local computation you think is useful. - -region,product,units,unit_price -north,widget,12,3.50 -north,gadget,5,9.00 -south,widget,7,3.50 -south,gadget,9,9.00 -west,widget,4,3.50 -west,gadget,11,9.00 - -Report the total revenue, the region with the most revenue, total units by -product, and a SHA-256 hash of the CSV content.""", + prefix="Use the workspace when local computation helps.", + user="Create report.txt containing the current UTC timestamp, then summarize it.", ), ), - RunLayerSpec( - name="shell", - type=DIFY_SHELL_LAYER_TYPE_ID, - deps={"execution_context": "execution_context"}, - config=DifyShellLayerConfig(), - ), RunLayerSpec( name="execution_context", type=DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, config=DifyExecutionContextLayerConfig( tenant_id="92cca973-2d6f-45e0-906e-0b7eda5f2ccf", - invoke_from="workflow_run", + agent_id="8d542564-159d-4168-985c-dde8d8ff6092", + agent_config_version_id="931a4cee-4434-4c1c-8fbd-0a3c7591095d", + agent_config_version_kind="snapshot", + agent_mode="workflow_run", + invoke_from="debugger", ), ), + RunLayerSpec( + name="runtime", + type=DIFY_RUNTIME_LAYER_TYPE_ID, + config=DifyRuntimeLayerConfig(backend_binding_ref="opaque-backend-binding-ref"), + ), + RunLayerSpec( + name="shell", + type=DIFY_SHELL_LAYER_TYPE_ID, + deps={"execution_context": "execution_context", "runtime": "runtime"}, + config=DifyShellLayerConfig(), + ), RunLayerSpec( name=DIFY_AGENT_MODEL_LAYER_ID, type="dify.plugin.llm", @@ -166,88 +196,58 @@ product, and a SHA-256 hash of the CSV content.""", config=DifyPluginLLMLayerConfig( plugin_id="langgenius/gemini", model_provider="google", - model="gemini-3.5-flash", + model="gemini-2.5-flash", credentials={"google_api_key": ""}, ), ), ] - ), - on_exit=LayerExitSignals(default=ExitIntent.DELETE), + ) ) ``` -The same request serialized as JSON has these important layer entries: +The resource part serializes as: ```json { - "composition": { - "schema_version": 1, - "layers": [ - {"name": "prompt", "type": "plain.prompt"}, - { - "name": "shell", - "type": "dify.shell", - "deps": {"execution_context": "execution_context"}, - "config": {} - }, - {"name": "execution_context", "type": "dify.execution_context"}, - { - "name": "llm", - "type": "dify.plugin.llm", - "deps": {"execution_context": "execution_context"} - } - ] - }, - "on_exit": {"default": "delete", "layers": {}} + "layers": [ + { + "name": "runtime", + "type": "dify.runtime", + "config": {"backend_binding_ref": "opaque-backend-binding-ref"} + }, + { + "name": "shell", + "type": "dify.shell", + "deps": {"execution_context": "execution_context", "runtime": "runtime"}, + "config": {} + } + ] } ``` -### Final answer +## Paths and persistence -The terminal `run_succeeded` output was: +`RuntimeLease.layout.home_dir` and `workspace_dir` are absolute paths inside the +backend execution namespace. They are not host filesystem paths and are not +sent in the run request. Shell commands start in `workspace_dir`, while `HOME` +is forced to `home_dir`; `~` therefore resolves to the current Binding's +materialized Home. -````markdown -Here is the analysis of the sales dataset: +Workspace content persists with the Workspace until Dify API retires and +collects it. Releasing a RuntimeLease ends only the current operation. Dify API can later +browse the Binding filesystem through Dify Agent's private +`/execution-bindings/files/list`, `/execution-bindings/files/read`, and +`/execution-bindings/files/download` routes, each of which acquires a fresh +lease. Relative paths start in the Workspace, while `~` starts in the Binding's +Home. Download uses the installed CLI to stream a ToolFile upload and returns +only its canonical reference to Dify API. -* **Total Revenue:** **$305.50** -* **Top Region:** **west** with **$113.00** -* **Total Units by Product:** gadget: 25 units, widget: 23 units -* **SHA-256 Hash:** `e86521a0d759037a09b059cb3cb2419f0a3f06e674db8151ccf2f93811dac0b8` -```` +On Local, multiple Bindings may share a Workspace while each receives a +separate materialized Home. Those directories may be siblings in one shellctl +namespace; path isolation restricts a lease to its Home and Workspace. On E2B, +one physical E2B resource currently represents both Binding and Workspace, so +shared Workspace attachment is unsupported. -## Running the local sandbox in Docker - -Build the local sandbox image from the Dify Agent package root: - -```bash -docker build -f docker/local-sandbox/Dockerfile -t dify-agent-local-sandbox:local . -``` - -Run it with a bearer token and publish the API on localhost: - -```bash -docker run --rm --name dify-agent-local-sandbox \ - -e SHELLCTL_AUTH_TOKEN=replace-with-a-token \ - -p 127.0.0.1:5004:5004 \ - dify-agent-local-sandbox:local -``` - -The image starts `shellctl serve --listen 0.0.0.0:5004` as the non-root -`dify` user. It also sets the fallback `DIFY_AGENT_STUB_DRIVE_BASE=/mnt/drive` -and pre-creates that directory with write access for the same user. - -## Docker image contents - -The provided `docker/local-sandbox/Dockerfile` installs: - -- `tmux`, required by `shellctl` to manage shell jobs; -- common shell workspace tools: `git`, `openssh-client`, `jq`, `ripgrep`, - `unzip`, `zip`, `file`, `procps`, and `less`; -- `dify-agent[grpc,shellctl-server]` as a standalone uv tool, which provides - both the Agent Stub client CLI and the built-in `shellctl` CLI/server; -- `uv`, so uv shebang scripts with PEP 723 metadata can run inside the shell - workspace and Python CLI tools can be installed with isolated tool - environments; -- `node==22.22.1` and `pnpm==11.9.0`, so JavaScript and TypeScript tooling can - run inside the shell workspace without per-job installation; -- a non-root default user named `dify`. +See [Runtime resources](../../concepts/runtime-resources/index.md) for the +ledger and lifecycle contract. The [Operations Guide](../../guide/index.md) +covers Local and E2B validation. diff --git a/dify-agent/mkdocs.yml b/dify-agent/mkdocs.yml index 3993b3618e5..34ba0e2c92d 100644 --- a/dify-agent/mkdocs.yml +++ b/dify-agent/mkdocs.yml @@ -16,6 +16,7 @@ nav: - Overview: dify-agent/index.md - Concepts: - Agent Run Lifecycle: dify-agent/concepts/run-lifecycle/index.md + - Runtime Resources: dify-agent/concepts/runtime-resources/index.md - User Manual: - Get Started: dify-agent/get-started/index.md - Prompt Layer: dify-agent/user-manual/prompt-layer/index.md diff --git a/dify-agent/proto/dify/agent/stub/v1/agent_stub.proto b/dify-agent/proto/dify/agent/stub/v1/agent_stub.proto deleted file mode 100644 index 608d6d0006c..00000000000 --- a/dify-agent/proto/dify/agent/stub/v1/agent_stub.proto +++ /dev/null @@ -1,49 +0,0 @@ -syntax = "proto3"; - -package dify.agent.stub.v1; - -option go_package = "dify/agent/stub/v1;stubv1"; - -service AgentStubService { - rpc Connect(ConnectRequest) returns (ConnectResponse); - rpc CreateFileUploadRequest(FileUploadRequest) returns (FileUploadResponse); - rpc CreateFileDownloadRequest(FileDownloadRequest) returns (FileDownloadResponse); -} - -message ConnectRequest { - int32 protocol_version = 1; - repeated string argv = 2; - string metadata_json = 3; -} - -message ConnectResponse { - string connection_id = 1; - string status = 2; -} - -message FileUploadRequest { - string filename = 1; - string mimetype = 2; -} - -message FileUploadResponse { - string upload_url = 1; -} - -message FileMapping { - string transfer_method = 1; - optional string reference = 2; - optional string url = 3; -} - -message FileDownloadRequest { - FileMapping file = 1; - optional bool for_external = 2; -} - -message FileDownloadResponse { - string filename = 1; - optional string mime_type = 2; - int64 size = 3; - string download_url = 4; -} diff --git a/dify-agent/pyproject.toml b/dify-agent/pyproject.toml index a496d19cdd1..a4512aa6c48 100644 --- a/dify-agent/pyproject.toml +++ b/dify-agent/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "dify-agent" -version = "1.16.0" +version = "1.16.1" description = "Add your description here" readme = "README.md" requires-python = ">=3.12,<4.0" @@ -17,8 +17,8 @@ dify-agent-stub-server = "dify_agent.agent_stub.server.cli:main" [project.optional-dependencies] -grpc = ["grpclib[protobuf]>=0.4.9,<0.5.0", "protobuf>=6.33.5,<7.0.0"] server = [ + "e2b>=2.34.0,<3.0.0", "fastapi==0.136.0", "graphon==0.5.2", "jsonschema>=4.23.0,<5.0.0", @@ -47,6 +47,7 @@ extraPaths = ["src", "examples/agenton", "examples/dify_agent"] [tool.pytest.ini_options] testpaths = ["tests"] python_files = ["test_*.py", "*_test.py"] +markers = ["integration: requires a real external service or exercises multiple concrete adapters"] [tool.ruff] line-length = 120 @@ -57,7 +58,6 @@ include = ["src/**/*.py", "examples/**/*.py", "tests/**/*.py", "docs/**/*.py"] dev = [ "basedpyright>=1.39.3", "coverage[toml]>=7.10.7", - "grpcio-tools>=1.81.0,<2.0.0", "pytest>=9.0.3", "pytest-examples>=0.0.18", "pytest-mock>=3.14.0", diff --git a/dify-agent/src/dify_agent/adapters/shell/__init__.py b/dify-agent/src/dify_agent/adapters/shell/__init__.py index faeb0cfb808..c09839b8d7b 100644 --- a/dify-agent/src/dify_agent/adapters/shell/__init__.py +++ b/dify-agent/src/dify_agent/adapters/shell/__init__.py @@ -9,23 +9,12 @@ from dify_agent.adapters.shell.protocols import ( ShellCommandProtocol, ShellCommandResult, ShellCommandStatus, - ShellFileTransferProtocol, ShellPromptObservation, ShellProviderError, - ShellProviderProtocol, - ShellResourceProtocol, ) def __getattr__(name: str) -> object: - if name == "ShellAdapterSettings": - from dify_agent.adapters.shell.config import ShellAdapterSettings - - return ShellAdapterSettings - if name == "create_shell_provider": - from dify_agent.adapters.shell.factory import create_shell_provider - - return create_shell_provider if name == "shellctl": from importlib import import_module @@ -35,14 +24,9 @@ def __getattr__(name: str) -> object: __all__ = [ "CompleteShellCommandResult", - "ShellAdapterSettings", "ShellCommandProtocol", "ShellCommandResult", "ShellCommandStatus", - "ShellFileTransferProtocol", "ShellPromptObservation", "ShellProviderError", - "ShellProviderProtocol", - "ShellResourceProtocol", - "create_shell_provider", ] diff --git a/dify-agent/src/dify_agent/adapters/shell/config.py b/dify-agent/src/dify_agent/adapters/shell/config.py deleted file mode 100644 index 3d6140a7fbc..00000000000 --- a/dify-agent/src/dify_agent/adapters/shell/config.py +++ /dev/null @@ -1,61 +0,0 @@ -from typing import ClassVar, Literal, Self -from urllib.parse import urlparse - -from pydantic import Field, model_validator -from pydantic_settings import BaseSettings, SettingsConfigDict - -DEFAULT_SHELL_PROVIDER = "shellctl" - - -class ShellAdapterSettings(BaseSettings): - """Env-backed settings used to construct a shell provider. - - ``shellctl_auth_token`` defaults to ``None``; the factory forwards an empty - string to the shellctl client so it does not fall back to ambient process - credentials. Deployments that enable shellctl bearer auth must set - ``DIFY_AGENT_SHELLCTL_AUTH_TOKEN`` explicitly. - """ - - shell_provider: Literal["shellctl", "enterprise"] = "shellctl" - shellctl_entrypoint: str | None = None - shellctl_auth_token: str | None = None - - enterprise_sandbox_gateway_endpoint: str | None = None - enterprise_sandbox_gateway_auth_token: str | None = None - enterprise_sandbox_gateway_timeout: float = Field(default=30.0, gt=0) - enterprise_sandbox_proxy_timeout: float = Field(default=60.0, gt=0) - - model_config: ClassVar[SettingsConfigDict] = SettingsConfigDict( - env_prefix="DIFY_AGENT_", - env_file=(".env", "dify-agent/.env"), - extra="ignore", - populate_by_name=True, - ) - - @model_validator(mode="after") - def validate_provider_fields(self) -> Self: - match self.shell_provider: - case "shellctl": - if not self.shellctl_entrypoint or not self.shellctl_entrypoint.strip(): - raise ValueError("shellctl_entrypoint is required when shell_provider is 'shellctl'.") - _validate_url(self.shellctl_entrypoint, field_name="shellctl_entrypoint") - case "enterprise": - if not self.enterprise_sandbox_gateway_endpoint or not self.enterprise_sandbox_gateway_endpoint.strip(): - raise ValueError( - "enterprise_sandbox_gateway_endpoint is required when shell_provider is 'enterprise'." - ) - _validate_url( - self.enterprise_sandbox_gateway_endpoint, field_name="enterprise_sandbox_gateway_endpoint" - ) - return self - - -def _validate_url(value: str, *, field_name: str) -> None: - parsed = urlparse(value.strip()) - if parsed.scheme not in ("http", "https") or not parsed.netloc: - raise ValueError(f"{field_name} must be a valid http(s) URL, got: {value!r}") - - -__all__ = [ - "ShellAdapterSettings", -] diff --git a/dify-agent/src/dify_agent/adapters/shell/enterprise/__init__.py b/dify-agent/src/dify_agent/adapters/shell/enterprise/__init__.py deleted file mode 100644 index 71e609c8c54..00000000000 --- a/dify-agent/src/dify_agent/adapters/shell/enterprise/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from dify_agent.adapters.shell.enterprise.enterprise import EnterpriseShellProvider - -__all__ = ["EnterpriseShellProvider"] diff --git a/dify-agent/src/dify_agent/adapters/shell/enterprise/enterprise.py b/dify-agent/src/dify_agent/adapters/shell/enterprise/enterprise.py deleted file mode 100644 index c56b4c32058..00000000000 --- a/dify-agent/src/dify_agent/adapters/shell/enterprise/enterprise.py +++ /dev/null @@ -1,235 +0,0 @@ -from __future__ import annotations - -import logging -from dataclasses import dataclass, field -from typing import TypedDict - -import httpx2 as httpx - -from dify_agent.adapters.shell.protocols import ( - SandboxExpiredError, - ShellCommandProtocol, - ShellFileTransferProtocol, - ShellProviderError, - ShellProviderProtocol, - ShellResourceProtocol, -) -from dify_agent.adapters.shell.shellctl import ( - ShellctlClientProtocol, - ShellctlCommands, - ShellctlFileTransfer, -) - -logger = logging.getLogger(__name__) - - -class _CreateSandboxReply(TypedDict): - sandboxId: str - status: str - - -# --------------------------------------------------------------------------- -# Gateway client (control plane) -# --------------------------------------------------------------------------- - - -@dataclass(slots=True) -class EnterpriseGatewayClient: - """HTTP client for the enterprise sandbox gateway control-plane API.""" - - endpoint: str - auth_token: str - timeout: float = 30.0 - _client: httpx.AsyncClient = field(init=False) - - def __post_init__(self) -> None: - headers: dict[str, str] = {} - if self.auth_token: - headers["X-Inner-Api-Key"] = self.auth_token - self._client = httpx.AsyncClient( - base_url=self.endpoint.rstrip("/"), - headers=headers, - timeout=httpx.Timeout(self.timeout), - ) - - async def create_sandbox(self, *, tenant_id: str | None = None, template: str | None = None) -> _CreateSandboxReply: - body: dict[str, str] = {} - if tenant_id: - body["tenantId"] = tenant_id - if template: - body["template"] = template - response = await self._request("POST", "/v1/sandboxes", json=body) - return response.json() - - async def delete_sandbox(self, sandbox_id: str) -> None: - await self._request("DELETE", f"/v1/sandboxes/{sandbox_id}") - - async def close(self) -> None: - await self._client.aclose() - - async def _request(self, method: str, path: str, **kwargs: object) -> httpx.Response: - try: - response = await self._client.request(method, path, **kwargs) - response.raise_for_status() - return response - except httpx.TimeoutException as exc: - raise ShellProviderError(f"Gateway request timed out: {method} {path}", code="timeout") from exc - except httpx.HTTPStatusError as exc: - raise ShellProviderError( - f"Gateway returned {exc.response.status_code}: {exc.response.text}", - code="gateway_error", - ) from exc - except httpx.RequestError as exc: - raise ShellProviderError(f"Gateway request failed: {exc}", code="request_error") from exc - - -# --------------------------------------------------------------------------- -# Resource & Provider -# --------------------------------------------------------------------------- - - -@dataclass(slots=True) -class EnterpriseResource(ShellResourceProtocol): - """A live enterprise sandbox session. - - Holds the gateway client, sandbox ID, and the shellctl client for data-plane - operations. ``suspend()`` closes the shellctl and gateway clients without - deleting the sandbox, so the same sandbox can be re-attached later. - ``delete()`` additionally calls the gateway to destroy the sandbox pod. - """ - - _sandbox_id: str - gateway: EnterpriseGatewayClient - shellctl_client: ShellctlClientProtocol - commands: ShellCommandProtocol - files: ShellFileTransferProtocol - - @property - def sandbox_id(self) -> str | None: - return self._sandbox_id - - async def suspend(self) -> None: - try: - await self.shellctl_client.close() - except Exception as exc: - logger.warning("Failed to close shellctl client for sandbox %s: %s", self._sandbox_id, exc) - try: - await self.gateway.close() - except Exception as exc: - logger.warning("Failed to close gateway client: %s", exc) - - async def delete(self) -> None: - try: - await self.shellctl_client.close() - except Exception as exc: - logger.warning("Failed to close shellctl client for sandbox %s: %s", self._sandbox_id, exc) - try: - logger.info("Deleting enterprise sandbox via gateway: id=%s", self._sandbox_id) - await self.gateway.delete_sandbox(self._sandbox_id) - logger.info("Enterprise sandbox deleted: id=%s", self._sandbox_id) - except ShellProviderError as exc: - logger.warning("Failed to delete sandbox %s via gateway: %s", self._sandbox_id, exc) - try: - await self.gateway.close() - except Exception as exc: - logger.warning("Failed to close gateway client: %s", exc) - - -@dataclass(slots=True) -class EnterpriseShellProvider(ShellProviderProtocol): - """Provisions enterprise sandboxes via the gateway, connects via shellctl. - - Lifecycle: - 1. ``create()`` calls the gateway to provision a new sandbox pod. - 2. ``attach(sandbox_id)`` builds a shellctl client pointed at the - gateway's ``/proxy/{sandboxId}`` route for an *existing* sandbox, - without provisioning a new one. - 3. Both return an ``EnterpriseResource`` with shellctl-backed - command/file adapters. - 4. ``suspend()`` on the resource closes the shellctl and gateway - clients but leaves the sandbox pod alive. - 5. ``delete()`` on the resource closes clients and deletes the - sandbox via the gateway. - """ - - gateway_endpoint: str - auth_token: str - tenant_id: str | None = None - template: str | None = None - gateway_timeout: float = 30.0 - proxy_timeout: float = 60.0 - - async def create(self) -> EnterpriseResource: - gateway = EnterpriseGatewayClient( - endpoint=self.gateway_endpoint, - auth_token=self.auth_token, - timeout=self.gateway_timeout, - ) - try: - logger.info("Creating enterprise sandbox via gateway %s", self.gateway_endpoint) - reply = await gateway.create_sandbox(tenant_id=self.tenant_id, template=self.template) - sandbox_id = reply["sandboxId"] - logger.info("Enterprise sandbox created: id=%s status=%s", sandbox_id, reply.get("status")) - except BaseException: - await gateway.close() - raise - return self._build_resource(sandbox_id=sandbox_id, gateway=gateway) - - async def attach(self, sandbox_id: str) -> EnterpriseResource: - gateway = EnterpriseGatewayClient( - endpoint=self.gateway_endpoint, - auth_token=self.auth_token, - timeout=self.gateway_timeout, - ) - logger.info("Attaching to existing enterprise sandbox: id=%s", sandbox_id) - resource = self._build_resource(sandbox_id=sandbox_id, gateway=gateway) - # Verify the sandbox is alive by running a trivial command through the proxy. - # If the sandbox has expired, the gateway returns a 404 with "sandbox_expired"; - # if the pod is already gone the gateway returns a 404 with reason "NOT_FOUND". - # Both surface as a ShellProviderError from the shellctl client and mean the - # sandbox must be re-created. - try: - await resource.commands.run("true", timeout=5.0) - except ShellProviderError as exc: - await resource.suspend() - message = str(exc).lower() - if "expired" in message or "not_found" in message: - raise SandboxExpiredError(sandbox_id, cause=exc) from exc - raise - return resource - - def _build_resource(self, *, sandbox_id: str, gateway: EnterpriseGatewayClient) -> EnterpriseResource: - proxy_base_url = f"{self.gateway_endpoint.rstrip('/')}/proxy/" - headers: dict[str, str] = {"X-Sandbox-Id": sandbox_id} - if self.auth_token: - headers["X-Inner-Api-Key"] = self.auth_token - proxy_http_client = httpx.AsyncClient( - base_url=proxy_base_url, - headers=headers, - follow_redirects=True, - timeout=httpx.Timeout(self.proxy_timeout), - transport=httpx.AsyncHTTPTransport(retries=3), - ) - - from shellctl.client import ShellctlClient - - client: ShellctlClientProtocol = ShellctlClient( - proxy_base_url, - token=self.auth_token, - client=proxy_http_client, - ) - - return EnterpriseResource( - _sandbox_id=sandbox_id, - gateway=gateway, - shellctl_client=client, - commands=ShellctlCommands(client=client), - files=ShellctlFileTransfer(client=client), - ) - - -__all__ = [ - "EnterpriseGatewayClient", - "EnterpriseResource", - "EnterpriseShellProvider", -] diff --git a/dify-agent/src/dify_agent/adapters/shell/factory.py b/dify-agent/src/dify_agent/adapters/shell/factory.py deleted file mode 100644 index d5177d96118..00000000000 --- a/dify-agent/src/dify_agent/adapters/shell/factory.py +++ /dev/null @@ -1,31 +0,0 @@ -from dify_agent.adapters.shell.config import ShellAdapterSettings -from dify_agent.adapters.shell.enterprise import EnterpriseShellProvider -from dify_agent.adapters.shell.protocols import ShellProviderProtocol -from dify_agent.adapters.shell.shellctl import ShellctlProvider - - -def create_shell_provider(settings: ShellAdapterSettings | None = None) -> ShellProviderProtocol: - """Return the shell provider selected by ``DIFY_AGENT_SHELL_PROVIDER``.""" - resolved = settings or ShellAdapterSettings() - provider = resolved.shell_provider - match provider: - case "shellctl": - entrypoint = (resolved.shellctl_entrypoint or "").strip() - if not entrypoint: - raise ValueError("DIFY_AGENT_SHELLCTL_ENTRYPOINT is required for the 'shellctl' shell provider.") - return ShellctlProvider( - entrypoint=entrypoint, - token=resolved.shellctl_auth_token or "", - ) - case "enterprise": - return EnterpriseShellProvider( - gateway_endpoint=(resolved.enterprise_sandbox_gateway_endpoint or "").strip(), - auth_token=resolved.enterprise_sandbox_gateway_auth_token or "", - gateway_timeout=resolved.enterprise_sandbox_gateway_timeout, - proxy_timeout=resolved.enterprise_sandbox_proxy_timeout, - ) - case _: - raise ValueError(f"Unknown shell provider: {resolved.shell_provider!r}.") - - -__all__ = ["create_shell_provider"] diff --git a/dify-agent/src/dify_agent/adapters/shell/protocols.py b/dify-agent/src/dify_agent/adapters/shell/protocols.py index 6c9988f0fbd..52b5436292a 100644 --- a/dify-agent/src/dify_agent/adapters/shell/protocols.py +++ b/dify-agent/src/dify_agent/adapters/shell/protocols.py @@ -47,19 +47,12 @@ class ShellPromptObservation: class ShellProviderError(RuntimeError): code: str | None + status_code: int | None - def __init__(self, message: str, *, code: str | None = None) -> None: + def __init__(self, message: str, *, code: str | None = None, status_code: int | None = None) -> None: super().__init__(message) self.code = code - - -class SandboxExpiredError(ShellProviderError): - """Raised by a shell provider when ``attach()`` targets a sandbox that no longer exists.""" - - def __init__(self, sandbox_id: str, *, cause: ShellProviderError) -> None: - super().__init__(str(cause), code=cause.code) - self.sandbox_id = sandbox_id - self.__cause__ = cause + self.status_code = status_code class ShellCommandProtocol(Protocol): @@ -112,50 +105,3 @@ class ShellCommandProtocol(Protocol): force: bool = False, grace_seconds: float | None = None, ) -> None: ... - - -class ShellFileTransferProtocol(Protocol): - async def upload(self, *, content: bytes, remote_path: str, cwd: str | None = None) -> None: ... - - async def download(self, *, remote_path: str, cwd: str | None = None) -> bytes: ... - - -class ShellResourceProtocol(Protocol): - @property - def commands(self) -> ShellCommandProtocol: ... - - @property - def files(self) -> ShellFileTransferProtocol: ... - - @property - def sandbox_id(self) -> str | None: ... - - async def suspend(self) -> None: - """Detach from the sandbox without destroying it. - - Called when the resource scope exits with suspend intent. The sandbox - remains alive and can be re-attached later via ``attach(sandbox_id)``. - """ - ... - - async def delete(self) -> None: - """Destroy the sandbox and release all resources. - - Called when the resource scope exits with delete intent. The sandbox is - permanently removed and cannot be re-attached. - """ - ... - - -class ShellProviderProtocol(Protocol): - async def create(self) -> ShellResourceProtocol: - """Provision a new sandbox and return a live resource.""" - ... - - async def attach(self, sandbox_id: str) -> ShellResourceProtocol: - """Connect to an existing sandbox without provisioning a new one. - - The returned resource carries the same ``sandbox_id`` so the caller can - persist it across runs and re-attach as needed. - """ - ... diff --git a/dify-agent/src/dify_agent/adapters/shell/shellctl.py b/dify-agent/src/dify_agent/adapters/shell/shellctl.py index f285483bda5..cea953a4282 100644 --- a/dify-agent/src/dify_agent/adapters/shell/shellctl.py +++ b/dify-agent/src/dify_agent/adapters/shell/shellctl.py @@ -1,51 +1,34 @@ -"""Shellctl-backed shell provider adapter for dify-agent. +"""Shellctl command adapter for RuntimeLease objects. The built-in shellctl SDK owns the HTTP timeout policy for long-polling -shellctl requests. This adapter stays narrowly focused on translating SDK and -transport failures into ``ShellProviderError`` so the shell layer can return -tool observations instead of aborting the agent loop. +shellctl requests. This adapter translates SDK and transport failures into +``ShellProviderError``. """ from __future__ import annotations -import base64 -import binascii -import logging -import re -import time -from collections.abc import Awaitable -from collections.abc import Callable +import posixpath +from collections.abc import Awaitable, Callable from dataclasses import dataclass from typing import Protocol, TypeVar, cast import httpx2 as httpx +from shellctl.client import ShellctlClientError +from shellctl.shared import HealthResponse from dify_agent.adapters.shell.protocols import ( ShellCommandProtocol, ShellCommandResult, ShellCommandStatus, - ShellFileTransferProtocol, ShellProviderError, - ShellProviderProtocol, - ShellResourceProtocol, ) -logger = logging.getLogger(__name__) - ResultT = TypeVar("ResultT") _DEFAULT_TIMEOUT_SECONDS = 30.0 _READ_OUTPUT_TIMEOUT_SECONDS = 0.0 _DEFAULT_TERMINATE_GRACE_SECONDS = 10.0 -_FILE_TRANSFER_TIMEOUT_SECONDS = 60.0 _SHELLCTL_OUTPUT_LIMIT_BYTES = 16 * 1024 -_TRANSFER_BEGIN = "<<>>" -_TRANSFER_END = "<<>>" -_DOWNLOAD_MISSING_EXIT_CODE = 66 - - -class ShellFileTransferError(RuntimeError): - """Raised when a file cannot be uploaded or downloaded through shellctl.""" class ShellctlJobResult(Protocol): @@ -68,6 +51,8 @@ class ShellctlJobStatus(Protocol): class ShellctlClientProtocol(Protocol): + async def health(self) -> HealthResponse: ... + async def run( self, script: str, @@ -119,6 +104,8 @@ type ShellctlClientFactory = Callable[[], ShellctlClientProtocol] @dataclass(slots=True) class ShellctlCommands(ShellCommandProtocol): client: ShellctlClientProtocol + home_dir: str | None = None + workspace_dir: str | None = None async def run( self, @@ -128,7 +115,15 @@ class ShellctlCommands(ShellCommandProtocol): env: dict[str, str] | None = None, timeout: float, ) -> ShellCommandResult: - return _from_job_result(await _run_client_call(self.client.run(script, cwd=cwd, env=env, timeout=timeout))) + resolved_cwd = _resolve_lease_cwd( + cwd, + home_dir=self.home_dir, + workspace_dir=self.workspace_dir, + ) + resolved_env = _lease_env(env, home_dir=self.home_dir) + return _from_job_result( + await _run_client_call(self.client.run(script, cwd=resolved_cwd, env=resolved_env, timeout=timeout)) + ) async def wait( self, @@ -185,119 +180,6 @@ class ShellctlCommands(ShellCommandProtocol): raise -@dataclass(slots=True) -class ShellctlFileTransfer(ShellFileTransferProtocol): - client: ShellctlClientProtocol - timeout: float = _FILE_TRANSFER_TIMEOUT_SECONDS - - async def upload(self, *, content: bytes, remote_path: str, cwd: str | None = None) -> None: - encoded = base64.b64encode(content).decode("ascii") - completed = await _run_to_completion( - self.client, - _upload_script(remote_path=remote_path, encoded=encoded), - cwd=cwd, - timeout=self.timeout, - ) - if completed.exit_code != 0: - raise ShellFileTransferError( - f"Failed to upload to {remote_path!r}: exit_code={completed.exit_code}, " - f"output={_output_tail(completed.output)!r}" - ) - - async def download(self, *, remote_path: str, cwd: str | None = None) -> bytes: - completed = await _run_to_completion( - self.client, - _download_script(remote_path=remote_path), - cwd=cwd, - timeout=self.timeout, - ) - if completed.exit_code == _DOWNLOAD_MISSING_EXIT_CODE: - raise ShellFileTransferError(f"Remote path not found: {remote_path!r}.") - if completed.exit_code != 0: - raise ShellFileTransferError( - f"Failed to download {remote_path!r}: exit_code={completed.exit_code}, " - f"output={_output_tail(completed.output)!r}" - ) - encoded = _extract_transfer_payload(completed.output) - try: - return base64.b64decode(encoded.encode("ascii"), validate=True) - except (ValueError, binascii.Error) as exc: - raise ShellFileTransferError(f"Downloaded payload for {remote_path!r} was not valid base64.") from exc - - -@dataclass(slots=True) -class ShellctlResource(ShellResourceProtocol): - """Live shellctl connection. - - For shellctl there is no separate sandbox lifecycle: the shellctl server is a - long-running process that persists across runs. The ``sandbox_id`` is the - shellctl entrypoint URL, used as a stable identifier so the layer can - distinguish ``create()`` (first run) from ``attach()`` (subsequent runs). - Both ``suspend()`` and ``delete()`` simply close the HTTP client; the server - and its filesystem remain intact either way. - """ - - client: ShellctlClientProtocol - _commands: ShellCommandProtocol - _files: ShellFileTransferProtocol - _sandbox_id: str | None = None - - @property - def commands(self) -> ShellCommandProtocol: - return self._commands - - @property - def files(self) -> ShellFileTransferProtocol: - return self._files - - @property - def sandbox_id(self) -> str | None: - return self._sandbox_id - - async def suspend(self) -> None: - await self._close_client() - - async def delete(self) -> None: - await self._close_client() - - async def _close_client(self) -> None: - try: - await self.client.close() - except RuntimeError as exc: - raise _map_error(exc) from exc - - -@dataclass(slots=True) -class ShellctlProvider(ShellProviderProtocol): - entrypoint: str - token: str - output_limit: int = _SHELLCTL_OUTPUT_LIMIT_BYTES - client_factory: ShellctlClientFactory | None = None - - async def create(self) -> ShellctlResource: - return self._build_resource(sandbox_id=self.entrypoint) - - async def attach(self, sandbox_id: str) -> ShellctlResource: - return self._build_resource(sandbox_id=sandbox_id) - - def _build_resource(self, *, sandbox_id: str | None) -> ShellctlResource: - client = ( - self.client_factory() - if self.client_factory is not None - else create_default_shellctl_client_factory( - entrypoint=self.entrypoint, - token=self.token, - output_limit=self.output_limit, - )() - ) - return ShellctlResource( - client=client, - _commands=ShellctlCommands(client=client), - _files=ShellctlFileTransfer(client=client), - _sandbox_id=sandbox_id, - ) - - def create_default_shellctl_client_factory( *, entrypoint: str, @@ -318,13 +200,6 @@ def create_default_shellctl_client_factory( return factory -@dataclass(frozen=True, slots=True) -class _CompletedShellctlJob: - job_id: str - exit_code: int | None - output: str - - async def _run_client_call(awaitable: Awaitable[ResultT]) -> ResultT: """Map shellctl client boundary failures into provider-layer errors.""" @@ -334,12 +209,16 @@ async def _run_client_call(awaitable: Awaitable[ResultT]) -> ResultT: raise ShellProviderError(str(exc), code="timeout") from exc except httpx.RequestError as exc: raise ShellProviderError(str(exc), code="request_error") from exc - except RuntimeError as exc: + except ShellctlClientError as exc: raise _map_error(exc) from exc -def _map_error(exc: RuntimeError) -> ShellProviderError: - return ShellProviderError(str(exc), code=getattr(exc, "code", None)) +def _map_error(exc: ShellctlClientError) -> ShellProviderError: + return ShellProviderError( + str(exc), + code=exc.code, + status_code=exc.status_code, + ) def _from_job_result(result: ShellctlJobResult) -> ShellCommandResult: @@ -374,84 +253,42 @@ def _status_name(status: object) -> str: return str(status) -async def _run_to_completion( - client: ShellctlClientProtocol, - script: str, - *, +def _lease_env(env: dict[str, str] | None, *, home_dir: str | None) -> dict[str, str] | None: + if home_dir is None: + return env + resolved = dict(env or {}) + resolved["HOME"] = home_dir + return resolved + + +def _resolve_lease_cwd( cwd: str | None, - timeout: float, -) -> _CompletedShellctlJob: - deadline = time.monotonic() + timeout - job_id: str | None = None - try: - result = await _run_client_call(client.run(script, cwd=cwd, env=None, timeout=_remaining_timeout(deadline))) - parts = [result.output] - job_id = result.job_id - while not result.done or result.truncated: - result = await _run_client_call( - client.wait(job_id, offset=result.offset, timeout=_remaining_timeout(deadline)) - ) - parts.append(result.output) - return _CompletedShellctlJob(job_id=job_id, exit_code=result.exit_code, output="".join(parts)) - finally: - if job_id is not None: - try: - await _run_client_call(client.delete(job_id, force=True)) - except RuntimeError as exc: - logger.warning("Failed to delete shellctl job %s: %s", job_id, exc) - - -def _upload_script(*, remote_path: str, encoded: str) -> str: - return ( - "set -eu\n" - f'mkdir -p "$(dirname -- {_shquote(remote_path)})"\n' - f"printf %s {_shquote(encoded)} | base64 -d > {_shquote(remote_path)}" - ) - - -def _download_script(*, remote_path: str) -> str: - return "\n".join( - [ - "set -eu", - f"path={_shquote(remote_path)}", - 'if [ ! -f "$path" ]; then exit 66; fi', - f"printf %s {_shquote(_TRANSFER_BEGIN)}", - 'base64 < "$path" | tr -d "\\n"', - f"printf %s {_shquote(_TRANSFER_END)}", - ] - ) - - -def _extract_transfer_payload(output: str) -> str: - pattern = re.escape(_TRANSFER_BEGIN) + r"(.*?)" + re.escape(_TRANSFER_END) - match = re.search(pattern, output, re.DOTALL) - if match is None: - raise ShellFileTransferError("Transfer payload markers were missing from shell output.") - return "".join(match.group(1).split()) - - -def _output_tail(output: str, *, limit: int = 256) -> str: - return output[-limit:] - - -def _remaining_timeout(deadline: float) -> float: - remaining = deadline - time.monotonic() - if remaining <= 0.0: - raise ShellProviderError("Shellctl command timed out before completion.", code="timeout") - return remaining - - -def _shquote(value: str) -> str: - return "'" + value.replace("'", "'\\''") + "'" + *, + home_dir: str | None, + workspace_dir: str | None, +) -> str | None: + if home_dir is None or workspace_dir is None: + return cwd + if cwd is None or cwd == "": + candidate = workspace_dir + elif cwd == "~": + candidate = home_dir + elif cwd.startswith("~/"): + candidate = posixpath.join(home_dir, cwd[2:]) + elif posixpath.isabs(cwd): + candidate = cwd + else: + candidate = posixpath.join(workspace_dir, cwd) + candidate = posixpath.normpath(candidate) + roots = (posixpath.normpath(home_dir), posixpath.normpath(workspace_dir)) + if not any(posixpath.commonpath((candidate, root)) == root for root in roots): + raise ValueError("shell cwd is outside this RuntimeLease Home and Workspace") + return candidate __all__ = [ - "ShellFileTransferError", "ShellctlClientFactory", "ShellctlClientProtocol", "ShellctlCommands", - "ShellctlFileTransfer", - "ShellctlProvider", - "ShellctlResource", "create_default_shellctl_client_factory", ] diff --git a/dify-agent/src/dify_agent/agent_stub/grpc/__init__.py b/dify-agent/src/dify_agent/agent_stub/grpc/__init__.py deleted file mode 100644 index d6c26981fe1..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/grpc/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -"""Internal gRPC helpers for the optional Agent Stub transport.""" - -__all__: list[str] = [] diff --git a/dify-agent/src/dify_agent/agent_stub/grpc/_generated/__init__.py b/dify-agent/src/dify_agent/agent_stub/grpc/_generated/__init__.py deleted file mode 100644 index 60221b58eaf..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/grpc/_generated/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -"""Generated protobuf and grpclib code for the Agent Stub gRPC contract.""" - -__all__: list[str] = [] diff --git a/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_grpc.py b/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_grpc.py deleted file mode 100644 index e8850813cc0..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_grpc.py +++ /dev/null @@ -1,80 +0,0 @@ -# Generated by the Protocol Buffers compiler. DO NOT EDIT! -# source: dify/agent/stub/v1/agent_stub.proto -# plugin: grpclib.plugin.main -import abc -import typing - -import grpclib.client -import grpclib.const - -from . import agent_stub_pb2 - -if typing.TYPE_CHECKING: - import grpclib.server - - -class AgentStubServiceBase(abc.ABC): - @abc.abstractmethod - async def Connect( - self, - stream: "grpclib.server.Stream[agent_stub_pb2.ConnectRequest, agent_stub_pb2.ConnectResponse]", - ) -> None: - pass - - @abc.abstractmethod - async def CreateFileUploadRequest( - self, - stream: "grpclib.server.Stream[agent_stub_pb2.FileUploadRequest, agent_stub_pb2.FileUploadResponse]", - ) -> None: - pass - - @abc.abstractmethod - async def CreateFileDownloadRequest( - self, - stream: "grpclib.server.Stream[agent_stub_pb2.FileDownloadRequest, agent_stub_pb2.FileDownloadResponse]", - ) -> None: - pass - - def __mapping__(self) -> typing.Dict[str, grpclib.const.Handler]: - return { - "/dify.agent.stub.v1.AgentStubService/Connect": grpclib.const.Handler( - self.Connect, - grpclib.const.Cardinality.UNARY_UNARY, - agent_stub_pb2.ConnectRequest, - agent_stub_pb2.ConnectResponse, - ), - "/dify.agent.stub.v1.AgentStubService/CreateFileUploadRequest": grpclib.const.Handler( - self.CreateFileUploadRequest, - grpclib.const.Cardinality.UNARY_UNARY, - agent_stub_pb2.FileUploadRequest, - agent_stub_pb2.FileUploadResponse, - ), - "/dify.agent.stub.v1.AgentStubService/CreateFileDownloadRequest": grpclib.const.Handler( - self.CreateFileDownloadRequest, - grpclib.const.Cardinality.UNARY_UNARY, - agent_stub_pb2.FileDownloadRequest, - agent_stub_pb2.FileDownloadResponse, - ), - } - - -class AgentStubServiceStub: - def __init__(self, channel: grpclib.client.Channel) -> None: - self.Connect = grpclib.client.UnaryUnaryMethod( - channel, - "/dify.agent.stub.v1.AgentStubService/Connect", - agent_stub_pb2.ConnectRequest, - agent_stub_pb2.ConnectResponse, - ) - self.CreateFileUploadRequest = grpclib.client.UnaryUnaryMethod( - channel, - "/dify.agent.stub.v1.AgentStubService/CreateFileUploadRequest", - agent_stub_pb2.FileUploadRequest, - agent_stub_pb2.FileUploadResponse, - ) - self.CreateFileDownloadRequest = grpclib.client.UnaryUnaryMethod( - channel, - "/dify.agent.stub.v1.AgentStubService/CreateFileDownloadRequest", - agent_stub_pb2.FileDownloadRequest, - agent_stub_pb2.FileDownloadResponse, - ) diff --git a/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_pb2.py b/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_pb2.py deleted file mode 100644 index db940989607..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_pb2.py +++ /dev/null @@ -1,47 +0,0 @@ -# -*- coding: utf-8 -*- -# Generated by the protocol buffer compiler. DO NOT EDIT! -# NO CHECKED-IN PROTOBUF GENCODE -# source: dify/agent/stub/v1/agent_stub.proto -# Protobuf Python Version: 6.33.5 -"""Generated protocol buffer code.""" - -from google.protobuf import descriptor as _descriptor -from google.protobuf import descriptor_pool as _descriptor_pool -from google.protobuf import runtime_version as _runtime_version -from google.protobuf import symbol_database as _symbol_database -from google.protobuf.internal import builder as _builder - -_runtime_version.ValidateProtobufRuntimeVersion( - _runtime_version.Domain.PUBLIC, 6, 33, 5, "", "dify/agent/stub/v1/agent_stub.proto" -) -# @@protoc_insertion_point(imports) - -_sym_db = _symbol_database.Default() - - -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( - b'\n#dify/agent/stub/v1/agent_stub.proto\x12\x12\x64ify.agent.stub.v1"O\n\x0e\x43onnectRequest\x12\x18\n\x10protocol_version\x18\x01 \x01(\x05\x12\x0c\n\x04\x61rgv\x18\x02 \x03(\t\x12\x15\n\rmetadata_json\x18\x03 \x01(\t"8\n\x0f\x43onnectResponse\x12\x15\n\rconnection_id\x18\x01 \x01(\t\x12\x0e\n\x06status\x18\x02 \x01(\t"7\n\x11\x46ileUploadRequest\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x10\n\x08mimetype\x18\x02 \x01(\t"(\n\x12\x46ileUploadResponse\x12\x12\n\nupload_url\x18\x01 \x01(\t"f\n\x0b\x46ileMapping\x12\x17\n\x0ftransfer_method\x18\x01 \x01(\t\x12\x16\n\treference\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x10\n\x03url\x18\x03 \x01(\tH\x01\x88\x01\x01\x42\x0c\n\n_referenceB\x06\n\x04_url"p\n\x13\x46ileDownloadRequest\x12-\n\x04\x66ile\x18\x01 \x01(\x0b\x32\x1f.dify.agent.stub.v1.FileMapping\x12\x19\n\x0c\x66or_external\x18\x02 \x01(\x08H\x00\x88\x01\x01\x42\x0f\n\r_for_external"r\n\x14\x46ileDownloadResponse\x12\x10\n\x08\x66ilename\x18\x01 \x01(\t\x12\x16\n\tmime_type\x18\x02 \x01(\tH\x00\x88\x01\x01\x12\x0c\n\x04size\x18\x03 \x01(\x03\x12\x14\n\x0c\x64ownload_url\x18\x04 \x01(\tB\x0c\n\n_mime_type2\xc0\x02\n\x10\x41gentStubService\x12R\n\x07\x43onnect\x12".dify.agent.stub.v1.ConnectRequest\x1a#.dify.agent.stub.v1.ConnectResponse\x12h\n\x17\x43reateFileUploadRequest\x12%.dify.agent.stub.v1.FileUploadRequest\x1a&.dify.agent.stub.v1.FileUploadResponse\x12n\n\x19\x43reateFileDownloadRequest\x12\'.dify.agent.stub.v1.FileDownloadRequest\x1a(.dify.agent.stub.v1.FileDownloadResponseb\x06proto3' -) - -_globals = globals() -_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) -_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, "dify.agent.stub.v1.agent_stub_pb2", _globals) -if not _descriptor._USE_C_DESCRIPTORS: - DESCRIPTOR._loaded_options = None - _globals["_CONNECTREQUEST"]._serialized_start = 59 - _globals["_CONNECTREQUEST"]._serialized_end = 138 - _globals["_CONNECTRESPONSE"]._serialized_start = 140 - _globals["_CONNECTRESPONSE"]._serialized_end = 196 - _globals["_FILEUPLOADREQUEST"]._serialized_start = 198 - _globals["_FILEUPLOADREQUEST"]._serialized_end = 253 - _globals["_FILEUPLOADRESPONSE"]._serialized_start = 255 - _globals["_FILEUPLOADRESPONSE"]._serialized_end = 295 - _globals["_FILEMAPPING"]._serialized_start = 297 - _globals["_FILEMAPPING"]._serialized_end = 399 - _globals["_FILEDOWNLOADREQUEST"]._serialized_start = 401 - _globals["_FILEDOWNLOADREQUEST"]._serialized_end = 513 - _globals["_FILEDOWNLOADRESPONSE"]._serialized_start = 515 - _globals["_FILEDOWNLOADRESPONSE"]._serialized_end = 629 - _globals["_AGENTSTUBSERVICE"]._serialized_start = 632 - _globals["_AGENTSTUBSERVICE"]._serialized_end = 952 -# @@protoc_insertion_point(module_scope) diff --git a/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_pb2.pyi b/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_pb2.pyi deleted file mode 100644 index 617956bd845..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/grpc/_generated/agent_stub_pb2.pyi +++ /dev/null @@ -1,69 +0,0 @@ -from google.protobuf.internal import containers as _containers -from google.protobuf import descriptor as _descriptor -from google.protobuf import message as _message -from collections.abc import Iterable as _Iterable, Mapping as _Mapping -from typing import ClassVar as _ClassVar, Optional as _Optional, Union as _Union - -DESCRIPTOR: _descriptor.FileDescriptor - -class ConnectRequest(_message.Message): - __slots__ = ("protocol_version", "argv", "metadata_json") - PROTOCOL_VERSION_FIELD_NUMBER: _ClassVar[int] - ARGV_FIELD_NUMBER: _ClassVar[int] - METADATA_JSON_FIELD_NUMBER: _ClassVar[int] - protocol_version: int - argv: _containers.RepeatedScalarFieldContainer[str] - metadata_json: str - def __init__(self, protocol_version: _Optional[int] = ..., argv: _Optional[_Iterable[str]] = ..., metadata_json: _Optional[str] = ...) -> None: ... - -class ConnectResponse(_message.Message): - __slots__ = ("connection_id", "status") - CONNECTION_ID_FIELD_NUMBER: _ClassVar[int] - STATUS_FIELD_NUMBER: _ClassVar[int] - connection_id: str - status: str - def __init__(self, connection_id: _Optional[str] = ..., status: _Optional[str] = ...) -> None: ... - -class FileUploadRequest(_message.Message): - __slots__ = ("filename", "mimetype") - FILENAME_FIELD_NUMBER: _ClassVar[int] - MIMETYPE_FIELD_NUMBER: _ClassVar[int] - filename: str - mimetype: str - def __init__(self, filename: _Optional[str] = ..., mimetype: _Optional[str] = ...) -> None: ... - -class FileUploadResponse(_message.Message): - __slots__ = ("upload_url",) - UPLOAD_URL_FIELD_NUMBER: _ClassVar[int] - upload_url: str - def __init__(self, upload_url: _Optional[str] = ...) -> None: ... - -class FileMapping(_message.Message): - __slots__ = ("transfer_method", "reference", "url") - TRANSFER_METHOD_FIELD_NUMBER: _ClassVar[int] - REFERENCE_FIELD_NUMBER: _ClassVar[int] - URL_FIELD_NUMBER: _ClassVar[int] - transfer_method: str - reference: str - url: str - def __init__(self, transfer_method: _Optional[str] = ..., reference: _Optional[str] = ..., url: _Optional[str] = ...) -> None: ... - -class FileDownloadRequest(_message.Message): - __slots__ = ("file", "for_external") - FILE_FIELD_NUMBER: _ClassVar[int] - FOR_EXTERNAL_FIELD_NUMBER: _ClassVar[int] - file: FileMapping - for_external: bool - def __init__(self, file: _Optional[_Union[FileMapping, _Mapping]] = ..., for_external: _Optional[bool] = ...) -> None: ... - -class FileDownloadResponse(_message.Message): - __slots__ = ("filename", "mime_type", "size", "download_url") - FILENAME_FIELD_NUMBER: _ClassVar[int] - MIME_TYPE_FIELD_NUMBER: _ClassVar[int] - SIZE_FIELD_NUMBER: _ClassVar[int] - DOWNLOAD_URL_FIELD_NUMBER: _ClassVar[int] - filename: str - mime_type: str - size: int - download_url: str - def __init__(self, filename: _Optional[str] = ..., mime_type: _Optional[str] = ..., size: _Optional[int] = ..., download_url: _Optional[str] = ...) -> None: ... diff --git a/dify-agent/src/dify_agent/agent_stub/grpc/conversions.py b/dify-agent/src/dify_agent/agent_stub/grpc/conversions.py deleted file mode 100644 index a79e446c85e..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/grpc/conversions.py +++ /dev/null @@ -1,174 +0,0 @@ -"""Conversions between Agent Stub protobuf messages and public DTOs.""" - -from __future__ import annotations - -import json -from collections.abc import Mapping -from typing import TYPE_CHECKING - -from pydantic import JsonValue - -from dify_agent.agent_stub.protocol.agent_stub import ( - AgentStubConnectRequest, - AgentStubConnectResponse, - AgentStubFileDownloadRequest, - AgentStubFileDownloadResponse, - AgentStubFileMapping, - AgentStubFileUploadRequest, - AgentStubFileUploadResponse, -) - -if TYPE_CHECKING: - from dify_agent.agent_stub.grpc._generated import agent_stub_pb2 - - -def connect_request_from_proto(message: agent_stub_pb2.ConnectRequest) -> AgentStubConnectRequest: - """Validate one protobuf connect request into the public DTO.""" - metadata: object = {} - if message.metadata_json: - metadata = json.loads(message.metadata_json) - return AgentStubConnectRequest.model_validate( - { - "protocol_version": message.protocol_version, - "argv": list(message.argv), - "metadata": metadata, - } - ) - - -def proto_connect_request( - pb2_module, - *, - argv: list[str], - metadata: Mapping[str, JsonValue] | None, -) -> agent_stub_pb2.ConnectRequest: - """Build one protobuf connect request from public client inputs.""" - request = AgentStubConnectRequest(argv=argv, metadata=dict(metadata or {})) - return pb2_module.ConnectRequest( - protocol_version=request.protocol_version, - argv=request.argv, - metadata_json=json.dumps(request.metadata, separators=(",", ":")), - ) - - -def connect_response_from_proto(message: agent_stub_pb2.ConnectResponse) -> AgentStubConnectResponse: - """Validate one protobuf connect response into the public DTO.""" - return AgentStubConnectResponse.model_validate( - { - "connection_id": message.connection_id, - "status": message.status, - } - ) - - -def proto_connect_response(response: AgentStubConnectResponse, *, pb2_module=None) -> agent_stub_pb2.ConnectResponse: - """Build one protobuf connect response from the public DTO.""" - resolved_pb2 = pb2_module or _require_pb2_module() - return resolved_pb2.ConnectResponse(connection_id=response.connection_id, status=response.status) - - -def file_upload_request_from_proto(message: agent_stub_pb2.FileUploadRequest) -> AgentStubFileUploadRequest: - """Validate one protobuf file-upload request into the public DTO.""" - return AgentStubFileUploadRequest.model_validate({"filename": message.filename, "mimetype": message.mimetype}) - - -def proto_file_upload_request(pb2_module, *, filename: str, mimetype: str) -> agent_stub_pb2.FileUploadRequest: - """Build one protobuf file-upload request from public client inputs.""" - request = AgentStubFileUploadRequest(filename=filename, mimetype=mimetype) - return pb2_module.FileUploadRequest(filename=request.filename, mimetype=request.mimetype) - - -def file_upload_response_from_proto(message: agent_stub_pb2.FileUploadResponse) -> AgentStubFileUploadResponse: - """Validate one protobuf file-upload response into the public DTO.""" - return AgentStubFileUploadResponse.model_validate({"upload_url": message.upload_url}) - - -def proto_file_upload_response( - response: AgentStubFileUploadResponse, *, pb2_module=None -) -> agent_stub_pb2.FileUploadResponse: - """Build one protobuf file-upload response from the public DTO.""" - resolved_pb2 = pb2_module or _require_pb2_module() - return resolved_pb2.FileUploadResponse(upload_url=response.upload_url) - - -def file_download_request_from_proto(message: agent_stub_pb2.FileDownloadRequest) -> AgentStubFileDownloadRequest: - """Validate one protobuf file-download request into the public DTO.""" - file_mapping_kwargs = { - "transfer_method": message.file.transfer_method, - "reference": message.file.reference if message.file.HasField("reference") else None, - "url": message.file.url if message.file.HasField("url") else None, - } - return AgentStubFileDownloadRequest.model_validate( - { - "file": file_mapping_kwargs, - "for_external": message.for_external if message.HasField("for_external") else True, - } - ) - - -def proto_file_download_request( - pb2_module, - *, - file: AgentStubFileMapping, - for_external: bool = True, -) -> agent_stub_pb2.FileDownloadRequest: - """Build one protobuf file-download request from the public DTO.""" - mapping = pb2_module.FileMapping(transfer_method=file.transfer_method) - if file.reference is not None: - mapping.reference = file.reference - if file.url is not None: - mapping.url = file.url - request = pb2_module.FileDownloadRequest(file=mapping) - request.for_external = for_external - return request - - -def file_download_response_from_proto(message: agent_stub_pb2.FileDownloadResponse) -> AgentStubFileDownloadResponse: - """Validate one protobuf file-download response into the public DTO.""" - return AgentStubFileDownloadResponse.model_validate( - { - "filename": message.filename, - "mime_type": message.mime_type if message.HasField("mime_type") else None, - "size": message.size, - "download_url": message.download_url, - } - ) - - -def proto_file_download_response( - response: AgentStubFileDownloadResponse, - *, - pb2_module=None, -) -> agent_stub_pb2.FileDownloadResponse: - """Build one protobuf file-download response from the public DTO.""" - resolved_pb2 = pb2_module or _require_pb2_module() - message = resolved_pb2.FileDownloadResponse( - filename=response.filename, - size=response.size, - download_url=response.download_url, - ) - if response.mime_type is not None: - message.mime_type = response.mime_type - return message - - -def _require_pb2_module(): - from dify_agent.agent_stub.grpc._generated import agent_stub_pb2 - - return agent_stub_pb2 - - -__all__ = [ - "connect_request_from_proto", - "connect_response_from_proto", - "file_download_request_from_proto", - "file_download_response_from_proto", - "file_upload_request_from_proto", - "file_upload_response_from_proto", - "proto_connect_request", - "proto_connect_response", - "proto_file_download_request", - "proto_file_download_response", - "proto_file_upload_request", - "proto_file_upload_response", -] diff --git a/dify-agent/src/dify_agent/agent_stub/protocol/__init__.py b/dify-agent/src/dify_agent/agent_stub/protocol/__init__.py index 6aedf24b08b..9bd7d1f994a 100644 --- a/dify-agent/src/dify_agent/agent_stub/protocol/__init__.py +++ b/dify-agent/src/dify_agent/agent_stub/protocol/__init__.py @@ -8,6 +8,7 @@ from .agent_stub import ( AGENT_STUB_API_BASE_URL_ENV_VAR, AgentStubConnectRequest, AgentStubConnectResponse, + AgentStubConfigDownloadSource, AgentStubConfigEnvUpdateRequest, AgentStubConfigFileItem, AgentStubConfigFileRef, @@ -33,12 +34,10 @@ from .agent_stub import ( AgentStubFileUploadResponse, AgentStubURLScheme, agent_stub_config_env_url, - agent_stub_config_file_pull_url, agent_stub_config_manifest_url, agent_stub_config_note_url, agent_stub_config_push_url, agent_stub_config_skill_inspect_url, - agent_stub_config_skill_pull_url, agent_stub_connections_url, agent_stub_drive_base_for_ref, agent_stub_drive_commit_url, @@ -58,6 +57,7 @@ __all__ = [ "DEFAULT_AGENT_STUB_DRIVE_BASE", "AgentStubConnectRequest", "AgentStubConnectResponse", + "AgentStubConfigDownloadSource", "AgentStubConfigEnvUpdateRequest", "AgentStubConfigFileItem", "AgentStubConfigFileRef", @@ -83,12 +83,10 @@ __all__ = [ "AgentStubFileUploadResponse", "AgentStubURLScheme", "agent_stub_config_env_url", - "agent_stub_config_file_pull_url", "agent_stub_config_manifest_url", "agent_stub_config_note_url", "agent_stub_config_push_url", "agent_stub_config_skill_inspect_url", - "agent_stub_config_skill_pull_url", "agent_stub_connections_url", "agent_stub_drive_base_for_ref", "agent_stub_drive_commit_url", diff --git a/dify-agent/src/dify_agent/agent_stub/protocol/agent_stub.py b/dify-agent/src/dify_agent/agent_stub/protocol/agent_stub.py index 734ca50d53e..3b0c8c419bb 100644 --- a/dify-agent/src/dify_agent/agent_stub/protocol/agent_stub.py +++ b/dify-agent/src/dify_agent/agent_stub/protocol/agent_stub.py @@ -1,7 +1,7 @@ -"""Client-safe DTOs and endpoint parsing for the Agent Stub protocol. +"""Client-safe DTOs and endpoint parsing for the Agent Stub HTTP protocol. -The Agent Stub contract is shared by the HTTP router, optional gRPC transport, -the sandbox-visible CLI, and tests. Control-plane requests always validate into +The Agent Stub contract is shared by the HTTP router, the sandbox-visible CLI, +and tests. Control-plane requests always validate into these Pydantic DTOs before business logic runs, while token issuance and JWE validation stay under ``dify_agent.agent_stub.server.tokens.agent_stub`` so the default package remains free of server-only crypto dependencies. @@ -11,11 +11,12 @@ from __future__ import annotations import base64 import json +import re from dataclasses import dataclass from typing import ClassVar, Final, Literal from urllib.parse import urlsplit, urlunsplit -from pydantic import BaseModel, ConfigDict, Field, JsonValue, model_validator +from pydantic import AliasChoices, BaseModel, ConfigDict, Field, JsonValue, model_validator from dify_agent.agent_stub._constants import AGENT_STUB_DRIVE_BASE_ENV_VAR, DEFAULT_AGENT_STUB_DRIVE_BASE @@ -24,7 +25,7 @@ AGENT_STUB_PROTOCOL_VERSION: Final[int] = 1 AGENT_STUB_API_BASE_URL_ENV_VAR: Final[str] = "DIFY_AGENT_STUB_API_BASE_URL" AGENT_STUB_AUTH_JWE_ENV_VAR: Final[str] = "DIFY_AGENT_STUB_AUTH_JWE" -type AgentStubURLScheme = Literal["http", "https", "grpc"] +type AgentStubURLScheme = Literal["http", "https"] @dataclass(frozen=True, slots=True) @@ -37,14 +38,6 @@ class AgentStubEndpoint: port: int | None path: str - @property - def is_http(self) -> bool: - return self.scheme in {"http", "https"} - - @property - def is_grpc(self) -> bool: - return self.scheme == "grpc" - def agent_stub_drive_base_for_ref(drive_ref: str | None) -> str: """Return the fixed sandbox-local Agent Stub drive base for one drive ref.""" @@ -58,20 +51,13 @@ def agent_stub_drive_base_for_ref(drive_ref: str | None) -> str: def parse_agent_stub_endpoint(url: str) -> AgentStubEndpoint: - """Parse one Agent Stub endpoint URL for HTTP or gRPC transport selection. - - HTTP(S) endpoints accept either the service root or the explicit - ``/agent-stub`` API root and normalize to the latter. gRPC endpoints must be - plain ``grpc://host:port`` targets with no path, query string, or fragment - because transport routing happens on the gRPC service name instead of an - HTTP URL path. - """ + """Parse an HTTP(S) Agent Stub endpoint and normalize its API root.""" stripped = url.strip() if not stripped: raise ValueError("Agent Stub URL must not be empty") parsed = urlsplit(stripped) - if parsed.scheme not in {"http", "https", "grpc"}: - raise ValueError("Agent Stub URL must use http, https, or grpc") + if parsed.scheme not in {"http", "https"}: + raise ValueError("Agent Stub URL must use http or https") if not parsed.netloc: raise ValueError("Agent Stub URL must include a host") if parsed.username is not None or parsed.password is not None: @@ -82,21 +68,6 @@ def parse_agent_stub_endpoint(url: str) -> AgentStubEndpoint: raise ValueError("Agent Stub URL must include a host") scheme = parsed.scheme - if scheme == "grpc": - if parsed.path not in {"", "/"}: - raise ValueError("gRPC Agent Stub URL must not include a path") - if parsed.port is None: - raise ValueError("gRPC Agent Stub URL must include an explicit port") - host = parsed.hostname - normalized_url = f"grpc://{_format_url_host(host)}:{parsed.port}" - return AgentStubEndpoint( - url=normalized_url, - scheme="grpc", - host=host, - port=parsed.port, - path="", - ) - normalized_path = parsed.path.rstrip("/") if normalized_path in {"", "/"}: normalized_path = "/agent-stub" @@ -147,18 +118,10 @@ def agent_stub_config_manifest_url(base_url: str) -> str: return f"{_require_http_base_url(base_url)}/config/manifest" -def agent_stub_config_skill_pull_url(base_url: str, name: str) -> str: - return f"{_require_http_base_url(base_url)}/config/skills/{name}/pull" - - def agent_stub_config_skill_inspect_url(base_url: str, name: str) -> str: return f"{_require_http_base_url(base_url)}/config/skills/{name}/inspect" -def agent_stub_config_file_pull_url(base_url: str, name: str) -> str: - return f"{_require_http_base_url(base_url)}/config/files/{name}/pull" - - def agent_stub_config_push_url(base_url: str) -> str: return f"{_require_http_base_url(base_url)}/config/push" @@ -245,14 +208,58 @@ class AgentStubFileMapping(BaseModel): return self -class AgentStubFileDownloadRequest(BaseModel): - """Request body for one signed download URL allocation.""" +class AgentStubConfigDownloadSource(BaseModel): + """Config asset selected by name within the authenticated Config target.""" - file: AgentStubFileMapping - for_external: bool = True + kind: Literal["file", "skill"] + name: str model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + @model_validator(mode="after") + def validate_name(self) -> "AgentStubConfigDownloadSource": + normalized = self.name.strip() + if ( + not normalized + or normalized in {".", ".."} + or "/" in normalized + or "\\" in normalized + or "\x00" in normalized + or any(ord(char) < 0x20 for char in normalized) + ): + raise ValueError("config asset name must be a safe path segment") + if self.kind == "skill" and re.fullmatch(r"[a-z0-9][a-z0-9_-]{0,63}", normalized) is None: + raise ValueError("config skill name is invalid") + self.name = normalized + return self + + +class AgentStubFileDownloadRequest(BaseModel): + """Request one file URL for a specific consumer audience. + + ``for_frontend=True`` allocates a frontend-display URL that the CLI only + returns to its caller. ``False`` allocates a Sandbox byte-transfer URL that + the CLI immediately fetches. The deprecated HTTP input name + ``for_external`` remains accepted for one compatibility cycle. + """ + + file: AgentStubFileMapping | None = None + config: AgentStubConfigDownloadSource | None = None + for_frontend: bool = Field( + default=True, + validation_alias=AliasChoices("for_frontend", "for_external"), + ) + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + @model_validator(mode="after") + def validate_source(self) -> "AgentStubFileDownloadRequest": + if (self.file is None) == (self.config is None): + raise ValueError("exactly one of file or config is required") + if self.config is not None and self.for_frontend: + raise ValueError("config downloads are available only to the Sandbox data plane") + return self + class AgentStubFileDownloadResponse(BaseModel): """Response body containing download metadata plus the signed URL.""" @@ -412,14 +419,7 @@ class AgentStubConfigNoteUpdateRequest(BaseModel): def _require_http_base_url(base_url: str) -> str: - endpoint = parse_agent_stub_endpoint(base_url) - if not endpoint.is_http: - raise ValueError("HTTP Agent Stub URLs must use http or https") - return endpoint.url - - -def _format_url_host(host: str) -> str: - return f"[{host}]" if ":" in host and not host.startswith("[") else host + return parse_agent_stub_endpoint(base_url).url __all__ = [ @@ -432,6 +432,7 @@ __all__ = [ "AgentStubConnectResponse", "AgentStubEndpoint", "AgentStubConfigEnvUpdateRequest", + "AgentStubConfigDownloadSource", "AgentStubConfigFileItem", "AgentStubConfigFileItemsResponse", "AgentStubConfigFileRef", @@ -457,12 +458,10 @@ __all__ = [ "AgentStubFileUploadResponse", "AgentStubURLScheme", "agent_stub_config_env_url", - "agent_stub_config_file_pull_url", "agent_stub_config_manifest_url", "agent_stub_config_note_url", "agent_stub_config_push_url", "agent_stub_config_skill_inspect_url", - "agent_stub_config_skill_pull_url", "agent_stub_connections_url", "agent_stub_drive_base_for_ref", "agent_stub_drive_commit_url", diff --git a/dify-agent/src/dify_agent/agent_stub/server/agent_stub_config.py b/dify-agent/src/dify_agent/agent_stub/server/agent_stub_config.py index 46f4a13bb77..993267c8c7f 100644 --- a/dify-agent/src/dify_agent/agent_stub/server/agent_stub_config.py +++ b/dify-agent/src/dify_agent/agent_stub/server/agent_stub_config.py @@ -2,8 +2,8 @@ Config requests are scoped entirely by the signed execution context carried in the Agent Stub token. Tenant, agent, user, and config-version identifiers come -only from that trusted context; sandbox request bodies contribute only mutable -content such as asset names, env text, and note text. +only from that trusted context; Sandbox request bodies provide only allowed +Config operation inputs and asset selectors. """ from __future__ import annotations @@ -16,10 +16,13 @@ import httpx from pydantic import ValidationError from dify_agent.agent_stub.protocol.agent_stub import ( + AgentStubConfigDownloadSource, AgentStubConfigManifestResponse, AgentStubConfigPushRequest, AgentStubConfigPushResponse, + AgentStubFileDownloadResponse, ) +from dify_agent.agent_stub.server.agent_stub_files import bind_sandbox_file_uri from dify_agent.agent_stub.server.tokens.agent_stub import AgentStubPrincipal from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig @@ -27,12 +30,15 @@ from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig class AgentStubConfigRequestHandler(Protocol): async def manifest(self, *, principal: AgentStubPrincipal) -> AgentStubConfigManifestResponse: ... - async def pull_skill(self, *, principal: AgentStubPrincipal, name: str) -> bytes: ... + async def create_download_request( + self, + *, + principal: AgentStubPrincipal, + source: AgentStubConfigDownloadSource, + ) -> AgentStubFileDownloadResponse: ... async def inspect_skill(self, *, principal: AgentStubPrincipal, name: str) -> dict[str, object]: ... - async def pull_file(self, *, principal: AgentStubPrincipal, name: str) -> bytes: ... - async def push( self, *, @@ -59,13 +65,16 @@ class AgentStubConfigRequestError(RuntimeError): class DifyApiAgentStubConfigRequestHandler: """Call Dify API inner config endpoints on behalf of authenticated sandboxes. - The sandbox never chooses tenant, agent, user, or config-version scope - directly. Those routing fields are derived from the signed execution - context, while request payloads only carry mutable config content. + The Sandbox never chooses tenant, agent, user, or config-version scope. + Those routing fields come from the signed execution context, while request + payloads carry only allowed Config operation inputs or asset selectors. A + download request forwards metadata only and binds the returned URI to the + Sandbox files base URL; it never proxies file bytes. """ inner_api_url: str inner_api_key: str + sandbox_files_base_url: str timeout: httpx.Timeout | float = 30.0 async def manifest(self, *, principal: AgentStubPrincipal) -> AgentStubConfigManifestResponse: @@ -79,12 +88,35 @@ class DifyApiAgentStubConfigRequestHandler: except ValidationError as exc: raise AgentStubConfigRequestError(502, "Dify API config manifest response is invalid") from exc - async def pull_skill(self, *, principal: AgentStubPrincipal, name: str) -> bytes: + async def create_download_request( + self, + *, + principal: AgentStubPrincipal, + source: AgentStubConfigDownloadSource, + ) -> AgentStubFileDownloadResponse: execution_context = self._require_config_context(principal.execution_context) - return await self._get_inner_api_bytes( - f"/inner/api/agent-config/{execution_context.agent_id}/skills/{name}/pull", - self._config_query_params(execution_context), + payload = await self._post_inner_api_json( + f"/inner/api/agent-config/{execution_context.agent_id}/download-request", + { + **self._config_query_params(execution_context), + "config": source.model_dump(mode="json"), + }, ) + if not isinstance(payload, dict): + raise AgentStubConfigRequestError(502, "Dify API config download response is invalid") + download_uri = payload.get("download_uri") + if not isinstance(download_uri, str) or not download_uri: + raise AgentStubConfigRequestError(502, "Dify API config download response is missing download_uri") + try: + download_url = bind_sandbox_file_uri( + sandbox_files_base_url=self.sandbox_files_base_url, + uri=download_uri, + ) + return AgentStubFileDownloadResponse.model_validate( + {key: value for key, value in payload.items() if key != "download_uri"} | {"download_url": download_url} + ) + except (ValueError, ValidationError) as exc: + raise AgentStubConfigRequestError(502, "Dify API config download response is invalid") from exc async def inspect_skill(self, *, principal: AgentStubPrincipal, name: str) -> dict[str, object]: execution_context = self._require_config_context(principal.execution_context) @@ -96,13 +128,6 @@ class DifyApiAgentStubConfigRequestHandler: raise AgentStubConfigRequestError(502, "Dify API config skill inspect response is invalid") return payload - async def pull_file(self, *, principal: AgentStubPrincipal, name: str) -> bytes: - execution_context = self._require_config_context(principal.execution_context) - return await self._get_inner_api_bytes( - f"/inner/api/agent-config/{execution_context.agent_id}/files/{name}/pull", - self._config_query_params(execution_context), - ) - async def push( self, *, @@ -190,18 +215,6 @@ class DifyApiAgentStubConfigRequestHandler: response, invalid_json_detail="Dify API config request returned invalid JSON" ) - async def _get_inner_api_bytes(self, path: str, params: Mapping[str, str]) -> bytes: - response = await self._request("GET", path, params=dict(params)) - if response.is_error: - detail = self._normalize_json_payload( - response, - invalid_json_detail="Dify API config request returned invalid JSON", - ) - raise AgentStubConfigRequestError( - response.status_code, detail.get("detail", detail) if isinstance(detail, dict) else detail - ) - return response.content - async def _post_inner_api_json(self, path: str, payload: Mapping[str, Any]) -> object: response = await self._request("POST", path, json=dict(payload)) return self._normalize_json_payload( diff --git a/dify-agent/src/dify_agent/agent_stub/server/agent_stub_files.py b/dify-agent/src/dify_agent/agent_stub/server/agent_stub_files.py index f3efe34f00a..636a88247d1 100644 --- a/dify-agent/src/dify_agent/agent_stub/server/agent_stub_files.py +++ b/dify-agent/src/dify_agent/agent_stub/server/agent_stub_files.py @@ -14,14 +14,17 @@ from __future__ import annotations from collections.abc import Mapping from dataclasses import dataclass +import posixpath from typing import Any, Protocol +from urllib.parse import unquote, urlsplit import httpx -from pydantic import BaseModel, ConfigDict, ValidationError +from pydantic import ValidationError from dify_agent.agent_stub.protocol.agent_stub import ( AgentStubFileDownloadRequest, AgentStubFileDownloadResponse, + AgentStubFileMapping, AgentStubFileUploadRequest, AgentStubFileUploadResponse, ) @@ -70,36 +73,52 @@ class AgentStubFileRequestError(RuntimeError): super().__init__(str(detail)) -class _BackwardsInvocationEnvelope(BaseModel): - """Minimal parser for Dify API plugin-style inner API envelopes.""" +def bind_sandbox_file_uri(*, sandbox_files_base_url: str, uri: str) -> str: + """Bind one validated origin-free Dify file URI to the Sandbox audience.""" - data: object | None = None - error: str | None = None + if not is_safe_dify_file_uri(uri): + raise ValueError("Dify API returned an unsafe Dify file URI") + return f"{sandbox_files_base_url.rstrip('/')}{uri}" - model_config = ConfigDict(extra="ignore") + +def is_safe_dify_file_uri(value: str) -> bool: + """Return whether a URI stays within Dify's signed ``/files/*`` data plane.""" + + parsed = urlsplit(value) + if parsed.scheme or parsed.netloc or parsed.fragment or value.startswith("//"): + return False + + decoded_path = parsed.path + for _ in range(2): + decoded_path = unquote(decoded_path) + if "\\" in decoded_path: + return False + normalized_path = posixpath.normpath(decoded_path) + return decoded_path.startswith("/files/") and normalized_path.startswith("/files/") @dataclass(slots=True) class DifyApiAgentStubFileRequestHandler: """Call Dify API inner file request endpoints on behalf of the sandbox. - The upload path calls ``/inner/api/upload/file/request`` and injects the - authenticated execution context's ``tenant_id``, ``user_id``, and optional + The upload path calls ``/inner/api/agent/files/upload-request`` and injects the + authenticated execution context's ``tenant_id``, ``user_id``, ``user_from``, and optional ``conversation_id`` along with the requested filename and mimetype. The download path calls - ``/inner/api/download/file/request`` and injects ``tenant_id``, + ``/inner/api/agent/files/download-request`` and injects ``tenant_id``, ``user_id``, ``user_from``, and ``invoke_from`` plus the validated public file mapping. ``user_id`` is mandatory for both operations. Missing user context is rejected before any network call with ``AgentStubFileRequestError(400, ...)``. - Timeouts, transport failures, non-2xx responses, invalid JSON, invalid - plugin-style envelopes, and invalid success schemas are all normalized into + Timeouts, transport failures, non-2xx responses, invalid JSON, and invalid + success schemas are all normalized into ``AgentStubFileRequestError`` so the stub routes can preserve a stable HTTP contract without exposing raw ``httpx`` or Pydantic exceptions. """ inner_api_url: str inner_api_key: str + sandbox_files_base_url: str timeout: httpx.Timeout | float = 30.0 async def create_upload_request( @@ -118,21 +137,24 @@ class DifyApiAgentStubFileRequestHandler: Raises: AgentStubFileRequestError: when user context is incomplete, the inner API times out or fails, the response is non-2xx, or the - success payload does not contain a non-empty ``url`` string. + success payload does not contain a valid ``upload_uri``. """ execution_context = self._require_user_context(principal.execution_context) payload = { "tenant_id": execution_context.tenant_id, "user_id": execution_context.user_id, + "user_from": execution_context.user_from, "filename": request.filename, "mimetype": request.mimetype, "conversation_id": execution_context.conversation_id, } - data = await self._post_inner_api("/inner/api/upload/file/request", payload) - upload_url = data.get("url") - if not isinstance(upload_url, str) or not upload_url: - raise AgentStubFileRequestError(502, "Dify API upload request response is missing url") - return AgentStubFileUploadResponse(upload_url=upload_url) + data = await self._post_inner_api("/inner/api/agent/files/upload-request", payload) + upload_uri = data.get("upload_uri") + if not isinstance(upload_uri, str) or not upload_uri: + raise AgentStubFileRequestError(502, "Dify API upload request response is missing upload_uri") + return AgentStubFileUploadResponse( + upload_url=self._bind_sandbox_files_base_url(upload_uri), + ) async def create_download_request( self, @@ -150,22 +172,31 @@ class DifyApiAgentStubFileRequestHandler: Raises: AgentStubFileRequestError: when user context is incomplete, the inner API times out or fails, the response is non-2xx, the - plugin-style envelope is malformed, or the success payload does - not match ``AgentStubFileDownloadResponse``. + success payload does not contain safe download metadata. """ execution_context = self._require_user_context(principal.execution_context) - payload = { + file_mapping = request.file + if file_mapping is None: + raise AgentStubFileRequestError(400, "file mapping is required for file download requests") + payload: dict[str, object] = { "tenant_id": execution_context.tenant_id, "user_id": execution_context.user_id, "user_from": execution_context.user_from, "invoke_from": execution_context.invoke_from, - "file": request.file.model_dump(mode="json", exclude_none=True), + "file": file_mapping.model_dump(mode="json", exclude_none=True), } - if request.for_external is False: - payload["for_external"] = False - data = await self._post_inner_api("/inner/api/download/file/request", payload) + payload["for_frontend"] = request.for_frontend + data = await self._post_inner_api("/inner/api/agent/files/download-request", payload) + download_uri = data.get("download_uri") + if not isinstance(download_uri, str) or not download_uri: + raise AgentStubFileRequestError(502, "Dify API download request response is missing download_uri") + download_url = self._resolve_download_url( + file_mapping=file_mapping, + for_frontend=request.for_frontend, + download_uri=download_uri, + ) try: - return AgentStubFileDownloadResponse.model_validate(data) + return AgentStubFileDownloadResponse.model_validate({**data, "download_url": download_url}) except ValidationError as exc: raise AgentStubFileRequestError(502, "Dify API download request response is invalid") from exc @@ -194,15 +225,39 @@ class DifyApiAgentStubFileRequestHandler: if response.is_error: detail = raw_payload.get("detail", raw_payload) if isinstance(raw_payload, dict) else raw_payload raise AgentStubFileRequestError(response.status_code, detail) + if not isinstance(raw_payload, dict): + raise AgentStubFileRequestError(502, "Dify API file request response is invalid") + return raw_payload + + def _resolve_download_url( + self, + *, + file_mapping: AgentStubFileMapping, + for_frontend: bool, + download_uri: str, + ) -> str: + if file_mapping.transfer_method == "remote_url": + if self._is_absolute_http_url(download_uri): + return download_uri + raise AgentStubFileRequestError(502, "Dify API returned an invalid remote download URL") + + if for_frontend: + if self._is_absolute_http_url(download_uri) or is_safe_dify_file_uri(download_uri): + return download_uri + raise AgentStubFileRequestError(502, "Dify API returned an unsafe frontend download URI") + + return self._bind_sandbox_files_base_url(download_uri) + + def _bind_sandbox_files_base_url(self, uri: str) -> str: try: - envelope = _BackwardsInvocationEnvelope.model_validate(raw_payload) - except ValidationError as exc: - raise AgentStubFileRequestError(502, "Dify API file request response is invalid") from exc - if envelope.error: - raise AgentStubFileRequestError(400, envelope.error) - if not isinstance(envelope.data, dict): - raise AgentStubFileRequestError(502, "Dify API file request response is missing data") - return dict(envelope.data) + return bind_sandbox_file_uri(sandbox_files_base_url=self.sandbox_files_base_url, uri=uri) + except ValueError as exc: + raise AgentStubFileRequestError(502, str(exc)) from exc + + @staticmethod + def _is_absolute_http_url(value: str) -> bool: + parsed = urlsplit(value) + return parsed.scheme in {"http", "https"} and bool(parsed.netloc) @staticmethod def _parse_json(response: httpx.Response) -> object: @@ -216,4 +271,6 @@ __all__ = [ "AgentStubFileRequestError", "AgentStubFileRequestHandler", "DifyApiAgentStubFileRequestHandler", + "bind_sandbox_file_uri", + "is_safe_dify_file_uri", ] diff --git a/dify-agent/src/dify_agent/agent_stub/server/cli.py b/dify-agent/src/dify_agent/agent_stub/server/cli.py index ed16c1fc43a..1b0989d866a 100644 --- a/dify-agent/src/dify_agent/agent_stub/server/cli.py +++ b/dify-agent/src/dify_agent/agent_stub/server/cli.py @@ -1,22 +1,14 @@ """Console entry point for the standalone Dify Agent stub server. -This module backs the ``dify-agent-stub-server`` console script introduced by -the stub-package move. HTTP(S) endpoints continue to run through Uvicorn while -``grpc://`` Agent Stub URLs switch the process into grpclib server mode. +This module backs the ``dify-agent-stub-server`` console script and serves the +Agent Stub HTTP API through Uvicorn. """ from __future__ import annotations import argparse -import asyncio - import uvicorn -from dify_agent.agent_stub.protocol.agent_stub import parse_agent_stub_endpoint -from dify_agent.agent_stub.server.grpc_bind import AgentStubGRPCBindTarget, derive_agent_stub_grpc_bind_target -from dify_agent.agent_stub.server.grpc_runtime import start_agent_stub_grpc_server -from dify_agent.server.settings import ServerSettings - def main(argv: list[str] | None = None) -> None: """Run the standalone stub server with parsed uvicorn bind options. @@ -26,22 +18,13 @@ def main(argv: list[str] | None = None) -> None: ``argparse`` reads the process command line. Side effects: - Starts either ``dify_agent.agent_stub.server.app:app`` via - ``uvicorn.run`` or the grpclib Agent Stub server depending on the - configured ``DIFY_AGENT_STUB_API_BASE_URL`` scheme. + Starts ``dify_agent.agent_stub.server.app:app`` via ``uvicorn.run``. """ parser = argparse.ArgumentParser(prog="dify-agent-stub-server") parser.add_argument("--host", default=None) parser.add_argument("--port", type=int, default=None) parser.add_argument("--reload", action="store_true") args = parser.parse_args(argv) - settings = ServerSettings() - if ( - settings.agent_stub_api_base_url is not None - and parse_agent_stub_endpoint(settings.agent_stub_api_base_url).is_grpc - ): - asyncio.run(_serve_grpc(settings=settings, host=args.host, port=args.port)) - return uvicorn.run( "dify_agent.agent_stub.server.app:app", host=args.host or "127.0.0.1", @@ -50,24 +33,4 @@ def main(argv: list[str] | None = None) -> None: ) -async def _serve_grpc(*, settings: ServerSettings, host: str | None, port: int | None) -> None: - bind_target = derive_agent_stub_grpc_bind_target( - public_url=settings.agent_stub_api_base_url or "", - bind_address=settings.agent_stub_grpc_bind_address, - ) - if host is not None or port is not None: - bind_target = AgentStubGRPCBindTarget(host=host or bind_target.host, port=port or bind_target.port) - - server = await start_agent_stub_grpc_server( - public_url=settings.agent_stub_api_base_url or "", - bind_address=bind_target.address, - token_codec=settings.create_agent_stub_token_codec(), - file_request_handler=settings.create_agent_stub_file_request_handler(), - ) - try: - await asyncio.Event().wait() - finally: - await server.aclose() - - __all__ = ["main"] diff --git a/dify-agent/src/dify_agent/agent_stub/server/control_plane.py b/dify-agent/src/dify_agent/agent_stub/server/control_plane.py index 5b887af8288..05eba042288 100644 --- a/dify-agent/src/dify_agent/agent_stub/server/control_plane.py +++ b/dify-agent/src/dify_agent/agent_stub/server/control_plane.py @@ -1,9 +1,7 @@ -"""Shared Agent Stub control-plane service used by HTTP and gRPC transports. +"""Shared Agent Stub HTTP control-plane service. -This layer owns the authenticated control-plane delegation for file, config, -and drive operations. Transport adapters validate transport DTOs first, then -call into this service so auth, handler lookup, and error mapping stay shared -across HTTP and gRPC. +This layer owns authenticated delegation for file, config, and drive operations. +The HTTP adapter validates transport DTOs before calling into this service. """ from __future__ import annotations @@ -55,9 +53,8 @@ class AgentStubConfigurationError(AgentStubControlPlaneError): class AgentStubControlPlaneService: """Shared business service for authenticated Agent Stub control-plane calls. - HTTP and gRPC adapters validate or decode transport payloads before calling - this service, so this layer focuses only on shared auth, connection-id - generation, plus file, config, and drive request delegation. + The HTTP adapter validates transport payloads before calling this service, + which focuses on auth, connection-id generation, and request delegation. """ token_codec: AgentStubTokenCodec | None @@ -93,6 +90,13 @@ class AgentStubControlPlaneService: ) -> AgentStubFileDownloadResponse: """Authenticate and delegate one already-validated file-download request.""" principal = self._authenticate(authorization) + if request.config is not None: + handler = self._require_config_request_handler() + try: + return await handler.create_download_request(principal=principal, source=request.config) + except AgentStubConfigRequestError as exc: + raise AgentStubControlPlaneError(exc.status_code, exc.detail) from exc + handler = self._require_file_request_handler() try: return await handler.create_download_request(principal=principal, request=request) @@ -130,19 +134,6 @@ class AgentStubControlPlaneService: except AgentStubConfigRequestError as exc: raise AgentStubControlPlaneError(exc.status_code, exc.detail) from exc - async def pull_config_skill( - self, - *, - name: str, - authorization: str | None, - ) -> bytes: - principal = self._authenticate(authorization) - handler = self._require_config_request_handler() - try: - return await handler.pull_skill(principal=principal, name=name) - except AgentStubConfigRequestError as exc: - raise AgentStubControlPlaneError(exc.status_code, exc.detail) from exc - async def inspect_config_skill( self, *, @@ -156,19 +147,6 @@ class AgentStubControlPlaneService: except AgentStubConfigRequestError as exc: raise AgentStubControlPlaneError(exc.status_code, exc.detail) from exc - async def pull_config_file( - self, - *, - name: str, - authorization: str | None, - ) -> bytes: - principal = self._authenticate(authorization) - handler = self._require_config_request_handler() - try: - return await handler.pull_file(principal=principal, name=name) - except AgentStubConfigRequestError as exc: - raise AgentStubControlPlaneError(exc.status_code, exc.detail) from exc - async def push_config( self, *, diff --git a/dify-agent/src/dify_agent/agent_stub/server/grpc_bind.py b/dify-agent/src/dify_agent/agent_stub/server/grpc_bind.py deleted file mode 100644 index 9f4ee9e0002..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/server/grpc_bind.py +++ /dev/null @@ -1,69 +0,0 @@ -"""Bind-address helpers for optional Agent Stub gRPC hosting.""" - -from __future__ import annotations - -from dataclasses import dataclass -from urllib.parse import urlsplit - -from dify_agent.agent_stub.protocol.agent_stub import parse_agent_stub_endpoint - - -@dataclass(frozen=True, slots=True) -class AgentStubGRPCBindTarget: - """Validated host/port bind target for a grpclib Agent Stub server.""" - - host: str - port: int - - @property - def address(self) -> str: - return f"{_format_host(self.host)}:{self.port}" - - -def normalize_agent_stub_grpc_bind_address(value: str) -> str: - """Normalize one ``host:port`` gRPC bind address.""" - target = parse_agent_stub_grpc_bind_address(value) - return target.address - - -def parse_agent_stub_grpc_bind_address(value: str) -> AgentStubGRPCBindTarget: - """Parse one explicit ``host:port`` gRPC bind override.""" - stripped = value.strip() - if not stripped: - raise ValueError("Agent Stub gRPC bind address must not be empty") - parsed = urlsplit(f"grpc://{stripped}") - if not parsed.netloc or parsed.hostname is None: - raise ValueError("Agent Stub gRPC bind address must include a host") - if parsed.username is not None or parsed.password is not None: - raise ValueError("Agent Stub gRPC bind address must not include user info") - if parsed.port is None: - raise ValueError("Agent Stub gRPC bind address must include an explicit port") - if parsed.path not in {"", "/"} or parsed.query or parsed.fragment: - raise ValueError("Agent Stub gRPC bind address must be in host:port form") - return AgentStubGRPCBindTarget(host=parsed.hostname, port=parsed.port) - - -def derive_agent_stub_grpc_bind_target( - *, - public_url: str, - bind_address: str | None = None, -) -> AgentStubGRPCBindTarget: - """Resolve the runtime gRPC bind target from public URL plus optional override.""" - if bind_address is not None: - return parse_agent_stub_grpc_bind_address(bind_address) - endpoint = parse_agent_stub_endpoint(public_url) - if not endpoint.is_grpc or endpoint.port is None: - raise ValueError("Agent Stub gRPC bind target requires a grpc://host:port public URL") - return AgentStubGRPCBindTarget(host="0.0.0.0", port=endpoint.port) - - -def _format_host(host: str) -> str: - return f"[{host}]" if ":" in host and not host.startswith("[") else host - - -__all__ = [ - "AgentStubGRPCBindTarget", - "derive_agent_stub_grpc_bind_target", - "normalize_agent_stub_grpc_bind_address", - "parse_agent_stub_grpc_bind_address", -] diff --git a/dify-agent/src/dify_agent/agent_stub/server/grpc_runtime.py b/dify-agent/src/dify_agent/agent_stub/server/grpc_runtime.py deleted file mode 100644 index 0ef28ae5dc6..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/server/grpc_runtime.py +++ /dev/null @@ -1,59 +0,0 @@ -"""Runtime helpers for starting and stopping the optional Agent Stub gRPC server.""" - -from __future__ import annotations - -from dataclasses import dataclass - -from dify_agent.agent_stub.server.agent_stub_files import AgentStubFileRequestHandler -from dify_agent.agent_stub.server.control_plane import AgentStubControlPlaneService -from dify_agent.agent_stub.server.grpc_bind import AgentStubGRPCBindTarget, derive_agent_stub_grpc_bind_target -from dify_agent.agent_stub.server.tokens.agent_stub import AgentStubTokenCodec - - -@dataclass(slots=True) -class RunningAgentStubGRPCServer: - """Handle for one started grpclib Agent Stub server.""" - - server: object - bind_target: AgentStubGRPCBindTarget - - async def aclose(self) -> None: - """Stop accepting new requests and wait for open RPCs to close.""" - close = getattr(self.server, "close") - wait_closed = getattr(self.server, "wait_closed") - close() - await wait_closed() - - -async def start_agent_stub_grpc_server( - *, - public_url: str, - bind_address: str | None, - token_codec: AgentStubTokenCodec | None, - file_request_handler: AgentStubFileRequestHandler | None, -) -> RunningAgentStubGRPCServer: - """Start the optional grpclib Agent Stub server for one process.""" - from dify_agent.agent_stub.server.grpc_service import create_agent_stub_grpc_service - - runtime = _require_runtime() - bind_target = derive_agent_stub_grpc_bind_target(public_url=public_url, bind_address=bind_address) - service = AgentStubControlPlaneService(token_codec, file_request_handler) - server = runtime.Server([create_agent_stub_grpc_service(service)]) - await server.start(bind_target.host, bind_target.port) - return RunningAgentStubGRPCServer(server=server, bind_target=bind_target) - - -@dataclass(frozen=True, slots=True) -class _GRPCRuntime: - Server: type - - -def _require_runtime() -> _GRPCRuntime: - try: - from grpclib.server import Server - except ImportError as exc: - raise RuntimeError("Agent Stub gRPC support requires the optional dify-agent[grpc] dependencies") from exc - return _GRPCRuntime(Server=Server) - - -__all__ = ["RunningAgentStubGRPCServer", "start_agent_stub_grpc_server"] diff --git a/dify-agent/src/dify_agent/agent_stub/server/grpc_service.py b/dify-agent/src/dify_agent/agent_stub/server/grpc_service.py deleted file mode 100644 index 48749a48436..00000000000 --- a/dify-agent/src/dify_agent/agent_stub/server/grpc_service.py +++ /dev/null @@ -1,244 +0,0 @@ -"""gRPC transport adapter for the Agent Stub control plane. - -This module owns the gRPC-specific semantics that differ from the HTTP router: - -- compact-JWE auth is read from inbound ``authorization`` metadata rather than - an HTTP header object; -- protobuf request-shape or JSON-decoding failures are mapped to - ``INVALID_ARGUMENT`` before the shared business layer runs; -- auth/configuration/downstream failures raised by - ``AgentStubControlPlaneService`` are translated into gRPC statuses; -- structured downstream error details are stringified because gRPC status - details are text, unlike the HTTP path which can preserve object-shaped - ``detail`` payloads. - -The shared control-plane service still owns auth policy, connection-id -generation, and file-handler delegation so HTTP and gRPC stay semantically -aligned outside these transport-specific mappings. -""" - -from __future__ import annotations - -from collections.abc import Mapping, Sequence -from dataclasses import dataclass -from typing import TYPE_CHECKING - -from pydantic import ValidationError - -from dify_agent.agent_stub.server.control_plane import ( - AgentStubAuthenticationError, - AgentStubConfigurationError, - AgentStubControlPlaneError, - AgentStubControlPlaneService, -) - -if TYPE_CHECKING: - from grpclib.server import Stream - - from dify_agent.agent_stub.grpc._generated import agent_stub_pb2 - from dify_agent.agent_stub.grpc._generated.agent_stub_grpc import AgentStubServiceBase - - -@dataclass(slots=True) -class AgentStubGRPCTransport: - """Shared gRPC adapter that converts protobuf messages around control-plane calls. - - Each method validates the protobuf request at the transport boundary, - extracts ``authorization`` metadata, and translates shared control-plane - failures into grpclib ``GRPCError`` instances. - """ - - service: AgentStubControlPlaneService - - async def connect( - self, - *, - request, - metadata: object, - ): - """Handle one gRPC connect request. - - Invalid protobuf/message-shape input maps to ``INVALID_ARGUMENT``. Auth - and configuration failures raised by the shared service are translated by - ``_grpc_error_from_control_plane_error``. - """ - authorization = _authorization_from_metadata(metadata) - conversions = _require_conversions() - try: - _ = conversions.connect_request_from_proto(request) - response = await self.service.connect(authorization=authorization) - return conversions.proto_connect_response(response) - except (ValidationError, ValueError, TypeError) as exc: - raise _grpc_error("INVALID_ARGUMENT", "invalid Agent Stub connect request") from exc - except AgentStubControlPlaneError as exc: - raise _grpc_error_from_control_plane_error(exc) from exc - - async def create_file_upload_request( - self, - *, - request, - metadata: object, - ): - """Handle one gRPC file-upload request. - - This transport validates the protobuf request before delegating to the - shared service, then stringifies any non-string downstream error detail - when mapping the resulting failure into gRPC status text. - """ - authorization = _authorization_from_metadata(metadata) - conversions = _require_conversions() - try: - validated_request = conversions.file_upload_request_from_proto(request) - response = await self.service.create_file_upload_request( - request=validated_request, - authorization=authorization, - ) - return conversions.proto_file_upload_response(response) - except (ValidationError, ValueError, TypeError) as exc: - raise _grpc_error("INVALID_ARGUMENT", "invalid Agent Stub file upload request") from exc - except AgentStubControlPlaneError as exc: - raise _grpc_error_from_control_plane_error(exc) from exc - - async def create_file_download_request( - self, - *, - request, - metadata: object, - ): - """Handle one gRPC file-download request. - - This transport validates the protobuf request before delegating to the - shared service, then stringifies any non-string downstream error detail - when mapping the resulting failure into gRPC status text. - """ - authorization = _authorization_from_metadata(metadata) - conversions = _require_conversions() - try: - validated_request = conversions.file_download_request_from_proto(request) - response = await self.service.create_file_download_request( - request=validated_request, - authorization=authorization, - ) - return conversions.proto_file_download_response(response) - except (ValidationError, ValueError, TypeError) as exc: - raise _grpc_error("INVALID_ARGUMENT", "invalid Agent Stub file download request") from exc - except AgentStubControlPlaneError as exc: - raise _grpc_error_from_control_plane_error(exc) from exc - - -def create_agent_stub_grpc_service( - service: AgentStubControlPlaneService, -) -> AgentStubServiceBase: - """Wrap the shared control-plane service in a grpclib-generated service base. - - The generated grpclib service methods are unary-only and reject missing - request messages with ``INVALID_ARGUMENT`` before handing control to the - transport adapter. - """ - try: - from dify_agent.agent_stub.grpc._generated.agent_stub_grpc import AgentStubServiceBase - except ImportError as exc: # pragma: no cover - exercised via runtime import guard - raise RuntimeError("Agent Stub gRPC support requires the optional grpc dependencies") from exc - - transport = AgentStubGRPCTransport(service) - - class _Service(AgentStubServiceBase): - async def Connect( - self, - stream: Stream[agent_stub_pb2.ConnectRequest, agent_stub_pb2.ConnectResponse], - ) -> None: # type: ignore[name-defined] - request = await stream.recv_message() - if request is None: - raise _grpc_error("INVALID_ARGUMENT", "missing Agent Stub request message") - await stream.send_message(await transport.connect(request=request, metadata=stream.metadata)) - - async def CreateFileUploadRequest( - self, - stream: Stream[agent_stub_pb2.FileUploadRequest, agent_stub_pb2.FileUploadResponse], - ) -> None: # type: ignore[name-defined] - request = await stream.recv_message() - if request is None: - raise _grpc_error("INVALID_ARGUMENT", "missing Agent Stub request message") - await stream.send_message( - await transport.create_file_upload_request(request=request, metadata=stream.metadata) - ) - - async def CreateFileDownloadRequest( - self, - stream: Stream[agent_stub_pb2.FileDownloadRequest, agent_stub_pb2.FileDownloadResponse], - ) -> None: # type: ignore[name-defined] - request = await stream.recv_message() - if request is None: - raise _grpc_error("INVALID_ARGUMENT", "missing Agent Stub request message") - await stream.send_message( - await transport.create_file_download_request(request=request, metadata=stream.metadata) - ) - - return _Service() - - -def _authorization_from_metadata(metadata: object) -> str | None: - """Extract the optional bearer token from grpclib metadata containers.""" - if metadata is None: - return None - if isinstance(metadata, Mapping): - value = metadata.get("authorization") - return value if isinstance(value, str) else None - if isinstance(metadata, Sequence): - for item in metadata: - if isinstance(item, tuple) and len(item) == 2: - key, value = item - if isinstance(key, str) and key.lower() == "authorization" and isinstance(value, str): - return value - return None - - -def _grpc_error_from_control_plane_error(exc: AgentStubControlPlaneError): - """Translate shared control-plane failures into transport-visible gRPC status. - - ``detail`` is normalized to text because gRPC status details are strings; - HTTP adapters can preserve richer object payloads. - """ - detail_text = _grpc_detail_text(exc.detail) - if isinstance(exc, AgentStubAuthenticationError): - return _grpc_error("UNAUTHENTICATED", detail_text) - if isinstance(exc, AgentStubConfigurationError): - return _grpc_error("UNAVAILABLE", detail_text) - if exc.status_code in {408, 504}: - return _grpc_error("DEADLINE_EXCEEDED", detail_text) - if exc.status_code == 429: - return _grpc_error("RESOURCE_EXHAUSTED", detail_text) - if exc.status_code == 404: - return _grpc_error("NOT_FOUND", detail_text) - if exc.status_code == 403: - return _grpc_error("PERMISSION_DENIED", detail_text) - if 400 <= exc.status_code < 500: - return _grpc_error("FAILED_PRECONDITION", detail_text) - if 500 <= exc.status_code < 600: - return _grpc_error("UNAVAILABLE", detail_text) - return _grpc_error("INTERNAL", "internal Agent Stub error") - - -def _grpc_error(status_name: str, detail: str): - from grpclib.const import Status - from grpclib.exceptions import GRPCError - - return GRPCError(getattr(Status, status_name), detail) - - -def _grpc_detail_text(detail: object) -> str: - """Return the text form used for gRPC status details.""" - if isinstance(detail, str): - return detail - return str(detail) - - -def _require_conversions(): - try: - from dify_agent.agent_stub.grpc import conversions - except ImportError as exc: # pragma: no cover - exercised by runtime import guard - raise RuntimeError("Agent Stub gRPC support requires the optional grpc dependencies") from exc - return conversions - - -__all__ = ["AgentStubGRPCTransport", "create_agent_stub_grpc_service"] diff --git a/dify-agent/src/dify_agent/agent_stub/server/routes/agent_stub.py b/dify-agent/src/dify_agent/agent_stub/server/routes/agent_stub.py index 411fcc4c6af..22e8fa9f94f 100644 --- a/dify-agent/src/dify_agent/agent_stub/server/routes/agent_stub.py +++ b/dify-agent/src/dify_agent/agent_stub/server/routes/agent_stub.py @@ -2,13 +2,12 @@ The router is a thin HTTP adapter around ``AgentStubControlPlaneService``. It keeps FastAPI-specific request parsing and HTTPException translation here while -sharing auth, DTO validation, connection-id generation, and file/config/drive -delegation with the gRPC transport. +the service owns auth and file/config/drive delegation. """ from __future__ import annotations -from fastapi import APIRouter, Header, HTTPException, Response +from fastapi import APIRouter, Header, HTTPException from dify_agent.agent_stub.protocol.agent_stub import ( AgentStubConnectRequest, @@ -42,7 +41,10 @@ def create_agent_stub_http_router( """Create HTTP routes bound to the application's Agent Stub dependencies.""" router = APIRouter(prefix="/agent-stub", tags=["agent-stub"]) service = AgentStubControlPlaneService( - token_codec, file_request_handler, config_request_handler, drive_request_handler + token_codec=token_codec, + file_request_handler=file_request_handler, + config_request_handler=config_request_handler, + drive_request_handler=drive_request_handler, ) @router.post("/connections", response_model=AgentStubConnectResponse) @@ -85,17 +87,6 @@ def create_agent_stub_http_router( except AgentStubControlPlaneError as exc: raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc - @router.get("/config/skills/{name}/pull") - async def pull_config_skill( - name: str, - authorization: str | None = Header(default=None, alias="Authorization"), - ) -> Response: - try: - payload = await service.pull_config_skill(name=name, authorization=authorization) - return Response(content=payload, media_type="application/zip") - except AgentStubControlPlaneError as exc: - raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc - @router.get("/config/skills/{name}/inspect") async def inspect_config_skill( name: str, @@ -106,17 +97,6 @@ def create_agent_stub_http_router( except AgentStubControlPlaneError as exc: raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc - @router.get("/config/files/{name}/pull") - async def pull_config_file( - name: str, - authorization: str | None = Header(default=None, alias="Authorization"), - ) -> Response: - try: - payload = await service.pull_config_file(name=name, authorization=authorization) - return Response(content=payload, media_type="application/octet-stream") - except AgentStubControlPlaneError as exc: - raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc - @router.post("/config/push", response_model=AgentStubConfigPushResponse) async def push_config( request: AgentStubConfigPushRequest, diff --git a/dify-agent/src/dify_agent/client/_client.py b/dify-agent/src/dify_agent/client/_client.py index 9c15d96cb59..0cd14a68711 100644 --- a/dify-agent/src/dify_agent/client/_client.py +++ b/dify-agent/src/dify_agent/client/_client.py @@ -1,7 +1,7 @@ """HTTPX-based client for the Dify Agent HTTP API. The client uses the public DTOs from ``dify_agent.protocol`` for request and -response parsing across both run-management and sandbox-file endpoints. It +response parsing across run-management and working-environment endpoints. It intentionally does not retry non-idempotent ``POST`` requests such as ``/runs``. SSE streams are the only operation with reconnect logic: transient stream, connect, or read failures, stream timeouts, and HTTP 5xx stream @@ -28,24 +28,30 @@ from pydantic_ai.messages import FunctionToolResultEvent from dify_agent.protocol import ( CancelRunRequest, CancelRunResponse, + BindingFileDownloadRequest, + BindingFileDownloadResponse, + BindingFileListRequest, + BindingFileListResponse, + BindingFileReadRequest, + BindingFileReadResponse, CreateRunRequest, CreateRunResponse, + CreateExecutionBindingRequest, + CreateExecutionBindingResponse, + CreateHomeSnapshotFromBindingRequest, + DeleteHomeSnapshotRequest, + DestroyExecutionBindingRequest, + HomeSnapshotResponse, RUN_EVENT_ADAPTER, RunEvent, RunEventsResponse, RunStatusResponse, - SandboxListRequest, - SandboxListResponse, - SandboxLocator, - SandboxReadRequest, - SandboxReadResponse, - SandboxUploadRequest, - SandboxUploadResponse, ) _ResponseModelT = TypeVar("_ResponseModelT", bound=BaseModel) _TERMINAL_EVENT_TYPES = {"run_succeeded", "run_failed", "run_cancelled"} _TERMINAL_RUN_STATUSES = {"succeeded", "failed", "cancelled"} +_BINDING_FILE_DOWNLOAD_TIMEOUT_SECONDS = 90.0 _function_tool_result_payload_key_cache: str | None = None @@ -255,7 +261,7 @@ class Client: headers, timeout settings, optional external HTTPX clients, and lazy-owned clients for whichever sync/async side is used. It is the shared transport boundary for both run-management endpoints (create/status/events/cancel) and - sandbox-file endpoints (list/read/upload). External clients are never closed + Binding-file endpoints (list/read/download). External clients are never closed by this wrapper. Owned sync clients close via ``close_sync`` or the sync context manager; owned async clients close via ``aclose`` or the async context manager. @@ -378,8 +384,8 @@ class Client: async def cancel_run(self, run_id: str, request: CancelRunRequest | None = None) -> CancelRunResponse: """Request explicit cancellation for ``run_id``. - The server may accept cancellation only for active runs; unsupported - deployments return an HTTP error rather than overloading ``run_failed``. + Acceptance atomically persists the cancelled state. The process executing + the run observes that state and performs runner cleanup asynchronously. """ request_model = request or CancelRunRequest() try: @@ -469,61 +475,123 @@ class Client: raise DifyAgentClientError(f"get_events_sync request failed: {exc}") from exc return _parse_model_response(response, RunEventsResponse) - async def list_sandbox_files(self, locator: SandboxLocator, path: str) -> SandboxListResponse: - """List a sandbox directory through ``POST /sandbox/files/list``.""" - request_model = _build_request_model(SandboxListRequest, locator=locator, path=path) - response = await self._post_async_json("list_sandbox_files", "/sandbox/files/list", request_model) - return _parse_model_response(response, SandboxListResponse) + async def list_binding_files(self, backend_binding_ref: str, path: str) -> BindingFileListResponse: + request_model = BindingFileListRequest(backend_binding_ref=backend_binding_ref, path=path) + response = await self._post_async_json("list_binding_files", "/execution-bindings/files/list", request_model) + return _parse_model_response(response, BindingFileListResponse) - def list_sandbox_files_sync(self, locator: SandboxLocator, path: str) -> SandboxListResponse: - """Synchronous variant of ``list_sandbox_files``.""" - request_model = _build_request_model(SandboxListRequest, locator=locator, path=path) - response = self._post_sync_json("list_sandbox_files_sync", "/sandbox/files/list", request_model) - return _parse_model_response(response, SandboxListResponse) + def list_binding_files_sync(self, backend_binding_ref: str, path: str) -> BindingFileListResponse: + request_model = BindingFileListRequest(backend_binding_ref=backend_binding_ref, path=path) + response = self._post_sync_json("list_binding_files_sync", "/execution-bindings/files/list", request_model) + return _parse_model_response(response, BindingFileListResponse) - async def read_sandbox_file( + async def read_binding_file( self, - locator: SandboxLocator, + backend_binding_ref: str, path: str, max_bytes: int = 262144, - ) -> SandboxReadResponse: - """Read a sandbox file preview through ``POST /sandbox/files/read``.""" - request_model = _build_request_model( - SandboxReadRequest, - locator=locator, - path=path, - max_bytes=max_bytes, - ) - response = await self._post_async_json("read_sandbox_file", "/sandbox/files/read", request_model) - return _parse_model_response(response, SandboxReadResponse) + ) -> BindingFileReadResponse: + request_model = BindingFileReadRequest(backend_binding_ref=backend_binding_ref, path=path, max_bytes=max_bytes) + response = await self._post_async_json("read_binding_file", "/execution-bindings/files/read", request_model) + return _parse_model_response(response, BindingFileReadResponse) - def read_sandbox_file_sync( + def read_binding_file_sync( self, - locator: SandboxLocator, + backend_binding_ref: str, path: str, max_bytes: int = 262144, - ) -> SandboxReadResponse: - """Synchronous variant of ``read_sandbox_file``.""" - request_model = _build_request_model( - SandboxReadRequest, - locator=locator, - path=path, - max_bytes=max_bytes, + ) -> BindingFileReadResponse: + request_model = BindingFileReadRequest(backend_binding_ref=backend_binding_ref, path=path, max_bytes=max_bytes) + response = self._post_sync_json("read_binding_file_sync", "/execution-bindings/files/read", request_model) + return _parse_model_response(response, BindingFileReadResponse) + + async def download_binding_file(self, request: BindingFileDownloadRequest) -> BindingFileDownloadResponse: + response = await self._post_async_json( + "download_binding_file", + "/execution-bindings/files/download", + request, + timeout=_BINDING_FILE_DOWNLOAD_TIMEOUT_SECONDS, ) - response = self._post_sync_json("read_sandbox_file_sync", "/sandbox/files/read", request_model) - return _parse_model_response(response, SandboxReadResponse) + return _parse_model_response(response, BindingFileDownloadResponse) - async def upload_sandbox_file(self, locator: SandboxLocator, path: str) -> SandboxUploadResponse: - """Upload a sandbox file mapping through ``POST /sandbox/files/upload``.""" - request_model = _build_request_model(SandboxUploadRequest, locator=locator, path=path) - response = await self._post_async_json("upload_sandbox_file", "/sandbox/files/upload", request_model) - return _parse_model_response(response, SandboxUploadResponse) + def download_binding_file_sync(self, request: BindingFileDownloadRequest) -> BindingFileDownloadResponse: + response = self._post_sync_json( + "download_binding_file_sync", + "/execution-bindings/files/download", + request, + timeout=_BINDING_FILE_DOWNLOAD_TIMEOUT_SECONDS, + ) + return _parse_model_response(response, BindingFileDownloadResponse) - def upload_sandbox_file_sync(self, locator: SandboxLocator, path: str) -> SandboxUploadResponse: - """Synchronous variant of ``upload_sandbox_file``.""" - request_model = _build_request_model(SandboxUploadRequest, locator=locator, path=path) - response = self._post_sync_json("upload_sandbox_file_sync", "/sandbox/files/upload", request_model) - return _parse_model_response(response, SandboxUploadResponse) + async def create_execution_binding(self, request: CreateExecutionBindingRequest) -> CreateExecutionBindingResponse: + response = await self._post_async_json("create_execution_binding", "/execution-bindings", request) + return _parse_model_response(response, CreateExecutionBindingResponse) + + def create_execution_binding_sync(self, request: CreateExecutionBindingRequest) -> CreateExecutionBindingResponse: + response = self._post_sync_json("create_execution_binding_sync", "/execution-bindings", request) + return _parse_model_response(response, CreateExecutionBindingResponse) + + async def destroy_execution_binding(self, request: DestroyExecutionBindingRequest) -> None: + response = await self._post_async_json("destroy_execution_binding", "/execution-bindings/destroy", request) + _raise_for_status(response) + + def destroy_execution_binding_sync(self, request: DestroyExecutionBindingRequest) -> None: + response = self._post_sync_json("destroy_execution_binding_sync", "/execution-bindings/destroy", request) + _raise_for_status(response) + + async def create_home_snapshot_from_binding( + self, + request: CreateHomeSnapshotFromBindingRequest, + ) -> HomeSnapshotResponse: + """Checkpoint Home from the exact Execution Binding identified by the request.""" + response = await self._post_async_json( + "create_home_snapshot_from_binding", + "/home-snapshots/from-binding", + request, + ) + return _parse_model_response(response, HomeSnapshotResponse) + + def create_home_snapshot_from_binding_sync( + self, + request: CreateHomeSnapshotFromBindingRequest, + ) -> HomeSnapshotResponse: + """Synchronous variant of ``create_home_snapshot_from_binding``.""" + response = self._post_sync_json( + "create_home_snapshot_from_binding_sync", + "/home-snapshots/from-binding", + request, + ) + return _parse_model_response(response, HomeSnapshotResponse) + + async def delete_home_snapshot(self, snapshot_ref: str) -> None: + """Idempotently delete one backend Home Snapshot.""" + try: + response = await self._get_async_http_client().post( + self._url("/home-snapshots/delete"), + content=DeleteHomeSnapshotRequest(snapshot_ref=snapshot_ref).model_dump_json(), + headers=self._merged_headers({"Content-Type": "application/json"}), + timeout=self._timeout, + ) + except httpx.TimeoutException as exc: + raise DifyAgentTimeoutError("delete_home_snapshot timed out") from exc + except httpx.RequestError as exc: + raise DifyAgentClientError(f"delete_home_snapshot request failed: {exc}") from exc + _raise_for_status(response) + + def delete_home_snapshot_sync(self, snapshot_ref: str) -> None: + """Synchronous variant of ``delete_home_snapshot``.""" + try: + response = self._get_sync_http_client().post( + self._url("/home-snapshots/delete"), + content=DeleteHomeSnapshotRequest(snapshot_ref=snapshot_ref).model_dump_json(), + headers=self._merged_headers({"Content-Type": "application/json"}), + timeout=self._timeout, + ) + except httpx.TimeoutException as exc: + raise DifyAgentTimeoutError("delete_home_snapshot_sync timed out") from exc + except httpx.RequestError as exc: + raise DifyAgentClientError(f"delete_home_snapshot_sync request failed: {exc}") from exc + _raise_for_status(response) async def stream_events( self, @@ -543,11 +611,16 @@ class Client: with an id, reconnects resume from that id using the ``after`` query parameter. HTTP 5xx stream responses are retried, but HTTP 4xx responses, DTO validation failures, and malformed SSE frames are not retried. By - default iteration stops after a succeeded, failed, or cancelled terminal event. + default, ``until_terminal=True`` returns immediately after yielding a + succeeded, failed, or cancelled terminal event. With + ``until_terminal=False``, iteration may consume the remainder of the current + response, but after observing a terminal event it will not reconnect when that + response ends normally or raises a reconnectable transport error. """ _validate_stream_options(max_reconnects, reconnect_delay_seconds, timeout_seconds) cursor = after or "0-0" reconnect_attempts = 0 + terminal_event_seen = False deadline = time.monotonic() + timeout_seconds if timeout_seconds is not None else None while True: _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) @@ -560,10 +633,14 @@ class Client: ): if event.id is not None: cursor = event.id + if event.type in _TERMINAL_EVENT_TYPES: + terminal_event_seen = True yield event - if until_terminal and event.type in _TERMINAL_EVENT_TYPES: + if until_terminal and terminal_event_seen: return except _ReconnectableStreamError as exc: + if terminal_event_seen: + return if not reconnect: raise exc.error from exc reconnect_attempts = _next_reconnect_attempt( @@ -574,6 +651,8 @@ class Client: _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) await _sleep_async(_bounded_sleep_seconds(reconnect_delay_seconds, deadline)) continue + if terminal_event_seen: + return if not reconnect: return reconnect_attempts = _next_reconnect_attempt( @@ -600,6 +679,7 @@ class Client: _validate_stream_options(max_reconnects, reconnect_delay_seconds, timeout_seconds) cursor = after or "0-0" reconnect_attempts = 0 + terminal_event_seen = False deadline = time.monotonic() + timeout_seconds if timeout_seconds is not None else None while True: _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) @@ -612,10 +692,14 @@ class Client: ): if event.id is not None: cursor = event.id + if event.type in _TERMINAL_EVENT_TYPES: + terminal_event_seen = True yield event - if until_terminal and event.type in _TERMINAL_EVENT_TYPES: + if until_terminal and terminal_event_seen: return except _ReconnectableStreamError as exc: + if terminal_event_seen: + return if not reconnect: raise exc.error from exc reconnect_attempts = _next_reconnect_attempt( @@ -626,6 +710,8 @@ class Client: _raise_if_stream_stopped(run_id, deadline=deadline, should_stop=should_stop) _sleep_sync(_bounded_sleep_seconds(reconnect_delay_seconds, deadline)) continue + if terminal_event_seen: + return if not reconnect: return reconnect_attempts = _next_reconnect_attempt( @@ -787,26 +873,40 @@ class Client: headers.update(extra) return headers - async def _post_async_json(self, operation: str, path: str, request_model: BaseModel) -> httpx.Response: + async def _post_async_json( + self, + operation: str, + path: str, + request_model: BaseModel, + *, + timeout: float | httpx.Timeout | None = None, + ) -> httpx.Response: try: return await self._get_async_http_client().post( self._url(path), content=request_model.model_dump_json(), headers=self._merged_headers({"Content-Type": "application/json"}), - timeout=self._timeout, + timeout=self._timeout if timeout is None else timeout, ) except httpx.TimeoutException as exc: raise DifyAgentTimeoutError(f"{operation} timed out") from exc except httpx.RequestError as exc: raise DifyAgentClientError(f"{operation} request failed: {exc}") from exc - def _post_sync_json(self, operation: str, path: str, request_model: BaseModel) -> httpx.Response: + def _post_sync_json( + self, + operation: str, + path: str, + request_model: BaseModel, + *, + timeout: float | httpx.Timeout | None = None, + ) -> httpx.Response: try: return self._get_sync_http_client().post( self._url(path), content=request_model.model_dump_json(), headers=self._merged_headers({"Content-Type": "application/json"}), - timeout=self._timeout, + timeout=self._timeout if timeout is None else timeout, ) except httpx.TimeoutException as exc: raise DifyAgentTimeoutError(f"{operation} timed out") from exc diff --git a/dify-agent/src/dify_agent/layers/_agent_cli_help.json b/dify-agent/src/dify_agent/layers/_agent_cli_help.json index 1361d1a4170..66578c3d9e0 100644 --- a/dify-agent/src/dify_agent/layers/_agent_cli_help.json +++ b/dify-agent/src/dify_agent/layers/_agent_cli_help.json @@ -21,5 +21,5 @@ "drive push": "Upload one local file or directory into the agent drive.\n\nUsage:\n dify-agent drive push LOCAL_PATH REMOTE_PATH [flags]\n\nFlags:\n -h, --help help for push\n --json Accepted for consistency; drive push output is already emitted as JSON.\n --kind string Directory upload kind: skill or dir.", "file": "Upload or download workflow files through the Agent Stub.\n\nUsage:\n dify-agent file [command]\n\nAvailable Commands:\n download Download one workflow file mapping into the local sandbox directory.\n upload Upload one sandbox-local file as a ToolFile output reference.\n\nFlags:\n -h, --help help for file\n\nUse \"dify-agent file [command] --help\" for more information about a command.", "file download": "Download one workflow file mapping into the local sandbox directory.\n\nUsage:\n dify-agent file download TRANSFER_METHOD REFERENCE_OR_URL [flags]\n\nFlags:\n -h, --help help for download\n --to string Local directory for the downloaded file.", - "file upload": "Upload one sandbox-local file as a ToolFile output reference.\n\nUsage:\n dify-agent file upload PATH [flags]\n\nFlags:\n -h, --help help for upload" + "file upload": "Upload one sandbox-local file as a ToolFile output reference.\n\nUsage:\n dify-agent file upload PATH [flags]\n\nFlags:\n -h, --help help for upload\n --no-download-link Skip creating a public download link after upload." } diff --git a/dify-agent/src/dify_agent/layers/_agent_file_cli_help.py b/dify-agent/src/dify_agent/layers/_agent_file_cli_help.py index 50deef3d851..cbe19a27bbe 100644 --- a/dify-agent/src/dify_agent/layers/_agent_file_cli_help.py +++ b/dify-agent/src/dify_agent/layers/_agent_file_cli_help.py @@ -3,5 +3,5 @@ AGENT_FILE_UPLOAD_REPLY_HINT = ( "When you want to provide a generated or sandbox-local file to the user in a " "natural-language reply, run the installed CLI command `dify-agent file upload PATH` and include the returned " - "`download_url` so the user can open or download the file." + "`public_download_url` so the user can open or download the file." ) diff --git a/dify-agent/src/dify_agent/layers/config/layer.py b/dify-agent/src/dify_agent/layers/config/layer.py index 7088179996f..7d68202d2a0 100644 --- a/dify-agent/src/dify_agent/layers/config/layer.py +++ b/dify-agent/src/dify_agent/layers/config/layer.py @@ -2,7 +2,7 @@ from __future__ import annotations -import asyncio +import json import shlex from dataclasses import dataclass from typing import ClassVar @@ -194,71 +194,101 @@ class DifyConfigLayer(PlainLayer[DifyConfigDeps, DifyConfigLayerConfig, DifyConf if not self.config.mentioned_skill_names and not self.config.mentioned_file_names: return - tasks = [ - *(self._pull_mentioned_skill(name) for name in self.config.mentioned_skill_names), - *(self._pull_mentioned_file(name) for name in self.config.mentioned_file_names), - ] - await asyncio.gather(*tasks) + if names := self.config.mentioned_skill_names: + output = await self._run_mentioned_pull( + script=self._build_shell_skill_pull_script(names), + target_kind="skill", + ) + self.runtime_state.pulled_skill_outputs = _parse_skill_pull_outputs(output, names) - async def _pull_mentioned_skill(self, name: str) -> None: + if names := self.config.mentioned_file_names: + output = await self._run_mentioned_pull( + script=self._build_shell_file_pull_script(names), + target_kind="file", + ) + self.runtime_state.pulled_file_outputs = _parse_file_pull_outputs(output, names) + + async def _run_mentioned_pull(self, *, script: str, target_kind: str) -> str: result = await self.deps.shell.run_remote_script( - self._build_shell_skill_pull_script(name), + script, inject_agent_stub_env=True, ) if result.exit_code != 0: raise DifyConfigLayerError( - "config mentioned skill pull failed in shell: " + f"config mentioned {target_kind} pull failed in shell: " f"{result.status} exit_code={result.exit_code}\n{result.output}" ) if not result.output_complete: reason = result.incomplete_reason or "unknown" raise DifyConfigLayerError( - f"config mentioned skill pull output was incomplete before the payload finished: {reason}" + f"config mentioned {target_kind} pull output was incomplete before the payload finished: {reason}" ) output = result.output.strip() if not output: + raise DifyConfigLayerError(f"missing pull output for mentioned config {target_kind}s") + return output + + def _build_shell_skill_pull_script(self, names: list[str]) -> str: + targets = " ".join(shlex.quote(name) for name in names) + return f"set -eu\ndify-agent config skills pull --json {targets}" + + def _build_shell_file_pull_script(self, names: list[str]) -> str: + targets = " ".join(shlex.quote(name) for name in names) + return f"set -eu\ndify-agent config files pull --json {targets}" + + +def _parse_skill_pull_outputs(output: str, expected_names: list[str]) -> dict[str, str]: + items = _parse_pull_items(output, target_kind="skill") + parsed: dict[str, str] = {} + for name in expected_names: + item = items.get(name) + if item is None: raise DifyConfigLayerError(f"missing pull output for mentioned config skill {name}") - self.runtime_state.pulled_skill_outputs = { - **self.runtime_state.pulled_skill_outputs, - name: output, - } + directory_path = item.get("directory_path") + skill_md = item.get("skill_md") + if not isinstance(directory_path, str) or not directory_path: + raise DifyConfigLayerError(f"invalid directory path in pull output for mentioned config skill {name}") + if not isinstance(skill_md, str): + raise DifyConfigLayerError(f"invalid skill content in pull output for mentioned config skill {name}") + parsed[name] = f"{directory_path}\n{skill_md}".strip() + return parsed - async def _pull_mentioned_file(self, name: str) -> None: - result = await self.deps.shell.run_remote_script( - self._build_shell_file_pull_script(name), - inject_agent_stub_env=True, - ) - if result.exit_code != 0: - raise DifyConfigLayerError( - "config mentioned file pull failed in shell: " - f"{result.status} exit_code={result.exit_code}\n{result.output}" - ) - if not result.output_complete: - reason = result.incomplete_reason or "unknown" - raise DifyConfigLayerError( - f"config mentioned file pull output was incomplete before the payload finished: {reason}" - ) - output = result.output.strip() - if not output: + +def _parse_file_pull_outputs(output: str, expected_names: list[str]) -> dict[str, str]: + items = _parse_pull_items(output, target_kind="file") + parsed: dict[str, str] = {} + for name in expected_names: + item = items.get(name) + if item is None: raise DifyConfigLayerError(f"missing pull output for mentioned config file {name}") - self.runtime_state.pulled_file_outputs = { - **self.runtime_state.pulled_file_outputs, - name: output, - } + path = item.get("path") + if not isinstance(path, str) or not path: + raise DifyConfigLayerError(f"invalid path in pull output for mentioned config file {name}") + parsed[name] = path + return parsed - def _build_shell_skill_pull_script(self, name: str) -> str: - lines = [ - "set -eu", - f"dify-agent config skills pull {shlex.quote(name)}", - ] - return "\n".join(lines) - def _build_shell_file_pull_script(self, name: str) -> str: - lines = [ - "set -eu", - f"dify-agent config files pull {shlex.quote(name)}", - ] - return "\n".join(lines) +def _parse_pull_items(output: str, *, target_kind: str) -> dict[str, dict[str, object]]: + try: + payload: object = json.loads(output) + except json.JSONDecodeError as exc: + raise DifyConfigLayerError(f"invalid JSON pull output for mentioned config {target_kind}s") from exc + if not isinstance(payload, dict): + raise DifyConfigLayerError(f"invalid pull output for mentioned config {target_kind}s") + raw_items = payload.get("items") + if not isinstance(raw_items, list): + raise DifyConfigLayerError(f"missing items in pull output for mentioned config {target_kind}s") + + items: dict[str, dict[str, object]] = {} + for raw_item in raw_items: + if not isinstance(raw_item, dict): + raise DifyConfigLayerError(f"invalid item in pull output for mentioned config {target_kind}s") + item = {str(key): value for key, value in raw_item.items()} + name = item.get("name") + if not isinstance(name, str) or not name: + raise DifyConfigLayerError(f"missing item name in pull output for mentioned config {target_kind}s") + items[name] = item + return items def _format_command_output(command: str, output: str) -> str: diff --git a/dify-agent/src/dify_agent/layers/runtime/__init__.py b/dify-agent/src/dify_agent/layers/runtime/__init__.py new file mode 100644 index 00000000000..872ed9120c4 --- /dev/null +++ b/dify-agent/src/dify_agent/layers/runtime/__init__.py @@ -0,0 +1,4 @@ +from .configs import DIFY_RUNTIME_LAYER_TYPE_ID, DifyRuntimeLayerConfig +from .layer import DifyRuntimeLayer + +__all__ = ["DIFY_RUNTIME_LAYER_TYPE_ID", "DifyRuntimeLayer", "DifyRuntimeLayerConfig"] diff --git a/dify-agent/src/dify_agent/layers/runtime/configs.py b/dify-agent/src/dify_agent/layers/runtime/configs.py new file mode 100644 index 00000000000..f7f3a576dbb --- /dev/null +++ b/dify-agent/src/dify_agent/layers/runtime/configs.py @@ -0,0 +1,17 @@ +"""Client-safe config for operation-scoped RuntimeLease acquisition.""" + +from typing import ClassVar, Final + +from agenton.layers import LayerConfig +from pydantic import ConfigDict, Field + +DIFY_RUNTIME_LAYER_TYPE_ID: Final[str] = "dify.runtime" + + +class DifyRuntimeLayerConfig(LayerConfig): + backend_binding_ref: str = Field(min_length=1) + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +__all__ = ["DIFY_RUNTIME_LAYER_TYPE_ID", "DifyRuntimeLayerConfig"] diff --git a/dify-agent/src/dify_agent/layers/runtime/layer.py b/dify-agent/src/dify_agent/layers/runtime/layer.py new file mode 100644 index 00000000000..e6fbaf70aff --- /dev/null +++ b/dify-agent/src/dify_agent/layers/runtime/layer.py @@ -0,0 +1,75 @@ +"""Agenton layer exposing one operation-scoped RuntimeLease.""" + +from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager +from dataclasses import dataclass, field +from typing import ClassVar + +from agenton.layers import EmptyRuntimeState, NoLayerDeps, PlainLayer +from typing_extensions import Self, override + +from dify_agent.layers.runtime.configs import DIFY_RUNTIME_LAYER_TYPE_ID, DifyRuntimeLayerConfig +from dify_agent.runtime_backend import ExecutionBindingBackend, RuntimeLease +from dify_agent.runtime_backend.leases import open_runtime_lease + + +@dataclass(slots=True) +class DifyRuntimeLayer(PlainLayer[NoLayerDeps, DifyRuntimeLayerConfig, EmptyRuntimeState]): + """Acquire/release a Binding without owning its product lifecycle.""" + + type_id: ClassVar[str | None] = DIFY_RUNTIME_LAYER_TYPE_ID + config: DifyRuntimeLayerConfig + backend: ExecutionBindingBackend + _lease: RuntimeLease | None = field(default=None, init=False) + + @classmethod + @override + def from_config(cls, config: DifyRuntimeLayerConfig) -> Self: + del config + raise TypeError("DifyRuntimeLayer requires a server-injected ExecutionBindingBackend") + + @classmethod + def from_config_with_backend( + cls, + config: DifyRuntimeLayerConfig, + *, + backend: ExecutionBindingBackend, + ) -> Self: + return cls(config=DifyRuntimeLayerConfig.model_validate(config), backend=backend) + + @property + def lease(self) -> RuntimeLease: + if self._lease is None: + raise RuntimeError("DifyRuntimeLayer lease is only available inside resource_context()") + return self._lease + + @override + @asynccontextmanager + async def resource_context(self) -> AsyncGenerator[None]: + if self._lease is not None: + raise RuntimeError("DifyRuntimeLayer resource_context() is already active") + async with open_runtime_lease(self.backend, self.config.backend_binding_ref) as lease: + self._lease = lease + try: + yield + finally: + self._lease = None + + @override + async def on_context_create(self) -> None: + _ = self.lease + + @override + async def on_context_resume(self) -> None: + _ = self.lease + + @override + async def on_context_suspend(self) -> None: + _ = self.lease + + @override + async def on_context_delete(self) -> None: + _ = self.lease + + +__all__ = ["DifyRuntimeLayer"] diff --git a/dify-agent/src/dify_agent/layers/shell/__init__.py b/dify-agent/src/dify_agent/layers/shell/__init__.py index 69af3150c76..34800d69471 100644 --- a/dify-agent/src/dify_agent/layers/shell/__init__.py +++ b/dify-agent/src/dify_agent/layers/shell/__init__.py @@ -1,8 +1,8 @@ """Client-safe exports for the Dify shell layer DTOs. -The runtime layer implementation lives in ``layer.py`` and imports shellctl -client code plus server-side lifecycle behavior. Keep this package root -import-safe for client code that only needs to build run requests. +The runtime layer implementation lives in ``layer.py`` and consumes a +server-injected active Sandbox lease. Keep this package root import-safe for +client code that only needs to build run requests. """ from dify_agent.layers.shell.configs import ( @@ -10,7 +10,6 @@ from dify_agent.layers.shell.configs import ( DifyShellCliToolConfig, DifyShellEnvVarConfig, DifyShellLayerConfig, - DifyShellSandboxConfig, DifyShellSecretRefConfig, ) @@ -19,6 +18,5 @@ __all__ = [ "DifyShellCliToolConfig", "DifyShellEnvVarConfig", "DifyShellLayerConfig", - "DifyShellSandboxConfig", "DifyShellSecretRefConfig", ] diff --git a/dify-agent/src/dify_agent/layers/shell/configs.py b/dify-agent/src/dify_agent/layers/shell/configs.py index 63f1faf5d5e..821e02cc89c 100644 --- a/dify-agent/src/dify_agent/layers/shell/configs.py +++ b/dify-agent/src/dify_agent/layers/shell/configs.py @@ -1,10 +1,11 @@ """Client-safe DTOs for the Dify shell Agenton layer. -Server-only shellctl connection settings are injected by the runtime provider -factory. Public config carries product-level Agent Soul settings that must affect -the sandbox workspace itself: CLI tool bootstrap commands, normal environment -variables, secret environment variable names, sandbox-provider metadata, and the -Agent Stub drive ref used by shell-visible drive commands. +Server-only Agent Stub and redaction settings are injected by the runtime +provider factory. The Sandbox dependency supplies the active shellctl data +plane. Public config carries product-level Agent Soul settings that affect the +workspace itself: CLI tool bootstrap commands, normal environment variables, +secret environment variable names, and the Agent Stub drive ref used by +shell-visible drive commands. Sandbox selection is a deployment concern. """ import re @@ -67,24 +68,8 @@ class DifyShellCliToolConfig(BaseModel): return [command for command in (item.strip() for item in value) if command] -class DifyShellSandboxConfig(BaseModel): - """Sandbox provider selection persisted in Agent Soul.""" - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - provider: str | None = Field(default=None, max_length=255) - config: dict[str, object] = Field(default_factory=dict) - - -class DifyShellEnterpriseSandboxConfig(DifyShellSandboxConfig): - """Enterprise sandbox provider configuration.""" - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - gateway_endpoint: str = Field(..., max_length=255) - - class DifyShellLayerConfig(LayerConfig): - """Public config for the shellctl-backed Dify shell layer.""" + """Public product behavior for the Sandbox-backed Dify Shell layer.""" model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") @@ -93,7 +78,6 @@ class DifyShellLayerConfig(LayerConfig): cli_tools: list[DifyShellCliToolConfig] = Field(default_factory=list) env: list[DifyShellEnvVarConfig] = Field(default_factory=list) secret_refs: list[DifyShellSecretRefConfig] = Field(default_factory=list) - sandbox: DifyShellSandboxConfig | None = None redact_patterns: list[str] = Field(default_factory=list) @@ -102,6 +86,5 @@ __all__ = [ "DifyShellCliToolConfig", "DifyShellEnvVarConfig", "DifyShellLayerConfig", - "DifyShellSandboxConfig", "DifyShellSecretRefConfig", ] diff --git a/dify-agent/src/dify_agent/layers/shell/layer.py b/dify-agent/src/dify_agent/layers/shell/layer.py index b38a151ddee..112a7def742 100644 --- a/dify-agent/src/dify_agent/layers/shell/layer.py +++ b/dify-agent/src/dify_agent/layers/shell/layer.py @@ -1,40 +1,13 @@ -"""Shell runtime layer backed by a live shell provider resource. - -Shell command execution requires a bound execution-context layer with a safe -``agent_id``. The layer uses the current bound execution context to run -commands with ``HOME=/`` and a home-rooted workspace path. The -persisted runtime state intentionally keeps the historical -``~/workspace/`` identity so existing session snapshots stay -compatible while live command execution no longer depends on the sandbox user's -ambient home directory. Entering or re-entering the layer re-ensures the live -home/workspace directories for the currently bound ``agent_id`` before user -commands are sent. - -Sandbox lifecycle: - The shell provider exposes four operations: ``create``, ``attach``, - ``suspend``, and ``delete``. On the first run (no ``sandbox_id`` in - runtime state) ``resource_context()`` calls ``create()`` to provision a - new sandbox and persists the returned ``sandbox_id``. On subsequent runs - it calls ``attach(sandbox_id)`` to re-connect to the existing sandbox. - If the sandbox has expired (the provider raises ``SandboxExpiredError``), - the error propagates to the caller — the user must start a new session. - On normal exit (suspend) the resource is detached via ``suspend()``, - keeping the sandbox alive. On final cleanup (``on_context_delete``) the - resource is destroyed via ``delete()``. This allows the enterprise - provider to reuse the same sandbox pod across conversation turns. -""" +"""Shell tools over the data plane exposed by the active Runtime layer.""" from __future__ import annotations -from collections.abc import AsyncGenerator, Sequence -from contextlib import asynccontextmanager +from collections.abc import Sequence import json import logging import re -import secrets -import time from dataclasses import dataclass, field -from typing import ClassVar, Literal, NotRequired, Protocol, TypedDict, runtime_checkable +from typing import ClassVar, NotRequired, Protocol, TypedDict, runtime_checkable from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt, field_validator, model_validator from pydantic_ai import Tool @@ -54,14 +27,15 @@ from dify_agent.adapters.shell.protocols import ( ShellCommandProtocol, ShellCommandResult, ShellPromptObservation, - ShellProviderProtocol, - ShellResourceProtocol, ) from dify_agent.agent_stub.protocol import AGENT_STUB_AUTH_JWE_ENV_VAR from dify_agent.agent_stub.shell_env import ShellAgentStubTokenFactory, build_shell_agent_stub_env from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig +from dify_agent.layers.runtime.layer import DifyRuntimeLayer from dify_agent.layers.shell.configs import DIFY_SHELL_LAYER_TYPE_ID, DifyShellLayerConfig from dify_agent.layers.shell.output_text import normalized_output_text, utf8_prefix, utf8_suffix +from dify_agent.runtime.command_runner import execute_complete_with_commands +from dify_agent.runtime_backend import RuntimeLease logger = logging.getLogger(__name__) @@ -74,15 +48,7 @@ class _HasErrorCode(Protocol): DEFAULT_TIMEOUT_SECONDS = 30.0 DEFAULT_TERMINATE_GRACE_SECONDS = 10.0 -_WORKSPACE_ROOT = "~/workspace" -_WORKSPACE_DIR_NAME = "workspace" -_WORKSPACE_COLLISION_EXIT_CODE = 17 -_SESSION_TIME_HEX_MASK = 0xFFFFF -_SESSION_RANDOM_HEX_LENGTH = 2 -_SESSION_ID_ATTEMPT_LIMIT = 256 -_SESSION_ID_PATTERN = re.compile(r"^[0-9a-f]{7}$") -_AGENT_HOME_SEGMENT_PATTERN = re.compile(r"^[A-Za-z0-9._-]+$") -_SHELL_OUTPUT_PROMPT_EDGE_BYTES = 8 * 1024 +_SHELL_OUTPUT_PROMPT_EDGE_BYTES = 4 * 1024 _SHELLCTL_OUTPUT_LIMIT_BYTES = 2 * _SHELL_OUTPUT_PROMPT_EDGE_BYTES _REMOTE_COMPLETE_OUTPUT_MAX_BYTES = 1024 * 1024 _REMOTE_COMMAND_TIMEOUT_SECONDS = 60.0 @@ -142,8 +108,8 @@ Installed CLI: Workspace persistence rules: -- The current workspace cwd is stable during this run, but it is temporary and may be deleted later. -- Do not treat files in the current workspace cwd as persisted state. +- The current workspace cwd is stable across runs for the current product session. +- Workspace files are working data, not Agent configuration, and are removed when that product session ends. - In build mode, config changes persist only after you run the matching `dify-agent config ...` mutation command. - Shell file edits alone do not save Agent config files, skills, env, or notes. - In non-build modes, local shell changes are not a persistence mechanism for Agent configuration. @@ -202,24 +168,15 @@ type ShellInterruptToolResult = str | ShellToolErrorObservation class DifyShellLayerDeps(LayerDeps): execution_context: PlainLayer[NoLayerDeps, DifyExecutionContextLayerConfig, EmptyRuntimeState] | None # pyright: ignore[reportUninitializedInstanceVariable] + runtime: DifyRuntimeLayer # pyright: ignore[reportUninitializedInstanceVariable] class DifyShellRuntimeState(BaseModel): - session_id: str | None = None - workspace_cwd: str | None = None - sandbox_id: str | None = None job_ids: list[str] = Field(default_factory=list) job_offsets: dict[str, NonNegativeInt] = Field(default_factory=dict) model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid", validate_assignment=True) - @field_validator("session_id") - @classmethod - def validate_session_id(cls, value: str | None) -> str | None: - if value is None: - return value - return _validated_session_id(value) - @field_validator("job_ids") @classmethod def validate_job_ids(cls, value: list[str]) -> list[str]: @@ -228,13 +185,7 @@ class DifyShellRuntimeState(BaseModel): return value @model_validator(mode="after") - def validate_workspace_and_offsets(self) -> Self: - if self.workspace_cwd is not None: - if self.session_id is None: - raise ValueError("workspace_cwd requires a matching session_id.") - expected_workspace = _workspace_cwd(self.session_id) - if self.workspace_cwd != expected_workspace: - raise ValueError(f"workspace_cwd must equal {expected_workspace!r} for session_id {self.session_id!r}.") + def validate_job_offsets(self) -> Self: unknown_offset_job_ids = set(self.job_offsets) - set(self.job_ids) if unknown_offset_job_ids: names = ", ".join(sorted(unknown_offset_job_ids)) @@ -247,46 +198,43 @@ CompleteRemoteCommandResult = CompleteShellCommandResult @dataclass(slots=True) class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerConfig, DifyShellRuntimeState]): + """Expose Shell tools over the active RuntimeLease without owning it. + + Create optionally bootstraps configured CLI tools in the lease's Workspace. + Suspend and delete best-effort remove tracked shellctl jobs, then clear job + ids and offsets so they do not persist across requests. Commands, files, + Home, and cwd come only from ``DifyRuntimeLayer.lease``. Persistent Binding + and Workspace lifecycle remains exclusively owned by Dify API. + """ + type_id: ClassVar[str | None] = DIFY_SHELL_LAYER_TYPE_ID config: DifyShellLayerConfig - shell_provider: ShellProviderProtocol - shell_home_root: str = "/home" shell_redact_patterns: list[str] = field(default_factory=list) agent_stub_api_base_url: str | None = None agent_stub_token_factory: ShellAgentStubTokenFactory | None = None - _shell_resource: ShellResourceProtocol | None = None - _resource_should_delete: bool = False @classmethod @override def from_config(cls, config: DifyShellLayerConfig) -> Self: del config - raise TypeError("DifyShellLayer requires a shell provider and must use a provider factory.") + raise TypeError("DifyShellLayer requires server-injected settings and must use a provider factory.") @classmethod def from_config_with_settings( cls, config: DifyShellLayerConfig, *, - shell_provider: ShellProviderProtocol | None, - shell_home_root: str = "/home", shell_redact_patterns: list[str] | None = None, agent_stub_api_base_url: str | None = None, agent_stub_token_factory: ShellAgentStubTokenFactory | None = None, ) -> Self: - if shell_provider is None: - raise ValueError("DifyShellLayer requires a non-null shell provider when the 'dify.shell' layer is used.") - layer = cls( + return cls( config=config, - shell_provider=shell_provider, - shell_home_root=_normalize_shell_home_root(shell_home_root), shell_redact_patterns=shell_redact_patterns or [], agent_stub_api_base_url=agent_stub_api_base_url, agent_stub_token_factory=agent_stub_token_factory, ) - layer.bind_deps({}) - return layer @property @override @@ -308,95 +256,31 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC Tool(self._tool_interrupt, name="shell_interrupt"), ] - @override - @asynccontextmanager - async def resource_context(self) -> AsyncGenerator[None]: - """Acquire the live shell resource for one run invocation. - - On the first run (no ``sandbox_id`` in runtime state) the provider's - ``create()`` is called to provision a new sandbox. On subsequent runs - ``attach(sandbox_id)`` re-connects to the existing sandbox. If the - sandbox has expired (``SandboxExpiredError``), the error propagates - to the caller — the user must start a new session. - - On exit, ``suspend()`` is called by default to keep the sandbox alive. - If ``on_context_delete()`` ran (setting ``_resource_should_delete``), - ``delete()`` is called instead to destroy the sandbox. - """ - if self._shell_resource is not None: - raise RuntimeError("DifyShellLayer resource_context() is already active for this layer instance.") - sandbox_id = self.runtime_state.sandbox_id - if sandbox_id is not None: - resource = await self.shell_provider.attach(sandbox_id) - else: - resource = await self.shell_provider.create() - self.runtime_state = DifyShellRuntimeState.model_validate( - { - **self.runtime_state.model_dump(mode="python"), - "sandbox_id": resource.sandbox_id, - } - ) - self._shell_resource = resource - self._resource_should_delete = False - try: - yield - finally: - self._shell_resource = None - if self._resource_should_delete: - await resource.delete() - else: - await resource.suspend() - @override async def on_context_create(self) -> None: - _ = self._require_resource() - session_id: str | None = None - try: - session_id, workspace_cwd = await self._allocate_workspace() - await self._bootstrap_workspace(session_id) - except BaseException: - if session_id is not None: - await self._cleanup_workspace_best_effort(session_id) - self._resource_should_delete = True - raise - self.runtime_state = DifyShellRuntimeState.model_validate( - { - **self.runtime_state.model_dump(mode="python"), - "session_id": session_id, - "workspace_cwd": workspace_cwd, - } - ) + bootstrap_script = _workspace_bootstrap_script(self.config) + if not bootstrap_script: + return + result = await self._run_internal_script_complete(bootstrap_script, cwd=self._require_workspace_cwd()) + if result.exit_code != 0 or not result.output_complete: + raise RuntimeError( + f"Failed to bootstrap shell workspace {self._require_workspace_cwd()}: " + f"{result.status} exit_code={result.exit_code}" + ) @override async def on_context_resume(self) -> None: _ = self._require_resource() - session_id, _workspace_cwd = self._require_session_identity() - await self._ensure_live_workspace_exists(session_id) @override async def on_context_suspend(self) -> None: - _ = self._require_resource() + await self._delete_tracked_jobs_best_effort(self.runtime_state.job_ids) + self._clear_tracked_jobs() @override async def on_context_delete(self) -> None: - _ = self._require_resource() - identity = self._try_session_identity() - if identity is not None: - session_id, _workspace_cwd = identity - result = await self._run_internal_script_complete( - _workspace_cleanup_script(session_id=session_id), cwd=None - ) - if result.exit_code != 0 or not result.output_complete: - logger.warning( - "Shell workspace cleanup for session %s ended with status=%s exit_code=%s output_complete=%s.", - session_id, - result.status, - result.exit_code, - result.output_complete, - ) await self._delete_tracked_jobs_best_effort(self.runtime_state.job_ids) self._clear_tracked_jobs() - self._resource_should_delete = True async def _tool_run(self, script: str, timeout: float = DEFAULT_TIMEOUT_SECONDS) -> ShellRunToolResult: try: @@ -426,7 +310,7 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC except (RuntimeError, ValueError) as exc: return _tool_error_from_exception(exc) except Exception as exc: - return _tool_unexpected_error("shell_run", exc, session_id=self.runtime_state.session_id) + return _tool_unexpected_error("shell_run", exc) async def _tool_wait(self, job_id: str, timeout: float = DEFAULT_TIMEOUT_SECONDS) -> ShellRunToolResult: try: @@ -452,7 +336,7 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC except (RuntimeError, ValueError) as exc: return _tool_error_from_exception(exc, job_id=job_id) except Exception as exc: - return _tool_unexpected_error("shell_wait", exc, session_id=self.runtime_state.session_id, job_id=job_id) + return _tool_unexpected_error("shell_wait", exc, job_id=job_id) async def _tool_input(self, job_id: str, text: str, timeout: float = DEFAULT_TIMEOUT_SECONDS) -> ShellRunToolResult: try: @@ -478,7 +362,7 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC except (RuntimeError, ValueError) as exc: return _tool_error_from_exception(exc, job_id=job_id) except Exception as exc: - return _tool_unexpected_error("shell_input", exc, session_id=self.runtime_state.session_id, job_id=job_id) + return _tool_unexpected_error("shell_input", exc, job_id=job_id) async def _tool_interrupt( self, @@ -498,16 +382,14 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC output_path = (await self._require_resource().commands.tail(job_id)).output_path except (RuntimeError, ValueError) as exc: logger.warning( - "Failed to fetch output path for interrupted shell job %s in session %s: %s", + "Failed to fetch output path for interrupted shell job %s: %s", job_id, - self.runtime_state.session_id, exc, ) except Exception: logger.exception( - "Failed to fetch output path for interrupted shell job %s in session %s", + "Failed to fetch output path for interrupted shell job %s", job_id, - self.runtime_state.session_id, ) return _tagged_shell_observation( _metadata_dict( @@ -522,9 +404,7 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC except (RuntimeError, ValueError) as exc: return _tool_error_from_exception(exc, job_id=job_id) except Exception as exc: - return _tool_unexpected_error( - "shell_interrupt", exc, session_id=self.runtime_state.session_id, job_id=job_id - ) + return _tool_unexpected_error("shell_interrupt", exc, job_id=job_id) async def run_remote_script_complete( self, @@ -559,45 +439,6 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC inject_agent_stub_env=inject_agent_stub_env, ) - async def _allocate_workspace(self) -> tuple[str, str]: - for _attempt in range(_SESSION_ID_ATTEMPT_LIMIT): - session_id = _generate_session_id() - result = await self._run_internal_script_complete(_workspace_mkdir_script(session_id=session_id), cwd=None) - if result.exit_code == _WORKSPACE_COLLISION_EXIT_CODE: - continue - if result.exit_code != 0 or not result.output_complete: - raise RuntimeError( - f"Failed to create shell workspace {_workspace_cwd(session_id)}: " - + f"{result.status} exit_code={result.exit_code}" - ) - return session_id, _workspace_cwd(session_id) - raise RuntimeError("Failed to allocate a unique shell workspace session id after 256 attempts.") - - async def _bootstrap_workspace(self, session_id: str) -> None: - bootstrap_script = _workspace_bootstrap_script(self.config) - if not bootstrap_script: - return - workspace_cwd = _workspace_cwd_for_home(self._shell_home_dir(), session_id) - result = await self._run_internal_script_complete(bootstrap_script, cwd=workspace_cwd) - if result.exit_code != 0 or not result.output_complete: - raise RuntimeError( - f"Failed to bootstrap shell workspace {workspace_cwd}: {result.status} exit_code={result.exit_code}" - ) - - async def _cleanup_workspace_best_effort(self, session_id: str) -> None: - try: - _ = await self._run_internal_script_complete(_workspace_cleanup_script(session_id=session_id), cwd=None) - except (RuntimeError, ValueError) as exc: - logger.warning("Failed to remove shell workspace for session %s after create failure: %s", session_id, exc) - - async def _ensure_live_workspace_exists(self, session_id: str) -> None: - result = await self._run_internal_script_complete(_workspace_ensure_script(session_id=session_id), cwd=None) - if result.exit_code != 0 or not result.output_complete: - raise RuntimeError( - f"Failed to ensure shell workspace {_workspace_cwd_for_home(self._shell_home_dir(), session_id)} " - + f"exists: {result.status} exit_code={result.exit_code}" - ) - async def _run_internal_script_complete( self, script: str, @@ -613,36 +454,11 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC max_output_bytes=_REMOTE_COMPLETE_OUTPUT_MAX_BYTES, ) - def _require_resource(self) -> ShellResourceProtocol: - if self._shell_resource is None: - raise RuntimeError( - "DifyShellLayer requires an active shell resource inside resource_context(); " - + "enter the layer through Agenton or wrap direct hook/tool usage in resource_context()." - ) - return self._shell_resource + def _require_resource(self) -> RuntimeLease: + return self.deps.runtime.lease def _require_workspace_cwd(self) -> str: - session_id, _workspace_cwd = self._require_session_identity() - return _workspace_cwd_for_home(self._shell_home_dir(), session_id) - - def _require_session_identity(self) -> tuple[str, str]: - identity = self._try_session_identity() - if identity is None: - raise ValueError("DifyShellLayer runtime state is missing session_id or workspace_cwd.") - session_id, workspace_cwd = identity - expected_workspace = _workspace_cwd(session_id) - if workspace_cwd != expected_workspace: - raise ValueError( - f"DifyShellLayer runtime state has inconsistent workspace_cwd {workspace_cwd!r}; expected {expected_workspace!r}." - ) - return session_id, workspace_cwd - - def _try_session_identity(self) -> tuple[str, str] | None: - session_id = self.runtime_state.session_id - workspace_cwd = self.runtime_state.workspace_cwd - if session_id is None or workspace_cwd is None: - return None - return session_id, workspace_cwd + return self._require_resource().layout.workspace_dir def _ensure_tracked_job(self, job_id: str) -> None: if job_id not in self.runtime_state.job_ids: @@ -669,9 +485,8 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC await commands.delete(job_id, force=True) except RuntimeError as exc: logger.warning( - "Failed to delete shell job %s for session %s: %s", + "Failed to delete shell job %s: %s", job_id, - self.runtime_state.session_id, exc, ) @@ -679,23 +494,6 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC self.runtime_state.job_offsets = {} self.runtime_state.job_ids = [] - def _shell_home_dir(self) -> str: - return _shell_home_dir_for_agent_id( - self._require_current_execution_agent_id(), - shell_home_root=self.shell_home_root, - ) - - def _current_execution_agent_id(self) -> str | None: - execution_context_layer = self.deps.execution_context - execution_context = execution_context_layer.config if execution_context_layer is not None else None - return execution_context.agent_id if execution_context is not None else None - - def _require_current_execution_agent_id(self) -> str: - agent_id = self._current_execution_agent_id() - if agent_id is None: - raise ValueError("ShellLayer command execution requires execution_context.agent_id.") - return _validated_agent_home_segment(agent_id) - def _build_shell_command_env( self, *, @@ -703,7 +501,7 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC require_agent_stub_env: bool = False, ) -> dict[str, str]: env = _shell_config_env(self.config) - env["HOME"] = self._shell_home_dir() + env["HOME"] = self._require_resource().layout.home_dir if not include_agent_stub_env: return env execution_context_layer = self.deps.execution_context @@ -713,7 +511,7 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC agent_stub_drive_ref=self.config.agent_stub_drive_ref, execution_context=execution_context, token_factory=self.agent_stub_token_factory, - session_id=self.runtime_state.session_id, + session_id=None, ) if agent_stub_env is None: if not require_agent_stub_env: @@ -747,77 +545,6 @@ class DifyShellLayer(PydanticAILayer[DifyShellLayerDeps, object, DifyShellLayerC return text -async def execute_complete_with_commands( - commands: ShellCommandProtocol, - script: str, - *, - cwd: str | None, - env: dict[str, str] | None, - timeout: float, - max_output_bytes: int, -) -> CompleteShellCommandResult: - deadline = time.monotonic() + timeout - job_id: str | None = None - result: ShellCommandResult | None = None - output_parts: list[str] = [] - captured_bytes = 0 - incomplete_reason: Literal["output_limit", "timeout"] | None = None - try: - result = await commands.run(script, cwd=cwd, env=env, timeout=_remaining_time(deadline)) - job_id = result.job_id - while True: - remaining_bytes = max(max_output_bytes - captured_bytes, 0) - limited_output = utf8_prefix(result.output, remaining_bytes) - output_parts.append(limited_output) - captured_bytes += len(limited_output.encode("utf-8")) - if limited_output != result.output: - incomplete_reason = "output_limit" - break - if captured_bytes >= max_output_bytes and (result.truncated or not result.done): - incomplete_reason = "output_limit" - break - if result.truncated: - result = await commands.read_output(result.job_id, offset=result.offset) - continue - if result.done: - break - remaining_time = _remaining_time(deadline) - if remaining_time <= 0.0: - incomplete_reason = "timeout" - break - result = await commands.wait(result.job_id, offset=result.offset, timeout=remaining_time) - - assert result is not None - final_status = result.status - final_done = result.done - final_exit_code = result.exit_code - final_offset = result.offset - final_output_path = result.output_path - if incomplete_reason is not None and not result.done: - terminal_status = await commands.interrupt(result.job_id, grace_seconds=DEFAULT_TERMINATE_GRACE_SECONDS) - final_status = terminal_status.status - final_done = terminal_status.done - final_exit_code = terminal_status.exit_code - final_offset = terminal_status.offset - return CompleteShellCommandResult( - job_id=result.job_id, - status=final_status, - done=final_done, - exit_code=final_exit_code, - output="".join(output_parts), - output_complete=incomplete_reason is None, - incomplete_reason=incomplete_reason, - offset=final_offset, - output_path=final_output_path, - ) - finally: - if job_id is not None: - try: - await commands.delete(job_id, force=True) - except RuntimeError as exc: - logger.warning("Failed to delete transient shell job %s: %s", job_id, exc) - - async def render_prompt_observation_from_result( commands: ShellCommandProtocol, result: ShellCommandResult, @@ -897,58 +624,19 @@ def _tool_unexpected_error( tool_name: str, exc: Exception, *, - session_id: str | None, job_id: str | None = None, ) -> ShellToolErrorObservation: # Unexpected Exception still becomes a tool observation so one shell tool # failure does not abort the agent loop, but it is logged with traceback for # debugging. BaseException is intentionally not caught by callers. logger.exception( - "Unexpected shell tool failure: tool=%s session_id=%s job_id=%s", + "Unexpected shell tool failure: tool=%s job_id=%s", tool_name, - session_id, job_id, ) return _tool_error_from_exception(exc, job_id=job_id) -def _generate_session_id() -> str: - time_component = int(time.time()) & _SESSION_TIME_HEX_MASK - random_component = secrets.token_hex(1) - if len(random_component) != _SESSION_RANDOM_HEX_LENGTH: - raise RuntimeError("Expected a one-byte random hex suffix for Dify shell session ids.") - return f"{time_component:05x}{random_component}" - - -def _workspace_cwd(session_id: str) -> str: - return f"{_WORKSPACE_ROOT}/{_validated_session_id(session_id)}" - - -def _normalize_shell_home_root(shell_home_root: str) -> str: - stripped = shell_home_root.strip().rstrip("/") - if not stripped: - raise ValueError("shell_home_root must not be empty") - if not stripped.startswith("/"): - raise ValueError("shell_home_root must be an absolute path") - return stripped - - -def _shell_home_dir_for_agent_id(agent_id: str | None, *, shell_home_root: str = "/home") -> str: - if agent_id is None: - raise ValueError("ShellLayer command execution requires execution_context.agent_id.") - return f"{_normalize_shell_home_root(shell_home_root)}/{_validated_agent_home_segment(agent_id)}" - - -def _validated_agent_home_segment(agent_id: str) -> str: - if agent_id in {".", ".."} or not _AGENT_HOME_SEGMENT_PATTERN.fullmatch(agent_id): - raise ValueError("execution_context.agent_id must be a safe single path segment for shell HOME.") - return agent_id - - -def _workspace_cwd_for_home(home_dir: str, session_id: str) -> str: - return f"{home_dir}/{_WORKSPACE_DIR_NAME}/{_validated_session_id(session_id)}" - - def _workspace_bootstrap_script(config: DifyShellLayerConfig) -> str: install_commands = [command for tool in config.cli_tools for command in tool.install_commands] if not install_commands: @@ -966,12 +654,6 @@ def _shell_config_export_lines(config: DifyShellLayerConfig) -> list[str]: for tool in config.cli_tools: for secret_ref in tool.secret_refs: lines.append(f'export {secret_ref.name}="${{{secret_ref.name}:-}}"') - if config.sandbox is not None: - if config.sandbox.provider: - lines.append(f"export DIFY_SANDBOX_PROVIDER={_shquote(config.sandbox.provider)}") - if config.sandbox.config: - sandbox_config = json.dumps(config.sandbox.config, ensure_ascii=True, sort_keys=True) - lines.append(f"export DIFY_SANDBOX_CONFIG_JSON={_shquote(sandbox_config)}") return lines @@ -992,35 +674,10 @@ def _wrap_user_script(script: str, config: DifyShellLayerConfig) -> str: return "\n".join([*lines, script]) -def _workspace_mkdir_script(*, session_id: str) -> str: - safe_session_id = _validated_session_id(session_id) - workspace_dir = f"$HOME/workspace/{safe_session_id}" - return ( - 'mkdir -p "$HOME/workspace"; ' - f'if mkdir "{workspace_dir}"; then exit 0; fi; ' - f'if [ -e "{workspace_dir}" ]; then exit {_WORKSPACE_COLLISION_EXIT_CODE}; fi; ' - "exit 1" - ) - - -def _workspace_cleanup_script(*, session_id: str) -> str: - return f'rm -rf -- "$HOME/workspace/{_validated_session_id(session_id)}"' - - -def _workspace_ensure_script(*, session_id: str) -> str: - return f'mkdir -p "$HOME/workspace/{_validated_session_id(session_id)}"' - - def _shquote(value: str) -> str: return "'" + value.replace("'", "'\\''") + "'" -def _validated_session_id(session_id: str) -> str: - if not _SESSION_ID_PATTERN.fullmatch(session_id): - raise ValueError("session_id must match the 5+2 lowercase hex format '<5 hex><2 hex>'.") - return session_id - - def _deduplicate_preserving_order(values: Sequence[str]) -> list[str]: seen: set[str] = set() result: list[str] = [] @@ -1037,10 +694,6 @@ def _tagged_shell_observation(metadata: dict[str, object], output: str) -> str: return f"\n{compact_metadata}\n\n\n\n{output}\n" -def _remaining_time(deadline: float) -> float: - return max(0.0, deadline - time.monotonic()) - - __all__ = [ "CompleteRemoteCommandResult", "DifyShellLayer", @@ -1048,6 +701,5 @@ __all__ = [ "DifyShellRuntimeState", "DEFAULT_TERMINATE_GRACE_SECONDS", "DEFAULT_TIMEOUT_SECONDS", - "execute_complete_with_commands", "render_prompt_observation_from_result", ] diff --git a/dify-agent/src/dify_agent/protocol/__init__.py b/dify-agent/src/dify_agent/protocol/__init__.py index 2fc5bd517ed..0eb18b57a36 100644 --- a/dify-agent/src/dify_agent/protocol/__init__.py +++ b/dify-agent/src/dify_agent/protocol/__init__.py @@ -28,6 +28,7 @@ from .schemas import ( RunEventsResponse, RunFailedEvent, RunFailedEventData, + RunFailureType, RunLayerSpec, RunStartedEvent, RunStatus, @@ -37,36 +38,53 @@ from .schemas import ( normalize_composition, utc_now, ) -from .sandbox import ( - RuntimeLayerSpec, - SandboxFileEntry, - SandboxListRequest, - SandboxListResponse, - SandboxLocator, - SandboxReadRequest, - SandboxReadResponse, - SandboxUploadRequest, - SandboxUploadResponse, - SandboxUploadedFile, - build_sandbox_locator_from_layer_specs, - build_sandbox_locator_from_run_request, - extract_runtime_layer_specs, +from .execution_binding import ( + CreateExecutionBindingRequest, + CreateExecutionBindingResponse, + DestroyExecutionBindingRequest, +) +from .binding_file import ( + BindingFileDownloadRequest, + BindingFileDownloadResponse, + BindingFileEntry, + BindingFileListRequest, + BindingFileListResponse, + BindingFileReadRequest, + BindingFileReadResponse, +) +from .home_snapshot import ( + CreateHomeSnapshotFromBindingRequest, + DeleteHomeSnapshotRequest, + HomeSnapshotResponse, ) __all__ = [ "BaseRunEvent", + "BindingFileDownloadRequest", + "BindingFileDownloadResponse", + "BindingFileEntry", + "BindingFileListRequest", + "BindingFileListResponse", + "BindingFileReadRequest", + "BindingFileReadResponse", "AgentRunUsage", "CancelRunRequest", "CancelRunResponse", "CreateRunRequest", "CreateRunResponse", + "CreateExecutionBindingRequest", + "CreateExecutionBindingResponse", + "CreateHomeSnapshotFromBindingRequest", + "DeleteHomeSnapshotRequest", "DeferredToolCallPayload", "DeferredToolResultsPayload", "DIFY_AGENT_HISTORY_LAYER_ID", "DIFY_AGENT_MODEL_LAYER_ID", "DIFY_AGENT_OUTPUT_LAYER_ID", + "DestroyExecutionBindingRequest", "EmptyRunEventData", "LayerExitSignals", + "HomeSnapshotResponse", "PydanticAIStreamRunEvent", "RUN_EVENT_ADAPTER", "RunCancelledEvent", @@ -77,25 +95,13 @@ __all__ = [ "RunEventsResponse", "RunFailedEvent", "RunFailedEventData", + "RunFailureType", "RunLayerSpec", "RunStartedEvent", "RunStatus", "RunStatusResponse", "RunSucceededEvent", "RunSucceededEventData", - "RuntimeLayerSpec", - "SandboxFileEntry", - "SandboxListRequest", - "SandboxListResponse", - "SandboxLocator", - "SandboxReadRequest", - "SandboxReadResponse", - "SandboxUploadRequest", - "SandboxUploadResponse", - "SandboxUploadedFile", - "build_sandbox_locator_from_layer_specs", - "build_sandbox_locator_from_run_request", - "extract_runtime_layer_specs", "normalize_composition", "utc_now", ] diff --git a/dify-agent/src/dify_agent/protocol/binding_file.py b/dify-agent/src/dify_agent/protocol/binding_file.py new file mode 100644 index 00000000000..0de2af7574e --- /dev/null +++ b/dify-agent/src/dify_agent/protocol/binding_file.py @@ -0,0 +1,80 @@ +"""Private Binding file DTOs resolved through an Execution Binding ref.""" + +from typing import ClassVar, Literal + +from pydantic import BaseModel, ConfigDict, Field + +from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig + +_BINDING_FILE_PREVIEW_MAX_BYTES = 262144 + + +class BindingFileEntry(BaseModel): + name: str + type: Literal["file", "dir", "symlink", "other"] + size: int | None = None + mtime: int | None = None + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class BindingFileListRequest(BaseModel): + backend_binding_ref: str = Field(min_length=1) + path: str = "." + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class BindingFileListResponse(BaseModel): + path: str + entries: list[BindingFileEntry] + truncated: bool + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class BindingFileReadRequest(BaseModel): + backend_binding_ref: str = Field(min_length=1) + path: str = Field(min_length=1) + max_bytes: int = Field( + default=_BINDING_FILE_PREVIEW_MAX_BYTES, + ge=1, + le=_BINDING_FILE_PREVIEW_MAX_BYTES, + ) + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class BindingFileReadResponse(BaseModel): + path: str + size: int | None = None + truncated: bool + binary: bool + text: str | None = None + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class BindingFileDownloadRequest(BaseModel): + backend_binding_ref: str = Field(min_length=1) + path: str = Field(min_length=1) + execution_context: DifyExecutionContextLayerConfig + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class BindingFileDownloadResponse(BaseModel): + reference: str + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +__all__ = [ + "BindingFileDownloadRequest", + "BindingFileDownloadResponse", + "BindingFileEntry", + "BindingFileListRequest", + "BindingFileListResponse", + "BindingFileReadRequest", + "BindingFileReadResponse", +] diff --git a/dify-agent/src/dify_agent/protocol/execution_binding.py b/dify-agent/src/dify_agent/protocol/execution_binding.py new file mode 100644 index 00000000000..fcb508d9d01 --- /dev/null +++ b/dify-agent/src/dify_agent/protocol/execution_binding.py @@ -0,0 +1,44 @@ +"""Private DTOs for persistent Execution Binding lifecycle operations.""" + +from typing import ClassVar + +from pydantic import BaseModel, ConfigDict, Field, model_validator + + +class CreateExecutionBindingRequest(BaseModel): + tenant_id: str = Field(min_length=1) + agent_id: str = Field(min_length=1) + binding_id: str = Field(min_length=1) + workspace_id: str = Field(min_length=1) + existing_workspace_ref: str | None = None + home_snapshot_ref: str | None = Field(default=None, min_length=1) + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class CreateExecutionBindingResponse(BaseModel): + binding_ref: str = Field(min_length=1) + workspace_ref: str = Field(min_length=1) + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class DestroyExecutionBindingRequest(BaseModel): + binding_ref: str = Field(min_length=1) + destroy_workspace: bool + workspace_ref: str | None = None + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + @model_validator(mode="after") + def validate_workspace_ref(self) -> "DestroyExecutionBindingRequest": + if self.destroy_workspace and not self.workspace_ref: + raise ValueError("workspace_ref is required when destroy_workspace is true") + return self + + +__all__ = [ + "CreateExecutionBindingRequest", + "CreateExecutionBindingResponse", + "DestroyExecutionBindingRequest", +] diff --git a/dify-agent/src/dify_agent/protocol/home_snapshot.py b/dify-agent/src/dify_agent/protocol/home_snapshot.py new file mode 100644 index 00000000000..5f306a2bb2c --- /dev/null +++ b/dify-agent/src/dify_agent/protocol/home_snapshot.py @@ -0,0 +1,33 @@ +"""Private DTOs for immutable Home Snapshot operations.""" + +from typing import ClassVar + +from pydantic import BaseModel, ConfigDict, Field + + +class CreateHomeSnapshotFromBindingRequest(BaseModel): + tenant_id: str = Field(min_length=1) + agent_id: str = Field(min_length=1) + home_snapshot_id: str = Field(min_length=1) + backend_binding_ref: str = Field(min_length=1) + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class DeleteHomeSnapshotRequest(BaseModel): + snapshot_ref: str = Field(min_length=1) + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +class HomeSnapshotResponse(BaseModel): + snapshot_ref: str = Field(min_length=1) + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") + + +__all__ = [ + "CreateHomeSnapshotFromBindingRequest", + "DeleteHomeSnapshotRequest", + "HomeSnapshotResponse", +] diff --git a/dify-agent/src/dify_agent/protocol/sandbox.py b/dify-agent/src/dify_agent/protocol/sandbox.py deleted file mode 100644 index db375e6a7b5..00000000000 --- a/dify-agent/src/dify_agent/protocol/sandbox.py +++ /dev/null @@ -1,254 +0,0 @@ -"""Public sandbox DTOs shared by the API and Dify Agent backends. - -The sandbox file APIs must rebuild only the minimum runtime needed to re-enter a -prior shell session: ``dify.execution_context`` for Dify-owned identity and -``dify.shell`` for the sandbox workspace itself. ``SandboxLocator`` therefore -contains a safe composition subset plus the matching filtered session snapshot. -Credential-bearing or runtime-only tool layers are intentionally excluded from -persisted runtime specs and from sandbox locators. -""" - -from __future__ import annotations - -from typing import ClassVar, Literal, cast - -from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot -from dify_agent.layers.dify_core_tools import DIFY_CORE_TOOLS_LAYER_TYPE_ID -from dify_agent.layers.dify_plugin import DIFY_PLUGIN_LLM_LAYER_TYPE_ID, DIFY_PLUGIN_TOOLS_LAYER_TYPE_ID -from pydantic import BaseModel, ConfigDict, Field, JsonValue - -from .schemas import CreateRunRequest, RunComposition, RunLayerSpec - -_SENSITIVE_LAYER_TYPES = frozenset( - { - DIFY_PLUGIN_LLM_LAYER_TYPE_ID, - DIFY_PLUGIN_TOOLS_LAYER_TYPE_ID, - DIFY_CORE_TOOLS_LAYER_TYPE_ID, - } -) - - -class RuntimeLayerSpec(BaseModel): - """Persistable non-sensitive layer spec derived from a run composition. - - API-side runtime-session rows store these specs so later cleanup or sandbox - requests can rebuild the minimal layer graph without persisting model or - tool credentials. - """ - - name: str - type: str - deps: dict[str, str] = Field(default_factory=dict) - metadata: dict[str, JsonValue] = Field(default_factory=dict) - config: JsonValue = None - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxLocator(BaseModel): - """Safe subset of one prior run request needed to re-enter a sandbox shell.""" - - composition: RunComposition - session_snapshot: CompositorSessionSnapshot - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxFileEntry(BaseModel): - """One directory entry returned by ``/sandbox/files/list``.""" - - name: str - type: Literal["file", "dir", "symlink", "other"] - size: int | None = None - mtime: int | None = None - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxListRequest(BaseModel): - """Request body for listing a sandbox directory.""" - - locator: SandboxLocator - path: str = "." - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxListResponse(BaseModel): - """Structured sandbox directory listing.""" - - path: str - entries: list[SandboxFileEntry] - truncated: bool - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxReadRequest(BaseModel): - """Request body for reading a sandbox file preview.""" - - locator: SandboxLocator - path: str - max_bytes: int = 262144 - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxReadResponse(BaseModel): - """Text preview returned by ``/sandbox/files/read``.""" - - path: str - size: int | None = None - truncated: bool - binary: bool - text: str | None = None - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxUploadedFile(BaseModel): - """Canonical ToolFile mapping returned after sandbox upload.""" - - transfer_method: Literal["tool_file"] = "tool_file" - reference: str - download_url: str - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxUploadRequest(BaseModel): - """Request body for uploading one sandbox file through the Agent Stub.""" - - locator: SandboxLocator - path: str - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -class SandboxUploadResponse(BaseModel): - """Result returned after sandbox upload creates a ToolFile mapping.""" - - path: str - file: SandboxUploadedFile - - model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") - - -def extract_runtime_layer_specs(composition: RunComposition) -> list[RuntimeLayerSpec]: - """Project a run composition into the persistable non-sensitive layer list.""" - specs: list[RuntimeLayerSpec] = [] - for layer in composition.layers: - if layer.type in _SENSITIVE_LAYER_TYPES: - continue - config_value: JsonValue = None - if isinstance(layer.config, BaseModel): - config_value = layer.config.model_dump(mode="json", warnings=False) - else: - config_value = cast(JsonValue, layer.config) - specs.append( - RuntimeLayerSpec( - name=layer.name, - type=layer.type, - deps=dict(layer.deps), - metadata=dict(layer.metadata), - config=config_value, - ) - ) - return specs - - -def build_sandbox_locator_from_run_request(request: CreateRunRequest) -> SandboxLocator: - """Build a safe sandbox locator from a full create-run request. - - Raises: - ValueError: if the request has no resumable session snapshot or lacks the - execution-context/shell layers needed for sandbox access. - """ - if request.session_snapshot is None: - raise ValueError("Sandbox locator requires a non-empty session_snapshot.") - return build_sandbox_locator_from_layer_specs( - layer_specs=extract_runtime_layer_specs(request.composition), - session_snapshot=request.session_snapshot, - ) - - -def build_sandbox_locator_from_layer_specs( - *, - layer_specs: list[RuntimeLayerSpec], - session_snapshot: CompositorSessionSnapshot, -) -> SandboxLocator: - """Build a sandbox locator from persisted runtime specs plus a saved snapshot.""" - if not layer_specs: - raise ValueError("Sandbox locator requires persisted runtime layer specs.") - - for spec in layer_specs: - if spec.type in _SENSITIVE_LAYER_TYPES: - raise ValueError(f"Sandbox locator runtime specs must not include sensitive layer type {spec.type!r}.") - - execution_context_index = next( - (index for index, spec in enumerate(layer_specs) if spec.name == "execution_context"), - None, - ) - shell_index = next((index for index, spec in enumerate(layer_specs) if spec.name == "shell"), None) - if execution_context_index is None: - raise ValueError("Sandbox locator requires an 'execution_context' runtime layer spec.") - if shell_index is None: - raise ValueError("Sandbox locator requires a 'shell' runtime layer spec.") - if execution_context_index > shell_index: - raise ValueError("Sandbox locator requires 'execution_context' to appear before 'shell'.") - - execution_context_spec = layer_specs[execution_context_index] - shell_spec = layer_specs[shell_index] - if shell_spec.deps.get("execution_context") != execution_context_spec.name: - raise ValueError("Sandbox shell layer must depend on the execution_context layer.") - - kept_specs = [execution_context_spec, shell_spec] - kept_names = [spec.name for spec in kept_specs] - snapshot_layers = [layer for layer in session_snapshot.layers if layer.name in set(kept_names)] - if [layer.name for layer in snapshot_layers] != kept_names: - raise ValueError("Sandbox locator session_snapshot must contain execution_context and shell layers in order.") - - return SandboxLocator( - composition=RunComposition( - schema_version=1, - layers=[ - RunLayerSpec( - name=spec.name, - type=spec.type, - deps=dict(spec.deps), - metadata=dict(spec.metadata), - config=spec.config, - ) - for spec in kept_specs - ], - ), - session_snapshot=CompositorSessionSnapshot( - schema_version=session_snapshot.schema_version, - layers=[ - LayerSessionSnapshot( - name=layer.name, - lifecycle_state=layer.lifecycle_state, - runtime_state=dict(layer.runtime_state), - ) - for layer in snapshot_layers - ], - ), - ) - - -__all__ = [ - "RuntimeLayerSpec", - "SandboxFileEntry", - "SandboxListRequest", - "SandboxListResponse", - "SandboxLocator", - "SandboxReadRequest", - "SandboxReadResponse", - "SandboxUploadRequest", - "SandboxUploadResponse", - "SandboxUploadedFile", - "build_sandbox_locator_from_layer_specs", - "build_sandbox_locator_from_run_request", - "extract_runtime_layer_specs", -] diff --git a/dify-agent/src/dify_agent/protocol/schemas.py b/dify-agent/src/dify_agent/protocol/schemas.py index b4ea3843918..ca23f8bfd59 100644 --- a/dify-agent/src/dify_agent/protocol/schemas.py +++ b/dify-agent/src/dify_agent/protocol/schemas.py @@ -23,13 +23,9 @@ by ``DIFY_AGENT_OUTPUT_LAYER_ID``. Request-level ``on_exit`` signals decide whether each active layer is suspended or deleted when the run exits, with suspend as the default so successful terminal events can include resumable snapshots. Successful runs always publish the resumable Agenton session snapshot -on the terminal ``run_succeeded`` event together with exactly one of the final -JSON-safe ``output`` or a deferred external ``deferred_tool_call`` payload. A -lifecycle-only run may also succeed with ``output = null`` and ``usage = null`` -when the composition intentionally omits the reserved model layer and only -replays layer enter/exit work from a supplied snapshot. That lets consumers -treat terminal success events as complete run summaries without a separate pause -protocol. Session snapshots carry only layer lifecycle/runtime state in +on the terminal ``run_succeeded`` event together with either the final JSON-safe +``output`` or a deferred external ``deferred_tool_call`` payload. Session +snapshots carry only layer lifecycle/runtime state in compositor order; they do not persist output-layer config. Resumed structured-output runs therefore must resubmit the same ``output`` layer in ``composition.layers[]`` so snapshot layer name/order still matches the @@ -40,6 +36,7 @@ from __future__ import annotations from datetime import datetime, timezone from decimal import Decimal +from enum import StrEnum from typing import Annotated, ClassVar, Final, Literal, TypeAlias from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, model_serializer, model_validator @@ -63,6 +60,16 @@ RunEventType = Literal[ ] +class RunFailureType(StrEnum): + """Stable machine-readable categories for failed Dify Agent runs. + + Run-limit failures cover execution budgets enforced by Dify Agent, not + provider, connection, or wall-clock timeouts. + """ + + AGENT_RUN_LIMIT_EXCEEDED = "agent_run_limit_exceeded" + + def utc_now() -> datetime: """Return the timezone-aware timestamp format used by public schemas.""" return datetime.now(timezone.utc) @@ -217,6 +224,7 @@ class RunStatusResponse(BaseModel): created_at: datetime updated_at: datetime error: str | None = None + error_type: RunFailureType | None = None model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") @@ -319,6 +327,7 @@ class RunFailedEventData(BaseModel): """Terminal failure payload shown to polling and SSE consumers.""" error: str + error_type: RunFailureType | None = None reason: str | None = None model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") @@ -366,7 +375,7 @@ class RunSucceededEvent(BaseRunEvent): class RunFailedEvent(BaseRunEvent): - """Terminal failure event emitted before the run status becomes failed.""" + """Terminal failure event atomically committed with the failed run status.""" type: Literal["run_failed"] = "run_failed" data: RunFailedEventData @@ -420,6 +429,7 @@ __all__ = [ "RunEventsResponse", "RunFailedEvent", "RunFailedEventData", + "RunFailureType", "RunStartedEvent", "RunStatus", "RunStatusResponse", diff --git a/dify-agent/src/dify_agent/runtime/command_runner.py b/dify-agent/src/dify_agent/runtime/command_runner.py new file mode 100644 index 00000000000..1ceec51a0cf --- /dev/null +++ b/dify-agent/src/dify_agent/runtime/command_runner.py @@ -0,0 +1,97 @@ +"""Bounded execution of one transient command through a RuntimeLease.""" + +from __future__ import annotations + +import logging +import time +from typing import Literal + +from dify_agent.adapters.shell.protocols import ( + CompleteShellCommandResult, + ShellCommandProtocol, + ShellCommandResult, +) +from dify_agent.layers.shell.output_text import utf8_prefix + +logger = logging.getLogger(__name__) + +_TERMINATE_GRACE_SECONDS = 10.0 + + +async def execute_complete_with_commands( + commands: ShellCommandProtocol, + script: str, + *, + cwd: str | None, + env: dict[str, str] | None, + timeout: float, + max_output_bytes: int, +) -> CompleteShellCommandResult: + """Run a command to completion with bounded output and deterministic cleanup.""" + + deadline = time.monotonic() + timeout + job_id: str | None = None + result: ShellCommandResult | None = None + output_parts: list[str] = [] + captured_bytes = 0 + incomplete_reason: Literal["output_limit", "timeout"] | None = None + try: + result = await commands.run(script, cwd=cwd, env=env, timeout=_remaining_time(deadline)) + job_id = result.job_id + while True: + remaining_bytes = max(max_output_bytes - captured_bytes, 0) + limited_output = utf8_prefix(result.output, remaining_bytes) + output_parts.append(limited_output) + captured_bytes += len(limited_output.encode("utf-8")) + if limited_output != result.output: + incomplete_reason = "output_limit" + break + if captured_bytes >= max_output_bytes and (result.truncated or not result.done): + incomplete_reason = "output_limit" + break + if result.truncated: + result = await commands.read_output(result.job_id, offset=result.offset) + continue + if result.done: + break + remaining_time = _remaining_time(deadline) + if remaining_time <= 0.0: + incomplete_reason = "timeout" + break + result = await commands.wait(result.job_id, offset=result.offset, timeout=remaining_time) + + final_status = result.status + final_done = result.done + final_exit_code = result.exit_code + final_offset = result.offset + final_output_path = result.output_path + if incomplete_reason is not None and not result.done: + terminal_status = await commands.interrupt(result.job_id, grace_seconds=_TERMINATE_GRACE_SECONDS) + final_status = terminal_status.status + final_done = terminal_status.done + final_exit_code = terminal_status.exit_code + final_offset = terminal_status.offset + return CompleteShellCommandResult( + job_id=result.job_id, + status=final_status, + done=final_done, + exit_code=final_exit_code, + output="".join(output_parts), + output_complete=incomplete_reason is None, + incomplete_reason=incomplete_reason, + offset=final_offset, + output_path=final_output_path, + ) + finally: + if job_id is not None: + try: + await commands.delete(job_id, force=True) + except RuntimeError as exc: + logger.warning("Failed to delete transient shell job %s: %s", job_id, exc) + + +def _remaining_time(deadline: float) -> float: + return max(0.0, deadline - time.monotonic()) + + +__all__ = ["execute_complete_with_commands"] diff --git a/dify-agent/src/dify_agent/runtime/compositor_factory.py b/dify-agent/src/dify_agent/runtime/compositor_factory.py index b81cf737478..429c7fce57a 100644 --- a/dify-agent/src/dify_agent/runtime/compositor_factory.py +++ b/dify-agent/src/dify_agent/runtime/compositor_factory.py @@ -3,12 +3,13 @@ Only explicitly allowed provider type ids are constructible here. The default provider set contains prompt layers, the optional pydantic-ai history layer, the state-free Dify structured output layer, the optional Dify ask-human layer, the -Dify execution-context layer, the stateful Dify shell layer, and the Dify -plugin/knowledge business-layer family: +Dify execution-context and runtime-resource layers, the stateful Dify shell +layer, and the Dify plugin/knowledge business-layer family: - ``dify.config`` for Agent Soul-backed config assets + eager pull, - ``dify.execution_context`` for shared tenant/user/run daemon context, -- ``dify.shell`` for shellctl-backed shell job control, +- ``dify.runtime`` for operation-scoped RuntimeLease acquisition, +- ``dify.shell`` for command/file capabilities from the active RuntimeLease, - ``dify.plugin.llm`` for plugin-backed model selection, - ``dify.plugin.tools`` for prepared plugin tool exposure, and - ``dify.core.tools`` for API-routed Dify tool exposure, and @@ -16,10 +17,9 @@ plugin/knowledge business-layer family: Public DTOs provide Dify context plus plugin/model/tool data, while server-only plugin daemon settings and Dify API inner settings are injected through provider -factories. An already-selected shell provider (shellctl or enterprise, built at -the runtime boundary) and the Agent Stub URL/token factory are injected for -``DifyShellLayer``; a ``None`` shell provider leaves the shell layer disabled. -The resulting ``Compositor`` +factories. The deployment-selected runtime backend profile is injected only +into the Runtime layer. Shell receives Agent Stub settings and consumes the +active lease's data plane. The resulting ``Compositor`` remains Agenton state-only at the snapshot boundary: live resources such as HTTP clients are injected by runtime-owned providers, may be held on active layer instances inside ``resource_context()``, and never enter session @@ -29,10 +29,7 @@ snapshots. from __future__ import annotations from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Any, cast - -if TYPE_CHECKING: - from dify_agent.adapters.shell.protocols import ShellProviderProtocol +from typing import Any, cast from pydantic_ai.messages import UserContent @@ -55,8 +52,11 @@ from dify_agent.layers.execution_context.layer import DifyExecutionContextLayer from dify_agent.layers.knowledge.configs import DifyKnowledgeBaseLayerConfig from dify_agent.layers.knowledge.layer import DifyKnowledgeBaseLayer from dify_agent.layers.output.output_layer import DifyOutputLayer +from dify_agent.layers.runtime.configs import DifyRuntimeLayerConfig +from dify_agent.layers.runtime.layer import DifyRuntimeLayer from dify_agent.layers.shell.configs import DifyShellLayerConfig from dify_agent.layers.shell.layer import DifyShellLayer +from dify_agent.runtime_backend import RuntimeBackendProfile type DifyAgentLayerProvider = LayerProvider[Any] @@ -67,14 +67,13 @@ def create_default_layer_providers( plugin_daemon_api_key: str = "", inner_api_url: str = "http://localhost:5001", inner_api_key: str = "", - shell_provider: ShellProviderProtocol | None = None, - shell_home_root: str = "/home", + runtime_backend_profile: RuntimeBackendProfile | None = None, shell_redact_patterns: list[str] | None = None, agent_stub_api_base_url: str | None = None, agent_stub_token_factory: ShellAgentStubTokenFactory | None = None, ) -> tuple[DifyAgentLayerProvider, ...]: """Return the server provider set of safe config-constructible layers.""" - return ( + providers: list[DifyAgentLayerProvider] = [ LayerProvider.from_layer_type(PromptLayer), LayerProvider.from_layer_type(PydanticAIHistoryLayer), LayerProvider.from_layer_type(DifyOutputLayer), @@ -93,8 +92,6 @@ def create_default_layer_providers( layer_type=DifyShellLayer, create=lambda config: DifyShellLayer.from_config_with_settings( DifyShellLayerConfig.model_validate(config), - shell_provider=shell_provider, - shell_home_root=shell_home_root, shell_redact_patterns=shell_redact_patterns or [], agent_stub_api_base_url=agent_stub_api_base_url, agent_stub_token_factory=agent_stub_token_factory, @@ -125,7 +122,20 @@ def create_default_layer_providers( inner_api_key=inner_api_key, ), ), - ) + ] + if runtime_backend_profile is not None: + providers.extend( + [ + LayerProvider.from_factory( + layer_type=DifyRuntimeLayer, + create=lambda config: DifyRuntimeLayer.from_config_with_backend( + DifyRuntimeLayerConfig.model_validate(config), + backend=runtime_backend_profile.execution_bindings, + ), + ), + ] + ) + return tuple(providers) def build_pydantic_ai_compositor( diff --git a/dify-agent/src/dify_agent/runtime/event_sink.py b/dify-agent/src/dify_agent/runtime/event_sink.py index 67babdc3754..121d58cecde 100644 --- a/dify-agent/src/dify_agent/runtime/event_sink.py +++ b/dify-agent/src/dify_agent/runtime/event_sink.py @@ -1,16 +1,15 @@ """Event sink contracts used by the runner and storage adapters. -The runner only needs append-only event writes and status transitions, so tests -can use ``InMemoryRunEventSink`` without Redis. Production storage implements the -same protocol with Redis streams in ``dify_agent.storage.redis_run_store``. The -terminal success helper writes either the final JSON-safe output or one deferred -tool request together with the resumable session snapshot in a single event so -consumers can stop at ``run_succeeded`` without correlating separate payload -events. +Non-terminal events remain append-only. Terminal events use ``finalize_run`` so +the event and matching run status are committed as one compare-and-set +transition. Tests can use ``InMemoryRunEventSink`` without Redis; production +storage implements the same contract with Redis streams in +``dify_agent.storage.redis_run_store``. """ from collections import defaultdict -from typing import Protocol, cast +from dataclasses import dataclass +from typing import Protocol, TypeAlias, cast from pydantic import JsonValue from pydantic_ai.messages import AgentStreamEvent @@ -26,6 +25,7 @@ from dify_agent.protocol.schemas import ( RunCancelledEventData, RunFailedEvent, RunFailedEventData, + RunFailureType, RunStartedEvent, RunStatus, RunSucceededEvent, @@ -35,17 +35,28 @@ from dify_agent.protocol.schemas import ( _UNSET = object() +TerminalRunEvent: TypeAlias = RunSucceededEvent | RunFailedEvent | RunCancelledEvent +NonTerminalRunEvent: TypeAlias = RunStartedEvent | PydanticAIStreamRunEvent + + +@dataclass(frozen=True, slots=True) +class RunFinalizationResult: + """Outcome of attempting the only terminal transition for one run.""" + + applied: bool + status: RunStatus + event_id: str | None = None class RunEventSink(Protocol): """Boundary used by runtime code to publish observable run progress.""" - async def append_event(self, event: RunEvent) -> str: - """Persist ``event`` and return its cursor id.""" + async def append_event(self, event: NonTerminalRunEvent) -> str: + """Persist a non-terminal event and return its cursor id.""" ... - async def update_status(self, run_id: str, status: RunStatus, error: str | None = None) -> None: - """Persist the current run status.""" + async def finalize_run(self, event: TerminalRunEvent) -> RunFinalizationResult: + """Atomically persist the first terminal event and matching status.""" ... @@ -55,31 +66,55 @@ class InMemoryRunEventSink: events: dict[str, list[RunEvent]] statuses: dict[str, RunStatus] errors: dict[str, str | None] + error_types: dict[str, RunFailureType | None] def __init__(self) -> None: self.events = defaultdict(list) self.statuses = {} self.errors = {} + self.error_types = {} - async def append_event(self, event: RunEvent) -> str: - """Store an event and assign a monotonic per-run cursor.""" + async def append_event(self, event: NonTerminalRunEvent) -> str: + """Store a non-terminal event and assign a monotonic per-run cursor.""" event_id = str(len(self.events[event.run_id]) + 1) stored = event.model_copy(update={"id": event_id}) self.events[event.run_id].append(stored) return event_id - async def update_status(self, run_id: str, status: RunStatus, error: str | None = None) -> None: - """Record the latest status; timestamps are owned by run stores.""" - self.statuses[run_id] = status - self.errors[run_id] = error + async def finalize_run(self, event: TerminalRunEvent) -> RunFinalizationResult: + """Store only the first terminal event and its derived status.""" + current_status = self.statuses.get(event.run_id, "running") + if current_status != "running": + return RunFinalizationResult(applied=False, status=current_status) + + status, error, error_type = terminal_event_status_fields(event) + event_id = str(len(self.events[event.run_id]) + 1) + self.events[event.run_id].append(event.model_copy(update={"id": event_id})) + self.statuses[event.run_id] = status + self.errors[event.run_id] = error + self.error_types[event.run_id] = error_type + return RunFinalizationResult(applied=True, status=status, event_id=event_id) + + +def terminal_event_status_fields( + event: TerminalRunEvent, +) -> tuple[RunStatus, str | None, RunFailureType | None]: + """Derive the persisted terminal status fields from one typed event.""" + match event: + case RunSucceededEvent(): + return "succeeded", None, None + case RunFailedEvent(): + return "failed", event.data.error, event.data.error_type + case RunCancelledEvent(): + return "cancelled", event.data.message or event.data.reason, None async def emit_run_event( sink: RunEventSink, *, - event: RunEvent, + event: NonTerminalRunEvent, ) -> str: - """Append an already typed public run event.""" + """Append an already typed non-terminal public run event.""" return await sink.append_event(event) @@ -118,8 +153,8 @@ async def emit_run_succeeded( deferred_tool_call: DeferredToolCallPayload | object = _UNSET, session_snapshot: CompositorSessionSnapshot, usage: AgentRunUsage | None = None, -) -> str: - """Emit the terminal success event with output or deferred continuation. +) -> RunFinalizationResult: + """Finalize a run as succeeded with output or deferred continuation. Callers must activate exactly one result branch. ``_UNSET`` is used instead of ``None`` to preserve the distinction between an omitted inactive branch @@ -137,9 +172,8 @@ async def emit_run_succeeded( if usage is not None: data["usage"] = usage - return await emit_run_event( - sink, - event=RunSucceededEvent( + return await sink.finalize_run( + RunSucceededEvent( run_id=run_id, data=RunSucceededEventData.model_validate(data), created_at=utc_now(), @@ -152,12 +186,16 @@ async def emit_run_failed( *, run_id: str, error: str, + error_type: RunFailureType | None = None, reason: str | None = None, -) -> str: - """Emit the terminal failure lifecycle event.""" - return await emit_run_event( - sink, - event=RunFailedEvent(run_id=run_id, data=RunFailedEventData(error=error, reason=reason), created_at=utc_now()), +) -> RunFinalizationResult: + """Finalize a run with a failed terminal event.""" + return await sink.finalize_run( + RunFailedEvent( + run_id=run_id, + data=RunFailedEventData(error=error, error_type=error_type, reason=reason), + created_at=utc_now(), + ), ) @@ -167,11 +205,10 @@ async def emit_run_cancelled( run_id: str, reason: str | None = None, message: str | None = None, -) -> str: - """Emit the terminal cancellation lifecycle event.""" - return await emit_run_event( - sink, - event=RunCancelledEvent( +) -> RunFinalizationResult: + """Finalize a run with a cancelled terminal event.""" + return await sink.finalize_run( + RunCancelledEvent( run_id=run_id, data=RunCancelledEventData(reason=reason, message=message), created_at=utc_now(), @@ -181,11 +218,15 @@ async def emit_run_cancelled( __all__ = [ "InMemoryRunEventSink", + "NonTerminalRunEvent", "RunEventSink", + "RunFinalizationResult", + "TerminalRunEvent", "emit_pydantic_ai_event", "emit_run_cancelled", "emit_run_event", "emit_run_failed", "emit_run_started", "emit_run_succeeded", + "terminal_event_status_fields", ] diff --git a/dify-agent/src/dify_agent/runtime/run_scheduler.py b/dify-agent/src/dify_agent/runtime/run_scheduler.py index c02799a66ec..8925edee292 100644 --- a/dify-agent/src/dify_agent/runtime/run_scheduler.py +++ b/dify-agent/src/dify_agent/runtime/run_scheduler.py @@ -1,10 +1,11 @@ """In-process scheduling for Dify Agent runs. The scheduler is intentionally process-local: it persists a run record, starts an -``asyncio.Task`` for ``AgentRunRunner.run()``, and keeps only a transient active -task registry. Redis remains the durable source for status and event streams, but -there is no Redis job queue or cross-process handoff. If the process crashes, -currently active runs are lost until an external operator marks or retries them. +``asyncio.Task`` supervisor for the local runner and cancellation observer, and +keeps only a transient active task registry. Redis remains the durable source for +status and event streams, but there is no Redis job queue or cross-process +handoff. If the process crashes, currently active runs are lost until an external +operator marks or retries them. Create-run requests are accepted once the scheduler is not stopping and storage can persist the run record. Request-shaped execution failures are left to ``AgentRunRunner`` so bad compositions, ``on_exit`` policies, prompts, @@ -34,7 +35,7 @@ class SchedulerStoppingError(RuntimeError): class RunCancellationConflictError(RuntimeError): - """Raised when a run exists but can no longer be cancelled by this scheduler.""" + """Raised when a run exists but a different terminal state already won.""" class RunStore(RunEventSink, Protocol): @@ -44,8 +45,8 @@ class RunStore(RunEventSink, Protocol): """Persist a new run record and return it with status ``running``.""" ... - async def get_run(self, run_id: str) -> RunRecord: - """Return the latest persisted run record.""" + async def wait_for_cancellation(self, run_id: str) -> bool: + """Wait for a terminal state and report whether cancellation won.""" ... @@ -63,19 +64,20 @@ type RunRunnerFactory = Callable[[RunRecord, CreateRunRequest], RunnableRun] class RunScheduler: """Owns process-local run tasks and best-effort graceful shutdown. - ``active_tasks`` is mutated only on the event loop that calls ``create_run`` - and ``shutdown``. The task registry is not durable; it exists so the lifespan - hook can wait for in-flight work and mark cancelled runs failed before Redis is - closed. A lock guards the stopping flag, run persistence, and task - registration so shutdown cannot begin after a request is admitted. + ``active_tasks`` contains local supervisor tasks and is mutated only on the + event loop that calls ``create_run`` and ``shutdown``. The task registry is not + durable; it exists so the lifespan hook can wait for in-flight work and mark + shutdown-cancelled runs failed before Redis is closed. It is not consulted + when accepting cancellation requests. A lock guards the stopping flag, run + persistence, and task registration so shutdown cannot begin after a request + is admitted. """ store: RunStore shutdown_grace_seconds: float active_tasks: dict[str, asyncio.Task[None]] - cancelled_run_ids: set[str] stopping: bool - runner_factory: RunRunnerFactory + runner_factory: RunRunnerFactory | None layer_providers: tuple[LayerProviderInput, ...] plugin_daemon_http_client: httpx.AsyncClient dify_api_http_client: httpx.AsyncClient @@ -94,12 +96,11 @@ class RunScheduler: self.store = store self.shutdown_grace_seconds = shutdown_grace_seconds self.active_tasks = {} - self.cancelled_run_ids = set() self.stopping = False self.plugin_daemon_http_client = plugin_daemon_http_client self.dify_api_http_client = dify_api_http_client self.layer_providers = layer_providers if layer_providers is not None else create_default_layer_providers() - self.runner_factory = runner_factory or self._default_runner_factory + self.runner_factory = runner_factory self._lifecycle_lock = asyncio.Lock() async def create_run(self, request: CreateRunRequest) -> RunRecord: @@ -120,37 +121,15 @@ class RunScheduler: return record async def cancel_run(self, run_id: str, request: CancelRunRequest) -> CancelRunResponse: - """Cancel one active task and persist an idempotent cancelled terminal state.""" - async with self._lifecycle_lock: - record = await self.store.get_run(run_id) - if record.status == "cancelled": - return CancelRunResponse(run_id=run_id, status="cancelled") - if record.status != "running": - raise RunCancellationConflictError(f"run already finished with status {record.status!r}") - - task = self.active_tasks.get(run_id) - if task is None: - raise RunCancellationConflictError("run is not active in this scheduler process") - self.cancelled_run_ids.add(run_id) - _ = task.cancel(request.message or request.reason) - _ = await emit_run_cancelled( - self.store, - run_id=run_id, - reason=request.reason, - message=request.message, - ) - await self.store.update_status(run_id, "cancelled", request.message or request.reason) - - # Some model/tool stacks can consume one CancelledError. Re-inject it - # after the terminal state is durable without making the HTTP request - # wait for arbitrary third-party cleanup. - for _attempt in range(2): - if task.done(): - break - _ = task.cancel(request.message or request.reason) - await asyncio.sleep(0) - if task.done(): - self._discard_active_run(run_id) + """Persist an idempotent cancellation without relying on local task ownership.""" + finalization = await emit_run_cancelled( + self.store, + run_id=run_id, + reason=request.reason, + message=request.message, + ) + if finalization.status != "cancelled": + raise RunCancellationConflictError(f"run already finished with status {finalization.status!r}") return CancelRunResponse(run_id=run_id, status="cancelled") async def shutdown(self) -> None: @@ -165,11 +144,7 @@ class RunScheduler: if not pending: return - pending_run_ids = [ - run_id - for run_id, task in tasks_by_run_id.items() - if task in pending and run_id not in self.cancelled_run_ids - ] + pending_run_ids = [run_id for run_id, task in tasks_by_run_id.items() if task in pending] for task in pending: _ = task.cancel() _ = await asyncio.gather(*pending, return_exceptions=True) @@ -177,15 +152,69 @@ class RunScheduler: await self._mark_cancelled_run_failed(run_id) async def _run_record(self, record: RunRecord, request: CreateRunRequest) -> None: - """Execute a stored run and log failures already reflected in events.""" + """Supervise one local runner and its durable cancellation observer.""" + cancel_requested = asyncio.Event() + runner = self._create_runner(record, request, is_cancelled=cancel_requested.is_set) + runner_task = asyncio.create_task(runner.run(), name=f"dify-agent-runner-{record.run_id}") + observer_task = asyncio.create_task( + self.store.wait_for_cancellation(record.run_id), + name=f"dify-agent-cancellation-observer-{record.run_id}", + ) try: - await self.runner_factory(record, request).run() + _ = await asyncio.wait((runner_task, observer_task), return_when=asyncio.FIRST_COMPLETED) + if observer_task.done(): + try: + cancellation_won = observer_task.result() + except Exception as exc: + cancel_requested.set() + await self._cancel_and_wait(runner_task, reinject=True) + _ = await emit_run_failed( + self.store, + run_id=record.run_id, + error=f"run cancellation observer failed: {exc}", + reason="cancellation_observer", + ) + raise + + if cancellation_won: + cancel_requested.set() + await self._cancel_and_wait(runner_task, reinject=True) + else: + await runner_task + else: + await runner_task except asyncio.CancelledError: + cancel_requested.set() + await self._cancel_and_wait(observer_task) + await self._cancel_and_wait(runner_task, reinject=True) raise except Exception: logger.exception("scheduled run failed", extra={"run_id": record.run_id}) + finally: + await self._cancel_and_wait(observer_task) + if not runner_task.done(): + cancel_requested.set() + await self._cancel_and_wait(runner_task, reinject=True) - def _default_runner_factory(self, record: RunRecord, request: CreateRunRequest) -> RunnableRun: + def _create_runner( + self, + record: RunRecord, + request: CreateRunRequest, + *, + is_cancelled: Callable[[], bool], + ) -> RunnableRun: + """Create a runner while keeping injected test runners source-compatible.""" + if self.runner_factory is not None: + return self.runner_factory(record, request) + return self._default_runner_factory(record, request, is_cancelled=is_cancelled) + + def _default_runner_factory( + self, + record: RunRecord, + request: CreateRunRequest, + *, + is_cancelled: Callable[[], bool], + ) -> RunnableRun: """Create the production runner for a stored run record.""" return AgentRunRunner( sink=self.store, @@ -194,19 +223,30 @@ class RunScheduler: plugin_daemon_http_client=self.plugin_daemon_http_client, dify_api_http_client=self.dify_api_http_client, layer_providers=self.layer_providers, - is_cancelled=lambda: record.run_id in self.cancelled_run_ids, + is_cancelled=is_cancelled, ) def _discard_active_run(self, run_id: str) -> None: _ = self.active_tasks.pop(run_id, None) - self.cancelled_run_ids.discard(run_id) + + @staticmethod + async def _cancel_and_wait(task: asyncio.Task[object], *, reinject: bool = False) -> None: + """Cancel and reap a child task, with bounded reinjection for runners.""" + if not task.done(): + _ = task.cancel() + if reinject: + for _attempt in range(2): + await asyncio.sleep(0) + if task.done(): + break + _ = task.cancel() + _ = await asyncio.gather(task, return_exceptions=True) async def _mark_cancelled_run_failed(self, run_id: str) -> None: """Best-effort failure event/status for shutdown-cancelled runs.""" message = "run cancelled during server shutdown" try: _ = await emit_run_failed(self.store, run_id=run_id, error=message, reason="shutdown") - await self.store.update_status(run_id, "failed", message) except Exception: logger.exception("failed to mark cancelled run failed", extra={"run_id": run_id}) diff --git a/dify-agent/src/dify_agent/runtime/runner.py b/dify-agent/src/dify_agent/runtime/runner.py index 8fec8d5fbf8..814447871e5 100644 --- a/dify-agent/src/dify_agent/runtime/runner.py +++ b/dify-agent/src/dify_agent/runtime/runner.py @@ -1,18 +1,14 @@ """Runtime execution for one scheduled Dify Agent run. The runner is storage-agnostic: it normalizes the public Dify composition into -Agenton's graph/config split and chooses one of two execution modes after the -composition is normalized and the ``on_exit`` policy is validated: +Agenton's graph/config split and executes one model run after the ``on_exit`` +policy is validated: - model runs: enter a fresh ``CompositorRun`` (or resume one from a snapshot), render the current Dify system prompts into temporary ``message_history``, run pydantic-ai with either the current ``run.user_prompts`` or deferred external tool results, emit raw stream events with agent-message delta annotations, apply request-level ``on_exit`` signals, and publish a terminal success or failure event; -- lifecycle-only runs: enter from a supplied snapshot, apply request-level - ``on_exit`` signals, exit without invoking a model, and succeed with explicit - ``output = null`` and ``usage = null``. - The Pydantic AI model is resolved from the active Agenton layer named by ``DIFY_AGENT_MODEL_LAYER_ID``. An optional history layer contributes stored message history only through session state; successful model runs append only @@ -40,10 +36,11 @@ from typing import Any, Literal, Protocol, cast, runtime_checkable import httpx from graphon.model_runtime.entities.llm_entities import LLMUsage from pydantic import JsonValue, TypeAdapter -from pydantic_ai.exceptions import ModelHTTPError +from pydantic_ai.exceptions import ModelHTTPError, UsageLimitExceeded from pydantic_ai.messages import AgentStreamEvent, PartDeltaEvent, PartStartEvent, TextPart, TextPartDelta from pydantic_ai.output import OutputSpec from pydantic_ai.tools import DeferredToolRequests, DeferredToolResults +from pydantic_ai.usage import UsageLimits from agenton.compositor import CompositorSessionSnapshot, LayerConfigInput, LayerProviderInput from agenton.layers.types import PydanticAITool @@ -58,12 +55,13 @@ from dify_agent.protocol.schemas import ( CreateRunRequest, DIFY_AGENT_MODEL_LAYER_ID, DeferredToolCallPayload, + RunFailureType, normalize_composition, ) from dify_agent.runtime.agent_factory import create_agent, normalize_user_input from dify_agent.runtime.agenton_validation import is_agenton_enter_validation_runtime_error from dify_agent.runtime.compositor_factory import build_pydantic_ai_compositor, create_default_layer_providers -from dify_agent.adapters.shell.protocols import SandboxExpiredError +from dify_agent.runtime_backend import BindingLostError from dify_agent.runtime.event_sink import ( RunEventSink, emit_pydantic_ai_event, @@ -83,6 +81,7 @@ from dify_agent.runtime.user_prompt_validation import EMPTY_USER_PROMPTS_ERROR, _AGENT_OUTPUT_ADAPTER = TypeAdapter(object) +_MAX_AGENT_STEPS_PER_RUN = 100 @runtime_checkable @@ -115,13 +114,16 @@ class AgentRunValidationError(ValueError): """Raised when a run request is valid JSON but cannot execute.""" -def _run_failed_error_payload(exc: Exception) -> tuple[str, str | None]: - """Return the public failed-run error text and structured reason.""" +def _run_failed_error_payload(exc: Exception) -> tuple[str, RunFailureType | None, str | None]: + """Return the public failed-run error text, type, and structured reason.""" message = str(exc) or type(exc).__name__ reason: str | None = None - if isinstance(exc, SandboxExpiredError): - return message, "sandbox_expired" + if isinstance(exc, UsageLimitExceeded): + return message, RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, None + + if isinstance(exc, BindingLostError): + return message, None, "binding_lost" if isinstance(exc, ModelHTTPError): body = exc.body @@ -140,7 +142,7 @@ def _run_failed_error_payload(exc: Exception) -> tuple[str, str | None]: if isinstance(exc, DifyKnowledgeBaseClientError): reason = exc.error_code or "DifyKnowledgeBaseClientError" - return message, reason + return message, None, reason def _has_model_layer(request: CreateRunRequest) -> bool: @@ -201,9 +203,6 @@ class AgentRunRunner: async def run(self) -> None: """Execute the run and emit the documented event sequence.""" - if self.is_cancelled(): - return - await self.sink.update_status(self.run_id, "running") if self.is_cancelled(): return _ = await emit_run_started(self.sink, run_id=self.run_id) @@ -213,10 +212,17 @@ class AgentRunRunner: except Exception as exc: if self.is_cancelled(): return - message, reason = _run_failed_error_payload(exc) - _ = await emit_run_failed(self.sink, run_id=self.run_id, error=message, reason=reason) - await self.sink.update_status(self.run_id, "failed", message) - raise + message, error_type, reason = _run_failed_error_payload(exc) + finalization = await emit_run_failed( + self.sink, + run_id=self.run_id, + error=message, + error_type=error_type, + reason=reason, + ) + if finalization.applied: + raise + return if self.is_cancelled(): return @@ -231,10 +237,9 @@ class AgentRunRunner: session_snapshot=outcome.session_snapshot, usage=outcome.usage, ) - await self.sink.update_status(self.run_id, "succeeded") async def _run_agent(self) -> RunSuccessOutcome: - """Run the normalized request in model or lifecycle-only mode. + """Run the normalized request through the model path. Known request-shaped Agenton enter-time failures are normalized to ``AgentRunValidationError``. That includes the existing small class of @@ -262,49 +267,9 @@ class AgentRunRunner: raise AgentRunValidationError(str(exc)) from exc if not _has_model_layer(self.request): - return await self._run_lifecycle_only(compositor=compositor, layer_configs=layer_configs) + raise AgentRunValidationError(f"Missing required '{DIFY_AGENT_MODEL_LAYER_ID}' layer.") return await self._run_model(compositor=compositor, layer_configs=layer_configs) - async def _run_lifecycle_only( - self, - *, - compositor: Any, - layer_configs: dict[str, LayerConfigInput], - ) -> RunSuccessOutcome: - """Replay only layer lifecycle work for a no-LLM composition plus snapshot.""" - if self.request.session_snapshot is None: - raise AgentRunValidationError( - f"Missing '{DIFY_AGENT_MODEL_LAYER_ID}' requires a session_snapshot for lifecycle-only runs." - ) - if self.request.deferred_tool_results is not None: - raise AgentRunValidationError( - f"Deferred tool results require the reserved '{DIFY_AGENT_MODEL_LAYER_ID}' layer." - ) - - entered_run = False - try: - async with compositor.enter(configs=layer_configs, session_snapshot=self.request.session_snapshot) as run: - entered_run = True - apply_layer_exit_signals(run, self.request.on_exit) - except RuntimeError as exc: - if not entered_run and is_agenton_enter_validation_runtime_error(exc): - raise AgentRunValidationError(str(exc)) from exc - raise - except ValueError as exc: - if not entered_run: - raise AgentRunValidationError(str(exc)) from exc - raise - - if run.session_snapshot is None: - raise RuntimeError("Agenton run did not produce a session snapshot after exit.") - return RunSuccessOutcome( - result_kind="output", - output=None, - deferred_tool_call=None, - session_snapshot=run.session_snapshot, - usage=None, - ) - async def _run_model( self, *, @@ -371,6 +336,7 @@ class AgentRunRunner: message_history=message_history, deferred_tool_results=deferred_tool_results, event_stream_handler=handle_events, + usage_limits=UsageLimits(request_limit=_MAX_AGENT_STEPS_PER_RUN), ) complete_usage = model.accumulated_usage if isinstance(model, _HasAccumulatedUsage) else None usage = _serialize_agent_usage(complete_usage if complete_usage is not None else _result_usage(result)) diff --git a/dify-agent/src/dify_agent/runtime_backend/__init__.py b/dify-agent/src/dify_agent/runtime_backend/__init__.py new file mode 100644 index 00000000000..bca6748166f --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/__init__.py @@ -0,0 +1,47 @@ +"""Public contracts for deployment-selected runtime backends.""" + +from .errors import ( + BindingAcquireError, + BindingCreateError, + BindingDestroyError, + BindingLostError, + HomeSnapshotCreateError, + HomeSnapshotNotFoundError, + RuntimeBackendError, + SharedWorkspaceUnsupportedError, + WorkspacePreservationUnsupportedError, + WorkspaceUnavailableError, +) +from .protocols import ( + ExecutionBindingAllocation, + ExecutionBindingBackend, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotBackend, + HomeSnapshotCreateSpec, + RuntimeBackendProfile, + RuntimeLayout, + RuntimeLease, +) + +__all__ = [ + "BindingAcquireError", + "BindingCreateError", + "BindingDestroyError", + "BindingLostError", + "ExecutionBindingAllocation", + "ExecutionBindingBackend", + "ExecutionBindingCreateSpec", + "ExecutionBindingDestroySpec", + "HomeSnapshotBackend", + "HomeSnapshotCreateError", + "HomeSnapshotCreateSpec", + "HomeSnapshotNotFoundError", + "RuntimeBackendError", + "RuntimeBackendProfile", + "RuntimeLayout", + "RuntimeLease", + "SharedWorkspaceUnsupportedError", + "WorkspacePreservationUnsupportedError", + "WorkspaceUnavailableError", +] diff --git a/dify-agent/src/dify_agent/runtime_backend/e2b.py b/dify-agent/src/dify_agent/runtime_backend/e2b.py new file mode 100644 index 00000000000..930ea317dcf --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/e2b.py @@ -0,0 +1,400 @@ +"""E2B backend adapters with shellctl as the command and file data plane. + +Dify API persists only opaque Home Snapshot, Binding, and Workspace backend +refs. This adapter maps those refs to E2B resources internally. API keys, +traffic tokens, SDK objects, shellctl clients, and ``RuntimeLease`` objects stay +operation-local and are never serialized into Agenton state. +""" + +from __future__ import annotations + +import asyncio +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Literal, Protocol, cast + +import httpx2 as httpx +from shellctl.client import ShellctlClientError + +from dify_agent.adapters.shell.protocols import ShellCommandProtocol +from dify_agent.adapters.shell.shellctl import ShellctlClientProtocol +from dify_agent.runtime_backend.errors import ( + BindingAcquireError, + BindingCreateError, + BindingDestroyError, + BindingLostError, + HomeSnapshotCreateError, + SharedWorkspaceUnsupportedError, + WorkspacePreservationUnsupportedError, +) +from dify_agent.runtime_backend.protocols import ( + ExecutionBindingAllocation, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotCreateSpec, + RuntimeLayout, + RuntimeLease, +) +from dify_agent.runtime_backend.shellctl import ShellctlRuntimeLease, create_owned_shellctl_lease + +if TYPE_CHECKING: + from e2b.connection_config import ApiParams + +# One RuntimeLease spans the complete Agent run, not one Shell tool call. +E2B_MAX_ACTIVE_TIMEOUT_SECONDS = 60 * 60 +_SHELLCTL_READY_MAX_ATTEMPTS = 3 +_SHELLCTL_READY_RETRY_INTERVAL_SECONDS = 0.5 + + +class _E2BControlPlaneNotFoundError(RuntimeError): + """Typed boundary error for SDK resources that no longer exist.""" + + +class _E2BFileSystem(Protocol): + async def make_dir(self, path: str) -> bool: ... + + async def exists(self, path: str) -> bool: ... + + async def remove(self, path: str) -> None: ... + + +class _E2BSnapshotInfo(Protocol): + snapshot_id: str + names: list[str] + + +class _E2BSandbox(Protocol): + sandbox_id: str + traffic_access_token: str | None + files: _E2BFileSystem + + def get_host(self, port: int) -> str: ... + + async def pause(self, keep_memory: bool = True) -> bool: ... + + async def kill(self) -> bool: ... + + async def create_snapshot(self, name: str | None = None) -> _E2BSnapshotInfo: ... + + +class E2BControlPlane(Protocol): + async def create( + self, + template: str, + *, + timeout: int, + metadata: dict[str, str], + on_timeout: Literal["kill", "pause"], + ) -> _E2BSandbox: ... + + async def connect(self, handle: str, *, timeout: int) -> _E2BSandbox: ... + + async def kill(self, handle: str) -> bool: ... + + async def delete_snapshot(self, snapshot_ref: str) -> bool: ... + + +@dataclass(frozen=True, slots=True) +class E2BSDKControlPlane: + """Stateless async E2B SDK boundary configured with one deployment API key. + + SDK Sandbox objects are returned only to operation-local backend adapters. + Native not-found exceptions are normalized so adapters can distinguish + confirmed resource loss from transient acquisition and cleanup failures. + """ + + api_key: str + + def _options(self) -> ApiParams: + return {"api_key": self.api_key} + + async def create( + self, + template: str, + *, + timeout: int, + metadata: dict[str, str], + on_timeout: Literal["kill", "pause"], + ) -> _E2BSandbox: + from e2b import AsyncSandbox, NotFoundException, SandboxNotFoundException + + try: + return cast( + _E2BSandbox, + cast( + object, + await AsyncSandbox.create( + template, + timeout=timeout, + metadata=metadata, + lifecycle={"on_timeout": on_timeout, "auto_resume": False}, + **self._options(), + ), + ), + ) + except (SandboxNotFoundException, NotFoundException) as exc: + raise _E2BControlPlaneNotFoundError(str(exc)) from exc + + async def connect(self, handle: str, *, timeout: int) -> _E2BSandbox: + from e2b import AsyncSandbox, NotFoundException, SandboxNotFoundException + + try: + return cast( + _E2BSandbox, + cast(object, await AsyncSandbox.connect(handle, timeout=timeout, **self._options())), + ) + except (SandboxNotFoundException, NotFoundException) as exc: + raise _E2BControlPlaneNotFoundError(str(exc)) from exc + + async def kill(self, handle: str) -> bool: + from e2b import AsyncSandbox, NotFoundException, SandboxNotFoundException + + try: + return await AsyncSandbox.kill(handle, **self._options()) + except (SandboxNotFoundException, NotFoundException) as exc: + raise _E2BControlPlaneNotFoundError(str(exc)) from exc + + async def delete_snapshot(self, snapshot_ref: str) -> bool: + from e2b import AsyncSandbox, NotFoundException, SandboxNotFoundException + + try: + return await AsyncSandbox.delete_snapshot(snapshot_ref, **self._options()) + except (SandboxNotFoundException, NotFoundException) as exc: + raise _E2BControlPlaneNotFoundError(str(exc)) from exc + + +@dataclass(slots=True) +class E2BHomeSnapshotBackend: + """Implement immutable Home Snapshot operations with E2B snapshots. + + Build Apply snapshots the E2B resource behind the supplied ``RuntimeLease``. + Dify API stores the returned value as an opaque backend ref; this adapter + keeps no cross-request state. + """ + + control_plane: E2BControlPlane + + async def create_from_runtime(self, *, spec: HomeSnapshotCreateSpec, source: RuntimeLease) -> str: + """Create an immutable E2B snapshot from the source Binding's active lease.""" + del spec + if not isinstance(source, E2BRuntimeLease): + raise HomeSnapshotCreateError("E2B Home Snapshot requires an E2B RuntimeLease") + try: + snapshot = await source.sandbox.create_snapshot() + return snapshot.snapshot_id + except BaseException as exc: + if isinstance(exc, Exception): + raise HomeSnapshotCreateError(str(exc)) from exc + raise + + async def delete(self, snapshot_ref: str) -> None: + try: + _ = await self.control_plane.delete_snapshot(snapshot_ref) + except _E2BControlPlaneNotFoundError: + return + except Exception as exc: + raise BindingDestroyError(str(exc)) from exc + + +@dataclass(slots=True) +class E2BExecutionBindingBackend: + """Implement Execution Binding operations with E2B and shellctl. + + In this backend one physical E2B resource represents both a Binding and its + Workspace, so their opaque refs have the same value. Materialized Home and + Workspace remain distinct logical resources even though E2B couples their + physical lifecycle. Active timeout pauses the resource and is not a + resource-age TTL. + """ + + control_plane: E2BControlPlane + template: str + active_timeout_seconds: int + shellctl_auth_token: str = "" + shellctl_port: int = 5004 + layout: RuntimeLayout = field( + default_factory=lambda: RuntimeLayout(home_dir="/home/dify", workspace_dir="/home/dify/workspace") + ) + + async def create_binding(self, spec: ExecutionBindingCreateSpec) -> ExecutionBindingAllocation: + """Create one paused E2B resource from a snapshot or deployment template.""" + if spec.existing_workspace_ref is not None: + raise SharedWorkspaceUnsupportedError("current E2B backend cannot attach to an existing Workspace") + sandbox: _E2BSandbox | None = None + try: + sandbox = await self.control_plane.create( + self.template if spec.home_snapshot_ref is None else spec.home_snapshot_ref, + timeout=self.active_timeout_seconds, + metadata={ + "dify.resource": "runtime-sandbox", + "dify.binding_id": spec.binding_id, + "dify.workspace_id": spec.workspace_id, + "dify.tenant_id": spec.tenant_id, + "dify.agent_id": spec.agent_id, + }, + on_timeout="pause", + ) + if await sandbox.files.exists(self.layout.workspace_dir): + await sandbox.files.remove(self.layout.workspace_dir) + _ = await sandbox.files.make_dir(self.layout.workspace_dir) + sandbox_id = sandbox.sandbox_id + _ = await sandbox.pause(keep_memory=True) + return ExecutionBindingAllocation(binding_ref=sandbox_id, workspace_ref=sandbox_id) + except BaseException as exc: + if sandbox is not None: + try: + _ = await sandbox.kill() + except BaseException: + pass + if isinstance(exc, Exception): + if isinstance(exc, BindingCreateError): + raise + raise BindingCreateError(str(exc)) from exc + raise + + async def acquire(self, binding_ref: str) -> RuntimeLease: + """Acquire operation-scoped shellctl access for an opaque Binding ref.""" + sandbox: _E2BSandbox | None = None + lease: E2BRuntimeLease | None = None + try: + sandbox = await self.control_plane.connect(binding_ref, timeout=self.active_timeout_seconds) + if not await sandbox.files.exists(self.layout.workspace_dir): + raise BindingLostError(f"E2B Binding {binding_ref!r} no longer contains its Workspace") + lease = await self._lease(sandbox) + await _wait_for_shellctl_ready(lease.data_plane.client) + return lease + except _E2BControlPlaneNotFoundError as exc: + raise BindingLostError(f"E2B Binding {binding_ref!r} no longer exists") from exc + except BindingLostError: + await _best_effort_pause(sandbox) + raise + except BaseException as exc: + await _best_effort_close_data_plane(lease) + await _best_effort_pause(sandbox) + if isinstance(exc, Exception): + raise BindingAcquireError(str(exc)) from exc + raise + + async def release(self, lease: RuntimeLease) -> None: + """Close operation-local transports and pause the physical E2B resource.""" + if not isinstance(lease, E2BRuntimeLease): + raise TypeError("E2BExecutionBindingBackend can only release its own RuntimeLease") + close_error: Exception | None = None + try: + await lease.data_plane.close() + except Exception as exc: + close_error = exc + try: + _ = await lease.sandbox.pause(keep_memory=True) + except Exception as exc: + raise BindingAcquireError(str(exc)) from exc + if close_error is not None: + raise BindingAcquireError(str(close_error)) from close_error + + async def destroy_binding(self, spec: ExecutionBindingDestroySpec) -> None: + """Destroy the coupled physical Binding and Workspace idempotently.""" + if not spec.destroy_workspace: + raise WorkspacePreservationUnsupportedError( + "current E2B backend cannot destroy a Binding while preserving its Workspace" + ) + if spec.workspace_ref != spec.binding_ref: + raise BindingDestroyError("E2B Workspace ref must equal its Binding ref") + try: + _ = await self.control_plane.kill(spec.binding_ref) + except _E2BControlPlaneNotFoundError: + return + except Exception as exc: + raise BindingDestroyError(str(exc)) from exc + + async def _lease(self, sandbox: _E2BSandbox) -> "E2BRuntimeLease": + entrypoint = f"https://{sandbox.get_host(self.shellctl_port)}" + traffic_token = sandbox.traffic_access_token + headers = {"X-Access-Token": traffic_token} if isinstance(traffic_token, str) and traffic_token else {} + http_client = httpx.AsyncClient( + base_url=entrypoint, + headers=headers, + follow_redirects=True, + timeout=httpx.Timeout(60.0), + ) + + def client_factory() -> ShellctlClientProtocol: + from shellctl.client import ShellctlClient + + return cast( + ShellctlClientProtocol, + cast( + object, + ShellctlClient(entrypoint, token=self.shellctl_auth_token, client=http_client), + ), + ) + + data_plane = await create_owned_shellctl_lease( + handle=sandbox.sandbox_id, + layout=self.layout, + entrypoint=entrypoint, + token=self.shellctl_auth_token, + client_factory=client_factory, + owned_transport=http_client, + ) + return E2BRuntimeLease(sandbox=sandbox, data_plane=data_plane) + + +@dataclass(slots=True) +class E2BRuntimeLease: + """Invocation-local E2B SDK object plus the owned shellctl data-plane lease.""" + + sandbox: _E2BSandbox + data_plane: ShellctlRuntimeLease + + @property + def handle(self) -> str: + return self.data_plane.handle + + @property + def layout(self) -> RuntimeLayout: + return self.data_plane.layout + + @property + def commands(self) -> ShellCommandProtocol: + return self.data_plane.commands + + +async def _wait_for_shellctl_ready(client: ShellctlClientProtocol) -> None: + for attempt in range(_SHELLCTL_READY_MAX_ATTEMPTS): + try: + _ = await client.health() + return + except (httpx.TimeoutException, httpx.RequestError): + if attempt == _SHELLCTL_READY_MAX_ATTEMPTS - 1: + raise + except ShellctlClientError as exc: + if not 500 <= exc.status_code < 600 or attempt == _SHELLCTL_READY_MAX_ATTEMPTS - 1: + raise + await asyncio.sleep(_SHELLCTL_READY_RETRY_INTERVAL_SECONDS) + + +async def _best_effort_close_data_plane(lease: E2BRuntimeLease | None) -> None: + if lease is None: + return + try: + await lease.data_plane.close() + except BaseException: + pass + + +async def _best_effort_pause(sandbox: _E2BSandbox | None) -> None: + if sandbox is None: + return + try: + _ = await sandbox.pause(keep_memory=True) + except BaseException: + pass + + +__all__ = [ + "E2B_MAX_ACTIVE_TIMEOUT_SECONDS", + "E2BControlPlane", + "E2BExecutionBindingBackend", + "E2BHomeSnapshotBackend", + "E2BSDKControlPlane", + "E2BRuntimeLease", +] diff --git a/dify-agent/src/dify_agent/runtime_backend/enterprise.py b/dify-agent/src/dify_agent/runtime_backend/enterprise.py new file mode 100644 index 00000000000..f86867d778e --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/enterprise.py @@ -0,0 +1,283 @@ +"""Enterprise Gateway adapter for the working-environment protocol. + +The existing Gateway can allocate, reconnect to, and delete a sandbox, but it +does not expose immutable Home Snapshot operations. One physical sandbox owns +both the materialized Home and Workspace, so their cleanup is coupled. Runtime +access remains operation-local and is routed through the Gateway's shellctl +proxy. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +import logging +import shlex +from typing import cast +from urllib.parse import quote + +import httpx2 as httpx + +from dify_agent.adapters.shell.protocols import ShellCommandProtocol, ShellProviderError +from dify_agent.adapters.shell.shellctl import ShellctlClientProtocol, ShellctlCommands +from dify_agent.runtime_backend.errors import ( + BindingAcquireError, + BindingCreateError, + BindingDestroyError, + BindingLostError, + SharedWorkspaceUnsupportedError, + WorkspacePreservationUnsupportedError, +) +from dify_agent.runtime_backend.protocols import ( + ExecutionBindingAllocation, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotCreateSpec, + RuntimeLayout, + RuntimeLease, +) +from dify_agent.runtime_backend.shellctl import ( + ShellctlRuntimeLease, + create_owned_shellctl_lease, + run_shellctl_control_command, +) + +logger = logging.getLogger(__name__) + + +def _not_implemented() -> NotImplementedError: + return NotImplementedError("Enterprise Gateway does not implement immutable Home Snapshot operations") + + +@dataclass(slots=True) +class EnterpriseHomeSnapshotBackend: + """Reject Home Snapshot operations until the Gateway exposes immutable snapshots.""" + + async def create_from_runtime(self, *, spec: HomeSnapshotCreateSpec, source: RuntimeLease) -> str: + del spec, source + raise _not_implemented() + + async def delete(self, snapshot_ref: str) -> None: + del snapshot_ref + raise _not_implemented() + + +@dataclass(slots=True) +class EnterpriseExecutionBindingBackend: + """Manage Gateway sandboxes as coupled physical Bindings and Workspaces.""" + + gateway_endpoint: str + auth_token: str + gateway_timeout: float = 30.0 + proxy_timeout: float = 60.0 + layout: RuntimeLayout = field( + default_factory=lambda: RuntimeLayout(home_dir="/home/dify", workspace_dir="/home/dify/workspace") + ) + + async def create_binding(self, spec: ExecutionBindingCreateSpec) -> ExecutionBindingAllocation: + """Create a default Gateway sandbox and initialize its canonical layout.""" + if spec.existing_workspace_ref is not None: + raise SharedWorkspaceUnsupportedError("current Enterprise backend cannot attach to an existing Workspace") + if spec.home_snapshot_ref is not None: + raise BindingCreateError("current Enterprise backend cannot materialize an immutable Home Snapshot") + + sandbox_id: str | None = None + data_plane: ShellctlRuntimeLease | None = None + headers = {"X-Inner-Api-Key": self.auth_token} if self.auth_token else {} + try: + async with httpx.AsyncClient( + base_url=self.gateway_endpoint.rstrip("/"), + headers=headers, + timeout=httpx.Timeout(self.gateway_timeout), + ) as client: + response = await client.post("/v1/sandboxes", json={"tenantId": spec.tenant_id}) + _ = response.raise_for_status() + payload = response.json() + sandbox_id_value = payload.get("sandboxId") if isinstance(payload, dict) else None + if not isinstance(sandbox_id_value, str) or not sandbox_id_value: + raise BindingCreateError("Enterprise Gateway returned an invalid sandbox id") + sandbox_id = sandbox_id_value + + data_plane = await self._create_data_plane(sandbox_id) + result = await run_shellctl_control_command( + ShellctlCommands(client=data_plane.client), + "\n".join( + [ + "set -eu", + f"mkdir -p {shlex.quote(self.layout.home_dir)}", + f"rm -rf -- {shlex.quote(self.layout.workspace_dir)}", + f"mkdir -p {shlex.quote(self.layout.workspace_dir)}", + f"chmod 700 {shlex.quote(self.layout.home_dir)} {shlex.quote(self.layout.workspace_dir)}", + ] + ), + ) + if result.exit_code != 0: + raise BindingCreateError(result.output) + await data_plane.close() + data_plane = None + return ExecutionBindingAllocation(binding_ref=sandbox_id, workspace_ref=sandbox_id) + except BaseException as exc: + await _close_best_effort(data_plane, binding_ref=sandbox_id or spec.binding_id) + if sandbox_id is not None: + await self._delete_sandbox_best_effort(sandbox_id) + if isinstance(exc, BindingCreateError): + raise + if isinstance(exc, Exception): + raise BindingCreateError(str(exc)) from exc + raise + + async def acquire(self, binding_ref: str) -> RuntimeLease: + """Reconnect to one existing Gateway sandbox without creating a replacement.""" + data_plane: ShellctlRuntimeLease | None = None + try: + data_plane = await self._create_data_plane(binding_ref) + validation_commands = ShellctlCommands(client=data_plane.client) + result = await run_shellctl_control_command( + validation_commands, + "\n".join( + [ + "set -eu", + f"test -d {shlex.quote(self.layout.home_dir)}", + f"test -d {shlex.quote(self.layout.workspace_dir)}", + ] + ), + timeout=5.0, + ) + if result.exit_code != 0: + raise BindingLostError(f"Enterprise Binding {binding_ref!r} no longer contains its Home or Workspace") + return EnterpriseRuntimeLease(data_plane=data_plane) + except ShellProviderError as exc: + await _close_best_effort(data_plane, binding_ref=binding_ref) + if _is_missing_sandbox(exc): + raise BindingLostError(f"Enterprise Binding {binding_ref!r} no longer exists") from exc + raise BindingAcquireError(str(exc)) from exc + except BindingLostError: + await _close_best_effort(data_plane, binding_ref=binding_ref) + raise + except BaseException as exc: + await _close_best_effort(data_plane, binding_ref=binding_ref) + if isinstance(exc, Exception): + raise BindingAcquireError(str(exc)) from exc + raise + + async def release(self, lease: RuntimeLease) -> None: + """Close operation-local shellctl resources without deleting the sandbox.""" + if not isinstance(lease, EnterpriseRuntimeLease): + raise TypeError("EnterpriseExecutionBindingBackend can only release its own RuntimeLease") + try: + await lease.data_plane.close() + except Exception as exc: + raise BindingAcquireError(str(exc)) from exc + + async def destroy_binding(self, spec: ExecutionBindingDestroySpec) -> None: + """Delete the coupled sandbox only when its Workspace is also retired.""" + if not spec.destroy_workspace: + raise WorkspacePreservationUnsupportedError( + "current Enterprise backend cannot destroy a Binding while preserving its Workspace" + ) + if spec.workspace_ref != spec.binding_ref: + raise BindingDestroyError("Enterprise Workspace ref must equal its Binding ref") + + try: + await self._delete_sandbox(spec.binding_ref) + except (httpx.TimeoutException, httpx.RequestError, httpx.HTTPStatusError) as exc: + raise BindingDestroyError(str(exc)) from exc + + async def _delete_sandbox(self, sandbox_id: str) -> None: + headers = {"X-Inner-Api-Key": self.auth_token} if self.auth_token else {} + encoded_sandbox_id = quote(sandbox_id, safe="") + async with httpx.AsyncClient( + base_url=self.gateway_endpoint.rstrip("/"), + headers=headers, + timeout=httpx.Timeout(self.gateway_timeout), + ) as client: + response = await client.delete(f"/v1/sandboxes/{encoded_sandbox_id}") + if response.status_code == 404: + return + _ = response.raise_for_status() + + async def _delete_sandbox_best_effort(self, sandbox_id: str) -> None: + try: + await self._delete_sandbox(sandbox_id) + except BaseException: + logger.warning( + "failed to delete Enterprise sandbox after Binding creation failed", + exc_info=True, + extra={"binding_ref": sandbox_id}, + ) + + async def _create_data_plane(self, binding_ref: str) -> ShellctlRuntimeLease: + proxy_base_url = f"{self.gateway_endpoint.rstrip('/')}/proxy/" + headers = {"X-Sandbox-Id": binding_ref} + if self.auth_token: + headers["X-Inner-Api-Key"] = self.auth_token + http_client = httpx.AsyncClient( + base_url=proxy_base_url, + headers=headers, + follow_redirects=True, + timeout=httpx.Timeout(self.proxy_timeout), + transport=httpx.AsyncHTTPTransport(retries=3), + ) + + def client_factory() -> ShellctlClientProtocol: + from shellctl.client import ShellctlClient + + return cast( + ShellctlClientProtocol, + cast(object, ShellctlClient(proxy_base_url, token=self.auth_token, client=http_client)), + ) + + return await create_owned_shellctl_lease( + handle=binding_ref, + layout=self.layout, + entrypoint=proxy_base_url, + token=self.auth_token, + client_factory=client_factory, + owned_transport=http_client, + ) + + +@dataclass(slots=True) +class EnterpriseRuntimeLease: + """Invocation-local Enterprise shellctl connection and canonical layout.""" + + data_plane: ShellctlRuntimeLease + + @property + def handle(self) -> str: + return self.data_plane.handle + + @property + def layout(self) -> RuntimeLayout: + return self.data_plane.layout + + @property + def commands(self) -> ShellCommandProtocol: + return self.data_plane.commands + + +def _is_missing_sandbox(exc: ShellProviderError) -> bool: + return exc.status_code == 404 or (exc.code or "").casefold() in { + "not_found", + "sandbox_expired", + "sandbox_not_found", + } + + +async def _close_best_effort(data_plane: ShellctlRuntimeLease | None, *, binding_ref: str) -> None: + if data_plane is None: + return + try: + await data_plane.close() + except BaseException: + logger.warning( + "failed to close Enterprise RuntimeLease after acquisition failed", + exc_info=True, + extra={"binding_ref": binding_ref}, + ) + + +__all__ = [ + "EnterpriseExecutionBindingBackend", + "EnterpriseHomeSnapshotBackend", + "EnterpriseRuntimeLease", +] diff --git a/dify-agent/src/dify_agent/runtime_backend/errors.py b/dify-agent/src/dify_agent/runtime_backend/errors.py new file mode 100644 index 00000000000..260844b976f --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/errors.py @@ -0,0 +1,55 @@ +"""Stable failures exposed above provider-specific runtime backends.""" + + +class RuntimeBackendError(RuntimeError): + """Base infrastructure failure for the selected runtime backend.""" + + +class HomeSnapshotCreateError(RuntimeBackendError): + pass + + +class HomeSnapshotNotFoundError(RuntimeBackendError): + pass + + +class BindingCreateError(RuntimeBackendError): + pass + + +class BindingAcquireError(RuntimeBackendError): + pass + + +class BindingLostError(RuntimeBackendError): + pass + + +class BindingDestroyError(RuntimeBackendError): + pass + + +class SharedWorkspaceUnsupportedError(RuntimeBackendError): + pass + + +class WorkspacePreservationUnsupportedError(RuntimeBackendError): + pass + + +class WorkspaceUnavailableError(RuntimeBackendError): + pass + + +__all__ = [ + "BindingAcquireError", + "BindingCreateError", + "BindingDestroyError", + "BindingLostError", + "HomeSnapshotCreateError", + "HomeSnapshotNotFoundError", + "RuntimeBackendError", + "SharedWorkspaceUnsupportedError", + "WorkspacePreservationUnsupportedError", + "WorkspaceUnavailableError", +] diff --git a/dify-agent/src/dify_agent/runtime_backend/leases.py b/dify-agent/src/dify_agent/runtime_backend/leases.py new file mode 100644 index 00000000000..82f9efed92e --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/leases.py @@ -0,0 +1,35 @@ +"""Operation-scoped RuntimeLease acquisition and release.""" + +from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager +import logging + +from dify_agent.runtime_backend.protocols import ExecutionBindingBackend, RuntimeLease + +logger = logging.getLogger(__name__) + + +@asynccontextmanager +async def open_runtime_lease( + backend: ExecutionBindingBackend, + binding_ref: str, +) -> AsyncGenerator[RuntimeLease, None]: + """Acquire one Binding and deterministically release its lease.""" + + lease = await backend.acquire(binding_ref) + primary_error: BaseException | None = None + try: + yield lease + except BaseException as exc: + primary_error = exc + raise + finally: + try: + await backend.release(lease) + except BaseException: + if primary_error is None: + raise + logger.warning("failed to release RuntimeLease after operation failed", exc_info=True) + + +__all__ = ["open_runtime_lease"] diff --git a/dify-agent/src/dify_agent/runtime_backend/local.py b/dify-agent/src/dify_agent/runtime_backend/local.py new file mode 100644 index 00000000000..f0a2e5c2e66 --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/local.py @@ -0,0 +1,308 @@ +"""Local working-environment backend over one shared shellctl daemon. + +Materialized Homes and Workspaces live under separate roots. A Binding ref +encodes both logical path segments so this stateless adapter can reacquire the +same pair without maintaining a resource catalog. +""" + +from __future__ import annotations + +import logging +import posixpath +import re +import shlex +from dataclasses import dataclass + +from dify_agent.adapters.shell.protocols import ShellCommandProtocol +from dify_agent.adapters.shell.shellctl import ShellctlClientFactory +from dify_agent.runtime_backend.errors import ( + BindingAcquireError, + BindingCreateError, + BindingDestroyError, + BindingLostError, + HomeSnapshotCreateError, +) +from dify_agent.runtime_backend.protocols import ( + ExecutionBindingAllocation, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotCreateSpec, + RuntimeLayout, + RuntimeLease, +) +from dify_agent.runtime_backend.shellctl import ( + ShellctlRuntimeLease, + create_shellctl_lease, + run_shellctl_control_command, +) + +_SAFE_REF_PART = re.compile(r"^[A-Za-z0-9._-]+$") +_BINDING_REF_SEPARATOR = ":" +logger = logging.getLogger(__name__) + + +@dataclass(slots=True) +class LocalHomeSnapshotBackend: + endpoint: str + auth_token: str + snapshot_root: str = "/home/dify/.dify-agent-home-snapshots" + client_factory: ShellctlClientFactory | None = None + + async def create_from_runtime(self, *, spec: HomeSnapshotCreateSpec, source: RuntimeLease) -> str: + snapshot_ref = _local_snapshot_ref(spec.home_snapshot_id) + target = self._snapshot_dir(snapshot_ref) + lease = self._control_lease(snapshot_ref, paths=(source.layout.home_dir, target)) + script = "\n".join( + [ + "set -eu", + f"test -d {shlex.quote(source.layout.home_dir)}", + f"mkdir -p {shlex.quote(target)}", + f"cp -a {shlex.quote(source.layout.home_dir)}/. {shlex.quote(target)}/", + f"chmod 700 {shlex.quote(target)}", + ] + ) + try: + result = await run_shellctl_control_command(lease.commands, script) + if result.exit_code != 0: + raise HomeSnapshotCreateError(result.output) + return snapshot_ref + except BaseException as exc: + await _remove_partial(lease.commands, target=target, resource_ref=snapshot_ref) + if isinstance(exc, HomeSnapshotCreateError): + raise + if isinstance(exc, Exception): + raise HomeSnapshotCreateError(str(exc)) from exc + raise + finally: + await _close_best_effort(lease, resource_ref=snapshot_ref) + + async def delete(self, snapshot_ref: str) -> None: + normalized = _validated_ref_part(snapshot_ref) + lease = self._control_lease(normalized) + try: + result = await run_shellctl_control_command( + lease.commands, + f"rm -rf -- {shlex.quote(self._snapshot_dir(normalized))}", + ) + if result.exit_code != 0: + raise BindingDestroyError(result.output) + except BaseException: + await _close_best_effort(lease, resource_ref=normalized) + raise + else: + await lease.close() + + def _snapshot_dir(self, snapshot_ref: str) -> str: + return f"{self.snapshot_root.rstrip('/')}/{snapshot_ref}" + + def _control_lease(self, handle: str, *, paths: tuple[str, ...] = ()) -> ShellctlRuntimeLease: + control_root = _control_root(paths or (self.snapshot_root,)) + layout = RuntimeLayout(home_dir=control_root, workspace_dir=control_root) + return create_shellctl_lease( + handle=handle, + layout=layout, + entrypoint=self.endpoint, + token=self.auth_token, + client_factory=self.client_factory, + ) + + +@dataclass(slots=True) +class LocalExecutionBindingBackend: + endpoint: str + auth_token: str + materialized_home_root: str = "/home/dify/.dify-agent-materialized-homes" + workspace_root: str = "/home/dify/.dify-agent-workspaces" + snapshot_root: str = "/home/dify/.dify-agent-home-snapshots" + client_factory: ShellctlClientFactory | None = None + + async def create_binding(self, spec: ExecutionBindingCreateSpec) -> ExecutionBindingAllocation: + binding_id = _validated_ref_part(spec.binding_id) + workspace_id = _validated_ref_part(spec.workspace_id) + workspace_ref = workspace_id + if spec.existing_workspace_ref is not None: + existing_workspace_ref = _validated_ref_part(spec.existing_workspace_ref) + if existing_workspace_ref != workspace_ref: + raise BindingCreateError("existing Workspace ref does not match workspace_id") + snapshot_dir: str | None = None + if spec.home_snapshot_ref is not None: + snapshot_ref = _validated_ref_part(spec.home_snapshot_ref) + snapshot_dir = f"{self.snapshot_root.rstrip('/')}/{snapshot_ref}" + binding_ref = _local_binding_ref(binding_id=binding_id, workspace_id=workspace_id) + home_dir = self._home_dir(binding_id) + workspace_dir = self._workspace_dir(workspace_id) + creates_workspace = spec.existing_workspace_ref is None + workspace_setup = ( + f"mkdir -p {shlex.quote(workspace_dir)}" if creates_workspace else f"test -d {shlex.quote(workspace_dir)}" + ) + setup = ["set -eu"] + if snapshot_dir is not None: + setup.append(f"test -d {shlex.quote(snapshot_dir)}") + setup.extend([workspace_setup, f"mkdir -p {shlex.quote(home_dir)}"]) + if snapshot_dir is not None: + setup.append(f"cp -a {shlex.quote(snapshot_dir)}/. {shlex.quote(home_dir)}/") + setup.append(f"chmod 700 {shlex.quote(home_dir)} {shlex.quote(workspace_dir)}") + script = "\n".join(setup) + lease = self._control_lease(binding_ref) + try: + result = await run_shellctl_control_command(lease.commands, script) + if result.exit_code != 0: + raise BindingCreateError(result.output) + return ExecutionBindingAllocation(binding_ref=binding_ref, workspace_ref=workspace_ref) + except BaseException as exc: + targets = [home_dir] + if creates_workspace: + targets.append(workspace_dir) + await _remove_partial( + lease.commands, + target=" ".join(shlex.quote(target) for target in targets), + resource_ref=binding_ref, + target_is_shell_words=True, + ) + if isinstance(exc, BindingCreateError): + raise + if isinstance(exc, Exception): + raise BindingCreateError(str(exc)) from exc + raise + finally: + await _close_best_effort(lease, resource_ref=binding_ref) + + async def acquire(self, binding_ref: str) -> RuntimeLease: + lease = self._lease(binding_ref) + try: + result = await run_shellctl_control_command( + lease.commands, + "\n".join( + [ + "set -eu", + f"test -d {shlex.quote(lease.layout.home_dir)}", + f"test -d {shlex.quote(lease.layout.workspace_dir)}", + ] + ), + timeout=10.0, + ) + if result.exit_code != 0: + raise BindingLostError(f"Local Binding {binding_ref!r} no longer exists") + return lease + except BaseException as exc: + await _close_best_effort(lease, resource_ref=binding_ref) + if isinstance(exc, BindingLostError): + raise + if isinstance(exc, Exception): + raise BindingAcquireError(str(exc)) from exc + raise + + async def release(self, lease: RuntimeLease) -> None: + if not isinstance(lease, ShellctlRuntimeLease): + raise TypeError("LocalExecutionBindingBackend can only release its own RuntimeLease") + try: + await lease.close() + except Exception as exc: + raise BindingAcquireError(str(exc)) from exc + + async def destroy_binding(self, spec: ExecutionBindingDestroySpec) -> None: + binding_id, workspace_id = _parse_local_binding_ref(spec.binding_ref) + if spec.destroy_workspace: + workspace_ref = _validated_ref_part(spec.workspace_ref or "") + if workspace_ref != workspace_id: + raise BindingDestroyError("Workspace ref does not match Binding ref") + lease = self._control_lease(spec.binding_ref) + targets = [self._home_dir(binding_id)] + if spec.destroy_workspace: + targets.append(self._workspace_dir(workspace_id)) + try: + result = await run_shellctl_control_command( + lease.commands, + "rm -rf -- " + " ".join(shlex.quote(target) for target in targets), + ) + if result.exit_code != 0: + raise BindingDestroyError(result.output) + except BaseException: + await _close_best_effort(lease, resource_ref=spec.binding_ref) + raise + else: + await lease.close() + + def _lease(self, binding_ref: str) -> ShellctlRuntimeLease: + binding_id, workspace_id = _parse_local_binding_ref(binding_ref) + return create_shellctl_lease( + handle=binding_ref, + layout=RuntimeLayout( + home_dir=self._home_dir(binding_id), + workspace_dir=self._workspace_dir(workspace_id), + ), + entrypoint=self.endpoint, + token=self.auth_token, + client_factory=self.client_factory, + ) + + def _control_lease(self, handle: str) -> ShellctlRuntimeLease: + control_root = _control_root((self.materialized_home_root, self.workspace_root, self.snapshot_root)) + return create_shellctl_lease( + handle=handle, + layout=RuntimeLayout(home_dir=control_root, workspace_dir=control_root), + entrypoint=self.endpoint, + token=self.auth_token, + client_factory=self.client_factory, + ) + + def _home_dir(self, binding_id: str) -> str: + return f"{self.materialized_home_root.rstrip('/')}/{binding_id}" + + def _workspace_dir(self, workspace_id: str) -> str: + return f"{self.workspace_root.rstrip('/')}/{workspace_id}" + + +def _local_snapshot_ref(home_snapshot_id: str) -> str: + return f"home-{_validated_ref_part(home_snapshot_id)}" + + +def _local_binding_ref(*, binding_id: str, workspace_id: str) -> str: + return f"{_validated_ref_part(binding_id)}{_BINDING_REF_SEPARATOR}{_validated_ref_part(workspace_id)}" + + +def _parse_local_binding_ref(binding_ref: str) -> tuple[str, str]: + parts = binding_ref.split(_BINDING_REF_SEPARATOR) + if len(parts) != 2: + raise ValueError("Local Binding ref is invalid") + return _validated_ref_part(parts[0]), _validated_ref_part(parts[1]) + + +def _validated_ref_part(value: str) -> str: + if value in {"", ".", ".."} or _SAFE_REF_PART.fullmatch(value) is None: + raise ValueError("runtime backend ref must be a safe path segment") + return value + + +def _control_root(paths: tuple[str, ...]) -> str: + normalized = tuple(posixpath.normpath(path) for path in paths) + common = posixpath.commonpath(normalized) + if len(normalized) == 1 or common in normalized: + common = posixpath.dirname(common) + return common or "/" + + +async def _remove_partial( + commands: ShellCommandProtocol, + *, + target: str, + resource_ref: str, + target_is_shell_words: bool = False, +) -> None: + try: + target_words = target if target_is_shell_words else shlex.quote(target) + result = await run_shellctl_control_command(commands, f"rm -rf -- {target_words}") + if result.exit_code != 0: + logger.warning("failed to remove partial local resource", extra={"resource_ref": resource_ref}) + except BaseException: + logger.warning("failed to remove partial local resource", exc_info=True, extra={"resource_ref": resource_ref}) + + +async def _close_best_effort(lease: ShellctlRuntimeLease, *, resource_ref: str) -> None: + try: + await lease.close() + except BaseException: + logger.warning("failed to close local RuntimeLease", exc_info=True, extra={"resource_ref": resource_ref}) + + +__all__ = ["LocalExecutionBindingBackend", "LocalHomeSnapshotBackend"] diff --git a/dify-agent/src/dify_agent/runtime_backend/profile.py b/dify-agent/src/dify_agent/runtime_backend/profile.py new file mode 100644 index 00000000000..28bcf6bd01e --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/profile.py @@ -0,0 +1,173 @@ +"""Deployment-selected coherent runtime backend profile construction.""" + +from __future__ import annotations + +import posixpath +from typing import ClassVar, Literal, Self +from urllib.parse import urlparse + +from pydantic import AliasChoices, Field, model_validator +from pydantic_settings import BaseSettings, SettingsConfigDict + +from dify_agent.runtime_backend.e2b import ( + E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + E2BHomeSnapshotBackend, + E2BSDKControlPlane, + E2BExecutionBindingBackend, +) +from dify_agent.runtime_backend.enterprise import EnterpriseExecutionBindingBackend, EnterpriseHomeSnapshotBackend +from dify_agent.runtime_backend.local import LocalExecutionBindingBackend, LocalHomeSnapshotBackend +from dify_agent.runtime_backend.protocols import RuntimeBackendProfile + +DEFAULT_E2B_TEMPLATE = "difys-default-team/dify-agent-local-sandbox" +DEFAULT_LOCAL_MATERIALIZED_HOME_ROOT = "/home/dify/.dify-agent-materialized-homes" +DEFAULT_LOCAL_WORKSPACE_ROOT = "/home/dify/.dify-agent-workspaces" +DEFAULT_LOCAL_HOME_SNAPSHOT_ROOT = "/home/dify/.dify-agent-home-snapshots" + + +class RuntimeBackendSettings(BaseSettings): + """Server-private credentials and endpoints for one coherent backend profile.""" + + runtime_backend: Literal["local", "enterprise", "e2b"] = "local" + + local_sandbox_endpoint: str | None = Field( + default=None, + validation_alias=AliasChoices( + "DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT", + "DIFY_AGENT_SHELLCTL_ENTRYPOINT", + ), + ) + local_sandbox_auth_token: str | None = Field( + default=None, + validation_alias=AliasChoices( + "DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN", + "DIFY_AGENT_SHELLCTL_AUTH_TOKEN", + ), + ) + local_sandbox_materialized_home_root: str = DEFAULT_LOCAL_MATERIALIZED_HOME_ROOT + local_sandbox_workspace_root: str = DEFAULT_LOCAL_WORKSPACE_ROOT + local_sandbox_home_snapshot_root: str = DEFAULT_LOCAL_HOME_SNAPSHOT_ROOT + + enterprise_sandbox_gateway_endpoint: str | None = None + enterprise_sandbox_gateway_auth_token: str | None = None + enterprise_sandbox_gateway_timeout: float = Field(default=30.0, gt=0) + enterprise_sandbox_proxy_timeout: float = Field(default=60.0, gt=0) + + e2b_api_key: str | None = None + e2b_template: str = DEFAULT_E2B_TEMPLATE + e2b_active_timeout_seconds: int = Field( + default=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ge=1, + le=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ) + e2b_shellctl_auth_token: str = "" + e2b_shellctl_port: int = Field(default=5004, ge=1, le=65535) + + model_config: ClassVar[SettingsConfigDict] = SettingsConfigDict( + env_prefix="DIFY_AGENT_", + env_file=(".env", "dify-agent/.env"), + extra="ignore", + populate_by_name=True, + ) + + @model_validator(mode="after") + def validate_selected_backend(self) -> Self: + match self.runtime_backend: + case "local": + if not self.local_sandbox_endpoint or not self.local_sandbox_endpoint.strip(): + raise ValueError("local_sandbox_endpoint is required for the local runtime backend") + _validate_http_url(self.local_sandbox_endpoint, field_name="local_sandbox_endpoint") + _validate_absolute_posix_path( + self.local_sandbox_materialized_home_root, + field_name="local_sandbox_materialized_home_root", + ) + _validate_absolute_posix_path( + self.local_sandbox_workspace_root, + field_name="local_sandbox_workspace_root", + ) + _validate_absolute_posix_path( + self.local_sandbox_home_snapshot_root, + field_name="local_sandbox_home_snapshot_root", + ) + case "enterprise": + endpoint = self.enterprise_sandbox_gateway_endpoint + if not endpoint or not endpoint.strip(): + raise ValueError( + "enterprise_sandbox_gateway_endpoint is required for the enterprise runtime backend" + ) + _validate_http_url(endpoint, field_name="enterprise_sandbox_gateway_endpoint") + case "e2b": + if not self.e2b_api_key or not self.e2b_api_key.strip(): + raise ValueError("e2b_api_key is required for the e2b runtime backend") + if not self.e2b_template.strip(): + raise ValueError("e2b_template must not be blank") + return self + + +def create_runtime_backend_profile(settings: RuntimeBackendSettings) -> RuntimeBackendProfile: + """Construct one driver pair selected exclusively by server deployment settings.""" + match settings.runtime_backend: + case "local": + endpoint = settings.local_sandbox_endpoint or "" + token = settings.local_sandbox_auth_token or "" + return RuntimeBackendProfile( + home_snapshots=LocalHomeSnapshotBackend( + endpoint=endpoint, + auth_token=token, + snapshot_root=settings.local_sandbox_home_snapshot_root, + ), + execution_bindings=LocalExecutionBindingBackend( + endpoint=endpoint, + auth_token=token, + materialized_home_root=settings.local_sandbox_materialized_home_root, + workspace_root=settings.local_sandbox_workspace_root, + snapshot_root=settings.local_sandbox_home_snapshot_root, + ), + ) + case "enterprise": + endpoint = settings.enterprise_sandbox_gateway_endpoint or "" + token = settings.enterprise_sandbox_gateway_auth_token or "" + return RuntimeBackendProfile( + home_snapshots=EnterpriseHomeSnapshotBackend(), + execution_bindings=EnterpriseExecutionBindingBackend( + gateway_endpoint=endpoint, + auth_token=token, + gateway_timeout=settings.enterprise_sandbox_gateway_timeout, + proxy_timeout=settings.enterprise_sandbox_proxy_timeout, + ), + ) + case "e2b": + control_plane = E2BSDKControlPlane(api_key=settings.e2b_api_key or "") + return RuntimeBackendProfile( + home_snapshots=E2BHomeSnapshotBackend( + control_plane=control_plane, + ), + execution_bindings=E2BExecutionBindingBackend( + control_plane=control_plane, + template=settings.e2b_template, + active_timeout_seconds=settings.e2b_active_timeout_seconds, + shellctl_auth_token=settings.e2b_shellctl_auth_token, + shellctl_port=settings.e2b_shellctl_port, + ), + ) + + +def _validate_http_url(value: str, *, field_name: str) -> None: + parsed = urlparse(value.strip()) + if parsed.scheme not in {"http", "https"} or not parsed.netloc: + raise ValueError(f"{field_name} must be a valid http(s) URL") + + +def _validate_absolute_posix_path(value: str, *, field_name: str) -> None: + if not value.strip() or not posixpath.isabs(value): + raise ValueError(f"{field_name} must be an absolute POSIX path") + + +__all__ = [ + "DEFAULT_E2B_TEMPLATE", + "DEFAULT_LOCAL_HOME_SNAPSHOT_ROOT", + "DEFAULT_LOCAL_MATERIALIZED_HOME_ROOT", + "DEFAULT_LOCAL_WORKSPACE_ROOT", + "RuntimeBackendSettings", + "create_runtime_backend_profile", +] diff --git a/dify-agent/src/dify_agent/runtime_backend/protocols.py b/dify-agent/src/dify_agent/runtime_backend/protocols.py new file mode 100644 index 00000000000..7acd7da3dbe --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/protocols.py @@ -0,0 +1,163 @@ +"""Backend-neutral contracts for persistent working environments. + +Dify API owns logical Home Snapshot, Workspace, and Agent Workspace Binding +records. Backends own their physical representations. ``RuntimeLease`` is the +only invocation-local object and must never be serialized or persisted. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Protocol + +from dify_agent.adapters.shell.protocols import ShellCommandProtocol + + +@dataclass(frozen=True, slots=True) +class HomeSnapshotCreateSpec: + tenant_id: str + agent_id: str + home_snapshot_id: str + + +@dataclass(frozen=True, slots=True) +class RuntimeLayout: + """Canonical Home and Workspace roots exposed for one operation.""" + + home_dir: str + workspace_dir: str + + +class RuntimeLease(Protocol): + """Invocation-local data-plane access to one persistent Binding.""" + + @property + def layout(self) -> RuntimeLayout: ... + + @property + def commands(self) -> ShellCommandProtocol: ... + + +@dataclass(frozen=True, slots=True) +class ExecutionBindingCreateSpec: + tenant_id: str + agent_id: str + binding_id: str + workspace_id: str + existing_workspace_ref: str | None + home_snapshot_ref: str | None = None + + +@dataclass(frozen=True, slots=True) +class ExecutionBindingAllocation: + binding_ref: str + workspace_ref: str + + +@dataclass(frozen=True, slots=True) +class ExecutionBindingDestroySpec: + binding_ref: str + destroy_workspace: bool + workspace_ref: str | None = None + + def __post_init__(self) -> None: + if self.destroy_workspace and not self.workspace_ref: + raise ValueError("workspace_ref is required when destroy_workspace is true") + + +class ExecutionBindingBackend(Protocol): + """Manage physical Binding, Materialized Home, and Workspace resources.""" + + async def create_binding(self, spec: ExecutionBindingCreateSpec) -> ExecutionBindingAllocation: + """Materialize a mutable Home and make the requested Workspace ready. + + With a non-null ``home_snapshot_ref``, implementations must initialize + the Home from that exact immutable snapshot and must fail rather than + fall back when it is unavailable. With ``None``, implementations must + create an independent mutable Home from their deployment default without + implicitly creating an immutable snapshot. + + Implementations must return stable opaque refs only after Home and + Workspace are usable. With an ``existing_workspace_ref``, they must + attach that Workspace without clearing or replacing its contents; + unsupported sharing must fail before mutating it. On failure, + implementations should clean up newly allocated partial resources and + must not damage a pre-existing Workspace. + """ + ... + + async def acquire(self, binding_ref: str) -> RuntimeLease: + """Open fresh operation-scoped access to an existing physical Binding. + + Implementations must reactivate or reconnect the resource as needed, + verify that its materialized Home and Workspace still exist, and must not + silently create a replacement. Confirmed resource loss must raise + ``BindingLostError``; other acquisition failures must raise + ``BindingAcquireError``. + """ + ... + + async def release(self, lease: RuntimeLease) -> None: + """End operation-local access without destroying persistent resources. + + Implementations must close clients and owned transports and may suspend + the physical runtime. They must not delete or retire the Binding, + materialized Home, or Workspace, and must reject leases created by a + different backend implementation. + """ + ... + + async def destroy_binding(self, spec: ExecutionBindingDestroySpec) -> None: + """Destroy a Binding's materialized Home and optionally its Workspace. + + Implementations must honor ``destroy_workspace`` and validate + ``workspace_ref`` before destructive work. A backend unable to preserve + the Workspace must raise ``WorkspacePreservationUnsupportedError`` before + deleting anything. Cleanup must be safe to retry and must never delete + the immutable Home Snapshot from which the Binding was materialized. + """ + ... + + +class HomeSnapshotBackend(Protocol): + """Manage immutable backend-native Home resources.""" + + async def create_from_runtime(self, *, spec: HomeSnapshotCreateSpec, source: RuntimeLease) -> str: + """Capture the source lease's current Home as a new immutable snapshot. + + Implementations must not mutate or release ``source``. Workspace content + is not logically part of a Home Snapshot; if a provider can only snapshot + a coupled runtime, later materialization must prevent captured Workspace + state from becoming the new Binding's Workspace. + """ + ... + + async def delete(self, snapshot_ref: str) -> None: + """Delete one physical snapshot without affecting derived resources. + + Implementations must make deletion safe to retry, treat an already absent + snapshot as a successful outcome, and must not delete Bindings, + materialized Homes, or Workspaces created from the snapshot. + """ + ... + + +@dataclass(frozen=True, slots=True) +class RuntimeBackendProfile: + """Coherent Home and Binding backends selected once per deployment.""" + + home_snapshots: HomeSnapshotBackend + execution_bindings: ExecutionBindingBackend + + +__all__ = [ + "ExecutionBindingAllocation", + "ExecutionBindingBackend", + "ExecutionBindingCreateSpec", + "ExecutionBindingDestroySpec", + "HomeSnapshotBackend", + "HomeSnapshotCreateSpec", + "RuntimeBackendProfile", + "RuntimeLayout", + "RuntimeLease", +] diff --git a/dify-agent/src/dify_agent/runtime_backend/shellctl.py b/dify-agent/src/dify_agent/runtime_backend/shellctl.py new file mode 100644 index 00000000000..f947ea0d9ba --- /dev/null +++ b/dify-agent/src/dify_agent/runtime_backend/shellctl.py @@ -0,0 +1,149 @@ +"""Shared shellctl RuntimeLease used by every runtime backend.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +import logging +from typing import Protocol + +from dify_agent.adapters.shell.protocols import CompleteShellCommandResult, ShellCommandProtocol +from dify_agent.adapters.shell.shellctl import ( + ShellctlClientFactory, + ShellctlClientProtocol, + ShellctlCommands, + create_default_shellctl_client_factory, +) +from dify_agent.runtime_backend.protocols import RuntimeLayout + +_CONTROL_COMMAND_OUTPUT_LIMIT = 256 * 1024 +logger = logging.getLogger(__name__) + + +class AsyncCloseable(Protocol): + async def aclose(self) -> None: ... + + +@dataclass(slots=True) +class ShellctlRuntimeLease: + """One invocation-local shellctl connection and its canonical layout.""" + + handle: str + layout: RuntimeLayout + client: ShellctlClientProtocol + commands: ShellCommandProtocol + owned_transport: AsyncCloseable | None = None + _closed: bool = field(default=False, init=False) + + async def close(self) -> None: + if self._closed: + return + self._closed = True + client_error: BaseException | None = None + try: + await self.client.close() + except BaseException as exc: + client_error = exc + try: + if self.owned_transport is not None: + await self.owned_transport.aclose() + except BaseException as exc: + if client_error is None: + raise + logger.warning("Failed to close owned shellctl transport after client close failed: %s", exc) + if client_error is not None: + raise client_error + + +def create_shellctl_lease( + *, + handle: str, + layout: RuntimeLayout, + entrypoint: str, + token: str, + client_factory: ShellctlClientFactory | None = None, + owned_transport: AsyncCloseable | None = None, +) -> ShellctlRuntimeLease: + """Create adapters around one new shellctl client without owning control-plane lifecycle.""" + factory = client_factory or create_default_shellctl_client_factory(entrypoint=entrypoint, token=token) + client = factory() + return ShellctlRuntimeLease( + handle=handle, + layout=layout, + client=client, + commands=ShellctlCommands( + client=client, + home_dir=layout.home_dir, + workspace_dir=layout.workspace_dir, + ), + owned_transport=owned_transport, + ) + + +async def create_owned_shellctl_lease( + *, + handle: str, + layout: RuntimeLayout, + entrypoint: str, + token: str, + client_factory: ShellctlClientFactory, + owned_transport: AsyncCloseable, +) -> ShellctlRuntimeLease: + """Create a lease that owns an injected transport, closing it if construction fails.""" + try: + return create_shellctl_lease( + handle=handle, + layout=layout, + entrypoint=entrypoint, + token=token, + client_factory=client_factory, + owned_transport=owned_transport, + ) + except BaseException: + try: + await owned_transport.aclose() + except BaseException as cleanup_exc: + logger.warning("Failed to close owned shellctl transport after lease construction failed: %s", cleanup_exc) + raise + + +async def run_shellctl_control_command( + commands: ShellCommandProtocol, + script: str, + *, + timeout: float = 30.0, +) -> CompleteShellCommandResult: + """Run one bounded driver control command and always delete its transient job.""" + result = await commands.run(script, cwd=None, env=None, timeout=timeout) + job_id = result.job_id + output_parts = [result.output] + try: + while result.truncated or not result.done: + result = await commands.wait(result.job_id, offset=result.offset, timeout=timeout) + output_parts.append(result.output) + if sum(len(part.encode("utf-8")) for part in output_parts) > _CONTROL_COMMAND_OUTPUT_LIMIT: + raise RuntimeError("shellctl control command exceeded its output limit") + return CompleteShellCommandResult( + job_id=result.job_id, + status=result.status, + done=result.done, + exit_code=result.exit_code, + output="".join(output_parts), + output_complete=True, + incomplete_reason=None, + offset=result.offset, + output_path=result.output_path, + ) + finally: + try: + await commands.delete(job_id, force=True) + except Exception as exc: + logger.warning("Failed to delete transient shellctl control job %s: %s", job_id, exc) + + +__all__ = [ + "AsyncCloseable", + "ShellctlRuntimeLease", + "create_owned_shellctl_lease", + "create_shellctl_lease", + "run_shellctl_control_command", +] diff --git a/dify-agent/src/dify_agent/server/app.py b/dify-agent/src/dify_agent/server/app.py index 1013cf89bb0..567397f0cdc 100644 --- a/dify-agent/src/dify_agent/server/app.py +++ b/dify-agent/src/dify_agent/server/app.py @@ -7,12 +7,11 @@ rather than request handlers, so client disconnects do not cancel the agent runtime. Redis persists run records and per-run event streams with configured retention only; it is not used as a job queue. Agenton layers and providers stay state-only: they borrow the lifespan-owned clients through the runner and -receive shell-layer server settings through provider construction rather than -reading environment variables themselves. The standard server always mounts the -HTTP Agent Stub router and additionally starts the optional grpclib Agent Stub -server when ``DIFY_AGENT_STUB_API_BASE_URL`` uses ``grpc://``. Process-level -Logfire instrumentation is configured at app construction time and only exports -remotely when Logfire's default environment configuration provides a token. +receive runtime-backend and Shell settings through provider construction rather +than reading environment variables themselves. The standard server mounts the +HTTP Agent Stub router. Process-level Logfire instrumentation is configured at +app construction time and only exports remotely when Logfire's default +environment configuration provides a token. """ from collections.abc import AsyncGenerator @@ -23,16 +22,19 @@ from fastapi import FastAPI from redis.asyncio import Redis from dify_agent.agent_stub.shell_env import ShellAgentStubTokenFactory -from dify_agent.agent_stub.protocol.agent_stub import parse_agent_stub_endpoint -from dify_agent.agent_stub.server.grpc_runtime import start_agent_stub_grpc_server from dify_agent.agent_stub.server.router import create_agent_stub_router from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig from dify_agent.runtime.compositor_factory import create_default_layer_providers from dify_agent.runtime.run_scheduler import RunScheduler +from dify_agent.server.auth import create_bearer_token_dependency from dify_agent.server.observability import configure_server_observability from dify_agent.server.routes.runs import create_runs_router -from dify_agent.server.routes.sandbox_files import create_sandbox_files_router -from dify_agent.server.sandbox_files import SandboxFileService +from dify_agent.server.routes.execution_bindings import create_execution_bindings_router +from dify_agent.server.routes.home_snapshots import create_home_snapshots_router +from dify_agent.server.routes.binding_files import create_binding_files_router +from dify_agent.server.execution_bindings import ExecutionBindingService +from dify_agent.server.binding_files import BindingFileService +from dify_agent.server.home_snapshots import HomeSnapshotService from dify_agent.server.settings import ServerSettings from dify_agent.storage.redis_run_store import RedisRunStore @@ -43,8 +45,8 @@ def create_app(settings: ServerSettings | None = None) -> FastAPI: agent_stub_token_codec = resolved_settings.create_agent_stub_token_codec() agent_stub_token_factory: ShellAgentStubTokenFactory | None = None if agent_stub_token_codec is not None: - # Runtime receives only this callable boundary; router and gRPC wiring - # keep the concrete token codec on the server side. + # Runtime receives only this callable boundary; the HTTP router keeps + # the concrete token codec on the server side. def issue_agent_stub_token( execution_context: DifyExecutionContextLayerConfig, *, @@ -59,19 +61,39 @@ def create_app(settings: ServerSettings | None = None) -> FastAPI: agent_stub_file_request_handler = resolved_settings.create_agent_stub_file_request_handler() agent_stub_config_request_handler = resolved_settings.create_agent_stub_config_request_handler() agent_stub_drive_request_handler = resolved_settings.create_agent_stub_drive_request_handler() - shell_provider = resolved_settings.build_shell_provider() + runtime_backend_profile = resolved_settings.build_runtime_backend_profile() layer_providers = create_default_layer_providers( plugin_daemon_url=resolved_settings.plugin_daemon_url, plugin_daemon_api_key=resolved_settings.plugin_daemon_api_key, inner_api_url=resolved_settings.inner_api_url, inner_api_key=resolved_settings.inner_api_key or "", - shell_provider=shell_provider, - shell_home_root=resolved_settings.shell_home_root, + runtime_backend_profile=runtime_backend_profile, shell_redact_patterns=resolved_settings.get_shell_redact_patterns(), agent_stub_api_base_url=resolved_settings.agent_stub_api_base_url, agent_stub_token_factory=agent_stub_token_factory, ) - sandbox_file_service = SandboxFileService(layer_providers=layer_providers) if shell_provider is not None else None + binding_file_service = ( + BindingFileService( + execution_bindings=runtime_backend_profile.execution_bindings, + agent_stub_api_base_url=resolved_settings.agent_stub_api_base_url, + agent_stub_token_factory=agent_stub_token_factory, + ) + if runtime_backend_profile is not None + else None + ) + home_snapshot_service = ( + HomeSnapshotService( + home_snapshots=runtime_backend_profile.home_snapshots, + execution_bindings=runtime_backend_profile.execution_bindings, + ) + if runtime_backend_profile is not None + else None + ) + execution_binding_service = ( + ExecutionBindingService(backend=runtime_backend_profile.execution_bindings) + if runtime_backend_profile is not None + else None + ) state: dict[str, object] = {} @asynccontextmanager @@ -91,24 +113,11 @@ def create_app(settings: ServerSettings | None = None) -> FastAPI: shutdown_grace_seconds=resolved_settings.shutdown_grace_seconds, layer_providers=layer_providers, ) - grpc_server = None - if ( - resolved_settings.agent_stub_api_base_url is not None - and parse_agent_stub_endpoint(resolved_settings.agent_stub_api_base_url).is_grpc - ): - grpc_server = await start_agent_stub_grpc_server( - public_url=resolved_settings.agent_stub_api_base_url, - bind_address=resolved_settings.agent_stub_grpc_bind_address, - token_codec=agent_stub_token_codec, - file_request_handler=agent_stub_file_request_handler, - ) state["store"] = store state["scheduler"] = scheduler try: yield finally: - if grpc_server is not None: - await grpc_server.aclose() await scheduler.shutdown() await dify_api_inner_http_client.aclose() await plugin_daemon_http_client.aclose() @@ -123,8 +132,16 @@ def create_app(settings: ServerSettings | None = None) -> FastAPI: def get_scheduler() -> RunScheduler: return state["scheduler"] # pyright: ignore[reportReturnType] - app.include_router(create_runs_router(get_store, get_scheduler)) - app.include_router(create_sandbox_files_router(lambda: sandbox_file_service)) + app.include_router( + create_runs_router( + get_store, + get_scheduler, + auth_dependency=create_bearer_token_dependency(resolved_settings.api_token), + ) + ) + app.include_router(create_execution_bindings_router(lambda: execution_binding_service)) + app.include_router(create_home_snapshots_router(lambda: home_snapshot_service)) + app.include_router(create_binding_files_router(lambda: binding_file_service)) app.include_router( create_agent_stub_router( token_codec=agent_stub_token_codec, diff --git a/dify-agent/src/dify_agent/server/auth.py b/dify-agent/src/dify_agent/server/auth.py new file mode 100644 index 00000000000..a991639a3ec --- /dev/null +++ b/dify-agent/src/dify_agent/server/auth.py @@ -0,0 +1,43 @@ +import hmac + +from fastapi import Depends, Header, HTTPException + + +def create_bearer_token_dependency(expected_token: str | None): + """Return a FastAPI dependency that validates Bearer token authentication. + + When ``expected_token`` is ``None``, the returned dependency permits all + requests without checking the header, supporting graceful migration for + existing deployments. + """ + + async def require_bearer_token( + authorization: str | None = Header(default=None, alias="Authorization"), + ) -> None: + if expected_token is None: + return + if authorization is None: + raise HTTPException( + status_code=401, + detail="missing authorization header", + headers={"WWW-Authenticate": "Bearer"}, + ) + scheme, _, token = authorization.partition(" ") + token = token.strip() + if scheme.lower() != "bearer" or not token: + raise HTTPException( + status_code=401, + detail="invalid authorization scheme", + headers={"WWW-Authenticate": "Bearer"}, + ) + if not hmac.compare_digest(token.encode(), expected_token.encode()): + raise HTTPException( + status_code=401, + detail="invalid bearer token", + headers={"WWW-Authenticate": "Bearer"}, + ) + + return Depends(require_bearer_token) + + +__all__ = ["create_bearer_token_dependency"] diff --git a/dify-agent/src/dify_agent/server/binding_files.py b/dify-agent/src/dify_agent/server/binding_files.py new file mode 100644 index 00000000000..acb24094eaf --- /dev/null +++ b/dify-agent/src/dify_agent/server/binding_files.py @@ -0,0 +1,347 @@ +"""Binding filesystem operations through an operation-scoped RuntimeLease.""" + +from __future__ import annotations + +import base64 +import json +import logging +import posixpath +import re +import shlex +from dataclasses import dataclass +from typing import ClassVar, Literal + +from pydantic import BaseModel, ConfigDict, ValidationError + +from dify_agent.adapters.shell.protocols import CompleteShellCommandResult, ShellProviderError +from dify_agent.agent_stub.protocol import is_canonical_dify_file_reference +from dify_agent.agent_stub.shell_env import ShellAgentStubTokenFactory, build_shell_agent_stub_env +from dify_agent.protocol import ( + BindingFileDownloadRequest, + BindingFileDownloadResponse, + BindingFileListRequest, + BindingFileListResponse, + BindingFileReadRequest, + BindingFileReadResponse, +) +from dify_agent.runtime.command_runner import execute_complete_with_commands +from dify_agent.runtime_backend import ( + BindingAcquireError, + BindingLostError, + ExecutionBindingBackend, + RuntimeLayout, + WorkspaceUnavailableError, +) +from dify_agent.runtime_backend.leases import open_runtime_lease + +logger = logging.getLogger(__name__) + +_LIST_MAX_ENTRIES = 1000 +_BROWSE_TIMEOUT_SECONDS = 60.0 +_BROWSE_OUTPUT_MAX_BYTES = 1024 * 1024 +_DOWNLOAD_TIMEOUT_SECONDS = 60.0 +_DOWNLOAD_OUTPUT_MAX_BYTES = 32 * 1024 +_PAYLOAD_BEGIN = "<<>>" +_PAYLOAD_END = "<<>>" + +_LIST_BINDING_FILES_SCRIPT = r""" +import base64 +import json +import os +import stat +import sys + +path = sys.argv[1] +response_path = sys.argv[2] +limit = int(sys.argv[3]) +directory_fd = None +try: + directory_fd = os.open(path, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW) + names = sorted(os.listdir(directory_fd)) + entries = [] + for name in names[:limit]: + child_stat = os.stat(name, dir_fd=directory_fd, follow_symlinks=False) + mode = child_stat.st_mode + entry_type = ( + "symlink" if stat.S_ISLNK(mode) else + "dir" if stat.S_ISDIR(mode) else + "file" if stat.S_ISREG(mode) else + "other" + ) + entries.append({ + "name": name, + "type": entry_type, + "size": int(child_stat.st_size), + "mtime": int(child_stat.st_mtime), + }) +finally: + if directory_fd is not None: + os.close(directory_fd) + +payload = {"path": response_path, "entries": entries, "truncated": len(names) > limit} +blob = base64.b64encode(json.dumps(payload, separators=(",", ":")).encode()).decode() +print("<<>>" + blob + "<<>>") +""" + +_READ_BINDING_FILE_SCRIPT = r""" +import base64 +import json +import os +import stat +import sys + +path = sys.argv[1] +response_path = sys.argv[2] +max_bytes = int(sys.argv[3]) +file_fd = None +try: + file_fd = os.open(path, os.O_RDONLY | os.O_NOFOLLOW) + file_stat = os.fstat(file_fd) + if not stat.S_ISREG(file_stat.st_mode): + raise FileNotFoundError(path) + size = int(file_stat.st_size) + data = os.read(file_fd, max_bytes + 1) +finally: + if file_fd is not None: + os.close(file_fd) + +truncated = len(data) > max_bytes +data = data[:max_bytes] +try: + text = data.decode("utf-8") + binary = False +except UnicodeDecodeError: + text = None + binary = True +payload = { + "path": response_path, + "size": size, + "truncated": truncated, + "binary": binary, + "text": text, +} +blob = base64.b64encode(json.dumps(payload, separators=(",", ":")).encode()).decode() +print("<<>>" + blob + "<<>>") +""" + + +class BindingFileError(Exception): + code: str + message: str + status_code: int + + def __init__(self, code: str, message: str, *, status_code: int = 400) -> None: + super().__init__(message) + self.code = code + self.message = message + self.status_code = status_code + + +class _CliUploadResult(BaseModel): + transfer_method: Literal["tool_file"] + reference: str + + model_config: ClassVar[ConfigDict] = ConfigDict(extra="ignore") + + +@dataclass(slots=True) +class BindingFileService: + execution_bindings: ExecutionBindingBackend + agent_stub_api_base_url: str | None + agent_stub_token_factory: ShellAgentStubTokenFactory | None + + async def list_files(self, request: BindingFileListRequest) -> BindingFileListResponse: + try: + async with open_runtime_lease(self.execution_bindings, request.backend_binding_ref) as lease: + resolved_path = resolve_binding_path(request.path, lease.layout) + result = await execute_complete_with_commands( + lease.commands, + _python_command( + _LIST_BINDING_FILES_SCRIPT, + resolved_path, + request.path, + str(_LIST_MAX_ENTRIES), + ), + cwd=lease.layout.workspace_dir, + env={"HOME": lease.layout.home_dir}, + timeout=_BROWSE_TIMEOUT_SECONDS, + max_output_bytes=_BROWSE_OUTPUT_MAX_BYTES, + ) + payload = _require_browse_payload(result, operation="list") + try: + return BindingFileListResponse.model_validate(payload) + except ValidationError as exc: + raise WorkspaceUnavailableError("Binding file list returned an invalid response") from exc + except BindingFileError: + raise + except Exception as exc: + raise _normalize_binding_file_error(exc) from exc + + async def read_file(self, request: BindingFileReadRequest) -> BindingFileReadResponse: + try: + async with open_runtime_lease(self.execution_bindings, request.backend_binding_ref) as lease: + resolved_path = resolve_binding_path(request.path, lease.layout) + result = await execute_complete_with_commands( + lease.commands, + _python_command( + _READ_BINDING_FILE_SCRIPT, + resolved_path, + request.path, + str(request.max_bytes), + ), + cwd=lease.layout.workspace_dir, + env={"HOME": lease.layout.home_dir}, + timeout=_BROWSE_TIMEOUT_SECONDS, + max_output_bytes=_BROWSE_OUTPUT_MAX_BYTES, + ) + payload = _require_browse_payload(result, operation="read") + try: + return BindingFileReadResponse.model_validate(payload) + except ValidationError as exc: + raise WorkspaceUnavailableError("Binding file read returned an invalid response") from exc + except BindingFileError: + raise + except Exception as exc: + raise _normalize_binding_file_error(exc) from exc + + async def download_file(self, request: BindingFileDownloadRequest) -> BindingFileDownloadResponse: + context = request.execution_context + if not context.user_id or not context.user_from: + raise BindingFileError( + "invalid_execution_context", + "Binding file download requires user_id and user_from", + status_code=400, + ) + if self.agent_stub_api_base_url is None or self.agent_stub_token_factory is None: + raise BindingFileError( + "agent_stub_upload_unavailable", + "Agent Stub file upload is not configured", + status_code=503, + ) + + try: + agent_stub_env = build_shell_agent_stub_env( + agent_stub_api_base_url=self.agent_stub_api_base_url, + execution_context=context, + token_factory=self.agent_stub_token_factory, + session_id=None, + ) + if agent_stub_env is None: + raise BindingFileError( + "agent_stub_upload_unavailable", + "Agent Stub file upload is not configured", + status_code=503, + ) + async with open_runtime_lease(self.execution_bindings, request.backend_binding_ref) as lease: + resolved_path = resolve_binding_path(request.path, lease.layout) + env = {"HOME": lease.layout.home_dir, **agent_stub_env} + try: + result = await execute_complete_with_commands( + lease.commands, + f"dify-agent file upload --no-download-link {shlex.quote(resolved_path)}", + cwd=lease.layout.workspace_dir, + env=env, + timeout=_DOWNLOAD_TIMEOUT_SECONDS, + max_output_bytes=_DOWNLOAD_OUTPUT_MAX_BYTES, + ) + except ShellProviderError as exc: + if exc.code == "timeout": + raise _download_failed() from exc + raise + if result.exit_code != 0 or not result.output_complete: + _log_download_failure(result.output, agent_stub_env) + raise _download_failed() + try: + payload = json.loads(result.output) + uploaded = _CliUploadResult.model_validate(payload) + except (json.JSONDecodeError, ValidationError, TypeError) as exc: + raise _download_failed() from exc + if not is_canonical_dify_file_reference(uploaded.reference): + raise _download_failed() + return BindingFileDownloadResponse(reference=uploaded.reference) + except BindingFileError: + raise + except Exception as exc: + normalized = _normalize_binding_file_error(exc) + if normalized.code in {"binding_not_found", "binding_unavailable"}: + raise normalized from exc + raise _download_failed() from exc + + +def resolve_binding_path(path: str, layout: RuntimeLayout) -> str: + """Resolve convenient Binding paths without adding a containment policy.""" + + if path == "~": + candidate = layout.home_dir + elif path.startswith("~/"): + candidate = posixpath.join(layout.home_dir, path[2:]) + elif posixpath.isabs(path): + candidate = path + else: + candidate = posixpath.join(layout.workspace_dir, path or ".") + return posixpath.normpath(candidate) + + +def _python_command(source: str, *args: str) -> str: + return " ".join(["python3", "-c", shlex.quote(source), *(shlex.quote(arg) for arg in args)]) + + +def _require_browse_payload(result: CompleteShellCommandResult, *, operation: str) -> dict[str, object]: + exit_code = result.exit_code + output_complete = result.output_complete + output = result.output + if exit_code != 0: + if any( + name in output + for name in ("FileNotFoundError", "NotADirectoryError", "IsADirectoryError", "PermissionError") + ): + raise BindingFileError( + "invalid_binding_path", + f"Binding file {operation} path is unavailable", + status_code=400, + ) + raise WorkspaceUnavailableError(f"Binding file {operation} command failed") + if not output_complete: + raise WorkspaceUnavailableError(f"Binding file {operation} output was incomplete") + match = re.search(re.escape(_PAYLOAD_BEGIN) + r"(.*?)" + re.escape(_PAYLOAD_END), output, flags=re.DOTALL) + if match is None: + raise WorkspaceUnavailableError(f"Binding file {operation} returned no framed payload") + try: + decoded = base64.b64decode("".join(match.group(1).split()), validate=True) + payload = json.loads(decoded) + except (ValueError, json.JSONDecodeError) as exc: + raise WorkspaceUnavailableError(f"Binding file {operation} returned an invalid payload") from exc + if not isinstance(payload, dict): + raise WorkspaceUnavailableError(f"Binding file {operation} returned a non-object payload") + return payload + + +def _normalize_binding_file_error(exc: Exception) -> BindingFileError: + if isinstance(exc, BindingFileError): + return exc + if isinstance(exc, ValueError): + return BindingFileError("invalid_binding_path", "Binding file path or payload is invalid", status_code=400) + if isinstance(exc, BindingLostError): + return BindingFileError("binding_not_found", "Execution Binding was not found", status_code=404) + if isinstance(exc, BindingAcquireError | WorkspaceUnavailableError): + return BindingFileError("binding_unavailable", "Execution Binding is unavailable", status_code=502) + return BindingFileError("binding_unavailable", "Execution Binding file operation failed", status_code=502) + + +def _download_failed() -> BindingFileError: + return BindingFileError( + "binding_file_download_failed", + "Binding file could not be converted to a ToolFile", + status_code=502, + ) + + +def _log_download_failure(output: str, env: dict[str, str]) -> None: + redacted = output + for secret in env.values(): + if len(secret) > 8: + redacted = redacted.replace(secret, "***") + logger.warning("Binding file upload command failed: %s", redacted[-1024:]) + + +__all__ = ["BindingFileError", "BindingFileService", "resolve_binding_path"] diff --git a/dify-agent/src/dify_agent/server/execution_bindings.py b/dify-agent/src/dify_agent/server/execution_bindings.py new file mode 100644 index 00000000000..fbecb129daa --- /dev/null +++ b/dify-agent/src/dify_agent/server/execution_bindings.py @@ -0,0 +1,73 @@ +"""Stateless application facade for Execution Binding lifecycle operations.""" + +from dataclasses import dataclass + +from dify_agent.protocol import ( + CreateExecutionBindingRequest, + CreateExecutionBindingResponse, + DestroyExecutionBindingRequest, +) +from dify_agent.runtime_backend import ( + BindingCreateError, + BindingDestroyError, + ExecutionBindingBackend, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + SharedWorkspaceUnsupportedError, + WorkspacePreservationUnsupportedError, +) + + +class ExecutionBindingServiceError(RuntimeError): + code: str + message: str + status_code: int + + def __init__(self, code: str, message: str, *, status_code: int) -> None: + super().__init__(message) + self.code = code + self.message = message + self.status_code = status_code + + +@dataclass(slots=True) +class ExecutionBindingService: + backend: ExecutionBindingBackend + + async def create_binding(self, request: CreateExecutionBindingRequest) -> CreateExecutionBindingResponse: + try: + allocation = await self.backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id=request.tenant_id, + agent_id=request.agent_id, + binding_id=request.binding_id, + workspace_id=request.workspace_id, + existing_workspace_ref=request.existing_workspace_ref, + home_snapshot_ref=request.home_snapshot_ref, + ) + ) + except SharedWorkspaceUnsupportedError as exc: + raise ExecutionBindingServiceError("shared_workspace_unsupported", str(exc), status_code=409) from exc + except BindingCreateError as exc: + raise ExecutionBindingServiceError("binding_create_failed", str(exc), status_code=502) from exc + return CreateExecutionBindingResponse( + binding_ref=allocation.binding_ref, + workspace_ref=allocation.workspace_ref, + ) + + async def destroy_binding(self, request: DestroyExecutionBindingRequest) -> None: + try: + await self.backend.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref=request.binding_ref, + destroy_workspace=request.destroy_workspace, + workspace_ref=request.workspace_ref, + ) + ) + except WorkspacePreservationUnsupportedError as exc: + raise ExecutionBindingServiceError("workspace_preservation_unsupported", str(exc), status_code=409) from exc + except BindingDestroyError as exc: + raise ExecutionBindingServiceError("binding_destroy_failed", str(exc), status_code=502) from exc + + +__all__ = ["ExecutionBindingService", "ExecutionBindingServiceError"] diff --git a/dify-agent/src/dify_agent/server/home_snapshots.py b/dify-agent/src/dify_agent/server/home_snapshots.py new file mode 100644 index 00000000000..27ac62fc44d --- /dev/null +++ b/dify-agent/src/dify_agent/server/home_snapshots.py @@ -0,0 +1,68 @@ +"""Stateless application facade for deployment-owned Home Snapshots.""" + +from dataclasses import dataclass + +from dify_agent.protocol.home_snapshot import ( + CreateHomeSnapshotFromBindingRequest, + DeleteHomeSnapshotRequest, + HomeSnapshotResponse, +) +from dify_agent.runtime_backend import ( + BindingAcquireError, + BindingLostError, + ExecutionBindingBackend, + HomeSnapshotBackend, + HomeSnapshotCreateError, + HomeSnapshotCreateSpec, + HomeSnapshotNotFoundError, +) +from dify_agent.runtime_backend.leases import open_runtime_lease + + +class HomeSnapshotServiceError(RuntimeError): + code: str + message: str + status_code: int + + def __init__(self, code: str, message: str, *, status_code: int) -> None: + super().__init__(message) + self.code = code + self.message = message + self.status_code = status_code + + +@dataclass(slots=True) +class HomeSnapshotService: + home_snapshots: HomeSnapshotBackend + execution_bindings: ExecutionBindingBackend + + async def create_from_binding( + self, + request: CreateHomeSnapshotFromBindingRequest, + ) -> HomeSnapshotResponse: + try: + async with open_runtime_lease(self.execution_bindings, request.backend_binding_ref) as lease: + snapshot_ref = await self.home_snapshots.create_from_runtime( + spec=HomeSnapshotCreateSpec( + tenant_id=request.tenant_id, + agent_id=request.agent_id, + home_snapshot_id=request.home_snapshot_id, + ), + source=lease, + ) + except BindingLostError as exc: + raise HomeSnapshotServiceError("binding_lost", str(exc), status_code=404) from exc + except (BindingAcquireError, HomeSnapshotCreateError) as exc: + raise HomeSnapshotServiceError("home_snapshot_create_failed", str(exc), status_code=502) from exc + return HomeSnapshotResponse(snapshot_ref=snapshot_ref) + + async def delete(self, request: DeleteHomeSnapshotRequest) -> None: + try: + await self.home_snapshots.delete(request.snapshot_ref) + except HomeSnapshotNotFoundError: + return + except RuntimeError as exc: + raise HomeSnapshotServiceError("home_snapshot_delete_failed", str(exc), status_code=502) from exc + + +__all__ = ["HomeSnapshotService", "HomeSnapshotServiceError"] diff --git a/dify-agent/src/dify_agent/server/routes/binding_files.py b/dify-agent/src/dify_agent/server/routes/binding_files.py new file mode 100644 index 00000000000..c8ab0c2bc24 --- /dev/null +++ b/dify-agent/src/dify_agent/server/routes/binding_files.py @@ -0,0 +1,107 @@ +"""Private Binding file routes used by Dify API.""" + +from collections.abc import Callable, Coroutine +from typing import Annotated, Any, override + +from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi.exceptions import RequestValidationError +from fastapi.responses import JSONResponse +from fastapi.routing import APIRoute +from starlette.responses import Response + +from dify_agent.protocol import ( + BindingFileDownloadRequest, + BindingFileDownloadResponse, + BindingFileListRequest, + BindingFileListResponse, + BindingFileReadRequest, + BindingFileReadResponse, +) +from dify_agent.server.binding_files import BindingFileError, BindingFileService + +_INVALID_BINDING_PATH_MESSAGE = "Binding file path or payload is invalid" +_BROWSE_PATHS = frozenset( + { + "/execution-bindings/files/list", + "/execution-bindings/files/read", + } +) + + +class _BindingFileValidationRoute(APIRoute): + @override + def get_route_handler(self) -> Callable[[Request], Coroutine[Any, Any, Response]]: + route_handler = super().get_route_handler() + + async def handle(request: Request) -> Response: + try: + return await route_handler(request) + except RequestValidationError: + if self.path not in _BROWSE_PATHS: + raise + return JSONResponse( + status_code=400, + content={ + "detail": { + "code": "invalid_binding_path", + "message": _INVALID_BINDING_PATH_MESSAGE, + } + }, + ) + + return handle + + +def create_binding_files_router(get_service: Callable[[], BindingFileService | None]) -> APIRouter: + router = APIRouter( + prefix="/execution-bindings/files", + tags=["execution-bindings"], + route_class=_BindingFileValidationRoute, + ) + + def service_dep() -> BindingFileService: + service = get_service() + if service is None: + raise HTTPException( + status_code=503, + detail={"code": "runtime_backend_unavailable", "message": "Binding file service is not configured"}, + ) + return service + + def raise_http(exc: BindingFileError) -> HTTPException: + return HTTPException(status_code=exc.status_code, detail={"code": exc.code, "message": exc.message}) + + @router.post("/list", response_model=BindingFileListResponse) + async def list_files( + request: BindingFileListRequest, + service: Annotated[BindingFileService, Depends(service_dep)], + ) -> BindingFileListResponse: + try: + return await service.list_files(request) + except BindingFileError as exc: + raise raise_http(exc) from exc + + @router.post("/read", response_model=BindingFileReadResponse) + async def read_file( + request: BindingFileReadRequest, + service: Annotated[BindingFileService, Depends(service_dep)], + ) -> BindingFileReadResponse: + try: + return await service.read_file(request) + except BindingFileError as exc: + raise raise_http(exc) from exc + + @router.post("/download", response_model=BindingFileDownloadResponse) + async def download_file( + request: BindingFileDownloadRequest, + service: Annotated[BindingFileService, Depends(service_dep)], + ) -> BindingFileDownloadResponse: + try: + return await service.download_file(request) + except BindingFileError as exc: + raise raise_http(exc) from exc + + return router + + +__all__ = ["create_binding_files_router"] diff --git a/dify-agent/src/dify_agent/server/routes/execution_bindings.py b/dify-agent/src/dify_agent/server/routes/execution_bindings.py new file mode 100644 index 00000000000..607574a36e1 --- /dev/null +++ b/dify-agent/src/dify_agent/server/routes/execution_bindings.py @@ -0,0 +1,55 @@ +"""Private Execution Binding control-plane routes used by Dify API.""" + +from collections.abc import Callable +from typing import Annotated + +from fastapi import APIRouter, Depends, HTTPException, Response, status + +from dify_agent.protocol import ( + CreateExecutionBindingRequest, + CreateExecutionBindingResponse, + DestroyExecutionBindingRequest, +) +from dify_agent.server.execution_bindings import ExecutionBindingService, ExecutionBindingServiceError + + +def create_execution_bindings_router(get_service: Callable[[], ExecutionBindingService | None]) -> APIRouter: + router = APIRouter(prefix="/execution-bindings", tags=["execution-bindings"]) + + def service_dep() -> ExecutionBindingService: + service = get_service() + if service is None: + raise HTTPException( + status_code=503, + detail={"code": "runtime_backend_unavailable", "message": "runtime backend is not configured"}, + ) + return service + + def raise_http(exc: ExecutionBindingServiceError) -> HTTPException: + return HTTPException(status_code=exc.status_code, detail={"code": exc.code, "message": exc.message}) + + @router.post("", response_model=CreateExecutionBindingResponse, status_code=status.HTTP_201_CREATED) + async def create_binding( + request: CreateExecutionBindingRequest, + service: Annotated[ExecutionBindingService, Depends(service_dep)], + ) -> CreateExecutionBindingResponse: + try: + return await service.create_binding(request) + except ExecutionBindingServiceError as exc: + raise raise_http(exc) from exc + + @router.post("/destroy", status_code=status.HTTP_204_NO_CONTENT) + async def destroy_binding( + request: DestroyExecutionBindingRequest, + service: Annotated[ExecutionBindingService, Depends(service_dep)], + ) -> Response: + try: + await service.destroy_binding(request) + except ExecutionBindingServiceError as exc: + raise raise_http(exc) from exc + return Response(status_code=status.HTTP_204_NO_CONTENT) + + return router + + +__all__ = ["create_execution_bindings_router"] diff --git a/dify-agent/src/dify_agent/server/routes/home_snapshots.py b/dify-agent/src/dify_agent/server/routes/home_snapshots.py new file mode 100644 index 00000000000..8f508bd41de --- /dev/null +++ b/dify-agent/src/dify_agent/server/routes/home_snapshots.py @@ -0,0 +1,58 @@ +"""Private Home Snapshot control-plane routes used by Dify API.""" + +from collections.abc import Callable +from typing import Annotated + +from fastapi import APIRouter, Depends, HTTPException, Response, status + +from dify_agent.protocol.home_snapshot import ( + CreateHomeSnapshotFromBindingRequest, + DeleteHomeSnapshotRequest, + HomeSnapshotResponse, +) +from dify_agent.server.home_snapshots import HomeSnapshotService, HomeSnapshotServiceError + + +def create_home_snapshots_router(get_service: Callable[[], HomeSnapshotService | None]) -> APIRouter: + router = APIRouter(prefix="/home-snapshots", tags=["home-snapshots"]) + + def service_dep() -> HomeSnapshotService: + service = get_service() + if service is None: + raise HTTPException( + status_code=503, + detail={"code": "runtime_backend_unavailable", "message": "runtime backend is not configured"}, + ) + return service + + @router.post("/from-binding", response_model=HomeSnapshotResponse, status_code=status.HTTP_201_CREATED) + async def create_snapshot_from_binding( + request: CreateHomeSnapshotFromBindingRequest, + service: Annotated[HomeSnapshotService, Depends(service_dep)], + ) -> HomeSnapshotResponse: + try: + return await service.create_from_binding(request) + except HomeSnapshotServiceError as exc: + raise HTTPException( + status_code=exc.status_code, + detail={"code": exc.code, "message": exc.message}, + ) from exc + + @router.post("/delete", status_code=status.HTTP_204_NO_CONTENT) + async def delete_snapshot( + request: DeleteHomeSnapshotRequest, + service: Annotated[HomeSnapshotService, Depends(service_dep)], + ) -> Response: + try: + await service.delete(request) + except HomeSnapshotServiceError as exc: + raise HTTPException( + status_code=exc.status_code, + detail={"code": exc.code, "message": exc.message}, + ) from exc + return Response(status_code=status.HTTP_204_NO_CONTENT) + + return router + + +__all__ = ["create_home_snapshots_router"] diff --git a/dify-agent/src/dify_agent/server/routes/runs.py b/dify-agent/src/dify_agent/server/routes/runs.py index f41567648e8..4f08e76aa9f 100644 --- a/dify-agent/src/dify_agent/server/routes/runs.py +++ b/dify-agent/src/dify_agent/server/routes/runs.py @@ -14,6 +14,7 @@ from collections.abc import Callable from typing import Annotated from fastapi import APIRouter, Depends, Header, HTTPException, Query +from fastapi.params import Depends as DependsInstance from fastapi.responses import StreamingResponse from dify_agent.protocol.schemas import ( @@ -32,9 +33,11 @@ from dify_agent.storage.redis_run_store import RedisRunStore, RunNotFoundError def create_runs_router( get_store: Callable[[], RedisRunStore], get_scheduler: Callable[[], RunScheduler], + *, + auth_dependency: DependsInstance | None = None, ) -> APIRouter: - """Create routes bound to the application's store dependency provider.""" - router = APIRouter(prefix="/runs", tags=["runs"]) + dependencies: list[DependsInstance] = [auth_dependency] if auth_dependency is not None else [] + router = APIRouter(prefix="/runs", tags=["runs"], dependencies=dependencies) async def store_dep() -> RedisRunStore: return get_store() @@ -65,6 +68,7 @@ def create_runs_router( created_at=record.created_at, updated_at=record.updated_at, error=record.error, + error_type=record.error_type, ) @router.post("/{run_id}/cancel", response_model=CancelRunResponse) @@ -73,7 +77,7 @@ def create_runs_router( request: CancelRunRequest, scheduler: Annotated[RunScheduler, Depends(scheduler_dep)], ) -> CancelRunResponse: - """Cancel a process-local run and publish its terminal event/status.""" + """Persist cancellation; the owner process observes it and stops its runner.""" try: return await scheduler.cancel_run(run_id, request) except RunNotFoundError as exc: diff --git a/dify-agent/src/dify_agent/server/routes/sandbox_files.py b/dify-agent/src/dify_agent/server/routes/sandbox_files.py deleted file mode 100644 index 10429620cc9..00000000000 --- a/dify-agent/src/dify_agent/server/routes/sandbox_files.py +++ /dev/null @@ -1,73 +0,0 @@ -"""FastAPI routes for sandbox file operations. - -The agent backend receives a structured ``SandboxLocator`` rather than a raw -shell session id. Routes stay private-network only like ``/runs`` and forward -all sandbox work to ``SandboxFileService``. -""" - -from collections.abc import Callable -from typing import Annotated - -from fastapi import APIRouter, Depends, HTTPException - -from dify_agent.protocol import ( - SandboxListRequest, - SandboxListResponse, - SandboxReadRequest, - SandboxReadResponse, - SandboxUploadRequest, - SandboxUploadResponse, -) -from dify_agent.server.sandbox_files import SandboxFileError, SandboxFileService - - -def create_sandbox_files_router(get_service: Callable[[], SandboxFileService | None]) -> APIRouter: - """Create sandbox file routes bound to the app's service provider.""" - router = APIRouter(prefix="/sandbox", tags=["sandbox"]) - - def service_dep() -> SandboxFileService: - service = get_service() - if service is None: - raise HTTPException( - status_code=503, - detail={"code": "sandbox_backend_unavailable", "message": "sandbox service is not configured"}, - ) - return service - - def raise_http(exc: SandboxFileError) -> HTTPException: - return HTTPException(status_code=exc.status_code, detail={"code": exc.code, "message": exc.message}) - - @router.post("/files/list", response_model=SandboxListResponse) - async def list_files( - request: SandboxListRequest, - service: Annotated[SandboxFileService, Depends(service_dep)], - ) -> SandboxListResponse: - try: - return await service.list_files(request) - except SandboxFileError as exc: - raise raise_http(exc) from exc - - @router.post("/files/read", response_model=SandboxReadResponse) - async def read_file( - request: SandboxReadRequest, - service: Annotated[SandboxFileService, Depends(service_dep)], - ) -> SandboxReadResponse: - try: - return await service.read_file(request) - except SandboxFileError as exc: - raise raise_http(exc) from exc - - @router.post("/files/upload", response_model=SandboxUploadResponse) - async def upload_file( - request: SandboxUploadRequest, - service: Annotated[SandboxFileService, Depends(service_dep)], - ) -> SandboxUploadResponse: - try: - return await service.upload_file(request) - except SandboxFileError as exc: - raise raise_http(exc) from exc - - return router - - -__all__ = ["create_sandbox_files_router"] diff --git a/dify-agent/src/dify_agent/server/sandbox_files.py b/dify-agent/src/dify_agent/server/sandbox_files.py deleted file mode 100644 index 9d712f2db3e..00000000000 --- a/dify-agent/src/dify_agent/server/sandbox_files.py +++ /dev/null @@ -1,428 +0,0 @@ -"""Sandbox file service that re-enters prior shell sessions through the shell layer. - -Unlike the removed workspace inspector, this service never talks to shellctl -directly and never reads sandbox files outside the shell layer. It rebuilds a -minimal compositor from ``SandboxLocator``, enters the saved -``execution_context`` + ``shell`` layers, and executes fixed scripts through -``DifyShellLayer.run_remote_script_complete()``. - -The scripts still frame their structured payloads with a PTY-safe -base64-between-sentinels envelope. shellctl jobs are tmux-backed, so raw JSON can -be wrapped or surrounded by prompt noise; the framing keeps list/read/upload -responses parseable without falling back to direct shellctl file access. Path -arguments resolve from the saved shell workspace cwd, and ``~`` resolves through -the shell layer's injected sandbox ``HOME``. The scripts do not re-impose a -workspace-root boundary, so callers can use ``../`` or ``~/...`` when the sandbox -filesystem layout expects it. -""" - -from __future__ import annotations - -import json -import base64 -import binascii -import shlex -import textwrap -from dataclasses import dataclass -from typing import TypeVar, cast - -from dify_agent.layers.shell.layer import CompleteRemoteCommandResult, DifyShellLayer -from dify_agent.layers.shell.output_text import utf8_suffix -from dify_agent.protocol import ( - SandboxListRequest, - SandboxListResponse, - SandboxLocator, - SandboxReadRequest, - SandboxReadResponse, - SandboxUploadRequest, - SandboxUploadResponse, - normalize_composition, -) -from pydantic import BaseModel, ValidationError -from dify_agent.runtime.compositor_factory import DifyAgentLayerProvider, build_pydantic_ai_compositor - -_LIST_MAX_ENTRIES = 1000 -_LIST_TIMEOUT_SECONDS = 10.0 -_READ_TIMEOUT_SECONDS = 15.0 -_UPLOAD_TIMEOUT_SECONDS = 30.0 -_OUTPUT_BEGIN = "<<>>" -_OUTPUT_END = "<<>>" -_SHELL_RESULT_OUTPUT_TAIL_BYTES = 8 * 1024 -ResponseModelT = TypeVar("ResponseModelT", bound=BaseModel) - -_LIST_SCRIPT = """ -import base64 -import json -import stat -import sys -from pathlib import Path - - -BEGIN = "<<>>" -END = "<<>>" - - -def emit(payload): - blob = base64.b64encode(json.dumps(payload, ensure_ascii=False).encode("utf-8")).decode("ascii") - print(BEGIN + blob + END) - - -raw_path = sys.argv[1] -limit = int(sys.argv[2]) -target = Path(raw_path).expanduser().resolve() - -if not target.exists(): - emit({"error": "sandbox_path_not_found", "message": "path not found in sandbox"}) - sys.exit(0) -if not target.is_dir(): - emit({"error": "sandbox_path_not_readable", "message": "path is not a directory"}) - sys.exit(0) - -entries = [] -for child in sorted(target.iterdir(), key=lambda item: item.name)[:limit]: - child_stat = child.lstat() - mode = child_stat.st_mode - if stat.S_ISLNK(mode): - entry_type = "symlink" - elif stat.S_ISDIR(mode): - entry_type = "dir" - elif stat.S_ISREG(mode): - entry_type = "file" - else: - entry_type = "other" - entries.append( - { - "name": child.name, - "type": entry_type, - "size": int(child_stat.st_size), - "mtime": int(child_stat.st_mtime), - } - ) - -emit( - { - "path": raw_path, - "entries": entries, - "truncated": len(list(target.iterdir())) > limit, - } -) -""" - -_READ_SCRIPT = """ -import base64 -import json -import sys -from pathlib import Path - - -BEGIN = "<<>>" -END = "<<>>" - - -def emit(payload): - blob = base64.b64encode(json.dumps(payload, ensure_ascii=False).encode("utf-8")).decode("ascii") - print(BEGIN + blob + END) - - -raw_path = sys.argv[1] -max_bytes = int(sys.argv[2]) -target = Path(raw_path).expanduser().resolve() -if not target.exists(): - emit({"error": "sandbox_path_not_found", "message": "path not found in sandbox"}) - sys.exit(0) -if not target.is_file(): - emit({"error": "sandbox_path_not_readable", "message": "path is not a readable file"}) - sys.exit(0) - -size = int(target.stat().st_size) -with target.open("rb") as file_obj: - data = file_obj.read(max_bytes + 1) - -truncated = len(data) > max_bytes -data = data[:max_bytes] -try: - text = data.decode("utf-8") -except UnicodeDecodeError: - emit( - { - "path": raw_path, - "size": size, - "truncated": truncated, - "binary": True, - "text": None, - } - ) - sys.exit(0) - -emit( - { - "path": raw_path, - "size": size, - "truncated": truncated, - "binary": False, - "text": text, - } -) -""" - -_UPLOAD_SCRIPT = """ -import base64 -import json -import subprocess -import sys -from pathlib import Path - - -BEGIN = "<<>>" -END = "<<>>" - - -def emit(payload): - blob = base64.b64encode(json.dumps(payload, ensure_ascii=False).encode("utf-8")).decode("ascii") - print(BEGIN + blob + END) - - -raw_path = sys.argv[1] -target = Path(raw_path).expanduser().resolve() -if not target.exists(): - emit({"error": "sandbox_path_not_found", "message": "path not found in sandbox"}) - sys.exit(0) -if not target.is_file(): - emit({"error": "sandbox_path_not_readable", "message": "path is not a readable file"}) - sys.exit(0) - -command = ["dify-agent", "file", "upload", raw_path] -completed = subprocess.run(command, capture_output=True, text=True, check=False) -if completed.returncode != 0: - emit( - { - "error": "agent_stub_upload_failed", - "message": (completed.stderr or completed.stdout or f"upload exited with code {completed.returncode}").strip(), - } - ) - sys.exit(0) - -try: - file_mapping = json.loads(completed.stdout) -except ValueError as exc: - emit({"error": "agent_stub_upload_failed", "message": f"upload returned invalid JSON: {exc}"}) - sys.exit(0) - -emit({"path": raw_path, "file": file_mapping}) -""" - - -class SandboxFileError(Exception): - """Sandbox file failure mapped to HTTP by the FastAPI route layer.""" - - code: str - message: str - status_code: int - - def __init__(self, code: str, message: str, *, status_code: int = 400) -> None: - super().__init__(message) - self.code = code - self.message = message - self.status_code = status_code - - -@dataclass(slots=True) -class SandboxFileService: - """Execute fixed sandbox file operations through the saved shell session.""" - - layer_providers: tuple[DifyAgentLayerProvider, ...] - - async def list_files(self, request: SandboxListRequest) -> SandboxListResponse: - normalized_path = _normalize_sandbox_path(request.path, allow_current_directory=True) - payload = await self._run_locator_script( - request.locator, - script_source=_LIST_SCRIPT, - args=[normalized_path, str(_LIST_MAX_ENTRIES)], - timeout=_LIST_TIMEOUT_SECONDS, - inject_agent_stub_env=False, - ) - return _validate_response_model(SandboxListResponse, payload) - - async def read_file(self, request: SandboxReadRequest) -> SandboxReadResponse: - normalized_path = _normalize_sandbox_path(request.path, allow_current_directory=False) - payload = await self._run_locator_script( - request.locator, - script_source=_READ_SCRIPT, - args=[normalized_path, str(request.max_bytes)], - timeout=_READ_TIMEOUT_SECONDS, - inject_agent_stub_env=False, - ) - return _validate_response_model(SandboxReadResponse, payload) - - async def upload_file(self, request: SandboxUploadRequest) -> SandboxUploadResponse: - normalized_path = _normalize_sandbox_path(request.path, allow_current_directory=False) - payload = await self._run_locator_script( - request.locator, - script_source=_UPLOAD_SCRIPT, - args=[normalized_path], - timeout=_UPLOAD_TIMEOUT_SECONDS, - inject_agent_stub_env=True, - ) - return _validate_response_model(SandboxUploadResponse, payload) - - async def _run_locator_script( - self, - locator: SandboxLocator, - *, - script_source: str, - args: list[str], - timeout: float, - inject_agent_stub_env: bool, - ) -> dict[str, object]: - try: - graph_config, layer_configs = normalize_composition(locator.composition) - compositor = build_pydantic_ai_compositor(graph_config, providers=self.layer_providers) - async with compositor.enter(configs=layer_configs, session_snapshot=locator.session_snapshot) as run: - run.suspend_on_exit() - shell_layer = run.get_layer("shell", DifyShellLayer) - result = await shell_layer.run_remote_script_complete( - _build_python_script_command(script_source=script_source, args=args), - timeout=timeout, - inject_agent_stub_env=inject_agent_stub_env, - ) - except (KeyError, TypeError, ValueError) as exc: - raise SandboxFileError("invalid_sandbox_locator", str(exc), status_code=400) from exc - except RuntimeError as exc: - raise SandboxFileError("sandbox_command_failed", str(exc), status_code=502) from exc - - return _decode_sandbox_payload(result) - - -def _normalize_sandbox_path(path: str, *, allow_current_directory: bool) -> str: - """Reject only syntactically unsafe paths and preserve relative traversal. - - The remote scripts run with the saved workspace cwd, so ``../`` remains a - valid sandbox-relative path when callers need to reach sibling directories. - ``~`` and ``~/...`` are also accepted and resolved by the embedded scripts - against the shell layer's sandbox ``HOME``. - """ - - normalized = (path or "").strip() - if normalized in {"", ".", "./"}: - if allow_current_directory: - return "." - raise SandboxFileError("invalid_sandbox_path", "path must not be blank", status_code=400) - if normalized.startswith("/"): - raise SandboxFileError( - "invalid_sandbox_path", "path must be relative to the sandbox workspace", status_code=400 - ) - if normalized.startswith("~") and normalized != "~" and not normalized.startswith("~/"): - raise SandboxFileError("invalid_sandbox_path", "path must use ~ or ~/ for the sandbox home", status_code=400) - if "\x00" in normalized or any(ord(ch) < 0x20 for ch in normalized): - raise SandboxFileError("invalid_sandbox_path", "path contains unsupported control characters", status_code=400) - return normalized - - -def _build_python_script_command(*, script_source: str, args: list[str]) -> str: - quoted_args = " ".join(shlex.quote(value) for value in args) - script = textwrap.dedent(script_source).strip() - return f"python3 - {quoted_args} <<'PY'\n{script}\nPY" - - -def _decode_sandbox_payload(result: CompleteRemoteCommandResult) -> dict[str, object]: - if result.exit_code not in (0, None): - raise SandboxFileError( - "sandbox_command_failed", - "sandbox command exited with code " + f"{result.exit_code}: {_shell_result_details(result)}", - status_code=502, - ) - begin = result.output.find(_OUTPUT_BEGIN) - end = result.output.find(_OUTPUT_END, begin + len(_OUTPUT_BEGIN)) if begin != -1 else -1 - if begin == -1 or end == -1: - if not result.output_complete: - raise SandboxFileError( - "sandbox_command_failed", - "sandbox command output incomplete before framed payload was captured: " - + _shell_result_details(result), - status_code=502, - ) - raise SandboxFileError( - "sandbox_command_failed", - "sandbox command returned no framed payload", - status_code=502, - ) - blob = result.output[begin + len(_OUTPUT_BEGIN) : end] - compact = "".join(blob.split()) - try: - decoded = base64.b64decode(compact, validate=True) - loaded = cast(object, json.loads(decoded.decode("utf-8"))) - except (binascii.Error, ValueError) as exc: - if not result.output_complete: - raise SandboxFileError( - "sandbox_command_failed", - "sandbox command output incomplete while decoding framed payload: " + _shell_result_details(result), - status_code=502, - ) from exc - raise SandboxFileError( - "sandbox_command_failed", - f"sandbox command returned invalid framed payload: {exc}", - status_code=502, - ) from exc - if not isinstance(loaded, dict): - if not result.output_complete: - raise SandboxFileError( - "sandbox_command_failed", - "sandbox command output incomplete while validating framed payload object: " - + _shell_result_details(result), - status_code=502, - ) - raise SandboxFileError( - "sandbox_command_failed", "sandbox command returned a non-object payload", status_code=502 - ) - payload = cast(dict[str, object], loaded) - error = payload.get("error") - if isinstance(error, str): - status_code = ( - 404 - if error in {"sandbox_not_found", "sandbox_path_not_found"} - else 502 - if error == "agent_stub_upload_failed" - else 400 - ) - if error in {"sandbox_command_failed", "agent_stub_upload_failed"}: - status_code = 502 - message = payload.get("message") - raise SandboxFileError( - error, - str(message) if isinstance(message, str) and message else error, - status_code=status_code, - ) - return payload - - -def _shell_result_details(result: CompleteRemoteCommandResult) -> str: - details = ( - f"output_complete={result.output_complete} " - + f"incomplete_reason={result.incomplete_reason} " - + f"output_path={result.output_path}" - ) - if not result.output: - return details - return details + "\n" + _bounded_output_tail(result.output) - - -def _bounded_output_tail(output: str) -> str: - tail = utf8_suffix(output, _SHELL_RESULT_OUTPUT_TAIL_BYTES) - if tail == output: - return output - return f"... (showing last {_SHELL_RESULT_OUTPUT_TAIL_BYTES} bytes of raw output) ...\n{tail}" - - -def _validate_response_model( - model_type: type[ResponseModelT], - payload: dict[str, object], -) -> ResponseModelT: - try: - return model_type.model_validate(payload) - except ValidationError as exc: - raise SandboxFileError( - "sandbox_command_failed", f"sandbox command returned invalid payload: {exc}", status_code=502 - ) from exc - - -__all__ = ["SandboxFileError", "SandboxFileService"] diff --git a/dify-agent/src/dify_agent/server/schemas.py b/dify-agent/src/dify_agent/server/schemas.py index 21e8e624a0b..879d1b211e2 100644 --- a/dify-agent/src/dify_agent/server/schemas.py +++ b/dify-agent/src/dify_agent/server/schemas.py @@ -33,6 +33,7 @@ class RunRecord(BaseModel): created_at: datetime = Field(default_factory=_protocol_schemas.utc_now) updated_at: datetime = Field(default_factory=_protocol_schemas.utc_now) error: str | None = None + error_type: _protocol_schemas.RunFailureType | None = None model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid") diff --git a/dify-agent/src/dify_agent/server/settings.py b/dify-agent/src/dify_agent/server/settings.py index d6e8dfbcb4f..177456dc47a 100644 --- a/dify-agent/src/dify_agent/server/settings.py +++ b/dify-agent/src/dify_agent/server/settings.py @@ -6,32 +6,37 @@ Dify API inner calls. Layers and Agenton providers do not own those clients, so these settings are process resource limits rather than per-run lifecycle knobs. Endpoint URLs and API keys stay service-specific. The Agent Stub also uses this settings model directly: the public Agent Stub API base URL, server secret, -optional gRPC bind override, and optional Dify inner API bridge settings all -live here under the ``DIFY_AGENT_...`` environment-variable namespace. +and optional Dify inner API bridge settings live here under the +``DIFY_AGENT_...`` environment-variable namespace. """ import httpx -from typing import TYPE_CHECKING, ClassVar, Literal +from typing import ClassVar, Literal, cast -from pydantic import AnyHttpUrl, Field, TypeAdapter, field_validator, model_validator +from pydantic import AliasChoices, AnyHttpUrl, Field, TypeAdapter, field_validator, model_validator from pydantic_settings import BaseSettings, SettingsConfigDict -if TYPE_CHECKING: - from dify_agent.adapters.shell.protocols import ShellProviderProtocol - -from dify_agent.agent_stub.protocol.agent_stub import normalize_agent_stub_api_base_url, parse_agent_stub_endpoint +from dify_agent.agent_stub.protocol.agent_stub import normalize_agent_stub_api_base_url from dify_agent.agent_stub.server.agent_stub_config import DifyApiAgentStubConfigRequestHandler from dify_agent.agent_stub.server.agent_stub_drive import DifyApiAgentStubDriveRequestHandler from dify_agent.agent_stub.server.agent_stub_files import DifyApiAgentStubFileRequestHandler -from dify_agent.agent_stub.server.grpc_bind import normalize_agent_stub_grpc_bind_address from dify_agent.agent_stub.server.tokens.agent_stub import AgentStubTokenCodec, decode_server_secret_key +from dify_agent.runtime_backend import RuntimeBackendProfile +from dify_agent.runtime_backend.e2b import E2B_MAX_ACTIVE_TIMEOUT_SECONDS +from dify_agent.runtime_backend.profile import ( + DEFAULT_LOCAL_HOME_SNAPSHOT_ROOT, + DEFAULT_LOCAL_MATERIALIZED_HOME_ROOT, + DEFAULT_LOCAL_WORKSPACE_ROOT, + RuntimeBackendSettings, + create_runtime_backend_profile, +) DEFAULT_RUN_RETENTION_SECONDS = 3 * 24 * 60 * 60 class ServerSettings(BaseSettings): - """Environment-backed settings for Redis, scheduling, outbound HTTP, and shell access.""" + """Environment settings for scheduling, outbound HTTP, and runtime resources.""" redis_url: str = "redis://localhost:6379/0" redis_prefix: str = "dify-agent" @@ -41,17 +46,38 @@ class ServerSettings(BaseSettings): plugin_daemon_api_key: str = "" inner_api_url: str = "http://localhost:5001" inner_api_key: str | None = None - shell_provider: Literal["shellctl", "enterprise"] = "shellctl" - shellctl_entrypoint: str | None = None - shellctl_auth_token: str | None = None - shell_home_root: str = "/home" + runtime_backend: Literal["local", "enterprise", "e2b"] = "local" + local_sandbox_endpoint: str | None = Field( + default=None, + validation_alias=AliasChoices("DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT", "DIFY_AGENT_SHELLCTL_ENTRYPOINT"), + ) + local_sandbox_auth_token: str | None = Field( + default=None, + validation_alias=AliasChoices("DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN", "DIFY_AGENT_SHELLCTL_AUTH_TOKEN"), + ) + local_sandbox_materialized_home_root: str = DEFAULT_LOCAL_MATERIALIZED_HOME_ROOT + local_sandbox_workspace_root: str = DEFAULT_LOCAL_WORKSPACE_ROOT + local_sandbox_home_snapshot_root: str = DEFAULT_LOCAL_HOME_SNAPSHOT_ROOT enterprise_sandbox_gateway_endpoint: str | None = None enterprise_sandbox_gateway_auth_token: str | None = None enterprise_sandbox_gateway_timeout: float = Field(default=30.0, gt=0) enterprise_sandbox_proxy_timeout: float = Field(default=60.0, gt=0) + e2b_api_key: str | None = None + e2b_template: str = "difys-default-team/dify-agent-local-sandbox" + e2b_active_timeout_seconds: int = Field( + default=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ge=1, + le=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ) + e2b_shellctl_auth_token: str = "" + e2b_shellctl_port: int = Field(default=5004, ge=1, le=65535) agent_stub_api_base_url: str | None = Field(default=None, validation_alias="DIFY_AGENT_STUB_API_BASE_URL") - agent_stub_grpc_bind_address: str | None = Field(default=None, validation_alias="DIFY_AGENT_STUB_GRPC_BIND_ADDRESS") + sandbox_files_base_url: str | None = Field( + default=None, + validation_alias="DIFY_AGENT_SANDBOX_FILES_BASE_URL", + ) server_secret_key: str | None = None + api_token: str | None = None shell_redact_patterns: str = "" outbound_http_connect_timeout: float = Field(default=10.0, ge=0) outbound_http_read_timeout: float = Field(default=600.0, ge=0) @@ -82,16 +108,21 @@ class ServerSettings(BaseSettings): return normalize_agent_stub_api_base_url(validated) return normalize_agent_stub_api_base_url(stripped) - @field_validator("agent_stub_grpc_bind_address") + @field_validator("sandbox_files_base_url") @classmethod - def normalize_agent_stub_grpc_bind_address_value(cls, value: str | None) -> str | None: - """Normalize the optional explicit Agent Stub gRPC bind override.""" + def normalize_sandbox_files_base_url_value(cls, value: str | None) -> str | None: + """Normalize the Dify API base URL reachable from the Sandbox.""" + if value is None: return None stripped = value.strip() if not stripped: return None - return normalize_agent_stub_grpc_bind_address(stripped) + validated = str(TypeAdapter(AnyHttpUrl).validate_python(stripped)) + parsed = validated.rstrip("/") + if "?" in parsed or "#" in parsed: + raise ValueError("DIFY_AGENT_SANDBOX_FILES_BASE_URL must not include a query string or fragment") + return parsed @field_validator("server_secret_key") @classmethod @@ -134,56 +165,47 @@ class ServerSettings(BaseSettings): return [] import json as _json - parsed = _json.loads(stripped) + parsed = cast(object, _json.loads(stripped)) if not isinstance(parsed, list): raise ValueError("DIFY_AGENT_SHELL_REDACT_PATTERNS must be a JSON array of strings") - return [str(p) for p in parsed] - - @field_validator("shell_home_root") - @classmethod - def normalize_shell_home_root(cls, value: str) -> str: - """Normalize the root used for per-Agent shell HOME directories.""" - stripped = value.strip().rstrip("/") - if not stripped: - raise ValueError("DIFY_AGENT_SHELL_HOME_ROOT must not be empty") - if not stripped.startswith("/"): - raise ValueError("DIFY_AGENT_SHELL_HOME_ROOT must be an absolute path") - return stripped + return TypeAdapter(list[str]).validate_python(parsed) @model_validator(mode="after") def validate_agent_stub_requirements(self) -> "ServerSettings": """Require Agent Stub settings while allowing deployments without inner API calls.""" if self.agent_stub_api_base_url is not None and self.server_secret_key is None: raise ValueError("DIFY_AGENT_SERVER_SECRET_KEY is required when DIFY_AGENT_STUB_API_BASE_URL is set.") - if self.agent_stub_grpc_bind_address is not None: - if self.agent_stub_api_base_url is None: - raise ValueError( - "DIFY_AGENT_STUB_API_BASE_URL is required when DIFY_AGENT_STUB_GRPC_BIND_ADDRESS is set." - ) - if not parse_agent_stub_endpoint(self.agent_stub_api_base_url).is_grpc: - raise ValueError("DIFY_AGENT_STUB_GRPC_BIND_ADDRESS requires a grpc:// DIFY_AGENT_STUB_API_BASE_URL.") + if ( + self.agent_stub_api_base_url is not None + and self.inner_api_key is not None + and self.sandbox_files_base_url is None + ): + raise ValueError( + "DIFY_AGENT_SANDBOX_FILES_BASE_URL is required for Agent Stub file transfers and Config downloads." + ) return self - def build_shell_provider(self) -> "ShellProviderProtocol | None": - from dify_agent.adapters.shell.config import ShellAdapterSettings - from dify_agent.adapters.shell.factory import create_shell_provider - - match self.shell_provider: - case "shellctl": - if not self.shellctl_entrypoint: - return None - case "enterprise": - if not self.enterprise_sandbox_gateway_endpoint: - return None - return create_shell_provider( - ShellAdapterSettings( - shell_provider=self.shell_provider, - shellctl_entrypoint=self.shellctl_entrypoint, - shellctl_auth_token=self.shellctl_auth_token, + def build_runtime_backend_profile(self) -> RuntimeBackendProfile | None: + """Build the deployment-selected resource backend without adding service state.""" + if self.runtime_backend == "local" and not self.local_sandbox_endpoint: + return None + return create_runtime_backend_profile( + RuntimeBackendSettings( + runtime_backend=self.runtime_backend, + local_sandbox_endpoint=self.local_sandbox_endpoint, + local_sandbox_auth_token=self.local_sandbox_auth_token, + local_sandbox_materialized_home_root=self.local_sandbox_materialized_home_root, + local_sandbox_workspace_root=self.local_sandbox_workspace_root, + local_sandbox_home_snapshot_root=self.local_sandbox_home_snapshot_root, enterprise_sandbox_gateway_endpoint=self.enterprise_sandbox_gateway_endpoint, enterprise_sandbox_gateway_auth_token=self.enterprise_sandbox_gateway_auth_token, enterprise_sandbox_gateway_timeout=self.enterprise_sandbox_gateway_timeout, enterprise_sandbox_proxy_timeout=self.enterprise_sandbox_proxy_timeout, + e2b_api_key=self.e2b_api_key, + e2b_template=self.e2b_template, + e2b_active_timeout_seconds=self.e2b_active_timeout_seconds, + e2b_shellctl_auth_token=self.e2b_shellctl_auth_token, + e2b_shellctl_port=self.e2b_shellctl_port, ) ) @@ -194,21 +216,24 @@ class ServerSettings(BaseSettings): return AgentStubTokenCodec.from_server_secret(self.server_secret_key) def create_agent_stub_file_request_handler(self) -> DifyApiAgentStubFileRequestHandler | None: - """Return the Dify API file bridge when both Dify API settings are configured.""" - if self.inner_api_key is None: + """Return the file bridge when inner API and Sandbox data-plane settings are configured.""" + if self.inner_api_key is None or self.sandbox_files_base_url is None: return None return DifyApiAgentStubFileRequestHandler( inner_api_url=self.inner_api_url, inner_api_key=self.inner_api_key, + sandbox_files_base_url=self.sandbox_files_base_url, + timeout=self.create_outbound_http_timeout(), ) def create_agent_stub_config_request_handler(self) -> DifyApiAgentStubConfigRequestHandler | None: - """Return the Dify API config bridge when both Dify API settings are configured.""" - if self.inner_api_key is None: + """Return the Config bridge when inner API and Sandbox data-plane settings are configured.""" + if self.inner_api_key is None or self.sandbox_files_base_url is None: return None return DifyApiAgentStubConfigRequestHandler( inner_api_url=self.inner_api_url, inner_api_key=self.inner_api_key, + sandbox_files_base_url=self.sandbox_files_base_url, timeout=self.create_outbound_http_timeout(), ) diff --git a/dify-agent/src/dify_agent/storage/redis_run_store.py b/dify-agent/src/dify_agent/storage/redis_run_store.py index ced155ba3b2..6bba2368b16 100644 --- a/dify-agent/src/dify_agent/storage/redis_run_store.py +++ b/dify-agent/src/dify_agent/storage/redis_run_store.py @@ -9,22 +9,63 @@ create-run payloads are never persisted because layer config may include model credentials. """ -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Awaitable from typing import cast from redis.asyncio import Redis -from dify_agent.protocol.schemas import RUN_EVENT_ADAPTER, RunEvent, RunEventsResponse, RunStatus, utc_now -from dify_agent.runtime.event_sink import RunEventSink +from dify_agent.protocol.schemas import RUN_EVENT_ADAPTER, RunEvent, RunEventsResponse, RunStatus +from dify_agent.runtime.event_sink import ( + NonTerminalRunEvent, + RunEventSink, + RunFinalizationResult, + TerminalRunEvent, + terminal_event_status_fields, +) from dify_agent.server.schemas import RunRecord, new_run_id from dify_agent.server.settings import DEFAULT_RUN_RETENTION_SECONDS from dify_agent.storage.redis_keys import run_events_key, run_record_key +_TERMINAL_RUN_EVENT_TYPES = {"run_succeeded", "run_failed", "run_cancelled"} + class RunNotFoundError(LookupError): """Raised when a requested run record does not exist.""" +_FINALIZE_RUN_SCRIPT = """ +local record_json = redis.call("GET", KEYS[1]) +if not record_json then + return {-1, "", ""} +end + +local record = cjson.decode(record_json) +if record.status ~= "running" then + return {0, tostring(record.status), ""} +end + +record.status = ARGV[1] +record.updated_at = ARGV[2] +if ARGV[3] == "1" then + record.error = ARGV[4] +else + record.error = cjson.null +end +if ARGV[5] == "1" then + record.error_type = ARGV[6] +else + record.error_type = cjson.null +end + +local ttl = tonumber(ARGV[8]) +local updated_record_json = cjson.encode(record) +local event_id = redis.call("XADD", KEYS[2], "*", "payload", ARGV[7]) +redis.call("EXPIRE", KEYS[2], ttl) +redis.call("SET", KEYS[1], updated_record_json, "EX", ttl) +return {1, ARGV[1], event_id} +""" + + class RedisRunStore(RunEventSink): """Async Redis implementation for run records and event logs. @@ -73,18 +114,8 @@ class RedisRunStore(RunEventSink): value = value.decode() return RunRecord.model_validate_json(value) - async def update_status(self, run_id: str, status: RunStatus, error: str | None = None) -> None: - """Update the status fields of an existing run record.""" - record = await self.get_run(run_id) - updated = record.model_copy(update={"status": status, "updated_at": utc_now(), "error": error}) - await self.redis.set( - run_record_key(self.prefix, run_id), - updated.model_dump_json(), - ex=self.run_retention_seconds, - ) - - async def append_event(self, event: RunEvent) -> str: - """Append an event JSON payload to the run's Redis stream with TTLs.""" + async def append_event(self, event: NonTerminalRunEvent) -> str: + """Append a non-terminal event JSON payload with refreshed TTLs.""" events_key = run_events_key(self.prefix, event.run_id) payload = RUN_EVENT_ADAPTER.dump_json(event, exclude={"id"}).decode() async with self.redis.pipeline(transaction=True) as pipeline: @@ -98,6 +129,67 @@ class RedisRunStore(RunEventSink): event_id = results[0] return event_id.decode() if isinstance(event_id, bytes) else str(event_id) + async def finalize_run(self, event: TerminalRunEvent) -> RunFinalizationResult: + """Atomically append the first terminal event and update its run record.""" + status, error, error_type = terminal_event_status_fields(event) + payload = RUN_EVENT_ADAPTER.dump_json(event, exclude={"id"}).decode() + evaluation = cast( + Awaitable[object], + self.redis.eval( + _FINALIZE_RUN_SCRIPT, + 2, + run_record_key(self.prefix, event.run_id), + run_events_key(self.prefix, event.run_id), + status, + event.created_at.isoformat(), + "1" if error is not None else "0", + error or "", + "1" if error_type is not None else "0", + error_type.value if error_type is not None else "", + payload, + str(self.run_retention_seconds), + ), + ) + raw_result = await evaluation + result = cast(list[object], raw_result) + applied = int(cast(int | bytes | str, result[0])) + if applied == -1: + raise RunNotFoundError(event.run_id) + + persisted_status = cast(RunStatus, _decode_redis_text(result[1])) + event_id = _decode_redis_text(result[2]) or None + return RunFinalizationResult( + applied=applied == 1, + status=persisted_status, + event_id=event_id, + ) + + async def wait_for_cancellation(self, run_id: str) -> bool: + """Wait until cancellation or another terminal state wins for one run. + + The stream cursor is captured before reading the record so a terminal + transition cannot fall between the initial status check and blocking + stream read. + """ + events_key = run_events_key(self.prefix, run_id) + latest_events = await self.redis.xrevrange(events_key, count=1) + cursor = _decode_redis_text(latest_events[0][0]) if latest_events else "0-0" + record = await self.get_run(run_id) + if record.status != "running": + return record.status == "cancelled" + + while True: + response = await self.redis.xread({events_key: cursor}, block=0, count=100) + for _stream_name, entries in response: + for raw_id, fields in entries: + event = self._decode_event(run_id, raw_id, fields) + if event.id is not None: + cursor = event.id + if event.type == "run_cancelled": + return True + if event.type in {"run_succeeded", "run_failed"}: + return False + async def get_events(self, run_id: str, *, after: str = "0-0", limit: int = 100) -> RunEventsResponse: """Read a bounded page of events after ``after`` cursor.""" await self.get_run(run_id) @@ -107,7 +199,7 @@ class RedisRunStore(RunEventSink): return RunEventsResponse(run_id=run_id, events=events, next_cursor=next_cursor) async def iter_events(self, run_id: str, *, after: str = "0-0") -> AsyncIterator[RunEvent]: - """Yield replayed and future events for SSE clients.""" + """Yield replayed and future events through the first terminal event.""" await self.get_run(run_id) cursor = after while True: @@ -116,6 +208,8 @@ class RedisRunStore(RunEventSink): if event.id is not None: cursor = event.id yield event + if event.type in _TERMINAL_RUN_EVENT_TYPES: + return if not page.events: break while True: @@ -128,6 +222,8 @@ class RedisRunStore(RunEventSink): if event.id is not None: cursor = event.id yield event + if event.type in _TERMINAL_RUN_EVENT_TYPES: + return @staticmethod def _decode_event(run_id: str, raw_id: object, fields: dict[object, object]) -> RunEvent: @@ -140,4 +236,8 @@ class RedisRunStore(RunEventSink): return event.model_copy(update={"id": event_id, "run_id": run_id}) +def _decode_redis_text(value: object) -> str: + return value.decode() if isinstance(value, bytes) else str(value) + + __all__ = ["DEFAULT_RUN_RETENTION_SECONDS", "RedisRunStore", "RunNotFoundError"] diff --git a/dify-agent/src/shellctl/__init__.py b/dify-agent/src/shellctl/__init__.py index 3f2237dc690..cb294b365cf 100644 --- a/dify-agent/src/shellctl/__init__.py +++ b/dify-agent/src/shellctl/__init__.py @@ -115,7 +115,7 @@ def __getattr__(name: str) -> Any: if name not in _EXPORTS: raise AttributeError(f"module {__name__!r} has no attribute {name!r}") module = import_module(_EXPORTS[name]) - value = getattr(module, name) # noqa: no-new-getattr lazy export proxy + value = getattr(module, name) # guard-ignore: no-new-getattr -- lazy export proxy globals()[name] = value return value diff --git a/dify-agent/src/shellctl/client/sdk.py b/dify-agent/src/shellctl/client/sdk.py index d007fff551d..cc8ee05ca9f 100644 --- a/dify-agent/src/shellctl/client/sdk.py +++ b/dify-agent/src/shellctl/client/sdk.py @@ -23,6 +23,7 @@ from shellctl.shared.constants import ( DEFAULT_OUTPUT_LIMIT_BYTES, DEFAULT_TERMINATE_GRACE_SECONDS, DEFAULT_TIMEOUT_SECONDS, + SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS, ) from shellctl.shared.schemas import ( DeleteJobResponse, @@ -71,7 +72,7 @@ class ShellctlClient: token: str | None = None, client: httpx.AsyncClient | None = None, transport: httpx.AsyncBaseTransport | None = None, - request_timeout_grace_seconds: float = 10.0, + request_timeout_grace_seconds: float = SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS, ) -> None: self.base_url = base_url.rstrip("/") self.output_limit = output_limit diff --git a/dify-agent/src/shellctl/shared/__init__.py b/dify-agent/src/shellctl/shared/__init__.py index 5f235b19d9f..e53a46d5e27 100644 --- a/dify-agent/src/shellctl/shared/__init__.py +++ b/dify-agent/src/shellctl/shared/__init__.py @@ -31,6 +31,9 @@ if TYPE_CHECKING: MAX_LIST_LIMIT, MAX_OUTPUT_LIMIT_BYTES, MAX_WAIT_TIMEOUT_SECONDS, + SHELL_TOOL_HARD_TIMEOUT_SECONDS, + SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS, + SHELL_TOOL_TIMEOUT_WITH_HTTP_GRACE_SECONDS, SESSION_NAME_PREFIX, ) from shellctl.shared.output import ( @@ -87,6 +90,9 @@ __all__ = [ "MAX_LIST_LIMIT", "MAX_OUTPUT_LIMIT_BYTES", "MAX_WAIT_TIMEOUT_SECONDS", + "SHELL_TOOL_HARD_TIMEOUT_SECONDS", + "SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS", + "SHELL_TOOL_TIMEOUT_WITH_HTTP_GRACE_SECONDS", "SESSION_NAME_PREFIX", "TERMINAL_JOB_STATUSES", "DeleteJobResponse", @@ -137,6 +143,9 @@ _EXPORTS = { "MAX_LIST_LIMIT": "shellctl.shared.constants", "MAX_OUTPUT_LIMIT_BYTES": "shellctl.shared.constants", "MAX_WAIT_TIMEOUT_SECONDS": "shellctl.shared.constants", + "SHELL_TOOL_HARD_TIMEOUT_SECONDS": "shellctl.shared.constants", + "SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS": "shellctl.shared.constants", + "SHELL_TOOL_TIMEOUT_WITH_HTTP_GRACE_SECONDS": "shellctl.shared.constants", "SESSION_NAME_PREFIX": "shellctl.shared.constants", "OutputWindow": "shellctl.shared.output", "read_output_window": "shellctl.shared.output", @@ -173,7 +182,7 @@ def __getattr__(name: str) -> Any: if name not in _EXPORTS: raise AttributeError(f"module {__name__!r} has no attribute {name!r}") module = import_module(_EXPORTS[name]) - value = getattr(module, name) # noqa: no-new-getattr lazy export proxy + value = getattr(module, name) # guard-ignore: no-new-getattr -- lazy export proxy globals()[name] = value return value diff --git a/dify-agent/src/shellctl/shared/constants.py b/dify-agent/src/shellctl/shared/constants.py index 83eb78160f5..36cf2112e14 100644 --- a/dify-agent/src/shellctl/shared/constants.py +++ b/dify-agent/src/shellctl/shared/constants.py @@ -12,7 +12,13 @@ DEFAULT_BASE_URL = "http://127.0.0.1:8765" DEFAULT_OUTPUT_LIMIT_BYTES = 1024 * 8 MAX_OUTPUT_LIMIT_BYTES = 1024 * 1024 DEFAULT_TIMEOUT_SECONDS = 30.0 -MAX_WAIT_TIMEOUT_SECONDS = 5.0 * 60.0 +# Single source of truth for the hard business timeout exposed by Shell tools. +SHELL_TOOL_HARD_TIMEOUT_SECONDS = 5.0 * 60.0 +# Transport grace lets the HTTP response arrive after the Shell wait budget expires. +SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS = 10.0 +SHELL_TOOL_TIMEOUT_WITH_HTTP_GRACE_SECONDS = SHELL_TOOL_HARD_TIMEOUT_SECONDS + SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS +# Backward-compatible shellctl name; the Shell tool constant above owns the value. +MAX_WAIT_TIMEOUT_SECONDS = SHELL_TOOL_HARD_TIMEOUT_SECONDS DEFAULT_IDLE_FLUSH_SECONDS = 0.5 DEFAULT_TERMINATE_GRACE_SECONDS = 5.0 DEFAULT_TERMINAL_COLS = 120 @@ -45,5 +51,8 @@ __all__ = [ "MAX_LIST_LIMIT", "MAX_OUTPUT_LIMIT_BYTES", "MAX_WAIT_TIMEOUT_SECONDS", + "SHELL_TOOL_HARD_TIMEOUT_SECONDS", + "SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS", + "SHELL_TOOL_TIMEOUT_WITH_HTTP_GRACE_SECONDS", "SESSION_NAME_PREFIX", ] diff --git a/dify-agent/src/shellctl/shared/schemas.py b/dify-agent/src/shellctl/shared/schemas.py index 0120b8b8a5f..0bbaf323a18 100644 --- a/dify-agent/src/shellctl/shared/schemas.py +++ b/dify-agent/src/shellctl/shared/schemas.py @@ -15,7 +15,7 @@ from shellctl.shared.constants import ( DEFAULT_TERMINATE_GRACE_SECONDS, DEFAULT_TIMEOUT_SECONDS, MAX_OUTPUT_LIMIT_BYTES, - MAX_WAIT_TIMEOUT_SECONDS, + SHELL_TOOL_HARD_TIMEOUT_SECONDS, ) @@ -139,7 +139,7 @@ class RunJobRequest(ShellctlModel): cwd: str | None = None env: dict[str, str] | None = None terminal: TerminalSize | None = None - timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, gt=0, le=MAX_WAIT_TIMEOUT_SECONDS) + timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, gt=0, le=SHELL_TOOL_HARD_TIMEOUT_SECONDS) output_limit: int = Field(default=DEFAULT_OUTPUT_LIMIT_BYTES, ge=1, le=MAX_OUTPUT_LIMIT_BYTES) idle_flush_seconds: float = Field(default=DEFAULT_IDLE_FLUSH_SECONDS, ge=0, le=30) @@ -171,7 +171,7 @@ class RunJobRequest(ShellctlModel): class WaitJobRequest(ShellctlModel): """HTTP request body for `POST /v1/jobs/{job_id}/wait`.""" - timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, ge=0, le=MAX_WAIT_TIMEOUT_SECONDS) + timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, ge=0, le=SHELL_TOOL_HARD_TIMEOUT_SECONDS) offset: int = Field(ge=0) output_limit: int = Field(default=DEFAULT_OUTPUT_LIMIT_BYTES, ge=1, le=MAX_OUTPUT_LIMIT_BYTES) idle_flush_seconds: float = Field(default=DEFAULT_IDLE_FLUSH_SECONDS, ge=0, le=30) @@ -181,7 +181,7 @@ class InputJobRequest(ShellctlModel): """HTTP request body for `POST /v1/jobs/{job_id}/input`.""" text: str - timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, gt=0, le=MAX_WAIT_TIMEOUT_SECONDS) + timeout: float = Field(default=DEFAULT_TIMEOUT_SECONDS, gt=0, le=SHELL_TOOL_HARD_TIMEOUT_SECONDS) offset: int = Field(ge=0) output_limit: int = Field(default=DEFAULT_OUTPUT_LIMIT_BYTES, ge=1, le=MAX_OUTPUT_LIMIT_BYTES) idle_flush_seconds: float = Field(default=DEFAULT_IDLE_FLUSH_SECONDS, ge=0, le=30) @@ -190,7 +190,7 @@ class InputJobRequest(ShellctlModel): class TerminateJobRequest(ShellctlModel): """HTTP request body for `POST /v1/jobs/{job_id}/terminate`.""" - grace_seconds: float = Field(default=DEFAULT_TERMINATE_GRACE_SECONDS, ge=0, le=300) + grace_seconds: float = Field(default=DEFAULT_TERMINATE_GRACE_SECONDS, ge=0, le=SHELL_TOOL_HARD_TIMEOUT_SECONDS) __all__ = [ diff --git a/dify-agent/tests/integration/dify_agent/runtime_backend/run_local_integration.sh b/dify-agent/tests/integration/dify_agent/runtime_backend/run_local_integration.sh new file mode 100755 index 00000000000..1247636b058 --- /dev/null +++ b/dify-agent/tests/integration/dify_agent/runtime_backend/run_local_integration.sh @@ -0,0 +1,40 @@ +#!/usr/bin/env sh +set -eu + +script_dir=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd) +project_dir=$(CDPATH= cd -- "$script_dir/../../../.." && pwd) +container_name="dify-agent-runtime-backend-integration-$$" +image="${DIFY_AGENT_TEST_LOCAL_SANDBOX_IMAGE:-langgenius/dify-agent-local-sandbox:1.16.0}" +token="${DIFY_AGENT_TEST_LOCAL_SHELLCTL_AUTH_TOKEN:-runtime-backend-integration}" + +cleanup() { + docker rm -f "$container_name" >/dev/null 2>&1 || true +} +trap cleanup EXIT INT TERM + +docker run --detach --rm \ + --name "$container_name" \ + --env SHELLCTL_ENABLE_PATH_ISOLATION=true \ + --env SHELLCTL_AUTH_TOKEN="$token" \ + --publish 127.0.0.1::5004 \ + "$image" >/dev/null + +published_address=$(docker port "$container_name" 5004/tcp | head -n 1) +endpoint="http://$published_address" +attempt=0 +until curl --fail --silent "$endpoint/healthz" >/dev/null; do + attempt=$((attempt + 1)) + if [ "$attempt" -ge 100 ]; then + docker logs "$container_name" + exit 1 + fi + sleep 0.1 +done + +cd "$project_dir" +NO_PROXY=127.0.0.1,localhost \ + DIFY_AGENT_TEST_LOCAL_SHELLCTL_ENDPOINT="$endpoint" \ + DIFY_AGENT_TEST_LOCAL_SHELLCTL_AUTH_TOKEN="$token" \ + pdm run pytest --import-mode=importlib \ + tests/integration/dify_agent/runtime_backend/test_runtime_backend_lifecycle.py \ + -k local -q -rs "$@" diff --git a/dify-agent/tests/integration/dify_agent/runtime_backend/test_working_environment.py b/dify-agent/tests/integration/dify_agent/runtime_backend/test_working_environment.py new file mode 100644 index 00000000000..93a67c2e816 --- /dev/null +++ b/dify-agent/tests/integration/dify_agent/runtime_backend/test_working_environment.py @@ -0,0 +1,221 @@ +"""Opt-in integration contracts for final Local and E2B working environments.""" + +from __future__ import annotations + +import os +import shlex +import sys +import uuid + +import pytest + +from dify_agent.runtime_backend import ( + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotCreateSpec, +) +from dify_agent.runtime_backend.e2b import ( + E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + E2BExecutionBindingBackend, + E2BHomeSnapshotBackend, + E2BSDKControlPlane, +) +from dify_agent.runtime_backend.local import LocalExecutionBindingBackend +from dify_agent.runtime.command_runner import execute_complete_with_commands + +pytestmark = pytest.mark.integration + + +async def _run(lease, script: str, *, cwd: str) -> str: + result = await execute_complete_with_commands( + lease.commands, + script, + cwd=cwd, + env={"HOME": lease.layout.home_dir}, + timeout=30.0, + max_output_bytes=4096, + ) + assert result.exit_code == 0 + assert result.output_complete + return result.output + + +def _required_env(name: str, purpose: str) -> str: + value = os.environ.get(name, "").strip() + if not value: + pytest.skip(f"set {name} to run the {purpose} integration contract") + return value + + +@pytest.mark.anyio +async def test_local_two_agents_share_workspace_but_not_home() -> None: + endpoint = _required_env("DIFY_AGENT_TEST_LOCAL_SHELLCTL_ENDPOINT", "real Local shellctl") + token = os.environ.get("DIFY_AGENT_TEST_LOCAL_SHELLCTL_AUTH_TOKEN", "") + marker = uuid.uuid4().hex + bindings = LocalExecutionBindingBackend(endpoint=endpoint, auth_token=token) + allocations = [] + active_leases = [] + try: + first = await bindings.create_binding( + ExecutionBindingCreateSpec( + tenant_id="integration-tenant", + agent_id="agent-a", + binding_id=f"binding-a-{marker}", + workspace_id=f"workspace-{marker}", + existing_workspace_ref=None, + home_snapshot_ref=None, + ) + ) + allocations.append(first) + first_lease = await bindings.acquire(first.binding_ref) + active_leases.append(first_lease) + await _run(first_lease, "printf shared > shared.txt", cwd=first_lease.layout.workspace_dir) + await bindings.release(first_lease) + active_leases.remove(first_lease) + + second = await bindings.create_binding( + ExecutionBindingCreateSpec( + tenant_id="integration-tenant", + agent_id="agent-b", + binding_id=f"binding-b-{marker}", + workspace_id=f"workspace-{marker}", + existing_workspace_ref=first.workspace_ref, + home_snapshot_ref=None, + ) + ) + allocations.append(second) + second_lease = await bindings.acquire(second.binding_ref) + active_leases.append(second_lease) + shared = await _run(second_lease, "cat shared.txt", cwd=second_lease.layout.workspace_dir) + assert shared == "shared" + assert second_lease.layout.home_dir != first_lease.layout.home_dir + assert second_lease.layout.workspace_dir == first_lease.layout.workspace_dir + await bindings.release(second_lease) + active_leases.remove(second_lease) + finally: + primary_error = sys.exc_info()[0] is not None + cleanup_errors: list[BaseException] = [] + for lease in active_leases: + try: + await bindings.release(lease) + except BaseException as exc: + cleanup_errors.append(exc) + for index, allocation in enumerate(allocations): + try: + await bindings.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref=allocation.binding_ref, + workspace_ref=allocation.workspace_ref if index == len(allocations) - 1 else None, + destroy_workspace=index == len(allocations) - 1, + ) + ) + except BaseException as exc: + cleanup_errors.append(exc) + if cleanup_errors and not primary_error: + raise cleanup_errors[0] + + +@pytest.mark.anyio +async def test_e2b_binding_checkpoint_and_collection() -> None: + api_key = _required_env("DIFY_AGENT_TEST_E2B_API_KEY", "real E2B") + template = os.environ.get( + "DIFY_AGENT_TEST_E2B_TEMPLATE", + "difys-default-team/dify-agent-local-sandbox", + ) + marker = uuid.uuid4().hex + control = E2BSDKControlPlane(api_key=api_key) + snapshots = E2BHomeSnapshotBackend(control_plane=control) + bindings = E2BExecutionBindingBackend( + control_plane=control, + template=template, + active_timeout_seconds=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ) + checkpoint_ref: str | None = None + allocation = None + checkpoint_allocation = None + lease = None + checkpoint_lease = None + try: + allocation = await bindings.create_binding( + ExecutionBindingCreateSpec( + tenant_id="integration-tenant", + agent_id="integration-agent", + binding_id=marker, + workspace_id=marker, + existing_workspace_ref=None, + home_snapshot_ref=None, + ) + ) + lease = await bindings.acquire(allocation.binding_ref) + await _run(lease, "printf e2b > probe.txt", cwd=lease.layout.workspace_dir) + await _run(lease, "printf checkpoint-home > .checkpoint-probe", cwd=lease.layout.home_dir) + assert await _run(lease, "cat probe.txt", cwd=lease.layout.workspace_dir) == "e2b" + checkpoint_ref = await snapshots.create_from_runtime( + spec=HomeSnapshotCreateSpec( + tenant_id="integration-tenant", + agent_id="integration-agent", + home_snapshot_id=f"checkpoint-{marker}", + ), + source=lease, + ) + await bindings.release(lease) + lease = None + + checkpoint_allocation = await bindings.create_binding( + ExecutionBindingCreateSpec( + tenant_id="integration-tenant", + agent_id="integration-agent", + binding_id=f"checkpoint-{marker}", + workspace_id=f"checkpoint-{marker}", + existing_workspace_ref=None, + home_snapshot_ref=checkpoint_ref, + ) + ) + checkpoint_lease = await bindings.acquire(checkpoint_allocation.binding_ref) + checkpoint_path = shlex.quote(f"{checkpoint_lease.layout.home_dir}/.checkpoint-probe") + restored = await _run(checkpoint_lease, f"cat {checkpoint_path}", cwd=checkpoint_lease.layout.workspace_dir) + assert restored == "checkpoint-home" + await bindings.release(checkpoint_lease) + checkpoint_lease = None + finally: + primary_error = sys.exc_info()[0] is not None + cleanup_errors: list[BaseException] = [] + if lease is not None: + try: + await bindings.release(lease) + except BaseException as exc: + cleanup_errors.append(exc) + if checkpoint_lease is not None: + try: + await bindings.release(checkpoint_lease) + except BaseException as exc: + cleanup_errors.append(exc) + if checkpoint_allocation is not None: + try: + await bindings.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref=checkpoint_allocation.binding_ref, + workspace_ref=checkpoint_allocation.workspace_ref, + destroy_workspace=True, + ) + ) + except BaseException as exc: + cleanup_errors.append(exc) + if allocation is not None: + try: + await bindings.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref=allocation.binding_ref, + workspace_ref=allocation.workspace_ref, + destroy_workspace=True, + ) + ) + except BaseException as exc: + cleanup_errors.append(exc) + if checkpoint_ref is not None: + try: + await snapshots.delete(checkpoint_ref) + except BaseException as exc: + cleanup_errors.append(exc) + if cleanup_errors and not primary_error: + raise cleanup_errors[0] diff --git a/dify-agent/tests/integration/dify_agent/storage/test_terminal_finalization.py b/dify-agent/tests/integration/dify_agent/storage/test_terminal_finalization.py new file mode 100644 index 00000000000..396e7736f68 --- /dev/null +++ b/dify-agent/tests/integration/dify_agent/storage/test_terminal_finalization.py @@ -0,0 +1,235 @@ +"""Real-Redis contracts for terminal finalization and cancellation observation.""" + +import asyncio +from collections.abc import Iterator +import shutil +import socket +import subprocess +import time +from uuid import uuid4 + +import httpx +import pytest +from redis.asyncio import Redis + +from agenton.compositor import CompositorSessionSnapshot +from dify_agent.protocol.schemas import ( + CancelRunRequest, + CreateRunRequest, + RunCancelledEvent, + RunCancelledEventData, + RunComposition, + RunFailedEvent, + RunFailedEventData, + RunFailureType, + RunSucceededEvent, + RunSucceededEventData, +) +from dify_agent.runtime.event_sink import TerminalRunEvent, terminal_event_status_fields +from dify_agent.runtime.run_scheduler import RunScheduler +from dify_agent.storage.redis_keys import run_events_key, run_record_key +from dify_agent.storage.redis_run_store import RedisRunStore + + +pytestmark = pytest.mark.integration + + +@pytest.fixture +def redis_url() -> Iterator[str]: + """Start an isolated Redis when the binary is available locally.""" + redis_server = shutil.which("redis-server") + if redis_server is None: + pytest.skip("redis-server is required for run terminal integration tests") + + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as probe: + probe.bind(("127.0.0.1", 0)) + port = probe.getsockname()[1] + + process = subprocess.Popen( # noqa: S603 + [ + redis_server, + "--bind", + "127.0.0.1", + "--port", + str(port), + "--save", + "", + "--appendonly", + "no", + ], + stdout=subprocess.DEVNULL, + stderr=subprocess.STDOUT, + ) + try: + deadline = time.monotonic() + 5 + while time.monotonic() < deadline: + if process.poll() is not None: + pytest.fail(f"redis-server exited during startup with code {process.returncode}") + try: + with socket.create_connection(("127.0.0.1", port), timeout=0.1): + break + except OSError: + time.sleep(0.05) + else: + pytest.fail("redis-server did not accept connections within 5 seconds") + yield f"redis://127.0.0.1:{port}/0" + finally: + process.terminate() + try: + process.wait(timeout=5) + except subprocess.TimeoutExpired: + process.kill() + process.wait(timeout=5) + + +def test_two_redis_clients_commit_exactly_one_matching_terminal(redis_url: str) -> None: + async def scenario() -> None: + first_client = Redis.from_url(redis_url) + second_client = Redis.from_url(redis_url) + prefix = f"terminal-finalization-{uuid4().hex}" + retention_seconds = 60 + first_store = RedisRunStore(first_client, prefix=prefix, run_retention_seconds=retention_seconds) + second_store = RedisRunStore(second_client, prefix=prefix, run_retention_seconds=retention_seconds) + try: + record = await first_store.create_run() + terminal_events: tuple[TerminalRunEvent, TerminalRunEvent] = ( + RunSucceededEvent( + run_id=record.run_id, + data=RunSucceededEventData( + output="done", + session_snapshot=CompositorSessionSnapshot(layers=[]), + ), + ), + RunCancelledEvent( + run_id=record.run_id, + data=RunCancelledEventData( + reason="concurrent_cancel", + message="cancel accepted", + ), + ), + ) + + results = await asyncio.gather( + first_store.finalize_run(terminal_events[0]), + second_store.finalize_run(terminal_events[1]), + ) + + assert sum(result.applied for result in results) == 1 + winner_index = next(index for index, result in enumerate(results) if result.applied) + winner_event = terminal_events[winner_index] + winner_result = results[winner_index] + expected_status, expected_error, expected_error_type = terminal_event_status_fields(winner_event) + + persisted = await first_store.get_run(record.run_id) + page = await second_store.get_events(record.run_id) + assert persisted.status == expected_status + assert persisted.error == expected_error + assert persisted.error_type == expected_error_type + assert persisted.updated_at == winner_event.created_at + assert len(page.events) == 1 + assert page.events[0].type == winner_event.type + assert page.events[0].created_at == winner_event.created_at + assert page.events[0].id == winner_result.event_id + + record_ttl = await first_client.ttl(run_record_key(prefix, record.run_id)) + events_ttl = await second_client.ttl(run_events_key(prefix, record.run_id)) + assert 0 < record_ttl <= retention_seconds + assert 0 < events_ttl <= retention_seconds + finally: + await first_client.aclose() + await second_client.aclose() + + asyncio.run(scenario()) + + +def test_classified_failure_persists_matching_record_and_event_error_type(redis_url: str) -> None: + async def scenario() -> None: + client = Redis.from_url(redis_url) + store = RedisRunStore(client, prefix=f"classified-failure-{uuid4().hex}", run_retention_seconds=60) + try: + record = await store.create_run() + event = RunFailedEvent( + run_id=record.run_id, + data=RunFailedEventData( + error="run limit reached", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + ), + ) + + result = await store.finalize_run(event) + persisted = await store.get_run(record.run_id) + page = await store.get_events(record.run_id) + + assert result.applied is True + assert persisted.error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + assert len(page.events) == 1 + persisted_event = page.events[0] + assert isinstance(persisted_event, RunFailedEvent) + assert persisted_event.data.error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + finally: + await client.aclose() + + asyncio.run(scenario()) + + +def test_non_owner_scheduler_cancellation_stops_owner_runner(redis_url: str) -> None: + class BlockingRunner: + def __init__(self, *, started: asyncio.Event, stopped: asyncio.Event) -> None: + self.started = started + self.stopped = stopped + + async def run(self) -> None: + self.started.set() + try: + await asyncio.Event().wait() + finally: + self.stopped.set() + + async def scenario() -> None: + owner_client = Redis.from_url(redis_url) + remote_client = Redis.from_url(redis_url) + prefix = f"route-independent-cancellation-{uuid4().hex}" + owner_store = RedisRunStore(owner_client, prefix=prefix, run_retention_seconds=60) + remote_store = RedisRunStore(remote_client, prefix=prefix, run_retention_seconds=60) + runner_started = asyncio.Event() + runner_stopped = asyncio.Event() + async with httpx.AsyncClient() as http_client: + owner_scheduler = RunScheduler( + store=owner_store, + plugin_daemon_http_client=http_client, + dify_api_http_client=http_client, + runner_factory=lambda _record, _request: BlockingRunner( + started=runner_started, + stopped=runner_stopped, + ), + ) + remote_scheduler = RunScheduler( + store=remote_store, + plugin_daemon_http_client=http_client, + dify_api_http_client=http_client, + ) + try: + record = await owner_scheduler.create_run(CreateRunRequest(composition=RunComposition(layers=[]))) + owner_task = owner_scheduler.active_tasks[record.run_id] + await asyncio.wait_for(runner_started.wait(), timeout=1) + + response = await remote_scheduler.cancel_run( + record.run_id, + CancelRunRequest(reason="remote_cancel"), + ) + + assert response.status == "cancelled" + assert remote_scheduler.active_tasks == {} + await asyncio.wait_for(runner_stopped.wait(), timeout=1) + await asyncio.wait_for(owner_task, timeout=1) + persisted = await owner_store.get_run(record.run_id) + events = await remote_store.get_events(record.run_id) + assert persisted.status == "cancelled" + assert [event.type for event in events.events] == ["run_cancelled"] + finally: + await owner_scheduler.shutdown() + await remote_scheduler.shutdown() + await owner_client.aclose() + await remote_client.aclose() + + asyncio.run(scenario()) diff --git a/dify-agent/tests/local/dify_agent/adapters/shell/test_config.py b/dify-agent/tests/local/dify_agent/adapters/shell/test_config.py index 19110694403..9bbeb84f3c9 100644 --- a/dify-agent/tests/local/dify_agent/adapters/shell/test_config.py +++ b/dify-agent/tests/local/dify_agent/adapters/shell/test_config.py @@ -1,78 +1,78 @@ -"""Tests for ShellAdapterSettings provider-based field validation.""" +"""Tests for deployment-selected runtime backend settings.""" from __future__ import annotations import pytest from pydantic import ValidationError -from dify_agent.adapters.shell.config import ShellAdapterSettings +from dify_agent.runtime_backend.profile import RuntimeBackendSettings -class TestShellctlValidation: +class TestLocalValidation: def test_valid_url_passes(self) -> None: - settings = ShellAdapterSettings( - shell_provider="shellctl", - shellctl_entrypoint="http://shellctl.example.com", + settings = RuntimeBackendSettings( + runtime_backend="local", + local_sandbox_endpoint="http://shellctl.example.com", ) - assert settings.shellctl_entrypoint == "http://shellctl.example.com" + assert settings.local_sandbox_endpoint == "http://shellctl.example.com" def test_https_url_passes(self) -> None: - settings = ShellAdapterSettings( - shell_provider="shellctl", - shellctl_entrypoint="https://shellctl.internal:8443/v1", + settings = RuntimeBackendSettings( + runtime_backend="local", + local_sandbox_endpoint="https://shellctl.internal:8443/v1", ) - assert settings.shellctl_entrypoint == "https://shellctl.internal:8443/v1" + assert settings.local_sandbox_endpoint == "https://shellctl.internal:8443/v1" def test_missing_entrypoint_raises(self) -> None: - with pytest.raises(ValidationError, match="shellctl_entrypoint is required"): - ShellAdapterSettings(shell_provider="shellctl", shellctl_entrypoint=None) + with pytest.raises(ValidationError, match="local_sandbox_endpoint is required"): + RuntimeBackendSettings(runtime_backend="local", local_sandbox_endpoint=None) def test_blank_entrypoint_raises(self) -> None: - with pytest.raises(ValidationError, match="shellctl_entrypoint is required"): - ShellAdapterSettings(shell_provider="shellctl", shellctl_entrypoint=" ") + with pytest.raises(ValidationError, match="local_sandbox_endpoint is required"): + RuntimeBackendSettings(runtime_backend="local", local_sandbox_endpoint=" ") def test_invalid_url_raises(self) -> None: with pytest.raises(ValidationError, match="must be a valid http\\(s\\) URL"): - ShellAdapterSettings(shell_provider="shellctl", shellctl_entrypoint="not-a-url") + RuntimeBackendSettings(runtime_backend="local", local_sandbox_endpoint="not-a-url") def test_ftp_scheme_rejected(self) -> None: with pytest.raises(ValidationError, match="must be a valid http\\(s\\) URL"): - ShellAdapterSettings(shell_provider="shellctl", shellctl_entrypoint="ftp://host/path") + RuntimeBackendSettings(runtime_backend="local", local_sandbox_endpoint="ftp://host/path") class TestEnterpriseValidation: def test_valid_url_passes(self) -> None: - settings = ShellAdapterSettings( - shell_provider="enterprise", + settings = RuntimeBackendSettings( + runtime_backend="enterprise", enterprise_sandbox_gateway_endpoint="http://gateway.internal:9000", ) assert settings.enterprise_sandbox_gateway_endpoint == "http://gateway.internal:9000" def test_https_url_passes(self) -> None: - settings = ShellAdapterSettings( - shell_provider="enterprise", + settings = RuntimeBackendSettings( + runtime_backend="enterprise", enterprise_sandbox_gateway_endpoint="https://gateway.prod.example.com/api", ) assert settings.enterprise_sandbox_gateway_endpoint == "https://gateway.prod.example.com/api" def test_missing_endpoint_raises(self) -> None: with pytest.raises(ValidationError, match="enterprise_sandbox_gateway_endpoint is required"): - ShellAdapterSettings(shell_provider="enterprise", enterprise_sandbox_gateway_endpoint=None) + RuntimeBackendSettings(runtime_backend="enterprise", enterprise_sandbox_gateway_endpoint=None) def test_blank_endpoint_raises(self) -> None: with pytest.raises(ValidationError, match="enterprise_sandbox_gateway_endpoint is required"): - ShellAdapterSettings(shell_provider="enterprise", enterprise_sandbox_gateway_endpoint="") + RuntimeBackendSettings(runtime_backend="enterprise", enterprise_sandbox_gateway_endpoint="") def test_invalid_url_raises(self) -> None: with pytest.raises(ValidationError, match="must be a valid http\\(s\\) URL"): - ShellAdapterSettings( - shell_provider="enterprise", + RuntimeBackendSettings( + runtime_backend="enterprise", enterprise_sandbox_gateway_endpoint="just-a-hostname", ) def test_ftp_scheme_rejected(self) -> None: with pytest.raises(ValidationError, match="must be a valid http\\(s\\) URL"): - ShellAdapterSettings( - shell_provider="enterprise", + RuntimeBackendSettings( + runtime_backend="enterprise", enterprise_sandbox_gateway_endpoint="ftp://gateway/path", ) diff --git a/dify-agent/tests/local/dify_agent/adapters/shell/test_enterprise.py b/dify-agent/tests/local/dify_agent/adapters/shell/test_enterprise.py deleted file mode 100644 index 86ac7e79172..00000000000 --- a/dify-agent/tests/local/dify_agent/adapters/shell/test_enterprise.py +++ /dev/null @@ -1,42 +0,0 @@ -from unittest.mock import patch - -import httpx2 -import pytest - -from dify_agent.adapters.shell.enterprise import EnterpriseShellProvider - - -@pytest.mark.anyio -async def test_attach_uses_shellctl_compatible_http_client(monkeypatch: pytest.MonkeyPatch) -> None: - def handler(request: httpx2.Request) -> httpx2.Response: - assert request.url.path == "/proxy/v1/jobs/run" - assert request.headers["X-Sandbox-Id"] == "sandbox-1" - return httpx2.Response( - 200, - json={ - "job_id": "job-1", - "done": True, - "status": "exited", - "exit_code": 0, - "output_path": "/tmp/output.log", - "output": "", - "offset": 0, - "truncated": False, - }, - ) - - transport = httpx2.MockTransport(handler) - monkeypatch.setattr(httpx2, "AsyncHTTPTransport", lambda **kwargs: transport) - with patch.object(httpx2, "AsyncClient", wraps=httpx2.AsyncClient) as async_client: - provider = EnterpriseShellProvider( - gateway_endpoint="http://gateway.example", - auth_token="secret", - proxy_timeout=90, - ) - - resource = await provider.attach("sandbox-1") - - timeout = async_client.call_args.kwargs["timeout"] - assert isinstance(timeout, httpx2.Timeout) - assert timeout.read == 90 - await resource.suspend() diff --git a/dify-agent/tests/local/dify_agent/adapters/shell/test_shellctl.py b/dify-agent/tests/local/dify_agent/adapters/shell/test_shellctl.py index 9af8a15dbb1..6984d3946ad 100644 --- a/dify-agent/tests/local/dify_agent/adapters/shell/test_shellctl.py +++ b/dify-agent/tests/local/dify_agent/adapters/shell/test_shellctl.py @@ -1,35 +1,26 @@ -"""Local tests for the shellctl shell adapter and env-driven provider factory.""" +"""Local tests for the shellctl command adapter.""" from __future__ import annotations import asyncio -import base64 -from collections.abc import Callable from dataclasses import dataclass, field from typing import cast import httpx2 as httpx import pytest -from pydantic import ValidationError +from shellctl.client import ShellctlClientError -from dify_agent.adapters.shell import shellctl -from dify_agent.adapters.shell.config import ShellAdapterSettings -from dify_agent.adapters.shell.factory import create_shell_provider -from dify_agent.adapters.shell.protocols import ShellCommandResult, ShellProviderError -from dify_agent.adapters.shell.shellctl import ( - ShellctlClientProtocol, - ShellFileTransferError, - ShellctlProvider, -) +from dify_agent.adapters.shell.protocols import ShellProviderError +from dify_agent.adapters.shell.shellctl import ShellctlClientProtocol, ShellctlCommands @dataclass(slots=True) class _Job: - job_id: str - status: str = "running" + job_id: str = "job-1" + status: str = "exited" done: bool = True - output: str = "" - offset: int = 0 + output: str = "ok" + offset: int = 2 truncated: bool = False exit_code: int | None = 0 output_path: str | None = "/tmp/output.log" @@ -37,410 +28,115 @@ class _Job: @dataclass(slots=True) class _Status: - job_id: str + job_id: str = "job-1" status: str = "terminated" done: bool = True - offset: int = 0 + offset: int = 2 exit_code: int | None = 130 @dataclass(slots=True) -class _RunCall: - script: str - cwd: str | None - env: dict[str, str] | None - timeout: float - - -type _RunHandler = Callable[[str, str | None, dict[str, str] | None, float], _Job] -type _WaitHandler = Callable[[str, int, float], _Job] -type _InputHandler = Callable[[str, str, int, float], _Job] -type _TerminateHandler = Callable[[str, float], _Status] - - -@dataclass(slots=True) -class FakeShellctlClient: - run_handler: _RunHandler | None = None - wait_handler: _WaitHandler | None = None - input_handler: _InputHandler | None = None - tail_handler: Callable[[str], _Job] | None = None - terminate_handler: _TerminateHandler | None = None - run_calls: list[_RunCall] = field(default_factory=list) +class _Client: + run_result: object = field(default_factory=_Job) + delete_error: Exception | None = None + run_calls: list[tuple[str, str | None, dict[str, str] | None, float]] = field(default_factory=list) wait_calls: list[tuple[str, int, float]] = field(default_factory=list) - input_calls: list[tuple[str, str, int, float]] = field(default_factory=list) - terminate_calls: list[tuple[str, float]] = field(default_factory=list) delete_calls: list[tuple[str, bool, float | None]] = field(default_factory=list) - closed: bool = False - async def run( - self, - script: str, - *, - cwd: str | None = None, - env: dict[str, str] | None = None, - timeout: float = 30.0, - ) -> _Job: - self.run_calls.append(_RunCall(script=script, cwd=cwd, env=env, timeout=timeout)) - if self.run_handler is not None: - return self.run_handler(script, cwd, env, timeout) - return _Job(job_id="job", status="exited", done=True, exit_code=0) + async def run(self, script: str, *, cwd=None, env=None, timeout=30.0): + self.run_calls.append((script, cwd, env, timeout)) + if isinstance(self.run_result, Exception): + raise self.run_result + return self.run_result - async def wait(self, job_id: str, *, offset: int, timeout: float = 30.0) -> _Job: + async def wait(self, job_id: str, *, offset: int, timeout=30.0): self.wait_calls.append((job_id, offset, timeout)) - if self.wait_handler is not None: - return self.wait_handler(job_id, offset, timeout) - return _Job(job_id=job_id, status="exited", done=True, offset=offset, exit_code=0) + return _Job(job_id=job_id) - async def input( - self, - job_id: str, - text: str, - *, - offset: int, - timeout: float = 30.0, - ) -> _Job: - self.input_calls.append((job_id, text, offset, timeout)) - if self.input_handler is not None: - return self.input_handler(job_id, text, offset, timeout) - return _Job(job_id=job_id, status="exited", done=True, offset=offset, exit_code=0) + async def input(self, job_id: str, text: str, *, offset: int, timeout=30.0): + return _Job(job_id=job_id) - async def tail(self, job_id: str) -> _Job: - if self.tail_handler is not None: - return self.tail_handler(job_id) - return _Job(job_id=job_id, status="exited", done=True, output="", exit_code=0) + async def tail(self, job_id: str): + return _Job(job_id=job_id) - async def terminate(self, job_id: str, grace_seconds: float = 10.0) -> _Status: - self.terminate_calls.append((job_id, grace_seconds)) - if self.terminate_handler is not None: - return self.terminate_handler(job_id, grace_seconds) + async def terminate(self, job_id: str, grace_seconds=10.0): return _Status(job_id=job_id) - async def delete( - self, - job_id: str, - *, - force: bool = False, - grace_seconds: float | None = None, - ) -> None: + async def delete(self, job_id: str, *, force=False, grace_seconds=None): self.delete_calls.append((job_id, force, grace_seconds)) - return None + if self.delete_error is not None: + raise self.delete_error + return object() async def close(self) -> None: - self.closed = True + return None -def _provider(client: FakeShellctlClient) -> ShellctlProvider: - return ShellctlProvider( - entrypoint="http://shellctl", - token="", - client_factory=lambda: _client_protocol(client), - ) - - -def _client_protocol(client: FakeShellctlClient) -> ShellctlClientProtocol: +def _client(client: _Client) -> ShellctlClientProtocol: return cast(ShellctlClientProtocol, cast(object, client)) -def test_factory_unknown_provider_raises() -> None: - with pytest.raises(ValidationError): - ShellAdapterSettings(shell_provider="nope") # type: ignore[arg-type] - - -def test_factory_shellctl_requires_entrypoint() -> None: - with pytest.raises(ValidationError, match="shellctl_entrypoint is required"): - ShellAdapterSettings(shell_provider="shellctl", shellctl_entrypoint=None) - - -def test_factory_builds_shellctl_provider_from_settings() -> None: - settings = ShellAdapterSettings(shell_provider="shellctl", shellctl_entrypoint="http://shellctl.example") - provider = create_shell_provider(settings) - assert isinstance(provider, ShellctlProvider) - assert provider.entrypoint == "http://shellctl.example" - assert provider.token == "" - - -def test_provider_create_opens_only_live_resource_and_suspend_closes_client() -> None: - client = FakeShellctlClient() +def test_commands_apply_runtime_layout_and_home_environment() -> None: + client = _Client() async def scenario() -> None: - resource = await _provider(client).create() - assert client.run_calls == [] - await resource.suspend() + commands = ShellctlCommands(_client(client), home_dir="/home/binding", workspace_dir="/workspace") + result = await commands.run("pwd", cwd="reports", env={"TOKEN": "value"}, timeout=2.5) + assert result.output == "ok" asyncio.run(scenario()) - assert client.closed is True + assert client.run_calls == [("pwd", "/workspace/reports", {"TOKEN": "value", "HOME": "/home/binding"}, 2.5)] -def test_commands_forward_parameters_and_map_metadata() -> None: - client = FakeShellctlClient( - run_handler=lambda script, cwd, env, timeout: _Job( - job_id="run-job", - status="running", - done=False, - output="abc", - offset=3, - truncated=True, - exit_code=None, - output_path="/tmp/run.log", - ), - wait_handler=lambda job_id, offset, timeout: _Job( - job_id=job_id, - status="running", - done=False, - output="def", - offset=6, - truncated=False, - exit_code=None, - output_path="/tmp/run.log", - ), - input_handler=lambda job_id, text, offset, timeout: _Job( - job_id=job_id, - status="exited", - done=True, - output="ghi", - offset=9, - truncated=False, - exit_code=0, - output_path="/tmp/run.log", - ), - tail_handler=lambda job_id: _Job( - job_id=job_id, - status="exited", - done=True, - output="tail", - offset=11, - truncated=False, - exit_code=0, - output_path="/tmp/tail.log", - ), - terminate_handler=lambda job_id, grace_seconds: _Status( - job_id=job_id, - status="terminated", - done=True, - offset=12, - exit_code=130, - ), - ) - +def test_commands_reject_cwd_outside_runtime_layout() -> None: async def scenario() -> None: - resource = await _provider(client).create() - run_result = await resource.commands.run("pwd", cwd="~/workspace/abc12ff", env={"FOO": "bar"}, timeout=2.5) - wait_result = await resource.commands.wait("run-job", offset=3, timeout=4.0) - read_result = await resource.commands.read_output("run-job", offset=6) - input_result = await resource.commands.input("run-job", "ls\n", offset=6, timeout=5.0) - interrupt_result = await resource.commands.interrupt("run-job", grace_seconds=1.5) - tail_result = await resource.commands.tail("run-job") - await resource.commands.delete("run-job", force=True, grace_seconds=2.0) - await resource.suspend() - - assert run_result == ShellCommandResult( - job_id="run-job", - status="running", - done=False, - exit_code=None, - output="abc", - offset=3, - truncated=True, - output_path="/tmp/run.log", - ) - assert wait_result.offset == 6 - assert read_result.offset == 6 - assert input_result.exit_code == 0 - assert interrupt_result.status == "terminated" - assert tail_result.output_path == "/tmp/tail.log" + commands = ShellctlCommands(_client(_Client()), home_dir="/home/binding", workspace_dir="/workspace") + with pytest.raises(ValueError, match="outside this RuntimeLease"): + await commands.run("pwd", cwd="/var/private", timeout=2.5) asyncio.run(scenario()) - assert client.run_calls == [_RunCall(script="pwd", cwd="~/workspace/abc12ff", env={"FOO": "bar"}, timeout=2.5)] - assert client.wait_calls == [ - ("run-job", 3, 4.0), - ("run-job", 6, 0.0), - ] - assert client.input_calls == [("run-job", "ls\n", 6, 5.0)] - assert client.terminate_calls == [("run-job", 1.5)] - assert client.delete_calls == [("run-job", True, 2.0)] + +def test_read_output_uses_nonblocking_wait() -> None: + client = _Client() + + async def scenario() -> None: + commands = ShellctlCommands(_client(client)) + result = await commands.read_output("job-1", offset=7) + assert result.job_id == "job-1" + + asyncio.run(scenario()) + assert client.wait_calls == [("job-1", 7, 0.0)] -def test_commands_map_http_timeout_to_shell_provider_error() -> None: +def test_commands_map_http_and_structured_errors() -> None: request = httpx.Request("POST", "http://shellctl.example/v1/jobs") - client = FakeShellctlClient( - run_handler=lambda script, cwd, env, timeout: (_ for _ in ()).throw( - httpx.ReadTimeout("timed out", request=request) + + async def scenario() -> None: + timeout_commands = ShellctlCommands( + _client(_Client(run_result=httpx.ReadTimeout("timed out", request=request))) ) - ) + with pytest.raises(ShellProviderError) as timeout_error: + await timeout_commands.run("pwd", timeout=2.5) + assert timeout_error.value.code == "timeout" - async def scenario() -> None: - resource = await _provider(client).create() - with pytest.raises(ShellProviderError, match="timed out") as exc_info: - await resource.commands.run("pwd", timeout=2.5) - assert exc_info.value.code == "timeout" - - asyncio.run(scenario()) - - -def test_commands_map_http_request_error_to_shell_provider_error() -> None: - request = httpx.Request("POST", "http://shellctl.example/v1/jobs/run") - client = FakeShellctlClient( - wait_handler=lambda job_id, offset, timeout: (_ for _ in ()).throw( - httpx.ConnectError("connection failed", request=request) + missing_commands = ShellctlCommands( + _client(_Client(run_result=ShellctlClientError(404, "sandbox_not_found", "expired"))) ) - ) - - async def scenario() -> None: - resource = await _provider(client).create() - with pytest.raises(ShellProviderError, match="connection failed") as exc_info: - await resource.commands.wait("run-job", offset=3, timeout=4.0) - assert exc_info.value.code == "request_error" + with pytest.raises(ShellProviderError) as missing_error: + await missing_commands.run("pwd", timeout=2.5) + assert missing_error.value.code == "sandbox_not_found" + assert missing_error.value.status_code == 404 asyncio.run(scenario()) -def test_delete_maps_http_timeout_to_shell_provider_error() -> None: - request = httpx.Request("DELETE", "http://shellctl.example/v1/jobs/run-job") - - @dataclass(slots=True) - class DeleteTimeoutClient(FakeShellctlClient): - async def delete(self, job_id, *, force=False, grace_seconds=None): - self.delete_calls.append((job_id, force, grace_seconds)) - raise httpx.ReadTimeout("delete timed out", request=request) - - client = DeleteTimeoutClient() +def test_delete_treats_missing_job_as_already_deleted() -> None: + client = _Client(delete_error=ShellctlClientError(404, "job_not_found", "missing")) async def scenario() -> None: - resource = await _provider(client).create() - with pytest.raises(ShellProviderError, match="delete timed out") as exc_info: - await resource.commands.delete("run-job", force=True, grace_seconds=2.0) - assert exc_info.value.code == "timeout" - - asyncio.run(scenario()) - assert client.delete_calls == [("run-job", True, 2.0)] - - -def test_delete_maps_http_request_error_to_shell_provider_error() -> None: - request = httpx.Request("DELETE", "http://shellctl.example/v1/jobs/run-job") - - @dataclass(slots=True) - class DeleteRequestErrorClient(FakeShellctlClient): - async def delete(self, job_id, *, force=False, grace_seconds=None): - self.delete_calls.append((job_id, force, grace_seconds)) - raise httpx.ConnectError("delete connection failed", request=request) - - client = DeleteRequestErrorClient() - - async def scenario() -> None: - resource = await _provider(client).create() - with pytest.raises(ShellProviderError, match="delete connection failed") as exc_info: - await resource.commands.delete("run-job", force=True, grace_seconds=2.0) - assert exc_info.value.code == "request_error" - - asyncio.run(scenario()) - assert client.delete_calls == [("run-job", True, 2.0)] - - -def test_files_upload_and_download_still_work() -> None: - content = b"hello \x00 world" - encoded = base64.b64encode(content).decode("ascii") - client = FakeShellctlClient( - run_handler=lambda script, cwd, env, timeout: ( - _Job(job_id="ul-job", status="exited", done=True, exit_code=0) - if "base64 -d" in script - else _Job( - job_id="dl-job", - status="exited", - done=True, - exit_code=0, - output=f"noise{shellctl._TRANSFER_BEGIN}{encoded}{shellctl._TRANSFER_END}tail", - ) - ) - ) - - async def scenario() -> None: - resource = await _provider(client).create() - await resource.files.upload(content=content, remote_path="out.bin", cwd="~/workspace/abc12ff") - downloaded = await resource.files.download(remote_path="report.txt", cwd="~/workspace/abc12ff") - assert downloaded == content - - asyncio.run(scenario()) - - -def test_file_transfer_timeout_is_an_end_to_end_budget(monkeypatch: pytest.MonkeyPatch) -> None: - clock = {"value": 100.0} - - def fake_monotonic() -> float: - return clock["value"] - - monkeypatch.setattr(shellctl.time, "monotonic", fake_monotonic) - - def run_handler(script: str, cwd: str | None, env: dict[str, str] | None, timeout: float) -> _Job: - del script, cwd, env - assert timeout == pytest.approx(5.0, rel=0, abs=0.01) - clock["value"] = 103.5 - return _Job(job_id="upload-job", status="running", done=False, output="part-1", offset=6, exit_code=None) - - def wait_handler(job_id: str, offset: int, timeout: float) -> _Job: - assert job_id == "upload-job" - assert offset == 6 - assert timeout == pytest.approx(1.5, rel=0, abs=0.01) - return _Job(job_id=job_id, status="exited", done=True, output="part-2", offset=12, exit_code=0) - - client = FakeShellctlClient(run_handler=run_handler, wait_handler=wait_handler) - - async def scenario() -> None: - transfer = shellctl.ShellctlFileTransfer( - client=_client_protocol(client), - timeout=5.0, - ) - await transfer.upload(content=b"payload", remote_path="out.bin") - - asyncio.run(scenario()) - assert client.delete_calls == [("upload-job", True, None)] - - -def test_file_transfer_timeout_exhaustion_raises_timeout_and_still_deletes_job( - monkeypatch: pytest.MonkeyPatch, -) -> None: - clock = {"value": 100.0} - - def fake_monotonic() -> float: - return clock["value"] - - monkeypatch.setattr(shellctl.time, "monotonic", fake_monotonic) - - def run_handler(script: str, cwd: str | None, env: dict[str, str] | None, timeout: float) -> _Job: - del script, cwd, env - assert timeout == pytest.approx(5.0, rel=0, abs=0.01) - clock["value"] = 106.0 - return _Job(job_id="upload-job", status="running", done=False, output="part-1", offset=6, exit_code=None) - - client = FakeShellctlClient(run_handler=run_handler) - - async def scenario() -> None: - transfer = shellctl.ShellctlFileTransfer( - client=_client_protocol(client), - timeout=5.0, - ) - with pytest.raises(ShellProviderError, match="timed out") as exc_info: - await transfer.upload(content=b"payload", remote_path="out.bin") - assert exc_info.value.code == "timeout" - - asyncio.run(scenario()) - assert client.delete_calls == [("upload-job", True, None)] - - -def test_download_missing_file_raises() -> None: - client = FakeShellctlClient( - run_handler=lambda script, cwd, env, timeout: _Job( - job_id="dl-job", - status="exited", - done=True, - output="", - exit_code=shellctl._DOWNLOAD_MISSING_EXIT_CODE, - ) - ) - - async def scenario() -> None: - resource = await _provider(client).create() - with pytest.raises(ShellFileTransferError, match="not found"): - await resource.files.download(remote_path="missing.txt") + commands = ShellctlCommands(_client(client)) + await commands.delete("job-1", force=True) asyncio.run(scenario()) + assert client.delete_calls == [("job-1", True, None)] diff --git a/dify-agent/tests/local/dify_agent/agent_stub/protocol/test_agent_stub_protocol.py b/dify-agent/tests/local/dify_agent/agent_stub/protocol/test_agent_stub_protocol.py index 0c16cf8bccf..8a99158b20f 100644 --- a/dify-agent/tests/local/dify_agent/agent_stub/protocol/test_agent_stub_protocol.py +++ b/dify-agent/tests/local/dify_agent/agent_stub/protocol/test_agent_stub_protocol.py @@ -12,6 +12,8 @@ from dify_agent.agent_stub.protocol.agent_stub import ( AgentStubDriveCommitRequest, AgentStubDriveFileRef, AgentStubDriveManifestResponse, + AgentStubConfigDownloadSource, + AgentStubFileDownloadRequest, AgentStubFileMapping, agent_stub_connections_url, agent_stub_drive_base_for_ref, @@ -20,7 +22,6 @@ from dify_agent.agent_stub.protocol.agent_stub import ( agent_stub_file_download_request_url, agent_stub_file_upload_request_url, normalize_agent_stub_api_base_url, - parse_agent_stub_endpoint, ) @@ -100,31 +101,16 @@ def test_normalize_agent_stub_api_base_url_accepts_service_root_or_agent_stub_ro def test_parse_agent_stub_endpoint_rejects_invalid_schemes_and_missing_host() -> None: - with pytest.raises(ValueError, match="http, https, or grpc"): + with pytest.raises(ValueError, match="http or https"): _ = normalize_agent_stub_api_base_url("not-a-url") - with pytest.raises(ValueError, match="http, https, or grpc"): + with pytest.raises(ValueError, match="http or https"): _ = normalize_agent_stub_api_base_url("ftp://agent.example.com/agent-stub") with pytest.raises(ValueError, match="include a host"): _ = normalize_agent_stub_api_base_url("https:///agent-stub") -def test_parse_agent_stub_endpoint_accepts_grpc_host_and_port() -> None: - endpoint = parse_agent_stub_endpoint("grpc://agent.example.com:9091") - - assert endpoint.url == "grpc://agent.example.com:9091" - assert endpoint.is_grpc is True - assert endpoint.host == "agent.example.com" - assert endpoint.port == 9091 - - -@pytest.mark.parametrize("invalid_url", ["grpc://agent.example.com", "grpc://agent.example.com:9091/path"]) -def test_parse_agent_stub_endpoint_rejects_invalid_grpc_urls(invalid_url: str) -> None: - with pytest.raises(ValueError): - _ = parse_agent_stub_endpoint(invalid_url) - - def test_agent_stub_file_mapping_validates_reference_and_url_by_transfer_method() -> None: reference = _reference("tool-file-1") assert AgentStubFileMapping(transfer_method="tool_file", reference=reference).reference == reference @@ -150,6 +136,60 @@ def test_agent_stub_file_mapping_rejects_remote_url_with_reference() -> None: ) +def test_agent_stub_file_download_request_accepts_legacy_http_audience_alias() -> None: + mapping = {"transfer_method": "tool_file", "reference": _reference("tool-file-1")} + + request = AgentStubFileDownloadRequest.model_validate({"file": mapping, "for_external": False}) + + assert request.for_frontend is False + assert request.model_dump() == { + "file": {"transfer_method": "tool_file", "reference": _reference("tool-file-1"), "url": None}, + "config": None, + "for_frontend": False, + } + + with pytest.raises(ValidationError): + _ = AgentStubFileDownloadRequest.model_validate({"file": mapping, "for_frontend": True, "for_external": False}) + + +def test_agent_stub_file_download_request_accepts_exactly_one_sandbox_config_source() -> None: + request = AgentStubFileDownloadRequest( + config=AgentStubConfigDownloadSource(kind="skill", name="alpha"), + for_frontend=False, + ) + + assert request.config == AgentStubConfigDownloadSource(kind="skill", name="alpha") + with pytest.raises(ValidationError, match="exactly one"): + _ = AgentStubFileDownloadRequest(for_frontend=False) + with pytest.raises(ValidationError, match="exactly one"): + _ = AgentStubFileDownloadRequest( + file=AgentStubFileMapping(transfer_method="tool_file", reference=_reference("tool-file-1")), + config=AgentStubConfigDownloadSource(kind="file", name="guide.txt"), + for_frontend=False, + ) + with pytest.raises(ValidationError, match="Sandbox data plane"): + _ = AgentStubFileDownloadRequest( + config=AgentStubConfigDownloadSource(kind="file", name="guide.txt"), + for_frontend=True, + ) + + +@pytest.mark.parametrize( + ("source", "message"), + [ + ({"kind": "file", "name": "../guide.txt"}, "safe path segment"), + ({"kind": "skill", "name": "Alpha"}, "skill name is invalid"), + ({"kind": "file", "name": "guide.txt", "tenant_id": "tenant-1"}, "extra_forbidden"), + ], +) +def test_agent_stub_config_download_source_rejects_invalid_names_and_identity_fields( + source: dict[str, object], + message: str, +) -> None: + with pytest.raises(ValidationError, match=message): + _ = AgentStubConfigDownloadSource.model_validate(source) + + def test_agent_stub_drive_commit_request_validates_file_refs() -> None: request = AgentStubDriveCommitRequest( items=[ diff --git a/dify-agent/tests/local/dify_agent/agent_stub/protocol/test_grpc_conversions.py b/dify-agent/tests/local/dify_agent/agent_stub/protocol/test_grpc_conversions.py deleted file mode 100644 index 8c061254151..00000000000 --- a/dify-agent/tests/local/dify_agent/agent_stub/protocol/test_grpc_conversions.py +++ /dev/null @@ -1,109 +0,0 @@ -from __future__ import annotations - -import base64 -import json - -import pytest - -pytest.importorskip("google.protobuf") - -from dify_agent.agent_stub.grpc._generated import agent_stub_pb2 -from dify_agent.agent_stub.grpc.conversions import ( - connect_request_from_proto, - connect_response_from_proto, - file_download_request_from_proto, - proto_connect_request, - proto_file_download_request, - proto_file_download_response, -) -from dify_agent.agent_stub.protocol.agent_stub import ( - AgentStubConnectResponse, - AgentStubFileDownloadResponse, - AgentStubFileMapping, -) - - -def _reference(record_id: str) -> str: - payload = base64.urlsafe_b64encode(json.dumps({"record_id": record_id}, separators=(",", ":")).encode()).decode() - return f"dify-file-ref:{payload}" - - -def test_connect_request_from_proto_round_trips_metadata_json() -> None: - message = proto_connect_request( - agent_stub_pb2, - argv=["connect", "--", "echo", "hello"], - metadata={"source": "cli"}, - ) - - request = connect_request_from_proto(message) - - assert request.argv == ["connect", "--", "echo", "hello"] - assert request.metadata == {"source": "cli"} - - -def test_file_download_request_from_proto_respects_optional_reference() -> None: - message = agent_stub_pb2.FileDownloadRequest( - file=agent_stub_pb2.FileMapping( - transfer_method="tool_file", - reference=_reference("tool-file-1"), - ) - ) - - request = file_download_request_from_proto(message) - - assert request.file.reference == _reference("tool-file-1") - assert request.file.url is None - assert request.for_external is True - - -def test_file_download_request_from_proto_preserves_explicit_internal_audience() -> None: - message = agent_stub_pb2.FileDownloadRequest( - file=agent_stub_pb2.FileMapping( - transfer_method="tool_file", - reference=_reference("tool-file-1"), - ), - for_external=False, - ) - - request = file_download_request_from_proto(message) - - assert request.for_external is False - - -def test_proto_file_download_request_preserves_selected_audience() -> None: - message = proto_file_download_request( - agent_stub_pb2, - file=AgentStubFileMapping(transfer_method="tool_file", reference=_reference("tool-file-1")), - for_external=False, - ) - - assert message.HasField("for_external") is True - assert message.for_external is False - - -def test_connect_request_from_proto_rejects_invalid_metadata_json() -> None: - with pytest.raises(json.JSONDecodeError): - _ = connect_request_from_proto( - agent_stub_pb2.ConnectRequest(protocol_version=1, argv=["connect"], metadata_json="not-json") - ) - - -def test_proto_responses_preserve_optional_fields() -> None: - message = proto_file_download_response( - AgentStubFileDownloadResponse( - filename="report.pdf", - mime_type="application/pdf", - size=123, - download_url="https://files.example.com/download", - ) - ) - - assert message.HasField("mime_type") is True - assert ( - connect_response_from_proto( - agent_stub_pb2.ConnectResponse( - **AgentStubConnectResponse(connection_id="conn-1", status="connected").model_dump() - ) - ).connection_id - == "conn-1" - ) diff --git a/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_app.py b/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_app.py index e683a7ea6db..8206a466e87 100644 --- a/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_app.py +++ b/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_app.py @@ -64,6 +64,7 @@ def test_create_agent_stub_app_wires_configured_file_handler_for_upload_requests server_secret_key=_base64url_secret(b"1" * 32), inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://files.example.com", ) token_codec = settings.create_agent_stub_token_codec() assert token_codec is not None @@ -72,9 +73,9 @@ def test_create_agent_stub_app_wires_configured_file_handler_for_upload_requests original_async_client = httpx.AsyncClient def handler(request: httpx.Request) -> httpx.Response: - assert str(request.url) == "https://api.example.com/inner/api/upload/file/request" + assert str(request.url) == "https://api.example.com/inner/api/agent/files/upload-request" assert request.headers["X-Inner-Api-Key"] == "inner-secret" - return httpx.Response(200, json={"data": {"url": "https://files.example.com/upload"}}) + return httpx.Response(200, json={"upload_uri": "/files/upload/for-plugin?sign=1"}) monkeypatch.setattr( "dify_agent.agent_stub.server.agent_stub_files.httpx.AsyncClient", @@ -89,7 +90,7 @@ def test_create_agent_stub_app_wires_configured_file_handler_for_upload_requests ) assert response.status_code == 200 - assert response.json() == {"upload_url": "https://files.example.com/upload"} + assert response.json() == {"upload_url": "https://files.example.com/files/upload/for-plugin?sign=1"} def test_create_agent_stub_app_wires_configured_drive_handler_for_manifest_requests(monkeypatch) -> None: @@ -98,6 +99,7 @@ def test_create_agent_stub_app_wires_configured_drive_handler_for_manifest_reque server_secret_key=_base64url_secret(b"1" * 32), inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://files.example.com", ) token_codec = settings.create_agent_stub_token_codec() assert token_codec is not None diff --git a/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_config.py b/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_config.py index 7846fcded42..6215192a98b 100644 --- a/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_config.py +++ b/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_config.py @@ -12,9 +12,12 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from dify_agent.agent_stub.protocol.agent_stub import ( + AgentStubConfigDownloadSource, AgentStubConfigManifestResponse, AgentStubConfigPushRequest, AgentStubConfigPushResponse, + AgentStubFileDownloadRequest, + AgentStubFileDownloadResponse, ) from dify_agent.agent_stub.server.agent_stub_config import ( AgentStubConfigRequestError, @@ -36,7 +39,7 @@ def _token_codec() -> AgentStubTokenCodec: def _execution_context(**updates: object) -> DifyExecutionContextLayerConfig: - payload = { + payload: dict[str, object] = { "tenant_id": "tenant-1", "user_id": "user-1", "user_from": "account", @@ -91,6 +94,7 @@ async def test_dify_api_handler_manifest_success(monkeypatch: pytest.MonkeyPatch response = await DifyApiAgentStubConfigRequestHandler( inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://sandbox-files.example.com", ).manifest(principal=_principal()) assert isinstance(response, AgentStubConfigManifestResponse) @@ -112,16 +116,29 @@ async def test_dify_api_handler_manifest_success(monkeypatch: pytest.MonkeyPatch @pytest.mark.anyio -async def test_dify_api_handler_pull_endpoints_return_bytes(monkeypatch: pytest.MonkeyPatch) -> None: +async def test_dify_api_handler_config_download_returns_sandbox_data_plane_url( + monkeypatch: pytest.MonkeyPatch, +) -> None: original_async_client = httpx.AsyncClient def handler(request: httpx.Request) -> httpx.Response: - assert request.url.params["user_id"] == "user-1" - if request.url.path.endswith("/skills/alpha/pull"): - return httpx.Response(200, content=b"zip-bytes") - if request.url.path.endswith("/files/guide.txt/pull"): - return httpx.Response(200, content=b"file-bytes") - raise AssertionError(f"unexpected path: {request.url.path}") + assert request.url.path == "/inner/api/agent-config/agent-1/download-request" + assert json.loads(request.content) == { + "tenant_id": "tenant-1", + "user_id": "user-1", + "config_version_id": "cfg-1", + "config_version_kind": "build_draft", + "config": {"kind": "skill", "name": "alpha"}, + } + return httpx.Response( + 200, + json={ + "filename": "alpha.zip", + "mime_type": "application/zip", + "size": 123, + "download_uri": "/files/tools/alpha.zip?sign=1", + }, + ) monkeypatch.setattr( "dify_agent.agent_stub.server.agent_stub_config.httpx.AsyncClient", @@ -131,10 +148,20 @@ async def test_dify_api_handler_pull_endpoints_return_bytes(monkeypatch: pytest. request_handler = DifyApiAgentStubConfigRequestHandler( inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://sandbox-files.example.com/dify", ) - assert await request_handler.pull_skill(principal=_principal(), name="alpha") == b"zip-bytes" - assert await request_handler.pull_file(principal=_principal(), name="guide.txt") == b"file-bytes" + response = await request_handler.create_download_request( + principal=_principal(), + source=AgentStubConfigDownloadSource(kind="skill", name="alpha"), + ) + + assert response == AgentStubFileDownloadResponse( + filename="alpha.zip", + mime_type="application/zip", + size=123, + download_url="https://sandbox-files.example.com/dify/files/tools/alpha.zip?sign=1", + ) @pytest.mark.anyio @@ -161,6 +188,7 @@ async def test_dify_api_handler_push_env_and_note_success(monkeypatch: pytest.Mo request_handler = DifyApiAgentStubConfigRequestHandler( inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://sandbox-files.example.com", ) push_response = await request_handler.push( @@ -221,8 +249,8 @@ async def test_dify_api_handler_push_env_and_note_success(monkeypatch: pytest.Mo "manifest", ), ( - "pull_skill", - "/inner/api/agent-config/agent-1/skills/alpha/pull", + "create_download_request", + "/inner/api/agent-config/agent-1/download-request", httpx.Response(404, json={"detail": "missing"}), 404, "missing", @@ -260,14 +288,18 @@ async def test_dify_api_handler_maps_error_cases( request_handler = DifyApiAgentStubConfigRequestHandler( inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://sandbox-files.example.com", ) with pytest.raises(AgentStubConfigRequestError, match=expected_message) as exc_info: match method_name: case "manifest": await request_handler.manifest(principal=_principal()) - case "pull_skill": - await request_handler.pull_skill(principal=_principal(), name="alpha") + case "create_download_request": + await request_handler.create_download_request( + principal=_principal(), + source=AgentStubConfigDownloadSource(kind="skill", name="alpha"), + ) case "push": await request_handler.push(principal=_principal(), request=AgentStubConfigPushRequest()) case "update_env": @@ -297,6 +329,7 @@ async def test_dify_api_handler_validates_required_execution_context_fields( request_handler = DifyApiAgentStubConfigRequestHandler( inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://sandbox-files.example.com", ) with pytest.raises(AgentStubConfigRequestError, match=expected_message) as exc_info: @@ -313,9 +346,8 @@ async def test_dify_api_handler_validates_required_execution_context_fields( "method_name", [ "get_config_manifest", - "pull_config_skill", + "create_file_download_request", "inspect_config_skill", - "pull_config_file", "push_config", "update_config_env", "update_config_note", @@ -330,18 +362,14 @@ async def test_control_plane_maps_config_request_errors(method_name: str) -> Non del principal raise AgentStubConfigRequestError(409, {"code": "conflict"}) - async def pull_skill(self, *, principal, name): - del principal, name + async def create_download_request(self, *, principal, source): + del principal, source raise AgentStubConfigRequestError(409, {"code": "conflict"}) async def inspect_skill(self, *, principal, name): del principal, name raise AgentStubConfigRequestError(409, {"code": "conflict"}) - async def pull_file(self, *, principal, name): - del principal, name - raise AgentStubConfigRequestError(409, {"code": "conflict"}) - async def push(self, *, principal, request): del principal, request raise AgentStubConfigRequestError(409, {"code": "conflict"}) @@ -362,12 +390,16 @@ async def test_control_plane_maps_config_request_errors(method_name: str) -> Non match method_name: case "get_config_manifest": await service.get_config_manifest(authorization=authorization) - case "pull_config_skill": - await service.pull_config_skill(name="alpha", authorization=authorization) + case "create_file_download_request": + await service.create_file_download_request( + request=AgentStubFileDownloadRequest( + config=AgentStubConfigDownloadSource(kind="skill", name="alpha"), + for_frontend=False, + ), + authorization=authorization, + ) case "inspect_config_skill": await service.inspect_config_skill(name="alpha", authorization=authorization) - case "pull_config_file": - await service.pull_config_file(name="guide.txt", authorization=authorization) case "push_config": await service.push_config(request=AgentStubConfigPushRequest(), authorization=authorization) case "update_config_env": @@ -385,9 +417,14 @@ async def test_control_plane_maps_config_request_errors(method_name: str) -> Non ("path", "method", "body", "expected_status", "expected_detail"), [ ("/agent-stub/config/manifest", "get", None, 200, {"agent_id": "agent-1"}), - ("/agent-stub/config/skills/alpha/pull", "get", None, 200, b"zip-bytes"), + ( + "/agent-stub/files/download-request", + "post", + {"config": {"kind": "skill", "name": "alpha"}, "for_frontend": False}, + 200, + {"download_url": "https://sandbox-files.example.com/files/alpha.zip"}, + ), ("/agent-stub/config/skills/alpha/inspect", "get", None, 200, {"name": "alpha", "files": ["SKILL.md"]}), - ("/agent-stub/config/files/guide.txt/pull", "get", None, 200, b"file-bytes"), ("/agent-stub/config/push", "post", {"note": "hello"}, 200, {"agent_id": "agent-1"}), ("/agent-stub/config/env", "patch", {"env_text": "API_KEY=value\n"}, 200, {"env_keys": ["API_KEY"]}), ("/agent-stub/config/note", "put", {"note": "hello"}, 200, {"note": "hello"}), @@ -398,7 +435,7 @@ def test_http_config_routes_forward_requests( method: str, body: dict[str, object] | None, expected_status: int, - expected_detail: dict[str, object] | bytes, + expected_detail: dict[str, object], ) -> None: codec = _token_codec() token = codec.encode_connection_token(_execution_context(), now=int(time.time()) - 1) @@ -409,21 +446,21 @@ def test_http_config_routes_forward_requests( captured["manifest_agent_id"] = principal.execution_context.agent_id return AgentStubConfigManifestResponse.model_validate(_manifest_payload()) - async def pull_skill(self, *, principal, name): - del principal - captured["skill_name"] = name - return b"zip-bytes" + async def create_download_request(self, *, principal, source): + captured["download_agent_id"] = principal.execution_context.agent_id + captured["download_source"] = source.model_dump() + return AgentStubFileDownloadResponse( + filename="alpha.zip", + mime_type="application/zip", + size=123, + download_url="https://sandbox-files.example.com/files/alpha.zip", + ) async def inspect_skill(self, *, principal, name): del principal captured["inspect_name"] = name return {"name": name, "files": ["SKILL.md"]} - async def pull_file(self, *, principal, name): - del principal - captured["file_name"] = name - return b"file-bytes" - async def push(self, *, principal, request): del principal captured["push_note"] = request.note @@ -441,29 +478,26 @@ def test_http_config_routes_forward_requests( app = FastAPI() app.include_router( - create_agent_stub_http_router(codec, config_request_handler=cast(AgentStubConfigRequestHandler, FakeHandler())) + create_agent_stub_http_router( + codec, + config_request_handler=cast(AgentStubConfigRequestHandler, cast(object, FakeHandler())), + ) ) client = TestClient(app) headers = {"Authorization": f"Bearer {token}"} response = client.request(method.upper(), path, headers=headers, json=body) assert response.status_code == expected_status - if isinstance(expected_detail, bytes): - assert response.content == expected_detail - else: - for key, value in expected_detail.items(): - assert response.json()[key] == value + for key, value in expected_detail.items(): + assert response.json()[key] == value if path.endswith("/manifest"): assert captured["manifest_agent_id"] == "agent-1" - elif path.endswith("/skills/alpha/pull"): - assert captured["skill_name"] == "alpha" - assert response.headers["content-type"] == "application/zip" + elif path.endswith("/files/download-request"): + assert captured["download_agent_id"] == "agent-1" + assert captured["download_source"] == {"kind": "skill", "name": "alpha"} elif path.endswith("/skills/alpha/inspect"): assert captured["inspect_name"] == "alpha" - elif path.endswith("/files/guide.txt/pull"): - assert captured["file_name"] == "guide.txt" - assert response.headers["content-type"] == "application/octet-stream" elif path.endswith("/config/push"): assert captured["push_note"] == "hello" elif path.endswith("/config/env"): @@ -481,18 +515,14 @@ def test_http_config_routes_map_handler_errors() -> None: del principal raise AgentStubConfigRequestError(422, {"code": "invalid_request"}) - async def pull_skill(self, *, principal, name): - del principal, name + async def create_download_request(self, *, principal, source): + del principal, source raise AssertionError("unexpected route") async def inspect_skill(self, *, principal, name): del principal, name raise AssertionError("unexpected route") - async def pull_file(self, *, principal, name): - del principal, name - raise AssertionError("unexpected route") - async def push(self, *, principal, request): del principal, request raise AssertionError("unexpected route") diff --git a/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_files.py b/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_files.py index 54e7863aeb4..2f204ac2b71 100644 --- a/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_files.py +++ b/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_files.py @@ -5,6 +5,7 @@ import base64 import json import httpx +import pytest from dify_agent.agent_stub.protocol.agent_stub import ( AgentStubFileDownloadRequest, @@ -33,7 +34,7 @@ def _principal() -> AgentStubPrincipal: ) -def _patch_async_client(monkeypatch, handler) -> None: +def _patch_async_client(monkeypatch: pytest.MonkeyPatch, handler) -> None: original_async_client = httpx.AsyncClient monkeypatch.setattr( "dify_agent.agent_stub.server.agent_stub_files.httpx.AsyncClient", @@ -41,236 +42,245 @@ def _patch_async_client(monkeypatch, handler) -> None: ) +def _file_handler(*, sandbox_files_base_url: str = "https://sandbox-files.example.com/dify"): + return DifyApiAgentStubFileRequestHandler( + inner_api_url="https://api.internal.example.com", + inner_api_key="inner-secret", + sandbox_files_base_url=sandbox_files_base_url, + ) + + def _reference(record_id: str) -> str: payload = base64.urlsafe_b64encode(json.dumps({"record_id": record_id}, separators=(",", ":")).encode()).decode() return f"dify-file-ref:{payload}" -def test_dify_api_agent_stub_file_handler_injects_execution_context_for_upload(monkeypatch) -> None: +def test_upload_request_uses_agent_inner_endpoint_and_binds_sandbox_base(monkeypatch: pytest.MonkeyPatch) -> None: def handler(request: httpx.Request) -> httpx.Response: - assert str(request.url) == "https://api.example.com/inner/api/upload/file/request" + assert str(request.url) == "https://api.internal.example.com/inner/api/agent/files/upload-request" assert request.headers["X-Inner-Api-Key"] == "inner-secret" assert json.loads(request.content) == { "tenant_id": "tenant-1", "user_id": "user-1", + "user_from": "account", "filename": "report.pdf", "mimetype": "application/pdf", "conversation_id": "conversation-1", } - return httpx.Response(200, json={"data": {"url": "https://files.example.com/upload"}}) + return httpx.Response(200, json={"upload_uri": "/files/upload/for-plugin?signed=yes"}) _patch_async_client(monkeypatch, handler) - file_handler = DifyApiAgentStubFileRequestHandler( - inner_api_url="https://api.example.com", - inner_api_key="inner-secret", - ) async def scenario() -> None: - response = await file_handler.create_upload_request( + response = await _file_handler().create_upload_request( principal=_principal(), request=AgentStubFileUploadRequest(filename="report.pdf", mimetype="application/pdf"), ) - assert response.upload_url == "https://files.example.com/upload" + assert response.upload_url == "https://sandbox-files.example.com/dify/files/upload/for-plugin?signed=yes" asyncio.run(scenario()) -def test_dify_api_agent_stub_file_handler_injects_execution_context_for_download(monkeypatch) -> None: +def test_sandbox_download_request_binds_origin_free_uri(monkeypatch: pytest.MonkeyPatch) -> None: + reference = _reference("tool-file-1") + def handler(request: httpx.Request) -> httpx.Response: - assert str(request.url) == "https://api.example.com/inner/api/download/file/request" + assert str(request.url) == "https://api.internal.example.com/inner/api/agent/files/download-request" assert json.loads(request.content) == { "tenant_id": "tenant-1", "user_id": "user-1", "user_from": "account", "invoke_from": "service-api", - "file": {"transfer_method": "tool_file", "reference": _reference("tool-file-1")}, + "file": {"transfer_method": "tool_file", "reference": reference}, + "for_frontend": False, } return httpx.Response( 200, json={ - "data": { - "filename": "report.pdf", - "mime_type": "application/pdf", - "size": 123, - "download_url": "https://files.example.com/download", - } + "filename": "report.pdf", + "mime_type": "application/pdf", + "size": 123, + "download_uri": "/files/tools/tool-file-1.pdf?sign=1", }, ) _patch_async_client(monkeypatch, handler) - file_handler = DifyApiAgentStubFileRequestHandler( - inner_api_url="https://api.example.com", - inner_api_key="inner-secret", - ) async def scenario() -> None: - response = await file_handler.create_download_request( + response = await _file_handler().create_download_request( principal=_principal(), request=AgentStubFileDownloadRequest( - file=AgentStubFileMapping(transfer_method="tool_file", reference=_reference("tool-file-1")) + file=AgentStubFileMapping(transfer_method="tool_file", reference=reference), + for_frontend=False, ), ) - assert response.download_url == "https://files.example.com/download" + assert response.download_url == "https://sandbox-files.example.com/dify/files/tools/tool-file-1.pdf?sign=1" asyncio.run(scenario()) -def test_dify_api_agent_stub_file_handler_forwards_internal_download_audience(monkeypatch) -> None: +@pytest.mark.parametrize( + "download_uri", + [ + "https://dify.example.com/files/tools/tool-file-1.pdf?sign=1", + "/files/tools/tool-file-1.pdf?sign=1", + ], +) +def test_frontend_download_request_preserves_public_or_relative_uri( + monkeypatch: pytest.MonkeyPatch, + download_uri: str, +) -> None: def handler(request: httpx.Request) -> httpx.Response: - assert str(request.url) == "https://api.example.com/inner/api/download/file/request" - assert json.loads(request.content)["for_external"] is False + assert json.loads(request.content)["for_frontend"] is True return httpx.Response( 200, - json={ - "data": { - "filename": "report.pdf", - "mime_type": "application/pdf", - "size": 123, - "download_url": "http://internal-files/report.pdf", - } - }, + json={"filename": "report.pdf", "mime_type": "application/pdf", "size": 123, "download_uri": download_uri}, ) _patch_async_client(monkeypatch, handler) - file_handler = DifyApiAgentStubFileRequestHandler( - inner_api_url="https://api.example.com", - inner_api_key="inner-secret", - ) async def scenario() -> None: - response = await file_handler.create_download_request( + response = await _file_handler().create_download_request( principal=_principal(), request=AgentStubFileDownloadRequest( file=AgentStubFileMapping(transfer_method="tool_file", reference=_reference("tool-file-1")), - for_external=False, + for_frontend=True, ), ) - assert response.download_url == "http://internal-files/report.pdf" + assert response.download_url == download_uri asyncio.run(scenario()) -def test_dify_api_agent_stub_file_handler_rejects_missing_user_id() -> None: - file_handler = DifyApiAgentStubFileRequestHandler( - inner_api_url="https://api.example.com", - inner_api_key="inner-secret", - ) +def test_remote_download_url_is_never_rewritten(monkeypatch: pytest.MonkeyPatch) -> None: + remote_url = "https://remote.example.com/report.pdf" + + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={"filename": "report.pdf", "mime_type": "application/pdf", "size": 123, "download_uri": remote_url}, + ) + + _patch_async_client(monkeypatch, handler) + + async def scenario() -> None: + response = await _file_handler().create_download_request( + principal=_principal(), + request=AgentStubFileDownloadRequest( + file=AgentStubFileMapping(transfer_method="remote_url", url=remote_url), + for_frontend=False, + ), + ) + assert response.download_url == remote_url + + asyncio.run(scenario()) + + +def test_remote_download_rejects_relative_uri_for_frontend_audience(monkeypatch: pytest.MonkeyPatch) -> None: + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={"filename": "report.pdf", "mime_type": "application/pdf", "size": 123, "download_uri": "/files/x"}, + ) + + _patch_async_client(monkeypatch, handler) + + async def scenario() -> None: + with pytest.raises(AgentStubFileRequestError, match="invalid remote download URL"): + await _file_handler().create_download_request( + principal=_principal(), + request=AgentStubFileDownloadRequest( + file=AgentStubFileMapping( + transfer_method="remote_url", url="https://remote.example.com/report.pdf" + ), + for_frontend=True, + ), + ) + + asyncio.run(scenario()) + + +@pytest.mark.parametrize( + "unsafe_uri", + [ + "//attacker.example/files/x", + "/files/../admin", + "/files/%252e%252e/admin", + "http://api:5001/files/tools/x", + "/not-files/x", + ], +) +def test_sandbox_download_rejects_unsafe_dify_file_uri( + monkeypatch: pytest.MonkeyPatch, + unsafe_uri: str, +) -> None: + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={"filename": "x", "mime_type": None, "size": 1, "download_uri": unsafe_uri}, + ) + + _patch_async_client(monkeypatch, handler) + + async def scenario() -> None: + with pytest.raises(AgentStubFileRequestError, match="unsafe Dify file URI"): + await _file_handler().create_download_request( + principal=_principal(), + request=AgentStubFileDownloadRequest( + file=AgentStubFileMapping(transfer_method="tool_file", reference=_reference("tool-file-1")), + for_frontend=False, + ), + ) + + asyncio.run(scenario()) + + +def test_handler_rejects_missing_execution_user_before_network() -> None: principal = _principal() principal.execution_context = principal.execution_context.model_copy(update={"user_id": None}) async def scenario() -> None: - try: - await file_handler.create_upload_request( + with pytest.raises(AgentStubFileRequestError, match="user_id"): + await _file_handler().create_upload_request( principal=principal, request=AgentStubFileUploadRequest(filename="report.pdf", mimetype="application/pdf"), ) - except AgentStubFileRequestError as exc: - assert "user_id" in str(exc) - else: - raise AssertionError("expected AgentStubFileRequestError") asyncio.run(scenario()) -def test_dify_api_agent_stub_file_handler_maps_non_2xx_response(monkeypatch) -> None: +def test_handler_preserves_inner_api_error_status(monkeypatch: pytest.MonkeyPatch) -> None: def handler(_request: httpx.Request) -> httpx.Response: return httpx.Response(403, json={"detail": "forbidden"}) _patch_async_client(monkeypatch, handler) - file_handler = DifyApiAgentStubFileRequestHandler( - inner_api_url="https://api.example.com", - inner_api_key="inner-secret", - ) async def scenario() -> None: - try: - await file_handler.create_upload_request( + with pytest.raises(AgentStubFileRequestError) as exc_info: + await _file_handler().create_upload_request( principal=_principal(), request=AgentStubFileUploadRequest(filename="report.pdf", mimetype="application/pdf"), ) - except AgentStubFileRequestError as exc: - assert exc.status_code == 403 - assert exc.detail == "forbidden" - else: - raise AssertionError("expected AgentStubFileRequestError") + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == "forbidden" asyncio.run(scenario()) -def test_dify_api_agent_stub_file_handler_maps_error_envelope(monkeypatch) -> None: +def test_handler_rejects_missing_download_uri(monkeypatch: pytest.MonkeyPatch) -> None: def handler(_request: httpx.Request) -> httpx.Response: - return httpx.Response(200, json={"error": "bad request"}) + return httpx.Response(200, json={"filename": "report.pdf", "size": 1}) _patch_async_client(monkeypatch, handler) - file_handler = DifyApiAgentStubFileRequestHandler( - inner_api_url="https://api.example.com", - inner_api_key="inner-secret", - ) async def scenario() -> None: - try: - await file_handler.create_download_request( + with pytest.raises(AgentStubFileRequestError, match="missing download_uri"): + await _file_handler().create_download_request( principal=_principal(), request=AgentStubFileDownloadRequest( - file=AgentStubFileMapping(transfer_method="tool_file", reference=_reference("tool-file-1")) + file=AgentStubFileMapping(transfer_method="tool_file", reference=_reference("tool-file-1")), + for_frontend=False, ), ) - except AgentStubFileRequestError as exc: - assert exc.status_code == 400 - assert exc.detail == "bad request" - else: - raise AssertionError("expected AgentStubFileRequestError") - - asyncio.run(scenario()) - - -def test_dify_api_agent_stub_file_handler_rejects_upload_response_missing_url(monkeypatch) -> None: - def handler(_request: httpx.Request) -> httpx.Response: - return httpx.Response(200, json={"data": {}}) - - _patch_async_client(monkeypatch, handler) - file_handler = DifyApiAgentStubFileRequestHandler( - inner_api_url="https://api.example.com", - inner_api_key="inner-secret", - ) - - async def scenario() -> None: - try: - await file_handler.create_upload_request( - principal=_principal(), - request=AgentStubFileUploadRequest(filename="report.pdf", mimetype="application/pdf"), - ) - except AgentStubFileRequestError as exc: - assert exc.status_code == 502 - assert exc.detail == "Dify API upload request response is missing url" - else: - raise AssertionError("expected AgentStubFileRequestError") - - asyncio.run(scenario()) - - -def test_dify_api_agent_stub_file_handler_rejects_invalid_download_response_schema(monkeypatch) -> None: - def handler(_request: httpx.Request) -> httpx.Response: - return httpx.Response(200, json={"data": {"filename": "report.pdf"}}) - - _patch_async_client(monkeypatch, handler) - file_handler = DifyApiAgentStubFileRequestHandler( - inner_api_url="https://api.example.com", - inner_api_key="inner-secret", - ) - - async def scenario() -> None: - try: - await file_handler.create_download_request( - principal=_principal(), - request=AgentStubFileDownloadRequest( - file=AgentStubFileMapping(transfer_method="tool_file", reference=_reference("tool-file-1")) - ), - ) - except AgentStubFileRequestError as exc: - assert exc.status_code == 502 - assert exc.detail == "Dify API download request response is invalid" - else: - raise AssertionError("expected AgentStubFileRequestError") asyncio.run(scenario()) diff --git a/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_routes.py b/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_routes.py index ab285cc681a..155835f5d25 100644 --- a/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_routes.py +++ b/dify-agent/tests/local/dify_agent/agent_stub/server/test_agent_stub_routes.py @@ -158,6 +158,7 @@ def test_agent_stub_file_download_route_forwards_authenticated_request() -> None class FakeHandler: async def create_download_request(self, *, principal, request): assert principal.execution_context.user_id == "user-1" + assert request.file is not None assert request.file.transfer_method == "tool_file" return AgentStubFileDownloadResponse( filename="report.pdf", diff --git a/dify-agent/tests/local/dify_agent/agent_stub/server/test_cli.py b/dify-agent/tests/local/dify_agent/agent_stub/server/test_cli.py index cb40a3b583b..610b5ec9d0f 100644 --- a/dify-agent/tests/local/dify_agent/agent_stub/server/test_cli.py +++ b/dify-agent/tests/local/dify_agent/agent_stub/server/test_cli.py @@ -1,10 +1,6 @@ from __future__ import annotations -import asyncio -from typing import cast - import dify_agent.agent_stub.server.cli as cli_module -from dify_agent.server.settings import ServerSettings def test_stub_server_cli_uses_default_uvicorn_settings(monkeypatch) -> None: @@ -41,101 +37,3 @@ def test_stub_server_cli_passes_explicit_uvicorn_settings(monkeypatch) -> None: "port": 9000, "reload": True, } - - -def test_stub_server_cli_switches_to_grpc_when_agent_stub_api_base_url_uses_grpc(monkeypatch) -> None: - captured: dict[str, object] = {} - - async def fake_serve_grpc(*, settings, host, port) -> None: - captured.update(settings=settings, host=host, port=port) - - monkeypatch.setattr(cli_module, "_serve_grpc", fake_serve_grpc) - monkeypatch.setattr( - cli_module, "ServerSettings", lambda: type("Settings", (), {"agent_stub_api_base_url": "grpc://agent:9091"})() - ) - - cli_module.main(["--host", "0.0.0.0", "--port", "9092"]) - - assert captured["host"] == "0.0.0.0" - assert captured["port"] == 9092 - - -def test_serve_grpc_derives_default_bind_target_and_closes_server(monkeypatch) -> None: - captured: dict[str, object] = {} - - class FakeServer: - def __init__(self) -> None: - self.closed = False - - async def aclose(self) -> None: - self.closed = True - - fake_server = FakeServer() - - async def fake_start_agent_stub_grpc_server(**kwargs): - captured.update(kwargs) - return fake_server - - class FakeEvent: - async def wait(self) -> None: - return None - - settings = type( - "Settings", - (), - { - "agent_stub_api_base_url": "grpc://agent.example.com:9091", - "agent_stub_grpc_bind_address": None, - "create_agent_stub_token_codec": lambda self: "token-codec", - "create_agent_stub_file_request_handler": lambda self: "file-handler", - }, - )() - - monkeypatch.setattr(cli_module, "start_agent_stub_grpc_server", fake_start_agent_stub_grpc_server) - monkeypatch.setattr(cli_module.asyncio, "Event", FakeEvent) - - asyncio.run(cli_module._serve_grpc(settings=cast(ServerSettings, cast(object, settings)), host=None, port=None)) - - assert captured == { - "public_url": "grpc://agent.example.com:9091", - "bind_address": "0.0.0.0:9091", - "token_codec": "token-codec", - "file_request_handler": "file-handler", - } - assert fake_server.closed is True - - -def test_serve_grpc_applies_cli_host_port_overrides(monkeypatch) -> None: - captured: dict[str, object] = {} - - class FakeServer: - async def aclose(self) -> None: - return None - - async def fake_start_agent_stub_grpc_server(**kwargs): - captured.update(kwargs) - return FakeServer() - - class FakeEvent: - async def wait(self) -> None: - return None - - settings = type( - "Settings", - (), - { - "agent_stub_api_base_url": "grpc://agent.example.com:9091", - "agent_stub_grpc_bind_address": "127.0.0.1:9191", - "create_agent_stub_token_codec": lambda self: None, - "create_agent_stub_file_request_handler": lambda self: None, - }, - )() - - monkeypatch.setattr(cli_module, "start_agent_stub_grpc_server", fake_start_agent_stub_grpc_server) - monkeypatch.setattr(cli_module.asyncio, "Event", FakeEvent) - - asyncio.run( - cli_module._serve_grpc(settings=cast(ServerSettings, cast(object, settings)), host="0.0.0.0", port=9292) - ) - - assert captured["bind_address"] == "0.0.0.0:9292" diff --git a/dify-agent/tests/local/dify_agent/agent_stub/server/test_grpc_bind.py b/dify-agent/tests/local/dify_agent/agent_stub/server/test_grpc_bind.py deleted file mode 100644 index 0fe7ad36661..00000000000 --- a/dify-agent/tests/local/dify_agent/agent_stub/server/test_grpc_bind.py +++ /dev/null @@ -1,37 +0,0 @@ -from __future__ import annotations - -import pytest - -from dify_agent.agent_stub.server.grpc_bind import ( - derive_agent_stub_grpc_bind_target, - parse_agent_stub_grpc_bind_address, -) - - -def test_derive_agent_stub_grpc_bind_target_defaults_to_all_interfaces() -> None: - target = derive_agent_stub_grpc_bind_target(public_url="grpc://agent.example.com:9091") - - assert target.host == "0.0.0.0" - assert target.port == 9091 - assert target.address == "0.0.0.0:9091" - - -def test_derive_agent_stub_grpc_bind_target_prefers_explicit_override() -> None: - target = derive_agent_stub_grpc_bind_target( - public_url="grpc://agent.example.com:9091", - bind_address="127.0.0.1:9191", - ) - - assert target.host == "127.0.0.1" - assert target.port == 9191 - - -def test_parse_agent_stub_grpc_bind_address_rejects_missing_port() -> None: - with pytest.raises(ValueError, match="explicit port"): - _ = parse_agent_stub_grpc_bind_address("127.0.0.1") - - -@pytest.mark.parametrize("value", ["user@0.0.0.0:9091", "user:password@0.0.0.0:9091"]) -def test_parse_agent_stub_grpc_bind_address_rejects_user_info(value: str) -> None: - with pytest.raises(ValueError, match="must not include user info"): - _ = parse_agent_stub_grpc_bind_address(value) diff --git a/dify-agent/tests/local/dify_agent/agent_stub/server/test_grpc_service.py b/dify-agent/tests/local/dify_agent/agent_stub/server/test_grpc_service.py deleted file mode 100644 index 149ce1cbffe..00000000000 --- a/dify-agent/tests/local/dify_agent/agent_stub/server/test_grpc_service.py +++ /dev/null @@ -1,279 +0,0 @@ -from __future__ import annotations - -import base64 -import json -import secrets -from types import SimpleNamespace -from typing import cast - -import pytest - -pytest.importorskip("grpclib") -pytest.importorskip("google.protobuf") - -from grpclib.const import Status -from grpclib.exceptions import GRPCError - -from dify_agent.agent_stub.grpc._generated import agent_stub_pb2 -from dify_agent.agent_stub.server.agent_stub_files import AgentStubFileRequestError, AgentStubFileRequestHandler -from dify_agent.agent_stub.server.control_plane import AgentStubControlPlaneService -from dify_agent.agent_stub.server.grpc_service import AgentStubGRPCTransport -from dify_agent.agent_stub.server.tokens.agent_stub import AgentStubTokenCodec -from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig - - -def _base64url_secret(value: bytes) -> str: - return base64.urlsafe_b64encode(value).rstrip(b"=").decode("ascii") - - -def _token_codec() -> AgentStubTokenCodec: - return AgentStubTokenCodec.from_server_secret(_base64url_secret(secrets.token_bytes(32))) - - -def _execution_context() -> DifyExecutionContextLayerConfig: - return DifyExecutionContextLayerConfig( - tenant_id="tenant-1", - user_id="user-1", - user_from="account", - agent_mode="workflow_run", - invoke_from="service-api", - ) - - -def _reference(record_id: str) -> str: - payload = base64.urlsafe_b64encode(json.dumps({"record_id": record_id}, separators=(",", ":")).encode()).decode() - return f"dify-file-ref:{payload}" - - -def test_agent_stub_grpc_transport_connects_with_bearer_metadata() -> None: - codec = _token_codec() - token = codec.encode_connection_token(_execution_context()) - transport = AgentStubGRPCTransport( - AgentStubControlPlaneService( - codec, - connection_id_factory=lambda: "conn-1", - ) - ) - - async def scenario() -> None: - response = await transport.connect( - request=agent_stub_pb2.ConnectRequest( - protocol_version=1, - argv=["connect"], - metadata_json='{"source":"cli"}', - ), - metadata=(("authorization", f"Bearer {token}"),), - ) - assert response.connection_id == "conn-1" - assert response.status == "connected" - - import asyncio - - asyncio.run(scenario()) - - -def test_agent_stub_grpc_transport_maps_missing_authorization_to_unauthenticated() -> None: - codec = _token_codec() - transport = AgentStubGRPCTransport(AgentStubControlPlaneService(codec)) - - async def scenario() -> None: - with pytest.raises(GRPCError) as exc_info: - await transport.connect( - request=agent_stub_pb2.ConnectRequest(protocol_version=1, argv=["connect"], metadata_json="{}"), - metadata=(), - ) - assert exc_info.value.status == Status.UNAUTHENTICATED - - import asyncio - - asyncio.run(scenario()) - - -def test_agent_stub_grpc_transport_delegates_file_upload_requests() -> None: - codec = _token_codec() - token = codec.encode_connection_token(_execution_context()) - - class FakeHandler: - async def create_upload_request(self, *, principal, request): - assert principal.execution_context.tenant_id == "tenant-1" - assert request.filename == "report.pdf" - return type("Response", (), {"upload_url": "https://files.example.com/upload"})() - - async def create_download_request(self, *, principal, request): - del principal, request - raise AssertionError("unexpected download request") - - transport = AgentStubGRPCTransport( - AgentStubControlPlaneService( - codec, - cast(AgentStubFileRequestHandler, cast(object, FakeHandler())), - ) - ) - - async def scenario() -> None: - response = await transport.create_file_upload_request( - request=agent_stub_pb2.FileUploadRequest(filename="report.pdf", mimetype="application/pdf"), - metadata=(("authorization", f"Bearer {token}"),), - ) - assert response.upload_url == "https://files.example.com/upload" - - import asyncio - - asyncio.run(scenario()) - - -def test_agent_stub_grpc_transport_delegates_file_download_requests() -> None: - codec = _token_codec() - token = codec.encode_connection_token(_execution_context()) - - class FakeHandler: - async def create_download_request(self, *, principal, request): - assert principal.execution_context.user_id == "user-1" - assert request.file.reference == _reference("tool-file-1") - assert request.for_external is False - return type( - "Response", - (), - { - "filename": "report.pdf", - "mime_type": "application/pdf", - "size": 123, - "download_url": "https://files.example.com/download", - }, - )() - - async def create_upload_request(self, *, principal, request): - del principal, request - raise AssertionError("unexpected upload request") - - transport = AgentStubGRPCTransport( - AgentStubControlPlaneService( - codec, - cast(AgentStubFileRequestHandler, cast(object, FakeHandler())), - ) - ) - - async def scenario() -> None: - response = await transport.create_file_download_request( - request=agent_stub_pb2.FileDownloadRequest( - file=agent_stub_pb2.FileMapping( - transfer_method="tool_file", - reference=_reference("tool-file-1"), - ), - for_external=False, - ), - metadata=(("authorization", f"Bearer {token}"),), - ) - assert response.download_url == "https://files.example.com/download" - - import asyncio - - asyncio.run(scenario()) - - -def test_agent_stub_grpc_transport_stringifies_structured_file_error_details() -> None: - codec = _token_codec() - token = codec.encode_connection_token(_execution_context()) - - class FakeHandler: - async def create_upload_request(self, *, principal, request): - del principal, request - raise AgentStubFileRequestError(400, {"detail": "bad request", "code": "inner_api_error"}) - - async def create_download_request(self, *, principal, request): - del principal, request - raise AssertionError("unexpected download request") - - transport = AgentStubGRPCTransport( - AgentStubControlPlaneService( - codec, - cast(AgentStubFileRequestHandler, cast(object, FakeHandler())), - ) - ) - - async def scenario() -> None: - with pytest.raises(GRPCError) as exc_info: - await transport.create_file_upload_request( - request=agent_stub_pb2.FileUploadRequest(filename="report.pdf", mimetype="application/pdf"), - metadata=(("authorization", f"Bearer {token}"),), - ) - assert exc_info.value.status == Status.FAILED_PRECONDITION - assert isinstance(exc_info.value.message, str) - assert "bad request" in exc_info.value.message - assert "inner_api_error" in exc_info.value.message - - import asyncio - - asyncio.run(scenario()) - - -def test_agent_stub_grpc_transport_maps_missing_token_codec_to_unavailable() -> None: - transport = AgentStubGRPCTransport(AgentStubControlPlaneService(None)) - - async def scenario() -> None: - with pytest.raises(GRPCError) as exc_info: - await transport.connect( - request=agent_stub_pb2.ConnectRequest(protocol_version=1, argv=["connect"], metadata_json="{}"), - metadata=(("authorization", "Bearer token"),), - ) - assert exc_info.value.status == Status.UNAVAILABLE - - import asyncio - - asyncio.run(scenario()) - - -def test_agent_stub_grpc_transport_maps_missing_file_handler_to_unavailable() -> None: - codec = _token_codec() - token = codec.encode_connection_token(_execution_context()) - transport = AgentStubGRPCTransport(AgentStubControlPlaneService(codec, None)) - - async def scenario() -> None: - with pytest.raises(GRPCError) as exc_info: - await transport.create_file_upload_request( - request=agent_stub_pb2.FileUploadRequest(filename="report.pdf", mimetype="application/pdf"), - metadata=(("authorization", f"Bearer {token}"),), - ) - assert exc_info.value.status == Status.UNAVAILABLE - - import asyncio - - asyncio.run(scenario()) - - -def test_agent_stub_grpc_transport_maps_invalid_upload_request_to_invalid_argument() -> None: - codec = _token_codec() - token = codec.encode_connection_token(_execution_context()) - transport = AgentStubGRPCTransport(AgentStubControlPlaneService(codec)) - - async def scenario() -> None: - with pytest.raises(GRPCError) as exc_info: - await transport.create_file_upload_request( - request=SimpleNamespace(filename=None, mimetype="application/pdf"), - metadata=(("authorization", f"Bearer {token}"),), - ) - assert exc_info.value.status == Status.INVALID_ARGUMENT - - import asyncio - - asyncio.run(scenario()) - - -def test_agent_stub_grpc_transport_maps_invalid_download_request_to_invalid_argument() -> None: - codec = _token_codec() - token = codec.encode_connection_token(_execution_context()) - transport = AgentStubGRPCTransport(AgentStubControlPlaneService(codec)) - - async def scenario() -> None: - with pytest.raises(GRPCError) as exc_info: - await transport.create_file_download_request( - request=agent_stub_pb2.FileDownloadRequest( - file=agent_stub_pb2.FileMapping(transfer_method="tool_file") - ), - metadata=(("authorization", f"Bearer {token}"),), - ) - assert exc_info.value.status == Status.INVALID_ARGUMENT - - import asyncio - - asyncio.run(scenario()) diff --git a/dify-agent/tests/local/dify_agent/client/test_client.py b/dify-agent/tests/local/dify_agent/client/test_client.py index 37d5c540987..e08c39d9009 100644 --- a/dify-agent/tests/local/dify_agent/client/test_client.py +++ b/dify-agent/tests/local/dify_agent/client/test_client.py @@ -2,7 +2,7 @@ from __future__ import annotations import asyncio import json -from collections.abc import Iterator +from collections.abc import AsyncIterator, Iterator from datetime import UTC, datetime from typing import cast, override @@ -11,6 +11,7 @@ import pytest from agenton.compositor import CompositorSessionSnapshot from agenton_collections.layers.plain import PLAIN_PROMPT_LAYER_TYPE_ID +from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig from dify_agent.client import _client as client_module from dify_agent.client import ( Client, @@ -21,9 +22,16 @@ from dify_agent.client import ( DifyAgentValidationError, ) from dify_agent.protocol import ( + BindingFileDownloadRequest, + BindingFileDownloadResponse, + BindingFileListResponse, + BindingFileReadResponse, CancelRunRequest, CancelRunResponse, + CreateExecutionBindingRequest, + CreateHomeSnapshotFromBindingRequest, CreateRunRequest, + DestroyExecutionBindingRequest, RUN_EVENT_ADAPTER, RunCancelledEvent, RunEvent, @@ -33,10 +41,6 @@ from dify_agent.protocol import ( RunStartedEvent, RunSucceededEvent, RunSucceededEventData, - SandboxListResponse, - SandboxLocator, - SandboxReadResponse, - SandboxUploadResponse, ) @@ -79,44 +83,25 @@ def _run_status_json(status: str) -> dict[str, object]: return {"run_id": "run-1", "status": status, "created_at": now, "updated_at": now, "error": None} -def _sandbox_locator() -> SandboxLocator: - return SandboxLocator.model_validate( - { - "composition": { - "schema_version": 1, - "layers": [ - { - "name": "execution_context", - "type": "dify.execution_context", - "config": { - "tenant_id": "tenant-1", - "user_from": "account", - "agent_mode": "agent_app", - "invoke_from": "service-api", - }, - }, - { - "name": "shell", - "type": "dify.shell", - "deps": {"execution_context": "execution_context"}, - "config": {}, - }, - ], - }, - "session_snapshot": { - "layers": [ - {"name": "execution_context", "lifecycle_state": "suspended", "runtime_state": {}}, - { - "name": "shell", - "lifecycle_state": "suspended", - "runtime_state": {"session_id": "abc12ff", "workspace_cwd": "~/workspace/abc12ff"}, - }, - ] - }, - } +def _binding_file_download_request(path: str = "report.txt") -> BindingFileDownloadRequest: + return BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path=path, + execution_context=DifyExecutionContextLayerConfig( + tenant_id="tenant-1", + user_id="account-1", + user_from="account", + agent_mode="agent_app", + invoke_from="debugger", + ), ) +def _assert_binding_download_timeout(request: httpx.Request) -> None: + timeout = cast(dict[str, float], request.extensions["timeout"]) + assert timeout == {"connect": 90.0, "read": 90.0, "write": 90.0, "pool": 90.0} + + def _function_tool_result_payload(key: str) -> dict[str, object]: return { "type": "pydantic_ai_event", @@ -146,6 +131,19 @@ class DisconnectingSyncStream(httpx.SyncByteStream): raise httpx.ReadError("stream disconnected") +class DisconnectingAsyncStream(httpx.AsyncByteStream): + chunks: list[bytes] + + def __init__(self, *chunks: str) -> None: + self.chunks = [chunk.encode() for chunk in chunks] + + @override + async def __aiter__(self) -> AsyncIterator[bytes]: + for chunk in self.chunks: + yield chunk + raise httpx.ReadError("stream disconnected") + + def test_sse_decoder_accepts_function_tool_result_part_alias(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(client_module, "_function_tool_result_payload_key_cache", "part") decoder = client_module._SSEDecoder() @@ -247,117 +245,213 @@ def test_async_methods_and_wait_run_parse_protocol_dtos() -> None: asyncio.run(scenario()) -def test_sync_sandbox_methods_post_dtos_and_parse_responses() -> None: - locator = _sandbox_locator() - +def test_sync_binding_file_methods_post_dtos_and_parse_responses() -> None: def handler(request: httpx.Request) -> httpx.Response: - if request.url.path == "/sandbox/files/list": + if request.url.path == "/execution-bindings/files/list": payload = cast(dict[str, object], json.loads(request.content)) assert payload["path"] == "." + assert payload["backend_binding_ref"] == "binding-ref" return httpx.Response(200, json={"path": ".", "entries": [], "truncated": False}) - if request.url.path == "/sandbox/files/read": + if request.url.path == "/execution-bindings/files/read": payload = cast(dict[str, object], json.loads(request.content)) assert payload["path"] == "note.txt" assert payload["max_bytes"] == 128 return httpx.Response( 200, json={"path": "note.txt", "size": 5, "truncated": False, "binary": False, "text": "hello"} ) - if request.url.path == "/sandbox/files/upload": + if request.url.path == "/execution-bindings/files/download": + _assert_binding_download_timeout(request) payload = cast(dict[str, object], json.loads(request.content)) assert payload["path"] == "report.txt" return httpx.Response( 200, - json={ - "path": "report.txt", - "file": { - "transfer_method": "tool_file", - "reference": "dify-file-ref:file-1", - "download_url": "https://files.example.com/report.txt", - }, - }, + json={"reference": "dify-file-ref:file-1"}, ) raise AssertionError(f"unexpected request: {request.method} {request.url}") client = Client(base_url="http://testserver", sync_http_client=httpx.Client(transport=httpx.MockTransport(handler))) - listing = client.list_sandbox_files_sync(locator, ".") - preview = client.read_sandbox_file_sync(locator, "note.txt", max_bytes=128) - uploaded = client.upload_sandbox_file_sync(locator, "report.txt") + listing = client.list_binding_files_sync("binding-ref", ".") + preview = client.read_binding_file_sync("binding-ref", "note.txt", max_bytes=128) + downloaded = client.download_binding_file_sync(_binding_file_download_request()) - assert isinstance(listing, SandboxListResponse) + assert isinstance(listing, BindingFileListResponse) assert listing.path == "." - assert isinstance(preview, SandboxReadResponse) + assert isinstance(preview, BindingFileReadResponse) assert preview.text == "hello" - assert isinstance(uploaded, SandboxUploadResponse) - assert uploaded.file.reference == "dify-file-ref:file-1" - assert uploaded.file.download_url == "https://files.example.com/report.txt" + assert isinstance(downloaded, BindingFileDownloadResponse) + assert downloaded.reference == "dify-file-ref:file-1" -def test_async_sandbox_methods_post_dtos_and_parse_responses() -> None: - locator = _sandbox_locator() - +def test_async_binding_file_methods_post_dtos_and_parse_responses() -> None: def handler(request: httpx.Request) -> httpx.Response: - if request.url.path == "/sandbox/files/list": + if request.url.path == "/execution-bindings/files/list": return httpx.Response(200, json={"path": ".", "entries": [], "truncated": False}) - if request.url.path == "/sandbox/files/read": + if request.url.path == "/execution-bindings/files/read": + payload = cast(dict[str, object], json.loads(request.content)) + assert payload["max_bytes"] == 262144 return httpx.Response( 200, json={"path": "note.txt", "size": 5, "truncated": False, "binary": False, "text": "hello"} ) - if request.url.path == "/sandbox/files/upload": - return httpx.Response( - 200, - json={ - "path": "report.txt", - "file": { - "transfer_method": "tool_file", - "reference": "dify-file-ref:file-1", - "download_url": "https://files.example.com/report.txt", - }, - }, - ) + if request.url.path == "/execution-bindings/files/download": + _assert_binding_download_timeout(request) + return httpx.Response(200, json={"reference": "dify-file-ref:file-1"}) raise AssertionError(f"unexpected request: {request.method} {request.url}") async def scenario() -> None: http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) client = Client(base_url="http://testserver", async_http_client=http_client) - listing = await client.list_sandbox_files(locator, ".") - preview = await client.read_sandbox_file(locator, "note.txt") - uploaded = await client.upload_sandbox_file(locator, "report.txt") + listing = await client.list_binding_files("binding-ref", ".") + preview = await client.read_binding_file("binding-ref", "note.txt") + downloaded = await client.download_binding_file(_binding_file_download_request()) assert listing.path == "." assert preview.text == "hello" - assert uploaded.file.reference == "dify-file-ref:file-1" - assert uploaded.file.download_url == "https://files.example.com/report.txt" + assert downloaded.reference == "dify-file-ref:file-1" await http_client.aclose() asyncio.run(scenario()) -def test_sync_upload_sandbox_file_rejects_missing_download_url() -> None: - locator = _sandbox_locator() - +def test_sync_execution_binding_client_uses_private_binding_routes() -> None: def handler(request: httpx.Request) -> httpx.Response: - if request.url.path != "/sandbox/files/upload": + payload = cast(dict[str, object], json.loads(request.content)) + if request.url.path == "/execution-bindings": + assert payload["binding_id"] == "binding-1" + assert payload["home_snapshot_ref"] == "home-ref" + return httpx.Response(201, json={"binding_ref": "backend-binding", "workspace_ref": "backend-workspace"}) + assert request.url.path == "/execution-bindings/destroy" + assert payload == { + "binding_ref": "backend-binding", + "destroy_workspace": True, + "workspace_ref": "backend-workspace", + } + return httpx.Response(204) + + client = Client( + base_url="http://testserver", + sync_http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) + allocation = client.create_execution_binding_sync( + CreateExecutionBindingRequest( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="home-ref", + ) + ) + client.destroy_execution_binding_sync( + DestroyExecutionBindingRequest( + binding_ref=allocation.binding_ref, + workspace_ref=allocation.workspace_ref, + destroy_workspace=True, + ) + ) + + assert allocation.binding_ref == "backend-binding" + + +def _create_home_snapshot_from_binding_request() -> CreateHomeSnapshotFromBindingRequest: + return CreateHomeSnapshotFromBindingRequest( + tenant_id="tenant-1", + agent_id="agent-1", + home_snapshot_id="home-2", + backend_binding_ref="binding-ref", + ) + + +def test_sync_home_snapshot_client_parses_checkpoint_and_delete() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.method == "POST": + if request.url.path == "/home-snapshots/from-binding": + assert json.loads(request.content) == _create_home_snapshot_from_binding_request().model_dump( + mode="json" + ) + return httpx.Response(201, json={"snapshot_ref": "team/home 1"}) + assert request.url.path == "/home-snapshots/delete" + assert json.loads(request.content) == {"snapshot_ref": "team/home 1"} + return httpx.Response(204) + raise AssertionError(request.url) + + http_client = httpx.Client(transport=httpx.MockTransport(handler)) + client = Client(base_url="http://testserver", sync_http_client=http_client) + + created = client.create_home_snapshot_from_binding_sync(_create_home_snapshot_from_binding_request()) + client.delete_home_snapshot_sync(created.snapshot_ref) + + assert created.snapshot_ref == "team/home 1" + http_client.close() + + +def test_async_home_snapshot_client_parses_checkpoint_and_delete() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.method == "POST": + if request.url.path == "/home-snapshots/from-binding": + return httpx.Response(201, json={"snapshot_ref": "team/home 1"}) + assert request.url.path == "/home-snapshots/delete" + return httpx.Response(204) + raise AssertionError(request.url) + + async def scenario() -> None: + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + client = Client(base_url="http://testserver", async_http_client=http_client) + + created = await client.create_home_snapshot_from_binding(_create_home_snapshot_from_binding_request()) + await client.delete_home_snapshot(created.snapshot_ref) + + assert created.snapshot_ref == "team/home 1" + await http_client.aclose() + + asyncio.run(scenario()) + + +def test_home_snapshot_client_maps_sync_validation_and_async_http_errors() -> None: + sync_http_client = httpx.Client(transport=httpx.MockTransport(lambda _request: httpx.Response(200, json={}))) + sync_client = Client( + base_url="http://testserver", + sync_http_client=sync_http_client, + ) + + with pytest.raises(DifyAgentValidationError): + _ = sync_client.create_home_snapshot_from_binding_sync(_create_home_snapshot_from_binding_request()) + sync_http_client.close() + + async def scenario() -> None: + http_client = httpx.AsyncClient( + transport=httpx.MockTransport( + lambda _request: httpx.Response(502, json={"detail": {"code": "backend_failed"}}) + ) + ) + client = Client(base_url="http://testserver", async_http_client=http_client) + + with pytest.raises(DifyAgentHTTPError) as exc_info: + _ = await client.create_home_snapshot_from_binding(_create_home_snapshot_from_binding_request()) + assert exc_info.value.status_code == 502 + assert exc_info.value.detail == {"code": "backend_failed"} + await http_client.aclose() + + asyncio.run(scenario()) + + +def test_sync_download_binding_file_rejects_missing_reference() -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path != "/execution-bindings/files/download": raise AssertionError(f"unexpected request: {request.method} {request.url}") return httpx.Response( 200, - json={ - "path": "report.txt", - "file": { - "transfer_method": "tool_file", - "reference": "dify-file-ref:file-1", - }, - }, + json={}, ) client = Client(base_url="http://testserver", sync_http_client=httpx.Client(transport=httpx.MockTransport(handler))) with pytest.raises(DifyAgentValidationError): - _ = client.upload_sandbox_file_sync(locator, "report.txt") + _ = client.download_binding_file_sync(_binding_file_download_request()) -def test_sync_sandbox_methods_map_invalid_json_to_validation_error() -> None: +def test_sync_binding_file_methods_map_invalid_json_to_validation_error() -> None: responses = iter([httpx.Response(200, text="not-json"), httpx.Response(404, json={"detail": "missing"})]) def handler(_request: httpx.Request) -> httpx.Response: @@ -366,10 +460,10 @@ def test_sync_sandbox_methods_map_invalid_json_to_validation_error() -> None: client = Client(base_url="http://testserver", sync_http_client=httpx.Client(transport=httpx.MockTransport(handler))) with pytest.raises(DifyAgentValidationError): - _ = client.list_sandbox_files_sync(_sandbox_locator(), ".") + _ = client.list_binding_files_sync("binding-ref", ".") with pytest.raises(DifyAgentHTTPError) as http_error: - _ = client.read_sandbox_file_sync(_sandbox_locator(), "missing.txt") + _ = client.read_binding_file_sync("binding-ref", "missing.txt") assert http_error.value.status_code == 404 @@ -528,6 +622,44 @@ def test_stream_events_stops_after_cancelled_terminal_event() -> None: assert calls == 1 +def test_stream_events_does_not_reconnect_after_terminal_when_until_terminal_is_false() -> None: + calls = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(200, content=_event_frame(_run_succeeded_event())) + + client = Client( + base_url="http://testserver", + sync_http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) + + events = list(client.stream_events_sync("run-1", until_terminal=False, reconnect_delay_seconds=0)) + + assert [event.type for event in events] == ["run_succeeded"] + assert calls == 1 + + +def test_stream_events_does_not_reconnect_after_terminal_transport_error() -> None: + calls = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(200, stream=DisconnectingSyncStream(_event_frame(_run_succeeded_event()))) + + client = Client( + base_url="http://testserver", + sync_http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) + + events = list(client.stream_events_sync("run-1", until_terminal=False, reconnect_delay_seconds=0)) + + assert [event.type for event in events] == ["run_succeeded"] + assert calls == 1 + + def test_stream_events_reconnects_from_latest_event_id() -> None: seen_after: list[str] = [] @@ -682,6 +814,101 @@ def test_async_stream_events_yields_terminal_event() -> None: asyncio.run(scenario()) +def test_async_stream_events_does_not_reconnect_after_terminal_when_until_terminal_is_false() -> None: + calls = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(200, content=_event_frame(_run_succeeded_event())) + + async def scenario() -> None: + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + client = Client(base_url="http://testserver", async_http_client=http_client) + + events = [event async for event in client.stream_events("run-1", until_terminal=False)] + + assert [event.type for event in events] == ["run_succeeded"] + assert calls == 1 + await http_client.aclose() + + asyncio.run(scenario()) + + +def test_async_stream_events_does_not_reconnect_after_terminal_transport_error() -> None: + calls = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(200, stream=DisconnectingAsyncStream(_event_frame(_run_succeeded_event()))) + + async def scenario() -> None: + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + client = Client(base_url="http://testserver", async_http_client=http_client) + + events = [event async for event in client.stream_events("run-1", until_terminal=False)] + + assert [event.type for event in events] == ["run_succeeded"] + assert calls == 1 + await http_client.aclose() + + asyncio.run(scenario()) + + +def test_async_stream_events_reconnects_from_latest_event_after_transport_error() -> None: + seen_after: list[str] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen_after.append(request.url.params["after"]) + if len(seen_after) == 1: + return httpx.Response( + 200, + stream=DisconnectingAsyncStream(_event_frame(RunStartedEvent(id="1-0", run_id="run-1"))), + ) + return httpx.Response(200, content=_event_frame(_run_succeeded_event())) + + async def scenario() -> None: + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + client = Client(base_url="http://testserver", async_http_client=http_client) + + events = [event async for event in client.stream_events("run-1", reconnect_delay_seconds=0)] + + assert seen_after == ["0-0", "1-0"] + assert [event.type for event in events] == ["run_started", "run_succeeded"] + await http_client.aclose() + + asyncio.run(scenario()) + + +def test_async_stream_events_reconnects_after_eof_before_terminal() -> None: + calls = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal calls + calls += 1 + return httpx.Response(200, content="") + + async def scenario() -> None: + http_client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + client = Client(base_url="http://testserver", async_http_client=http_client) + + with pytest.raises(DifyAgentStreamError, match="reconnect attempts exhausted"): + _ = [ + event + async for event in client.stream_events( + "run-1", + max_reconnects=1, + reconnect_delay_seconds=0, + ) + ] + + assert calls == 2 + await http_client.aclose() + + asyncio.run(scenario()) + + def test_async_sse_parser_preserves_unicode_line_separators() -> None: error = "next-line:\x85line-separator:\u2028paragraph-separator:\u2029done" body = _event_frame(_run_failed_event(error)) diff --git a/dify-agent/tests/local/dify_agent/layers/config/test_layer.py b/dify-agent/tests/local/dify_agent/layers/config/test_layer.py index f085f6cc51b..990a19eb477 100644 --- a/dify-agent/tests/local/dify_agent/layers/config/test_layer.py +++ b/dify-agent/tests/local/dify_agent/layers/config/test_layer.py @@ -2,12 +2,11 @@ from __future__ import annotations -import asyncio +import json from typing import Literal import pytest -from dify_agent.adapters.shell.shellctl import ShellctlProvider from dify_agent.layers.config import DifyConfigLayerConfig from dify_agent.layers.config.layer import ( DifyConfigLayer, @@ -18,14 +17,9 @@ from dify_agent.layers.shell import DifyShellLayerConfig from dify_agent.layers.shell.layer import CompleteRemoteCommandResult, DifyShellLayer -def _unused_client_factory(): - raise AssertionError("shellctl client should not be used by these config-layer tests") - - def _shell_layer() -> DifyShellLayer: return DifyShellLayer.from_config_with_settings( DifyShellLayerConfig(agent_stub_drive_ref="agent-1"), - shell_provider=ShellctlProvider(entrypoint="http://shellctl", token="", client_factory=_unused_client_factory), ) @@ -68,32 +62,40 @@ def _remote_result( ) -def _skill_pull_output(*, include_skill: bool = True) -> str: - if not include_skill: - return "" - return "/workspace/.dify_conf/skills/alpha\n# Alpha\nUse it.\n" +def _skill_pull_output(*names: str, include_skill: bool = True) -> str: + items = [] + if include_skill: + items = [ + { + "name": name, + "archive_path": f"/workspace/.dify_conf/skills/{name}.zip", + "directory_path": f"/workspace/.dify_conf/skills/{name}", + "skill_md": "# Alpha\nUse it.\n", + } + for name in names or ("alpha",) + ] + return json.dumps({"items": items}) -def _file_pull_output(*, include_file: bool = True) -> str: - if not include_file: - return "" - return "/workspace/.dify_conf/files/guide.txt\n" +def _file_pull_output(*names: str, include_file: bool = True) -> str: + items = [] + if include_file: + items = [{"name": name, "path": f"/workspace/.dify_conf/files/{name}"} for name in names or ("guide.txt",)] + return json.dumps({"items": items}) def test_build_shell_pull_scripts_include_targets() -> None: layer = _build_layer() - skill_script = layer._build_shell_skill_pull_script("alpha") - file_script = layer._build_shell_file_pull_script("guide.txt") + skill_script = layer._build_shell_skill_pull_script(["alpha", "skill with space"]) + file_script = layer._build_shell_file_pull_script(["guide.txt", "file with space.txt"]) - assert skill_script == "set -eu\ndify-agent config skills pull alpha" - assert "__DIFY_CONFIG_SKILLS_BEGIN__" not in skill_script - assert file_script == "set -eu\ndify-agent config files pull guide.txt" - assert "__DIFY_CONFIG_FILES_BEGIN__" not in file_script + assert skill_script == "set -eu\ndify-agent config skills pull --json alpha 'skill with space'" + assert file_script == "set -eu\ndify-agent config files pull --json guide.txt 'file with space.txt'" @pytest.mark.anyio -async def test_on_context_create_computes_runtime_fields_and_pulls_mentioned_assets_in_parallel( +async def test_on_context_create_computes_runtime_fields_and_pulls_mentioned_assets_in_batches( monkeypatch: pytest.MonkeyPatch, ) -> None: layer = _build_layer() @@ -110,7 +112,6 @@ async def test_on_context_create_computes_runtime_fields_and_pulls_mentioned_ass captured_scripts.append(script) active_commands += 1 max_active_commands = max(max_active_commands, active_commands) - await asyncio.sleep(0) active_commands -= 1 if "skills pull" in script: return _remote_result(_skill_pull_output()) @@ -122,11 +123,11 @@ async def test_on_context_create_computes_runtime_fields_and_pulls_mentioned_ass await layer.on_context_create() - assert max_active_commands > 1 + assert max_active_commands == 1 assert len(captured_scripts) == 2 - assert sorted(captured_scripts) == [ - "set -eu\ndify-agent config files pull guide.txt", - "set -eu\ndify-agent config skills pull alpha", + assert captured_scripts == [ + "set -eu\ndify-agent config skills pull --json alpha", + "set -eu\ndify-agent config files pull --json guide.txt", ] assert layer.runtime_state.pulled_skill_outputs == {"alpha": "/workspace/.dify_conf/skills/alpha\n# Alpha\nUse it."} assert layer.runtime_state.pulled_file_outputs == {"guide.txt": "/workspace/.dify_conf/files/guide.txt"} @@ -146,6 +147,39 @@ async def test_on_context_create_computes_runtime_fields_and_pulls_mentioned_ass assert _AGENT_FILE_UPLOAD_REPLY_HINT in suffix_prompt +@pytest.mark.anyio +async def test_on_context_create_batches_all_mentioned_assets_into_two_serial_jobs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + layer = _build_layer() + layer.config = layer.config.model_copy( + update={ + "mentioned_skill_names": [f"skill-{index}" for index in range(4)], + "mentioned_file_names": [f"file-{index}.txt" for index in range(4)], + } + ) + captured_scripts: list[str] = [] + + async def fake_run_remote_script(self, script: str, *, inject_agent_stub_env: bool = False, timeout: float = 10.0): + del self, timeout + assert inject_agent_stub_env is True + captured_scripts.append(script) + if "skills pull" in script: + return _remote_result(_skill_pull_output(*(f"skill-{index}" for index in range(4)))) + return _remote_result(_file_pull_output(*(f"file-{index}.txt" for index in range(4)))) + + monkeypatch.setattr(DifyShellLayer, "run_remote_script", fake_run_remote_script) + + await layer.on_context_create() + + assert captured_scripts == [ + "set -eu\ndify-agent config skills pull --json skill-0 skill-1 skill-2 skill-3", + "set -eu\ndify-agent config files pull --json file-0.txt file-1.txt file-2.txt file-3.txt", + ] + assert set(layer.runtime_state.pulled_skill_outputs) == {f"skill-{index}" for index in range(4)} + assert set(layer.runtime_state.pulled_file_outputs) == {f"file-{index}.txt" for index in range(4)} + + @pytest.mark.anyio async def test_on_context_resume_does_not_recompute_or_pull(monkeypatch: pytest.MonkeyPatch) -> None: layer = _build_layer() diff --git a/dify-agent/tests/local/dify_agent/layers/drive/test_layer.py b/dify-agent/tests/local/dify_agent/layers/drive/test_layer.py index afcef5cc396..c4cfb74346f 100644 --- a/dify-agent/tests/local/dify_agent/layers/drive/test_layer.py +++ b/dify-agent/tests/local/dify_agent/layers/drive/test_layer.py @@ -6,21 +6,15 @@ from typing import Literal import pytest -from dify_agent.adapters.shell.shellctl import ShellctlProvider from dify_agent.layers.drive import DifyDriveLayerConfig, DifyDriveSkillConfig from dify_agent.layers.drive.layer import DifyDriveLayer, DifyDriveLayerError, _AGENT_FILE_UPLOAD_REPLY_HINT from dify_agent.layers.shell import DifyShellLayerConfig from dify_agent.layers.shell.layer import CompleteRemoteCommandResult, DifyShellLayer -def _unused_client_factory(): - raise AssertionError("shellctl client should not be used by these drive-layer tests") - - def _shell_layer() -> DifyShellLayer: return DifyShellLayer.from_config_with_settings( DifyShellLayerConfig(agent_stub_drive_ref="agent-1"), - shell_provider=ShellctlProvider(entrypoint="http://shellctl", token="", client_factory=_unused_client_factory), ) diff --git a/dify-agent/tests/local/dify_agent/layers/shell/test_configs.py b/dify-agent/tests/local/dify_agent/layers/shell/test_configs.py index b6ad74965b9..30405baff66 100644 --- a/dify-agent/tests/local/dify_agent/layers/shell/test_configs.py +++ b/dify-agent/tests/local/dify_agent/layers/shell/test_configs.py @@ -7,7 +7,6 @@ from dify_agent.layers.shell import ( DifyShellCliToolConfig, DifyShellEnvVarConfig, DifyShellLayerConfig, - DifyShellSandboxConfig, DifyShellSecretRefConfig, ) @@ -18,7 +17,6 @@ def test_shell_package_exports_client_safe_config_symbols_only() -> None: "DifyShellCliToolConfig", "DifyShellEnvVarConfig", "DifyShellLayerConfig", - "DifyShellSandboxConfig", "DifyShellSecretRefConfig", ] assert DIFY_SHELL_LAYER_TYPE_ID == "dify.shell" @@ -33,7 +31,6 @@ def test_shell_layer_config_defaults_and_forbids_unknown_fields() -> None: "cli_tools": [], "env": [], "secret_refs": [], - "sandbox": None, "redact_patterns": [], } @@ -54,7 +51,6 @@ def test_shell_layer_config_accepts_agent_soul_shell_settings() -> None: env=[DifyShellEnvVarConfig(name="PROJECT_NAME", value="demo")], secret_refs=[DifyShellSecretRefConfig(name="OPENAI_API_KEY", ref="credential-1")], agent_stub_drive_ref="agent-1", - sandbox=DifyShellSandboxConfig(provider="independent", config={"cpu": 2}), ) assert config.cli_tools[0].install_commands == ["apt-get update", "apt-get install -y ripgrep"] @@ -63,8 +59,6 @@ def test_shell_layer_config_accepts_agent_soul_shell_settings() -> None: assert config.env[0].name == "PROJECT_NAME" assert config.secret_refs[0].ref == "credential-1" assert config.agent_stub_drive_ref == "agent-1" - assert config.sandbox is not None - assert config.sandbox.config == {"cpu": 2} def test_shell_layer_config_rejects_invalid_env_names() -> None: diff --git a/dify-agent/tests/local/dify_agent/layers/shell/test_layer.py b/dify-agent/tests/local/dify_agent/layers/shell/test_layer.py index 3d1e8f292cd..d8b07395774 100644 --- a/dify-agent/tests/local/dify_agent/layers/shell/test_layer.py +++ b/dify-agent/tests/local/dify_agent/layers/shell/test_layer.py @@ -4,29 +4,39 @@ import asyncio import json from collections.abc import Callable, Mapping from dataclasses import dataclass, field -from pathlib import Path from typing import cast import pytest import dify_agent.layers.shell.layer as shell_layer_module -from dify_agent.layers.shell import DIFY_SHELL_LAYER_TYPE_ID, DifyShellEnvVarConfig, DifyShellLayerConfig +import dify_agent.runtime.command_runner as command_runner_module +from dify_agent.layers.shell import ( + DIFY_SHELL_LAYER_TYPE_ID, + DifyShellCliToolConfig, + DifyShellEnvVarConfig, + DifyShellLayerConfig, +) from dify_agent.layers.shell.layer import ( CompleteRemoteCommandResult, DEFAULT_TERMINATE_GRACE_SECONDS, DifyShellLayer, + DifyShellLayerDeps, DifyShellRuntimeState, ) from dify_agent.adapters.shell.protocols import ( ShellCommandResult, ShellCommandStatus, - ShellFileTransferProtocol, ShellProviderError, - SandboxExpiredError, - ShellProviderProtocol, - ShellResourceProtocol, ) from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig +from dify_agent.layers.execution_context.layer import DifyExecutionContextLayer +from dify_agent.layers.runtime import DifyRuntimeLayerConfig +from dify_agent.layers.runtime.layer import DifyRuntimeLayer +from dify_agent.runtime_backend import ( + ExecutionBindingBackend, + RuntimeLayout, + RuntimeLease, +) def _command_result( @@ -142,14 +152,6 @@ class _UnexpectedToolError(Exception): pass -class FakeFiles(ShellFileTransferProtocol): - async def upload(self, *, content: bytes, remote_path: str, cwd: str | None = None) -> None: - raise AssertionError("resource.files should not be used by production shell layer logic") - - async def download(self, *, remote_path: str, cwd: str | None = None) -> bytes: - raise AssertionError("resource.files should not be used by production shell layer logic") - - @dataclass(slots=True) class FakeCommands: run_handler: Callable[[str, str | None, Mapping[str, str] | None, float], ShellCommandResult] | None = None @@ -208,41 +210,22 @@ class FakeCommands: @dataclass(slots=True) -class FakeResource(ShellResourceProtocol): +class FakeResource: commands: FakeCommands - files: FakeFiles = field(default_factory=FakeFiles) - suspended: bool = False - deleted: bool = False - _sandbox_id: str | None = None - - @property - def sandbox_id(self) -> str | None: - return self._sandbox_id - - async def suspend(self) -> None: - self.suspended = True - - async def delete(self) -> None: - self.deleted = True + handle: str = "sandbox-1" + layout: RuntimeLayout = field( + default_factory=lambda: RuntimeLayout( + home_dir="/home/agent-1", + workspace_dir="/home/agent-1/workspace/abc12ff", + ) + ) @dataclass(slots=True) -class FakeProvider(ShellProviderProtocol): +class FakeProvider: + """Test fixture retaining the old name while representing an active lease.""" + resource: FakeResource - create_calls: int = 0 - attach_calls: list[str] = field(default_factory=list) - attach_error: ShellProviderError | None = None - - async def create(self) -> ShellResourceProtocol: - self.create_calls += 1 - return self.resource - - async def attach(self, sandbox_id: str) -> ShellResourceProtocol: - self.attach_calls.append(sandbox_id) - if self.attach_error is not None: - raise self.attach_error - self.resource._sandbox_id = sandbox_id - return self.resource def _layer( @@ -252,22 +235,31 @@ def _layer( shell_home_root: str = "/home", ) -> tuple[DifyShellLayer, FakeProvider]: provider = FakeProvider(resource=FakeResource(commands=commands)) + root = shell_home_root.rstrip("/") + provider.resource.layout = RuntimeLayout( + home_dir=f"{root}/agent-1", + workspace_dir=f"{root}/agent-1/workspace/abc12ff", + ) layer = DifyShellLayer.from_config_with_settings( config or DifyShellLayerConfig(), - shell_provider=provider, - shell_home_root=shell_home_root, + ) + runtime = DifyRuntimeLayer.from_config_with_backend( + DifyRuntimeLayerConfig(backend_binding_ref="binding-1"), + backend=cast(ExecutionBindingBackend, object()), + ) + runtime._lease = cast(RuntimeLease, cast(object, provider.resource)) + layer.deps = DifyShellLayerDeps( + runtime=runtime, + execution_context=None, ) return layer, provider -@dataclass(slots=True) -class _ExecutionContextStub: - config: DifyExecutionContextLayerConfig - - def _bind_execution_context(layer: DifyShellLayer, *, agent_id: str | None = "agent-1") -> None: - layer.deps.execution_context = cast( - object, _ExecutionContextStub(config=_execution_context_config(agent_id=agent_id)) + layer.deps.execution_context = DifyExecutionContextLayer.from_config_with_settings( + _execution_context_config(agent_id=agent_id), + daemon_url="http://plugin-daemon", + daemon_api_key="", ) @@ -290,10 +282,8 @@ def _runtime_state( job_ids: list[str] | None = None, job_offsets: dict[str, int] | None = None, ) -> DifyShellRuntimeState: + del session_id, workspace_cwd, sandbox_id return DifyShellRuntimeState( - session_id=session_id, - workspace_cwd=workspace_cwd, - sandbox_id=sandbox_id, job_ids=[] if job_ids is None else job_ids, job_offsets={} if job_offsets is None else job_offsets, ) @@ -303,31 +293,12 @@ def test_shell_type_id_constant_matches_implementation_class() -> None: assert DIFY_SHELL_LAYER_TYPE_ID == DifyShellLayer.type_id -def test_resource_context_calls_provider_create_and_suspends_on_exit() -> None: - layer, provider = _layer(commands=FakeCommands()) - - async def scenario() -> None: - async with layer.resource_context(): - assert provider.create_calls == 1 - assert provider.resource.suspended is False - assert provider.resource.deleted is False - assert provider.resource.suspended is True - assert provider.resource.deleted is False - - asyncio.run(scenario()) - - -def test_shell_layer_create_allocates_workspace_and_bootstraps(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(shell_layer_module.time, "time", lambda: int("abc12", 16)) - monkeypatch.setattr(shell_layer_module.secrets, "token_hex", lambda _nbytes: "ff") +def test_shell_layer_create_bootstraps_inside_sandbox_workspace() -> None: expected_home = "/home/agent-1" expected_workspace_cwd = "/home/agent-1/workspace/abc12ff" def run_handler(script: str, cwd: str | None, env: Mapping[str, str] | None, timeout: float) -> ShellCommandResult: assert env == {"HOME": expected_home} - if cwd is None: - assert 'mkdir -p "$HOME/workspace"' in script - return _command_result("mkdir-job", status="exited", done=True, exit_code=0) assert cwd == expected_workspace_cwd assert "apt-get install -y ripgrep" in script return _command_result("bootstrap-job", status="exited", done=True, exit_code=0) @@ -335,7 +306,7 @@ def test_shell_layer_create_allocates_workspace_and_bootstraps(monkeypatch: pyte layer, provider = _layer( commands=FakeCommands(run_handler=run_handler), config=DifyShellLayerConfig( - cli_tools=[{"name": "ripgrep", "install_commands": ["apt-get install -y ripgrep"]}], + cli_tools=[DifyShellCliToolConfig(name="ripgrep", install_commands=["apt-get install -y ripgrep"])], ), ) _bind_execution_context(layer) @@ -343,31 +314,18 @@ def test_shell_layer_create_allocates_workspace_and_bootstraps(monkeypatch: pyte async def scenario() -> None: async with layer.resource_context(): await layer.on_context_create() - assert provider.resource.suspended is False asyncio.run(scenario()) - assert layer.runtime_state.session_id == "abc12ff" - assert layer.runtime_state.workspace_cwd == "~/workspace/abc12ff" - assert [call.job_id for call in provider.resource.commands.delete_calls] == ["mkdir-job", "bootstrap-job"] - assert provider.resource.suspended is True + assert [call.job_id for call in provider.resource.commands.delete_calls] == ["bootstrap-job"] -def test_shell_layer_uses_agent_specific_home_and_workspace_cwd( - monkeypatch: pytest.MonkeyPatch, -) -> None: - monkeypatch.setattr(shell_layer_module.time, "time", lambda: int("abc12", 16)) - monkeypatch.setattr(shell_layer_module.secrets, "token_hex", lambda _nbytes: "ff") - +def test_shell_layer_uses_sandbox_layout_for_home_and_workspace_cwd() -> None: expected_home = "/home/agent-1" expected_workspace_cwd = "/home/agent-1/workspace/abc12ff" def run_handler(script: str, cwd: str | None, env: Mapping[str, str] | None, timeout: float) -> ShellCommandResult: del timeout - if script.startswith('mkdir -p "$HOME/workspace";'): - assert cwd is None - assert env == {"HOME": expected_home} - return _command_result("mkdir-job", status="exited", done=True, exit_code=0) if script == "pwd": assert cwd == expected_workspace_cwd assert env == {"HOME": expected_home} @@ -381,7 +339,6 @@ def test_shell_layer_uses_agent_specific_home_and_workspace_cwd( async def scenario() -> None: async with layer.resource_context(): - await layer.on_context_create() run_result = await tools["shell_run"].function_schema.call({"script": "pwd"}, None) # pyright: ignore[reportArgumentType] metadata, output = _parse_tagged_observation(run_result) assert metadata["job_id"] == "user-job" @@ -389,75 +346,26 @@ def test_shell_layer_uses_agent_specific_home_and_workspace_cwd( asyncio.run(scenario()) - assert layer.runtime_state.session_id == "abc12ff" - assert layer.runtime_state.workspace_cwd == "~/workspace/abc12ff" assert layer.runtime_state.job_ids == ["user-job"] - assert [call.job_id for call in commands.delete_calls] == ["mkdir-job"] + assert commands.delete_calls == [] -def test_shell_layer_uses_configured_home_root_for_local_development( - monkeypatch: pytest.MonkeyPatch, - tmp_path: Path, -) -> None: - monkeypatch.setattr(shell_layer_module.time, "time", lambda: int("abc12", 16)) - monkeypatch.setattr(shell_layer_module.secrets, "token_hex", lambda _nbytes: "ff") - - shell_home_root = tmp_path / "shell-home" - expected_home = f"{shell_home_root}/agent-1" - expected_workspace_cwd = f"{expected_home}/workspace/abc12ff" - - def run_handler(script: str, cwd: str | None, env: Mapping[str, str] | None, timeout: float) -> ShellCommandResult: - del timeout - if script.startswith('mkdir -p "$HOME/workspace";'): - assert cwd is None - assert env == {"HOME": expected_home} - return _command_result("mkdir-job", status="exited", done=True, exit_code=0) - if script == "pwd": - assert cwd == expected_workspace_cwd - assert env == {"HOME": expected_home} - return _command_result("user-job", status="exited", done=True, exit_code=0, output=expected_home, offset=13) - raise AssertionError(f"Unexpected script: {script!r}") - - layer, _provider = _layer(commands=FakeCommands(run_handler=run_handler), shell_home_root=f"{shell_home_root}/") - _bind_execution_context(layer) - tools = {tool.name: tool for tool in layer.tools} - - async def scenario() -> None: - async with layer.resource_context(): - await layer.on_context_create() - await tools["shell_run"].function_schema.call({"script": "pwd"}, None) # pyright: ignore[reportArgumentType] - - asyncio.run(scenario()) - - assert layer.shell_home_root == str(shell_home_root) - assert layer.runtime_state.workspace_cwd == "~/workspace/abc12ff" - - -def test_shell_layer_suspend_does_not_close_before_resource_context_exits() -> None: - layer, provider = _layer(commands=FakeCommands()) +def test_shell_layer_suspend_cleans_tracked_jobs_without_owning_sandbox() -> None: + commands = FakeCommands() + layer, _provider = _layer(commands=commands) layer.runtime_state = _runtime_state() async def scenario() -> None: async with layer.resource_context(): await layer.on_context_suspend() - assert provider.resource.suspended is False - assert provider.resource.suspended is True asyncio.run(scenario()) + assert commands.delete_calls == [] -def test_shell_layer_resume_recreates_live_home_and_workspace() -> None: - expected_home = "/home/agent-1" - - def run_handler(script: str, cwd: str | None, env: Mapping[str, str] | None, timeout: float) -> ShellCommandResult: - del timeout - assert script == 'mkdir -p "$HOME/workspace/abc12ff"' - assert cwd is None - assert env == {"HOME": expected_home} - return _command_result("resume-job", status="exited", done=True, exit_code=0) - - commands = FakeCommands(run_handler=run_handler) - layer, provider = _layer(commands=commands) +def test_shell_layer_resume_requires_active_sandbox_lease_only() -> None: + commands = FakeCommands() + layer, _provider = _layer(commands=commands) _bind_execution_context(layer) layer.runtime_state = _runtime_state() @@ -467,34 +375,12 @@ def test_shell_layer_resume_recreates_live_home_and_workspace() -> None: asyncio.run(scenario()) - assert [call.job_id for call in commands.delete_calls] == ["resume-job"] - assert provider.resource.suspended is True - - -def test_shell_layer_resume_requires_agent_id_for_live_workspace_reentry() -> None: - commands = FakeCommands() - layer, _provider = _layer(commands=commands) - _bind_execution_context(layer, agent_id=None) - layer.runtime_state = _runtime_state() - - async def scenario() -> None: - async with layer.resource_context(): - with pytest.raises(ValueError, match="requires execution_context\\.agent_id"): - await layer.on_context_resume() - - asyncio.run(scenario()) assert commands.run_calls == [] -def test_shell_layer_delete_cleans_workspace_and_tracked_jobs() -> None: - def run_handler(script: str, cwd: str | None, env: Mapping[str, str] | None, timeout: float) -> ShellCommandResult: - del env, timeout - assert cwd is None - assert script == 'rm -rf -- "$HOME/workspace/abc12ff"' - return _command_result("cleanup-job", status="exited", done=True, exit_code=0) - - commands = FakeCommands(run_handler=run_handler) - layer, provider = _layer(commands=commands) +def test_shell_layer_delete_cleans_tracked_jobs_without_deleting_workspace() -> None: + commands = FakeCommands() + layer, _provider = _layer(commands=commands) _bind_execution_context(layer) layer.runtime_state = _runtime_state(job_ids=["user-job"], job_offsets={"user-job": 9}) @@ -504,10 +390,9 @@ def test_shell_layer_delete_cleans_workspace_and_tracked_jobs() -> None: asyncio.run(scenario()) - assert [call.job_id for call in commands.delete_calls] == ["cleanup-job", "user-job"] + assert [call.job_id for call in commands.delete_calls] == ["user-job"] assert layer.runtime_state.job_ids == [] assert layer.runtime_state.job_offsets == {} - assert provider.resource.deleted is True def test_shell_layer_tools_map_inputs_and_maintain_offsets_with_tail_end() -> None: @@ -559,7 +444,7 @@ def test_shell_layer_tools_map_inputs_and_maintain_offsets_with_tail_end() -> No tail_handler=tail_handler, interrupt_handler=interrupt_handler, ) - layer, provider = _layer(commands=commands) + layer, _provider = _layer(commands=commands) _bind_execution_context(layer) tools = {tool.name: tool for tool in layer.tools} layer.runtime_state = _runtime_state() @@ -611,7 +496,6 @@ def test_shell_layer_tools_map_inputs_and_maintain_offsets_with_tail_end() -> No assert layer.runtime_state.job_offsets == {"user-job": 34} assert commands.tail_calls == [TailCall(job_id="user-job"), TailCall(job_id="user-job")] - assert provider.resource.suspended is True def test_shell_run_keeps_original_offset_when_tail_lookup_fails_for_truncated_output() -> None: @@ -633,7 +517,7 @@ def test_shell_run_keeps_original_offset_when_tail_lookup_fails_for_truncated_ou raise RuntimeError(f"tail unavailable for {job_id}") commands = FakeCommands(run_handler=run_handler, tail_handler=tail_handler) - layer, provider = _layer(commands=commands) + layer, _provider = _layer(commands=commands) _bind_execution_context(layer) tools = {tool.name: tool for tool in layer.tools} layer.runtime_state = _runtime_state() @@ -656,7 +540,6 @@ def test_shell_run_keeps_original_offset_when_tail_lookup_fails_for_truncated_ou assert layer.runtime_state.job_offsets == {"user-job": 10} assert commands.tail_calls == [TailCall(job_id="user-job")] - assert provider.resource.suspended is True def test_shell_run_formats_large_non_truncated_output_without_tail_lookup() -> None: @@ -683,6 +566,7 @@ def test_shell_run_formats_large_non_truncated_output_without_tail_lookup() -> N metadata, output = _parse_tagged_observation(result) assert metadata["output_path"] == "/tmp/large.log" assert output.startswith("head-y") + assert "max output size is limited to 8192 bytes" in output assert output.endswith("(check the /tmp/large.log for full output)") assert "-tail" in output @@ -831,7 +715,10 @@ def test_shell_run_returns_error_observation_and_logs_unexpected_exception( asyncio.run(scenario()) assert logged == [ - ("Unexpected shell tool failure: tool=%s session_id=%s job_id=%s", ("shell_run", "abc12ff", None)) + ( + "Unexpected shell tool failure: tool=%s job_id=%s", + ("shell_run", None), + ) ] @@ -856,7 +743,10 @@ def test_shell_wait_returns_error_observation_and_logs_unexpected_exception( asyncio.run(scenario()) assert logged == [ - ("Unexpected shell tool failure: tool=%s session_id=%s job_id=%s", ("shell_wait", "abc12ff", "user-job")) + ( + "Unexpected shell tool failure: tool=%s job_id=%s", + ("shell_wait", "user-job"), + ) ] @@ -884,8 +774,8 @@ def test_shell_input_returns_error_observation_and_logs_unexpected_exception( asyncio.run(scenario()) assert logged == [ ( - "Unexpected shell tool failure: tool=%s session_id=%s job_id=%s", - ("shell_input", "abc12ff", "user-job"), + "Unexpected shell tool failure: tool=%s job_id=%s", + ("shell_input", "user-job"), ) ] @@ -914,8 +804,8 @@ def test_shell_interrupt_returns_error_observation_and_logs_unexpected_exception asyncio.run(scenario()) assert logged == [ ( - "Unexpected shell tool failure: tool=%s session_id=%s job_id=%s", - ("shell_interrupt", "abc12ff", "user-job"), + "Unexpected shell tool failure: tool=%s job_id=%s", + ("shell_interrupt", "user-job"), ) ] @@ -947,12 +837,48 @@ def test_shell_interrupt_logs_unexpected_tail_failure_but_still_succeeds( asyncio.run(scenario()) assert logged == [ ( - "Failed to fetch output path for interrupted shell job %s in session %s", - ("user-job", "abc12ff"), + "Failed to fetch output path for interrupted shell job %s", + ("user-job",), ) ] +def test_agent_stub_token_omits_workspace_session_identity() -> None: + seen_session_ids: list[str | None] = [] + + def token_factory( + execution_context: DifyExecutionContextLayerConfig, + *, + session_id: str | None, + ) -> str: + del execution_context + seen_session_ids.append(session_id) + return "stub-token" + + def run_handler( + _script: str, + _cwd: str | None, + env: Mapping[str, str] | None, + _timeout: float, + ) -> ShellCommandResult: + assert env is not None + assert env["DIFY_AGENT_STUB_AUTH_JWE"] == "stub-token" + return _command_result("remote-job", status="exited", done=True, exit_code=0) + + layer, _provider = _layer(commands=FakeCommands(run_handler=run_handler)) + _bind_execution_context(layer) + layer.agent_stub_api_base_url = "http://localhost:5050/agent-stub" + layer.agent_stub_token_factory = token_factory + + async def scenario() -> None: + async with layer.resource_context(): + _ = await layer.run_remote_script_complete("true", inject_agent_stub_env=True) + + asyncio.run(scenario()) + + assert seen_session_ids == [None] + + def test_shell_run_propagates_cancelled_error() -> None: commands = FakeCommands( run_handler=lambda script, cwd, env, timeout: (_ for _ in ()).throw(asyncio.CancelledError()) @@ -1151,7 +1077,7 @@ def test_run_remote_script_complete_returns_incomplete_reason_for_timeout( def fake_monotonic() -> float: return now - monkeypatch.setattr(shell_layer_module.time, "monotonic", fake_monotonic) + monkeypatch.setattr(command_runner_module.time, "monotonic", fake_monotonic) def run_handler(script: str, cwd: str | None, env: Mapping[str, str] | None, timeout: float) -> ShellCommandResult: nonlocal now @@ -1217,59 +1143,22 @@ def test_shell_layer_rejects_untracked_job_ids_without_provider_calls() -> None: assert commands.interrupt_calls == [] -def test_shell_layer_requires_agent_id_for_live_command_execution(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(shell_layer_module.time, "time", lambda: int("abc12", 16)) - monkeypatch.setattr(shell_layer_module.secrets, "token_hex", lambda _nbytes: "ff") - commands = FakeCommands() - layer, _provider = _layer(commands=commands) - _bind_execution_context(layer, agent_id=None) - - async def scenario() -> None: - async with layer.resource_context(): - with pytest.raises(ValueError, match="requires execution_context\\.agent_id"): - await layer.on_context_create() - - asyncio.run(scenario()) - assert commands.run_calls == [] - - def test_shell_layer_hooks_and_tools_fail_clearly_outside_active_resource_context() -> None: layer, _provider = _layer(commands=FakeCommands()) + layer.deps.runtime._lease = None tools = {tool.name: tool for tool in layer.tools} layer.runtime_state = _runtime_state() async def scenario() -> None: result = await tools["shell_run"].function_schema.call({"script": "pwd"}, None) # pyright: ignore[reportArgumentType] - _assert_error_observation(result, includes="shell resource") + _assert_error_observation(result, includes="resource_context") asyncio.run(scenario()) -def test_shell_runtime_state_validates_workspace_identity_and_offset_keys() -> None: - with pytest.raises(ValueError, match="5\\+2 lowercase hex format"): - _ = DifyShellRuntimeState.model_validate( - { - "session_id": "../../tmp", - "workspace_cwd": "~/workspace/../../tmp", - "job_ids": [], - "job_offsets": {}, - } - ) - - with pytest.raises(ValueError, match="workspace_cwd must equal"): - _ = DifyShellRuntimeState.model_validate( - { - "session_id": "abc12ff", - "workspace_cwd": "~/workspace/def34aa", - "job_ids": [], - "job_offsets": {}, - } - ) - +def test_shell_runtime_state_validates_offset_keys() -> None: state = DifyShellRuntimeState.model_validate( { - "session_id": "abc12ff", - "workspace_cwd": "~/workspace/abc12ff", "job_ids": ['job"bad with spaces'], "job_offsets": {'job"bad with spaces': 0}, } @@ -1278,108 +1167,12 @@ def test_shell_runtime_state_validates_workspace_identity_and_offset_keys() -> N with pytest.raises(ValueError, match="unknown job ids"): _ = DifyShellRuntimeState.model_validate( { - "session_id": "abc12ff", - "workspace_cwd": "~/workspace/abc12ff", "job_ids": ["job-1"], "job_offsets": {"job-2": 3}, } ) -def test_resource_context_attaches_when_sandbox_id_is_present() -> None: - layer, provider = _layer(commands=FakeCommands()) - layer.runtime_state = _runtime_state(sandbox_id="existing-sandbox-1") - - async def scenario() -> None: - async with layer.resource_context(): - assert provider.create_calls == 0 - assert provider.attach_calls == ["existing-sandbox-1"] - assert provider.resource.suspended is False - assert provider.resource.suspended is True - - asyncio.run(scenario()) - - -def test_resource_context_deletes_on_context_delete() -> None: - commands = FakeCommands( - run_handler=lambda script, cwd, env, timeout: _command_result( - "cleanup-job", status="exited", done=True, exit_code=0 - ) - ) - layer, provider = _layer(commands=commands) - _bind_execution_context(layer) - layer.runtime_state = _runtime_state() - - async def scenario() -> None: - async with layer.resource_context(): - await layer.on_context_delete() - assert provider.resource.deleted is True - assert provider.resource.suspended is False - - asyncio.run(scenario()) - - -def test_resource_context_suspends_on_context_suspend() -> None: - layer, provider = _layer(commands=FakeCommands()) - layer.runtime_state = _runtime_state() - - async def scenario() -> None: - async with layer.resource_context(): - await layer.on_context_suspend() - assert provider.resource.suspended is True - assert provider.resource.deleted is False - - asyncio.run(scenario()) - - -def test_resource_context_persists_sandbox_id_from_provider() -> None: - layer, provider = _layer(commands=FakeCommands()) - provider.resource._sandbox_id = "new-sandbox-42" - - async def scenario() -> None: - async with layer.resource_context(): - pass - - asyncio.run(scenario()) - assert layer.runtime_state.sandbox_id == "new-sandbox-42" - - -def test_resource_context_propagates_sandbox_expired() -> None: - """When attach() raises SandboxExpiredError, the error propagates to the caller. - The user must start a new session — no in-place recovery is attempted.""" - layer, provider = _layer(commands=FakeCommands()) - provider.attach_error = SandboxExpiredError( - "stale-sandbox-1", - cause=ShellProviderError( - 'request_failed (404): {"reason":"NOT_FOUND","message":"error: code = 400 reason = sandbox_expired message = sandbox has expired"}', - code="request_failed", - ), - ) - layer.runtime_state = _runtime_state(sandbox_id="stale-sandbox-1") - - async def scenario() -> None: - with pytest.raises(SandboxExpiredError): - async with layer.resource_context(): - pass - - asyncio.run(scenario()) - assert provider.attach_calls == ["stale-sandbox-1"] - assert provider.create_calls == 0 - - -def test_resource_context_reraises_non_expired_attach_error() -> None: - layer, provider = _layer(commands=FakeCommands()) - provider.attach_error = ShellProviderError("some other error", code="request_failed") - layer.runtime_state = _runtime_state(sandbox_id="sandbox-1") - - async def scenario() -> None: - async with layer.resource_context(): - pass - - with pytest.raises(ShellProviderError, match="some other error"): - asyncio.run(scenario()) - - # --------------------------------------------------------------------------- # Output redaction tests # --------------------------------------------------------------------------- @@ -1393,15 +1186,10 @@ def _layer_with_redaction( token_value: str = "eyJhbGciOiJkaXIiLCJlbmMiOiJBMjU2R0NNIn0.fake-long-jwe-token-value", ) -> tuple[DifyShellLayer, FakeProvider]: """Create a layer with agent_stub env injection and optional redaction patterns.""" - provider = FakeProvider(resource=FakeResource(commands=commands)) - layer = DifyShellLayer.from_config_with_settings( - config or DifyShellLayerConfig(), - shell_provider=provider, - shell_home_root="/home", - shell_redact_patterns=shell_redact_patterns, - agent_stub_api_base_url="http://localhost:5050/agent-stub", - agent_stub_token_factory=lambda execution_context, session_id: token_value, - ) + layer, provider = _layer(commands=commands, config=config) + layer.shell_redact_patterns = shell_redact_patterns or [] + layer.agent_stub_api_base_url = "http://localhost:5050/agent-stub" + layer.agent_stub_token_factory = lambda execution_context, session_id: token_value return layer, provider diff --git a/dify-agent/tests/local/dify_agent/layers/test_runtime_layer.py b/dify-agent/tests/local/dify_agent/layers/test_runtime_layer.py new file mode 100644 index 00000000000..de6ed027b7a --- /dev/null +++ b/dify-agent/tests/local/dify_agent/layers/test_runtime_layer.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import cast + +import pytest + +from dify_agent.layers.runtime import DifyRuntimeLayer, DifyRuntimeLayerConfig +from dify_agent.runtime_backend import RuntimeLayout, RuntimeLease + + +@dataclass(slots=True) +class _Backend: + lease: RuntimeLease + acquired: list[str] = field(default_factory=list) + released: list[RuntimeLease] = field(default_factory=list) + + async def acquire(self, binding_ref: str) -> RuntimeLease: + self.acquired.append(binding_ref) + return self.lease + + async def release(self, lease: RuntimeLease) -> None: + self.released.append(lease) + + +@pytest.mark.anyio +async def test_runtime_layer_acquires_and_releases_operation_scoped_lease() -> None: + lease = cast( + RuntimeLease, + cast( + object, + type( + "Lease", + (), + { + "layout": RuntimeLayout(home_dir="/home/agent", workspace_dir="/workspace"), + "commands": object(), + }, + )(), + ), + ) + backend = _Backend(lease=lease) + layer = DifyRuntimeLayer.from_config_with_backend( + DifyRuntimeLayerConfig(backend_binding_ref="binding-1"), + backend=backend, # pyright: ignore[reportArgumentType] + ) + + async with layer.resource_context(): + assert layer.lease is lease + await layer.on_context_create() + await layer.on_context_suspend() + await layer.on_context_delete() + + assert backend.acquired == ["binding-1"] + assert backend.released == [lease] + with pytest.raises(RuntimeError, match="resource_context"): + _ = layer.lease + + +def test_runtime_layer_config_contains_only_backend_binding_ref() -> None: + config = DifyRuntimeLayerConfig(backend_binding_ref="binding-1") + + assert config.model_dump() == {"backend_binding_ref": "binding-1"} diff --git a/dify-agent/tests/local/dify_agent/protocol/test_protocol_schemas.py b/dify-agent/tests/local/dify_agent/protocol/test_protocol_schemas.py index 9aa8cf8af0f..5bc3915ffad 100644 --- a/dify-agent/tests/local/dify_agent/protocol/test_protocol_schemas.py +++ b/dify-agent/tests/local/dify_agent/protocol/test_protocol_schemas.py @@ -25,6 +25,7 @@ from dify_agent.protocol.schemas import ( RunComposition, RunFailedEvent, RunFailedEventData, + RunFailureType, RunLayerSpec, RunStartedEvent, RunSucceededEvent, @@ -68,7 +69,14 @@ def test_run_event_adapter_round_trips_typed_variants() -> None: session_snapshot=CompositorSessionSnapshot(layers=[]), ), ), - RunFailedEvent(run_id="run-1", data=RunFailedEventData(error="boom", reason="shutdown")), + RunFailedEvent( + run_id="run-1", + data=RunFailedEventData( + error="boom", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + reason="shutdown", + ), + ), RunCancelledEvent(run_id="run-1", data=RunCancelledEventData(reason="user_cancelled")), ] @@ -80,6 +88,31 @@ def test_run_event_adapter_round_trips_typed_variants() -> None: assert decoded.run_id == event.run_id +def test_run_failed_event_error_type_is_optional_and_round_trips() -> None: + legacy = RUN_EVENT_ADAPTER.validate_python( + { + "run_id": "legacy-run", + "type": "run_failed", + "data": {"error": "legacy failure", "reason": None}, + } + ) + classified = RunFailedEvent( + run_id="classified-run", + data=RunFailedEventData( + error="run limit reached", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + ), + ) + + decoded = RUN_EVENT_ADAPTER.validate_json(RUN_EVENT_ADAPTER.dump_json(classified)) + + assert isinstance(legacy, RunFailedEvent) + assert legacy.data.error_type is None + assert isinstance(decoded, RunFailedEvent) + assert decoded.data.error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + assert protocol_exports.RunFailureType is RunFailureType + + def test_pydantic_ai_event_data_uses_agent_stream_event_model() -> None: event = RUN_EVENT_ADAPTER.validate_python( { diff --git a/dify-agent/tests/local/dify_agent/protocol/test_sandbox_locator.py b/dify-agent/tests/local/dify_agent/protocol/test_sandbox_locator.py deleted file mode 100644 index 6de05187708..00000000000 --- a/dify-agent/tests/local/dify_agent/protocol/test_sandbox_locator.py +++ /dev/null @@ -1,178 +0,0 @@ -from __future__ import annotations - -import pytest -from agenton.compositor import CompositorSessionSnapshot -from agenton.compositor.schemas import LayerSessionSnapshot -from agenton.layers.base import LifecycleState -from agenton_collections.layers.plain import PromptLayerConfig -from dify_agent.layers.dify_core_tools import DIFY_CORE_TOOLS_LAYER_TYPE_ID -from dify_agent.layers.dify_plugin import DIFY_PLUGIN_LLM_LAYER_TYPE_ID -from dify_agent.layers.execution_context import DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, DifyExecutionContextLayerConfig -from dify_agent.layers.shell import DIFY_SHELL_LAYER_TYPE_ID, DifyShellLayerConfig -from dify_agent.protocol import ( - CreateRunRequest, - RunComposition, - RunLayerSpec, - RuntimeLayerSpec, - build_sandbox_locator_from_layer_specs, - build_sandbox_locator_from_run_request, - extract_runtime_layer_specs, -) - - -def _request() -> CreateRunRequest: - composition = RunComposition( - layers=[ - RunLayerSpec(name="prompt", type="plain.prompt", config=PromptLayerConfig(prefix="hi")), - RunLayerSpec( - name="execution_context", - type=DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, - config=DifyExecutionContextLayerConfig( - tenant_id="tenant-1", - user_from="account", - agent_mode="workflow_run", - invoke_from="service-api", - ), - ), - RunLayerSpec(name="llm", type=DIFY_PLUGIN_LLM_LAYER_TYPE_ID), - RunLayerSpec( - name="shell", - type=DIFY_SHELL_LAYER_TYPE_ID, - deps={"execution_context": "execution_context"}, - config=DifyShellLayerConfig(), - ), - ] - ) - snapshot = CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot(name="prompt", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={}), - LayerSessionSnapshot( - name="execution_context", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ), - LayerSessionSnapshot(name="llm", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={}), - LayerSessionSnapshot( - name="shell", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={"session_id": "abc12ff", "workspace_cwd": "~/workspace/abc12ff"}, - ), - ] - ) - return CreateRunRequest(composition=composition, session_snapshot=snapshot) - - -def test_build_sandbox_locator_from_run_request_filters_to_execution_context_and_shell() -> None: - locator = build_sandbox_locator_from_run_request(_request()) - - assert [layer.name for layer in locator.composition.layers] == ["execution_context", "shell"] - assert [layer.type for layer in locator.composition.layers] == [ - DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, - DIFY_SHELL_LAYER_TYPE_ID, - ] - assert [layer.name for layer in locator.session_snapshot.layers] == ["execution_context", "shell"] - - -def test_build_sandbox_locator_from_run_request_rejects_missing_session_snapshot() -> None: - request = _request() - request.session_snapshot = None - - with pytest.raises(ValueError, match="session_snapshot"): - build_sandbox_locator_from_run_request(request) - - -def test_extract_runtime_layer_specs_drops_sensitive_plugin_layers() -> None: - specs = extract_runtime_layer_specs(_request().composition) - - assert [spec.name for spec in specs] == ["prompt", "execution_context", "shell"] - - -def test_build_sandbox_locator_from_layer_specs_rejects_missing_shell() -> None: - with pytest.raises(ValueError, match="shell"): - build_sandbox_locator_from_layer_specs( - layer_specs=[RuntimeLayerSpec(name="execution_context", type=DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID)], - session_snapshot=CompositorSessionSnapshot(layers=[]), - ) - - -def test_build_sandbox_locator_from_layer_specs_rejects_missing_snapshot_layer() -> None: - specs = [ - RuntimeLayerSpec(name="execution_context", type=DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID), - RuntimeLayerSpec(name="shell", type=DIFY_SHELL_LAYER_TYPE_ID, deps={"execution_context": "execution_context"}), - ] - - with pytest.raises(ValueError, match="session_snapshot"): - build_sandbox_locator_from_layer_specs( - layer_specs=specs, - session_snapshot=CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot( - name="execution_context", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={} - ) - ] - ), - ) - - -def test_build_sandbox_locator_from_layer_specs_rejects_shell_dep_mismatch() -> None: - specs = [ - RuntimeLayerSpec(name="execution_context", type=DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID), - RuntimeLayerSpec(name="shell", type=DIFY_SHELL_LAYER_TYPE_ID, deps={"execution_context": "wrong-layer"}), - ] - - with pytest.raises(ValueError, match="depend on the execution_context"): - build_sandbox_locator_from_layer_specs( - layer_specs=specs, - session_snapshot=CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot( - name="execution_context", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ), - LayerSessionSnapshot( - name="shell", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ), - ] - ), - ) - - -def test_build_sandbox_locator_from_layer_specs_rejects_order_mismatch() -> None: - specs = [ - RuntimeLayerSpec(name="execution_context", type=DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID), - RuntimeLayerSpec(name="shell", type=DIFY_SHELL_LAYER_TYPE_ID, deps={"execution_context": "execution_context"}), - ] - - with pytest.raises(ValueError, match="in order"): - build_sandbox_locator_from_layer_specs( - layer_specs=specs, - session_snapshot=CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot(name="shell", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={}), - LayerSessionSnapshot( - name="execution_context", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={}, - ), - ] - ), - ) - - -def test_build_sandbox_locator_from_layer_specs_rejects_sensitive_runtime_specs() -> None: - with pytest.raises(ValueError, match="sensitive"): - build_sandbox_locator_from_layer_specs( - layer_specs=[RuntimeLayerSpec(name="llm", type=DIFY_PLUGIN_LLM_LAYER_TYPE_ID)], - session_snapshot=CompositorSessionSnapshot(layers=[]), - ) - - -def test_build_sandbox_locator_from_layer_specs_rejects_sensitive_core_tool_runtime_specs() -> None: - with pytest.raises(ValueError, match="sensitive"): - build_sandbox_locator_from_layer_specs( - layer_specs=[RuntimeLayerSpec(name="core_tools", type=DIFY_CORE_TOOLS_LAYER_TYPE_ID)], - session_snapshot=CompositorSessionSnapshot(layers=[]), - ) diff --git a/dify-agent/tests/local/dify_agent/protocol/test_working_environment.py b/dify-agent/tests/local/dify_agent/protocol/test_working_environment.py new file mode 100644 index 00000000000..cdb217c1e29 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/protocol/test_working_environment.py @@ -0,0 +1,90 @@ +import pytest +from pydantic import ValidationError + +from dify_agent.protocol import ( + BindingFileListRequest, + BindingFileReadRequest, + CreateExecutionBindingRequest, + CreateHomeSnapshotFromBindingRequest, + DestroyExecutionBindingRequest, +) + + +def test_execution_binding_request_uses_opaque_backend_refs() -> None: + request = CreateExecutionBindingRequest( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref="opaque-workspace", + home_snapshot_ref="opaque-home", + ) + + assert request.model_dump() == { + "tenant_id": "tenant-1", + "agent_id": "agent-1", + "binding_id": "binding-1", + "workspace_id": "workspace-1", + "existing_workspace_ref": "opaque-workspace", + "home_snapshot_ref": "opaque-home", + } + + +def test_execution_binding_request_accepts_missing_or_null_home_snapshot_ref() -> None: + fields = { + "tenant_id": "tenant-1", + "agent_id": "agent-1", + "binding_id": "binding-1", + "workspace_id": "workspace-1", + } + + assert CreateExecutionBindingRequest(**fields).home_snapshot_ref is None + assert CreateExecutionBindingRequest(**fields, home_snapshot_ref=None).home_snapshot_ref is None + + +def test_execution_binding_request_rejects_empty_home_snapshot_ref() -> None: + with pytest.raises(ValidationError, match="home_snapshot_ref"): + CreateExecutionBindingRequest( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + home_snapshot_ref="", + ) + + +def test_destroy_workspace_requires_workspace_ref() -> None: + with pytest.raises(ValidationError, match="workspace_ref"): + DestroyExecutionBindingRequest(binding_ref="binding-1", destroy_workspace=True) + + +def test_snapshot_and_file_requests_locate_binding_directly() -> None: + snapshot = CreateHomeSnapshotFromBindingRequest( + tenant_id="tenant-1", + agent_id="agent-1", + home_snapshot_id="home-2", + backend_binding_ref="binding-ref", + ) + listing = BindingFileListRequest(backend_binding_ref="binding-ref", path="~/files") + + assert snapshot.backend_binding_ref == "binding-ref" + assert listing.path == "~/files" + + +def test_binding_file_read_preview_uses_bounded_default_and_limit() -> None: + assert BindingFileReadRequest(backend_binding_ref="binding-ref", path="report.txt").max_bytes == 262144 + assert ( + BindingFileReadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + max_bytes=262144, + ).max_bytes + == 262144 + ) + + with pytest.raises(ValidationError, match="max_bytes"): + BindingFileReadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + max_bytes=262145, + ) diff --git a/dify-agent/tests/local/dify_agent/runtime/test_command_runner.py b/dify-agent/tests/local/dify_agent/runtime/test_command_runner.py new file mode 100644 index 00000000000..2a8e9f6b658 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/runtime/test_command_runner.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +import asyncio +from contextlib import suppress +from dataclasses import dataclass, field + +import pytest + +from dify_agent.adapters.shell.protocols import ShellCommandResult +from dify_agent.runtime.command_runner import execute_complete_with_commands + + +@dataclass(slots=True) +class _BlockingCommands: + wait_started: asyncio.Event = field(default_factory=asyncio.Event) + wait_forever: asyncio.Event = field(default_factory=asyncio.Event) + deletes: list[tuple[str, bool]] = field(default_factory=list) + + async def run(self, script: str, *, cwd: str | None, env: dict[str, str] | None, timeout: float): + assert script == "long-running" + assert cwd == "/workspace" + assert env == {"HOME": "/home/agent"} + assert timeout > 0 + return ShellCommandResult( + job_id="job-1", + status="running", + done=False, + exit_code=None, + output="started", + offset=7, + truncated=False, + ) + + async def wait(self, job_id: str, *, offset: int, timeout: float): + assert (job_id, offset) == ("job-1", 7) + assert timeout > 0 + self.wait_started.set() + await self.wait_forever.wait() + raise AssertionError("wait must remain blocked until cancellation") + + async def read_output(self, job_id: str, *, offset: int): + raise AssertionError("unexpected read_output") + + async def input(self, job_id: str, text: str, *, offset: int, timeout: float): + raise AssertionError("unexpected input") + + async def interrupt(self, job_id: str, *, grace_seconds: float): + raise AssertionError("unexpected interrupt") + + async def tail(self, job_id: str): + raise AssertionError("unexpected tail") + + async def delete(self, job_id: str, *, force: bool = False, grace_seconds: float | None = None) -> None: + assert grace_seconds is None + self.deletes.append((job_id, force)) + + +@pytest.mark.anyio +async def test_cancellation_deletes_job_returned_before_blocking_wait() -> None: + commands = _BlockingCommands() + task = asyncio.create_task( + execute_complete_with_commands( + commands, # pyright: ignore[reportArgumentType] + "long-running", + cwd="/workspace", + env={"HOME": "/home/agent"}, + timeout=60.0, + max_output_bytes=4096, + ) + ) + try: + await asyncio.wait_for(commands.wait_started.wait(), timeout=1) + + task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=1) + finally: + if not task.done(): + task.cancel() + with suppress(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=1) + + assert commands.deletes == [("job-1", True)] diff --git a/dify-agent/tests/local/dify_agent/runtime/test_compositor_factory.py b/dify-agent/tests/local/dify_agent/runtime/test_compositor_factory.py index e0ce4e7a50f..11b40b20bc7 100644 --- a/dify-agent/tests/local/dify_agent/runtime/test_compositor_factory.py +++ b/dify-agent/tests/local/dify_agent/runtime/test_compositor_factory.py @@ -2,15 +2,6 @@ import sys import types from typing import cast -from pydantic import BaseModel - - -if "pydantic_settings" not in sys.modules: - pydantic_settings = types.ModuleType("pydantic_settings") - pydantic_settings.BaseSettings = BaseModel - pydantic_settings.SettingsConfigDict = dict - sys.modules["pydantic_settings"] = pydantic_settings - if "graphon.model_runtime.entities.llm_entities" not in sys.modules: graphon_module = types.ModuleType("graphon") model_runtime_module = types.ModuleType("graphon.model_runtime") @@ -84,13 +75,15 @@ if "jsonschema" not in sys.modules: sys.modules["jsonschema.protocols"] = jsonschema_protocols_module sys.modules["jsonschema.validators"] = jsonschema_validators_module -from dify_agent.adapters.shell.protocols import ShellProviderProtocol from dify_agent.layers.dify_core_tools import DIFY_CORE_TOOLS_LAYER_TYPE_ID, DifyCoreToolsLayerConfig from dify_agent.layers.dify_core_tools.layer import DifyCoreToolsLayer +from dify_agent.layers.runtime import DIFY_RUNTIME_LAYER_TYPE_ID, DifyRuntimeLayerConfig +from dify_agent.layers.runtime.layer import DifyRuntimeLayer from dify_agent.layers.shell import DIFY_SHELL_LAYER_TYPE_ID, DifyShellLayerConfig from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig from dify_agent.layers.shell.layer import DifyShellLayer from dify_agent.runtime.compositor_factory import create_default_layer_providers +from dify_agent.runtime_backend import ExecutionBindingBackend, HomeSnapshotBackend, RuntimeBackendProfile class FakeProvider: @@ -100,19 +93,28 @@ class FakeProvider: raise AssertionError("create should not be called by these tests") -def test_default_layer_providers_wire_provided_shell_provider() -> None: - fake_provider = FakeProvider() +def _runtime_backend_profile() -> RuntimeBackendProfile: + return RuntimeBackendProfile( + home_snapshots=cast(HomeSnapshotBackend, FakeProvider()), + execution_bindings=cast(ExecutionBindingBackend, FakeProvider()), + ) + + +def test_default_layer_providers_register_runtime_layer() -> None: + profile = _runtime_backend_profile() providers = create_default_layer_providers( - shell_provider=cast(ShellProviderProtocol, fake_provider), - shell_home_root="/tmp/dify-agent-home", + runtime_backend_profile=profile, ) shell_provider = next(provider for provider in providers if provider.type_id == DIFY_SHELL_LAYER_TYPE_ID) shell_layer = shell_provider.create_layer(DifyShellLayerConfig()) + runtime_provider = next(provider for provider in providers if provider.type_id == DIFY_RUNTIME_LAYER_TYPE_ID) + runtime_layer = runtime_provider.create_layer(DifyRuntimeLayerConfig(backend_binding_ref="binding-1")) assert isinstance(shell_layer, DifyShellLayer) - assert shell_layer.shell_provider is fake_provider - assert shell_layer.shell_home_root == "/tmp/dify-agent-home" + assert isinstance(runtime_layer, DifyRuntimeLayer) + assert runtime_layer.backend is profile.execution_bindings + assert {provider.type_id for provider in providers} >= {"dify.runtime", "dify.shell"} def test_default_layer_providers_forward_agent_stub_token_factory() -> None: @@ -127,7 +129,7 @@ def test_default_layer_providers_forward_agent_stub_token_factory() -> None: return f"token-for:{execution_context.tenant_id}:{session_id}" providers = create_default_layer_providers( - shell_provider=cast(ShellProviderProtocol, FakeProvider()), + runtime_backend_profile=_runtime_backend_profile(), agent_stub_api_base_url="https://agent.example.com/agent-stub", agent_stub_token_factory=build_agent_stub_token, ) diff --git a/dify-agent/tests/local/dify_agent/runtime/test_run_scheduler.py b/dify-agent/tests/local/dify_agent/runtime/test_run_scheduler.py index ac72d0502f1..a22ce101d2b 100644 --- a/dify-agent/tests/local/dify_agent/runtime/test_run_scheduler.py +++ b/dify-agent/tests/local/dify_agent/runtime/test_run_scheduler.py @@ -8,8 +8,10 @@ import pytest from agenton.compositor import CompositorSessionSnapshot, LayerSessionSnapshot from agenton.layers import LifecycleState from agenton_collections.layers.plain import PromptLayerConfig +from dify_agent.layers.dify_plugin import DifyPluginLLMLayerConfig +from dify_agent.layers.execution_context import DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, DifyExecutionContextLayerConfig from dify_agent.layers.output import DIFY_OUTPUT_LAYER_TYPE_ID, DifyOutputLayerConfig -from dify_agent.protocol import DIFY_AGENT_OUTPUT_LAYER_ID +from dify_agent.protocol import DIFY_AGENT_MODEL_LAYER_ID, DIFY_AGENT_OUTPUT_LAYER_ID, RunFailureType from dify_agent.protocol.schemas import ( CancelRunRequest, CreateRunRequest, @@ -18,6 +20,14 @@ from dify_agent.protocol.schemas import ( RunLayerSpec, RunStatus, ) +from dify_agent.runtime.event_sink import ( + NonTerminalRunEvent, + RunFinalizationResult, + TerminalRunEvent, + emit_run_failed, + emit_run_succeeded, + terminal_event_status_fields, +) from dify_agent.runtime.run_scheduler import RunCancellationConflictError, RunScheduler, SchedulerStoppingError from dify_agent.server.schemas import RunRecord @@ -27,7 +37,30 @@ def _request( *, output_config: Mapping[str, object] | DifyOutputLayerConfig | None = None, ) -> CreateRunRequest: - layers = [RunLayerSpec(name="prompt", type="plain.prompt", config=PromptLayerConfig(user=user))] + layers = [ + RunLayerSpec(name="prompt", type="plain.prompt", config=PromptLayerConfig(user=user)), + RunLayerSpec( + name="execution_context", + type=DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, + config=DifyExecutionContextLayerConfig( + tenant_id="tenant-1", + user_from="account", + agent_mode="workflow_run", + invoke_from="service-api", + ), + ), + RunLayerSpec( + name=DIFY_AGENT_MODEL_LAYER_ID, + type="dify.plugin.llm", + deps={"execution_context": "execution_context"}, + config=DifyPluginLLMLayerConfig( + plugin_id="langgenius/openai", + model_provider="openai", + model="demo-model", + credentials={"api_key": "secret"}, + ), + ), + ] if output_config is not None: layers.append( RunLayerSpec( @@ -60,33 +93,57 @@ class FakeStore: events: dict[str, list[RunEvent]] statuses: dict[str, RunStatus] errors: dict[str, str | None] + error_types: dict[str, RunFailureType | None] + terminal_changes: dict[str, asyncio.Event] def __init__(self) -> None: self.records = {} self.events = defaultdict(list) self.statuses = {} self.errors = {} + self.error_types = {} + self.terminal_changes = {} async def create_run(self) -> RunRecord: run_id = f"run-{len(self.records) + 1}" record = RunRecord(run_id=run_id, status="running") self.records[run_id] = record self.statuses[run_id] = "running" + self.terminal_changes[run_id] = asyncio.Event() return record - async def append_event(self, event: RunEvent) -> str: + async def append_event(self, event: NonTerminalRunEvent) -> str: event_id = str(len(self.events[event.run_id]) + 1) self.events[event.run_id].append(event.model_copy(update={"id": event_id})) return event_id async def get_run(self, run_id: str) -> RunRecord: return self.records[run_id].model_copy( - update={"status": self.statuses[run_id], "error": self.errors.get(run_id)}, + update={ + "status": self.statuses[run_id], + "error": self.errors.get(run_id), + "error_type": self.error_types.get(run_id), + }, ) - async def update_status(self, run_id: str, status: RunStatus, error: str | None = None) -> None: - self.statuses[run_id] = status - self.errors[run_id] = error + async def finalize_run(self, event: TerminalRunEvent) -> RunFinalizationResult: + current_status = self.statuses[event.run_id] + if current_status != "running": + return RunFinalizationResult(applied=False, status=current_status) + + status, error, error_type = terminal_event_status_fields(event) + event_id = str(len(self.events[event.run_id]) + 1) + self.events[event.run_id].append(event.model_copy(update={"id": event_id})) + self.statuses[event.run_id] = status + self.errors[event.run_id] = error + self.error_types[event.run_id] = error_type + self.terminal_changes[event.run_id].set() + return RunFinalizationResult(applied=True, status=status, event_id=event_id) + + async def wait_for_cancellation(self, run_id: str) -> bool: + while self.statuses[run_id] == "running": + await self.terminal_changes[run_id].wait() + return self.statuses[run_id] == "cancelled" class SlowCreateStore(FakeStore): @@ -104,17 +161,69 @@ class SlowCreateStore(FakeStore): return await super().create_run() +class TrackingStore(FakeStore): + observer_started: asyncio.Event + observer_finished: asyncio.Event + release_observer: asyncio.Event + + def __init__(self, *, pause_observer: bool = False) -> None: + super().__init__() + self.observer_started = asyncio.Event() + self.observer_finished = asyncio.Event() + self.release_observer = asyncio.Event() + if not pause_observer: + self.release_observer.set() + + async def wait_for_cancellation(self, run_id: str) -> bool: + self.observer_started.set() + try: + await self.release_observer.wait() + return await super().wait_for_cancellation(run_id) + finally: + self.observer_finished.set() + + +class FailingObserverStore(FakeStore): + fail_observer: asyncio.Event + observer_finished: asyncio.Event + + def __init__(self, *, fail_observer: asyncio.Event) -> None: + super().__init__() + self.fail_observer = fail_observer + self.observer_finished = asyncio.Event() + + async def wait_for_cancellation(self, run_id: str) -> bool: + del run_id + try: + await self.fail_observer.wait() + raise RuntimeError("redis read failed") + finally: + self.observer_finished.set() + + class ControlledRunner: started: asyncio.Event release: asyncio.Event + finished: asyncio.Event | None - def __init__(self, *, started: asyncio.Event, release: asyncio.Event) -> None: + def __init__( + self, + *, + started: asyncio.Event, + release: asyncio.Event, + finished: asyncio.Event | None = None, + ) -> None: self.started = started self.release = release + self.finished = finished async def run(self) -> None: _ = self.started.set() - await self.release.wait() + try: + await self.release.wait() + finally: + if self.finished is not None: + self.finished.set() class SwallowOneCancellationRunner: @@ -134,6 +243,145 @@ class SwallowOneCancellationRunner: await asyncio.Event().wait() +class SuccessThenWaitRunner: + def __init__( + self, + *, + store: FakeStore, + run_id: str, + finalized: asyncio.Event, + release: asyncio.Event, + ) -> None: + self.store = store + self.run_id = run_id + self.finalized = finalized + self.release = release + + async def run(self) -> None: + result = await emit_run_succeeded( + self.store, + run_id=self.run_id, + output="done", + session_snapshot=CompositorSessionSnapshot(layers=[]), + ) + assert result.applied is True + self.finalized.set() + await self.release.wait() + + +class IgnoreCancellationThenSucceedRunner: + def __init__( + self, + *, + store: FakeStore, + run_id: str, + started: asyncio.Event, + release: asyncio.Event, + finished: asyncio.Event, + ) -> None: + self.store = store + self.run_id = run_id + self.started = started + self.release = release + self.finished = finished + + async def run(self) -> None: + try: + self.started.set() + while not self.release.is_set(): + try: + await self.release.wait() + except asyncio.CancelledError: + continue + result = await emit_run_succeeded( + self.store, + run_id=self.run_id, + output="late success", + session_snapshot=CompositorSessionSnapshot(layers=[]), + ) + assert result.applied is False + assert result.status == "cancelled" + finally: + self.finished.set() + + +class ReleaseThenSucceedRunner: + def __init__( + self, + *, + store: FakeStore, + run_id: str, + started: asyncio.Event, + release: asyncio.Event, + finished: asyncio.Event, + ) -> None: + self.store = store + self.run_id = run_id + self.started = started + self.release = release + self.finished = finished + + async def run(self) -> None: + self.started.set() + try: + await self.release.wait() + result = await emit_run_succeeded( + self.store, + run_id=self.run_id, + output="done", + session_snapshot=CompositorSessionSnapshot(layers=[]), + ) + assert result.applied is True + finally: + self.finished.set() + + +class CompetingFailureRunner: + def __init__( + self, + *, + store: FakeStore, + run_id: str, + started: asyncio.Event, + release: asyncio.Event, + failure_attempted: asyncio.Event, + ) -> None: + self.store = store + self.run_id = run_id + self.started = started + self.release = release + self.failure_attempted = failure_attempted + + async def run(self) -> None: + self.started.set() + try: + await self.release.wait() + except asyncio.CancelledError: + pass + _ = await emit_run_failed(self.store, run_id=self.run_id, error="runner failed", reason="model_error") + self.failure_attempted.set() + + +class FinalizeSuccessOnCancellationRunner: + def __init__(self, *, store: FakeStore, run_id: str, started: asyncio.Event) -> None: + self.store = store + self.run_id = run_id + self.started = started + + async def run(self) -> None: + self.started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + result = await emit_run_succeeded( + self.store, + run_id=self.run_id, + output="completed during shutdown", + session_snapshot=CompositorSessionSnapshot(layers=[]), + ) + assert result.applied is True + + def test_create_run_starts_background_task_and_returns_running() -> None: async def scenario() -> None: store = FakeStore() @@ -186,39 +434,91 @@ def test_shutdown_marks_unfinished_runs_failed_and_appends_event() -> None: asyncio.run(scenario()) -def test_cancel_run_stops_task_and_persists_cancelled_terminal() -> None: +def test_cancellation_observer_failure_stops_runner_and_finalizes_failed() -> None: async def scenario() -> None: - store = FakeStore() - started = asyncio.Event() + fail_observer = asyncio.Event() + store = FailingObserverStore(fail_observer=fail_observer) + runner_started = asyncio.Event() + runner_finished = asyncio.Event() async with httpx.AsyncClient() as client: scheduler = RunScheduler( store=store, plugin_daemon_http_client=client, dify_api_http_client=client, - runner_factory=lambda _record, _request: ControlledRunner(started=started, release=asyncio.Event()), + runner_factory=lambda _record, _request: ControlledRunner( + started=runner_started, + release=asyncio.Event(), + finished=runner_finished, + ), ) record = await scheduler.create_run(_request()) - await asyncio.wait_for(started.wait(), timeout=1) + supervisor_task = scheduler.active_tasks[record.run_id] + await asyncio.wait_for(runner_started.wait(), timeout=1) - response = await scheduler.cancel_run( + fail_observer.set() + await asyncio.wait_for(supervisor_task, timeout=1) + + assert store.statuses[record.run_id] == "failed" + assert store.errors[record.run_id] == "run cancellation observer failed: redis read failed" + assert [event.type for event in store.events[record.run_id]] == ["run_failed"] + assert runner_finished.is_set() + assert store.observer_finished.is_set() + await asyncio.sleep(0) + assert scheduler.active_tasks == {} + + asyncio.run(scenario()) + + +def test_non_owner_cancel_run_stops_owner_task_and_persists_cancelled_terminal() -> None: + async def scenario() -> None: + store = TrackingStore() + started = asyncio.Event() + runner_finished = asyncio.Event() + async with httpx.AsyncClient() as client: + owner_scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + runner_factory=lambda _record, _request: ControlledRunner( + started=started, + release=asyncio.Event(), + finished=runner_finished, + ), + ) + remote_scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + ) + record = await owner_scheduler.create_run(_request()) + owner_task = owner_scheduler.active_tasks[record.run_id] + await asyncio.wait_for(started.wait(), timeout=1) + await asyncio.wait_for(store.observer_started.wait(), timeout=1) + + response = await remote_scheduler.cancel_run( record.run_id, CancelRunRequest(reason="workflow_aborted", message="outer workflow stopped"), ) assert response.status == "cancelled" - assert scheduler.active_tasks == {} + assert remote_scheduler.active_tasks == {} assert store.statuses[record.run_id] == "cancelled" assert store.errors[record.run_id] == "outer workflow stopped" assert [event.type for event in store.events[record.run_id]] == ["run_cancelled"] + await asyncio.wait_for(owner_task, timeout=1) + assert runner_finished.is_set() + assert store.observer_finished.is_set() + await asyncio.sleep(0) + assert owner_scheduler.active_tasks == {} - repeated = await scheduler.cancel_run(record.run_id, CancelRunRequest(reason="duplicate")) + repeated = await remote_scheduler.cancel_run(record.run_id, CancelRunRequest(reason="duplicate")) assert repeated.status == "cancelled" assert [event.type for event in store.events[record.run_id]] == ["run_cancelled"] asyncio.run(scenario()) -def test_cancel_run_reinjects_cancellation_without_waiting_for_runner_cleanup() -> None: +def test_owner_observer_reinjects_cancellation_consumed_by_runner() -> None: async def scenario() -> None: store = FakeStore() started = asyncio.Event() @@ -234,6 +534,7 @@ def test_cancel_run_reinjects_cancellation_without_waiting_for_runner_cleanup() ), ) record = await scheduler.create_run(_request()) + supervisor_task = scheduler.active_tasks[record.run_id] await asyncio.wait_for(started.wait(), timeout=1) response = await asyncio.wait_for( @@ -242,22 +543,228 @@ def test_cancel_run_reinjects_cancellation_without_waiting_for_runner_cleanup() ) assert response.status == "cancelled" - assert first_cancellation.is_set() assert store.statuses[record.run_id] == "cancelled" assert [event.type for event in store.events[record.run_id]] == ["run_cancelled"] + await asyncio.wait_for(first_cancellation.wait(), timeout=1) + await asyncio.wait_for(supervisor_task, timeout=1) await asyncio.sleep(0) assert scheduler.active_tasks == {} asyncio.run(scenario()) +def test_cancel_run_does_not_override_successful_terminal() -> None: + async def scenario() -> None: + store = FakeStore() + finalized = asyncio.Event() + release = asyncio.Event() + async with httpx.AsyncClient() as client: + scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + runner_factory=lambda record, _request: SuccessThenWaitRunner( + store=store, + run_id=record.run_id, + finalized=finalized, + release=release, + ), + ) + record = await scheduler.create_run(_request()) + await asyncio.wait_for(finalized.wait(), timeout=1) + task = scheduler.active_tasks[record.run_id] + + with pytest.raises(RunCancellationConflictError, match="already finished with status 'succeeded'"): + await scheduler.cancel_run(record.run_id, CancelRunRequest(reason="late_cancel")) + + assert task.done() is False + assert store.statuses[record.run_id] == "succeeded" + assert [event.type for event in store.events[record.run_id]] == ["run_succeeded"] + release.set() + await asyncio.wait_for(task, timeout=1) + + asyncio.run(scenario()) + + +def test_cancelled_terminal_survives_shutdown_while_runner_cleanup_is_pending() -> None: + async def scenario() -> None: + store = TrackingStore() + started = asyncio.Event() + release = asyncio.Event() + runner_finished = asyncio.Event() + async with httpx.AsyncClient() as client: + scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + shutdown_grace_seconds=0, + runner_factory=lambda record, _request: IgnoreCancellationThenSucceedRunner( + store=store, + run_id=record.run_id, + started=started, + release=release, + finished=runner_finished, + ), + ) + record = await scheduler.create_run(_request()) + supervisor_task = scheduler.active_tasks[record.run_id] + await asyncio.wait_for(started.wait(), timeout=1) + await asyncio.wait_for(store.observer_started.wait(), timeout=1) + + response = await scheduler.cancel_run(record.run_id, CancelRunRequest(reason="workflow_aborted")) + + assert response.status == "cancelled" + assert store.statuses[record.run_id] == "cancelled" + assert [event.type for event in store.events[record.run_id]] == ["run_cancelled"] + await asyncio.wait_for(store.observer_finished.wait(), timeout=1) + assert supervisor_task.done() is False + shutdown_task = asyncio.create_task(scheduler.shutdown()) + await asyncio.sleep(0) + assert shutdown_task.done() is False + release.set() + await asyncio.wait_for(shutdown_task, timeout=1) + + assert supervisor_task.done() + assert runner_finished.is_set() + assert store.observer_finished.is_set() + assert scheduler.active_tasks == {} + assert store.statuses[record.run_id] == "cancelled" + assert [event.type for event in store.events[record.run_id]] == ["run_cancelled"] + + asyncio.run(scenario()) + + +@pytest.mark.parametrize( + ("winner", "expected_event_type"), + [ + pytest.param("failed", "run_failed", id="failure-first"), + pytest.param("cancelled", "run_cancelled", id="cancellation-first"), + ], +) +def test_failure_and_cancellation_keep_the_first_terminal( + winner: RunStatus, + expected_event_type: str, +) -> None: + async def scenario() -> None: + store = FakeStore() + runner_started = asyncio.Event() + release_runner = asyncio.Event() + failure_attempted = asyncio.Event() + async with httpx.AsyncClient() as client: + scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + runner_factory=lambda record, _request: CompetingFailureRunner( + store=store, + run_id=record.run_id, + started=runner_started, + release=release_runner, + failure_attempted=failure_attempted, + ), + ) + record = await scheduler.create_run(_request()) + supervisor_task = scheduler.active_tasks[record.run_id] + await asyncio.wait_for(runner_started.wait(), timeout=1) + + if winner == "failed": + release_runner.set() + await asyncio.wait_for(failure_attempted.wait(), timeout=1) + with pytest.raises(RunCancellationConflictError, match="already finished with status 'failed'"): + await scheduler.cancel_run(record.run_id, CancelRunRequest(reason="late_cancel")) + else: + response = await scheduler.cancel_run( + record.run_id, + CancelRunRequest(reason="cancel_before_failure"), + ) + assert response.run_id == record.run_id + assert response.status == "cancelled" + release_runner.set() + + await asyncio.wait_for(failure_attempted.wait(), timeout=1) + await asyncio.wait_for(supervisor_task, timeout=1) + + assert store.statuses[record.run_id] == winner + assert [event.type for event in store.events[record.run_id]] == [expected_event_type] + + asyncio.run(scenario()) + + +def test_shutdown_grace_allows_runner_first_completion_and_reaps_children() -> None: + async def scenario() -> None: + store = TrackingStore(pause_observer=True) + runner_started = asyncio.Event() + release_runner = asyncio.Event() + runner_finished = asyncio.Event() + async with httpx.AsyncClient() as client: + scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + shutdown_grace_seconds=1, + runner_factory=lambda record, _request: ReleaseThenSucceedRunner( + store=store, + run_id=record.run_id, + started=runner_started, + release=release_runner, + finished=runner_finished, + ), + ) + record = await scheduler.create_run(_request()) + supervisor_task = scheduler.active_tasks[record.run_id] + await asyncio.wait_for(runner_started.wait(), timeout=1) + await asyncio.wait_for(store.observer_started.wait(), timeout=1) + + shutdown_task = asyncio.create_task(scheduler.shutdown()) + await asyncio.sleep(0) + assert shutdown_task.done() is False + release_runner.set() + await asyncio.wait_for(shutdown_task, timeout=1) + + assert supervisor_task.done() + assert runner_finished.is_set() + assert store.observer_finished.is_set() + assert scheduler.active_tasks == {} + assert store.statuses[record.run_id] == "succeeded" + assert [event.type for event in store.events[record.run_id]] == ["run_succeeded"] + + asyncio.run(scenario()) + + +def test_shutdown_does_not_append_failed_after_success_wins() -> None: + async def scenario() -> None: + store = FakeStore() + started = asyncio.Event() + async with httpx.AsyncClient() as client: + scheduler = RunScheduler( + store=store, + plugin_daemon_http_client=client, + dify_api_http_client=client, + shutdown_grace_seconds=0, + runner_factory=lambda record, _request: FinalizeSuccessOnCancellationRunner( + store=store, + run_id=record.run_id, + started=started, + ), + ) + record = await scheduler.create_run(_request()) + await asyncio.wait_for(started.wait(), timeout=1) + + await scheduler.shutdown() + + assert store.statuses[record.run_id] == "succeeded" + assert [event.type for event in store.events[record.run_id]] == ["run_succeeded"] + + asyncio.run(scenario()) + + def test_cancel_run_rejects_finished_run() -> None: async def scenario() -> None: store = FakeStore() async with httpx.AsyncClient() as client: scheduler = RunScheduler(store=store, plugin_daemon_http_client=client, dify_api_http_client=client) record = await store.create_run() - await store.update_status(record.run_id, "succeeded") + store.statuses[record.run_id] = "succeeded" with pytest.raises(RunCancellationConflictError, match="already finished"): await scheduler.cancel_run(record.run_id, CancelRunRequest()) @@ -339,7 +846,17 @@ def test_create_run_accepts_closed_session_snapshot_and_runner_fails_asynchronou name="prompt", lifecycle_state=LifecycleState.CLOSED, runtime_state={}, - ) + ), + LayerSessionSnapshot( + name="execution_context", + lifecycle_state=LifecycleState.SUSPENDED, + runtime_state={}, + ), + LayerSessionSnapshot( + name=DIFY_AGENT_MODEL_LAYER_ID, + lifecycle_state=LifecycleState.SUSPENDED, + runtime_state={}, + ), ] ) diff --git a/dify-agent/tests/local/dify_agent/runtime/test_runner.py b/dify-agent/tests/local/dify_agent/runtime/test_runner.py index dceab28ab98..e2fd7a4099c 100644 --- a/dify-agent/tests/local/dify_agent/runtime/test_runner.py +++ b/dify-agent/tests/local/dify_agent/runtime/test_runner.py @@ -8,7 +8,7 @@ import pytest from graphon.model_runtime.entities.llm_entities import LLMUsage from pydantic import JsonValue from pydantic_ai import Tool -from pydantic_ai.exceptions import ModelHTTPError, UnexpectedModelBehavior +from pydantic_ai.exceptions import ModelHTTPError, UnexpectedModelBehavior, UsageLimitExceeded from pydantic_ai.messages import ( ToolReturnPart, ModelMessage, @@ -22,6 +22,7 @@ from pydantic_ai.messages import ( from pydantic_ai.models import ModelRequestParameters from pydantic_ai.models.test import TestModel from pydantic_ai.tools import DeferredToolRequests, DeferredToolResults +from pydantic_ai.usage import UsageLimits from pydantic_ai.settings import ModelSettings from agenton.compositor import CompositorSessionSnapshot, LayerProvider, LayerSessionSnapshot @@ -30,9 +31,8 @@ from agenton_collections.layers.pydantic_ai import PYDANTIC_AI_HISTORY_LAYER_TYP from agenton_collections.layers.plain import PromptLayerConfig, ToolsLayer from dify_agent.layers.ask_human import DIFY_ASK_HUMAN_LAYER_TYPE_ID, DifyAskHumanLayerConfig from dify_agent.layers.execution_context import DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, DifyExecutionContextLayerConfig +from dify_agent.layers.runtime import DIFY_RUNTIME_LAYER_TYPE_ID, DifyRuntimeLayerConfig from dify_agent.layers.shell import DIFY_SHELL_LAYER_TYPE_ID, DifyShellLayerConfig -from dify_agent.adapters.shell.shellctl import ShellctlProvider -from dify_agent.layers.shell.layer import DifyShellLayer from dify_agent.layers.dify_plugin.configs import ( DIFY_PLUGIN_TOOLS_LAYER_TYPE_ID, DifyPluginLLMLayerConfig, @@ -61,10 +61,12 @@ from dify_agent.protocol.schemas import ( LayerExitSignals, PydanticAIStreamRunEvent, RunComposition, + RunFailedEvent, + RunFailureType, RunLayerSpec, RunSucceededEvent, ) -from dify_agent.runtime.event_sink import InMemoryRunEventSink +from dify_agent.runtime.event_sink import InMemoryRunEventSink, emit_run_cancelled from dify_agent.runtime.compositor_factory import create_default_layer_providers from dify_agent.runtime.runner import ( AgentRunRunner, @@ -72,6 +74,16 @@ from dify_agent.runtime.runner import ( RunSuccessOutcome, _run_failed_error_payload, ) +from dify_agent.runtime_backend import ( + ExecutionBindingAllocation, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotBackend, + RuntimeBackendProfile, + RuntimeLayout, + RuntimeLease, +) +from dify_agent.runtime_backend.shellctl import ShellctlRuntimeLease, create_shellctl_lease from shellctl.shared import DeleteJobResponse, JobResult, JobStatusName, JobStatusView @@ -135,6 +147,33 @@ class FakeRunnerShellctlClient: return DeleteJobResponse(job_id=job_id) +class FakeRunnerExecutionBindingBackend: + def __init__(self, client: FakeRunnerShellctlClient) -> None: + self.client = client + + async def create_binding(self, spec: ExecutionBindingCreateSpec) -> ExecutionBindingAllocation: + return ExecutionBindingAllocation(binding_ref=spec.binding_id, workspace_ref=spec.workspace_id) + + async def acquire(self, binding_ref: str) -> RuntimeLease: + return create_shellctl_lease( + handle=binding_ref, + layout=RuntimeLayout( + home_dir="/home/agent-1", + workspace_dir="/home/agent-1/workspace/abc12ff", + ), + entrypoint="http://shellctl", + token="", + client_factory=lambda: self.client, # pyright: ignore[reportArgumentType] + ) + + async def release(self, lease: RuntimeLease) -> None: + assert isinstance(lease, ShellctlRuntimeLease) + await lease.close() + + async def destroy_binding(self, spec: ExecutionBindingDestroySpec) -> None: + del spec + + def test_run_failed_error_payload_preserves_plugin_rate_limit_error() -> None: exc = ModelHTTPError( 429, @@ -142,18 +181,20 @@ def test_run_failed_error_payload_preserves_plugin_rate_limit_error() -> None: {"error_type": "InvokeRateLimitError", "message": "quota exceeded"}, ) - message, reason = _run_failed_error_payload(exc) + message, error_type, reason = _run_failed_error_payload(exc) assert message == "quota exceeded" + assert error_type is None assert reason == "InvokeRateLimitError" def test_run_failed_error_payload_infers_rate_limit_reason_from_status_code() -> None: exc = ModelHTTPError(429, "gpt-4o-mini", {"message": "too many requests"}) - message, reason = _run_failed_error_payload(exc) + message, error_type, reason = _run_failed_error_payload(exc) assert message == "too many requests" + assert error_type is None assert reason == "InvokeRateLimitError" @@ -165,12 +206,23 @@ def test_run_failed_error_payload_preserves_knowledge_error_code() -> None: retryable=False, ) - message, reason = _run_failed_error_payload(exc) + message, error_type, reason = _run_failed_error_payload(exc) assert message == "Knowledge base search failed with HTTP 400 (dataset_not_found): Dataset not found" + assert error_type is None assert reason == "dataset_not_found" +def test_run_failed_error_payload_classifies_usage_limit() -> None: + exc = UsageLimitExceeded("The next request would exceed the request_limit of 100") + + message, error_type, reason = _run_failed_error_payload(exc) + + assert message == "The next request would exceed the request_limit of 100" + assert error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + assert reason is None + + def test_cancelled_runner_does_not_overwrite_cancelled_status_with_late_failure( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -190,7 +242,12 @@ def test_cancelled_runner_does_not_overwrite_cancelled_status_with_late_failure( async def fail_after_cancel() -> RunSuccessOutcome: nonlocal cancelled cancelled = True - await sink.update_status("run-cancelled", "cancelled", "workflow stopped") + _ = await emit_run_cancelled( + sink, + run_id="run-cancelled", + reason="workflow_aborted", + message="workflow stopped", + ) raise RuntimeError("late model failure") monkeypatch.setattr(runner, "_run_agent", fail_after_cancel) @@ -198,7 +255,7 @@ def test_cancelled_runner_does_not_overwrite_cancelled_status_with_late_failure( assert sink.statuses["run-cancelled"] == "cancelled" assert sink.errors["run-cancelled"] == "workflow stopped" - assert [event.type for event in sink.events["run-cancelled"]] == ["run_started"] + assert [event.type for event in sink.events["run-cancelled"]] == ["run_started", "run_cancelled"] asyncio.run(scenario()) @@ -273,44 +330,6 @@ def _request( ) -def _lifecycle_only_request( - *, - on_exit: LayerExitSignals | None = None, - session_snapshot: CompositorSessionSnapshot | None = None, - deferred_tool_results: DeferredToolResultsPayload | None = None, -) -> CreateRunRequest: - snapshot = session_snapshot or CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot(name="prompt", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={}), - LayerSessionSnapshot(name="execution_context", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={}), - ] - ) - return CreateRunRequest( - composition=RunComposition( - layers=[ - RunLayerSpec( - name="prompt", - type="plain.prompt", - config=PromptLayerConfig(prefix="system", user="hello"), - ), - RunLayerSpec( - name="execution_context", - type=DIFY_EXECUTION_CONTEXT_LAYER_TYPE_ID, - config=DifyExecutionContextLayerConfig( - tenant_id="tenant-1", - user_from="account", - agent_mode="workflow_run", - invoke_from="service-api", - ), - ), - ] - ), - session_snapshot=snapshot, - deferred_tool_results=deferred_tool_results, - on_exit=on_exit or LayerExitSignals(default=ExitIntent.DELETE), - ) - - def _recursive_output_schema() -> dict[str, object]: return { "type": "object", @@ -605,6 +624,36 @@ def test_runner_preserves_explicit_json_null_output(monkeypatch: pytest.MonkeyPa assert sink.statuses["run-null-output"] == "succeeded" +def test_runner_passes_explicit_step_limit_to_agent(monkeypatch: pytest.MonkeyPatch) -> None: + def fake_get_model(_self: DifyPluginLLMLayer, *, http_client: httpx.AsyncClient): + assert http_client.is_closed is False + return TestModel(custom_output_text="unused") # pyright: ignore[reportReturnType] + + class FakeAgent: + async def run(self, *_args: object, **kwargs: object) -> FakeAgentRunResult: + usage_limits = cast(UsageLimits, kwargs["usage_limits"]) + assert usage_limits.request_limit == 100 + return FakeAgentRunResult("done", []) + + monkeypatch.setattr(DifyPluginLLMLayer, "get_model", fake_get_model) + monkeypatch.setattr("dify_agent.runtime.runner.create_agent", lambda *_args, **_kwargs: FakeAgent()) + sink = InMemoryRunEventSink() + + async def scenario() -> None: + async with httpx.AsyncClient() as client: + await AgentRunRunner( + sink=sink, + request=_request(), + run_id="run-explicit-step-limit", + plugin_daemon_http_client=client, + dify_api_http_client=client, + ).run() + + asyncio.run(scenario()) + + assert sink.statuses["run-explicit-step-limit"] == "succeeded" + + def test_runner_emits_deferred_tool_call_and_persists_pending_history(monkeypatch: pytest.MonkeyPatch) -> None: captured_output_types: list[object] = [] captured_user_prompts: list[object] = [] @@ -1655,20 +1704,11 @@ def test_runner_rejects_duplicate_tool_names_between_shell_and_other_layers( monkeypatch.setattr(DifyPluginToolsLayer, "get_tools", fake_get_tools) monkeypatch.setattr("dify_agent.runtime.runner.create_agent", fake_create_agent) - shell_provider = LayerProvider.from_factory( - layer_type=DifyShellLayer, - create=lambda config: DifyShellLayer.from_config_with_settings( - DifyShellLayerConfig.model_validate(config), - shell_provider=ShellctlProvider( - entrypoint="http://shellctl", - token="", - client_factory=lambda: shell_client, - ), - ), + runtime_backend_profile = RuntimeBackendProfile( + home_snapshots=cast(HomeSnapshotBackend, object()), + execution_bindings=FakeRunnerExecutionBindingBackend(shell_client), ) - layer_providers = tuple( - provider for provider in create_default_layer_providers() if provider.type_id != DIFY_SHELL_LAYER_TYPE_ID - ) + (shell_provider,) + layer_providers = create_default_layer_providers(runtime_backend_profile=runtime_backend_profile) request = CreateRunRequest( composition=RunComposition( @@ -1684,15 +1724,21 @@ def test_runner_rejects_duplicate_tool_names_between_shell_and_other_layers( config=DifyExecutionContextLayerConfig( tenant_id="tenant-1", agent_id="agent-1", + agent_config_version_id="config-1", user_from="account", agent_mode="workflow_run", invoke_from="service-api", ), ), + RunLayerSpec( + name="runtime", + type=DIFY_RUNTIME_LAYER_TYPE_ID, + config=DifyRuntimeLayerConfig(backend_binding_ref="binding-1"), + ), RunLayerSpec( name="shell", type=DIFY_SHELL_LAYER_TYPE_ID, - deps={"execution_context": "execution_context"}, + deps={"execution_context": "execution_context", "runtime": "runtime"}, config=DifyShellLayerConfig(), ), RunLayerSpec( @@ -1746,7 +1792,7 @@ def test_runner_rejects_duplicate_tool_names_between_shell_and_other_layers( asyncio.run(scenario()) assert create_agent_called is False - assert shell_client.delete_calls == [("mkdir-job", True, None)] + assert shell_client.delete_calls == [] assert shell_client.closed is True assert [event.type for event in sink.events["run-shell-duplicate-tools"]] == ["run_started", "run_failed"] assert sink.statuses["run-shell-duplicate-tools"] == "failed" @@ -1926,6 +1972,37 @@ def test_runner_failure_with_history_layer_emits_failed_terminal_event_without_s assert _history_messages_from_snapshot(request.session_snapshot) == stored_history +def test_runner_persists_usage_limit_failure_type_in_event_and_status( + monkeypatch: pytest.MonkeyPatch, +) -> None: + sink = InMemoryRunEventSink() + + async def scenario() -> None: + async with httpx.AsyncClient() as client: + runner = AgentRunRunner( + sink=sink, + request=_request(), + run_id="run-limit", + plugin_daemon_http_client=client, + dify_api_http_client=client, + ) + + async def exceed_limit() -> RunSuccessOutcome: + raise UsageLimitExceeded("The next request would exceed the request_limit of 100") + + monkeypatch.setattr(runner, "_run_agent", exceed_limit) + with pytest.raises(UsageLimitExceeded): + await runner.run() + + asyncio.run(scenario()) + + terminal = sink.events["run-limit"][-1] + assert isinstance(terminal, RunFailedEvent) + assert terminal.data.error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + assert sink.statuses["run-limit"] == "failed" + assert sink.error_types["run-limit"] is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + + def test_runner_applies_on_exit_overrides_to_success_snapshot(monkeypatch: pytest.MonkeyPatch) -> None: def fake_get_model(_self: DifyPluginLLMLayer, *, http_client: httpx.AsyncClient): assert http_client.is_closed is False @@ -1961,110 +2038,25 @@ def test_runner_applies_on_exit_overrides_to_success_snapshot(monkeypatch: pytes } -def test_runner_lifecycle_only_cleanup_succeeds_without_model_and_emits_no_pydantic_ai_events() -> None: - request = _lifecycle_only_request() - sink = InMemoryRunEventSink() - - async def scenario() -> None: - async with httpx.AsyncClient() as client: - await AgentRunRunner( - sink=sink, - request=request, - run_id="run-lifecycle-only", - plugin_daemon_http_client=client, - dify_api_http_client=client, - ).run() - - asyncio.run(scenario()) - - events = sink.events["run-lifecycle-only"] - assert [event.type for event in events] == ["run_started", "run_succeeded"] - terminal = events[-1] - assert isinstance(terminal, RunSucceededEvent) - assert terminal.data.output is None - assert terminal.data.usage is None - assert {layer.name: layer.lifecycle_state for layer in terminal.data.session_snapshot.layers} == { - "prompt": LifecycleState.CLOSED, - "execution_context": LifecycleState.CLOSED, - } - - -def test_runner_lifecycle_only_requires_session_snapshot() -> None: +def test_runner_requires_model_layer() -> None: request = _request(llm_layer_name="not-llm") sink = InMemoryRunEventSink() async def scenario() -> None: async with httpx.AsyncClient() as client: - with pytest.raises(AgentRunValidationError, match="session_snapshot"): + with pytest.raises(AgentRunValidationError, match="Missing required"): await AgentRunRunner( sink=sink, request=request, - run_id="run-lifecycle-only-missing-snapshot", + run_id="run-missing-model", plugin_daemon_http_client=client, dify_api_http_client=client, ).run() asyncio.run(scenario()) - assert [event.type for event in sink.events["run-lifecycle-only-missing-snapshot"]] == ["run_started", "run_failed"] - assert sink.statuses["run-lifecycle-only-missing-snapshot"] == "failed" - - -def test_runner_lifecycle_only_rejects_deferred_tool_results() -> None: - request = _lifecycle_only_request( - deferred_tool_results=DeferredToolResultsPayload.model_validate({"calls": {"tool-call-1": {"ok": True}}}) - ) - sink = InMemoryRunEventSink() - - async def scenario() -> None: - async with httpx.AsyncClient() as client: - with pytest.raises(AgentRunValidationError, match="Deferred tool results"): - await AgentRunRunner( - sink=sink, - request=request, - run_id="run-lifecycle-only-deferred-results", - plugin_daemon_http_client=client, - dify_api_http_client=client, - ).run() - - asyncio.run(scenario()) - - assert [event.type for event in sink.events["run-lifecycle-only-deferred-results"]] == [ - "run_started", - "run_failed", - ] - assert sink.statuses["run-lifecycle-only-deferred-results"] == "failed" - - -def test_runner_lifecycle_only_exit_hook_failure_emits_run_failed_not_validation_error( - monkeypatch: pytest.MonkeyPatch, -) -> None: - request = _lifecycle_only_request() - sink = InMemoryRunEventSink() - - def _explode(_run: object, _signals: LayerExitSignals) -> None: - raise RuntimeError("delete hook failed") - - monkeypatch.setattr("dify_agent.runtime.runner.apply_layer_exit_signals", _explode) - - async def scenario() -> None: - async with httpx.AsyncClient() as client: - with pytest.raises(RuntimeError, match="delete hook failed"): - await AgentRunRunner( - sink=sink, - request=request, - run_id="run-lifecycle-only-exit-hook-failure", - plugin_daemon_http_client=client, - dify_api_http_client=client, - ).run() - - asyncio.run(scenario()) - - assert [event.type for event in sink.events["run-lifecycle-only-exit-hook-failure"]] == [ - "run_started", - "run_failed", - ] - assert sink.statuses["run-lifecycle-only-exit-hook-failure"] == "failed" + assert [event.type for event in sink.events["run-missing-model"]] == ["run_started", "run_failed"] + assert sink.statuses["run-missing-model"] == "failed" def test_runner_passes_output_layer_spec_to_agent_and_serializes_structured_result( @@ -2757,7 +2749,7 @@ def test_runner_rejects_closed_session_snapshot_as_validation_error() -> None: assert sink.statuses["run-closed-snapshot"] == "failed" -def test_runner_treats_missing_shell_entrypoint_as_validation_error() -> None: +def test_runner_treats_missing_runtime_dependency_as_validation_error() -> None: request = CreateRunRequest( composition=RunComposition( layers=[ @@ -2795,7 +2787,7 @@ def test_runner_treats_missing_shell_entrypoint_as_validation_error() -> None: async def scenario() -> None: async with httpx.AsyncClient() as client: - with pytest.raises(AgentRunValidationError, match="non-null shell provider"): + with pytest.raises(AgentRunValidationError, match="Dependency 'runtime' is required"): await AgentRunRunner( sink=sink, request=request, @@ -2880,9 +2872,7 @@ def test_runner_treats_invalid_shell_snapshot_offsets_as_validation_error() -> N run_id="run-invalid-shell-offset", plugin_daemon_http_client=client, dify_api_http_client=client, - layer_providers=create_default_layer_providers( - shell_provider=ShellctlProvider(entrypoint="http://shellctl", token=""), - ), + layer_providers=create_default_layer_providers(), ).run() asyncio.run(scenario()) diff --git a/dify-agent/tests/local/dify_agent/runtime_backend/test_e2b.py b/dify-agent/tests/local/dify_agent/runtime_backend/test_e2b.py new file mode 100644 index 00000000000..247491b0b9c --- /dev/null +++ b/dify-agent/tests/local/dify_agent/runtime_backend/test_e2b.py @@ -0,0 +1,432 @@ +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import cast + +import httpx2 as httpx +import pytest +from shellctl.client import ShellctlClientError + +from dify_agent.runtime_backend import ( + BindingAcquireError, + BindingCreateError, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotCreateSpec, + SharedWorkspaceUnsupportedError, + WorkspacePreservationUnsupportedError, +) +from dify_agent.runtime_backend import e2b as e2b_module +from dify_agent.runtime_backend.e2b import ( + E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + E2BExecutionBindingBackend, + E2BHomeSnapshotBackend, + E2BRuntimeLease, +) +from dify_agent.runtime_backend.shellctl import ShellctlRuntimeLease + + +@dataclass(slots=True) +class _Files: + paths: set[str] = field(default_factory=set) + + async def make_dir(self, path: str) -> bool: + self.paths.add(path) + return True + + async def exists(self, path: str) -> bool: + return path in self.paths + + async def remove(self, path: str) -> None: + self.paths.discard(path) + + +@dataclass(slots=True) +class _Snapshot: + snapshot_id: str + names: list[str] = field(default_factory=list) + + +@dataclass(slots=True) +class _Sandbox: + sandbox_id: str + files: _Files = field(default_factory=_Files) + traffic_access_token: str | None = "traffic-token" + pauses: list[bool] = field(default_factory=list) + killed: int = 0 + snapshots: int = 0 + pause_error: Exception | None = None + + def get_host(self, port: int) -> str: + return f"{self.sandbox_id}-{port}.example.test" + + async def pause(self, keep_memory: bool = True) -> bool: + self.pauses.append(keep_memory) + if self.pause_error is not None: + raise self.pause_error + return True + + async def kill(self) -> bool: + self.killed += 1 + return True + + async def create_snapshot(self, name: str | None = None) -> _Snapshot: + del name + self.snapshots += 1 + return _Snapshot(snapshot_id=f"snapshot-{self.sandbox_id}-{self.snapshots}") + + +@dataclass(slots=True) +class _ControlPlane: + created: list[tuple[str, str]] = field(default_factory=list) + sandboxes: dict[str, _Sandbox] = field(default_factory=dict) + killed: list[str] = field(default_factory=list) + deleted_snapshots: list[str] = field(default_factory=list) + pause_error: Exception | None = None + + async def create(self, template: str, *, timeout: int, metadata: dict[str, str], on_timeout: str) -> _Sandbox: + del timeout + sandbox_id = f"sandbox-{len(self.sandboxes) + 1}" + sandbox = _Sandbox(sandbox_id=sandbox_id, pause_error=self.pause_error) + self.sandboxes[sandbox_id] = sandbox + self.created.append((template, on_timeout)) + assert metadata["dify.resource"] == "runtime-sandbox" + return sandbox + + async def connect(self, handle: str, *, timeout: int) -> _Sandbox: + del timeout + return self.sandboxes[handle] + + async def kill(self, handle: str) -> bool: + self.killed.append(handle) + return True + + async def delete_snapshot(self, snapshot_ref: str) -> bool: + self.deleted_snapshots.append(snapshot_ref) + return True + + +def _mock_http( + monkeypatch: pytest.MonkeyPatch, + handler: Callable[[httpx.Request], httpx.Response], +) -> list[httpx.AsyncClient]: + original_async_client = httpx.AsyncClient + transport = httpx.MockTransport(handler) + clients: list[httpx.AsyncClient] = [] + + def create_client(*args: object, **kwargs: object) -> httpx.AsyncClient: + _ = kwargs.setdefault("transport", transport) + client = original_async_client(*args, **kwargs) + clients.append(client) + return client + + monkeypatch.setattr(httpx, "AsyncClient", create_client) + return clients + + +def _connected_backend(*, pause_error: Exception | None = None) -> tuple[E2BExecutionBindingBackend, _Sandbox]: + control = _ControlPlane() + sandbox = _Sandbox(sandbox_id="sandbox-1", pause_error=pause_error) + sandbox.files.paths.add("/home/dify/workspace") + control.sandboxes[sandbox.sandbox_id] = sandbox + return ( + E2BExecutionBindingBackend( + control_plane=control, # pyright: ignore[reportArgumentType] + template="prepared-template", + active_timeout_seconds=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ), + sandbox, + ) + + +@pytest.mark.anyio +async def test_e2b_binding_uses_default_template_or_exact_snapshot_and_couples_refs() -> None: + control = _ControlPlane() + snapshots = E2BHomeSnapshotBackend( + control_plane=control, # pyright: ignore[reportArgumentType] + ) + bindings = E2BExecutionBindingBackend( + control_plane=control, # pyright: ignore[reportArgumentType] + template="prepared-template", + active_timeout_seconds=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ) + + default_allocation = await bindings.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref=None, + ) + ) + snapshot_allocation = await bindings.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-2", + workspace_id="workspace-2", + existing_workspace_ref=None, + home_snapshot_ref="snapshot-1", + ) + ) + + assert control.created == [("prepared-template", "pause"), ("snapshot-1", "pause")] + assert default_allocation.binding_ref == default_allocation.workspace_ref + assert snapshot_allocation.binding_ref == snapshot_allocation.workspace_ref + runtime = control.sandboxes[default_allocation.binding_ref] + assert runtime.files.paths == {"/home/dify/workspace"} + assert runtime.pauses == [True] + + for allocation in (default_allocation, snapshot_allocation): + await bindings.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref=allocation.binding_ref, + workspace_ref=allocation.workspace_ref, + destroy_workspace=True, + ) + ) + await snapshots.delete("snapshot-1") + + assert control.killed == [default_allocation.binding_ref, snapshot_allocation.binding_ref] + assert control.deleted_snapshots == ["snapshot-1"] + + +@pytest.mark.anyio +async def test_e2b_rejects_shared_workspace_and_binding_only_destroy() -> None: + control = _ControlPlane() + backend = E2BExecutionBindingBackend( + control_plane=control, # pyright: ignore[reportArgumentType] + template="prepared-template", + active_timeout_seconds=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ) + spec = ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-2", + binding_id="binding-2", + workspace_id="workspace-1", + existing_workspace_ref="sandbox-1", + home_snapshot_ref="snapshot-1", + ) + + with pytest.raises(SharedWorkspaceUnsupportedError): + await backend.create_binding(spec) + assert control.created == [] + assert control.sandboxes == {} + with pytest.raises(WorkspacePreservationUnsupportedError): + await backend.destroy_binding(ExecutionBindingDestroySpec(binding_ref="sandbox-1", destroy_workspace=False)) + + +@pytest.mark.anyio +async def test_e2b_binding_create_kills_sandbox_when_initialization_fails() -> None: + control = _ControlPlane(pause_error=RuntimeError("pause failed")) + backend = E2BExecutionBindingBackend( + control_plane=control, # pyright: ignore[reportArgumentType] + template="prepared-template", + active_timeout_seconds=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ) + + with pytest.raises(BindingCreateError, match="pause failed"): + await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="snapshot-1", + ) + ) + + sandbox = next(iter(control.sandboxes.values())) + assert sandbox.killed == 1 + + +@pytest.mark.anyio +async def test_e2b_missing_explicit_snapshot_does_not_fall_back_to_template() -> None: + class _FailingControlPlane(_ControlPlane): + async def create(self, template: str, *, timeout: int, metadata: dict[str, str], on_timeout: str) -> _Sandbox: + del timeout, metadata + self.created.append((template, on_timeout)) + raise RuntimeError("snapshot unavailable") + + control = _FailingControlPlane() + backend = E2BExecutionBindingBackend( + control_plane=control, # pyright: ignore[reportArgumentType] + template="prepared-template", + active_timeout_seconds=E2B_MAX_ACTIVE_TIMEOUT_SECONDS, + ) + + with pytest.raises(BindingCreateError, match="snapshot unavailable"): + await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="missing-snapshot", + ) + ) + + assert control.created == [("missing-snapshot", "pause")] + + +@pytest.mark.anyio +async def test_e2b_checkpoint_uses_exact_source_runtime() -> None: + control = _ControlPlane() + source_sandbox = _Sandbox(sandbox_id="source") + source = E2BRuntimeLease( + sandbox=source_sandbox, + data_plane=cast(ShellctlRuntimeLease, object()), + ) + backend = E2BHomeSnapshotBackend( + control_plane=control, # pyright: ignore[reportArgumentType] + ) + + snapshot_ref = await backend.create_from_runtime( + spec=HomeSnapshotCreateSpec(tenant_id="tenant-1", agent_id="agent-1", home_snapshot_id="home-2"), + source=source, + ) + + assert snapshot_ref == "snapshot-source-1" + assert source_sandbox.snapshots == 1 + + +@pytest.mark.anyio +async def test_e2b_acquire_retries_transient_shellctl_failures_until_ready( + monkeypatch: pytest.MonkeyPatch, +) -> None: + attempts = 0 + + def handler(request: httpx.Request) -> httpx.Response: + nonlocal attempts + attempts += 1 + if attempts == 1: + raise httpx.ReadTimeout("shellctl starting", request=request) + if attempts == 2: + raise httpx.ConnectError("shellctl starting", request=request) + return httpx.Response(200, json={"status": "ok"}) + + clients = _mock_http(monkeypatch, handler) + sleeps: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleeps.append(delay) + + monkeypatch.setattr(e2b_module.asyncio, "sleep", record_sleep) + backend, sandbox = _connected_backend() + + lease = await backend.acquire(sandbox.sandbox_id) + + assert attempts == 3 + assert sleeps == [0.5, 0.5] + assert not clients[0].is_closed + await backend.release(lease) + assert clients[0].is_closed + + +@pytest.mark.anyio +async def test_e2b_acquire_closes_transport_and_pauses_after_readiness_retries_exhausted( + monkeypatch: pytest.MonkeyPatch, +) -> None: + attempts = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal attempts + attempts += 1 + return httpx.Response(503, json={"error": {"code": "starting", "message": "not ready"}}) + + clients = _mock_http(monkeypatch, handler) + sleeps: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleeps.append(delay) + + monkeypatch.setattr(e2b_module.asyncio, "sleep", record_sleep) + backend, sandbox = _connected_backend() + + with pytest.raises(BindingAcquireError, match="not ready"): + _ = await backend.acquire(sandbox.sandbox_id) + + assert attempts == 3 + assert sleeps == [0.5, 0.5] + assert clients[0].is_closed + assert sandbox.pauses == [True] + + +@pytest.mark.anyio +async def test_e2b_acquire_does_not_retry_shellctl_4xx_and_preserves_error_when_pause_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + attempts = 0 + + def handler(_request: httpx.Request) -> httpx.Response: + nonlocal attempts + attempts += 1 + return httpx.Response(401, json={"error": {"code": "unauthorized", "message": "bad token"}}) + + clients = _mock_http(monkeypatch, handler) + + async def fail_sleep(_delay: float) -> None: + raise AssertionError("non-retryable health failures must not sleep") + + monkeypatch.setattr(e2b_module.asyncio, "sleep", fail_sleep) + backend, sandbox = _connected_backend(pause_error=RuntimeError("pause failed")) + + with pytest.raises(BindingAcquireError, match="bad token"): + _ = await backend.acquire(sandbox.sandbox_id) + + assert attempts == 1 + assert clients[0].is_closed + assert sandbox.pauses == [True] + + +@pytest.mark.anyio +async def test_e2b_acquire_preserves_health_failure_when_close_and_pause_fail( + monkeypatch: pytest.MonkeyPatch, +) -> None: + @dataclass(slots=True) + class _UnavailableHealthClient: + calls: int = 0 + + async def health(self) -> object: + self.calls += 1 + raise ShellctlClientError(503, "starting", "primary health failure") + + @dataclass(slots=True) + class _FailingCloseDataPlane: + client: _UnavailableHealthClient + close_calls: int = 0 + + async def close(self) -> None: + self.close_calls += 1 + raise RuntimeError("close failed") + + client = _UnavailableHealthClient() + data_plane = _FailingCloseDataPlane(client=client) + + async def create_lease( + _self: E2BExecutionBindingBackend, + sandbox: _Sandbox, + ) -> E2BRuntimeLease: + return E2BRuntimeLease( + sandbox=sandbox, # pyright: ignore[reportArgumentType] + data_plane=cast(ShellctlRuntimeLease, cast(object, data_plane)), + ) + + async def skip_sleep(_delay: float) -> None: + return None + + monkeypatch.setattr(E2BExecutionBindingBackend, "_lease", create_lease) + monkeypatch.setattr(e2b_module.asyncio, "sleep", skip_sleep) + backend, sandbox = _connected_backend(pause_error=RuntimeError("pause failed")) + + with pytest.raises(BindingAcquireError, match="primary health failure"): + _ = await backend.acquire(sandbox.sandbox_id) + + assert client.calls == 3 + assert data_plane.close_calls == 1 + assert sandbox.pauses == [True] diff --git a/dify-agent/tests/local/dify_agent/runtime_backend/test_enterprise_backend.py b/dify-agent/tests/local/dify_agent/runtime_backend/test_enterprise_backend.py new file mode 100644 index 00000000000..841e6d63a91 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/runtime_backend/test_enterprise_backend.py @@ -0,0 +1,403 @@ +from __future__ import annotations + +import json +from collections.abc import Callable +from typing import cast + +import httpx2 as httpx +import pytest + +from dify_agent.runtime_backend import ( + BindingAcquireError, + BindingCreateError, + BindingDestroyError, + BindingLostError, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotCreateSpec, + RuntimeLease, + SharedWorkspaceUnsupportedError, + WorkspacePreservationUnsupportedError, +) +from dify_agent.runtime_backend.enterprise import ( + EnterpriseExecutionBindingBackend, + EnterpriseHomeSnapshotBackend, +) + + +def _mock_http( + monkeypatch: pytest.MonkeyPatch, + handler: Callable[[httpx.Request], httpx.Response], +) -> list[httpx.AsyncClient]: + original_async_client = httpx.AsyncClient + transport = httpx.MockTransport(handler) + clients: list[httpx.AsyncClient] = [] + + def create_transport(*, retries: int = 0) -> httpx.AsyncBaseTransport: + del retries + return transport + + monkeypatch.setattr(httpx, "AsyncHTTPTransport", create_transport) + + def create_client(*args: object, **kwargs: object) -> httpx.AsyncClient: + _ = kwargs.setdefault("transport", transport) + client = original_async_client(*args, **kwargs) + clients.append(client) + return client + + monkeypatch.setattr(httpx, "AsyncClient", create_client) + return clients + + +def _job_response(*, exit_code: int = 0) -> httpx.Response: + return httpx.Response( + 200, + json={ + "job_id": "job-1", + "done": True, + "status": "exited", + "exit_code": exit_code, + "output_path": "/tmp/output.log", + "output": "", + "offset": 0, + "truncated": False, + }, + ) + + +@pytest.mark.anyio +async def test_enterprise_acquire_exposes_canonical_layout_through_gateway_proxy( + monkeypatch: pytest.MonkeyPatch, +) -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + assert request.headers["X-Sandbox-Id"] == "sandbox-1" + assert request.headers["X-Inner-Api-Key"] == "secret" + if request.method == "POST": + payload = cast(dict[str, object], json.loads(request.content)) + script = payload["script"] + assert isinstance(script, str) + assert "test -d /home/dify" in script + assert "test -d /home/dify/workspace" in script + return _job_response() + return httpx.Response(200, json={"job_id": "job-1"}) + + clients = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend( + gateway_endpoint="http://gateway.example", + auth_token="secret", + proxy_timeout=90, + ) + + lease = await backend.acquire("sandbox-1") + + assert lease.layout.home_dir == "/home/dify" + assert lease.layout.workspace_dir == "/home/dify/workspace" + assert [request.url.path for request in requests] == [ + "/proxy/v1/jobs/run", + "/proxy/v1/jobs/job-1", + ] + assert clients[0].timeout.read == 90 + await backend.release(lease) + assert clients[0].is_closed + + +@pytest.mark.anyio +@pytest.mark.parametrize( + ("status_code", "code", "expected_error"), + [ + (404, "sandbox_expired", BindingLostError), + (502, "upstream_failure", BindingAcquireError), + ], +) +async def test_enterprise_acquire_maps_proxy_failures_and_closes_transport( + monkeypatch: pytest.MonkeyPatch, + status_code: int, + code: str, + expected_error: type[Exception], +) -> None: + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response(status_code, json={"error": {"code": code, "message": "unavailable"}}) + + clients = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend( + gateway_endpoint="http://gateway.example", + auth_token="secret", + ) + + with pytest.raises(expected_error): + _ = await backend.acquire("sandbox-1") + + assert clients[0].is_closed + + +@pytest.mark.anyio +async def test_enterprise_acquire_treats_missing_runtime_directories_as_lost( + monkeypatch: pytest.MonkeyPatch, +) -> None: + def handler(request: httpx.Request) -> httpx.Response: + if request.method == "POST": + return _job_response(exit_code=1) + return httpx.Response(200, json={"job_id": "job-1"}) + + clients = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend( + gateway_endpoint="http://gateway.example", + auth_token="secret", + ) + + with pytest.raises(BindingLostError, match="Home or Workspace"): + _ = await backend.acquire("sandbox-1") + + assert clients[0].is_closed + + +@pytest.mark.anyio +async def test_enterprise_release_rejects_foreign_runtime_lease() -> None: + backend = EnterpriseExecutionBindingBackend( + gateway_endpoint="http://gateway.example", + auth_token="secret", + ) + + with pytest.raises(TypeError, match="only release its own RuntimeLease"): + await backend.release(cast(RuntimeLease, object())) + + +@pytest.mark.anyio +async def test_enterprise_destroy_validates_coupled_workspace_before_gateway_call( + monkeypatch: pytest.MonkeyPatch, +) -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(204) + + _ = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend( + gateway_endpoint="http://gateway.example", + auth_token="secret", + ) + + with pytest.raises(WorkspacePreservationUnsupportedError): + await backend.destroy_binding(ExecutionBindingDestroySpec(binding_ref="sandbox-1", destroy_workspace=False)) + with pytest.raises(BindingDestroyError, match="must equal"): + await backend.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref="sandbox-1", + workspace_ref="workspace-1", + destroy_workspace=True, + ) + ) + + assert requests == [] + + +@pytest.mark.anyio +@pytest.mark.parametrize("status_code", [204, 404]) +async def test_enterprise_destroy_is_authenticated_and_idempotent( + monkeypatch: pytest.MonkeyPatch, + status_code: int, +) -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(status_code) + + clients = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend( + gateway_endpoint="http://gateway.example", + auth_token="secret", + ) + + await backend.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref="sandbox-1", + workspace_ref="sandbox-1", + destroy_workspace=True, + ) + ) + + assert len(requests) == 1 + assert requests[0].method == "DELETE" + assert requests[0].url.path == "/v1/sandboxes/sandbox-1" + assert requests[0].headers["X-Inner-Api-Key"] == "secret" + assert clients[0].is_closed + + +@pytest.mark.anyio +async def test_enterprise_destroy_encodes_binding_ref_as_one_path_segment( + monkeypatch: pytest.MonkeyPatch, +) -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(204) + + _ = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend( + gateway_endpoint="http://gateway.example", + auth_token="secret", + ) + opaque_ref = "../admin?x=1" + + await backend.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref=opaque_ref, + workspace_ref=opaque_ref, + destroy_workspace=True, + ) + ) + + assert len(requests) == 1 + assert requests[0].url.raw_path == b"/v1/sandboxes/..%2Fadmin%3Fx%3D1" + + +@pytest.mark.anyio +async def test_enterprise_destroy_propagates_gateway_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + def handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response(502, text="gateway failed") + + clients = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend( + gateway_endpoint="http://gateway.example", + auth_token="secret", + ) + + with pytest.raises(BindingDestroyError, match="502"): + await backend.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref="sandbox-1", + workspace_ref="sandbox-1", + destroy_workspace=True, + ) + ) + + assert clients[0].is_closed + + +@pytest.mark.anyio +async def test_enterprise_default_binding_creates_gateway_sandbox_and_layout( + monkeypatch: pytest.MonkeyPatch, +) -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.url.path == "/v1/sandboxes": + assert json.loads(request.content) == {"tenantId": "tenant-1"} + return httpx.Response(201, json={"sandboxId": "sandbox-1", "status": "running"}) + if request.url.path == "/proxy/v1/jobs/run": + assert request.headers["X-Sandbox-Id"] == "sandbox-1" + payload = cast(dict[str, object], json.loads(request.content)) + script = payload["script"] + assert isinstance(script, str) + assert "mkdir -p /home/dify" in script + assert "rm -rf -- /home/dify/workspace" in script + return _job_response() + return httpx.Response(200, json={"job_id": "job-1"}) + + clients = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend(gateway_endpoint="http://gateway.example", auth_token="secret") + + allocation = await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref=None, + ) + ) + + assert allocation.binding_ref == allocation.workspace_ref == "sandbox-1" + assert requests[0].headers["X-Inner-Api-Key"] == "secret" + assert all(client.is_closed for client in clients) + + +@pytest.mark.anyio +async def test_enterprise_binding_rejects_snapshot_and_shared_workspace_before_gateway_call( + monkeypatch: pytest.MonkeyPatch, +) -> None: + requests: list[httpx.Request] = [] + _ = _mock_http(monkeypatch, lambda request: requests.append(request) or httpx.Response(500)) + backend = EnterpriseExecutionBindingBackend(gateway_endpoint="http://gateway.example", auth_token="secret") + + with pytest.raises(BindingCreateError, match="immutable Home Snapshot"): + await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="snapshot-1", + ) + ) + with pytest.raises(SharedWorkspaceUnsupportedError): + await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref="workspace-1", + home_snapshot_ref=None, + ) + ) + + assert requests == [] + + +@pytest.mark.anyio +async def test_enterprise_binding_create_deletes_new_sandbox_when_layout_setup_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.url.path == "/v1/sandboxes": + return httpx.Response(201, json={"sandboxId": "sandbox-1"}) + if request.url.path == "/proxy/v1/jobs/run": + return _job_response(exit_code=1) + if request.method == "DELETE": + return httpx.Response(204) + return httpx.Response(200, json={"job_id": "job-1"}) + + _ = _mock_http(monkeypatch, handler) + backend = EnterpriseExecutionBindingBackend(gateway_endpoint="http://gateway.example", auth_token="secret") + + with pytest.raises(BindingCreateError): + await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref=None, + ) + ) + + assert any(request.method == "DELETE" and request.url.path == "/v1/sandboxes/sandbox-1" for request in requests) + + +@pytest.mark.anyio +async def test_enterprise_home_snapshots_remain_explicitly_not_implemented() -> None: + snapshots = EnterpriseHomeSnapshotBackend() + + with pytest.raises(NotImplementedError, match="immutable Home Snapshot"): + _ = await snapshots.create_from_runtime( + spec=HomeSnapshotCreateSpec(tenant_id="tenant-1", agent_id="agent-1", home_snapshot_id="home-2"), + source=cast(RuntimeLease, object()), + ) + with pytest.raises(NotImplementedError, match="immutable Home Snapshot"): + await snapshots.delete("snapshot-1") diff --git a/dify-agent/tests/local/dify_agent/runtime_backend/test_leases.py b/dify-agent/tests/local/dify_agent/runtime_backend/test_leases.py new file mode 100644 index 00000000000..9faecc90786 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/runtime_backend/test_leases.py @@ -0,0 +1,33 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from dify_agent.runtime_backend.leases import open_runtime_lease + + +@pytest.mark.anyio +async def test_runtime_lease_body_error_survives_release_error() -> None: + lease = MagicMock() + backend = MagicMock() + backend.acquire = AsyncMock(return_value=lease) + backend.release = AsyncMock(side_effect=RuntimeError("release failed")) + + with pytest.raises(ValueError, match="body failed"): + async with open_runtime_lease(backend, "binding-ref"): + raise ValueError("body failed") + + backend.release.assert_awaited_once_with(lease) + + +@pytest.mark.anyio +async def test_runtime_lease_release_error_propagates_after_successful_body() -> None: + lease = MagicMock() + backend = MagicMock() + backend.acquire = AsyncMock(return_value=lease) + backend.release = AsyncMock(side_effect=RuntimeError("release failed")) + + with pytest.raises(RuntimeError, match="release failed"): + async with open_runtime_lease(backend, "binding-ref"): + pass + + backend.release.assert_awaited_once_with(lease) diff --git a/dify-agent/tests/local/dify_agent/runtime_backend/test_local.py b/dify-agent/tests/local/dify_agent/runtime_backend/test_local.py new file mode 100644 index 00000000000..a2f8190e32f --- /dev/null +++ b/dify-agent/tests/local/dify_agent/runtime_backend/test_local.py @@ -0,0 +1,436 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +import shlex +from typing import Mapping + +import pytest +from shellctl.shared import DeleteJobResponse, JobResult, JobStatusName, JobStatusView + +from dify_agent.runtime_backend import ( + BindingCreateError, + BindingDestroyError, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, + HomeSnapshotCreateSpec, +) +from dify_agent.runtime_backend.local import LocalExecutionBindingBackend, LocalHomeSnapshotBackend + + +@dataclass(slots=True) +class _RunCall: + commands: tuple[tuple[str, ...], ...] + cwd: str | None + env: Mapping[str, str] | None + + +@dataclass(slots=True) +class _Client: + runs: list[_RunCall] + closed: bool = False + exit_code: int = 0 + output: str = "" + close_error: Exception | None = None + exit_codes: list[int] = field(default_factory=list) + + async def run( + self, + script: str, + *, + cwd: str | None = None, + env: Mapping[str, str] | None = None, + timeout: float = 10.0, + ) -> JobResult: + del timeout + commands = tuple( + tuple(shlex.split(line)) for line in script.splitlines() if line.strip() and line.strip() != "set -eu" + ) + self.runs.append(_RunCall(commands=commands, cwd=cwd, env=env)) + return JobResult( + job_id=f"job-{len(self.runs)}", + status=JobStatusName.EXITED, + done=True, + exit_code=self.exit_codes.pop(0) if self.exit_codes else self.exit_code, + output_path="/tmp/output.log", + output=self.output, + offset=0, + truncated=False, + ) + + async def wait(self, job_id: str, *, offset: int, timeout: float = 10.0) -> JobResult: + raise AssertionError((job_id, offset, timeout)) + + async def input(self, job_id: str, text: str, *, offset: int, timeout: float = 10.0) -> JobResult: + raise AssertionError((job_id, text, offset, timeout)) + + async def tail(self, job_id: str) -> JobResult: + raise AssertionError(job_id) + + async def terminate(self, job_id: str, grace_seconds: float = 2.0) -> JobStatusView: + raise AssertionError((job_id, grace_seconds)) + + async def delete( + self, + job_id: str, + *, + force: bool = False, + grace_seconds: float | None = None, + ) -> DeleteJobResponse: + del force, grace_seconds + return DeleteJobResponse(job_id=job_id) + + async def close(self) -> None: + self.closed = True + if self.close_error is not None: + raise self.close_error + + +@dataclass(slots=True) +class _Factory: + clients: list[_Client] = field(default_factory=list) + runs: list[_RunCall] = field(default_factory=list) + + def __call__(self) -> _Client: + client = _Client(runs=self.runs) + self.clients.append(client) + return client + + @property + def commands(self) -> tuple[tuple[str, ...], ...]: + return tuple(command for run in self.runs for command in run.commands) + + +@dataclass(slots=True) +class _FailingFactory: + clients: list[_Client] = field(default_factory=list) + runs: list[_RunCall] = field(default_factory=list) + + def __call__(self) -> _Client: + client = _Client( + runs=self.runs, + exit_code=1, + output="primary shellctl failure", + close_error=RuntimeError("secondary close failure"), + ) + self.clients.append(client) + return client + + +@dataclass(slots=True) +class _FailThenSucceedFactory: + clients: list[_Client] = field(default_factory=list) + runs: list[_RunCall] = field(default_factory=list) + + def __call__(self) -> _Client: + client = _Client( + runs=self.runs, + output="primary shellctl failure", + exit_codes=[1, 0], + ) + self.clients.append(client) + return client + + @property + def commands(self) -> tuple[tuple[str, ...], ...]: + return tuple(command for run in self.runs for command in run.commands) + + +@pytest.mark.anyio +async def test_local_binding_create_materializes_home_and_new_workspace() -> None: + factory = _Factory() + backend = LocalExecutionBindingBackend( + endpoint="http://shellctl", + auth_token="", + materialized_home_root="/homes", + workspace_root="/workspaces", + snapshot_root="/snapshots", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + allocation = await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="home-home-1", + ) + ) + + assert allocation.binding_ref == "binding-1:workspace-1" + assert allocation.workspace_ref == "workspace-1" + assert factory.commands[0] == ("test", "-d", "/snapshots/home-home-1") + assert ("mkdir", "-p", "/workspaces/workspace-1") in factory.commands + assert ("mkdir", "-p", "/homes/binding-1") in factory.commands + assert ("cp", "-a", "/snapshots/home-home-1/.", "/homes/binding-1/") in factory.commands + assert ("chmod", "700", "/homes/binding-1", "/workspaces/workspace-1") in factory.commands + + +@pytest.mark.anyio +async def test_local_binding_create_uses_empty_default_home_without_snapshot_access() -> None: + factory = _Factory() + backend = LocalExecutionBindingBackend( + endpoint="http://shellctl", + auth_token="", + materialized_home_root="/homes", + workspace_root="/workspaces", + snapshot_root="/snapshots", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + allocation = await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref=None, + ) + ) + + assert allocation.binding_ref == "binding-1:workspace-1" + assert ("mkdir", "-p", "/homes/binding-1") in factory.commands + assert all("/snapshots" not in part for command in factory.commands for part in command) + + +@pytest.mark.anyio +async def test_local_binding_create_failure_removes_partial_home_and_workspace() -> None: + factory = _FailThenSucceedFactory() + backend = LocalExecutionBindingBackend( + endpoint="http://shellctl", + auth_token="", + materialized_home_root="/homes", + workspace_root="/workspaces", + snapshot_root="/snapshots", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + with pytest.raises(BindingCreateError, match="primary shellctl failure"): + await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="home-home-1", + ) + ) + + assert ("rm", "-rf", "--", "/homes/binding-1", "/workspaces/workspace-1") in factory.commands + + +@pytest.mark.anyio +async def test_local_binding_acquire_scopes_commands_to_materialized_home_and_workspace() -> None: + factory = _Factory() + backend = LocalExecutionBindingBackend( + endpoint="http://shellctl", + auth_token="", + materialized_home_root="/homes", + workspace_root="/workspaces", + snapshot_root="/snapshots", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + lease = await backend.acquire("binding-1:workspace-1") + + assert lease.layout.home_dir == "/homes/binding-1" + assert lease.layout.workspace_dir == "/workspaces/workspace-1" + assert ("test", "-d", "/homes/binding-1") in factory.commands + assert ("test", "-d", "/workspaces/workspace-1") in factory.commands + + await lease.commands.run("pwd", cwd=None, env={"HOME": "/homes/other"}, timeout=10.0) + pwd_run = next(run for run in factory.runs if run.commands == (("pwd",),)) + assert pwd_run.cwd == "/workspaces/workspace-1" + assert pwd_run.env == {"HOME": "/homes/binding-1"} + with pytest.raises(ValueError, match="outside this RuntimeLease"): + await lease.commands.run("cat secret", cwd="/homes/other", timeout=10.0) + await backend.release(lease) + + +@pytest.mark.anyio +async def test_local_snapshot_checkpoint_copies_only_materialized_home() -> None: + factory = _Factory() + snapshots = LocalHomeSnapshotBackend( + endpoint="http://shellctl", + auth_token="", + snapshot_root="/snapshots", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + bindings = LocalExecutionBindingBackend( + endpoint="http://shellctl", + auth_token="", + materialized_home_root="/homes", + workspace_root="/workspaces", + snapshot_root="/snapshots", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + lease = await bindings.acquire("binding-1:workspace-1") + + snapshot_ref = await snapshots.create_from_runtime( + spec=HomeSnapshotCreateSpec(tenant_id="tenant-1", agent_id="agent-1", home_snapshot_id="home-2"), + source=lease, + ) + await bindings.release(lease) + + assert snapshot_ref == "home-home-2" + assert ("test", "-d", "/homes/binding-1") in factory.commands + assert ("mkdir", "-p", "/snapshots/home-home-2") in factory.commands + assert ("cp", "-a", "/homes/binding-1/.", "/snapshots/home-home-2/") in factory.commands + assert ("cp", "-a", "/workspaces/workspace-1/.", "/snapshots/home-home-2/") not in factory.commands + + +@pytest.mark.anyio +async def test_local_binding_destroy_removes_home_and_requested_workspace() -> None: + factory = _Factory() + backend = LocalExecutionBindingBackend( + endpoint="http://shellctl", + auth_token="", + materialized_home_root="/homes", + workspace_root="/workspaces", + snapshot_root="/snapshots", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + await backend.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref="binding-1:workspace-1", + workspace_ref="workspace-1", + destroy_workspace=True, + ) + ) + + assert ("rm", "-rf", "--", "/homes/binding-1", "/workspaces/workspace-1") in factory.commands + + +@pytest.mark.anyio +async def test_local_snapshot_delete_removes_snapshot_directory() -> None: + factory = _Factory() + backend = LocalHomeSnapshotBackend( + endpoint="http://shellctl", + auth_token="", + snapshot_root="/snapshots", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + await backend.delete("home-home-2") + + assert ("rm", "-rf", "--", "/snapshots/home-home-2") in factory.commands + + +@pytest.mark.anyio +async def test_local_backend_materializes_same_agent_twice_in_one_workspace() -> None: + factory = _Factory() + backend = LocalExecutionBindingBackend( + endpoint="http://shellctl", + auth_token="", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + first = await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="home-home-1", + ) + ) + second = await backend.create_binding( + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-2", + workspace_id="workspace-1", + existing_workspace_ref=first.workspace_ref, + home_snapshot_ref="home-home-1", + ) + ) + + assert first.binding_ref == "binding-1:workspace-1" + assert second.binding_ref == "binding-2:workspace-1" + assert first.workspace_ref == second.workspace_ref == "workspace-1" + first_lease = await backend.acquire(first.binding_ref) + second_lease = await backend.acquire(second.binding_ref) + assert first_lease.layout.home_dir != second_lease.layout.home_dir + assert first_lease.layout.workspace_dir == second_lease.layout.workspace_dir + await backend.release(first_lease) + await backend.release(second_lease) + + await backend.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref=first.binding_ref, + destroy_workspace=False, + ) + ) + surviving_lease = await backend.acquire(second.binding_ref) + assert surviving_lease.layout.home_dir == "/home/dify/.dify-agent-materialized-homes/binding-2" + assert surviving_lease.layout.workspace_dir == "/home/dify/.dify-agent-workspaces/workspace-1" + await backend.release(surviving_lease) + + workspace_dir = "/home/dify/.dify-agent-workspaces/workspace-1" + assert ("test", "-d", workspace_dir) in factory.commands + assert factory.commands.count(("mkdir", "-p", workspace_dir)) == 1 + assert ( + "rm", + "-rf", + "--", + "/home/dify/.dify-agent-materialized-homes/binding-1", + ) in factory.commands + assert ( + "rm", + "-rf", + "--", + "/home/dify/.dify-agent-materialized-homes/binding-1", + workspace_dir, + ) not in factory.commands + assert ( + "cp", + "-a", + "/home/dify/.dify-agent-home-snapshots/home-home-1/.", + "/home/dify/.dify-agent-materialized-homes/binding-1/", + ) in factory.commands + assert ( + "cp", + "-a", + "/home/dify/.dify-agent-home-snapshots/home-home-1/.", + "/home/dify/.dify-agent-materialized-homes/binding-2/", + ) in factory.commands + + +@pytest.mark.anyio +async def test_local_snapshot_delete_preserves_shellctl_error_when_close_fails() -> None: + factory = _FailingFactory() + backend = LocalHomeSnapshotBackend( + endpoint="http://shellctl", + auth_token="", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + with pytest.raises(BindingDestroyError, match="primary shellctl failure"): + await backend.delete("home-1") + + assert factory.clients and all(client.closed for client in factory.clients) + + +@pytest.mark.anyio +async def test_local_binding_destroy_preserves_shellctl_error_when_close_fails() -> None: + factory = _FailingFactory() + backend = LocalExecutionBindingBackend( + endpoint="http://shellctl", + auth_token="", + client_factory=factory, # pyright: ignore[reportArgumentType] + ) + + with pytest.raises(BindingDestroyError, match="primary shellctl failure"): + await backend.destroy_binding( + ExecutionBindingDestroySpec( + binding_ref="binding-1:workspace-1", + destroy_workspace=False, + ) + ) + + assert factory.clients and all(client.closed for client in factory.clients) diff --git a/dify-agent/tests/local/dify_agent/runtime_backend/test_profile.py b/dify-agent/tests/local/dify_agent/runtime_backend/test_profile.py new file mode 100644 index 00000000000..fe6704ac969 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/runtime_backend/test_profile.py @@ -0,0 +1,68 @@ +from __future__ import annotations + +import pytest +from pydantic import ValidationError + +from dify_agent.runtime_backend.e2b import E2B_MAX_ACTIVE_TIMEOUT_SECONDS +from dify_agent.runtime_backend.local import LocalExecutionBindingBackend, LocalHomeSnapshotBackend +from dify_agent.runtime_backend.profile import ( + DEFAULT_E2B_TEMPLATE, + RuntimeBackendSettings, + create_runtime_backend_profile, +) + + +def test_e2b_backend_uses_prepared_dify_template_and_one_hour_lease_by_default() -> None: + settings = RuntimeBackendSettings(runtime_backend="e2b", e2b_api_key="secret") + + assert settings.e2b_template == "difys-default-team/dify-agent-local-sandbox" + assert settings.e2b_template == DEFAULT_E2B_TEMPLATE + assert E2B_MAX_ACTIVE_TIMEOUT_SECONDS == 60 * 60 + assert settings.e2b_active_timeout_seconds == E2B_MAX_ACTIVE_TIMEOUT_SECONDS + + +def test_e2b_backend_requires_api_key() -> None: + with pytest.raises(ValidationError, match="e2b_api_key"): + _ = RuntimeBackendSettings(runtime_backend="e2b") + + +def test_e2b_backend_rejects_active_timeout_above_platform_limit() -> None: + with pytest.raises(ValidationError, match="less than or equal"): + _ = RuntimeBackendSettings( + runtime_backend="e2b", + e2b_api_key="secret", + e2b_active_timeout_seconds=E2B_MAX_ACTIVE_TIMEOUT_SECONDS + 1, + ) + + +def test_local_backend_requires_shellctl_endpoint() -> None: + with pytest.raises(ValidationError, match="local_sandbox_endpoint"): + _ = RuntimeBackendSettings(runtime_backend="local") + + +def test_local_backend_passes_configured_roots_to_drivers() -> None: + settings = RuntimeBackendSettings( + runtime_backend="local", + local_sandbox_endpoint="http://shellctl.example", + local_sandbox_materialized_home_root="/tmp/dify/homes", + local_sandbox_workspace_root="/tmp/dify/workspaces", + local_sandbox_home_snapshot_root="/tmp/dify/snapshots", + ) + + profile = create_runtime_backend_profile(settings) + + assert isinstance(profile.execution_bindings, LocalExecutionBindingBackend) + assert isinstance(profile.home_snapshots, LocalHomeSnapshotBackend) + assert profile.execution_bindings.materialized_home_root == "/tmp/dify/homes" + assert profile.execution_bindings.workspace_root == "/tmp/dify/workspaces" + assert profile.execution_bindings.snapshot_root == "/tmp/dify/snapshots" + assert profile.home_snapshots.snapshot_root == "/tmp/dify/snapshots" + + +def test_local_backend_rejects_relative_roots() -> None: + with pytest.raises(ValidationError, match="absolute POSIX path"): + _ = RuntimeBackendSettings( + runtime_backend="local", + local_sandbox_endpoint="http://shellctl.example", + local_sandbox_workspace_root="relative/workspaces", + ) diff --git a/dify-agent/tests/local/dify_agent/runtime_backend/test_shellctl_backend.py b/dify-agent/tests/local/dify_agent/runtime_backend/test_shellctl_backend.py new file mode 100644 index 00000000000..9d07ccfb027 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/runtime_backend/test_shellctl_backend.py @@ -0,0 +1,191 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import cast + +import pytest + +from dify_agent.adapters.shell.protocols import ShellCommandResult, ShellCommandStatus +from dify_agent.adapters.shell.shellctl import ShellctlClientProtocol +from dify_agent.runtime_backend.protocols import RuntimeLayout +from dify_agent.runtime_backend.shellctl import ( + create_owned_shellctl_lease, + create_shellctl_lease, + run_shellctl_control_command, +) + + +@dataclass(slots=True) +class _FakeClient: + close_error: Exception | None = None + close_calls: int = 0 + + async def close(self) -> None: + self.close_calls += 1 + if self.close_error is not None: + raise self.close_error + + +@dataclass(slots=True) +class _FakeTransport: + close_error: Exception | None = None + close_calls: int = 0 + + async def aclose(self) -> None: + self.close_calls += 1 + if self.close_error is not None: + raise self.close_error + + +@dataclass(slots=True) +class _FakeCommands: + initial: ShellCommandResult + wait_error: Exception | None = None + delete_error: Exception | None = None + delete_calls: list[tuple[str, bool]] = field(default_factory=list) + + async def run( + self, + script: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout: float, + ) -> ShellCommandResult: + del script, cwd, env, timeout + return self.initial + + async def wait(self, job_id: str, *, offset: int, timeout: float) -> ShellCommandResult: + del job_id, offset, timeout + if self.wait_error is not None: + raise self.wait_error + raise AssertionError("wait was not expected") + + async def read_output(self, job_id: str, *, offset: int) -> ShellCommandResult: + raise AssertionError((job_id, offset)) + + async def input(self, job_id: str, text: str, *, offset: int, timeout: float) -> ShellCommandResult: + raise AssertionError((job_id, text, offset, timeout)) + + async def interrupt(self, job_id: str, *, grace_seconds: float) -> ShellCommandStatus: + raise AssertionError((job_id, grace_seconds)) + + async def tail(self, job_id: str) -> ShellCommandResult: + raise AssertionError(job_id) + + async def delete( + self, + job_id: str, + *, + force: bool = False, + grace_seconds: float | None = None, + ) -> None: + del grace_seconds + self.delete_calls.append((job_id, force)) + if self.delete_error is not None: + raise self.delete_error + + +def _result(*, done: bool = True) -> ShellCommandResult: + return ShellCommandResult( + job_id="job-1", + status="exited" if done else "running", + done=done, + exit_code=0 if done else None, + output="ok", + offset=2, + truncated=False, + ) + + +@pytest.mark.anyio +async def test_owned_transport_is_closed_exactly_once() -> None: + client = _FakeClient() + transport = _FakeTransport() + lease = create_shellctl_lease( + handle="sandbox-1", + layout=RuntimeLayout(home_dir="/home/dify", workspace_dir="/home/dify/workspace"), + entrypoint="http://shellctl", + token="secret", + client_factory=lambda: cast(ShellctlClientProtocol, cast(object, client)), + owned_transport=transport, + ) + + await lease.close() + await lease.close() + + assert client.close_calls == 1 + assert transport.close_calls == 1 + + +@pytest.mark.anyio +async def test_owned_transport_closes_when_client_close_fails_without_double_close() -> None: + client = _FakeClient(close_error=RuntimeError("client close failed")) + transport = _FakeTransport() + lease = create_shellctl_lease( + handle="sandbox-1", + layout=RuntimeLayout(home_dir="/home/dify", workspace_dir="/home/dify/workspace"), + entrypoint="http://shellctl", + token="secret", + client_factory=lambda: cast(ShellctlClientProtocol, cast(object, client)), + owned_transport=transport, + ) + + with pytest.raises(RuntimeError, match="client close failed"): + await lease.close() + await lease.close() + + assert client.close_calls == 1 + assert transport.close_calls == 1 + + +@pytest.mark.anyio +async def test_owned_transport_closes_when_client_construction_fails() -> None: + transport = _FakeTransport() + + def fail_factory() -> ShellctlClientProtocol: + raise RuntimeError("client construction failed") + + with pytest.raises(RuntimeError, match="client construction failed"): + _ = await create_owned_shellctl_lease( + handle="sandbox-1", + layout=RuntimeLayout(home_dir="/home/dify", workspace_dir="/home/dify/workspace"), + entrypoint="http://shellctl", + token="secret", + client_factory=fail_factory, + owned_transport=transport, + ) + + assert transport.close_calls == 1 + + +@pytest.mark.anyio +async def test_control_command_success_is_preserved_when_delete_fails( + caplog: pytest.LogCaptureFixture, +) -> None: + commands = _FakeCommands(initial=_result(), delete_error=RuntimeError("delete failed")) + + with caplog.at_level("WARNING", logger="dify_agent.runtime_backend.shellctl"): + result = await run_shellctl_control_command(commands, "true") + + assert result.output == "ok" + assert commands.delete_calls == [("job-1", True)] + assert "delete failed" in caplog.text + + +@pytest.mark.anyio +async def test_control_command_error_is_preserved_when_delete_also_fails( + caplog: pytest.LogCaptureFixture, +) -> None: + commands = _FakeCommands( + initial=_result(done=False), + wait_error=RuntimeError("command failed"), + delete_error=RuntimeError("delete failed"), + ) + + with caplog.at_level("WARNING", logger="dify_agent.runtime_backend.shellctl"): + with pytest.raises(RuntimeError, match="command failed"): + _ = await run_shellctl_control_command(commands, "false") + + assert commands.delete_calls == [("job-1", True)] + assert "delete failed" in caplog.text diff --git a/dify-agent/tests/local/dify_agent/server/test_app.py b/dify-agent/tests/local/dify_agent/server/test_app.py index f8f0bb57ffa..c68252d6c2b 100644 --- a/dify-agent/tests/local/dify_agent/server/test_app.py +++ b/dify-agent/tests/local/dify_agent/server/test_app.py @@ -8,8 +8,6 @@ import httpx import pytest from fastapi.testclient import TestClient -from dify_agent.adapters.shell.shellctl import ShellctlProvider - import dify_agent.server.app as app_module from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig from dify_agent.layers.execution_context.layer import DifyExecutionContextLayer @@ -17,6 +15,9 @@ from dify_agent.layers.knowledge.configs import DifyKnowledgeBaseLayerConfig from dify_agent.layers.knowledge.layer import DifyKnowledgeBaseLayer from dify_agent.layers.shell import DifyShellLayerConfig from dify_agent.layers.shell.layer import DifyShellLayer +from dify_agent.layers.runtime import DifyRuntimeLayerConfig +from dify_agent.layers.runtime.layer import DifyRuntimeLayer +from dify_agent.runtime_backend.local import LocalExecutionBindingBackend from dify_agent.runtime.compositor_factory import DifyAgentLayerProvider from dify_agent.server.app import create_app, create_dify_api_inner_http_client, create_plugin_daemon_http_client from dify_agent.server.settings import ServerSettings @@ -115,16 +116,6 @@ class FakePluginDaemonHttpClient: self.is_closed = True -class FakeAgentStubGRPCServer: - closed: bool - - def __init__(self) -> None: - self.closed = False - - async def aclose(self) -> None: - self.closed = True - - class FakeTimeout: connect: float read: float @@ -191,8 +182,9 @@ def test_create_app_creates_scheduler_and_closes_after_shutdown(monkeypatch: pyt plugin_daemon_api_key="daemon-secret", inner_api_url="http://dify-api", inner_api_key="inner-secret", - shellctl_entrypoint="http://shellctl", - shellctl_auth_token="shell-secret", + sandbox_files_base_url="http://api:5001", + local_sandbox_endpoint="http://shellctl", + local_sandbox_auth_token="shell-secret", agent_stub_api_base_url="https://agent.example.com/agent-stub", server_secret_key=_base64url_secret(b"1" * 32), outbound_http_connect_timeout=1, @@ -229,7 +221,9 @@ def test_create_app_creates_scheduler_and_closes_after_shutdown(monkeypatch: pyt assert execution_context_layer.daemon_api_key == "daemon-secret" assert shell_layer.agent_stub_token_factory is not None token = shell_layer.agent_stub_token_factory(_execution_context(), session_id="abc12ff") - decoded = settings.create_agent_stub_token_codec().decode_token(token) + token_codec = settings.create_agent_stub_token_codec() + assert token_codec is not None + decoded = token_codec.decode_token(token) assert decoded.execution_context == _execution_context() assert decoded.session_id == "abc12ff" knowledge_provider = next(provider for provider in layer_providers if provider.type_id == "dify.knowledge_base") @@ -251,7 +245,10 @@ def test_create_app_creates_scheduler_and_closes_after_shutdown(monkeypatch: pyt assert isinstance(knowledge_layer, DifyKnowledgeBaseLayer) assert knowledge_layer.inner_api_url == "http://dify-api" assert knowledge_layer.inner_api_key == "inner-secret" - assert isinstance(shell_layer.shell_provider, ShellctlProvider) + runtime_provider = next(provider for provider in layer_providers if provider.type_id == "dify.runtime") + runtime_layer = runtime_provider.create_layer(DifyRuntimeLayerConfig(backend_binding_ref="binding-1")) + assert isinstance(runtime_layer, DifyRuntimeLayer) + assert isinstance(runtime_layer.backend, LocalExecutionBindingBackend) assert shell_layer.agent_stub_api_base_url == "https://agent.example.com/agent-stub" http_client = scheduler.plugin_daemon_http_client assert http_client is fake_http_client @@ -273,6 +270,15 @@ def test_create_app_creates_scheduler_and_closes_after_shutdown(monkeypatch: pyt getattr(route, "path", None) == "/agent-stub/drive/manifest" for route in create_app(settings).routes ) assert any(getattr(route, "path", None) == "/agent-stub/drive/commit" for route in create_app(settings).routes) + route_paths = create_app(settings).openapi()["paths"] + assert { + "/execution-bindings/files/list", + "/execution-bindings/files/read", + "/execution-bindings/files/download", + }.issubset(route_paths) + assert "/workspace/files/list" not in route_paths + assert "/workspace/files/read" not in route_paths + assert "/workspace/files/upload" not in route_paths assert FakeRunScheduler.created[0].shutdown_called is True assert FakeRunScheduler.created[0].dify_api_http_client.is_closed is True @@ -314,6 +320,7 @@ def test_create_app_wires_authenticated_agent_stub_file_upload_route(monkeypatch server_secret_key=_base64url_secret(b"1" * 32), inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://files.example.com", ) token_codec = settings.create_agent_stub_token_codec() assert token_codec is not None @@ -322,9 +329,9 @@ def test_create_app_wires_authenticated_agent_stub_file_upload_route(monkeypatch original_async_client = httpx.AsyncClient def handler(request: httpx.Request) -> httpx.Response: - assert str(request.url) == "https://api.example.com/inner/api/upload/file/request" + assert str(request.url) == "https://api.example.com/inner/api/agent/files/upload-request" assert request.headers["X-Inner-Api-Key"] == "inner-secret" - return httpx.Response(200, json={"data": {"url": "https://files.example.com/upload"}}) + return httpx.Response(200, json={"upload_uri": "/files/upload/for-plugin?sign=1"}) monkeypatch.setattr( "dify_agent.agent_stub.server.agent_stub_files.httpx.AsyncClient", @@ -339,7 +346,7 @@ def test_create_app_wires_authenticated_agent_stub_file_upload_route(monkeypatch ) assert response.status_code == 200 - assert response.json() == {"upload_url": "https://files.example.com/upload"} + assert response.json() == {"upload_url": "https://files.example.com/files/upload/for-plugin?sign=1"} assert FakeRunScheduler.created[0].shutdown_called is True assert fake_http_client.is_closed is True assert fake_redis.closed is True @@ -353,6 +360,7 @@ def test_create_app_wires_authenticated_agent_stub_drive_manifest_route(monkeypa server_secret_key=_base64url_secret(b"1" * 32), inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://files.example.com", ) token_codec = settings.create_agent_stub_token_codec() assert token_codec is not None @@ -403,34 +411,6 @@ def test_create_app_wires_authenticated_agent_stub_drive_manifest_route(monkeypa assert fake_redis.closed is True -def test_create_app_starts_and_stops_agent_stub_grpc_server_for_grpc_url(monkeypatch: pytest.MonkeyPatch) -> None: - fake_redis, fake_http_client = _patch_app_lifecycle(monkeypatch) - started: dict[str, object] = {} - fake_grpc_server = FakeAgentStubGRPCServer() - - async def fake_start_agent_stub_grpc_server(**kwargs): - started.update(kwargs) - return fake_grpc_server - - monkeypatch.setattr(app_module, "start_agent_stub_grpc_server", fake_start_agent_stub_grpc_server) - - settings = ServerSettings( - redis_url="redis://example.invalid/0", - agent_stub_api_base_url="grpc://agent.example.com:9091", - agent_stub_grpc_bind_address="0.0.0.0:9191", - server_secret_key=_base64url_secret(b"1" * 32), - ) - - with TestClient(create_app(settings)): - assert started["public_url"] == "grpc://agent.example.com:9091" - assert started["bind_address"] == "0.0.0.0:9191" - - assert fake_grpc_server.closed is True - assert FakeRunScheduler.created[0].shutdown_called is True - assert fake_http_client.is_closed is True - assert fake_redis.closed is True - - def test_create_plugin_daemon_http_client_uses_generic_outbound_httpx_construction_args( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/dify-agent/tests/local/dify_agent/server/test_auth.py b/dify-agent/tests/local/dify_agent/server/test_auth.py new file mode 100644 index 00000000000..fa561b3bf56 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/server/test_auth.py @@ -0,0 +1,57 @@ +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from dify_agent.server.auth import create_bearer_token_dependency + + +def _build_app(expected_token: str | None) -> FastAPI: + app = FastAPI() + dep = create_bearer_token_dependency(expected_token) + + @app.get("/protected", dependencies=[dep]) + async def protected() -> dict[str, str]: + return {"status": "ok"} + + return app + + +class TestBearerTokenAuthEnabled: + """Auth is enforced when a non-None token is configured.""" + + def test_missing_header_returns_401(self) -> None: + client = TestClient(_build_app("secret-token")) + response = client.get("/protected") + assert response.status_code == 401 + assert "missing" in response.json()["detail"] + + def test_invalid_scheme_returns_401(self) -> None: + client = TestClient(_build_app("secret-token")) + response = client.get("/protected", headers={"Authorization": "Basic abc"}) + assert response.status_code == 401 + assert "scheme" in response.json()["detail"] + + def test_wrong_token_returns_401(self) -> None: + client = TestClient(_build_app("secret-token")) + response = client.get("/protected", headers={"Authorization": "Bearer wrong"}) + assert response.status_code == 401 + assert "invalid bearer token" in response.json()["detail"] + + def test_correct_token_passes(self) -> None: + client = TestClient(_build_app("secret-token")) + response = client.get("/protected", headers={"Authorization": "Bearer secret-token"}) + assert response.status_code == 200 + assert response.json() == {"status": "ok"} + + +class TestBearerTokenAuthDisabled: + """Auth is a no-op when expected_token is None (backward compatibility).""" + + def test_no_header_passes_when_token_unconfigured(self) -> None: + client = TestClient(_build_app(None)) + response = client.get("/protected") + assert response.status_code == 200 + + def test_any_header_passes_when_token_unconfigured(self) -> None: + client = TestClient(_build_app(None)) + response = client.get("/protected", headers={"Authorization": "Bearer anything"}) + assert response.status_code == 200 diff --git a/dify-agent/tests/local/dify_agent/server/test_binding_files.py b/dify-agent/tests/local/dify_agent/server/test_binding_files.py new file mode 100644 index 00000000000..21e93ea334b --- /dev/null +++ b/dify-agent/tests/local/dify_agent/server/test_binding_files.py @@ -0,0 +1,765 @@ +from __future__ import annotations + +import asyncio +import base64 +import json +import os +import shlex +from contextlib import suppress +from dataclasses import dataclass, field +from pathlib import Path +from typing import Literal, cast +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from dify_agent.adapters.shell.protocols import ShellCommandResult, ShellProviderError +from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig +from dify_agent.protocol import BindingFileDownloadRequest, BindingFileListRequest, BindingFileReadRequest +from dify_agent.runtime_backend import BindingAcquireError, BindingLostError, RuntimeLayout, RuntimeLease +from dify_agent.server import binding_files as binding_files_module +from dify_agent.server.binding_files import BindingFileError, BindingFileService, resolve_binding_path +from dify_agent.server.routes.binding_files import create_binding_files_router + +_REFERENCE = "dify-file-ref:eyJyZWNvcmRfaWQiOiJ0b29sLTEifQ==" + + +def _framed(payload: dict[str, object]) -> str: + encoded = base64.b64encode(json.dumps(payload).encode()).decode() + return f"<<>>{encoded}<<>>" + + +@dataclass(slots=True) +class _Commands: + outputs: list[str] + exit_codes: list[int] = field(default_factory=list) + calls: list[tuple[str, str | None, dict[str, str] | None, float]] = field(default_factory=list) + deletes: list[str] = field(default_factory=list) + + async def run(self, script: str, *, cwd: str | None = None, env=None, timeout: float) -> ShellCommandResult: + self.calls.append((script, cwd, env, timeout)) + output = self.outputs.pop(0) + exit_code = self.exit_codes.pop(0) if self.exit_codes else 0 + return ShellCommandResult( + job_id=f"job-{len(self.calls)}", + status="exited", + done=True, + exit_code=exit_code, + output=output, + offset=len(output), + truncated=False, + ) + + async def wait(self, job_id: str, *, offset: int, timeout: float) -> ShellCommandResult: + raise AssertionError("unexpected wait") + + async def read_output(self, job_id: str, *, offset: int): + raise AssertionError("unexpected read_output") + + async def input(self, job_id: str, text: str, *, offset: int, timeout: float): + raise AssertionError("unexpected input") + + async def interrupt(self, job_id: str, *, grace_seconds: float): + raise AssertionError("unexpected interrupt") + + async def tail(self, job_id: str): + raise AssertionError("unexpected tail") + + async def delete(self, job_id: str, *, force: bool = False, grace_seconds: float | None = None) -> None: + assert force is True + self.deletes.append(job_id) + + +@dataclass(slots=True) +class _ProviderErrorCommands(_Commands): + phase: Literal["run", "wait"] = "run" + error_code: str = "timeout" + + async def run(self, script: str, *, cwd: str | None = None, env=None, timeout: float) -> ShellCommandResult: + self.calls.append((script, cwd, env, timeout)) + if self.phase == "run": + raise ShellProviderError("shell provider failed", code=self.error_code) + return ShellCommandResult( + job_id="job-1", + status="running", + done=False, + exit_code=None, + output="", + offset=0, + truncated=False, + ) + + async def wait(self, job_id: str, *, offset: int, timeout: float) -> ShellCommandResult: + assert (job_id, offset) == ("job-1", 0) + assert timeout > 0 + raise ShellProviderError("shell provider failed", code=self.error_code) + + +@dataclass(slots=True) +class _LocalCommands(_Commands): + async def run(self, script: str, *, cwd: str | None = None, env=None, timeout: float) -> ShellCommandResult: + self.calls.append((script, cwd, env, timeout)) + process = await asyncio.create_subprocess_shell( + script, + cwd=cwd, + env={**os.environ, **(env or {})}, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.STDOUT, + ) + output, _ = await asyncio.wait_for(process.communicate(), timeout=timeout) + text = output.decode(errors="replace") + return ShellCommandResult( + job_id=f"local-job-{len(self.calls)}", + status="exited", + done=True, + exit_code=process.returncode, + output=text, + offset=len(text), + truncated=False, + ) + + +@dataclass(slots=True) +class _Lease: + commands: _Commands + layout: RuntimeLayout = RuntimeLayout(home_dir="/home/agent", workspace_dir="/workspace") + + +@dataclass(slots=True) +class _Backend: + lease: RuntimeLease + acquired: list[str] = field(default_factory=list) + releases: int = 0 + + async def acquire(self, binding_ref: str) -> RuntimeLease: + self.acquired.append(binding_ref) + return self.lease + + async def release(self, lease: RuntimeLease) -> None: + assert lease is self.lease + self.releases += 1 + + +def _context() -> DifyExecutionContextLayerConfig: + return DifyExecutionContextLayerConfig( + tenant_id="tenant-1", + user_id="account-1", + user_from="account", + agent_mode="agent_app", + invoke_from="debugger", + ) + + +def _service(commands: _Commands, *, configured: bool = True) -> tuple[BindingFileService, _Backend]: + backend = _Backend(lease=cast(RuntimeLease, _Lease(commands=commands))) + service = BindingFileService( + execution_bindings=backend, # pyright: ignore[reportArgumentType] + agent_stub_api_base_url="http://stub/agent-stub" if configured else None, + agent_stub_token_factory=(lambda execution_context, *, session_id: "secret-jwe") if configured else None, + ) + return service, backend + + +def _local_service(tmp_path: Path) -> tuple[BindingFileService, _Backend, _LocalCommands, Path, Path]: + workspace = tmp_path / "workspace" + home = tmp_path / "home" + workspace.mkdir() + home.mkdir() + commands = _LocalCommands(outputs=[]) + backend = _Backend( + lease=cast( + RuntimeLease, + _Lease( + commands=commands, + layout=RuntimeLayout(home_dir=str(home), workspace_dir=str(workspace)), + ), + ) + ) + service = BindingFileService( + execution_bindings=backend, # pyright: ignore[reportArgumentType] + agent_stub_api_base_url=None, + agent_stub_token_factory=None, + ) + return service, backend, commands, workspace, home + + +def test_resolve_binding_path_supports_workspace_home_absolute_and_parent_paths() -> None: + layout = RuntimeLayout(home_dir="/home/agent", workspace_dir="/workspace") + + assert resolve_binding_path("", layout) == "/workspace" + assert resolve_binding_path("reports/out.csv", layout) == "/workspace/reports/out.csv" + assert resolve_binding_path("~", layout) == "/home/agent" + assert resolve_binding_path("~/outputs/out.csv", layout) == "/home/agent/outputs/out.csv" + assert resolve_binding_path("/var/data/out.csv", layout) == "/var/data/out.csv" + assert resolve_binding_path("../shared/out.csv", layout) == "/shared/out.csv" + + +@pytest.mark.anyio +async def test_list_and_read_use_commands_with_consistent_binding_paths_and_release_leases() -> None: + commands = _Commands( + outputs=[ + _framed( + { + "path": "reports", + "entries": [{"name": "note.txt", "type": "file", "size": 4, "mtime": 1}], + "truncated": False, + } + ), + _framed( + { + "path": "~/note.txt", + "size": 4, + "truncated": False, + "binary": False, + "text": "note", + } + ), + ] + ) + service, backend = _service(commands) + + listing = await service.list_files(BindingFileListRequest(backend_binding_ref="binding-ref", path="reports")) + preview = await service.read_file(BindingFileReadRequest(backend_binding_ref="binding-ref", path="~/note.txt")) + + assert listing.entries[0].name == "note.txt" + assert preview.text == "note" + assert "/workspace/reports" in commands.calls[0][0] + assert "/home/agent/note.txt" in commands.calls[1][0] + assert all(call[1] == "/workspace" for call in commands.calls) + assert all(call[2] == {"HOME": "/home/agent"} for call in commands.calls) + assert backend.acquired == ["binding-ref", "binding-ref"] + assert backend.releases == 2 + assert commands.deletes == ["job-1", "job-2"] + + +@pytest.mark.anyio +async def test_real_list_script_caps_1001_entries_at_1000(tmp_path: Path) -> None: + service, backend, commands, workspace, _ = _local_service(tmp_path) + for index in range(1001): + (workspace / f"{index:04d}.txt").write_bytes(b"x") + + listing = await service.list_files(BindingFileListRequest(backend_binding_ref="binding-ref", path=".")) + + assert len(listing.entries) == 1000 + assert listing.entries[0].name == "0000.txt" + assert listing.entries[-1].name == "0999.txt" + assert listing.truncated is True + assert commands.deletes == ["local-job-1"] + assert backend.releases == 1 + + +@pytest.mark.anyio +async def test_real_read_script_handles_boundary_truncation_and_binary(tmp_path: Path) -> None: + service, backend, commands, workspace, _ = _local_service(tmp_path) + (workspace / "boundary.txt").write_bytes(b"a" * 262144) + (workspace / "truncated.txt").write_bytes(b"b" * 262145) + (workspace / "binary.bin").write_bytes(b"\xff\x00") + + boundary = await service.read_file( + BindingFileReadRequest(backend_binding_ref="binding-ref", path="boundary.txt", max_bytes=262144) + ) + truncated = await service.read_file( + BindingFileReadRequest(backend_binding_ref="binding-ref", path="truncated.txt", max_bytes=262144) + ) + binary = await service.read_file( + BindingFileReadRequest(backend_binding_ref="binding-ref", path="binary.bin", max_bytes=262144) + ) + + assert boundary.size == 262144 + assert boundary.truncated is False + assert boundary.binary is False + assert boundary.text == "a" * 262144 + assert truncated.size == 262145 + assert truncated.truncated is True + assert truncated.binary is False + assert truncated.text == "b" * 262144 + assert binary.size == 2 + assert binary.truncated is False + assert binary.binary is True + assert binary.text is None + assert commands.deletes == ["local-job-1", "local-job-2", "local-job-3"] + assert backend.releases == 3 + + +@pytest.mark.anyio +async def test_real_browse_script_output_over_command_cap_normalizes_to_unavailable( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + service, backend, commands, workspace, _ = _local_service(tmp_path) + (workspace / "report.txt").write_bytes(b"report") + monkeypatch.setattr(binding_files_module, "_BROWSE_OUTPUT_MAX_BYTES", 64) + + with pytest.raises(BindingFileError) as exc_info: + await service.list_files(BindingFileListRequest(backend_binding_ref="binding-ref", path=".")) + + assert exc_info.value.code == "binding_unavailable" + assert exc_info.value.status_code == 502 + assert commands.deletes == ["local-job-1"] + assert backend.releases == 1 + + +@pytest.mark.parametrize("operation", ["list", "read"]) +@pytest.mark.anyio +async def test_browse_preserves_binding_file_error_and_releases_lease(operation: str) -> None: + commands = _Commands(outputs=["FileNotFoundError: missing"], exit_codes=[1]) + service, backend = _service(commands) + + with pytest.raises(BindingFileError) as exc_info: + if operation == "list": + await service.list_files(BindingFileListRequest(backend_binding_ref="binding-ref", path="missing")) + else: + await service.read_file(BindingFileReadRequest(backend_binding_ref="binding-ref", path="missing")) + + assert exc_info.value.code == "invalid_binding_path" + assert exc_info.value.status_code == 400 + assert backend.releases == 1 + + +@pytest.mark.parametrize("operation", ["list", "read"]) +@pytest.mark.anyio +async def test_browse_maps_malformed_backend_response_to_unavailable_and_releases_lease(operation: str) -> None: + commands = _Commands(outputs=[_framed({"path": "."})]) + service, backend = _service(commands) + + with pytest.raises(BindingFileError) as exc_info: + if operation == "list": + await service.list_files(BindingFileListRequest(backend_binding_ref="binding-ref", path=".")) + else: + await service.read_file(BindingFileReadRequest(backend_binding_ref="binding-ref", path="report.txt")) + + assert exc_info.value.code == "binding_unavailable" + assert exc_info.value.status_code == 502 + assert backend.releases == 1 + + +@pytest.mark.parametrize("missing_field", ["user_id", "user_from"]) +@pytest.mark.anyio +async def test_download_rejects_each_missing_identity_before_token_or_lease(missing_field: str) -> None: + context = _context().model_copy(update={missing_field: None}) + request = BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + execution_context=context, + ) + commands = _Commands(outputs=[]) + service, backend = _service(commands) + issued_tokens: list[tuple[DifyExecutionContextLayerConfig, str | None]] = [] + + def issue_token(execution_context: DifyExecutionContextLayerConfig, *, session_id: str | None) -> str: + issued_tokens.append((execution_context, session_id)) + return "must-not-be-issued" + + service.agent_stub_token_factory = issue_token + + with pytest.raises(BindingFileError) as identity_error: + await service.download_file(request) + + assert identity_error.value.code == "invalid_execution_context" + assert identity_error.value.status_code == 400 + assert issued_tokens == [] + assert backend.acquired == [] + assert commands.calls == [] + + +@pytest.mark.anyio +async def test_download_rejects_missing_configuration_before_acquiring_lease() -> None: + request = BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + execution_context=_context(), + ) + commands = _Commands(outputs=[]) + unavailable_service, unavailable_backend = _service(commands, configured=False) + + with pytest.raises(BindingFileError) as unavailable_error: + await unavailable_service.download_file(request) + + assert unavailable_error.value.code == "agent_stub_upload_unavailable" + assert unavailable_error.value.status_code == 503 + assert unavailable_backend.acquired == [] + assert commands.calls == [] + + +@pytest.mark.anyio +async def test_download_shell_quotes_resolved_path_and_returns_only_reference_inside_lease() -> None: + commands = _Commands( + outputs=[json.dumps({"transfer_method": "tool_file", "reference": _REFERENCE, "public_download_url": "bad"})] + ) + service, backend = _service(commands) + issued_tokens: list[tuple[DifyExecutionContextLayerConfig, str | None]] = [] + + def issue_token(execution_context: DifyExecutionContextLayerConfig, *, session_id: str | None) -> str: + issued_tokens.append((execution_context, session_id)) + return "secret-jwe" + + service.agent_stub_token_factory = issue_token + context = _context() + + result = await service.download_file( + BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path="../shared/report $(touch should-not-run); final.txt", + execution_context=context, + ) + ) + + assert result.reference == _REFERENCE + script, cwd, env, timeout = commands.calls[0] + assert shlex.split(script) == [ + "dify-agent", + "file", + "upload", + "--no-download-link", + "/shared/report $(touch should-not-run); final.txt", + ] + assert cwd == "/workspace" + assert env == { + "HOME": "/home/agent", + "DIFY_AGENT_STUB_API_BASE_URL": "http://stub/agent-stub", + "DIFY_AGENT_STUB_AUTH_JWE": "secret-jwe", + "DIFY_AGENT_STUB_DRIVE_BASE": "/mnt/drive", + } + assert timeout == pytest.approx(60.0, rel=0, abs=0.01) + assert issued_tokens == [(context, None)] + assert issued_tokens[0][0].model_dump() == context.model_dump() + assert backend.releases == 1 + + +@pytest.mark.parametrize( + ("output", "exit_code"), + [ + ("upload failed", 1), + ("not-json", 0), + (json.dumps({"transfer_method": "url", "reference": _REFERENCE}), 0), + (json.dumps({"transfer_method": "tool_file", "reference": "raw-id"}), 0), + ("x" * (32 * 1024 + 1), 0), + ], +) +@pytest.mark.anyio +async def test_download_normalizes_cli_failures_and_releases_lease(output: str, exit_code: int) -> None: + commands = _Commands(outputs=[output], exit_codes=[exit_code]) + service, backend = _service(commands) + + with pytest.raises(BindingFileError) as exc_info: + await service.download_file( + BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + execution_context=_context(), + ) + ) + + assert exc_info.value.code == "binding_file_download_failed" + assert backend.releases == 1 + + +@pytest.mark.anyio +async def test_download_command_timeout_releases_lease_and_returns_download_failed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + commands = _Commands(outputs=[]) + service, backend = _service(commands) + + async def timed_out(*_args, **_kwargs): + return type( + "TimedOutResult", + (), + {"exit_code": None, "output_complete": False, "output": "", "incomplete_reason": "timeout"}, + )() + + monkeypatch.setattr(binding_files_module, "execute_complete_with_commands", timed_out) + + with pytest.raises(BindingFileError) as exc_info: + await service.download_file( + BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + execution_context=_context(), + ) + ) + + assert exc_info.value.code == "binding_file_download_failed" + assert exc_info.value.status_code == 502 + assert backend.releases == 1 + + +@pytest.mark.parametrize( + ("phase", "error_code", "expected_code", "expected_deletes"), + [ + ("run", "timeout", "binding_file_download_failed", []), + ("wait", "timeout", "binding_file_download_failed", ["job-1"]), + ("wait", "request_error", "binding_unavailable", ["job-1"]), + ], +) +@pytest.mark.anyio +async def test_download_maps_only_shell_provider_timeout_to_download_failed_and_cleans_up( + phase: Literal["run", "wait"], + error_code: str, + expected_code: str, + expected_deletes: list[str], +) -> None: + commands = _ProviderErrorCommands(outputs=[], phase=phase, error_code=error_code) + service, backend = _service(commands) + + with pytest.raises(BindingFileError) as exc_info: + await service.download_file( + BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + execution_context=_context(), + ) + ) + + assert exc_info.value.code == expected_code + assert exc_info.value.status_code == 502 + assert commands.deletes == expected_deletes + assert backend.acquired == ["binding-ref"] + assert backend.releases == 1 + + +@pytest.mark.anyio +async def test_download_cancellation_releases_lease(monkeypatch: pytest.MonkeyPatch) -> None: + commands = _Commands(outputs=[]) + service, backend = _service(commands) + command_started = asyncio.Event() + never_finishes = asyncio.Event() + + async def block(*_args, **_kwargs): + command_started.set() + await never_finishes.wait() + raise AssertionError("unreachable") + + monkeypatch.setattr(binding_files_module, "execute_complete_with_commands", block) + task = asyncio.create_task( + service.download_file( + BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + execution_context=_context(), + ) + ) + ) + try: + await asyncio.wait_for(command_started.wait(), timeout=1) + task.cancel() + + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=1) + finally: + if not task.done(): + task.cancel() + with suppress(asyncio.CancelledError): + await asyncio.wait_for(task, timeout=1) + + assert backend.releases == 1 + + +@pytest.mark.parametrize( + ("backend_error", "expected_code", "expected_status"), + [ + (BindingLostError("lost"), "binding_not_found", 404), + (BindingAcquireError("unavailable"), "binding_unavailable", 502), + ], +) +@pytest.mark.anyio +async def test_download_maps_binding_acquire_errors( + backend_error: Exception, + expected_code: str, + expected_status: int, +) -> None: + class FailingBackend: + releases = 0 + + async def acquire(self, _binding_ref: str) -> RuntimeLease: + raise backend_error + + async def release(self, _lease: RuntimeLease) -> None: + self.releases += 1 + + backend = FailingBackend() + service = BindingFileService( + execution_bindings=backend, # pyright: ignore[reportArgumentType] + agent_stub_api_base_url="http://stub/agent-stub", + agent_stub_token_factory=lambda execution_context, *, session_id: "secret-jwe", + ) + + with pytest.raises(BindingFileError) as exc_info: + await service.download_file( + BindingFileDownloadRequest( + backend_binding_ref="binding-ref", + path="report.txt", + execution_context=_context(), + ) + ) + + assert exc_info.value.code == expected_code + assert exc_info.value.status_code == expected_status + assert backend.releases == 0 + + +@pytest.mark.parametrize( + ("error", "expected_status", "expected_code"), + [ + (BindingFileError("binding_not_found", "missing", status_code=404), 404, "binding_not_found"), + (BindingFileError("binding_unavailable", "unavailable", status_code=502), 502, "binding_unavailable"), + ], +) +def test_binding_file_route_preserves_structured_error_status_and_code( + error: BindingFileError, + expected_status: int, + expected_code: str, +) -> None: + class FailingService: + async def download_file(self, _request): + raise error + + app = FastAPI() + app.include_router(create_binding_files_router(lambda: cast(BindingFileService, cast(object, FailingService())))) + + response = TestClient(app).post( + "/execution-bindings/files/download", + json={ + "backend_binding_ref": "binding-ref", + "path": "report.txt", + "execution_context": { + "tenant_id": "tenant-1", + "user_id": "account-1", + "user_from": "account", + "agent_mode": "agent_app", + "invoke_from": "debugger", + }, + }, + ) + + assert response.status_code == expected_status + assert response.json() == {"detail": {"code": expected_code, "message": error.message}} + + +@pytest.mark.parametrize( + ("path", "payload"), + [ + ("/execution-bindings/files/list", {"backend_binding_ref": "", "path": "."}), + ( + "/execution-bindings/files/read", + {"backend_binding_ref": "binding-ref", "path": "report.txt", "max_bytes": 262145}, + ), + ], +) +def test_list_and_read_route_validation_returns_structured_400_without_calling_service( + path: str, + payload: object, +) -> None: + service, backend = _service(_Commands(outputs=[])) + app = FastAPI() + app.include_router(create_binding_files_router(lambda: service)) + + with ( + patch.object(BindingFileService, "list_files", new_callable=AsyncMock) as list_files, + patch.object(BindingFileService, "read_file", new_callable=AsyncMock) as read_file, + ): + response = TestClient(app).post(path, json=payload) + + assert response.status_code == 400 + assert response.json() == { + "detail": { + "code": "invalid_binding_path", + "message": "Binding file path or payload is invalid", + } + } + list_files.assert_not_awaited() + read_file.assert_not_awaited() + assert backend.acquired == [] + assert backend.releases == 0 + + +def test_list_route_redacts_unexpected_field_value_from_validation_error() -> None: + service, backend = _service(_Commands(outputs=[])) + app = FastAPI() + app.include_router(create_binding_files_router(lambda: service)) + + with patch.object(BindingFileService, "list_files", new_callable=AsyncMock) as list_files: + response = TestClient(app).post( + "/execution-bindings/files/list", + json={"backend_binding_ref": "binding-ref", "path": ".", "unexpected": "top-secret"}, + ) + + assert response.status_code == 400 + assert response.json() == { + "detail": { + "code": "invalid_binding_path", + "message": "Binding file path or payload is invalid", + } + } + assert "top-secret" not in response.text + list_files.assert_not_awaited() + assert backend.acquired == [] + assert backend.releases == 0 + + +def test_download_route_validation_returns_422_without_calling_service() -> None: + service, backend = _service(_Commands(outputs=[])) + app = FastAPI() + app.include_router(create_binding_files_router(lambda: service)) + + with patch.object(BindingFileService, "download_file", new_callable=AsyncMock) as download_file: + download_response = TestClient(app).post( + "/execution-bindings/files/download", + json={"backend_binding_ref": "", "path": "", "execution_context": {}}, + ) + + assert download_response.status_code == 422 + assert isinstance(download_response.json()["detail"], list) + download_file.assert_not_awaited() + assert backend.acquired == [] + assert backend.releases == 0 + + +def test_read_route_accepts_preview_size_limit() -> None: + commands = _Commands( + outputs=[ + _framed( + { + "path": "report.txt", + "size": 262144, + "truncated": False, + "binary": False, + "text": "preview", + } + ) + ] + ) + service, backend = _service(commands) + app = FastAPI() + app.include_router(create_binding_files_router(lambda: service)) + + response = TestClient(app).post( + "/execution-bindings/files/read", + json={"backend_binding_ref": "binding-ref", "path": "report.txt", "max_bytes": 262144}, + ) + + assert response.status_code == 200 + assert response.json()["text"] == "preview" + assert "262144" in commands.calls[0][0] + assert backend.acquired == ["binding-ref"] + assert backend.releases == 1 + + +def test_binding_file_validation_override_preserves_openapi_request_schemas() -> None: + app = FastAPI() + app.include_router(create_binding_files_router(lambda: None)) + openapi = app.openapi() + paths = openapi["paths"] + + assert paths["/execution-bindings/files/list"]["post"]["requestBody"]["content"]["application/json"]["schema"] == { + "$ref": "#/components/schemas/BindingFileListRequest" + } + assert paths["/execution-bindings/files/read"]["post"]["requestBody"]["content"]["application/json"]["schema"] == { + "$ref": "#/components/schemas/BindingFileReadRequest" + } + max_bytes_schema = openapi["components"]["schemas"]["BindingFileReadRequest"]["properties"]["max_bytes"] + assert max_bytes_schema["default"] == 262144 + assert max_bytes_schema["minimum"] == 1 + assert max_bytes_schema["maximum"] == 262144 diff --git a/dify-agent/tests/local/dify_agent/server/test_execution_bindings.py b/dify-agent/tests/local/dify_agent/server/test_execution_bindings.py new file mode 100644 index 00000000000..4ff354be6c9 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/server/test_execution_bindings.py @@ -0,0 +1,66 @@ +from dataclasses import dataclass, field + +import pytest + +from dify_agent.protocol import CreateExecutionBindingRequest, DestroyExecutionBindingRequest +from dify_agent.runtime_backend import ( + ExecutionBindingAllocation, + ExecutionBindingCreateSpec, + ExecutionBindingDestroySpec, +) +from dify_agent.server.execution_bindings import ExecutionBindingService + + +@dataclass(slots=True) +class _Backend: + created: list[ExecutionBindingCreateSpec] = field(default_factory=list) + destroyed: list[ExecutionBindingDestroySpec] = field(default_factory=list) + + async def create_binding(self, spec: ExecutionBindingCreateSpec) -> ExecutionBindingAllocation: + self.created.append(spec) + return ExecutionBindingAllocation(binding_ref="opaque-binding", workspace_ref="opaque-workspace") + + async def destroy_binding(self, spec: ExecutionBindingDestroySpec) -> None: + self.destroyed.append(spec) + + +@pytest.mark.anyio +async def test_execution_binding_service_forwards_final_contract() -> None: + backend = _Backend() + service = ExecutionBindingService(backend=backend) # pyright: ignore[reportArgumentType] + + response = await service.create_binding( + CreateExecutionBindingRequest( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="home-ref", + ) + ) + await service.destroy_binding( + DestroyExecutionBindingRequest( + binding_ref=response.binding_ref, + workspace_ref=response.workspace_ref, + destroy_workspace=True, + ) + ) + + assert backend.created == [ + ExecutionBindingCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + binding_id="binding-1", + workspace_id="workspace-1", + existing_workspace_ref=None, + home_snapshot_ref="home-ref", + ) + ] + assert backend.destroyed == [ + ExecutionBindingDestroySpec( + binding_ref="opaque-binding", + workspace_ref="opaque-workspace", + destroy_workspace=True, + ) + ] diff --git a/dify-agent/tests/local/dify_agent/server/test_home_snapshots.py b/dify-agent/tests/local/dify_agent/server/test_home_snapshots.py new file mode 100644 index 00000000000..c82e06bfce0 --- /dev/null +++ b/dify-agent/tests/local/dify_agent/server/test_home_snapshots.py @@ -0,0 +1,106 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import cast + +import pytest + +from dify_agent.protocol import ( + CreateHomeSnapshotFromBindingRequest, + DeleteHomeSnapshotRequest, +) +from dify_agent.runtime_backend import HomeSnapshotCreateSpec, RuntimeLease +from dify_agent.server.home_snapshots import HomeSnapshotService + + +@dataclass(slots=True) +class _HomeBackend: + checkpointed: list[tuple[HomeSnapshotCreateSpec, RuntimeLease]] = field(default_factory=list) + deleted: list[str] = field(default_factory=list) + + async def create_from_runtime(self, *, spec: HomeSnapshotCreateSpec, source: RuntimeLease) -> str: + self.checkpointed.append((spec, source)) + return "snapshot-build" + + async def delete(self, snapshot_ref: str) -> None: + self.deleted.append(snapshot_ref) + + +@dataclass(slots=True) +class _BindingBackend: + lease: RuntimeLease + acquired: list[str] = field(default_factory=list) + released: list[RuntimeLease] = field(default_factory=list) + + async def acquire(self, binding_ref: str) -> RuntimeLease: + self.acquired.append(binding_ref) + return self.lease + + async def release(self, lease: RuntimeLease) -> None: + self.released.append(lease) + + +@pytest.mark.anyio +async def test_home_snapshot_service_checkpoints_exact_binding() -> None: + lease = cast(RuntimeLease, object()) + homes = _HomeBackend() + bindings = _BindingBackend(lease=lease) + service = HomeSnapshotService( + home_snapshots=homes, # pyright: ignore[reportArgumentType] + execution_bindings=bindings, # pyright: ignore[reportArgumentType] + ) + + checkpoint = await service.create_from_binding( + CreateHomeSnapshotFromBindingRequest( + tenant_id="tenant-1", + agent_id="agent-1", + home_snapshot_id="home-2", + backend_binding_ref="binding-ref", + ) + ) + + assert checkpoint.snapshot_ref == "snapshot-build" + assert bindings.acquired == ["binding-ref"] + assert bindings.released == [lease] + assert homes.checkpointed == [ + ( + HomeSnapshotCreateSpec( + tenant_id="tenant-1", + agent_id="agent-1", + home_snapshot_id="home-2", + ), + lease, + ) + ] + + await service.delete(DeleteHomeSnapshotRequest(snapshot_ref="snapshot-build")) + assert homes.deleted == ["snapshot-build"] + + +@pytest.mark.anyio +async def test_snapshot_checkpoint_releases_binding_when_create_fails() -> None: + lease = cast(RuntimeLease, object()) + + class _FailingHomeBackend(_HomeBackend): + async def create_from_runtime(self, *, spec: HomeSnapshotCreateSpec, source: RuntimeLease) -> str: + del spec, source + raise RuntimeError("checkpoint failed") + + homes = _FailingHomeBackend() + bindings = _BindingBackend(lease=lease) + service = HomeSnapshotService( + home_snapshots=homes, # pyright: ignore[reportArgumentType] + execution_bindings=bindings, # pyright: ignore[reportArgumentType] + ) + + with pytest.raises(RuntimeError, match="checkpoint failed"): + await service.create_from_binding( + CreateHomeSnapshotFromBindingRequest( + tenant_id="tenant-1", + agent_id="agent-1", + home_snapshot_id="home-2", + backend_binding_ref="binding-ref", + ) + ) + + assert bindings.released == [lease] diff --git a/dify-agent/tests/local/dify_agent/server/test_runs_routes.py b/dify-agent/tests/local/dify_agent/server/test_runs_routes.py index d90f82a152a..ce2c6247a7c 100644 --- a/dify-agent/tests/local/dify_agent/server/test_runs_routes.py +++ b/dify-agent/tests/local/dify_agent/server/test_runs_routes.py @@ -1,9 +1,10 @@ from fastapi.testclient import TestClient -from dify_agent.protocol import CancelRunResponse, DIFY_AGENT_MODEL_LAYER_ID +from dify_agent.protocol import CancelRunResponse, DIFY_AGENT_MODEL_LAYER_ID, RunFailureType from dify_agent.runtime.run_scheduler import RunCancellationConflictError, SchedulerStoppingError from dify_agent.server.routes.runs import create_runs_router from dify_agent.server.schemas import RunRecord +from dify_agent.storage.redis_run_store import RunNotFoundError class FakeScheduler: @@ -20,6 +21,29 @@ class FakeStore: pass +def test_get_run_status_returns_failure_type() -> None: + from fastapi import FastAPI + + class FailedRunStore: + async def get_run(self, run_id: str) -> RunRecord: + return RunRecord( + run_id=run_id, + status="failed", + error="run limit reached", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + ) + + app = FastAPI() + app.include_router( + create_runs_router(lambda: FailedRunStore(), lambda: FakeScheduler()) # pyright: ignore[reportArgumentType] + ) + response = TestClient(app).get("/runs/run-1") + + assert response.status_code == 200 + assert response.json()["error"] == "run limit reached" + assert response.json()["error_type"] == "agent_run_limit_exceeded" + + def test_create_run_accepts_effectively_blank_user_prompt_list() -> None: from fastapi import FastAPI @@ -106,6 +130,26 @@ def test_cancel_run_endpoint_maps_conflict() -> None: assert "already finished" in response.json()["detail"] +def test_cancel_run_endpoint_maps_missing_run() -> None: + from fastapi import FastAPI + + class MissingRunScheduler(FakeScheduler): + async def cancel_run(self, run_id: str, request: object) -> CancelRunResponse: + del request + raise RunNotFoundError(run_id) + + app = FastAPI() + app.include_router( + create_runs_router(lambda: FakeStore(), lambda: MissingRunScheduler()) # pyright: ignore[reportArgumentType] + ) + client = TestClient(app) + + response = client.post("/runs/missing/cancel", json={}) + + assert response.status_code == 404 + assert response.json() == {"detail": "run not found"} + + def test_create_run_accepts_valid_full_plugin_graph() -> None: from fastapi import FastAPI diff --git a/dify-agent/tests/local/dify_agent/server/test_sandbox_files.py b/dify-agent/tests/local/dify_agent/server/test_sandbox_files.py deleted file mode 100644 index 6bbe954ad7f..00000000000 --- a/dify-agent/tests/local/dify_agent/server/test_sandbox_files.py +++ /dev/null @@ -1,628 +0,0 @@ -from __future__ import annotations - -import asyncio -import base64 -import json -import os -from pathlib import Path -import subprocess -import sys -import types -from collections.abc import Callable, Mapping -from dataclasses import dataclass -from typing import Literal, cast - -import pytest -from agenton.compositor import CompositorSessionSnapshot, LayerProvider -from agenton.compositor.schemas import LayerSessionSnapshot -from agenton.layers.base import LifecycleState -from dify_agent.adapters.shell.shellctl import ShellctlClientProtocol, ShellctlProvider -from dify_agent.agent_stub.shell_env import ( - AGENT_STUB_API_BASE_URL_ENV_VAR, - AGENT_STUB_AUTH_JWE_ENV_VAR, - AGENT_STUB_DRIVE_BASE_ENV_VAR, -) - -if "graphon.model_runtime.entities.llm_entities" not in sys.modules: - graphon_module = types.ModuleType("graphon") - model_runtime_module = types.ModuleType("graphon.model_runtime") - entities_module = types.ModuleType("graphon.model_runtime.entities") - llm_entities_module = types.ModuleType("graphon.model_runtime.entities.llm_entities") - message_entities_module = types.ModuleType("graphon.model_runtime.entities.message_entities") - - llm_entities_module.LLMResultChunk = type("LLMResultChunk", (), {}) - llm_entities_module.LLMUsage = type("LLMUsage", (), {}) - - for name in ( - "AssistantPromptMessage", - "AudioPromptMessageContent", - "DocumentPromptMessageContent", - "ImagePromptMessageContent", - "PromptMessage", - "PromptMessageContentUnionTypes", - "PromptMessageTool", - "SystemPromptMessage", - "TextPromptMessageContent", - "ToolPromptMessage", - "UserPromptMessage", - "VideoPromptMessageContent", - ): - setattr(message_entities_module, name, type(name, (), {})) - - sys.modules["graphon"] = graphon_module - sys.modules["graphon.model_runtime"] = model_runtime_module - sys.modules["graphon.model_runtime.entities"] = entities_module - sys.modules["graphon.model_runtime.entities.llm_entities"] = llm_entities_module - sys.modules["graphon.model_runtime.entities.message_entities"] = message_entities_module - - graphon_module.model_runtime = model_runtime_module - model_runtime_module.entities = entities_module - entities_module.llm_entities = llm_entities_module - entities_module.message_entities = message_entities_module - -if "jsonschema" not in sys.modules: - jsonschema_module = types.ModuleType("jsonschema") - jsonschema_exceptions_module = types.ModuleType("jsonschema.exceptions") - jsonschema_protocols_module = types.ModuleType("jsonschema.protocols") - jsonschema_validators_module = types.ModuleType("jsonschema.validators") - - class _SchemaError(Exception): - pass - - class _ValidationError(Exception): - path: tuple[object, ...] = () - - class _Validator: - @staticmethod - def check_schema(schema): - return None - - def __init__(self, schema): - self.schema = schema - - def iter_errors(self, value): - return iter(()) - - def _validator_for(schema): - return _Validator - - jsonschema_module.SchemaError = _SchemaError - jsonschema_exceptions_module.ValidationError = _ValidationError - jsonschema_protocols_module.Validator = _Validator - jsonschema_validators_module.validator_for = _validator_for - - sys.modules["jsonschema"] = jsonschema_module - sys.modules["jsonschema.exceptions"] = jsonschema_exceptions_module - sys.modules["jsonschema.protocols"] = jsonschema_protocols_module - sys.modules["jsonschema.validators"] = jsonschema_validators_module - -from dify_agent.layers.execution_context import DifyExecutionContextLayerConfig -from dify_agent.layers.execution_context.layer import DifyExecutionContextLayer -from dify_agent.layers.shell import DifyShellLayerConfig -from dify_agent.layers.shell.layer import CompleteRemoteCommandResult, DifyShellLayer -from dify_agent.protocol import ( - CreateRunRequest, - RunComposition, - RunLayerSpec, - SandboxListRequest, - SandboxLocator, - SandboxReadRequest, - SandboxUploadRequest, - build_sandbox_locator_from_run_request, -) -from dify_agent.server.sandbox_files import ( - _LIST_SCRIPT, - SandboxFileError, - SandboxFileService, - _OUTPUT_BEGIN, - _OUTPUT_END, - _READ_SCRIPT, - _UPLOAD_SCRIPT, - _decode_sandbox_payload, - _shell_result_details, -) - - -@dataclass(slots=True) -class _Job: - job_id: str - status: str = "exited" - done: bool = True - exit_code: int | None = 0 - output: str = "" - offset: int = 0 - truncated: bool = False - output_path: str | None = "/tmp/sandbox-job.out" - - -@dataclass(slots=True) -class RunCall: - script: str - cwd: str | None - env: dict[str, str] | None - timeout: float - - -class FakeShellctlClient: - def __init__(self, *, run_handler: Callable[[str, str | None, dict[str, str] | None, float], _Job]) -> None: - self.run_handler = run_handler - self.run_calls: list[RunCall] = [] - self.delete_calls: list[str] = [] - - async def run( - self, script: str, *, cwd: str | None = None, env: dict[str, str] | None = None, timeout: float = 10.0 - ) -> _Job: - self.run_calls.append(RunCall(script=script, cwd=cwd, env=env, timeout=timeout)) - return self.run_handler(script, cwd, env, timeout) - - async def wait(self, job_id: str, *, offset: int, timeout: float = 10.0) -> _Job: - raise AssertionError(f"Unexpected wait() call for {job_id} offset={offset} timeout={timeout}") - - async def input(self, job_id: str, text: str, *, offset: int, timeout: float = 10.0) -> _Job: - raise AssertionError(f"Unexpected input() call for {job_id} text={text!r}") - - async def tail(self, job_id: str) -> _Job: - raise AssertionError(f"Unexpected tail() call for {job_id}") - - async def terminate(self, job_id: str, grace_seconds: float = 10.0) -> _Job: - raise AssertionError(f"Unexpected terminate() call for {job_id} grace={grace_seconds}") - - async def delete(self, job_id: str, *, force: bool = False, grace_seconds: float | None = None) -> None: - del force, grace_seconds - self.delete_calls.append(job_id) - return None - - async def close(self) -> None: - return None - - -def _wrap(payload: dict[str, object], *, pty_wrap: int = 0, noise: bool = False) -> str: - blob = base64.b64encode(json.dumps(payload).encode("utf-8")).decode("ascii") - if pty_wrap: - blob = "\n".join(blob[index : index + pty_wrap] for index in range(0, len(blob), pty_wrap)) - framed = f"{_OUTPUT_BEGIN}{blob}{_OUTPUT_END}\n" - if noise: - framed = f"user@host$ python3 - ...\r\n{framed}user@host$ \r\n" - return framed - - -def _complete_result( - *, - output: str, - exit_code: int | None = 0, - output_complete: bool = True, - incomplete_reason: Literal["output_limit", "timeout"] | None = None, - job_id: str = "sandbox-job", -) -> CompleteRemoteCommandResult: - return CompleteRemoteCommandResult( - job_id=job_id, - status="exited", - done=True, - exit_code=exit_code, - output=output, - output_complete=output_complete, - incomplete_reason=incomplete_reason, - offset=len(output), - output_path="/tmp/sandbox-job.out", - ) - - -def _run_embedded_script( - script_source: str, - *, - args: list[str], - cwd: Path, - env: Mapping[str, str] | None = None, -) -> dict[str, object]: - merged_env = dict(os.environ) - if env is not None: - merged_env.update(env) - completed = subprocess.run( - [sys.executable, "-", *args], - input=script_source, - text=True, - capture_output=True, - cwd=cwd, - env=merged_env, - check=False, - ) - return _decode_sandbox_payload(_complete_result(output=completed.stdout, exit_code=completed.returncode)) - - -def _execution_context() -> DifyExecutionContextLayerConfig: - return DifyExecutionContextLayerConfig( - tenant_id="tenant-1", - user_id="user-1", - user_from="account", - app_id="app-1", - conversation_id="conv-1", - agent_id="agent-1", - agent_config_version_id="snapshot-1", - agent_mode="agent_app", - invoke_from="service-api", - ) - - -def _locator() -> SandboxLocator: - request = CreateRunRequest( - composition=RunComposition( - layers=[ - RunLayerSpec(name="execution_context", type="dify.execution_context", config=_execution_context()), - RunLayerSpec( - name="shell", - type="dify.shell", - deps={"execution_context": "execution_context"}, - config=DifyShellLayerConfig(agent_stub_drive_ref="agent-1"), - ), - ] - ), - session_snapshot=CompositorSessionSnapshot( - layers=[ - LayerSessionSnapshot( - name="execution_context", lifecycle_state=LifecycleState.SUSPENDED, runtime_state={} - ), - LayerSessionSnapshot( - name="shell", - lifecycle_state=LifecycleState.SUSPENDED, - runtime_state={"session_id": "abc12ff", "workspace_cwd": "~/workspace/abc12ff"}, - ), - ] - ), - ) - return build_sandbox_locator_from_run_request(request) - - -def _service( - run_handler: Callable[[str, str | None, dict[str, str] | None, float], _Job], -) -> tuple[SandboxFileService, FakeShellctlClient]: - client = FakeShellctlClient(run_handler=run_handler) - execution_context_provider = LayerProvider.from_factory( - layer_type=DifyExecutionContextLayer, - create=lambda config: DifyExecutionContextLayer.from_config_with_settings( - DifyExecutionContextLayerConfig.model_validate(config), - daemon_url="http://plugin-daemon", - daemon_api_key="daemon-secret", - ), - ) - shell_provider = LayerProvider.from_factory( - layer_type=DifyShellLayer, - create=lambda config: DifyShellLayer.from_config_with_settings( - DifyShellLayerConfig.model_validate(config), - shell_provider=ShellctlProvider( - entrypoint="http://shellctl", - token="", - client_factory=lambda: cast(ShellctlClientProtocol, cast(object, client)), - ), - agent_stub_api_base_url="https://agent.example.com/agent-stub", - agent_stub_token_factory=lambda execution_context, *, session_id: ( - f"token-for:{execution_context.tenant_id}:{session_id}" - ), - ), - ) - return SandboxFileService(layer_providers=(execution_context_provider, shell_provider)), client - - -def _sandbox_python_run_call(client: FakeShellctlClient) -> RunCall: - for run_call in reversed(client.run_calls): - if run_call.script.startswith("python3 - "): - return run_call - raise AssertionError("sandbox python script was not executed") - - -def _sandbox_list_entries(payload: dict[str, object]) -> list[object]: - entries = payload.get("entries") - assert isinstance(entries, list) - return entries - - -def test_list_files_runs_fixed_script_and_parses_response() -> None: - service, client = _service( - lambda script, cwd, env, timeout: _Job( - job_id="sandbox-job", - output=_wrap( - { - "path": ".", - "entries": [{"name": "notes.txt", "type": "file", "size": 5, "mtime": 1}], - "truncated": False, - } - ), - ) - ) - - result = asyncio.run(service.list_files(SandboxListRequest(locator=_locator(), path="."))) - - assert result.entries[0].name == "notes.txt" - script_call = _sandbox_python_run_call(client) - assert script_call.cwd == "/home/agent-1/workspace/abc12ff" - assert script_call.env == {"HOME": "/home/agent-1"} - assert "python3 - . 1000 <<'PY'" in script_call.script - assert client.delete_calls[-1] == "sandbox-job" - - -def test_list_files_allows_parent_relative_paths() -> None: - service, client = _service( - lambda script, cwd, env, timeout: _Job( - job_id="sandbox-job", - output=_wrap({"path": "../shared", "entries": [], "truncated": False}), - ) - ) - - result = asyncio.run(service.list_files(SandboxListRequest(locator=_locator(), path="../shared"))) - - assert result.path == "../shared" - assert "python3 - ../shared 1000 <<'PY'" in _sandbox_python_run_call(client).script - - -@pytest.mark.parametrize( - ("path", "expected_command"), - [ - ("~", "python3 - '~' 1000 <<'PY'"), - ("~/shared", "python3 - '~/shared' 1000 <<'PY'"), - ], -) -def test_list_files_allows_home_relative_paths(path: str, expected_command: str) -> None: - service, client = _service( - lambda script, cwd, env, timeout: _Job( - job_id="sandbox-job", - output=_wrap({"path": path, "entries": [], "truncated": False}), - ) - ) - - result = asyncio.run(service.list_files(SandboxListRequest(locator=_locator(), path=path))) - - assert result.path == path - assert expected_command in _sandbox_python_run_call(client).script - - -def test_embedded_scripts_allow_parent_relative_paths(tmp_path: Path) -> None: - workspace_dir = tmp_path / "workspace" - cwd = workspace_dir / "run" - shared_dir = workspace_dir / "shared" - cwd.mkdir(parents=True) - shared_dir.mkdir() - notes_path = shared_dir / "notes.txt" - notes_path.write_text("hello", encoding="utf-8") - - bin_dir = tmp_path / "bin" - bin_dir.mkdir() - fake_dify_agent = bin_dir / "dify-agent" - fake_dify_agent.write_text( - "\n".join( - [ - "#!/usr/bin/env python3", - "import json", - "import sys", - 'if sys.argv[1:] != ["file", "upload", "../shared/notes.txt"]:', - ' raise SystemExit(f"unexpected args: {sys.argv[1:]!r}")', - 'print(json.dumps({"transfer_method": "tool_file", "reference": "file-ref", "download_url": "https://files.example.com/notes.txt"}))', - ] - ) - + "\n", - encoding="utf-8", - ) - fake_dify_agent.chmod(0o755) - - list_payload = _run_embedded_script(_LIST_SCRIPT, args=["../shared", "1000"], cwd=cwd) - read_payload = _run_embedded_script(_READ_SCRIPT, args=["../shared/notes.txt", "8"], cwd=cwd) - upload_payload = _run_embedded_script( - _UPLOAD_SCRIPT, - args=["../shared/notes.txt"], - cwd=cwd, - env={"PATH": f"{bin_dir}{os.pathsep}{os.environ.get('PATH', os.defpath)}"}, - ) - - assert list_payload["path"] == "../shared" - assert any( - isinstance(entry, dict) and entry.get("name") == "notes.txt" for entry in _sandbox_list_entries(list_payload) - ) - assert read_payload == { - "path": "../shared/notes.txt", - "size": 5, - "truncated": False, - "binary": False, - "text": "hello", - } - assert upload_payload == { - "path": "../shared/notes.txt", - "file": { - "transfer_method": "tool_file", - "reference": "file-ref", - "download_url": "https://files.example.com/notes.txt", - }, - } - - -def test_embedded_scripts_expand_home_relative_paths(tmp_path: Path) -> None: - cwd = tmp_path / "workspace" / "run" - cwd.mkdir(parents=True) - home_dir = tmp_path / "home" / "agent-1" - shared_dir = home_dir / "shared" - shared_dir.mkdir(parents=True) - notes_path = shared_dir / "notes.txt" - notes_path.write_text("hello", encoding="utf-8") - - bin_dir = tmp_path / "bin" - bin_dir.mkdir() - fake_dify_agent = bin_dir / "dify-agent" - fake_dify_agent.write_text( - "\n".join( - [ - "#!/usr/bin/env python3", - "import json", - "import sys", - 'if sys.argv[1:] != ["file", "upload", "~/shared/notes.txt"]:', - ' raise SystemExit(f"unexpected args: {sys.argv[1:]!r}")', - 'print(json.dumps({"transfer_method": "tool_file", "reference": "file-ref", "download_url": "https://files.example.com/notes.txt"}))', - ] - ) - + "\n", - encoding="utf-8", - ) - fake_dify_agent.chmod(0o755) - - script_env = {"HOME": str(home_dir), "PATH": f"{bin_dir}{os.pathsep}{os.environ.get('PATH', os.defpath)}"} - list_payload = _run_embedded_script(_LIST_SCRIPT, args=["~/shared", "1000"], cwd=cwd, env=script_env) - read_payload = _run_embedded_script(_READ_SCRIPT, args=["~/shared/notes.txt", "8"], cwd=cwd, env=script_env) - upload_payload = _run_embedded_script(_UPLOAD_SCRIPT, args=["~/shared/notes.txt"], cwd=cwd, env=script_env) - - assert list_payload["path"] == "~/shared" - assert any( - isinstance(entry, dict) and entry.get("name") == "notes.txt" for entry in _sandbox_list_entries(list_payload) - ) - assert read_payload == { - "path": "~/shared/notes.txt", - "size": 5, - "truncated": False, - "binary": False, - "text": "hello", - } - assert upload_payload == { - "path": "~/shared/notes.txt", - "file": { - "transfer_method": "tool_file", - "reference": "file-ref", - "download_url": "https://files.example.com/notes.txt", - }, - } - - -@pytest.mark.parametrize("bad_path", ["/etc/passwd", "~other/secret-dir", "bad\x00path"]) -def test_list_files_rejects_invalid_paths_before_shell_execution(bad_path: str) -> None: - service, client = _service(lambda script, cwd, env, timeout: _Job(job_id="sandbox-job", output="unused")) - - with pytest.raises(SandboxFileError, match="path"): - asyncio.run(service.list_files(SandboxListRequest(locator=_locator(), path=bad_path))) - - assert client.run_calls == [] - - -def test_decode_payload_reports_incomplete_capture_when_frame_is_missing() -> None: - with pytest.raises(SandboxFileError, match="incomplete before framed payload was captured"): - _decode_sandbox_payload( - _complete_result(output="partial", output_complete=False, incomplete_reason="output_limit") - ) - - -def test_decode_payload_reports_incomplete_capture_when_frame_is_corrupt() -> None: - broken = f"{_OUTPUT_BEGIN}%%%%{_OUTPUT_END}" - with pytest.raises(SandboxFileError, match="incomplete while decoding framed payload"): - _decode_sandbox_payload(_complete_result(output=broken, output_complete=False, incomplete_reason="timeout")) - - -def test_upload_injects_agent_stub_env_and_returns_mapping() -> None: - service, client = _service( - lambda script, cwd, env, timeout: _Job( - job_id="sandbox-job", - output=_wrap( - { - "path": "report.txt", - "file": { - "transfer_method": "tool_file", - "reference": "file-ref", - "download_url": "https://files.example.com/report.txt", - }, - }, - noise=True, - ), - ) - ) - - result = asyncio.run(service.upload_file(SandboxUploadRequest(locator=_locator(), path="report.txt"))) - - assert result.file.transfer_method == "tool_file" - assert result.file.reference == "file-ref" - assert result.file.download_url == "https://files.example.com/report.txt" - script_call = _sandbox_python_run_call(client) - assert script_call.cwd == "/home/agent-1/workspace/abc12ff" - assert script_call.env == { - "HOME": "/home/agent-1", - AGENT_STUB_API_BASE_URL_ENV_VAR: "https://agent.example.com/agent-stub", - AGENT_STUB_AUTH_JWE_ENV_VAR: "token-for:tenant-1:abc12ff", - AGENT_STUB_DRIVE_BASE_ENV_VAR: "/mnt/drive/agent-1", - } - - -def test_upload_rejects_missing_download_url_in_shell_payload() -> None: - service, _client = _service( - lambda script, cwd, env, timeout: _Job( - job_id="sandbox-job", - output=_wrap( - { - "path": "report.txt", - "file": { - "transfer_method": "tool_file", - "reference": "file-ref", - }, - } - ), - ) - ) - - with pytest.raises(SandboxFileError, match="sandbox command returned invalid payload"): - _ = asyncio.run(service.upload_file(SandboxUploadRequest(locator=_locator(), path="report.txt"))) - - -def test_shell_result_details_include_output_metadata_and_tail() -> None: - details = _shell_result_details( - _complete_result(output="hello", output_complete=False, incomplete_reason="output_limit") - ) - assert "output_complete=False" in details - assert "incomplete_reason=output_limit" in details - assert "output_path=/tmp/sandbox-job.out" in details - assert details.endswith("hello") - - -def test_read_file_uses_complete_mode_and_parses_response() -> None: - service, client = _service( - lambda script, cwd, env, timeout: _Job( - job_id="sandbox-job", - output=_wrap({"path": "notes.txt", "size": 5, "truncated": False, "binary": False, "text": "hello"}), - ) - ) - - result = asyncio.run(service.read_file(SandboxReadRequest(locator=_locator(), path="notes.txt", max_bytes=8))) - - assert result.text == "hello" - assert "python3 - notes.txt 8 <<'PY'" in _sandbox_python_run_call(client).script - - -@pytest.mark.parametrize( - ("sandbox_request", "expected_command"), - [ - (SandboxReadRequest(locator=_locator(), path="../notes.txt", max_bytes=8), "python3 - ../notes.txt 8 <<'PY'"), - (SandboxUploadRequest(locator=_locator(), path="../report.txt"), "python3 - ../report.txt <<'PY'"), - (SandboxReadRequest(locator=_locator(), path="~/notes.txt", max_bytes=8), "python3 - '~/notes.txt' 8 <<'PY'"), - (SandboxUploadRequest(locator=_locator(), path="~/report.txt"), "python3 - '~/report.txt' <<'PY'"), - ], -) -def test_read_and_upload_allow_relative_paths( - sandbox_request: SandboxReadRequest | SandboxUploadRequest, expected_command: str -) -> None: - expected_path = sandbox_request.path - service, client = _service( - lambda script, cwd, env, timeout: _Job( - job_id="sandbox-job", - output=_wrap( - {"path": expected_path, "size": 5, "truncated": False, "binary": False, "text": "hello"} - if isinstance(sandbox_request, SandboxReadRequest) - else { - "path": expected_path, - "file": { - "transfer_method": "tool_file", - "reference": "file-ref", - "download_url": "https://files.example.com/report.txt", - }, - } - ), - ) - ) - - if isinstance(sandbox_request, SandboxReadRequest): - result = asyncio.run(service.read_file(sandbox_request)) - assert result.path == expected_path - else: - result = asyncio.run(service.upload_file(sandbox_request)) - assert result.path == expected_path - assert result.file.download_url == "https://files.example.com/report.txt" - - assert expected_command in _sandbox_python_run_call(client).script diff --git a/dify-agent/tests/local/dify_agent/server/test_settings.py b/dify-agent/tests/local/dify_agent/server/test_settings.py index 25ae693963b..282542b4cca 100644 --- a/dify-agent/tests/local/dify_agent/server/test_settings.py +++ b/dify-agent/tests/local/dify_agent/server/test_settings.py @@ -8,12 +8,13 @@ import httpx import pytest from pydantic import ValidationError -from dify_agent.adapters.shell.enterprise import EnterpriseShellProvider -from dify_agent.adapters.shell.shellctl import ShellctlProvider from dify_agent.agent_stub.server.agent_stub_drive import DifyApiAgentStubDriveRequestHandler from dify_agent.agent_stub.server.agent_stub_files import DifyApiAgentStubFileRequestHandler from dify_agent.agent_stub.server.tokens.agent_stub import AgentStubTokenCodec from dify_agent.server.settings import ServerSettings +from dify_agent.runtime_backend.e2b import E2BExecutionBindingBackend +from dify_agent.runtime_backend.enterprise import EnterpriseExecutionBindingBackend +from dify_agent.runtime_backend.local import LocalExecutionBindingBackend, LocalHomeSnapshotBackend def _base64url_secret(value: bytes) -> str: @@ -27,7 +28,7 @@ def test_server_settings_reads_shellctl_entrypoint_from_env(monkeypatch: pytest. settings = ServerSettings() - assert settings.shellctl_entrypoint == "http://shellctl.example" + assert settings.local_sandbox_endpoint == "http://shellctl.example" def test_server_settings_reads_shellctl_auth_token_from_env(monkeypatch: pytest.MonkeyPatch) -> None: @@ -35,22 +36,7 @@ def test_server_settings_reads_shellctl_auth_token_from_env(monkeypatch: pytest. settings = ServerSettings() - assert settings.shellctl_auth_token == "shell-secret" - - -def test_server_settings_reads_shell_home_root_from_env(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("DIFY_AGENT_SHELL_HOME_ROOT", "/tmp/dify-agent-home/") - - settings = ServerSettings() - - assert settings.shell_home_root == "/tmp/dify-agent-home" - - -def test_server_settings_rejects_relative_shell_home_root(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setenv("DIFY_AGENT_SHELL_HOME_ROOT", "relative/path") - - with pytest.raises(ValidationError, match="DIFY_AGENT_SHELL_HOME_ROOT must be an absolute path"): - ServerSettings() + assert settings.local_sandbox_auth_token == "shell-secret" def test_server_settings_reads_enterprise_timeouts_from_env(monkeypatch: pytest.MonkeyPatch) -> None: @@ -63,6 +49,14 @@ def test_server_settings_reads_enterprise_timeouts_from_env(monkeypatch: pytest. assert settings.enterprise_sandbox_proxy_timeout == 90 +def test_server_settings_reads_e2b_active_timeout_from_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS", "900") + + settings = ServerSettings() + + assert settings.e2b_active_timeout_seconds == 900 + + def test_server_settings_defaults_shellctl_auth_token_to_none( monkeypatch: pytest.MonkeyPatch, tmp_path: Path, @@ -72,16 +66,18 @@ def test_server_settings_defaults_shellctl_auth_token_to_none( settings = ServerSettings() - assert settings.shellctl_auth_token is None + assert settings.local_sandbox_auth_token is None def test_server_settings_reads_agent_stub_settings_from_env(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("DIFY_AGENT_STUB_API_BASE_URL", "https://agent.example.com/agent-stub/") + monkeypatch.setenv("DIFY_AGENT_SANDBOX_FILES_BASE_URL", "https://dify.example.com/prefix/") monkeypatch.setenv("DIFY_AGENT_SERVER_SECRET_KEY", _base64url_secret(secrets.token_bytes(32))) settings = ServerSettings() assert settings.agent_stub_api_base_url == "https://agent.example.com/agent-stub" + assert settings.sandbox_files_base_url == "https://dify.example.com/prefix" def test_server_settings_normalizes_agent_stub_service_root_from_env(monkeypatch: pytest.MonkeyPatch) -> None: @@ -131,26 +127,23 @@ def test_server_settings_rejects_public_agent_stub_api_base_url_without_secret_k _ = ServerSettings(agent_stub_api_base_url="https://agent.example.com/agent-stub") -def test_server_settings_accepts_grpc_agent_stub_api_base_url_and_bind_override() -> None: - settings = ServerSettings( - agent_stub_api_base_url="grpc://agent.example.com:9091", - agent_stub_grpc_bind_address="0.0.0.0:9191", - server_secret_key=_base64url_secret(secrets.token_bytes(32)), - ) - - assert settings.agent_stub_api_base_url == "grpc://agent.example.com:9091" - assert settings.agent_stub_grpc_bind_address == "0.0.0.0:9191" - - -def test_server_settings_rejects_grpc_bind_override_without_grpc_url() -> None: - with pytest.raises(ValidationError, match="grpc://"): +def test_server_settings_requires_sandbox_files_base_url_for_agent_stub_file_operations() -> None: + with pytest.raises(ValidationError, match="DIFY_AGENT_SANDBOX_FILES_BASE_URL"): _ = ServerSettings( agent_stub_api_base_url="https://agent.example.com/agent-stub", - agent_stub_grpc_bind_address="0.0.0.0:9191", + inner_api_key="inner-secret", server_secret_key=_base64url_secret(secrets.token_bytes(32)), ) +def test_server_settings_rejects_sandbox_files_base_url_query_or_fragment() -> None: + with pytest.raises(ValidationError, match="query string or fragment"): + _ = ServerSettings(sandbox_files_base_url="https://dify.example.com?x=1") + + with pytest.raises(ValidationError, match="query string or fragment"): + _ = ServerSettings(sandbox_files_base_url="https://dify.example.com#fragment") + + def test_server_settings_rejects_invalid_server_secret_key() -> None: with pytest.raises(ValidationError, match="32 decoded bytes"): _ = ServerSettings(server_secret_key=_base64url_secret(b"short")) @@ -216,6 +209,7 @@ def test_server_settings_create_agent_stub_file_request_handler_returns_handler_ settings = ServerSettings( inner_api_url="https://api.example.com", inner_api_key="inner-secret", + sandbox_files_base_url="https://sandbox-files.example.com/dify", ) handler = settings.create_agent_stub_file_request_handler() @@ -223,6 +217,7 @@ def test_server_settings_create_agent_stub_file_request_handler_returns_handler_ assert isinstance(handler, DifyApiAgentStubFileRequestHandler) assert handler.inner_api_url == "https://api.example.com" assert handler.inner_api_key == "inner-secret" + assert handler.sandbox_files_base_url == "https://sandbox-files.example.com/dify" def test_server_settings_create_agent_stub_drive_request_handler_returns_none_without_full_settings() -> None: @@ -251,61 +246,87 @@ def test_server_settings_create_agent_stub_drive_request_handler_returns_handler assert timeout.pool == 44 -def test_build_shell_provider_returns_none_when_shellctl_entrypoint_is_unset( +def test_build_runtime_backend_profile_returns_none_when_local_endpoint_is_unset( monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: monkeypatch.delenv("DIFY_AGENT_SHELLCTL_ENTRYPOINT", raising=False) monkeypatch.chdir(tmp_path) - assert ServerSettings().build_shell_provider() is None + assert ServerSettings().build_runtime_backend_profile() is None -def test_build_shell_provider_returns_shellctl_provider_when_configured() -> None: +def test_build_runtime_backend_profile_returns_local_drivers_when_configured() -> None: settings = ServerSettings( - shell_provider="shellctl", - shellctl_entrypoint="http://shellctl.example", - shellctl_auth_token="shell-secret", + runtime_backend="local", + local_sandbox_endpoint="http://shellctl.example", + local_sandbox_auth_token="shell-secret", + local_sandbox_materialized_home_root="/tmp/dify/homes", + local_sandbox_workspace_root="/tmp/dify/workspaces", + local_sandbox_home_snapshot_root="/tmp/dify/snapshots", ) - provider = settings.build_shell_provider() + profile = settings.build_runtime_backend_profile() - assert isinstance(provider, ShellctlProvider) - assert provider.entrypoint == "http://shellctl.example" - assert provider.token == "shell-secret" + assert profile is not None + assert isinstance(profile.execution_bindings, LocalExecutionBindingBackend) + assert isinstance(profile.home_snapshots, LocalHomeSnapshotBackend) + assert profile.execution_bindings.endpoint == "http://shellctl.example" + assert profile.execution_bindings.auth_token == "shell-secret" + assert profile.execution_bindings.materialized_home_root == "/tmp/dify/homes" + assert profile.execution_bindings.workspace_root == "/tmp/dify/workspaces" + assert profile.execution_bindings.snapshot_root == "/tmp/dify/snapshots" + assert profile.home_snapshots.snapshot_root == "/tmp/dify/snapshots" -def test_build_shell_provider_returns_enterprise_provider_when_selected() -> None: +def test_build_runtime_backend_profile_returns_enterprise_drivers_when_selected() -> None: settings = ServerSettings( - shell_provider="enterprise", + runtime_backend="enterprise", enterprise_sandbox_gateway_endpoint="https://gateway.example", enterprise_sandbox_gateway_auth_token="gateway-secret", enterprise_sandbox_gateway_timeout=45, enterprise_sandbox_proxy_timeout=90, ) - provider = settings.build_shell_provider() + profile = settings.build_runtime_backend_profile() - assert isinstance(provider, EnterpriseShellProvider) - assert provider.gateway_endpoint == "https://gateway.example" - assert provider.auth_token == "gateway-secret" - assert provider.gateway_timeout == 45 - assert provider.proxy_timeout == 90 + assert profile is not None + assert isinstance(profile.execution_bindings, EnterpriseExecutionBindingBackend) + assert profile.execution_bindings.gateway_endpoint == "https://gateway.example" + assert profile.execution_bindings.auth_token == "gateway-secret" + assert profile.execution_bindings.gateway_timeout == 45 + assert profile.execution_bindings.proxy_timeout == 90 -def test_build_shell_provider_returns_none_when_enterprise_endpoint_is_unset( +def test_build_runtime_backend_profile_passes_e2b_active_timeout() -> None: + settings = ServerSettings( + runtime_backend="e2b", + e2b_api_key="e2b-secret", + e2b_active_timeout_seconds=900, + ) + + profile = settings.build_runtime_backend_profile() + + assert profile is not None + assert isinstance(profile.execution_bindings, E2BExecutionBindingBackend) + assert profile.execution_bindings.active_timeout_seconds == 900 + assert profile.execution_bindings.template == "difys-default-team/dify-agent-local-sandbox" + + +def test_build_runtime_backend_profile_rejects_missing_enterprise_endpoint( monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: monkeypatch.delenv("DIFY_AGENT_ENTERPRISE_SANDBOX_GATEWAY_ENDPOINT", raising=False) monkeypatch.chdir(tmp_path) - assert ServerSettings(shell_provider="enterprise").build_shell_provider() is None + with pytest.raises(ValidationError, match="enterprise_sandbox_gateway_endpoint is required"): + _ = ServerSettings(runtime_backend="enterprise").build_runtime_backend_profile() -def test_build_shell_provider_rejects_blank_shellctl_entrypoint() -> None: - with pytest.raises(ValidationError, match="shellctl_entrypoint is required"): - _ = ServerSettings(shell_provider="shellctl", shellctl_entrypoint=" ").build_shell_provider() +def test_build_runtime_backend_profile_rejects_blank_local_endpoint() -> None: + with pytest.raises(ValidationError, match="local_sandbox_endpoint is required"): + _ = ServerSettings(runtime_backend="local", local_sandbox_endpoint=" ").build_runtime_backend_profile() def test_server_settings_parses_shell_redact_patterns_json_array(monkeypatch: pytest.MonkeyPatch) -> None: @@ -342,4 +363,4 @@ def test_server_settings_rejects_non_array_shell_redact_patterns(monkeypatch: py settings = ServerSettings() with pytest.raises(ValueError, match="must be a JSON array"): - settings.get_shell_redact_patterns() + _ = settings.get_shell_redact_patterns() diff --git a/dify-agent/tests/local/dify_agent/server/test_sse.py b/dify-agent/tests/local/dify_agent/server/test_sse.py index 8f146ad396c..274725ce3cb 100644 --- a/dify-agent/tests/local/dify_agent/server/test_sse.py +++ b/dify-agent/tests/local/dify_agent/server/test_sse.py @@ -3,6 +3,8 @@ import json from collections.abc import AsyncGenerator from typing import cast +import pytest + from dify_agent.protocol.schemas import RunFailedEvent, RunFailedEventData, RunStartedEvent from dify_agent.server.sse import format_sse_event, sse_event_stream @@ -49,3 +51,16 @@ def test_sse_event_stream_emits_heartbeats_while_waiting() -> None: await stream.aclose() asyncio.run(scenario()) + + +def test_sse_event_stream_ends_after_finite_terminal_event_iterator() -> None: + async def scenario() -> None: + async def events(): + yield RunFailedEvent(id="2-0", run_id="run-1", data=RunFailedEventData(error="model failed")) + + stream = cast(AsyncGenerator[str, None], sse_event_stream(events(), heartbeat_interval_seconds=0.001)) + assert (await anext(stream)).startswith("id: 2-0\nevent: run_failed") + with pytest.raises(StopAsyncIteration): + _ = await asyncio.wait_for(anext(stream), timeout=0.1) + + asyncio.run(scenario()) diff --git a/dify-agent/tests/local/dify_agent/storage/test_redis_run_store.py b/dify-agent/tests/local/dify_agent/storage/test_redis_run_store.py index 48825bf32b5..e33e1ec8188 100644 --- a/dify-agent/tests/local/dify_agent/storage/test_redis_run_store.py +++ b/dify-agent/tests/local/dify_agent/storage/test_redis_run_store.py @@ -1,13 +1,26 @@ import asyncio from collections.abc import Mapping +import json from typing import cast +import pytest from pydantic import JsonValue from agenton.compositor import CompositorSessionSnapshot, LayerSessionSnapshot from agenton.layers import LifecycleState -from dify_agent.protocol.schemas import RunStartedEvent, RunSucceededEvent, RunSucceededEventData -from dify_agent.storage.redis_run_store import DEFAULT_RUN_RETENTION_SECONDS, RedisRunStore +from dify_agent.protocol.schemas import ( + RunCancelledEvent, + RunCancelledEventData, + RunFailedEvent, + RunFailedEventData, + RunFailureType, + RunStartedEvent, + RunStatus, + RunSucceededEvent, + RunSucceededEventData, +) +from dify_agent.runtime.event_sink import RunFinalizationResult +from dify_agent.storage.redis_run_store import DEFAULT_RUN_RETENTION_SECONDS, RedisRunStore, RunNotFoundError class FakeRedis: @@ -19,6 +32,7 @@ class FakeRedis: self.commands = [] self.values = {} self.streams = {} + self.stream_changed = asyncio.Event() async def set(self, key: str, value: object, *, ex: int | None = None) -> None: self.commands.append(("set", key, value, ex)) @@ -40,8 +54,46 @@ class FakeRedis: entries = self.streams.setdefault(key, []) event_id = f"{len(entries) + 1}-0" entries.append((event_id, dict(fields))) + self.stream_changed.set() return event_id + async def xrevrange( + self, + key: str, + max: str = "+", + min: str = "-", + *, + count: int | None = None, + ) -> list[tuple[str, dict[str, object]]]: + self.commands.append(("xrevrange", key, max, min, count)) + entries = list(reversed(self.streams.get(key, []))) + return entries[:count] if count is not None else entries + + async def xread( + self, + streams: Mapping[str, str], + *, + count: int | None = None, + block: int | None = None, + ) -> list[tuple[str, list[tuple[str, dict[str, object]]]]]: + self.commands.append(("xread", dict(streams), count, block)) + while True: + response: list[tuple[str, list[tuple[str, dict[str, object]]]]] = [] + for key, cursor in streams.items(): + entries = [ + entry + for entry in self.streams.get(key, []) + if self._stream_id_value(entry[0]) > self._stream_id_value(cursor) + ] + if count is not None: + entries = entries[:count] + if entries: + response.append((key, entries)) + if response: + return response + self.stream_changed.clear() + await self.stream_changed.wait() + async def xrange( self, key: str, *, min: str = "-", count: int | None = None ) -> list[tuple[str, dict[str, object]]]: @@ -55,6 +107,32 @@ class FakeRedis: self.commands.append(("expire", key, seconds)) return True + async def eval(self, script: str, numkeys: int, *keys_and_args: object) -> list[object]: + self.commands.append(("eval", script, numkeys, *keys_and_args)) + assert numkeys == 2 + record_key = str(keys_and_args[0]) + events_key = str(keys_and_args[1]) + status = str(keys_and_args[2]) + updated_at = str(keys_and_args[3]) + has_error = str(keys_and_args[4]) == "1" + error = str(keys_and_args[5]) if has_error else None + has_error_type = str(keys_and_args[6]) == "1" + error_type = str(keys_and_args[7]) if has_error_type else None + payload = str(keys_and_args[8]) + record_json = self.values.get(record_key) + if record_json is None: + return [-1, "", ""] + if isinstance(record_json, bytes): + record_json = record_json.decode() + record = json.loads(cast(str, record_json)) + if record["status"] != "running": + return [0, record["status"], ""] + + record.update({"status": status, "updated_at": updated_at, "error": error, "error_type": error_type}) + event_id = self._append_stream_entry(events_key, {"payload": payload}) + self.values[record_key] = json.dumps(record, separators=(",", ":")) + return [1, status, event_id] + @staticmethod def _is_after_min(event_id: str, min_id: str) -> bool: if min_id == "-": @@ -100,6 +178,25 @@ class FakeRedisPipeline: return list(self.results) +def _terminal_event( + event_type: str, + run_id: str, +) -> RunSucceededEvent | RunFailedEvent | RunCancelledEvent: + if event_type == "run_succeeded": + return RunSucceededEvent( + run_id=run_id, + data=RunSucceededEventData( + output="done", + session_snapshot=CompositorSessionSnapshot(layers=[]), + ), + ) + if event_type == "run_failed": + return RunFailedEvent(run_id=run_id, data=RunFailedEventData(error="model failed")) + if event_type == "run_cancelled": + return RunCancelledEvent(run_id=run_id, data=RunCancelledEventData(reason="cancelled")) + raise AssertionError(f"unexpected terminal event type: {event_type}") + + def test_create_run_writes_running_record_without_job_queue_and_with_retention() -> None: redis = FakeRedis() store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] @@ -113,17 +210,284 @@ def test_create_run_writes_running_record_without_job_queue_and_with_retention() assert "request" not in str(redis.commands[0][2]) -def test_update_status_refreshes_record_retention() -> None: +def test_get_run_accepts_legacy_record_without_error_type() -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + record = asyncio.run(store.create_run()) + record_key = f"test:runs:{record.run_id}:record" + payload = json.loads(cast(str, redis.values[record_key])) + del payload["error_type"] + redis.values[record_key] = json.dumps(payload) + + loaded = asyncio.run(store.get_run(record.run_id)) + + assert loaded.error_type is None + + +def test_finalize_run_atomically_writes_terminal_event_and_status() -> None: redis = FakeRedis() store = RedisRunStore(redis, prefix="test", run_retention_seconds=60) # pyright: ignore[reportArgumentType] record = asyncio.run(store.create_run()) redis.commands.clear() + event = RunCancelledEvent( + run_id=record.run_id, + data=RunCancelledEventData(reason="workflow_aborted", message="workflow stopped"), + ) - asyncio.run(store.update_status(record.run_id, "succeeded")) + result = asyncio.run(store.finalize_run(event)) + updated = asyncio.run(store.get_run(record.run_id)) - assert [command[0] for command in redis.commands] == ["get", "set"] - assert redis.commands[1][1] == f"test:runs:{record.run_id}:record" - assert redis.commands[1][3] == 60 + assert result.applied is True + assert result.status == "cancelled" + assert result.event_id == "1-0" + assert updated.status == "cancelled" + assert updated.error == "workflow stopped" + assert updated.error_type is None + assert updated.updated_at == event.created_at + stream_entry_id, stream_fields = redis.streams[f"test:runs:{record.run_id}:events"][0] + assert stream_entry_id == result.event_id + payload = json.loads(cast(str, stream_fields["payload"])) + assert "id" not in payload + assert payload["type"] == "run_cancelled" + assert payload["data"] == {"reason": "workflow_aborted", "message": "workflow stopped"} + assert payload["created_at"] == event.created_at.isoformat().replace("+00:00", "Z") + eval_command = redis.commands[0] + assert eval_command[0] == "eval" + assert eval_command[2] == 2 + assert eval_command[-1] == "60" + + +def test_finalize_run_rejects_a_second_terminal_without_appending_event() -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test", run_retention_seconds=60) # pyright: ignore[reportArgumentType] + record = asyncio.run(store.create_run()) + snapshot = CompositorSessionSnapshot(layers=[]) + + first = asyncio.run( + store.finalize_run( + RunSucceededEvent( + run_id=record.run_id, + data=RunSucceededEventData(output="done", session_snapshot=snapshot), + ) + ) + ) + second = asyncio.run( + store.finalize_run( + RunCancelledEvent( + run_id=record.run_id, + data=RunCancelledEventData(reason="late_cancel"), + ) + ) + ) + + assert first.applied is True + assert second.applied is False + assert second.status == "succeeded" + assert second.event_id is None + assert len(redis.streams[f"test:runs:{record.run_id}:events"]) == 1 + + +def test_finalize_failed_run_derives_error_and_timestamp_from_event() -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + record = asyncio.run(store.create_run()) + event = RunFailedEvent( + run_id=record.run_id, + data=RunFailedEventData( + error="model failed", + error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED, + reason="model_error", + ), + ) + + result = asyncio.run(store.finalize_run(event)) + updated = asyncio.run(store.get_run(record.run_id)) + + assert result.applied is True + assert result.status == "failed" + assert updated.status == "failed" + assert updated.error == "model failed" + assert updated.error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED + assert updated.updated_at == event.created_at + stream_entry = redis.streams[f"test:runs:{record.run_id}:events"][0] + payload = json.loads(cast(str, stream_entry[1]["payload"])) + assert payload["data"]["error_type"] == "agent_run_limit_exceeded" + + +def test_two_store_instances_choose_exactly_one_terminal_winner() -> None: + redis = FakeRedis() + first_store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + second_store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + async def scenario() -> tuple[list[RunFinalizationResult], RunStatus, list[str]]: + record = await first_store.create_run() + snapshot = CompositorSessionSnapshot(layers=[]) + results = await asyncio.gather( + first_store.finalize_run( + RunSucceededEvent( + run_id=record.run_id, + data=RunSucceededEventData(output="done", session_snapshot=snapshot), + ) + ), + second_store.finalize_run( + RunCancelledEvent( + run_id=record.run_id, + data=RunCancelledEventData(reason="concurrent_cancel"), + ) + ), + ) + persisted = await first_store.get_run(record.run_id) + page = await second_store.get_events(record.run_id) + return list(results), persisted.status, [event.type for event in page.events] + + results, status, event_types = asyncio.run(scenario()) + + assert sum(result.applied for result in results) == 1 + assert len(event_types) == 1 + assert (status, event_types[0]) in { + ("succeeded", "run_succeeded"), + ("cancelled", "run_cancelled"), + } + + +def test_failure_and_cancellation_compete_for_one_terminal() -> None: + redis = FakeRedis() + failure_store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + cancellation_store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + async def scenario() -> tuple[list[RunFinalizationResult], RunStatus, list[str]]: + record = await failure_store.create_run() + results = await asyncio.gather( + failure_store.finalize_run( + RunFailedEvent( + run_id=record.run_id, + data=RunFailedEventData(error="model failed", reason="model_error"), + ) + ), + cancellation_store.finalize_run( + RunCancelledEvent( + run_id=record.run_id, + data=RunCancelledEventData(reason="concurrent_cancel"), + ) + ), + ) + persisted = await failure_store.get_run(record.run_id) + page = await cancellation_store.get_events(record.run_id) + return list(results), persisted.status, [event.type for event in page.events] + + results, status, event_types = asyncio.run(scenario()) + + assert sum(result.applied for result in results) == 1 + assert len(event_types) == 1 + assert (status, event_types[0]) in { + ("failed", "run_failed"), + ("cancelled", "run_cancelled"), + } + + +def test_finalize_run_raises_when_record_is_missing() -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + with pytest.raises(RunNotFoundError): + asyncio.run( + store.finalize_run(RunCancelledEvent(run_id="missing", data=RunCancelledEventData(reason="cancelled"))) + ) + + +def test_wait_for_cancellation_observes_terminal_record_before_starting() -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + async def scenario() -> bool: + record = await store.create_run() + _ = await store.finalize_run( + RunCancelledEvent(run_id=record.run_id, data=RunCancelledEventData(reason="cancelled")) + ) + redis.commands.clear() + return await store.wait_for_cancellation(record.run_id) + + assert asyncio.run(scenario()) is True + assert [command[0] for command in redis.commands] == ["xrevrange", "get"] + + +def test_wait_for_cancellation_covers_terminal_transition_during_initialization() -> None: + class PausingRecordReadRedis(FakeRedis): + record_read_started: asyncio.Event + release_record_read: asyncio.Event + pause_next_record_read: bool + + def __init__(self) -> None: + super().__init__() + self.record_read_started = asyncio.Event() + self.release_record_read = asyncio.Event() + self.pause_next_record_read = True + + async def get(self, key: str) -> object | None: + if self.pause_next_record_read and key.endswith(":record"): + self.pause_next_record_read = False + self.record_read_started.set() + await self.release_record_read.wait() + return await super().get(key) + + redis = PausingRecordReadRedis() + observer_store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + cancelling_store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + async def scenario() -> bool: + record = await observer_store.create_run() + observer = asyncio.create_task(observer_store.wait_for_cancellation(record.run_id)) + await asyncio.wait_for(redis.record_read_started.wait(), timeout=1) + _ = await cancelling_store.finalize_run( + RunCancelledEvent(run_id=record.run_id, data=RunCancelledEventData(reason="cancelled")) + ) + redis.release_record_read.set() + return await asyncio.wait_for(observer, timeout=1) + + assert asyncio.run(scenario()) is True + + +def test_wait_for_cancellation_advances_past_non_terminal_events() -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + async def scenario() -> bool: + record = await store.create_run() + observer = asyncio.create_task(store.wait_for_cancellation(record.run_id)) + await asyncio.sleep(0) + _ = await store.append_event(RunStartedEvent(run_id=record.run_id)) + await asyncio.sleep(0) + _ = await store.finalize_run( + RunCancelledEvent(run_id=record.run_id, data=RunCancelledEventData(reason="cancelled")) + ) + return await asyncio.wait_for(observer, timeout=1) + + assert asyncio.run(scenario()) is True + cursors = [command[1] for command in redis.commands if command[0] == "xread"] + assert any("0-0" in streams.values() for streams in cursors if isinstance(streams, dict)) + assert any("1-0" in streams.values() for streams in cursors if isinstance(streams, dict)) + + +def test_wait_for_cancellation_returns_false_when_success_wins() -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + async def scenario() -> bool: + record = await store.create_run() + observer = asyncio.create_task(store.wait_for_cancellation(record.run_id)) + await asyncio.sleep(0) + _ = await store.finalize_run( + RunSucceededEvent( + run_id=record.run_id, + data=RunSucceededEventData( + output="done", + session_snapshot=CompositorSessionSnapshot(layers=[]), + ), + ) + ) + return await asyncio.wait_for(observer, timeout=1) + + assert asyncio.run(scenario()) is False def test_append_event_serializes_typed_event_without_id_and_expires_run_keys() -> None: @@ -166,21 +530,66 @@ def test_get_events_round_trips_run_succeeded_output_and_session_snapshot() -> N async def scenario() -> tuple[str, RunSucceededEvent]: record = await store.create_run() - event_id = await store.append_event( + result = await store.finalize_run( RunSucceededEvent( id="local-only", run_id=record.run_id, data=RunSucceededEventData(output=output, session_snapshot=session_snapshot), ) ) + assert result.event_id is not None page = await store.get_events(record.run_id, after="0-0", limit=10) decoded = page.events[0] assert isinstance(decoded, RunSucceededEvent) - assert page.next_cursor == event_id - return event_id, decoded + assert page.next_cursor == result.event_id + return result.event_id, decoded event_id, decoded = asyncio.run(scenario()) assert decoded.id == event_id assert decoded.data.output == output assert decoded.data.session_snapshot == session_snapshot + + +@pytest.mark.parametrize("terminal_type", ["run_succeeded", "run_failed", "run_cancelled"]) +def test_iter_events_ends_after_replaying_terminal_event(terminal_type: str) -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + async def scenario() -> list[str]: + record = await store.create_run() + _ = await store.append_event(RunStartedEvent(run_id=record.run_id)) + _ = await store.finalize_run(_terminal_event(terminal_type, record.run_id)) + redis.commands.clear() + + async def collect_events() -> list[str]: + return [event.type async for event in store.iter_events(record.run_id)] + + return await asyncio.wait_for(collect_events(), timeout=1) + + event_types = asyncio.run(scenario()) + + assert event_types == ["run_started", terminal_type] + assert "xread" not in [command[0] for command in redis.commands] + + +@pytest.mark.parametrize("terminal_type", ["run_succeeded", "run_failed", "run_cancelled"]) +def test_iter_events_ends_after_live_terminal_event(terminal_type: str) -> None: + redis = FakeRedis() + store = RedisRunStore(redis, prefix="test") # pyright: ignore[reportArgumentType] + + async def scenario() -> str: + record = await store.create_run() + events = store.iter_events(record.run_id) + next_event = asyncio.ensure_future(anext(events)) + await asyncio.sleep(0) + assert not next_event.done() + assert "xread" in [command[0] for command in redis.commands] + + _ = await store.finalize_run(_terminal_event(terminal_type, record.run_id)) + event = await asyncio.wait_for(next_event, timeout=1) + with pytest.raises(StopAsyncIteration): + _ = await anext(events) + return event.type + + assert asyncio.run(scenario()) == terminal_type diff --git a/dify-agent/tests/local/dify_agent/test_client_safe_exports.py b/dify-agent/tests/local/dify_agent/test_client_safe_exports.py index 175eb61aa5e..3ece328472a 100644 --- a/dify-agent/tests/local/dify_agent/test_client_safe_exports.py +++ b/dify-agent/tests/local/dify_agent/test_client_safe_exports.py @@ -54,11 +54,7 @@ def test_client_public_exports_work_with_default_dependencies_only(tmp_path: Pat requirement_name(requirement) for requirement in pyproject["project"].get("optional-dependencies", {}).get("server", []) } - grpc_dependency_names = { - requirement_name(requirement) - for requirement in pyproject["project"].get("optional-dependencies", {}).get("grpc", []) - } - server_only_dependency_names = (server_dependency_names | grpc_dependency_names) - default_dependency_names + server_only_dependency_names = server_dependency_names - default_dependency_names agenton_layers = importlib.import_module("agenton.layers") agenton_compositor = importlib.import_module("agenton.compositor") diff --git a/dify-agent/tests/local/dify_agent/test_import_boundaries.py b/dify-agent/tests/local/dify_agent/test_import_boundaries.py index 961dab1b8d2..4407d408f41 100644 --- a/dify-agent/tests/local/dify_agent/test_import_boundaries.py +++ b/dify-agent/tests/local/dify_agent/test_import_boundaries.py @@ -152,8 +152,6 @@ def test_agent_cli_help_import_is_client_safe() -> None: "dify_agent.server", "dify_agent.agent_stub.server", "fastapi", - "google.protobuf", - "grpclib", "jwcrypto", "pydantic_settings", "redis", @@ -229,8 +227,6 @@ def test_agent_cli_help_render_does_not_load_server_or_cli_modules() -> None: "dify_agent.server", "dify_agent.agent_stub.server", "fastapi", - "google.protobuf", - "grpclib", "jwcrypto", "pydantic_settings", "redis", diff --git a/dify-agent/tests/local/shellctl/test_shellctl_shared.py b/dify-agent/tests/local/shellctl/test_shellctl_shared.py index 8034ea3d9ef..38fd14987c4 100644 --- a/dify-agent/tests/local/shellctl/test_shellctl_shared.py +++ b/dify-agent/tests/local/shellctl/test_shellctl_shared.py @@ -7,13 +7,27 @@ from pydantic import ValidationError from shellctl.shared import ( JOB_ID_ALPHABET, + MAX_WAIT_TIMEOUT_SECONDS, RunJobRequest, + SHELL_TOOL_HARD_TIMEOUT_SECONDS, + SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS, + SHELL_TOOL_TIMEOUT_WITH_HTTP_GRACE_SECONDS, generate_job_id, read_output_window, tail_output_window, ) +def test_shell_tool_timeout_budget_has_one_source_of_truth() -> None: + assert MAX_WAIT_TIMEOUT_SECONDS == SHELL_TOOL_HARD_TIMEOUT_SECONDS == 300 + assert SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS == 10 + assert ( + SHELL_TOOL_TIMEOUT_WITH_HTTP_GRACE_SECONDS + == SHELL_TOOL_HARD_TIMEOUT_SECONDS + SHELL_TOOL_HTTP_TIMEOUT_GRACE_SECONDS + == 310 + ) + + def test_generate_job_id_matches_proposal_format() -> None: job_id = generate_job_id(now=datetime(2026, 5, 21, 15, 30, tzinfo=UTC)) diff --git a/dify-agent/tests/local/test_packaging.py b/dify-agent/tests/local/test_packaging.py index b11bcb650a8..13cdc058b49 100644 --- a/dify-agent/tests/local/test_packaging.py +++ b/dify-agent/tests/local/test_packaging.py @@ -15,6 +15,7 @@ CLIENT_SHARED_DTO_DEPENDENCIES = { } SERVER_RUNTIME_DEPENDENCIES = { + "e2b>=2.34.0,<3.0.0", "fastapi==0.136.0", "graphon==0.5.2", "jsonschema>=4.23.0,<5.0.0", @@ -26,15 +27,9 @@ SERVER_RUNTIME_DEPENDENCIES = { "uvicorn[standard]==0.46.0", } -GRPC_RUNTIME_DEPENDENCIES = { - "grpclib[protobuf]>=0.4.9,<0.5.0", - "protobuf>=6.33.5,<7.0.0", -} - DEV_DEPENDENCIES = { "basedpyright>=1.39.3", "coverage[toml]>=7.10.7", - "grpcio-tools>=1.81.0,<2.0.0", "pytest>=9.0.3", "pytest-examples>=0.0.18", "pytest-mock>=3.14.0", @@ -51,7 +46,6 @@ def test_project_dependencies_split_client_and_server_requirements() -> None: project = pyproject["project"] assert set(project["dependencies"]) == CLIENT_SHARED_DTO_DEPENDENCIES - assert set(project["optional-dependencies"]["grpc"]) == GRPC_RUNTIME_DEPENDENCIES assert set(project["optional-dependencies"]["server"]) == SERVER_RUNTIME_DEPENDENCIES assert set(pyproject["dependency-groups"]["dev"]) == DEV_DEPENDENCIES diff --git a/dify-agent/uv.lock b/dify-agent/uv.lock index c355ba99dca..4f6e9523f44 100644 --- a/dify-agent/uv.lock +++ b/dify-agent/uv.lock @@ -198,6 +198,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/0a/de/acae8e9f9a1f4bb393d41c8265898b0f29772e38eac14e9f69d191e2c006/blis-1.3.3-cp314-cp314-win_amd64.whl", hash = "sha256:9e5fdf4211b1972400f8ff6dafe87cb689c5d84f046b4a76b207c0bd2270faaf", size = 6324695, upload-time = "2025-11-17T12:28:28.401Z" }, ] +[[package]] +name = "bracex" +version = "3.0.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/ac/01/5f394b8bcd6e5b92f73130990960423bbb19711f906bd9fe9ea5557c667c/bracex-3.0.1.tar.gz", hash = "sha256:4e38e32392e4a4780fe15d644bfc7c8514057cfc3861e060b11814ce829c25e4", size = 44019, upload-time = "2026-07-20T13:43:00.335Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b8/8f/6f7273a7adb8d73fc8d21ede4376a3e475e52f98435c6007f69100dec8ca/bracex-3.0.1-py3-none-any.whl", hash = "sha256:6523ad83aeb5098a4ee597cff0f964442ff74e460bd3fafaffab6a013ff2288c", size = 11940, upload-time = "2026-07-20T13:42:59.268Z" }, +] + [[package]] name = "catalogue" version = "2.0.10" @@ -471,55 +480,52 @@ wheels = [ [[package]] name = "cryptography" -version = "46.0.7" +version = "50.0.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "cffi", marker = "platform_python_implementation != 'PyPy'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/47/93/ac8f3d5ff04d54bc814e961a43ae5b0b146154c89c61b47bb07557679b18/cryptography-46.0.7.tar.gz", hash = "sha256:e4cfd68c5f3e0bfdad0d38e023239b96a2fe84146481852dffbcca442c245aa5", size = 750652, upload-time = "2026-04-08T01:57:54.692Z" } +sdist = { url = "https://files.pythonhosted.org/packages/de/41/6cbdcf9142d00fe82836fbb51e503e58088575cf7a0fe1dbff6695bf0840/cryptography-50.0.0.tar.gz", hash = "sha256:eeac2acb5a20ed25e0ad6d1df9891a520b78b404266b6d11778f25d5d691a6c9", size = 880201, upload-time = "2026-07-31T14:25:10.11Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/0b/5d/4a8f770695d73be252331e60e526291e3df0c9b27556a90a6b47bccca4c2/cryptography-46.0.7-cp311-abi3-macosx_10_9_universal2.whl", hash = "sha256:ea42cbe97209df307fdc3b155f1b6fa2577c0defa8f1f7d3be7d31d189108ad4", size = 7179869, upload-time = "2026-04-08T01:56:17.157Z" }, - { url = "https://files.pythonhosted.org/packages/5f/45/6d80dc379b0bbc1f9d1e429f42e4cb9e1d319c7a8201beffd967c516ea01/cryptography-46.0.7-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:b36a4695e29fe69215d75960b22577197aca3f7a25b9cf9d165dcfe9d80bc325", size = 4275492, upload-time = "2026-04-08T01:56:19.36Z" }, - { url = "https://files.pythonhosted.org/packages/4a/9a/1765afe9f572e239c3469f2cb429f3ba7b31878c893b246b4b2994ffe2fe/cryptography-46.0.7-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:5ad9ef796328c5e3c4ceed237a183f5d41d21150f972455a9d926593a1dcb308", size = 4426670, upload-time = "2026-04-08T01:56:21.415Z" }, - { url = "https://files.pythonhosted.org/packages/8f/3e/af9246aaf23cd4ee060699adab1e47ced3f5f7e7a8ffdd339f817b446462/cryptography-46.0.7-cp311-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:73510b83623e080a2c35c62c15298096e2a5dc8d51c3b4e1740211839d0dea77", size = 4280275, upload-time = "2026-04-08T01:56:23.539Z" }, - { url = "https://files.pythonhosted.org/packages/0f/54/6bbbfc5efe86f9d71041827b793c24811a017c6ac0fd12883e4caa86b8ed/cryptography-46.0.7-cp311-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:cbd5fb06b62bd0721e1170273d3f4d5a277044c47ca27ee257025146c34cbdd1", size = 4928402, upload-time = "2026-04-08T01:56:25.624Z" }, - { url = "https://files.pythonhosted.org/packages/2d/cf/054b9d8220f81509939599c8bdbc0c408dbd2bdd41688616a20731371fe0/cryptography-46.0.7-cp311-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:420b1e4109cc95f0e5700eed79908cef9268265c773d3a66f7af1eef53d409ef", size = 4459985, upload-time = "2026-04-08T01:56:27.309Z" }, - { url = "https://files.pythonhosted.org/packages/f9/46/4e4e9c6040fb01c7467d47217d2f882daddeb8828f7df800cb806d8a2288/cryptography-46.0.7-cp311-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:24402210aa54baae71d99441d15bb5a1919c195398a87b563df84468160a65de", size = 3990652, upload-time = "2026-04-08T01:56:29.095Z" }, - { url = "https://files.pythonhosted.org/packages/36/5f/313586c3be5a2fbe87e4c9a254207b860155a8e1f3cca99f9910008e7d08/cryptography-46.0.7-cp311-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:8a469028a86f12eb7d2fe97162d0634026d92a21f3ae0ac87ed1c4a447886c83", size = 4279805, upload-time = "2026-04-08T01:56:30.928Z" }, - { url = "https://files.pythonhosted.org/packages/69/33/60dfc4595f334a2082749673386a4d05e4f0cf4df8248e63b2c3437585f2/cryptography-46.0.7-cp311-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:9694078c5d44c157ef3162e3bf3946510b857df5a3955458381d1c7cfc143ddb", size = 4892883, upload-time = "2026-04-08T01:56:32.614Z" }, - { url = "https://files.pythonhosted.org/packages/c7/0b/333ddab4270c4f5b972f980adef4faa66951a4aaf646ca067af597f15563/cryptography-46.0.7-cp311-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:42a1e5f98abb6391717978baf9f90dc28a743b7d9be7f0751a6f56a75d14065b", size = 4459756, upload-time = "2026-04-08T01:56:34.306Z" }, - { url = "https://files.pythonhosted.org/packages/d2/14/633913398b43b75f1234834170947957c6b623d1701ffc7a9600da907e89/cryptography-46.0.7-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:91bbcb08347344f810cbe49065914fe048949648f6bd5c2519f34619142bbe85", size = 4410244, upload-time = "2026-04-08T01:56:35.977Z" }, - { url = "https://files.pythonhosted.org/packages/10/f2/19ceb3b3dc14009373432af0c13f46aa08e3ce334ec6eff13492e1812ccd/cryptography-46.0.7-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:5d1c02a14ceb9148cc7816249f64f623fbfee39e8c03b3650d842ad3f34d637e", size = 4674868, upload-time = "2026-04-08T01:56:38.034Z" }, - { url = "https://files.pythonhosted.org/packages/1a/bb/a5c213c19ee94b15dfccc48f363738633a493812687f5567addbcbba9f6f/cryptography-46.0.7-cp311-abi3-win32.whl", hash = "sha256:d23c8ca48e44ee015cd0a54aeccdf9f09004eba9fc96f38c911011d9ff1bd457", size = 3026504, upload-time = "2026-04-08T01:56:39.666Z" }, - { url = "https://files.pythonhosted.org/packages/2b/02/7788f9fefa1d060ca68717c3901ae7fffa21ee087a90b7f23c7a603c32ae/cryptography-46.0.7-cp311-abi3-win_amd64.whl", hash = "sha256:397655da831414d165029da9bc483bed2fe0e75dde6a1523ec2fe63f3c46046b", size = 3488363, upload-time = "2026-04-08T01:56:41.893Z" }, - { url = "https://files.pythonhosted.org/packages/7b/56/15619b210e689c5403bb0540e4cb7dbf11a6bf42e483b7644e471a2812b3/cryptography-46.0.7-cp314-cp314t-macosx_10_9_universal2.whl", hash = "sha256:d151173275e1728cf7839aaa80c34fe550c04ddb27b34f48c232193df8db5842", size = 7119671, upload-time = "2026-04-08T01:56:44Z" }, - { url = "https://files.pythonhosted.org/packages/74/66/e3ce040721b0b5599e175ba91ab08884c75928fbeb74597dd10ef13505d2/cryptography-46.0.7-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:db0f493b9181c7820c8134437eb8b0b4792085d37dbb24da050476ccb664e59c", size = 4268551, upload-time = "2026-04-08T01:56:46.071Z" }, - { url = "https://files.pythonhosted.org/packages/03/11/5e395f961d6868269835dee1bafec6a1ac176505a167f68b7d8818431068/cryptography-46.0.7-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:ebd6daf519b9f189f85c479427bbd6e9c9037862cf8fe89ee35503bd209ed902", size = 4408887, upload-time = "2026-04-08T01:56:47.718Z" }, - { url = "https://files.pythonhosted.org/packages/40/53/8ed1cf4c3b9c8e611e7122fb56f1c32d09e1fff0f1d77e78d9ff7c82653e/cryptography-46.0.7-cp314-cp314t-manylinux_2_28_aarch64.whl", hash = "sha256:b7b412817be92117ec5ed95f880defe9cf18a832e8cafacf0a22337dc1981b4d", size = 4271354, upload-time = "2026-04-08T01:56:49.312Z" }, - { url = "https://files.pythonhosted.org/packages/50/46/cf71e26025c2e767c5609162c866a78e8a2915bbcfa408b7ca495c6140c4/cryptography-46.0.7-cp314-cp314t-manylinux_2_28_ppc64le.whl", hash = "sha256:fbfd0e5f273877695cb93baf14b185f4878128b250cc9f8e617ea0c025dfb022", size = 4905845, upload-time = "2026-04-08T01:56:50.916Z" }, - { url = "https://files.pythonhosted.org/packages/c0/ea/01276740375bac6249d0a971ebdf6b4dc9ead0ee0a34ef3b5a88c1a9b0d4/cryptography-46.0.7-cp314-cp314t-manylinux_2_28_x86_64.whl", hash = "sha256:ffca7aa1d00cf7d6469b988c581598f2259e46215e0140af408966a24cf086ce", size = 4444641, upload-time = "2026-04-08T01:56:52.882Z" }, - { url = "https://files.pythonhosted.org/packages/3d/4c/7d258f169ae71230f25d9f3d06caabcff8c3baf0978e2b7d65e0acac3827/cryptography-46.0.7-cp314-cp314t-manylinux_2_31_armv7l.whl", hash = "sha256:60627cf07e0d9274338521205899337c5d18249db56865f943cbe753aa96f40f", size = 3967749, upload-time = "2026-04-08T01:56:54.597Z" }, - { url = "https://files.pythonhosted.org/packages/b5/2a/2ea0767cad19e71b3530e4cad9605d0b5e338b6a1e72c37c9c1ceb86c333/cryptography-46.0.7-cp314-cp314t-manylinux_2_34_aarch64.whl", hash = "sha256:80406c3065e2c55d7f49a9550fe0c49b3f12e5bfff5dedb727e319e1afb9bf99", size = 4270942, upload-time = "2026-04-08T01:56:56.416Z" }, - { url = "https://files.pythonhosted.org/packages/41/3d/fe14df95a83319af25717677e956567a105bb6ab25641acaa093db79975d/cryptography-46.0.7-cp314-cp314t-manylinux_2_34_ppc64le.whl", hash = "sha256:c5b1ccd1239f48b7151a65bc6dd54bcfcc15e028c8ac126d3fada09db0e07ef1", size = 4871079, upload-time = "2026-04-08T01:56:58.31Z" }, - { url = "https://files.pythonhosted.org/packages/9c/59/4a479e0f36f8f378d397f4eab4c850b4ffb79a2f0d58704b8fa0703ddc11/cryptography-46.0.7-cp314-cp314t-manylinux_2_34_x86_64.whl", hash = "sha256:d5f7520159cd9c2154eb61eb67548ca05c5774d39e9c2c4339fd793fe7d097b2", size = 4443999, upload-time = "2026-04-08T01:57:00.508Z" }, - { url = "https://files.pythonhosted.org/packages/28/17/b59a741645822ec6d04732b43c5d35e4ef58be7bfa84a81e5ae6f05a1d33/cryptography-46.0.7-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:fcd8eac50d9138c1d7fc53a653ba60a2bee81a505f9f8850b6b2888555a45d0e", size = 4399191, upload-time = "2026-04-08T01:57:02.654Z" }, - { url = "https://files.pythonhosted.org/packages/59/6a/bb2e166d6d0e0955f1e9ff70f10ec4b2824c9cfcdb4da772c7dd69cc7d80/cryptography-46.0.7-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:65814c60f8cc400c63131584e3e1fad01235edba2614b61fbfbfa954082db0ee", size = 4655782, upload-time = "2026-04-08T01:57:04.592Z" }, - { url = "https://files.pythonhosted.org/packages/95/b6/3da51d48415bcb63b00dc17c2eff3a651b7c4fed484308d0f19b30e8cb2c/cryptography-46.0.7-cp314-cp314t-win32.whl", hash = "sha256:fdd1736fed309b4300346f88f74cd120c27c56852c3838cab416e7a166f67298", size = 3002227, upload-time = "2026-04-08T01:57:06.91Z" }, - { url = "https://files.pythonhosted.org/packages/32/a8/9f0e4ed57ec9cebe506e58db11ae472972ecb0c659e4d52bbaee80ca340a/cryptography-46.0.7-cp314-cp314t-win_amd64.whl", hash = "sha256:e06acf3c99be55aa3b516397fe42f5855597f430add9c17fa46bf2e0fb34c9bb", size = 3475332, upload-time = "2026-04-08T01:57:08.807Z" }, - { url = "https://files.pythonhosted.org/packages/a7/7f/cd42fc3614386bc0c12f0cb3c4ae1fc2bbca5c9662dfed031514911d513d/cryptography-46.0.7-cp38-abi3-macosx_10_9_universal2.whl", hash = "sha256:462ad5cb1c148a22b2e3bcc5ad52504dff325d17daf5df8d88c17dda1f75f2a4", size = 7165618, upload-time = "2026-04-08T01:57:10.645Z" }, - { url = "https://files.pythonhosted.org/packages/a5/d0/36a49f0262d2319139d2829f773f1b97ef8aef7f97e6e5bd21455e5a8fb5/cryptography-46.0.7-cp38-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:84d4cced91f0f159a7ddacad249cc077e63195c36aac40b4150e7a57e84fffe7", size = 4270628, upload-time = "2026-04-08T01:57:12.885Z" }, - { url = "https://files.pythonhosted.org/packages/8a/6c/1a42450f464dda6ffbe578a911f773e54dd48c10f9895a23a7e88b3e7db5/cryptography-46.0.7-cp38-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:128c5edfe5e5938b86b03941e94fac9ee793a94452ad1365c9fc3f4f62216832", size = 4415405, upload-time = "2026-04-08T01:57:14.923Z" }, - { url = "https://files.pythonhosted.org/packages/9a/92/4ed714dbe93a066dc1f4b4581a464d2d7dbec9046f7c8b7016f5286329e2/cryptography-46.0.7-cp38-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:5e51be372b26ef4ba3de3c167cd3d1022934bc838ae9eaad7e644986d2a3d163", size = 4272715, upload-time = "2026-04-08T01:57:16.638Z" }, - { url = "https://files.pythonhosted.org/packages/b7/e6/a26b84096eddd51494bba19111f8fffe976f6a09f132706f8f1bf03f51f7/cryptography-46.0.7-cp38-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:cdf1a610ef82abb396451862739e3fc93b071c844399e15b90726ef7470eeaf2", size = 4918400, upload-time = "2026-04-08T01:57:19.021Z" }, - { url = "https://files.pythonhosted.org/packages/c7/08/ffd537b605568a148543ac3c2b239708ae0bd635064bab41359252ef88ed/cryptography-46.0.7-cp38-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:1d25aee46d0c6f1a501adcddb2d2fee4b979381346a78558ed13e50aa8a59067", size = 4450634, upload-time = "2026-04-08T01:57:21.185Z" }, - { url = "https://files.pythonhosted.org/packages/16/01/0cd51dd86ab5b9befe0d031e276510491976c3a80e9f6e31810cce46c4ad/cryptography-46.0.7-cp38-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:cdfbe22376065ffcf8be74dc9a909f032df19bc58a699456a21712d6e5eabfd0", size = 3985233, upload-time = "2026-04-08T01:57:22.862Z" }, - { url = "https://files.pythonhosted.org/packages/92/49/819d6ed3a7d9349c2939f81b500a738cb733ab62fbecdbc1e38e83d45e12/cryptography-46.0.7-cp38-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:abad9dac36cbf55de6eb49badd4016806b3165d396f64925bf2999bcb67837ba", size = 4271955, upload-time = "2026-04-08T01:57:24.814Z" }, - { url = "https://files.pythonhosted.org/packages/80/07/ad9b3c56ebb95ed2473d46df0847357e01583f4c52a85754d1a55e29e4d0/cryptography-46.0.7-cp38-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:935ce7e3cfdb53e3536119a542b839bb94ec1ad081013e9ab9b7cfd478b05006", size = 4879888, upload-time = "2026-04-08T01:57:26.88Z" }, - { url = "https://files.pythonhosted.org/packages/b8/c7/201d3d58f30c4c2bdbe9b03844c291feb77c20511cc3586daf7edc12a47b/cryptography-46.0.7-cp38-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:35719dc79d4730d30f1c2b6474bd6acda36ae2dfae1e3c16f2051f215df33ce0", size = 4449961, upload-time = "2026-04-08T01:57:29.068Z" }, - { url = "https://files.pythonhosted.org/packages/a5/ef/649750cbf96f3033c3c976e112265c33906f8e462291a33d77f90356548c/cryptography-46.0.7-cp38-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:7bbc6ccf49d05ac8f7d7b5e2e2c33830d4fe2061def88210a126d130d7f71a85", size = 4401696, upload-time = "2026-04-08T01:57:31.029Z" }, - { url = "https://files.pythonhosted.org/packages/41/52/a8908dcb1a389a459a29008c29966c1d552588d4ae6d43f3a1a4512e0ebe/cryptography-46.0.7-cp38-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:a1529d614f44b863a7b480c6d000fe93b59acee9c82ffa027cfadc77521a9f5e", size = 4664256, upload-time = "2026-04-08T01:57:33.144Z" }, - { url = "https://files.pythonhosted.org/packages/4b/fa/f0ab06238e899cc3fb332623f337a7364f36f4bb3f2534c2bb95a35b132c/cryptography-46.0.7-cp38-abi3-win32.whl", hash = "sha256:f247c8c1a1fb45e12586afbb436ef21ff1e80670b2861a90353d9b025583d246", size = 3013001, upload-time = "2026-04-08T01:57:34.933Z" }, - { url = "https://files.pythonhosted.org/packages/d2/f1/00ce3bde3ca542d1acd8f8cfa38e446840945aa6363f9b74746394b14127/cryptography-46.0.7-cp38-abi3-win_amd64.whl", hash = "sha256:506c4ff91eff4f82bdac7633318a526b1d1309fc07ca76a3ad182cb5b686d6d3", size = 3472985, upload-time = "2026-04-08T01:57:36.714Z" }, + { url = "https://files.pythonhosted.org/packages/c5/5c/59086b4aac5e879d38ddbcf74e4be7ade89cebc3eb199a55da998c3bb46a/cryptography-50.0.0-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:031e2d5dd4bb9caa3ca9c82e5a197fd8ae680232cee62603d1a813f3f07e3d03", size = 4001252, upload-time = "2026-07-31T14:23:33.331Z" }, + { url = "https://files.pythonhosted.org/packages/57/ef/8f2df13c7216bcad3e1c74e07f6e193d93e998e114f524a53877c9af27ad/cryptography-50.0.0-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:fd9192b7b70c573d7f214eb1ae35e00d359f6f5e4b27c7e21e30de1fc6204645", size = 4719554, upload-time = "2026-07-31T14:23:35.611Z" }, + { url = "https://files.pythonhosted.org/packages/d9/41/029086c34d91052fc3b88bcc8056f709a7c915c7a23b235a54eb800b1c97/cryptography-50.0.0-cp311-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:06a32a980526a6ab9a4b9bf8f7385800791e2bb960903cb6b530e4817509a3b7", size = 4702130, upload-time = "2026-07-31T14:23:37.635Z" }, + { url = "https://files.pythonhosted.org/packages/7d/ff/b6ce0954962e7f7b969f850a883744197bb3910bdfd7b6da162eab7d9f68/cryptography-50.0.0-cp311-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:a1b30560f2acc95aa8b2e06e716a13dbfc97314747b80d9707e307f77b40d6b3", size = 4725244, upload-time = "2026-07-31T14:23:39.471Z" }, + { url = "https://files.pythonhosted.org/packages/06/1e/63a1027cb7fec360a182208e1b7767d5aa1fe57be3d6aa856e69a321edc0/cryptography-50.0.0-cp311-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:8d89f3976b10b4ce31118de72329025f70d2c6ead14a8217c5514dd2c6d5a78f", size = 5342265, upload-time = "2026-07-31T14:23:41.286Z" }, + { url = "https://files.pythonhosted.org/packages/6b/72/a1116d683a6d7ece94590013882515de087edf9ef0e6292aae615a44df73/cryptography-50.0.0-cp311-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:b42a28c1844fd9de8f3f7d540e36b66f3a9c83fceac7170ebc7a6a19edd9dcae", size = 4734609, upload-time = "2026-07-31T14:23:43.139Z" }, + { url = "https://files.pythonhosted.org/packages/15/37/36a9c479bbe49acea2636c7fd3360d20f7b7e079c300352011c44850b181/cryptography-50.0.0-cp311-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:900131fafd8aead39ac7dd3a7e833be754c17a95cfd91221636949fe4eb0aa8a", size = 4356517, upload-time = "2026-07-31T14:23:44.939Z" }, + { url = "https://files.pythonhosted.org/packages/32/98/8a151d64367204cbc63ec65d37502f1d9c53cf4bfc6ec3c532614dbec60d/cryptography-50.0.0-cp311-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:07949c449a1abcf60d1ee6e88956d89404c7df3c8258f46589e912988e551987", size = 4724529, upload-time = "2026-07-31T14:23:46.93Z" }, + { url = "https://files.pythonhosted.org/packages/22/f6/ec13b470172126464a86bf54d2294a46d29837fc51ba3e45d4047946fb5e/cryptography-50.0.0-cp311-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:f89831ef99dd7dd169ab06d63a831adb9e20a87aac6d380266bbda5823349169", size = 5299852, upload-time = "2026-07-31T14:23:48.851Z" }, + { url = "https://files.pythonhosted.org/packages/da/3a/f05e32c99d440c9bb891ea0e36c9091891e36be5a9a87ab2ee6ea20729f6/cryptography-50.0.0-cp311-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:82148ec5bddac30b51a5b3c1945075f896fa022cb93f8e4a01e9f6ee95292c5f", size = 4734462, upload-time = "2026-07-31T14:23:50.861Z" }, + { url = "https://files.pythonhosted.org/packages/ca/dc/bd72b26be8953f80625f63151efd38eee71c76ca6cf591c08ff34615a79e/cryptography-50.0.0-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:1489e263a8048bb8b6a8bac662eb2d402ea5d2b7b4699b72f385f1e2772db105", size = 4852708, upload-time = "2026-07-31T14:23:52.715Z" }, + { url = "https://files.pythonhosted.org/packages/27/20/c930314a2ab476d15dec966ec87e2e9637bb02b06106b12c0396c57bb603/cryptography-50.0.0-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:7cec5b856506da6defb290f30c9ee687d5f5e8cb0bd3f6459dde43b0b4fa40ef", size = 5004179, upload-time = "2026-07-31T14:23:54.887Z" }, + { url = "https://files.pythonhosted.org/packages/32/2e/c9db68a0c4bfa28e310707527c0ee3a2bd254104d2e02e68f368e197aa4c/cryptography-50.0.0-cp311-abi3-win_amd64.whl", hash = "sha256:bd1c592e4d5974f0d08d4888e432157adba757c66da0246918e43677fafa2d30", size = 3840395, upload-time = "2026-07-31T14:23:56.677Z" }, + { url = "https://files.pythonhosted.org/packages/c3/fb/951032a3bf22a5697c83183fb6294a4843772947a70e616c57b3ff5f522e/cryptography-50.0.0-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:49e7d93abdbd2990caced757e5fade25302f719c3c8fb6e6fff2dde98999fc41", size = 3989258, upload-time = "2026-07-31T14:23:58.881Z" }, + { url = "https://files.pythonhosted.org/packages/d4/67/91eb047e69c5e845f2f14b8a2e4a1aab0f283cb885531e9e22c8adb176bc/cryptography-50.0.0-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:19736989797678c6af1e55cd49055cdbcb55d8f6b5583ac5335f933aba9101dc", size = 4700648, upload-time = "2026-07-31T14:24:00.702Z" }, + { url = "https://files.pythonhosted.org/packages/30/82/85f0f7425c856b9f96459411eb12e74ef72df9caf6f8f15bf23a33ff131f/cryptography-50.0.0-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:80b63928fa35083b33966ce1efb70e5b9607181e49dcd1c22c8c005e319f667f", size = 4682442, upload-time = "2026-07-31T14:24:02.538Z" }, + { url = "https://files.pythonhosted.org/packages/1a/28/b555a365adff1cca2fbe7b9e487d68a40de6bc67ff2cb587473eb43de0e7/cryptography-50.0.0-cp314-cp314t-manylinux_2_28_aarch64.whl", hash = "sha256:d58c3db7cd6eed54e6c06744db55456b65ebd7492ddeae9c1e93cfca7aa857d3", size = 4707596, upload-time = "2026-07-31T14:24:04.394Z" }, + { url = "https://files.pythonhosted.org/packages/72/d8/f52538140cc719df62a01cf87d1c7142318d235817109d6f4054d7c352d6/cryptography-50.0.0-cp314-cp314t-manylinux_2_28_ppc64le.whl", hash = "sha256:df2a58a472f332225671c35b0a830208b86d004f82baa8530fa3782c85646533", size = 5314552, upload-time = "2026-07-31T14:24:06.31Z" }, + { url = "https://files.pythonhosted.org/packages/38/14/6120e5bd7c5aa022ad15424ba4d5c5269d0d9448ed4d55e492ea91e3c1c4/cryptography-50.0.0-cp314-cp314t-manylinux_2_28_x86_64.whl", hash = "sha256:11b74db56cdbe3cdee6e3f6982ecb70334fa10dce99ed58bf7894aaaa3b2a037", size = 4717113, upload-time = "2026-07-31T14:24:08.349Z" }, + { url = "https://files.pythonhosted.org/packages/fa/71/190bf38c3ee2e0f8efc9860ae100c9df4169742eef274b91e7aa1cb133b9/cryptography-50.0.0-cp314-cp314t-manylinux_2_31_armv7l.whl", hash = "sha256:f59e38625469987d7ef6d495323c55e7db6c212eaf6112267e0d3b565a2e9c9f", size = 4338580, upload-time = "2026-07-31T14:24:10.227Z" }, + { url = "https://files.pythonhosted.org/packages/3a/63/504ccfbbe61fd8aa983f7f146399cdf034c72c2fc55f5b2dfdcdcdb20c99/cryptography-50.0.0-cp314-cp314t-manylinux_2_34_aarch64.whl", hash = "sha256:ecfed7367f965a0328cfbdd70da860f15441f002f613185668c6e6ebf5a0ac11", size = 4707038, upload-time = "2026-07-31T14:24:12.169Z" }, + { url = "https://files.pythonhosted.org/packages/01/77/2cf79bbfc4d12ca106437a6e170d6aaa01a373e93093118aaaef0e801bd4/cryptography-50.0.0-cp314-cp314t-manylinux_2_34_ppc64le.whl", hash = "sha256:9aa87839c383bdbab6ef865787a1fb877af8dd03464c4400322726feaaadfc6d", size = 5273110, upload-time = "2026-07-31T14:24:14.38Z" }, + { url = "https://files.pythonhosted.org/packages/e5/45/8aae2972c520145377ea3559a605a899bebe227bf070b33cdb445929a9b9/cryptography-50.0.0-cp314-cp314t-manylinux_2_34_x86_64.whl", hash = "sha256:6ba6a53445bd3cfa809ef3ef5f1589aa6ba08784a1d962bf47d0940e871dab1c", size = 4716439, upload-time = "2026-07-31T14:24:16.415Z" }, + { url = "https://files.pythonhosted.org/packages/7b/20/4fe50b619a48c2525cc46e2dbc1ac490708d704be5d467bdaac6dc955682/cryptography-50.0.0-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:3f5735ffe4996d28b809371756219f5354864902a3b9e7c0b9ee87041209fc9c", size = 4837383, upload-time = "2026-07-31T14:24:18.553Z" }, + { url = "https://files.pythonhosted.org/packages/92/91/3a31366e183343d3703f8995c095f5734676bd6938118047e50fcf279eb4/cryptography-50.0.0-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:1b4a266766514614f8aa60416e71f2fc6e575d36e7bdc90f644fadb2f4b75b95", size = 4985772, upload-time = "2026-07-31T14:24:20.385Z" }, + { url = "https://files.pythonhosted.org/packages/74/9a/02ffe35b2853d121689871eb5dce862092562b3a1ed5cc98f1aaed441506/cryptography-50.0.0-cp314-cp314t-win_amd64.whl", hash = "sha256:12b9c6996425c76ea6c457ace4f3073e715b8c545add07cd1a8f3a4f90691269", size = 3816291, upload-time = "2026-07-31T14:24:22.125Z" }, + { url = "https://files.pythonhosted.org/packages/03/37/73d005be173aff344af30e9fd2a576575cb2391a7101d9cd3842e1fa8cce/cryptography-50.0.0-cp39-abi3-macosx_11_0_arm64.whl", hash = "sha256:ccdc4a71a4dabae05de219404f9f4abc38e3b58422177ff93d0da05967dafa07", size = 4036009, upload-time = "2026-07-31T14:24:24.122Z" }, + { url = "https://files.pythonhosted.org/packages/ff/c6/7a6202a534e32103a285b7834a120869557fe198d51d7cfe59754c8bda9c/cryptography-50.0.0-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:910e1d2668e7de9648f2bcee30e180db2a6b15c30f887d7c4c93ddf96e3992e3", size = 4745252, upload-time = "2026-07-31T14:24:26.118Z" }, + { url = "https://files.pythonhosted.org/packages/85/4f/0fa8c2f4428198f15d9ff8d63400e27afbf94ce833f6108da1eb3753f945/cryptography-50.0.0-cp39-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:a91296cb61e8df6f86d0c19cc4068228da256bf59bf86049fbd821084565327f", size = 4728939, upload-time = "2026-07-31T14:24:27.994Z" }, + { url = "https://files.pythonhosted.org/packages/d1/63/54dd723490ba2dc09b299682c10b38db38f159728bcaae8c591b8af2f22d/cryptography-50.0.0-cp39-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:e722f16708d854fe924790e051061f6704a472c3bac347b6fd88033ea8dd0dc5", size = 4748483, upload-time = "2026-07-31T14:24:30.254Z" }, + { url = "https://files.pythonhosted.org/packages/1d/dd/7c77d26285cc7f6991efce64a0f5b4f9383bfa5dd8c5033003eaf7db4cdb/cryptography-50.0.0-cp39-abi3-manylinux_2_28_ppc64le.whl", hash = "sha256:d764dcf130c428ef66786f866dd750f53182bc608813489915e9fc106bb0c82f", size = 5367599, upload-time = "2026-07-31T14:24:32.457Z" }, + { url = "https://files.pythonhosted.org/packages/46/c9/f60aed34c013f317f92817b6c171c2d22a78270fa41109bd4b08af26b194/cryptography-50.0.0-cp39-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:105110f43a471dbd0060b9c9516cb8a6a79233631a04cc2ba16f28323ac6e025", size = 4762647, upload-time = "2026-07-31T14:24:34.599Z" }, + { url = "https://files.pythonhosted.org/packages/be/f3/f9a0173b139372c3a48ed98154b45cc6b9de17c789d5ab552e621c293609/cryptography-50.0.0-cp39-abi3-manylinux_2_31_armv7l.whl", hash = "sha256:828743d939e9629bc267b8e2d08d8bb67cd4319c771a33d4b18b22dd8fb7440a", size = 4385197, upload-time = "2026-07-31T14:24:36.647Z" }, + { url = "https://files.pythonhosted.org/packages/d8/36/83bb81f6e569bc38e1e4a7bc80f29b46bb9601920bc455fc8e888f5d5742/cryptography-50.0.0-cp39-abi3-manylinux_2_34_aarch64.whl", hash = "sha256:2a8183b489dc1f7f80f135780fadc1108f14b31b8a40411c7a5b17425f65f28b", size = 4748095, upload-time = "2026-07-31T14:24:39.493Z" }, + { url = "https://files.pythonhosted.org/packages/6b/16/d3008eff98c764979865834c3d386d4fd041b5f52e7f34fc29ac1a5eb515/cryptography-50.0.0-cp39-abi3-manylinux_2_34_ppc64le.whl", hash = "sha256:6e7d61120573a7f2cd94cc095f9e81f6967c61ccdf194285aa143ecec8e0b708", size = 5325948, upload-time = "2026-07-31T14:24:41.556Z" }, + { url = "https://files.pythonhosted.org/packages/9c/f8/d97f9603efda3888187bfdb893f26c41be4735c10631d05d284ee6b047c4/cryptography-50.0.0-cp39-abi3-manylinux_2_34_x86_64.whl", hash = "sha256:37fdb0d0111f1e2ff07139dfb79f1b49531f8e213c46f1163dd7642979b58c47", size = 4762400, upload-time = "2026-07-31T14:24:43.636Z" }, + { url = "https://files.pythonhosted.org/packages/64/a2/4615c8f7d81a00b1d6e6afe19f694e1543582349fb5f4076f6cb5dc36485/cryptography-50.0.0-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:c87f62a3d3b9888ed0fdde100ec06aa61ca9cd44bad9057d1dff9a516b5f5bb9", size = 4878208, upload-time = "2026-07-31T14:24:45.522Z" }, + { url = "https://files.pythonhosted.org/packages/d2/1a/efcfb02f91407149a0dacffffab791f7e19bf6385f63b3666dc8b5e5c9c8/cryptography-50.0.0-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:65c2c3add92b45fd0709db8594536aea39c2a67af0e27ffcf049c498501140b7", size = 5037050, upload-time = "2026-07-31T14:24:47.697Z" }, + { url = "https://files.pythonhosted.org/packages/57/30/4a22984d4f1bdfb8c054f07a92bc176b97a3134cc1d6c4b3bffb1f3688b4/cryptography-50.0.0-cp39-abi3-win_amd64.whl", hash = "sha256:d24fead1d4d076e1bfb006dcec392074a3cd8d7b4fc8a595aa64073b2b7a96ba", size = 3874135, upload-time = "2026-07-31T14:24:50.085Z" }, ] [[package]] @@ -581,7 +587,7 @@ wheels = [ [[package]] name = "dify-agent" -version = "1.16.0" +version = "1.16.1" source = { editable = "." } dependencies = [ { name = "httpx" }, @@ -592,11 +598,8 @@ dependencies = [ ] [package.optional-dependencies] -grpc = [ - { name = "grpclib", extra = ["protobuf"] }, - { name = "protobuf" }, -] server = [ + { name = "e2b" }, { name = "fastapi" }, { name = "graphon" }, { name = "jsonschema" }, @@ -612,7 +615,6 @@ server = [ dev = [ { name = "basedpyright" }, { name = "coverage" }, - { name = "grpcio-tools" }, { name = "pytest" }, { name = "pytest-examples" }, { name = "pytest-mock" }, @@ -627,15 +629,14 @@ docs = [ [package.metadata] requires-dist = [ + { name = "e2b", marker = "extra == 'server'", specifier = ">=2.34.0,<3.0.0" }, { name = "fastapi", marker = "extra == 'server'", specifier = "==0.136.0" }, { name = "graphon", marker = "extra == 'server'", specifier = "==0.5.2" }, - { name = "grpclib", extras = ["protobuf"], marker = "extra == 'grpc'", specifier = ">=0.4.9,<0.5.0" }, { name = "httpx", specifier = "==0.28.1" }, { name = "httpx2", specifier = ">=2.5.0,<3.0.0" }, { name = "jsonschema", marker = "extra == 'server'", specifier = ">=4.23.0,<5.0.0" }, { name = "jwcrypto", marker = "extra == 'server'", specifier = ">=1.5.6,<2" }, { name = "logfire", extras = ["fastapi", "httpx", "redis"], marker = "extra == 'server'", specifier = ">=4.37.0,<5.0.0" }, - { name = "protobuf", marker = "extra == 'grpc'", specifier = ">=6.33.5,<7.0.0" }, { name = "pydantic", specifier = ">=2.12.5,<2.13" }, { name = "pydantic-ai-slim", specifier = ">=1.102.0,<2.0.0" }, { name = "pydantic-ai-slim", extras = ["anthropic", "google", "openai"], marker = "extra == 'server'", specifier = ">=1.85.1,<2.0.0" }, @@ -644,13 +645,12 @@ requires-dist = [ { name = "typing-extensions", specifier = ">=4.12.2,<5.0.0" }, { name = "uvicorn", extras = ["standard"], marker = "extra == 'server'", specifier = "==0.46.0" }, ] -provides-extras = ["grpc", "server"] +provides-extras = ["server"] [package.metadata.requires-dev] dev = [ { name = "basedpyright", specifier = ">=1.39.3" }, { name = "coverage", extras = ["toml"], specifier = ">=7.10.7" }, - { name = "grpcio-tools", specifier = ">=1.81.0,<2.0.0" }, { name = "pytest", specifier = ">=9.0.3" }, { name = "pytest-examples", specifier = ">=0.0.18" }, { name = "pytest-mock", specifier = ">=3.14.0" }, @@ -672,6 +672,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/12/b3/231ffd4ab1fc9d679809f356cebee130ac7daa00d6d6f3206dd4fd137e9e/distro-1.9.0-py3-none-any.whl", hash = "sha256:7bffd925d65168f85027d8da9af6bddab658135b840670a223589bc0c8ef02b2", size = 20277, upload-time = "2023-12-24T09:54:30.421Z" }, ] +[[package]] +name = "dockerfile-parse" +version = "2.0.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/92/df/929ee0b5d2c8bd8d713c45e71b94ab57c7e11e322130724d54f469b2cd48/dockerfile-parse-2.0.1.tar.gz", hash = "sha256:3184ccdc513221983e503ac00e1aa504a2aa8f84e5de673c46b0b6eee99ec7bc", size = 24556, upload-time = "2023-07-18T13:36:07.897Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7a/6c/79cd5bc1b880d8c1a9a5550aa8dacd57353fa3bb2457227e1fb47383eb49/dockerfile_parse-2.0.1-py2.py3-none-any.whl", hash = "sha256:bdffd126d2eb26acf1066acb54cb2e336682e1d72b974a40894fac76a4df17f6", size = 14845, upload-time = "2023-07-18T13:36:06.052Z" }, +] + [[package]] name = "docstring-parser" version = "0.18.0" @@ -681,6 +690,28 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a7/5f/ed01f9a3cdffbd5a008556fc7b2a08ddb1cc6ace7effa7340604b1d16699/docstring_parser-0.18.0-py3-none-any.whl", hash = "sha256:b3fcbed555c47d8479be0796ef7e19c2670d428d72e96da63f3a40122860374b", size = 22484, upload-time = "2026-04-14T04:09:18.638Z" }, ] +[[package]] +name = "e2b" +version = "2.34.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "attrs" }, + { name = "dockerfile-parse" }, + { name = "h2" }, + { name = "httpcore" }, + { name = "httpx" }, + { name = "packaging" }, + { name = "protobuf" }, + { name = "python-dateutil" }, + { name = "rich" }, + { name = "typing-extensions" }, + { name = "wcmatch" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/5b/22/5bcf34304b62e6984df9089eb996a4b2aa63618f901befe1f7cc73a24437/e2b-2.34.0.tar.gz", hash = "sha256:0dcc2694b3ead62e87b3d1070a05f66842dd2fa320fc1d7bd84c8fcfadeaf1b1", size = 189367, upload-time = "2026-07-17T09:59:23.844Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/6d/ec/b2f9297afc25de7c27ab5a85db4060783680e2b132285e2498fdbad484cf/e2b-2.34.0-py3-none-any.whl", hash = "sha256:873323571d18bf633be45e59fc6271410b30dfbc81e8df85e711f4f184c03fea", size = 343447, upload-time = "2026-07-17T09:59:22.553Z" }, +] + [[package]] name = "emoji" version = "2.15.0" @@ -864,108 +895,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/11/8c/c9138d881c79aa0ea9ed83cbd58d5ca75624378b38cee225dcf5c42cc91f/griffelib-2.0.2-py3-none-any.whl", hash = "sha256:925c857658fb1ba40c0772c37acbc2ab650bd794d9c1b9726922e36ea4117ea1", size = 142357, upload-time = "2026-03-27T11:34:46.275Z" }, ] -[[package]] -name = "grpcio" -version = "1.81.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "typing-extensions" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/15/f3/23f47b24f8d8c2028eba501db3acfbb2f592cbb5995eaa6e363a627b74d7/grpcio-1.81.0.tar.gz", hash = "sha256:a5acd7efd3b1fe9b4eb0bcaaa1507eed68a0ad0678b654c3f7b464df9ba9dca5", size = 13032272, upload-time = "2026-06-01T05:56:22.827Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/82/d5/896a3aaf07068d707d88b282a04914b872db4d32d3c7e6d88e43a3b911fa/grpcio-1.81.0-cp312-cp312-linux_armv7l.whl", hash = "sha256:57b3b0e73a518fa286959b40c3eddd02703504ca186e8b7b2945954519bd8b2c", size = 6053538, upload-time = "2026-06-01T05:54:58.965Z" }, - { url = "https://files.pythonhosted.org/packages/68/6a/7e3eafa4727cd405ff917605ed2949e2af162f233f5cbdd773723a5fea7d/grpcio-1.81.0-cp312-cp312-macosx_11_0_universal2.whl", hash = "sha256:8bb1789c94322a13336a2b6c58d9c14d68f8628b6e24205a799c69f5bf8516ce", size = 12053447, upload-time = "2026-06-01T05:55:01.862Z" }, - { url = "https://files.pythonhosted.org/packages/16/79/a4302aa82428de48a922421f522b027a1a727ab4d0926368454aa953d36d/grpcio-1.81.0-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:e4d053900a0d24b75d7521139a3872150301b3d6bde3bed5e12318fb25791e4d", size = 6595872, upload-time = "2026-06-01T05:55:04.946Z" }, - { url = "https://files.pythonhosted.org/packages/b4/1f/7ff2850eaefbecf99af3f624dbb28dd1ad6c5fd4c1d8c26909ed6482673b/grpcio-1.81.0-cp312-cp312-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:db217c2e52931719f9937bd12082cd4d7b495b35803d5760686975c285924bf8", size = 7303857, upload-time = "2026-06-01T05:55:07.205Z" }, - { url = "https://files.pythonhosted.org/packages/e2/98/1f3896a9baae1f2aedf4e99c55291d6fa1f30ad9603d63bc18bda967b53e/grpcio-1.81.0-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:19f201da7b4e5c0559198abe5a97157e726f3abe6e8f5e832d4a50740f6dcc22", size = 6809676, upload-time = "2026-06-01T05:55:09.513Z" }, - { url = "https://files.pythonhosted.org/packages/34/8b/3441983718095208c5d797fd3239882e97ea89a629f41c8df94b4eef4df9/grpcio-1.81.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:275144b0115353339dbb8a6f28a9cf8997b5bf40e37f8f66ac0b0ea57e95b43f", size = 7412654, upload-time = "2026-06-01T05:55:12.777Z" }, - { url = "https://files.pythonhosted.org/packages/3c/98/1eddf07df6e4fe85cf67502a793f7b05468b2dca3d1ef35b972cf5d54468/grpcio-1.81.0-cp312-cp312-musllinux_1_2_i686.whl", hash = "sha256:5192857589f223e5a98ff0e31f6e551b19040e647d17bfe10116c8a2ce3b8696", size = 8408026, upload-time = "2026-06-01T05:55:15.514Z" }, - { url = "https://files.pythonhosted.org/packages/5c/73/3860341e6a1f5347be6ab35c6c0e1e3a8eb59d010388207fd561dcf01a88/grpcio-1.81.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:c6ff087cb1f563f47b504b4e29e684129fc5ae4863faf3ebca08a327764ee6cb", size = 7849498, upload-time = "2026-06-01T05:55:18.078Z" }, - { url = "https://files.pythonhosted.org/packages/ae/3f/0ea06bd85c701966aa3f8f37314f2ed83520d2b7590f42d643d445d8bc8b/grpcio-1.81.0-cp312-cp312-win32.whl", hash = "sha256:98c6240f563178fc5877bd50e6ff274463e53e1472128f4110742450739659fa", size = 4184161, upload-time = "2026-06-01T05:55:20.127Z" }, - { url = "https://files.pythonhosted.org/packages/39/e3/a7c387406827a86f99ad7838b995bf9b4a182ffe2d2c439ed2873efec952/grpcio-1.81.0-cp312-cp312-win_amd64.whl", hash = "sha256:87e33b7afcfb3585121b5f007d2c52b8c534104d18f556e840d35193ca2a9141", size = 4929958, upload-time = "2026-06-01T05:55:22.736Z" }, - { url = "https://files.pythonhosted.org/packages/f3/29/779ee53c931d0fd55c1d459fde43e485172caa3ac87cbd43d003a13a0185/grpcio-1.81.0-cp313-cp313-linux_armv7l.whl", hash = "sha256:62bbe463c9f0f2ff24e31bd25f8dd8b4bae78900e315915a3195a0ef1471a855", size = 6054973, upload-time = "2026-06-01T05:55:25.043Z" }, - { url = "https://files.pythonhosted.org/packages/9e/b6/7211807926b5a17f8d9a5d47c739a163d6812fefe3e4714e81cf92945ed7/grpcio-1.81.0-cp313-cp313-macosx_11_0_universal2.whl", hash = "sha256:43c121e135ae44d1559b430db2b2dfad7421cbbe40e1deba506c7dc62b439719", size = 12048662, upload-time = "2026-06-01T05:55:28.453Z" }, - { url = "https://files.pythonhosted.org/packages/64/89/b1b93ef6b34bd20bbaf707fa99133bc9cc302139d5ec6f77a165c7169796/grpcio-1.81.0-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:f345de40ef2e65f63645d53d251824e6070e07804827c5b00ec2e44555f9f901", size = 6599116, upload-time = "2026-06-01T05:55:31.185Z" }, - { url = "https://files.pythonhosted.org/packages/eb/bc/c89f9b9d1c22895715356a1e009554dae66319e97826bb4d30bcda7d29e8/grpcio-1.81.0-cp313-cp313-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:8c0855a350886f713b9e458e2a10d208009dcaa849f574e39cd6067db1fe1279", size = 7307591, upload-time = "2026-06-01T05:55:33.463Z" }, - { url = "https://files.pythonhosted.org/packages/65/4a/1df2a4cb4a1386e066ab7e4175e34bb884b35ccb60d3621c09c84af6aabb/grpcio-1.81.0-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:a524cd530900bd24511fcb7f2ed144da4ea37711c4b094475d0bceca7a93a170", size = 6811797, upload-time = "2026-06-01T05:55:36.731Z" }, - { url = "https://files.pythonhosted.org/packages/8d/dc/fa189d20601a1be25b08850cfb733879bbb1047b62a8feec3a60e3e1a87b/grpcio-1.81.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:e7746ba3e6efc9e2b748eff59470a2b8684d5a9ec607c6580bcaa5be175820bc", size = 7415131, upload-time = "2026-06-01T05:55:39.451Z" }, - { url = "https://files.pythonhosted.org/packages/ad/a3/5625c48cb48d23c6631b3e5294f88e4c751f22a52591ae78859fab96dca1/grpcio-1.81.0-cp313-cp313-musllinux_1_2_i686.whl", hash = "sha256:aaaa4f7f2057d795952e4eacf3f342be8b5b156992f6ac85023c8b98794ebd47", size = 8408398, upload-time = "2026-06-01T05:55:42.219Z" }, - { url = "https://files.pythonhosted.org/packages/75/34/0f8202c6809a46c2b4d69125ef3667c40b1c211f8e19930e5fa1f1197039/grpcio-1.81.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:0fba53cb96004b2b7fb758b46b2288cb49d0b658316a4e73f3ef67230616ee65", size = 7844481, upload-time = "2026-06-01T05:55:44.849Z" }, - { url = "https://files.pythonhosted.org/packages/c0/95/c3366b5b5edf4c4adc90f2e29ca16e57965a8e56dc8d2ee89565ba1905bb/grpcio-1.81.0-cp313-cp313-win32.whl", hash = "sha256:c197e2ef75a442528072b29e9755da299110e8610e8bcbb59a6b4cf55384f005", size = 4182777, upload-time = "2026-06-01T05:55:47.459Z" }, - { url = "https://files.pythonhosted.org/packages/a9/a7/932f2f748511a32e641a2aba0d30dded3ed6e8bc330e0924e4d5d86853e6/grpcio-1.81.0-cp313-cp313-win_amd64.whl", hash = "sha256:194eddfacc84d80f50512e9fd4ee851d5f2499f18f299c95aa8fb4748f0537e0", size = 4928085, upload-time = "2026-06-01T05:55:50.158Z" }, - { url = "https://files.pythonhosted.org/packages/c5/1d/28b231333857deb840bc3d182ae087510170ea6d68f21393aeb0fe499530/grpcio-1.81.0-cp314-cp314-linux_armv7l.whl", hash = "sha256:a9351055f52660b58f3d4890ea66188b5134399f82b11aa0c55bd4b99eff5390", size = 6055712, upload-time = "2026-06-01T05:55:52.889Z" }, - { url = "https://files.pythonhosted.org/packages/e8/b8/999c14f9dff0fc47549d2e827cba1343ddc18e1d1bf0d06d2cf628eecbd9/grpcio-1.81.0-cp314-cp314-macosx_11_0_universal2.whl", hash = "sha256:300f3337b6425fd16ead9a4f9b2ac25801acb64aa5bc0b99eb69901645b2b1d2", size = 12057189, upload-time = "2026-06-01T05:55:55.952Z" }, - { url = "https://files.pythonhosted.org/packages/1e/3d/1fbde079572562af65351151d840525a13879eb7b481d35b55cd64c6127a/grpcio-1.81.0-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:97bbd623f7ded558fd4f7cb5a4f600c4d4de65c5dd364c83a5b14b2a10a2d3b5", size = 6608136, upload-time = "2026-06-01T05:55:59.069Z" }, - { url = "https://files.pythonhosted.org/packages/32/89/1f17cb6882abfd8e5a303a25d5d1665abef5a8c499a96198c65a651d1b85/grpcio-1.81.0-cp314-cp314-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:ff83d889e3ebf6341c8c7864ad8031591ad5ca61599072fc511644d1eb962d2b", size = 7307045, upload-time = "2026-06-01T05:56:02.376Z" }, - { url = "https://files.pythonhosted.org/packages/48/5a/f98e91b2e755652e637ea2144318b0229b290062199f761b445fe1fa6015/grpcio-1.81.0-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:c4fe218c5a35e1d87a5a26544237f1fa41dfd9cbd3c856b0810a30061f8b0aaf", size = 6812794, upload-time = "2026-06-01T05:56:05.777Z" }, - { url = "https://files.pythonhosted.org/packages/0a/0c/77892d715ac41e7ec0ace2a50080ffb64e189188056f607a66fe0014d1ee/grpcio-1.81.0-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:b8b025b6af43ee0ad4a70307025d77bcab5adde7c4597786010d802c203e9fc5", size = 7422767, upload-time = "2026-06-01T05:56:08.524Z" }, - { url = "https://files.pythonhosted.org/packages/3f/b8/aa04590c6564714d94954515f15a236e59d4b9b3ad01e615f1b706d7792d/grpcio-1.81.0-cp314-cp314-musllinux_1_2_i686.whl", hash = "sha256:3d4e0ce5a40a998cf608c8ba60ecfe18fdf364a9aa193ae4ac3faeecd0e86757", size = 8408551, upload-time = "2026-06-01T05:56:11.283Z" }, - { url = "https://files.pythonhosted.org/packages/43/3d/4f4a3450a1973568910c6909cb74abbf2126f68aefae5976962f9f7ad50d/grpcio-1.81.0-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:aa948712c8e5fa40ec250870bda14bc7578e1bb832a8912d9d2a0f720518edbe", size = 7846468, upload-time = "2026-06-01T05:56:14.536Z" }, - { url = "https://files.pythonhosted.org/packages/88/f4/5827fd248221ad3b44161c23ce9b5f4ee405b04fc6da5fd402a9aa87a84a/grpcio-1.81.0-cp314-cp314-win32.whl", hash = "sha256:fbbe81314a9d92156abce8b62c09364eb8bafc0ca2a19919a45ec64b5c6cb664", size = 4264427, upload-time = "2026-06-01T05:56:17.192Z" }, - { url = "https://files.pythonhosted.org/packages/0c/e8/127dc2b246096ad50ef7c8d9b7b31d757787aeb796368bcdd4454e4204c4/grpcio-1.81.0-cp314-cp314-win_amd64.whl", hash = "sha256:b93cee313cae4e113fbb3a0ce1ea5633db6f63cfde2b2dc1d817429026b2a50b", size = 5070848, upload-time = "2026-06-01T05:56:19.735Z" }, -] - -[[package]] -name = "grpcio-tools" -version = "1.81.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "grpcio" }, - { name = "protobuf" }, - { name = "setuptools" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/f3/b5/72f688670ce56ea59b05ea13430f06cbb728dd354dac508544fc7d4b5c95/grpcio_tools-1.81.0.tar.gz", hash = "sha256:0733d773eca8cb461f4f2a1b79c64c123db9661be41b08184b81497b2b991ccb", size = 6235718, upload-time = "2026-06-01T05:58:34.191Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/1c/e3/3c4f9da489413ef3f3dc9f7bc49a270270ec99fb3a00fd4302a2f59a7be2/grpcio_tools-1.81.0-cp312-cp312-linux_armv7l.whl", hash = "sha256:374fe0435447f283e3c69f719ef1bc66f2e187a239ce25444b2de45cb3a6a744", size = 2585927, upload-time = "2026-06-01T05:57:25.397Z" }, - { url = "https://files.pythonhosted.org/packages/d2/b2/de7aba18f87d722c215ae168add975b9e7729cfaf7a1292be43f87685fa1/grpcio_tools-1.81.0-cp312-cp312-macosx_11_0_universal2.whl", hash = "sha256:dce33d09851bc15dead814bd9d21023bffdf0f838ecf65995a2456e5692831fd", size = 5815566, upload-time = "2026-06-01T05:57:27.714Z" }, - { url = "https://files.pythonhosted.org/packages/eb/13/8f71b4830f129d896560c66964a3a8f4e33fbd59854396015e7449b75d3a/grpcio_tools-1.81.0-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:58e65dd13f7ffc25f5a9cd9890fddfd39b3c51e5c3c1acd987813b1dc1173704", size = 2635519, upload-time = "2026-06-01T05:57:29.848Z" }, - { url = "https://files.pythonhosted.org/packages/a2/e0/3ad58f1791c346a1fefc69ef3fcd19d63e3778736d4746f12b39f900b78b/grpcio_tools-1.81.0-cp312-cp312-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:760775f79bbefa321cd327fe1019a9d6ad0d93acb1ede7b7905c712679542fd7", size = 2958250, upload-time = "2026-06-01T05:57:31.836Z" }, - { url = "https://files.pythonhosted.org/packages/79/8a/3212db57815df0fa2a02e857e402c1abea15ec6b5fb63ebf306d90f2fb07/grpcio_tools-1.81.0-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:7af4faf34376c57f4f3c42aa05f065d3b10e774b8a8d8b27d659d5cc351e5c75", size = 2698437, upload-time = "2026-06-01T05:57:33.827Z" }, - { url = "https://files.pythonhosted.org/packages/ec/d5/4245bbb4c14b54ac539b5b59f5298750d75211e144e1e8b35e1af5144d6d/grpcio_tools-1.81.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:e9945945edc14022abab07caec0ebc16bf51219e586b4f008d09334cc479655c", size = 3152159, upload-time = "2026-06-01T05:57:36.065Z" }, - { url = "https://files.pythonhosted.org/packages/59/b1/59500d9fe41209e0887c66d80879ba80d0cc8e1327e24cc783eb879fe7a7/grpcio_tools-1.81.0-cp312-cp312-musllinux_1_2_i686.whl", hash = "sha256:61d9ad6b5c0f3857663701bdc2cbb3da7f7d835ce9a8d597ffca443124a96894", size = 3710468, upload-time = "2026-06-01T05:57:38.367Z" }, - { url = "https://files.pythonhosted.org/packages/0c/14/86fc8b64db62851bf5cb1c945b22da7ab0dec6e0b7002ec374247482d404/grpcio_tools-1.81.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:ffa74b84d201bea407f22e00ac0367f12df48a0073b0ffd9f3be9d8126a55245", size = 3370795, upload-time = "2026-06-01T05:57:40.724Z" }, - { url = "https://files.pythonhosted.org/packages/94/95/4fd57d9f948adadbe5b3e8e3b0d3c121ae5fe8276721097ab03361ce7adf/grpcio_tools-1.81.0-cp312-cp312-win32.whl", hash = "sha256:f1f407697873acbf1d961c6fb9223114a3e679938469d4186623b8b872dbdae0", size = 1008449, upload-time = "2026-06-01T05:57:42.491Z" }, - { url = "https://files.pythonhosted.org/packages/33/94/7567ddd3a13e24bbc5e146c6ac735004ab7303048115ccb032d61b5c2305/grpcio_tools-1.81.0-cp312-cp312-win_amd64.whl", hash = "sha256:283bb3465331a4034b14dce35425c47b0cfbd287b09a6e9d15c9f26fbb17e799", size = 1174889, upload-time = "2026-06-01T05:57:44.405Z" }, - { url = "https://files.pythonhosted.org/packages/f2/05/f0606a1b2e830d5fddfcd77c5d8e928f26dc221ced386fccf31a6efda57e/grpcio_tools-1.81.0-cp313-cp313-linux_armv7l.whl", hash = "sha256:33c579bf24040dbdce751e05b5bcedc13dafffa8f2dfa07193bc05960dc95b49", size = 2586070, upload-time = "2026-06-01T05:57:46.677Z" }, - { url = "https://files.pythonhosted.org/packages/c4/4a/3b6817547d65d9f7a106ea6a2352125d08b44ce1d120b64ca0c565d896e6/grpcio_tools-1.81.0-cp313-cp313-macosx_11_0_universal2.whl", hash = "sha256:2a369a9e27fcb6279387744778d67a64551de82db0e9cf7e1e9451c22f719e07", size = 5813211, upload-time = "2026-06-01T05:57:49.159Z" }, - { url = "https://files.pythonhosted.org/packages/bf/fc/078422558ebb337233379cdb0e4cc0d3d3218933d105003bae2790ff976f/grpcio_tools-1.81.0-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:b2b26111795d0e7a72fa483d377a75b57203edaac8bbbac1e9b8f174e7b01ff3", size = 2634663, upload-time = "2026-06-01T05:57:51.41Z" }, - { url = "https://files.pythonhosted.org/packages/53/2b/043f2d62d6f28a0962f29f180ae24770020a06985b14a2d8f0d489531d71/grpcio_tools-1.81.0-cp313-cp313-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:a741007631248dd903f880b575a63c0662433ed4b966262ab0f437ea860c9b7e", size = 2957926, upload-time = "2026-06-01T05:57:54.103Z" }, - { url = "https://files.pythonhosted.org/packages/76/d3/be8c1f7c5ca6adccba66ef787b4bba304a3247c1319a33ed330719931ba3/grpcio_tools-1.81.0-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:bac1c6ddc5eb3762257e773829b4e00ea7d0592f02498f27e5bdb75cd3182d69", size = 2697761, upload-time = "2026-06-01T05:57:56.494Z" }, - { url = "https://files.pythonhosted.org/packages/19/9d/c1650c72059f7d20d94597430ecbe0139d92c7e409cf007d0b5d765e40ec/grpcio_tools-1.81.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:e5d241c694a226bada06dae9accc47c5c25e586f49464823498e84cff63eedea", size = 3151460, upload-time = "2026-06-01T05:57:58.756Z" }, - { url = "https://files.pythonhosted.org/packages/6b/82/2aad863738dc4f749df93d97f81a8aebdd0a7c5daee7157c259ecea7dbbc/grpcio_tools-1.81.0-cp313-cp313-musllinux_1_2_i686.whl", hash = "sha256:475286558410e4d0394fc8129559080c22835ece3c23cac194b46c5ad7d0fad1", size = 3710466, upload-time = "2026-06-01T05:58:01.327Z" }, - { url = "https://files.pythonhosted.org/packages/47/7d/b79a5b132bd5db76ce48e686579ded0311bfa7ea19b8cb9ca88aea3e8d64/grpcio_tools-1.81.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:98b57a592a75553f4dd38db3993b037953245991b8df9f06927a8978962dcd82", size = 3370487, upload-time = "2026-06-01T05:58:03.769Z" }, - { url = "https://files.pythonhosted.org/packages/96/6f/a6a74a51b71aca801f4dcce6bd8bb7d3fd8d556ff191f0bd0de631b97dd1/grpcio_tools-1.81.0-cp313-cp313-win32.whl", hash = "sha256:78c2514bf172b20631685840fea0c5d1bec5518b649d0a498f6dd8f91dfce56a", size = 1008234, upload-time = "2026-06-01T05:58:05.668Z" }, - { url = "https://files.pythonhosted.org/packages/5a/91/02a4c529dd0a77c2d768b415ac334f6a9ea9d61c6f4c38e54d1a4394530c/grpcio_tools-1.81.0-cp313-cp313-win_amd64.whl", hash = "sha256:a87ea8056beea56b24353d27b7f0ab814daabb372aa517d2e179470e66fd8f6b", size = 1174523, upload-time = "2026-06-01T05:58:07.711Z" }, - { url = "https://files.pythonhosted.org/packages/13/1f/7885e23074d813ab71ba3ea689ecf5cb3bb3c76c51cc01bd393451f10257/grpcio_tools-1.81.0-cp314-cp314-linux_armv7l.whl", hash = "sha256:bf2ffd98c754d54a5affddfe3a18d4822b972c5c37ca6a33a456c94e6f3dd82b", size = 2585943, upload-time = "2026-06-01T05:58:10.077Z" }, - { url = "https://files.pythonhosted.org/packages/40/f4/a88116147d377a88fccef9b43667235835d504d60ebe0e0dcc81bc9c4b20/grpcio_tools-1.81.0-cp314-cp314-macosx_11_0_universal2.whl", hash = "sha256:4a8d9fdf42537cd1924f7e95013771aa73c89501ccf112b7517c22ca825be9f5", size = 5813367, upload-time = "2026-06-01T05:58:12.545Z" }, - { url = "https://files.pythonhosted.org/packages/2a/1e/01f310f0427dcddaf0097e4101f041d437ba8199a7ed62621788f2601042/grpcio_tools-1.81.0-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:564624654f5a0377adc69f226f93ec5ff52715030c2f846af6c54ccc4ca2d225", size = 2634992, upload-time = "2026-06-01T05:58:14.92Z" }, - { url = "https://files.pythonhosted.org/packages/fd/ee/10cec754cb89cf45039ef4a8ff500ab5c567278fc4f6333347dba20c99fb/grpcio_tools-1.81.0-cp314-cp314-manylinux2014_i686.manylinux_2_17_i686.whl", hash = "sha256:3f1ac691debbdbf00e7635ecacf28647f48f67fb4b14120a4541c8bb0f8333ae", size = 2957912, upload-time = "2026-06-01T05:58:17.596Z" }, - { url = "https://files.pythonhosted.org/packages/d2/71/cc273059fa3620d424df91a8faabd5875d4c20a0300875e721d15decd74b/grpcio_tools-1.81.0-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:f35da7b3b9537ecce9a5cfd967b1316a815378d5a3729e9ea556b22a14f6e315", size = 2697709, upload-time = "2026-06-01T05:58:19.653Z" }, - { url = "https://files.pythonhosted.org/packages/a1/d5/c72f9e7d18425586bbc5b4fba78570fab23cb9810fb34a8841f7baaef22b/grpcio_tools-1.81.0-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:1e21049aeedd9d62e5e58146e5d00f3b75674d8d8d3cc69709ab950082be8421", size = 3151885, upload-time = "2026-06-01T05:58:22.366Z" }, - { url = "https://files.pythonhosted.org/packages/9e/f9/7a3d7b3bce72fe22611a79d3790440c16af9d524a8bd1d38a89a44c65570/grpcio_tools-1.81.0-cp314-cp314-musllinux_1_2_i686.whl", hash = "sha256:019f05a17b5495d561603f75a74a4a76ad22456a95dc7623c7be4e4b44391b88", size = 3710403, upload-time = "2026-06-01T05:58:25.014Z" }, - { url = "https://files.pythonhosted.org/packages/58/2d/b41fe47b83eb197a48fdcbf48d04f5923c5fd62d4b1a7f2820720562d7ae/grpcio_tools-1.81.0-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:3401d6d4668c064c9ed344825fe97ff076cd8ba719a24e3d3c7169b604cbf00f", size = 3370524, upload-time = "2026-06-01T05:58:27.096Z" }, - { url = "https://files.pythonhosted.org/packages/94/d3/a4565146ab83232ddf86a3de497937662dc06649739f978763379c256311/grpcio_tools-1.81.0-cp314-cp314-win32.whl", hash = "sha256:5783b6758244f6eaceb41a0e651828824b0d0c92724ddf4b68879ced9bfd9b50", size = 1030574, upload-time = "2026-06-01T05:58:28.972Z" }, - { url = "https://files.pythonhosted.org/packages/35/72/f2102f3737b94e14b8ce56394918c0a4303e80d5d8b0627fc4ee85927f79/grpcio_tools-1.81.0-cp314-cp314-win_amd64.whl", hash = "sha256:69f8355b723db7b5e26a3bff76f9deb3a407d22fe289bca486ccf95d6133cad0", size = 1207499, upload-time = "2026-06-01T05:58:31.606Z" }, -] - -[[package]] -name = "grpclib" -version = "0.4.9" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "h2" }, - { name = "multidict" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/5b/28/5a2c299ec82a876a252c5919aa895a6f1d1d35c96417c5ce4a4660dc3a80/grpclib-0.4.9.tar.gz", hash = "sha256:cc589c330fa81004c6400a52a566407574498cb5b055fa927013361e21466c46", size = 84798, upload-time = "2025-12-14T22:23:14.349Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/5c/90/b0cbbd9efcc82816c58f31a34963071aa19fb792a212a5d9caf8e0fc3097/grpclib-0.4.9-py3-none-any.whl", hash = "sha256:7762ec1c8ed94dfad597475152dd35cbd11aecaaca2f243e29702435ca24cf0e", size = 77063, upload-time = "2025-12-14T22:23:13.224Z" }, -] - -[package.optional-dependencies] -protobuf = [ - { name = "protobuf" }, -] - [[package]] name = "h11" version = "0.16.0" @@ -1708,105 +1637,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/32/28/79f0f8de97cce916d5ae88a7bee1ad724855e83e6019c0b4d5b3fabc80f3/mkdocstrings_python-2.0.3-py3-none-any.whl", hash = "sha256:0b83513478bdfd803ff05aa43e9b1fca9dd22bcd9471f09ca6257f009bc5ee12", size = 104779, upload-time = "2026-02-20T10:38:34.517Z" }, ] -[[package]] -name = "multidict" -version = "6.7.1" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/1a/c2/c2d94cbe6ac1753f3fc980da97b3d930efe1da3af3c9f5125354436c073d/multidict-6.7.1.tar.gz", hash = "sha256:ec6652a1bee61c53a3e5776b6049172c53b6aaba34f18c9ad04f82712bac623d", size = 102010, upload-time = "2026-01-26T02:46:45.979Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/8d/9c/f20e0e2cf80e4b2e4b1c365bf5fe104ee633c751a724246262db8f1a0b13/multidict-6.7.1-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:a90f75c956e32891a4eda3639ce6dd86e87105271f43d43442a3aedf3cddf172", size = 76893, upload-time = "2026-01-26T02:43:52.754Z" }, - { url = "https://files.pythonhosted.org/packages/fe/cf/18ef143a81610136d3da8193da9d80bfe1cb548a1e2d1c775f26b23d024a/multidict-6.7.1-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:3fccb473e87eaa1382689053e4a4618e7ba7b9b9b8d6adf2027ee474597128cd", size = 45456, upload-time = "2026-01-26T02:43:53.893Z" }, - { url = "https://files.pythonhosted.org/packages/a9/65/1caac9d4cd32e8433908683446eebc953e82d22b03d10d41a5f0fefe991b/multidict-6.7.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:b0fa96985700739c4c7853a43c0b3e169360d6855780021bfc6d0f1ce7c123e7", size = 43872, upload-time = "2026-01-26T02:43:55.041Z" }, - { url = "https://files.pythonhosted.org/packages/cf/3b/d6bd75dc4f3ff7c73766e04e705b00ed6dbbaccf670d9e05a12b006f5a21/multidict-6.7.1-cp312-cp312-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:cb2a55f408c3043e42b40cc8eecd575afa27b7e0b956dfb190de0f8499a57a53", size = 251018, upload-time = "2026-01-26T02:43:56.198Z" }, - { url = "https://files.pythonhosted.org/packages/fd/80/c959c5933adedb9ac15152e4067c702a808ea183a8b64cf8f31af8ad3155/multidict-6.7.1-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:eb0ce7b2a32d09892b3dd6cc44877a0d02a33241fafca5f25c8b6b62374f8b75", size = 258883, upload-time = "2026-01-26T02:43:57.499Z" }, - { url = "https://files.pythonhosted.org/packages/86/85/7ed40adafea3d4f1c8b916e3b5cc3a8e07dfcdcb9cd72800f4ed3ca1b387/multidict-6.7.1-cp312-cp312-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:c3a32d23520ee37bf327d1e1a656fec76a2edd5c038bf43eddfa0572ec49c60b", size = 242413, upload-time = "2026-01-26T02:43:58.755Z" }, - { url = "https://files.pythonhosted.org/packages/d2/57/b8565ff533e48595503c785f8361ff9a4fde4d67de25c207cd0ba3befd03/multidict-6.7.1-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:9c90fed18bffc0189ba814749fdcc102b536e83a9f738a9003e569acd540a733", size = 268404, upload-time = "2026-01-26T02:44:00.216Z" }, - { url = "https://files.pythonhosted.org/packages/e0/50/9810c5c29350f7258180dfdcb2e52783a0632862eb334c4896ac717cebcb/multidict-6.7.1-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:da62917e6076f512daccfbbde27f46fed1c98fee202f0559adec8ee0de67f71a", size = 269456, upload-time = "2026-01-26T02:44:02.202Z" }, - { url = "https://files.pythonhosted.org/packages/f3/8d/5e5be3ced1d12966fefb5c4ea3b2a5b480afcea36406559442c6e31d4a48/multidict-6.7.1-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:bfde23ef6ed9db7eaee6c37dcec08524cb43903c60b285b172b6c094711b3961", size = 256322, upload-time = "2026-01-26T02:44:03.56Z" }, - { url = "https://files.pythonhosted.org/packages/31/6e/d8a26d81ac166a5592782d208dd90dfdc0a7a218adaa52b45a672b46c122/multidict-6.7.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:3758692429e4e32f1ba0df23219cd0b4fc0a52f476726fff9337d1a57676a582", size = 253955, upload-time = "2026-01-26T02:44:04.845Z" }, - { url = "https://files.pythonhosted.org/packages/59/4c/7c672c8aad41534ba619bcd4ade7a0dc87ed6b8b5c06149b85d3dd03f0cd/multidict-6.7.1-cp312-cp312-musllinux_1_2_armv7l.whl", hash = "sha256:398c1478926eca669f2fd6a5856b6de9c0acf23a2cb59a14c0ba5844fa38077e", size = 251254, upload-time = "2026-01-26T02:44:06.133Z" }, - { url = "https://files.pythonhosted.org/packages/7b/bd/84c24de512cbafbdbc39439f74e967f19570ce7924e3007174a29c348916/multidict-6.7.1-cp312-cp312-musllinux_1_2_i686.whl", hash = "sha256:c102791b1c4f3ab36ce4101154549105a53dc828f016356b3e3bcae2e3a039d3", size = 252059, upload-time = "2026-01-26T02:44:07.518Z" }, - { url = "https://files.pythonhosted.org/packages/fa/ba/f5449385510825b73d01c2d4087bf6d2fccc20a2d42ac34df93191d3dd03/multidict-6.7.1-cp312-cp312-musllinux_1_2_ppc64le.whl", hash = "sha256:a088b62bd733e2ad12c50dad01b7d0166c30287c166e137433d3b410add807a6", size = 263588, upload-time = "2026-01-26T02:44:09.382Z" }, - { url = "https://files.pythonhosted.org/packages/d7/11/afc7c677f68f75c84a69fe37184f0f82fce13ce4b92f49f3db280b7e92b3/multidict-6.7.1-cp312-cp312-musllinux_1_2_s390x.whl", hash = "sha256:3d51ff4785d58d3f6c91bdbffcb5e1f7ddfda557727043aa20d20ec4f65e324a", size = 259642, upload-time = "2026-01-26T02:44:10.73Z" }, - { url = "https://files.pythonhosted.org/packages/2b/17/ebb9644da78c4ab36403739e0e6e0e30ebb135b9caf3440825001a0bddcb/multidict-6.7.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:fc5907494fccf3e7d3f94f95c91d6336b092b5fc83811720fae5e2765890dfba", size = 251377, upload-time = "2026-01-26T02:44:12.042Z" }, - { url = "https://files.pythonhosted.org/packages/ca/a4/840f5b97339e27846c46307f2530a2805d9d537d8b8bd416af031cad7fa0/multidict-6.7.1-cp312-cp312-win32.whl", hash = "sha256:28ca5ce2fd9716631133d0e9a9b9a745ad7f60bac2bccafb56aa380fc0b6c511", size = 41887, upload-time = "2026-01-26T02:44:14.245Z" }, - { url = "https://files.pythonhosted.org/packages/80/31/0b2517913687895f5904325c2069d6a3b78f66cc641a86a2baf75a05dcbb/multidict-6.7.1-cp312-cp312-win_amd64.whl", hash = "sha256:fcee94dfbd638784645b066074b338bc9cc155d4b4bffa4adce1615c5a426c19", size = 46053, upload-time = "2026-01-26T02:44:15.371Z" }, - { url = "https://files.pythonhosted.org/packages/0c/5b/aba28e4ee4006ae4c7df8d327d31025d760ffa992ea23812a601d226e682/multidict-6.7.1-cp312-cp312-win_arm64.whl", hash = "sha256:ba0a9fb644d0c1a2194cf7ffb043bd852cea63a57f66fbd33959f7dae18517bf", size = 43307, upload-time = "2026-01-26T02:44:16.852Z" }, - { url = "https://files.pythonhosted.org/packages/f2/22/929c141d6c0dba87d3e1d38fbdf1ba8baba86b7776469f2bc2d3227a1e67/multidict-6.7.1-cp313-cp313-macosx_10_13_universal2.whl", hash = "sha256:2b41f5fed0ed563624f1c17630cb9941cf2309d4df00e494b551b5f3e3d67a23", size = 76174, upload-time = "2026-01-26T02:44:18.509Z" }, - { url = "https://files.pythonhosted.org/packages/c7/75/bc704ae15fee974f8fccd871305e254754167dce5f9e42d88a2def741a1d/multidict-6.7.1-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:84e61e3af5463c19b67ced91f6c634effb89ef8bfc5ca0267f954451ed4bb6a2", size = 45116, upload-time = "2026-01-26T02:44:19.745Z" }, - { url = "https://files.pythonhosted.org/packages/79/76/55cd7186f498ed080a18440c9013011eb548f77ae1b297206d030eb1180a/multidict-6.7.1-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:935434b9853c7c112eee7ac891bc4cb86455aa631269ae35442cb316790c1445", size = 43524, upload-time = "2026-01-26T02:44:21.571Z" }, - { url = "https://files.pythonhosted.org/packages/e9/3c/414842ef8d5a1628d68edee29ba0e5bcf235dbfb3ccd3ea303a7fe8c72ff/multidict-6.7.1-cp313-cp313-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:432feb25a1cb67fe82a9680b4d65fb542e4635cb3166cd9c01560651ad60f177", size = 249368, upload-time = "2026-01-26T02:44:22.803Z" }, - { url = "https://files.pythonhosted.org/packages/f6/32/befed7f74c458b4a525e60519fe8d87eef72bb1e99924fa2b0f9d97a221e/multidict-6.7.1-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:e82d14e3c948952a1a85503817e038cba5905a3352de76b9a465075d072fba23", size = 256952, upload-time = "2026-01-26T02:44:24.306Z" }, - { url = "https://files.pythonhosted.org/packages/03/d6/c878a44ba877f366630c860fdf74bfb203c33778f12b6ac274936853c451/multidict-6.7.1-cp313-cp313-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:4cfb48c6ea66c83bcaaf7e4dfa7ec1b6bbcf751b7db85a328902796dfde4c060", size = 240317, upload-time = "2026-01-26T02:44:25.772Z" }, - { url = "https://files.pythonhosted.org/packages/68/49/57421b4d7ad2e9e60e25922b08ceb37e077b90444bde6ead629095327a6f/multidict-6.7.1-cp313-cp313-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:1d540e51b7e8e170174555edecddbd5538105443754539193e3e1061864d444d", size = 267132, upload-time = "2026-01-26T02:44:27.648Z" }, - { url = "https://files.pythonhosted.org/packages/b7/fe/ec0edd52ddbcea2a2e89e174f0206444a61440b40f39704e64dc807a70bd/multidict-6.7.1-cp313-cp313-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:273d23f4b40f3dce4d6c8a821c741a86dec62cded82e1175ba3d99be128147ed", size = 268140, upload-time = "2026-01-26T02:44:29.588Z" }, - { url = "https://files.pythonhosted.org/packages/b0/73/6e1b01cbeb458807aa0831742232dbdd1fa92bfa33f52a3f176b4ff3dc11/multidict-6.7.1-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:9d624335fd4fa1c08a53f8b4be7676ebde19cd092b3895c421045ca87895b429", size = 254277, upload-time = "2026-01-26T02:44:30.902Z" }, - { url = "https://files.pythonhosted.org/packages/6a/b2/5fb8c124d7561a4974c342bc8c778b471ebbeb3cc17df696f034a7e9afe7/multidict-6.7.1-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:12fad252f8b267cc75b66e8fc51b3079604e8d43a75428ffe193cd9e2195dfd6", size = 252291, upload-time = "2026-01-26T02:44:32.31Z" }, - { url = "https://files.pythonhosted.org/packages/5a/96/51d4e4e06bcce92577fcd488e22600bd38e4fd59c20cb49434d054903bd2/multidict-6.7.1-cp313-cp313-musllinux_1_2_armv7l.whl", hash = "sha256:03ede2a6ffbe8ef936b92cb4529f27f42be7f56afcdab5ab739cd5f27fb1cbf9", size = 250156, upload-time = "2026-01-26T02:44:33.734Z" }, - { url = "https://files.pythonhosted.org/packages/db/6b/420e173eec5fba721a50e2a9f89eda89d9c98fded1124f8d5c675f7a0c0f/multidict-6.7.1-cp313-cp313-musllinux_1_2_i686.whl", hash = "sha256:90efbcf47dbe33dcf643a1e400d67d59abeac5db07dc3f27d6bdeae497a2198c", size = 249742, upload-time = "2026-01-26T02:44:35.222Z" }, - { url = "https://files.pythonhosted.org/packages/44/a3/ec5b5bd98f306bc2aa297b8c6f11a46714a56b1e6ef5ebda50a4f5d7c5fb/multidict-6.7.1-cp313-cp313-musllinux_1_2_ppc64le.whl", hash = "sha256:5c4b9bfc148f5a91be9244d6264c53035c8a0dcd2f51f1c3c6e30e30ebaa1c84", size = 262221, upload-time = "2026-01-26T02:44:36.604Z" }, - { url = "https://files.pythonhosted.org/packages/cd/f7/e8c0d0da0cd1e28d10e624604e1a36bcc3353aaebdfdc3a43c72bc683a12/multidict-6.7.1-cp313-cp313-musllinux_1_2_s390x.whl", hash = "sha256:401c5a650f3add2472d1d288c26deebc540f99e2fb83e9525007a74cd2116f1d", size = 258664, upload-time = "2026-01-26T02:44:38.008Z" }, - { url = "https://files.pythonhosted.org/packages/52/da/151a44e8016dd33feed44f730bd856a66257c1ee7aed4f44b649fb7edeb3/multidict-6.7.1-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:97891f3b1b3ffbded884e2916cacf3c6fc87b66bb0dde46f7357404750559f33", size = 249490, upload-time = "2026-01-26T02:44:39.386Z" }, - { url = "https://files.pythonhosted.org/packages/87/af/a3b86bf9630b732897f6fc3f4c4714b90aa4361983ccbdcd6c0339b21b0c/multidict-6.7.1-cp313-cp313-win32.whl", hash = "sha256:e1c5988359516095535c4301af38d8a8838534158f649c05dd1050222321bcb3", size = 41695, upload-time = "2026-01-26T02:44:41.318Z" }, - { url = "https://files.pythonhosted.org/packages/b2/35/e994121b0e90e46134673422dd564623f93304614f5d11886b1b3e06f503/multidict-6.7.1-cp313-cp313-win_amd64.whl", hash = "sha256:960c83bf01a95b12b08fd54324a4eb1d5b52c88932b5cba5d6e712bb3ed12eb5", size = 45884, upload-time = "2026-01-26T02:44:42.488Z" }, - { url = "https://files.pythonhosted.org/packages/ca/61/42d3e5dbf661242a69c97ea363f2d7b46c567da8eadef8890022be6e2ab0/multidict-6.7.1-cp313-cp313-win_arm64.whl", hash = "sha256:563fe25c678aaba333d5399408f5ec3c383ca5b663e7f774dd179a520b8144df", size = 43122, upload-time = "2026-01-26T02:44:43.664Z" }, - { url = "https://files.pythonhosted.org/packages/6d/b3/e6b21c6c4f314bb956016b0b3ef2162590a529b84cb831c257519e7fde44/multidict-6.7.1-cp313-cp313t-macosx_10_13_universal2.whl", hash = "sha256:c76c4bec1538375dad9d452d246ca5368ad6e1c9039dadcf007ae59c70619ea1", size = 83175, upload-time = "2026-01-26T02:44:44.894Z" }, - { url = "https://files.pythonhosted.org/packages/fb/76/23ecd2abfe0957b234f6c960f4ade497f55f2c16aeb684d4ecdbf1c95791/multidict-6.7.1-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:57b46b24b5d5ebcc978da4ec23a819a9402b4228b8a90d9c656422b4bdd8a963", size = 48460, upload-time = "2026-01-26T02:44:46.106Z" }, - { url = "https://files.pythonhosted.org/packages/c4/57/a0ed92b23f3a042c36bc4227b72b97eca803f5f1801c1ab77c8a212d455e/multidict-6.7.1-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:e954b24433c768ce78ab7929e84ccf3422e46deb45a4dc9f93438f8217fa2d34", size = 46930, upload-time = "2026-01-26T02:44:47.278Z" }, - { url = "https://files.pythonhosted.org/packages/b5/66/02ec7ace29162e447f6382c495dc95826bf931d3818799bbef11e8f7df1a/multidict-6.7.1-cp313-cp313t-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:3bd231490fa7217cc832528e1cd8752a96f0125ddd2b5749390f7c3ec8721b65", size = 242582, upload-time = "2026-01-26T02:44:48.604Z" }, - { url = "https://files.pythonhosted.org/packages/58/18/64f5a795e7677670e872673aca234162514696274597b3708b2c0d276cce/multidict-6.7.1-cp313-cp313t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:253282d70d67885a15c8a7716f3a73edf2d635793ceda8173b9ecc21f2fb8292", size = 250031, upload-time = "2026-01-26T02:44:50.544Z" }, - { url = "https://files.pythonhosted.org/packages/c8/ed/e192291dbbe51a8290c5686f482084d31bcd9d09af24f63358c3d42fd284/multidict-6.7.1-cp313-cp313t-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:0b4c48648d7649c9335cf1927a8b87fa692de3dcb15faa676c6a6f1f1aabda43", size = 228596, upload-time = "2026-01-26T02:44:51.951Z" }, - { url = "https://files.pythonhosted.org/packages/1e/7e/3562a15a60cf747397e7f2180b0a11dc0c38d9175a650e75fa1b4d325e15/multidict-6.7.1-cp313-cp313t-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:98bc624954ec4d2c7cb074b8eefc2b5d0ce7d482e410df446414355d158fe4ca", size = 257492, upload-time = "2026-01-26T02:44:53.902Z" }, - { url = "https://files.pythonhosted.org/packages/24/02/7d0f9eae92b5249bb50ac1595b295f10e263dd0078ebb55115c31e0eaccd/multidict-6.7.1-cp313-cp313t-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:1b99af4d9eec0b49927b4402bcbb58dea89d3e0db8806a4086117019939ad3dd", size = 255899, upload-time = "2026-01-26T02:44:55.316Z" }, - { url = "https://files.pythonhosted.org/packages/00/e3/9b60ed9e23e64c73a5cde95269ef1330678e9c6e34dd4eb6b431b85b5a10/multidict-6.7.1-cp313-cp313t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:6aac4f16b472d5b7dc6f66a0d49dd57b0e0902090be16594dc9ebfd3d17c47e7", size = 247970, upload-time = "2026-01-26T02:44:56.783Z" }, - { url = "https://files.pythonhosted.org/packages/3e/06/538e58a63ed5cfb0bd4517e346b91da32fde409d839720f664e9a4ae4f9d/multidict-6.7.1-cp313-cp313t-musllinux_1_2_aarch64.whl", hash = "sha256:21f830fe223215dffd51f538e78c172ed7c7f60c9b96a2bf05c4848ad49921c3", size = 245060, upload-time = "2026-01-26T02:44:58.195Z" }, - { url = "https://files.pythonhosted.org/packages/b2/2f/d743a3045a97c895d401e9bd29aaa09b94f5cbdf1bd561609e5a6c431c70/multidict-6.7.1-cp313-cp313t-musllinux_1_2_armv7l.whl", hash = "sha256:f5dd81c45b05518b9aa4da4aa74e1c93d715efa234fd3e8a179df611cc85e5f4", size = 235888, upload-time = "2026-01-26T02:44:59.57Z" }, - { url = "https://files.pythonhosted.org/packages/38/83/5a325cac191ab28b63c52f14f1131f3b0a55ba3b9aa65a6d0bf2a9b921a0/multidict-6.7.1-cp313-cp313t-musllinux_1_2_i686.whl", hash = "sha256:eb304767bca2bb92fb9c5bd33cedc95baee5bb5f6c88e63706533a1c06ad08c8", size = 243554, upload-time = "2026-01-26T02:45:01.054Z" }, - { url = "https://files.pythonhosted.org/packages/20/1f/9d2327086bd15da2725ef6aae624208e2ef828ed99892b17f60c344e57ed/multidict-6.7.1-cp313-cp313t-musllinux_1_2_ppc64le.whl", hash = "sha256:c9035dde0f916702850ef66460bc4239d89d08df4d02023a5926e7446724212c", size = 252341, upload-time = "2026-01-26T02:45:02.484Z" }, - { url = "https://files.pythonhosted.org/packages/e8/2c/2a1aa0280cf579d0f6eed8ee5211c4f1730bd7e06c636ba2ee6aafda302e/multidict-6.7.1-cp313-cp313t-musllinux_1_2_s390x.whl", hash = "sha256:af959b9beeb66c822380f222f0e0a1889331597e81f1ded7f374f3ecb0fd6c52", size = 246391, upload-time = "2026-01-26T02:45:03.862Z" }, - { url = "https://files.pythonhosted.org/packages/e5/03/7ca022ffc36c5a3f6e03b179a5ceb829be9da5783e6fe395f347c0794680/multidict-6.7.1-cp313-cp313t-musllinux_1_2_x86_64.whl", hash = "sha256:41f2952231456154ee479651491e94118229844dd7226541788be783be2b5108", size = 243422, upload-time = "2026-01-26T02:45:05.296Z" }, - { url = "https://files.pythonhosted.org/packages/dc/1d/b31650eab6c5778aceed46ba735bd97f7c7d2f54b319fa916c0f96e7805b/multidict-6.7.1-cp313-cp313t-win32.whl", hash = "sha256:df9f19c28adcb40b6aae30bbaa1478c389efd50c28d541d76760199fc1037c32", size = 47770, upload-time = "2026-01-26T02:45:06.754Z" }, - { url = "https://files.pythonhosted.org/packages/ac/5b/2d2d1d522e51285bd61b1e20df8f47ae1a9d80839db0b24ea783b3832832/multidict-6.7.1-cp313-cp313t-win_amd64.whl", hash = "sha256:d54ecf9f301853f2c5e802da559604b3e95bb7a3b01a9c295c6ee591b9882de8", size = 53109, upload-time = "2026-01-26T02:45:08.044Z" }, - { url = "https://files.pythonhosted.org/packages/3d/a3/cc409ba012c83ca024a308516703cf339bdc4b696195644a7215a5164a24/multidict-6.7.1-cp313-cp313t-win_arm64.whl", hash = "sha256:5a37ca18e360377cfda1d62f5f382ff41f2b8c4ccb329ed974cc2e1643440118", size = 45573, upload-time = "2026-01-26T02:45:09.349Z" }, - { url = "https://files.pythonhosted.org/packages/91/cc/db74228a8be41884a567e88a62fd589a913708fcf180d029898c17a9a371/multidict-6.7.1-cp314-cp314-macosx_10_15_universal2.whl", hash = "sha256:8f333ec9c5eb1b7105e3b84b53141e66ca05a19a605368c55450b6ba208cb9ee", size = 75190, upload-time = "2026-01-26T02:45:10.651Z" }, - { url = "https://files.pythonhosted.org/packages/d5/22/492f2246bb5b534abd44804292e81eeaf835388901f0c574bac4eeec73c5/multidict-6.7.1-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:a407f13c188f804c759fc6a9f88286a565c242a76b27626594c133b82883b5c2", size = 44486, upload-time = "2026-01-26T02:45:11.938Z" }, - { url = "https://files.pythonhosted.org/packages/f1/4f/733c48f270565d78b4544f2baddc2fb2a245e5a8640254b12c36ac7ac68e/multidict-6.7.1-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:0e161ddf326db5577c3a4cc2d8648f81456e8a20d40415541587a71620d7a7d1", size = 43219, upload-time = "2026-01-26T02:45:14.346Z" }, - { url = "https://files.pythonhosted.org/packages/24/bb/2c0c2287963f4259c85e8bcbba9182ced8d7fca65c780c38e99e61629d11/multidict-6.7.1-cp314-cp314-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:1e3a8bb24342a8201d178c3b4984c26ba81a577c80d4d525727427460a50c22d", size = 245132, upload-time = "2026-01-26T02:45:15.712Z" }, - { url = "https://files.pythonhosted.org/packages/a7/f9/44d4b3064c65079d2467888794dea218d1601898ac50222ab8a9a8094460/multidict-6.7.1-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:97231140a50f5d447d3164f994b86a0bed7cd016e2682f8650d6a9158e14fd31", size = 252420, upload-time = "2026-01-26T02:45:17.293Z" }, - { url = "https://files.pythonhosted.org/packages/8b/13/78f7275e73fa17b24c9a51b0bd9d73ba64bb32d0ed51b02a746eb876abe7/multidict-6.7.1-cp314-cp314-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:6b10359683bd8806a200fd2909e7c8ca3a7b24ec1d8132e483d58e791d881048", size = 233510, upload-time = "2026-01-26T02:45:19.356Z" }, - { url = "https://files.pythonhosted.org/packages/4b/25/8167187f62ae3cbd52da7893f58cb036b47ea3fb67138787c76800158982/multidict-6.7.1-cp314-cp314-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:283ddac99f7ac25a4acadbf004cb5ae34480bbeb063520f70ce397b281859362", size = 264094, upload-time = "2026-01-26T02:45:20.834Z" }, - { url = "https://files.pythonhosted.org/packages/a1/e7/69a3a83b7b030cf283fb06ce074a05a02322359783424d7edf0f15fe5022/multidict-6.7.1-cp314-cp314-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:538cec1e18c067d0e6103aa9a74f9e832904c957adc260e61cd9d8cf0c3b3d37", size = 260786, upload-time = "2026-01-26T02:45:22.818Z" }, - { url = "https://files.pythonhosted.org/packages/fe/3b/8ec5074bcfc450fe84273713b4b0a0dd47c0249358f5d82eb8104ffe2520/multidict-6.7.1-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:7eee46ccb30ff48a1e35bb818cc90846c6be2b68240e42a78599166722cea709", size = 248483, upload-time = "2026-01-26T02:45:24.368Z" }, - { url = "https://files.pythonhosted.org/packages/48/5a/d5a99e3acbca0e29c5d9cba8f92ceb15dce78bab963b308ae692981e3a5d/multidict-6.7.1-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:fa263a02f4f2dd2d11a7b1bb4362aa7cb1049f84a9235d31adf63f30143469a0", size = 248403, upload-time = "2026-01-26T02:45:25.982Z" }, - { url = "https://files.pythonhosted.org/packages/35/48/e58cd31f6c7d5102f2a4bf89f96b9cf7e00b6c6f3d04ecc44417c00a5a3c/multidict-6.7.1-cp314-cp314-musllinux_1_2_armv7l.whl", hash = "sha256:2e1425e2f99ec5bd36c15a01b690a1a2456209c5deed58f95469ffb46039ccbb", size = 240315, upload-time = "2026-01-26T02:45:27.487Z" }, - { url = "https://files.pythonhosted.org/packages/94/33/1cd210229559cb90b6786c30676bb0c58249ff42f942765f88793b41fdce/multidict-6.7.1-cp314-cp314-musllinux_1_2_i686.whl", hash = "sha256:497394b3239fc6f0e13a78a3e1b61296e72bf1c5f94b4c4eb80b265c37a131cd", size = 245528, upload-time = "2026-01-26T02:45:28.991Z" }, - { url = "https://files.pythonhosted.org/packages/64/f2/6e1107d226278c876c783056b7db43d800bb64c6131cec9c8dfb6903698e/multidict-6.7.1-cp314-cp314-musllinux_1_2_ppc64le.whl", hash = "sha256:233b398c29d3f1b9676b4b6f75c518a06fcb2ea0b925119fb2c1bc35c05e1601", size = 258784, upload-time = "2026-01-26T02:45:30.503Z" }, - { url = "https://files.pythonhosted.org/packages/4d/c1/11f664f14d525e4a1b5327a82d4de61a1db604ab34c6603bb3c2cc63ad34/multidict-6.7.1-cp314-cp314-musllinux_1_2_s390x.whl", hash = "sha256:93b1818e4a6e0930454f0f2af7dfce69307ca03cdcfb3739bf4d91241967b6c1", size = 251980, upload-time = "2026-01-26T02:45:32.603Z" }, - { url = "https://files.pythonhosted.org/packages/e1/9f/75a9ac888121d0c5bbd4ecf4eead45668b1766f6baabfb3b7f66a410e231/multidict-6.7.1-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:f33dc2a3abe9249ea5d8360f969ec7f4142e7ac45ee7014d8f8d5acddf178b7b", size = 243602, upload-time = "2026-01-26T02:45:34.043Z" }, - { url = "https://files.pythonhosted.org/packages/9a/e7/50bf7b004cc8525d80dbbbedfdc7aed3e4c323810890be4413e589074032/multidict-6.7.1-cp314-cp314-win32.whl", hash = "sha256:3ab8b9d8b75aef9df299595d5388b14530839f6422333357af1339443cff777d", size = 40930, upload-time = "2026-01-26T02:45:36.278Z" }, - { url = "https://files.pythonhosted.org/packages/e0/bf/52f25716bbe93745595800f36fb17b73711f14da59ed0bb2eba141bc9f0f/multidict-6.7.1-cp314-cp314-win_amd64.whl", hash = "sha256:5e01429a929600e7dab7b166062d9bb54a5eed752384c7384c968c2afab8f50f", size = 45074, upload-time = "2026-01-26T02:45:37.546Z" }, - { url = "https://files.pythonhosted.org/packages/97/ab/22803b03285fa3a525f48217963da3a65ae40f6a1b6f6cf2768879e208f9/multidict-6.7.1-cp314-cp314-win_arm64.whl", hash = "sha256:4885cb0e817aef5d00a2e8451d4665c1808378dc27c2705f1bf4ef8505c0d2e5", size = 42471, upload-time = "2026-01-26T02:45:38.889Z" }, - { url = "https://files.pythonhosted.org/packages/e0/6d/f9293baa6146ba9507e360ea0292b6422b016907c393e2f63fc40ab7b7b5/multidict-6.7.1-cp314-cp314t-macosx_10_15_universal2.whl", hash = "sha256:0458c978acd8e6ea53c81eefaddbbee9c6c5e591f41b3f5e8e194780fe026581", size = 82401, upload-time = "2026-01-26T02:45:40.254Z" }, - { url = "https://files.pythonhosted.org/packages/7a/68/53b5494738d83558d87c3c71a486504d8373421c3e0dbb6d0db48ad42ee0/multidict-6.7.1-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:c0abd12629b0af3cf590982c0b413b1e7395cd4ec026f30986818ab95bfaa94a", size = 48143, upload-time = "2026-01-26T02:45:41.635Z" }, - { url = "https://files.pythonhosted.org/packages/37/e8/5284c53310dcdc99ce5d66563f6e5773531a9b9fe9ec7a615e9bc306b05f/multidict-6.7.1-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:14525a5f61d7d0c94b368a42cff4c9a4e7ba2d52e2672a7b23d84dc86fb02b0c", size = 46507, upload-time = "2026-01-26T02:45:42.99Z" }, - { url = "https://files.pythonhosted.org/packages/e4/fc/6800d0e5b3875568b4083ecf5f310dcf91d86d52573160834fb4bfcf5e4f/multidict-6.7.1-cp314-cp314t-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:17307b22c217b4cf05033dabefe68255a534d637c6c9b0cc8382718f87be4262", size = 239358, upload-time = "2026-01-26T02:45:44.376Z" }, - { url = "https://files.pythonhosted.org/packages/41/75/4ad0973179361cdf3a113905e6e088173198349131be2b390f9fa4da5fc6/multidict-6.7.1-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:7a7e590ff876a3eaf1c02a4dfe0724b6e69a9e9de6d8f556816f29c496046e59", size = 246884, upload-time = "2026-01-26T02:45:47.167Z" }, - { url = "https://files.pythonhosted.org/packages/c3/9c/095bb28b5da139bd41fb9a5d5caff412584f377914bd8787c2aa98717130/multidict-6.7.1-cp314-cp314t-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:5fa6a95dfee63893d80a34758cd0e0c118a30b8dcb46372bf75106c591b77889", size = 225878, upload-time = "2026-01-26T02:45:48.698Z" }, - { url = "https://files.pythonhosted.org/packages/07/d0/c0a72000243756e8f5a277b6b514fa005f2c73d481b7d9e47cd4568aa2e4/multidict-6.7.1-cp314-cp314t-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:a0543217a6a017692aa6ae5cc39adb75e587af0f3a82288b1492eb73dd6cc2a4", size = 253542, upload-time = "2026-01-26T02:45:50.164Z" }, - { url = "https://files.pythonhosted.org/packages/c0/6b/f69da15289e384ecf2a68837ec8b5ad8c33e973aa18b266f50fe55f24b8c/multidict-6.7.1-cp314-cp314t-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:f99fe611c312b3c1c0ace793f92464d8cd263cc3b26b5721950d977b006b6c4d", size = 252403, upload-time = "2026-01-26T02:45:51.779Z" }, - { url = "https://files.pythonhosted.org/packages/a2/76/b9669547afa5a1a25cd93eaca91c0da1c095b06b6d2d8ec25b713588d3a1/multidict-6.7.1-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:9004d8386d133b7e6135679424c91b0b854d2d164af6ea3f289f8f2761064609", size = 244889, upload-time = "2026-01-26T02:45:53.27Z" }, - { url = "https://files.pythonhosted.org/packages/7e/a9/a50d2669e506dad33cfc45b5d574a205587b7b8a5f426f2fbb2e90882588/multidict-6.7.1-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:e628ef0e6859ffd8273c69412a2465c4be4a9517d07261b33334b5ec6f3c7489", size = 241982, upload-time = "2026-01-26T02:45:54.919Z" }, - { url = "https://files.pythonhosted.org/packages/c5/bb/1609558ad8b456b4827d3c5a5b775c93b87878fd3117ed3db3423dfbce1b/multidict-6.7.1-cp314-cp314t-musllinux_1_2_armv7l.whl", hash = "sha256:841189848ba629c3552035a6a7f5bf3b02eb304e9fea7492ca220a8eda6b0e5c", size = 232415, upload-time = "2026-01-26T02:45:56.981Z" }, - { url = "https://files.pythonhosted.org/packages/d8/59/6f61039d2aa9261871e03ab9dc058a550d240f25859b05b67fd70f80d4b3/multidict-6.7.1-cp314-cp314t-musllinux_1_2_i686.whl", hash = "sha256:ce1bbd7d780bb5a0da032e095c951f7014d6b0a205f8318308140f1a6aba159e", size = 240337, upload-time = "2026-01-26T02:45:58.698Z" }, - { url = "https://files.pythonhosted.org/packages/a1/29/fdc6a43c203890dc2ae9249971ecd0c41deaedfe00d25cb6564b2edd99eb/multidict-6.7.1-cp314-cp314t-musllinux_1_2_ppc64le.whl", hash = "sha256:b26684587228afed0d50cf804cc71062cc9c1cdf55051c4c6345d372947b268c", size = 248788, upload-time = "2026-01-26T02:46:00.862Z" }, - { url = "https://files.pythonhosted.org/packages/a9/14/a153a06101323e4cf086ecee3faadba52ff71633d471f9685c42e3736163/multidict-6.7.1-cp314-cp314t-musllinux_1_2_s390x.whl", hash = "sha256:9f9af11306994335398293f9958071019e3ab95e9a707dc1383a35613f6abcb9", size = 242842, upload-time = "2026-01-26T02:46:02.824Z" }, - { url = "https://files.pythonhosted.org/packages/41/5f/604ae839e64a4a6efc80db94465348d3b328ee955e37acb24badbcd24d83/multidict-6.7.1-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:b4938326284c4f1224178a560987b6cf8b4d38458b113d9b8c1db1a836e640a2", size = 240237, upload-time = "2026-01-26T02:46:05.898Z" }, - { url = "https://files.pythonhosted.org/packages/5f/60/c3a5187bf66f6fb546ff4ab8fb5a077cbdd832d7b1908d4365c7f74a1917/multidict-6.7.1-cp314-cp314t-win32.whl", hash = "sha256:98655c737850c064a65e006a3df7c997cd3b220be4ec8fe26215760b9697d4d7", size = 48008, upload-time = "2026-01-26T02:46:07.468Z" }, - { url = "https://files.pythonhosted.org/packages/0c/f7/addf1087b860ac60e6f382240f64fb99f8bfb532bb06f7c542b83c29ca61/multidict-6.7.1-cp314-cp314t-win_amd64.whl", hash = "sha256:497bde6223c212ba11d462853cfa4f0ae6ef97465033e7dc9940cdb3ab5b48e5", size = 53542, upload-time = "2026-01-26T02:46:08.809Z" }, - { url = "https://files.pythonhosted.org/packages/4c/81/4629d0aa32302ef7b2ec65c75a728cc5ff4fa410c50096174c1632e70b3e/multidict-6.7.1-cp314-cp314t-win_arm64.whl", hash = "sha256:2bbd113e0d4af5db41d5ebfe9ccaff89de2120578164f86a5d17d5a576d1e5b2", size = 44719, upload-time = "2026-01-26T02:46:11.146Z" }, - { url = "https://files.pythonhosted.org/packages/81/08/7036c080d7117f28a4af526d794aab6a84463126db031b007717c1a6676e/multidict-6.7.1-py3-none-any.whl", hash = "sha256:55d97cc6dae627efa6a6e548885712d4864b81110ac76fa4e534c03819fa4a56", size = 12319, upload-time = "2026-01-26T02:46:44.004Z" }, -] - [[package]] name = "murmurhash" version = "1.0.15" @@ -2704,15 +2534,15 @@ wheels = [ [[package]] name = "pymdown-extensions" -version = "10.21.2" +version = "11.0.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "markdown" }, { name = "pyyaml" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/df/08/f1c908c581fd11913da4711ea7ba32c0eee40b0190000996bb863b0c9349/pymdown_extensions-10.21.2.tar.gz", hash = "sha256:c3f55a5b8a1d0edf6699e35dcbea71d978d34ff3fa79f3d807b8a5b3fa90fbdc", size = 853922, upload-time = "2026-03-29T15:01:55.233Z" } +sdist = { url = "https://files.pythonhosted.org/packages/21/a9/5f0c535ba3b08fe09270c16808e053a968868242ecbd5676d4e3a488bf28/pymdown_extensions-11.0.1.tar.gz", hash = "sha256:dd2905ae6fc5b75582fafb139a1266ffc754705efa902aa50067fa7ff4f94ec0", size = 857113, upload-time = "2026-07-02T17:59:22.955Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/f7/27/a2fc51a4a122dfd1015e921ae9d22fee3d20b0b8080d9a704578bf9deece/pymdown_extensions-10.21.2-py3-none-any.whl", hash = "sha256:5c0fd2a2bea14eb39af8ff284f1066d898ab2187d81b889b75d46d4348c01638", size = 268901, upload-time = "2026-03-29T15:01:53.244Z" }, + { url = "https://files.pythonhosted.org/packages/d6/54/da572c98c0b77626a91b5d3b89f0231d8bff5125c225420908632f8b342d/pymdown_extensions-11.0.1-py3-none-any.whl", hash = "sha256:db3943a62bab7e03af1364f0c4083e64b91fb097675a4b6cceccfbe9a77e5eb2", size = 269455, upload-time = "2026-07-02T17:59:21.271Z" }, ] [[package]] @@ -2740,11 +2570,11 @@ wheels = [ [[package]] name = "pypdf" -version = "6.10.2" +version = "6.14.2" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/7b/3f/9f2167401c2e94833ca3b69535bad89e533b5de75fefe4197a2c224baec2/pypdf-6.10.2.tar.gz", hash = "sha256:7d09ce108eff6bf67465d461b6ef352dcb8d84f7a91befc02f904455c6eea11d", size = 5315679, upload-time = "2026-04-15T16:37:36.978Z" } +sdist = { url = "https://files.pythonhosted.org/packages/03/72/7dfd5ff1c9c37de97a731701f51af091325f123d9d4270361c9c69e4431f/pypdf-6.14.2.tar.gz", hash = "sha256:7873f502fe4385e79539b21d872392dc0c4e3714327c15881cbc7fbfd1f95b25", size = 6491182, upload-time = "2026-06-23T14:18:30.859Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/0c/d6/1d5c60cc17bbdf37c1552d9c03862fc6d32c5836732a0415b2d637edc2d0/pypdf-6.10.2-py3-none-any.whl", hash = "sha256:aa53be9826655b51c96741e5d7983ca224d898ac0a77896e64636810517624aa", size = 336308, upload-time = "2026-04-15T16:37:34.851Z" }, + { url = "https://files.pythonhosted.org/packages/49/e6/136aa8993a2ae7214e0b0ef2edaa0d2e08d1d4e4982635b08a835ff31ec8/pypdf-6.14.2-py3-none-any.whl", hash = "sha256:3f07891af76dc002657e04993ab9b4de81de29f9013b9761d0b7968bff12e946", size = 349514, upload-time = "2026-06-23T14:18:28.867Z" }, ] [[package]] @@ -3507,11 +3337,11 @@ wheels = [ [[package]] name = "soupsieve" -version = "2.8.3" +version = "2.9.1" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/7b/ae/2d9c981590ed9999a0d91755b47fc74f74de286b0f5cee14c9269041e6c4/soupsieve-2.8.3.tar.gz", hash = "sha256:3267f1eeea4251fb42728b6dfb746edc9acaffc4a45b27e19450b676586e8349", size = 118627, upload-time = "2026-01-20T04:27:02.457Z" } +sdist = { url = "https://files.pythonhosted.org/packages/d9/38/e12680bbe6b4f8f3d17adcaf38d26850aa756c85cf4a80e79fc12a018fe8/soupsieve-2.9.1.tar.gz", hash = "sha256:c33e6605bbc71dd628b00c632d58ae607c22bade247e52553928f83bbb75b4ba", size = 122261, upload-time = "2026-07-21T16:57:17.452Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/46/2c/1462b1d0a634697ae9e55b3cecdcb64788e8b7d63f54d923fcd0bb140aed/soupsieve-2.8.3-py3-none-any.whl", hash = "sha256:ed64f2ba4eebeab06cc4962affce381647455978ffc1e36bb79a545b91f45a95", size = 37016, upload-time = "2026-01-20T04:27:01.012Z" }, + { url = "https://files.pythonhosted.org/packages/0f/2c/437fe806897c2d6cfdc3ee43a18da8bf8e568530a4ae9bac781541ca9896/soupsieve-2.9.1-py3-none-any.whl", hash = "sha256:4f4477399246b7a0c720a88ca2454b11cd6bb9ae4c9d170140786e916776c14c", size = 37404, upload-time = "2026-07-21T16:57:16.421Z" }, ] [[package]] @@ -3620,15 +3450,15 @@ wheels = [ [[package]] name = "starlette" -version = "1.0.1" +version = "1.3.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "anyio" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/08/a3/84e821cc54b4ab50ae6dbc6ac3800a651b65ec35f045cc73785380654057/starlette-1.0.1.tar.gz", hash = "sha256:512399c5f1de7fac99c88572212ded9ddeddef2fb32afa82d724000e88b38f4f", size = 2659596, upload-time = "2026-05-21T21:58:58.433Z" } +sdist = { url = "https://files.pythonhosted.org/packages/eb/e3/7c1dc7381d9f8ab7d854328ebfa884e62cb3f3d8549ddfd37c7814f42afa/starlette-1.3.1.tar.gz", hash = "sha256:05d0213193f2fbaae60e2ecb593b4add4262ad4e46536b54abe36f11a71724e0", size = 2703240, upload-time = "2026-06-12T09:23:11.602Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/ec/e1/b2df4bc09a1e51ff664c1e17018a4274b42e5e9352e4a478ea540512dc88/starlette-1.0.1-py3-none-any.whl", hash = "sha256:7c0e69b2ee1c848bd54669d908500117a3ee13de603a21427e5c6fc1adf98dcd", size = 72802, upload-time = "2026-05-21T21:58:56.551Z" }, + { url = "https://files.pythonhosted.org/packages/ec/bb/2799cc2ede3ed41131f8975621e7213dfc7ef4acbbaadfa440f32500c370/starlette-1.3.1-py3-none-any.whl", hash = "sha256:c7372aae11c3c3f26a42df7bd626cec2f47d03483d261d369516a615a53714c6", size = 73632, upload-time = "2026-06-12T09:23:10.017Z" }, ] [[package]] @@ -4085,6 +3915,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/e3/bd/fa9bb053192491b3867ba07d2343d9f2252e00811567d30ae8d0f78136fe/watchfiles-1.1.1-cp314-cp314t-musllinux_1_1_x86_64.whl", hash = "sha256:a916a2932da8f8ab582f242c065f5c81bed3462849ca79ee357dd9551b0e9b01", size = 622112, upload-time = "2025-10-14T15:05:50.941Z" }, ] +[[package]] +name = "wcmatch" +version = "10.2.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "bracex" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/11/15/dc61746d8c0852f6d711ad09c774b63cf7c8211aa49e30871ac3d342b7e2/wcmatch-10.2.1.tar.gz", hash = "sha256:ecac70a5c70e62ba854b78318d3a1408e8651f8f1c96e5837743b71aa6a4fb92", size = 132497, upload-time = "2026-07-02T17:21:48.484Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/82/ba/20b48eedeab5316bf9a502bb9eb7e3b1588bd61d0f565822fefa8f06e10b/wcmatch-10.2.1-py3-none-any.whl", hash = "sha256:2d775395b93f233af66690f62cb9d52b084ec159a31cc4084f4069d72f437acd", size = 39763, upload-time = "2026-07-02T17:21:47.134Z" }, +] + [[package]] name = "weasel" version = "1.0.0" diff --git a/docker/.env.example b/docker/.env.example index 3b3de9cd976..d44703af976 100644 --- a/docker/.env.example +++ b/docker/.env.example @@ -41,9 +41,6 @@ FILES_ACCESS_TIMEOUT=300 # Remove `collaboration` from COMPOSE_PROFILES to stop the dedicated websocket service. ENABLE_COLLABORATION_MODE=true -# Learn app feature toggle -ENABLE_LEARN_APP=true - # Logging and server workers LOG_LEVEL=INFO LOG_OUTPUT_FORMAT=text @@ -140,9 +137,8 @@ API_SENTRY_TRACES_SAMPLE_RATE=1.0 API_SENTRY_PROFILES_SAMPLE_RATE=1.0 WEB_SENTRY_DSN= AMPLITUDE_API_KEY= -COOKIEYES_SITE_KEY= +TURNSTILE_SITE_KEY= TEXT_GENERATION_TIMEOUT_MS=60000 -WORKFLOW_GENERATION_TIMEOUT_MS=180000 CSP_WHITELIST= ALLOW_EMBED=false ALLOW_UNSAFE_DATA_SCHEME=false @@ -157,9 +153,6 @@ ENABLE_WEBSITE_JINAREADER=true ENABLE_WEBSITE_FIRECRAWL=true ENABLE_WEBSITE_WATERCRAWL=true NEXT_PUBLIC_ENABLE_SINGLE_DOLLAR_LATEX=false -# Enable preview features still in development (currently the /create and -# /refine slash commands in the "Go to Anything" command palette). -NEXT_PUBLIC_ENABLE_FEATURE_PREVIEW=true NEXT_PUBLIC_ENABLE_AGENT_V2=true EXPERIMENTAL_ENABLE_VINEXT=false @@ -217,6 +210,13 @@ SSRF_DEFAULT_WRITE_TIME_OUT=5 SSRF_POOL_MAX_CONNECTIONS=100 SSRF_POOL_MAX_KEEPALIVE_CONNECTIONS=20 SSRF_POOL_KEEPALIVE_EXPIRY=5.0 +# Comma-separated CIDR ranges that the SSRF proxy should allow even when they +# resolve to private, loopback, link-local, or otherwise non-public addresses. +# Leave empty (the default) to keep the deny-by-default policy. Required when +# Dify needs to reach internal HTTP endpoints (e.g. an internal API on +# http://172.21.x.x) from an HTTP Request node, a tool, or a plugin. +# Example: SSRF_PROXY_ALLOW_PRIVATE_IPS=172.21.0.0/16,10.0.0.0/8 +SSRF_PROXY_ALLOW_PRIVATE_IPS= # Plugin daemon DB_PLUGIN_DATABASE=dify_plugin @@ -225,6 +225,7 @@ PLUGIN_DAEMON_PORT=5002 PLUGIN_DAEMON_KEY=lYkiYYT6owG+71oLerGzA7GXCgOT++6ovaezWAjpCjf+Sjc3ZtU+qUEi PLUGIN_DAEMON_URL=http://plugin_daemon:5002 PLUGIN_MAX_PACKAGE_SIZE=52428800 +PLUGIN_MAX_FILE_SIZE=52428800 PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600 PLUGIN_PPROF_ENABLED=false PLUGIN_DEBUGGING_HOST=0.0.0.0 @@ -253,6 +254,10 @@ MARKETPLACE_URL= # Dify Agent backend AGENT_BACKEND_BASE_URL=http://agent_backend:5050 +# Bearer token for the Agent backend /runs API. +# Replace this development default in production. +# Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))' +DIFY_AGENT_API_TOKEN=dify-agent-run-token-for-dev-only AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS=30 AGENT_BACKEND_STREAM_MAX_RECONNECTS=3 AGENT_BACKEND_RUN_TIMEOUT_SECONDS=1200 @@ -268,8 +273,20 @@ DIFY_AGENT_PLUGIN_DAEMON_API_KEY= # DIFY_AGENT_INNER_API_KEY must match API/worker INNER_API_KEY_FOR_PLUGIN, not INNER_API_KEY. DIFY_AGENT_INNER_API_URL= DIFY_AGENT_INNER_API_KEY= -DIFY_AGENT_SHELLCTL_ENTRYPOINT=http://local_sandbox:5004 -DIFY_AGENT_SHELLCTL_AUTH_TOKEN= +# Select exactly one coherent Home Snapshot + Sandbox backend. +DIFY_AGENT_RUNTIME_BACKEND=local +DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT=http://local_sandbox:5004 +DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN= +# E2B_API_KEY and E2B_API_TOKEN remain accepted as deployment-level fallbacks. +DIFY_AGENT_E2B_API_KEY= +DIFY_AGENT_E2B_TEMPLATE=difys-default-team/dify-agent-local-sandbox +# One-hour RuntimeLease limit spanning a complete Agent run. +DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS=3600 +DIFY_AGENT_E2B_SHELLCTL_AUTH_TOKEN= +DIFY_AGENT_E2B_SHELLCTL_PORT=5004 +# Sandbox-reachable Dify API base for dify-agent CLI /files/* transfers. +# Remote Sandboxes should use the public Dify ingress; local Compose uses api via agent_ssrf_proxy. +DIFY_AGENT_SANDBOX_FILES_BASE_URL=http://api:5001 DIFY_AGENT_STUB_API_BASE_URL=http://agent_backend:5050/agent-stub # This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens. # Replace this development default in production. diff --git a/docker/README.md b/docker/README.md index 0dedf718ef8..7d66e62632f 100644 --- a/docker/README.md +++ b/docker/README.md @@ -86,8 +86,8 @@ The root `.env.example` file contains the essential startup settings. Optional a - `CONSOLE_API_URL`, `CONSOLE_WEB_URL`, `SERVICE_API_URL`, `APP_API_URL`, `APP_WEB_URL`: public URLs for the API and frontend services. - `SERVER_CONSOLE_API_URL`: internal API origin used by web server-side requests and, when `INTERNAL_FILES_URL` is unset, as the default fallback for API-side internal file URL generation. Keep the default `http://api:5001` for standard Docker Compose deployments, and only change it when services must reach the API through a different internal address. - `FILES_URL`, `INTERNAL_FILES_URL`: Public and internal base URLs for file downloads and previews. - - `ENDPOINT_URL_TEMPLATE`, `NEXT_PUBLIC_SOCKET_URL`, `TRIGGER_URL`: Additional service URLs. - + - `ENDPOINT_URL_TEMPLATE`, `NEXT_PUBLIC_SOCKET_URL`, `TRIGGER_URL`: Additional service URLs. + See `.env.example` for the full list. 2. **Server Configuration**: diff --git a/docker/docker-compose-template.yaml b/docker/docker-compose-template.yaml index 37d7185e4a6..05717bd872a 100644 --- a/docker/docker-compose-template.yaml +++ b/docker/docker-compose-template.yaml @@ -220,7 +220,7 @@ services: # API service api: <<: *shared-api-worker-config - image: langgenius/dify-api:1.16.0 + image: langgenius/dify-api:1.16.1 environment: MODE: api SENTRY_DSN: ${API_SENTRY_DSN:-} @@ -232,6 +232,7 @@ services: PLUGIN_DAEMON_TIMEOUT: ${PLUGIN_DAEMON_TIMEOUT:-600.0} INNER_API_KEY_FOR_PLUGIN: ${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1} AGENT_BACKEND_BASE_URL: ${AGENT_BACKEND_BASE_URL:-http://agent_backend:5050} + AGENT_BACKEND_API_TOKEN: ${DIFY_AGENT_API_TOKEN:-dify-agent-run-token-for-dev-only} AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: ${AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS:-30} AGENT_BACKEND_STREAM_MAX_RECONNECTS: ${AGENT_BACKEND_STREAM_MAX_RECONNECTS:-3} AGENT_BACKEND_RUN_TIMEOUT_SECONDS: ${AGENT_BACKEND_RUN_TIMEOUT_SECONDS:-1200} @@ -270,7 +271,7 @@ services: # WebSocket service for workflow collaboration. api_websocket: <<: *shared-api-worker-config - image: langgenius/dify-api:1.16.0 + image: langgenius/dify-api:1.16.1 profiles: - collaboration environment: @@ -280,6 +281,8 @@ services: SERVER_WORKER_CONNECTIONS: ${API_WEBSOCKET_WORKER_CONNECTIONS:-1000} GUNICORN_TIMEOUT: ${API_WEBSOCKET_GUNICORN_TIMEOUT:-360} depends_on: + init_permissions: + condition: service_completed_successfully db_postgres: condition: service_healthy required: false @@ -288,6 +291,9 @@ services: required: false redis: condition: service_started + volumes: + # Mount the storage directory to the container, for storing user files. + - ./volumes/app/storage:/app/api/storage networks: - ssrf_proxy_network - default @@ -296,7 +302,7 @@ services: # The Celery worker for processing all queues (dataset, workflow, mail, etc.) worker: <<: *shared-worker-config - image: langgenius/dify-api:1.16.0 + image: langgenius/dify-api:1.16.1 environment: MODE: worker SENTRY_DSN: ${API_SENTRY_DSN:-} @@ -305,6 +311,7 @@ services: PLUGIN_MAX_PACKAGE_SIZE: ${PLUGIN_MAX_PACKAGE_SIZE:-52428800} INNER_API_KEY_FOR_PLUGIN: ${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1} AGENT_BACKEND_BASE_URL: ${AGENT_BACKEND_BASE_URL:-http://agent_backend:5050} + AGENT_BACKEND_API_TOKEN: ${DIFY_AGENT_API_TOKEN:-dify-agent-run-token-for-dev-only} AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: ${AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS:-30} AGENT_BACKEND_STREAM_MAX_RECONNECTS: ${AGENT_BACKEND_STREAM_MAX_RECONNECTS:-3} AGENT_BACKEND_RUN_TIMEOUT_SECONDS: ${AGENT_BACKEND_RUN_TIMEOUT_SECONDS:-1200} @@ -345,7 +352,7 @@ services: # Celery beat for scheduling periodic tasks. worker_beat: <<: *shared-worker-beat-config - image: langgenius/dify-api:1.16.0 + image: langgenius/dify-api:1.16.1 environment: MODE: beat depends_on: @@ -378,7 +385,7 @@ services: # Frontend web application. web: - image: langgenius/dify-web:1.16.0 + image: langgenius/dify-web:1.16.1 restart: always env_file: - path: ./envs/core-services/web.env @@ -391,14 +398,13 @@ services: SERVER_CONSOLE_API_URL: ${SERVER_CONSOLE_API_URL:-http://api:5001} APP_API_URL: ${APP_API_URL:-} AMPLITUDE_API_KEY: ${AMPLITUDE_API_KEY:-} - COOKIEYES_SITE_KEY: ${COOKIEYES_SITE_KEY:-} + TURNSTILE_SITE_KEY: ${TURNSTILE_SITE_KEY:-} NEXT_PUBLIC_COOKIE_DOMAIN: ${NEXT_PUBLIC_COOKIE_DOMAIN:-} NEXT_PUBLIC_SOCKET_URL: ${NEXT_PUBLIC_SOCKET_URL:-ws://localhost} SENTRY_DSN: ${WEB_SENTRY_DSN:-} NEXT_TELEMETRY_DISABLED: ${NEXT_TELEMETRY_DISABLED:-0} EXPERIMENTAL_ENABLE_VINEXT: ${EXPERIMENTAL_ENABLE_VINEXT:-false} TEXT_GENERATION_TIMEOUT_MS: ${TEXT_GENERATION_TIMEOUT_MS:-60000} - WORKFLOW_GENERATION_TIMEOUT_MS: ${WORKFLOW_GENERATION_TIMEOUT_MS:-180000} CSP_WHITELIST: ${CSP_WHITELIST:-} ALLOW_EMBED: ${ALLOW_EMBED:-false} ALLOW_UNSAFE_DATA_SCHEME: ${ALLOW_UNSAFE_DATA_SCHEME:-false} @@ -531,14 +537,25 @@ services: - ssrf_proxy_network # Local sandbox for Dify Agent shell workspaces. + # Network isolation: local_sandbox has NO direct route to `api`. Its only + # networks are `agent_sandbox_network` (so agent_backend can reach it on 5004 + # for shellctl, and it can reach agent_backend directly) and + # `local_sandbox_proxy_network` (so its egress is forced through + # agent_ssrf_proxy). + # All non-agent_backend/localhost traffic goes through the Squid forward proxy + # on port 3128, which only allows agent_backend /agent-stub/ and the Dify API + # /files/* endpoints (see ssrf_proxy/squid-agent.conf.template). local_sandbox: - image: langgenius/dify-agent-local-sandbox:1.16.0 + image: langgenius/dify-agent-local-sandbox:1.16.1 restart: always env_file: - path: ./envs/core-services/local-sandbox.env required: false environment: - - SHELLCTL_AUTH_TOKEN=${DIFY_AGENT_SHELLCTL_AUTH_TOKEN:-} + - SHELLCTL_AUTH_TOKEN=${DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN:-${DIFY_AGENT_SHELLCTL_AUTH_TOKEN:-}} + - HTTP_PROXY=http://agent_ssrf_proxy:3128 + - HTTPS_PROXY=http://agent_ssrf_proxy:3128 + - NO_PROXY=localhost,127.0.0.1 healthcheck: test: ["CMD", "curl", "-f", "http://localhost:5004/healthz"] interval: 30s @@ -546,11 +563,12 @@ services: retries: 3 start_period: 10s networks: - - default + - agent_sandbox_network + - local_sandbox_proxy_network # plugin daemon plugin_daemon: - image: langgenius/dify-plugin-daemon:0.6.3-local + image: langgenius/dify-plugin-daemon:0.6.10-local restart: always env_file: - path: ./envs/core-services/shared.env @@ -637,7 +655,7 @@ services: # Dify Agent backend service. agent_backend: - image: langgenius/dify-agent-backend:1.16.0 + image: langgenius/dify-agent-backend:1.16.1 restart: always env_file: - path: ./envs/core-services/dify-agent.env @@ -654,13 +672,22 @@ services: DIFY_AGENT_PLUGIN_DAEMON_API_KEY: ${DIFY_AGENT_PLUGIN_DAEMON_API_KEY:-${PLUGIN_DAEMON_KEY:-lYkiYYT6owG+71oLerGzA7GXCgOT++6ovaezWAjpCjf+Sjc3ZtU+qUEi}} DIFY_AGENT_INNER_API_URL: ${DIFY_AGENT_INNER_API_URL:-${PLUGIN_DIFY_INNER_API_URL:-http://api:5001}} DIFY_AGENT_INNER_API_KEY: ${DIFY_AGENT_INNER_API_KEY:-${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1}} - DIFY_AGENT_SHELLCTL_ENTRYPOINT: ${DIFY_AGENT_SHELLCTL_ENTRYPOINT:-http://local_sandbox:5004} - DIFY_AGENT_SHELLCTL_AUTH_TOKEN: ${DIFY_AGENT_SHELLCTL_AUTH_TOKEN:-} + DIFY_AGENT_RUNTIME_BACKEND: ${DIFY_AGENT_RUNTIME_BACKEND:-local} + DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT: ${DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT:-${DIFY_AGENT_SHELLCTL_ENTRYPOINT:-http://local_sandbox:5004}} + DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN: ${DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN:-${DIFY_AGENT_SHELLCTL_AUTH_TOKEN:-}} + DIFY_AGENT_E2B_API_KEY: ${DIFY_AGENT_E2B_API_KEY:-${E2B_API_KEY:-${E2B_API_TOKEN:-}}} + DIFY_AGENT_E2B_TEMPLATE: ${DIFY_AGENT_E2B_TEMPLATE:-difys-default-team/dify-agent-local-sandbox} + # One-hour RuntimeLease limit spanning a complete Agent run. + DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS: ${DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS:-3600} + DIFY_AGENT_E2B_SHELLCTL_AUTH_TOKEN: ${DIFY_AGENT_E2B_SHELLCTL_AUTH_TOKEN:-} + DIFY_AGENT_E2B_SHELLCTL_PORT: ${DIFY_AGENT_E2B_SHELLCTL_PORT:-5004} DIFY_AGENT_STUB_API_BASE_URL: ${DIFY_AGENT_STUB_API_BASE_URL:-http://agent_backend:5050/agent-stub} + DIFY_AGENT_SANDBOX_FILES_BASE_URL: ${DIFY_AGENT_SANDBOX_FILES_BASE_URL:-http://api:5001} # This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens. # Replace this development default in production. # Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))' DIFY_AGENT_SERVER_SECRET_KEY: ${DIFY_AGENT_SERVER_SECRET_KEY:-MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY} + DIFY_AGENT_API_TOKEN: ${DIFY_AGENT_API_TOKEN:-dify-agent-run-token-for-dev-only} DIFY_AGENT_SHUTDOWN_GRACE_SECONDS: ${DIFY_AGENT_SHUTDOWN_GRACE_SECONDS:-30} DIFY_AGENT_RUN_RETENTION_SECONDS: ${DIFY_AGENT_RUN_RETENTION_SECONDS:-259200} depends_on: @@ -668,10 +695,35 @@ services: condition: service_started plugin_daemon: condition: service_started - local_sandbox: - condition: service_started networks: - default + # Shared internal network with local_sandbox so agent_backend can reach it + # on port 5004 (shellctl entrypoint) while local_sandbox stays off `default`. + - agent_sandbox_network + + # Dedicated SSRF proxy for the dify-agent local_sandbox. + agent_ssrf_proxy: + image: ubuntu/squid:latest + restart: always + volumes: + - ./ssrf_proxy/squid-agent.conf.template:/etc/squid/squid.conf.template + - ./ssrf_proxy/squid-common.conf.template:/etc/squid/dify_common.conf.template + - ./ssrf_proxy/docker-agent-entrypoint.sh:/docker-entrypoint-mount.sh + entrypoint: + [ + "sh", + "-c", + "cp /docker-entrypoint-mount.sh /docker-entrypoint.sh && sed -i 's/\r$$//' /docker-entrypoint.sh && chmod +x /docker-entrypoint.sh && /docker-entrypoint.sh", + ] + environment: + HTTP_PORT: ${SSRF_HTTP_PORT:-3128} + COREDUMP_DIR: ${SSRF_COREDUMP_DIR:-/var/spool/squid} + networks: + # Needs to reach api and agent_backend as forward-proxy destinations. + - default + # Only agent_ssrf_proxy and local_sandbox share this internal network, so + # the local_sandbox can reach Squid without gaining a direct route to `api`. + - local_sandbox_proxy_network # ssrf_proxy server # for more information, please refer to @@ -681,6 +733,7 @@ services: restart: always volumes: - ./ssrf_proxy/squid.conf.template:/etc/squid/squid.conf.template + - ./ssrf_proxy/squid-common.conf.template:/etc/squid/dify_common.conf.template - ./ssrf_proxy/docker-entrypoint.sh:/docker-entrypoint-mount.sh entrypoint: [ @@ -1253,6 +1306,19 @@ networks: ssrf_proxy_network: driver: bridge internal: true + # Internal network shared only by agent_ssrf_proxy and local_sandbox. + local_sandbox_proxy_network: + driver: bridge + internal: true + # shellctl control channel (agent_backend -> local_sandbox:5004). + # sandbox can access agent backend through this network, this is + # a known limitation. + # + # The agent runtime respects HTTP(S)_PROXY, but arbitrary code execution + # is still possible through the shellctl channel. + agent_sandbox_network: + driver: bridge + internal: true milvus: driver: bridge opensearch-net: diff --git a/docker/docker-compose.e2b.yaml b/docker/docker-compose.e2b.yaml new file mode 100644 index 00000000000..2bfb5d566b5 --- /dev/null +++ b/docker/docker-compose.e2b.yaml @@ -0,0 +1,50 @@ +x-api-build: &api-build + context: .. + dockerfile: api/Dockerfile + +services: + api: + image: dify-api:e2b-local + build: *api-build + + api_websocket: + image: dify-api:e2b-local + build: *api-build + + worker: + image: dify-api:e2b-local + build: *api-build + + worker_beat: + image: dify-api:e2b-local + build: *api-build + + agent_backend: + image: dify-agent-backend:e2b-local + build: + context: .. + dockerfile: dify-agent/Dockerfile + environment: + DIFY_AGENT_RUNTIME_BACKEND: e2b + DIFY_AGENT_E2B_API_KEY: ${DIFY_AGENT_E2B_API_KEY:-${E2B_API_KEY:-${E2B_API_TOKEN:-}}} + DIFY_AGENT_E2B_TEMPLATE: ${DIFY_AGENT_E2B_TEMPLATE:-difys-default-team/dify-agent-local-sandbox} + ports: + - "${EXPOSE_AGENT_BACKEND_PORT:-15050}:5050" + + plugin_daemon: + ports: !override [] + + # Keep the normal local sandbox out of an E2B deployment without changing + # the default compose stack. + local_sandbox: + profiles: + - local-agent-sandbox + + # This overlay is intended for branch validation and deliberately starts + # from an isolated PostgreSQL data volume. + db_postgres: + volumes: !override + - dify_e2b_postgres_data:/var/lib/postgresql/data + +volumes: + dify_e2b_postgres_data: diff --git a/docker/docker-compose.middleware.yaml b/docker/docker-compose.middleware.yaml index 01f8fbde3e3..23c964a38fc 100644 --- a/docker/docker-compose.middleware.yaml +++ b/docker/docker-compose.middleware.yaml @@ -129,7 +129,7 @@ services: # plugin daemon plugin_daemon: - image: langgenius/dify-plugin-daemon:0.6.3-local + image: langgenius/dify-plugin-daemon:0.6.10-local restart: always env_file: - ./middleware.env @@ -199,6 +199,7 @@ services: restart: always volumes: - ./ssrf_proxy/squid.conf.template:/etc/squid/squid.conf.template + - ./ssrf_proxy/squid-common.conf.template:/etc/squid/dify_common.conf.template - ./ssrf_proxy/docker-entrypoint.sh:/docker-entrypoint-mount.sh entrypoint: [ diff --git a/docker/docker-compose.yaml b/docker/docker-compose.yaml index 908e979c8cf..09b33e62e94 100644 --- a/docker/docker-compose.yaml +++ b/docker/docker-compose.yaml @@ -226,7 +226,7 @@ services: # API service api: <<: *shared-api-worker-config - image: langgenius/dify-api:1.16.0 + image: langgenius/dify-api:1.16.1 environment: MODE: api SENTRY_DSN: ${API_SENTRY_DSN:-} @@ -238,6 +238,7 @@ services: PLUGIN_DAEMON_TIMEOUT: ${PLUGIN_DAEMON_TIMEOUT:-600.0} INNER_API_KEY_FOR_PLUGIN: ${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1} AGENT_BACKEND_BASE_URL: ${AGENT_BACKEND_BASE_URL:-http://agent_backend:5050} + AGENT_BACKEND_API_TOKEN: ${DIFY_AGENT_API_TOKEN:-dify-agent-run-token-for-dev-only} AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: ${AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS:-30} AGENT_BACKEND_STREAM_MAX_RECONNECTS: ${AGENT_BACKEND_STREAM_MAX_RECONNECTS:-3} AGENT_BACKEND_RUN_TIMEOUT_SECONDS: ${AGENT_BACKEND_RUN_TIMEOUT_SECONDS:-1200} @@ -276,7 +277,7 @@ services: # WebSocket service for workflow collaboration. api_websocket: <<: *shared-api-worker-config - image: langgenius/dify-api:1.16.0 + image: langgenius/dify-api:1.16.1 profiles: - collaboration environment: @@ -286,6 +287,8 @@ services: SERVER_WORKER_CONNECTIONS: ${API_WEBSOCKET_WORKER_CONNECTIONS:-1000} GUNICORN_TIMEOUT: ${API_WEBSOCKET_GUNICORN_TIMEOUT:-360} depends_on: + init_permissions: + condition: service_completed_successfully db_postgres: condition: service_healthy required: false @@ -294,6 +297,9 @@ services: required: false redis: condition: service_started + volumes: + # Mount the storage directory to the container, for storing user files. + - ./volumes/app/storage:/app/api/storage networks: - ssrf_proxy_network - default @@ -302,7 +308,7 @@ services: # The Celery worker for processing all queues (dataset, workflow, mail, etc.) worker: <<: *shared-worker-config - image: langgenius/dify-api:1.16.0 + image: langgenius/dify-api:1.16.1 environment: MODE: worker SENTRY_DSN: ${API_SENTRY_DSN:-} @@ -311,6 +317,7 @@ services: PLUGIN_MAX_PACKAGE_SIZE: ${PLUGIN_MAX_PACKAGE_SIZE:-52428800} INNER_API_KEY_FOR_PLUGIN: ${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1} AGENT_BACKEND_BASE_URL: ${AGENT_BACKEND_BASE_URL:-http://agent_backend:5050} + AGENT_BACKEND_API_TOKEN: ${DIFY_AGENT_API_TOKEN:-dify-agent-run-token-for-dev-only} AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS: ${AGENT_BACKEND_STREAM_READ_TIMEOUT_SECONDS:-30} AGENT_BACKEND_STREAM_MAX_RECONNECTS: ${AGENT_BACKEND_STREAM_MAX_RECONNECTS:-3} AGENT_BACKEND_RUN_TIMEOUT_SECONDS: ${AGENT_BACKEND_RUN_TIMEOUT_SECONDS:-1200} @@ -351,7 +358,7 @@ services: # Celery beat for scheduling periodic tasks. worker_beat: <<: *shared-worker-beat-config - image: langgenius/dify-api:1.16.0 + image: langgenius/dify-api:1.16.1 environment: MODE: beat depends_on: @@ -384,7 +391,7 @@ services: # Frontend web application. web: - image: langgenius/dify-web:1.16.0 + image: langgenius/dify-web:1.16.1 restart: always env_file: - path: ./envs/core-services/web.env @@ -397,14 +404,13 @@ services: SERVER_CONSOLE_API_URL: ${SERVER_CONSOLE_API_URL:-http://api:5001} APP_API_URL: ${APP_API_URL:-} AMPLITUDE_API_KEY: ${AMPLITUDE_API_KEY:-} - COOKIEYES_SITE_KEY: ${COOKIEYES_SITE_KEY:-} + TURNSTILE_SITE_KEY: ${TURNSTILE_SITE_KEY:-} NEXT_PUBLIC_COOKIE_DOMAIN: ${NEXT_PUBLIC_COOKIE_DOMAIN:-} NEXT_PUBLIC_SOCKET_URL: ${NEXT_PUBLIC_SOCKET_URL:-ws://localhost} SENTRY_DSN: ${WEB_SENTRY_DSN:-} NEXT_TELEMETRY_DISABLED: ${NEXT_TELEMETRY_DISABLED:-0} EXPERIMENTAL_ENABLE_VINEXT: ${EXPERIMENTAL_ENABLE_VINEXT:-false} TEXT_GENERATION_TIMEOUT_MS: ${TEXT_GENERATION_TIMEOUT_MS:-60000} - WORKFLOW_GENERATION_TIMEOUT_MS: ${WORKFLOW_GENERATION_TIMEOUT_MS:-180000} CSP_WHITELIST: ${CSP_WHITELIST:-} ALLOW_EMBED: ${ALLOW_EMBED:-false} ALLOW_UNSAFE_DATA_SCHEME: ${ALLOW_UNSAFE_DATA_SCHEME:-false} @@ -537,14 +543,25 @@ services: - ssrf_proxy_network # Local sandbox for Dify Agent shell workspaces. + # Network isolation: local_sandbox has NO direct route to `api`. Its only + # networks are `agent_sandbox_network` (so agent_backend can reach it on 5004 + # for shellctl, and it can reach agent_backend directly) and + # `local_sandbox_proxy_network` (so its egress is forced through + # agent_ssrf_proxy). + # All non-agent_backend/localhost traffic goes through the Squid forward proxy + # on port 3128, which only allows agent_backend /agent-stub/ and the Dify API + # /files/* endpoints (see ssrf_proxy/squid-agent.conf.template). local_sandbox: - image: langgenius/dify-agent-local-sandbox:1.16.0 + image: langgenius/dify-agent-local-sandbox:1.16.1 restart: always env_file: - path: ./envs/core-services/local-sandbox.env required: false environment: - - SHELLCTL_AUTH_TOKEN=${DIFY_AGENT_SHELLCTL_AUTH_TOKEN:-} + - SHELLCTL_AUTH_TOKEN=${DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN:-${DIFY_AGENT_SHELLCTL_AUTH_TOKEN:-}} + - HTTP_PROXY=http://agent_ssrf_proxy:3128 + - HTTPS_PROXY=http://agent_ssrf_proxy:3128 + - NO_PROXY=localhost,127.0.0.1 healthcheck: test: ["CMD", "curl", "-f", "http://localhost:5004/healthz"] interval: 30s @@ -552,11 +569,12 @@ services: retries: 3 start_period: 10s networks: - - default + - agent_sandbox_network + - local_sandbox_proxy_network # plugin daemon plugin_daemon: - image: langgenius/dify-plugin-daemon:0.6.3-local + image: langgenius/dify-plugin-daemon:0.6.10-local restart: always env_file: - path: ./envs/core-services/shared.env @@ -643,7 +661,7 @@ services: # Dify Agent backend service. agent_backend: - image: langgenius/dify-agent-backend:1.16.0 + image: langgenius/dify-agent-backend:1.16.1 restart: always env_file: - path: ./envs/core-services/dify-agent.env @@ -660,13 +678,22 @@ services: DIFY_AGENT_PLUGIN_DAEMON_API_KEY: ${DIFY_AGENT_PLUGIN_DAEMON_API_KEY:-${PLUGIN_DAEMON_KEY:-lYkiYYT6owG+71oLerGzA7GXCgOT++6ovaezWAjpCjf+Sjc3ZtU+qUEi}} DIFY_AGENT_INNER_API_URL: ${DIFY_AGENT_INNER_API_URL:-${PLUGIN_DIFY_INNER_API_URL:-http://api:5001}} DIFY_AGENT_INNER_API_KEY: ${DIFY_AGENT_INNER_API_KEY:-${PLUGIN_DIFY_INNER_API_KEY:-QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1}} - DIFY_AGENT_SHELLCTL_ENTRYPOINT: ${DIFY_AGENT_SHELLCTL_ENTRYPOINT:-http://local_sandbox:5004} - DIFY_AGENT_SHELLCTL_AUTH_TOKEN: ${DIFY_AGENT_SHELLCTL_AUTH_TOKEN:-} + DIFY_AGENT_RUNTIME_BACKEND: ${DIFY_AGENT_RUNTIME_BACKEND:-local} + DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT: ${DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT:-${DIFY_AGENT_SHELLCTL_ENTRYPOINT:-http://local_sandbox:5004}} + DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN: ${DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN:-${DIFY_AGENT_SHELLCTL_AUTH_TOKEN:-}} + DIFY_AGENT_E2B_API_KEY: ${DIFY_AGENT_E2B_API_KEY:-${E2B_API_KEY:-${E2B_API_TOKEN:-}}} + DIFY_AGENT_E2B_TEMPLATE: ${DIFY_AGENT_E2B_TEMPLATE:-difys-default-team/dify-agent-local-sandbox} + # One-hour RuntimeLease limit spanning a complete Agent run. + DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS: ${DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS:-3600} + DIFY_AGENT_E2B_SHELLCTL_AUTH_TOKEN: ${DIFY_AGENT_E2B_SHELLCTL_AUTH_TOKEN:-} + DIFY_AGENT_E2B_SHELLCTL_PORT: ${DIFY_AGENT_E2B_SHELLCTL_PORT:-5004} DIFY_AGENT_STUB_API_BASE_URL: ${DIFY_AGENT_STUB_API_BASE_URL:-http://agent_backend:5050/agent-stub} + DIFY_AGENT_SANDBOX_FILES_BASE_URL: ${DIFY_AGENT_SANDBOX_FILES_BASE_URL:-http://api:5001} # This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens. # Replace this development default in production. # Generate one with: python -c 'import secrets; print(secrets.token_urlsafe(32))' DIFY_AGENT_SERVER_SECRET_KEY: ${DIFY_AGENT_SERVER_SECRET_KEY:-MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY} + DIFY_AGENT_API_TOKEN: ${DIFY_AGENT_API_TOKEN:-dify-agent-run-token-for-dev-only} DIFY_AGENT_SHUTDOWN_GRACE_SECONDS: ${DIFY_AGENT_SHUTDOWN_GRACE_SECONDS:-30} DIFY_AGENT_RUN_RETENTION_SECONDS: ${DIFY_AGENT_RUN_RETENTION_SECONDS:-259200} depends_on: @@ -674,10 +701,35 @@ services: condition: service_started plugin_daemon: condition: service_started - local_sandbox: - condition: service_started networks: - default + # Shared internal network with local_sandbox so agent_backend can reach it + # on port 5004 (shellctl entrypoint) while local_sandbox stays off `default`. + - agent_sandbox_network + + # Dedicated SSRF proxy for the dify-agent local_sandbox. + agent_ssrf_proxy: + image: ubuntu/squid:latest + restart: always + volumes: + - ./ssrf_proxy/squid-agent.conf.template:/etc/squid/squid.conf.template + - ./ssrf_proxy/squid-common.conf.template:/etc/squid/dify_common.conf.template + - ./ssrf_proxy/docker-agent-entrypoint.sh:/docker-entrypoint-mount.sh + entrypoint: + [ + "sh", + "-c", + "cp /docker-entrypoint-mount.sh /docker-entrypoint.sh && sed -i 's/\r$$//' /docker-entrypoint.sh && chmod +x /docker-entrypoint.sh && /docker-entrypoint.sh", + ] + environment: + HTTP_PORT: ${SSRF_HTTP_PORT:-3128} + COREDUMP_DIR: ${SSRF_COREDUMP_DIR:-/var/spool/squid} + networks: + # Needs to reach api and agent_backend as forward-proxy destinations. + - default + # Only agent_ssrf_proxy and local_sandbox share this internal network, so + # the local_sandbox can reach Squid without gaining a direct route to `api`. + - local_sandbox_proxy_network # ssrf_proxy server # for more information, please refer to @@ -687,6 +739,7 @@ services: restart: always volumes: - ./ssrf_proxy/squid.conf.template:/etc/squid/squid.conf.template + - ./ssrf_proxy/squid-common.conf.template:/etc/squid/dify_common.conf.template - ./ssrf_proxy/docker-entrypoint.sh:/docker-entrypoint-mount.sh entrypoint: [ @@ -1259,6 +1312,19 @@ networks: ssrf_proxy_network: driver: bridge internal: true + # Internal network shared only by agent_ssrf_proxy and local_sandbox. + local_sandbox_proxy_network: + driver: bridge + internal: true + # shellctl control channel (agent_backend -> local_sandbox:5004). + # sandbox can access agent backend through this network, this is + # a known limitation. + # + # The agent runtime respects HTTP(S)_PROXY, but arbitrary code execution + # is still possible through the shellctl channel. + agent_sandbox_network: + driver: bridge + internal: true milvus: driver: bridge opensearch-net: diff --git a/docker/envs/core-services/api.env.example b/docker/envs/core-services/api.env.example index 538c554070d..82962444c52 100644 --- a/docker/envs/core-services/api.env.example +++ b/docker/envs/core-services/api.env.example @@ -16,3 +16,7 @@ KNOWLEDGE_FS_BASE_URL= KNOWLEDGE_FS_JWT_SECRET= KNOWLEDGE_FS_SSE_READ_TIMEOUT_SECONDS=300 KNOWLEDGE_FS_TIMEOUT_SECONDS=10 + +# Cloudflare Turnstile server-side verification for Dify Cloud sign-in +TURNSTILE_SECRET_KEY= +TURNSTILE_ALLOWED_HOSTNAMES= diff --git a/docker/envs/core-services/dify-agent.env.example b/docker/envs/core-services/dify-agent.env.example index f2d7155f521..46b950bc9f7 100644 --- a/docker/envs/core-services/dify-agent.env.example +++ b/docker/envs/core-services/dify-agent.env.example @@ -22,8 +22,19 @@ DIFY_AGENT_PLUGIN_DAEMON_API_KEY= DIFY_AGENT_INNER_API_URL= DIFY_AGENT_INNER_API_KEY= -DIFY_AGENT_SHELLCTL_ENTRYPOINT=http://local_sandbox:5004 -DIFY_AGENT_SHELLCTL_AUTH_TOKEN= +# Select exactly one coherent Home Snapshot + Sandbox backend. +DIFY_AGENT_RUNTIME_BACKEND=local +DIFY_AGENT_LOCAL_SANDBOX_ENDPOINT=http://local_sandbox:5004 +DIFY_AGENT_LOCAL_SANDBOX_AUTH_TOKEN= +# E2B_API_KEY and E2B_API_TOKEN remain accepted as deployment-level fallbacks. +DIFY_AGENT_E2B_API_KEY= +DIFY_AGENT_E2B_TEMPLATE=difys-default-team/dify-agent-local-sandbox +# One-hour RuntimeLease limit spanning a complete Agent run. +DIFY_AGENT_E2B_ACTIVE_TIMEOUT_SECONDS=3600 +DIFY_AGENT_E2B_SHELLCTL_AUTH_TOKEN= +DIFY_AGENT_E2B_SHELLCTL_PORT=5004 +# Sandbox-reachable Dify API base for signed /files/* transfers. +DIFY_AGENT_SANDBOX_FILES_BASE_URL=http://api:5001 DIFY_AGENT_STUB_API_BASE_URL=http://agent_backend:5050/agent-stub # This is security-sensitive: it derives the JWE encryption key for Agent Stub bearer tokens. # Replace this development default in production. diff --git a/docker/envs/core-services/shared.env.example b/docker/envs/core-services/shared.env.example index 5fe6ab974e1..ed34a894809 100644 --- a/docker/envs/core-services/shared.env.example +++ b/docker/envs/core-services/shared.env.example @@ -31,6 +31,7 @@ ENABLE_EXPLORE_BANNER=false ENABLE_LEARN_APP=true ENABLE_STEP_BY_STEP_TOUR=false RBAC_ENABLED=false +ENABLE_LICENSE_EXPIRY_NOTICE=true CELERY_BROKER_URL=redis://:difyai123456@redis:6379/1 CELERY_TASK_ANNOTATIONS=null AZURE_BLOB_ACCOUNT_URL=https://.blob.core.windows.net @@ -48,6 +49,7 @@ LINDORM_URL=http://localhost:30070 LINDORM_USERNAME=admin UPSTASH_VECTOR_URL=https://xxx-vector.upstash.io UPLOAD_FILE_SIZE_LIMIT=15 +KNOWLEDGE_UPLOAD_FILE_SIZE_LIMIT_FOR_PAID_PLAN=15 UPLOAD_FILE_BATCH_LIMIT=5 UPLOAD_FILE_EXTENSION_BLACKLIST= SINGLE_CHUNK_ATTACHMENT_LIMIT=10 @@ -60,6 +62,7 @@ MULTIMODAL_SEND_FORMAT=base64 UPLOAD_IMAGE_FILE_SIZE_LIMIT=10 UPLOAD_VIDEO_FILE_SIZE_LIMIT=100 UPLOAD_AUDIO_FILE_SIZE_LIMIT=50 +UPLOAD_SKILL_FILE_SIZE_LIMIT=50 API_SENTRY_DSN= API_SENTRY_TRACES_SAMPLE_RATE=1.0 API_SENTRY_PROFILES_SAMPLE_RATE=1.0 @@ -73,6 +76,7 @@ SSRF_PROXY_HTTPS_URL=http://ssrf_proxy:3128 PGDATA=/var/lib/postgresql/data/pgdata PLUGIN_MAX_PACKAGE_SIZE=52428800 PLUGIN_MODEL_SCHEMA_CACHE_TTL=3600 +PLUGIN_MODEL_PROVIDERS_CACHE_ENABLED=true PLUGIN_MODEL_PROVIDERS_CACHE_TTL=86400 # Comma-separated marketplace plugin IDs whose latest versions are installed for newly registered users. # Example: langgenius/openai,langgenius/gemini @@ -100,8 +104,7 @@ WORKFLOW_LOG_CLEANUP_SPECIFIC_WORKFLOW_IDS= EXPOSE_PLUGIN_DEBUGGING_HOST=localhost EXPOSE_PLUGIN_DEBUGGING_PORT=5003 DEPLOY_ENV=PRODUCTION -EDITION=SELF_HOSTED -ENTERPRISE_ENABLED=false +DEPLOYMENT_EDITION=COMMUNITY ACCESS_TOKEN_EXPIRE_MINUTES=60 REFRESH_TOKEN_EXPIRE_DAYS=30 APP_DEFAULT_ACTIVE_REQUESTS=0 @@ -200,11 +203,10 @@ WORKFLOW_MAX_EXECUTION_TIME=1200 WORKFLOW_CALL_MAX_DEPTH=5 MAX_VARIABLE_SIZE=204800 WORKFLOW_GENERATOR_NODE_BUILDER_MAX_WORKERS=6 -WORKFLOW_GENERATION_TIMEOUT_MS=180000 WORKFLOW_FILE_UPLOAD_LIMIT=10 GRAPH_ENGINE_MIN_WORKERS=3 GRAPH_ENGINE_MAX_WORKERS=10 -GRAPH_ENGINE_SCALE_UP_THRESHOLD=3 +GRAPH_ENGINE_SCALE_UP_THRESHOLD=0 GRAPH_ENGINE_SCALE_DOWN_IDLE_TIME=5.0 ALIYUN_SLS_ACCESS_KEY_ID= ALIYUN_SLS_ACCESS_KEY_SECRET= @@ -317,6 +319,9 @@ ARCHIVE_STORAGE_EXPORT_BUCKET= ARCHIVE_STORAGE_REGION=auto AZURE_BLOB_ACCOUNT_NAME=difyai AZURE_BLOB_CONTAINER_NAME=difyai-container +AZURE_KEYVAULT_VAULT_URL= +AZURE_KEYVAULT_KEY_SIZE=2048 +AZURE_KEYVAULT_ROTATION_INTERVAL_DAYS= GOOGLE_STORAGE_BUCKET_NAME=your-bucket-name GOOGLE_STORAGE_SERVICE_ACCOUNT_JSON_BASE64= ALIYUN_OSS_BUCKET_NAME=your-bucket-name @@ -415,9 +420,11 @@ TIDB_VECTOR_HOST=tidb TIDB_VECTOR_PORT=4000 TIDB_VECTOR_USER= TIDB_VECTOR_PASSWORD= +TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH=false TIDB_ON_QDRANT_CLIENT_TIMEOUT=20 TIDB_ON_QDRANT_GRPC_ENABLED=false TIDB_ON_QDRANT_GRPC_PORT=6334 +TIDB_ON_QDRANT_ESTIMATED_STORAGE_LIMITS_MB=sandbox:60,professional:6400,team:25600 TIDB_PUBLIC_KEY=dify TIDB_PRIVATE_KEY=dify RELYT_HOST=db @@ -448,6 +455,7 @@ MAX_SUBMIT_COUNT=100 # Vector Store Configuration STORAGE_TYPE=opendal +KEY_PROVIDER_TYPE=local VECTOR_STORE=weaviate VECTOR_INDEX_NAME_PREFIX=Vector_index WEAVIATE_ENDPOINT=http://weaviate:8080 diff --git a/docker/envs/core-services/web.env.example b/docker/envs/core-services/web.env.example index 5c952f4af4c..d4366bca7d7 100644 --- a/docker/envs/core-services/web.env.example +++ b/docker/envs/core-services/web.env.example @@ -12,13 +12,17 @@ MAX_TOOLS_NUM=10 MAX_PARALLEL_LIMIT=10 MAX_ITERATIONS_NUM=99 TEXT_GENERATION_TIMEOUT_MS=60000 +WORKFLOW_GENERATION_TIMEOUT_MS=180000 ALLOW_INLINE_STYLES=false +# Example: ()!*&()!*&-。.;;+=— +MARKDOWN_FORM_FIELD_NAME_EXTRA_CHARS= ALLOW_UNSAFE_DATA_SCHEME=false MAX_TREE_DEPTH=50 MARKETPLACE_API_URL=https://marketplace.dify.ai INDEXING_MAX_SEGMENTATION_TOKENS_LENGTH=4000 ALLOW_EMBED=false AMPLITUDE_API_KEY= +TURNSTILE_SITE_KEY= COOKIEYES_SITE_KEY= ENABLE_WEBSITE_JINAREADER=true ENABLE_WEBSITE_FIRECRAWL=true diff --git a/docker/envs/infrastructure/ssrf-proxy.env.example b/docker/envs/infrastructure/ssrf-proxy.env.example index 3624bd9fbb4..9dedcc35c90 100644 --- a/docker/envs/infrastructure/ssrf-proxy.env.example +++ b/docker/envs/infrastructure/ssrf-proxy.env.example @@ -6,7 +6,16 @@ SSRF_PROXY_HTTP_URL=http://ssrf_proxy:3128 SSRF_PROXY_HTTPS_URL=http://ssrf_proxy:3128 SSRF_HTTP_PORT=3128 SSRF_COREDUMP_DIR=/var/spool/squid +# Comma-separated CIDR ranges that the SSRF proxy should allow even when they +# resolve to private, loopback, link-local, or otherwise non-public addresses. +# Leave empty (the default) to keep the deny-by-default policy. Required when +# Dify needs to reach internal HTTP endpoints (e.g. an internal API on +# http://172.21.x.x) from an HTTP Request node, a tool, or a plugin. +# Example: 172.21.0.0/16,10.0.0.0/8 SSRF_PROXY_ALLOW_PRIVATE_IPS= +# Comma-separated domain suffixes that the SSRF proxy should allow even when +# they would otherwise resolve to a denied range. Leave empty to keep the +# deny-by-default policy. SSRF_PROXY_ALLOW_PRIVATE_DOMAINS= SSRF_DEFAULT_TIME_OUT=5 SSRF_DEFAULT_CONNECT_TIME_OUT=5 diff --git a/docker/nginx/conf.d/default.conf.template b/docker/nginx/conf.d/default.conf.template index 6f17d0a37bb..feee1cb7e67 100644 --- a/docker/nginx/conf.d/default.conf.template +++ b/docker/nginx/conf.d/default.conf.template @@ -3,19 +3,21 @@ server { listen ${NGINX_PORT}; server_name ${NGINX_SERVER_NAME}; + resolver 127.0.0.11 valid=30s ipv6=off; location /console/api { - proxy_pass http://api:5001; + set $api_upstream api:5001; + proxy_pass http://$api_upstream; include proxy.conf; } location /api { - proxy_pass http://api:5001; + set $api_upstream api:5001; + proxy_pass http://$api_upstream; include proxy.conf; } location /socket.io/ { - resolver 127.0.0.11 valid=30s ipv6=off; set $socket_io_upstream ${NGINX_SOCKET_IO_UPSTREAM}; proxy_pass http://$socket_io_upstream; include proxy.conf; @@ -25,43 +27,51 @@ server { } location /v1 { - proxy_pass http://api:5001; + set $api_upstream api:5001; + proxy_pass http://$api_upstream; include proxy.conf; } location /openapi { - proxy_pass http://api:5001; + set $api_upstream api:5001; + proxy_pass http://$api_upstream; include proxy.conf; } location /files { - proxy_pass http://api:5001; + set $api_upstream api:5001; + proxy_pass http://$api_upstream; include proxy.conf; } location /explore { - proxy_pass http://web:3000; + set $web_upstream web:3000; + proxy_pass http://$web_upstream; include proxy.conf; } location /e/ { - proxy_pass http://plugin_daemon:5002; + set $plugin_daemon_upstream plugin_daemon:5002; + proxy_pass http://$plugin_daemon_upstream; proxy_set_header Dify-Hook-Url $scheme://$host$request_uri; include proxy.conf; } location / { - proxy_pass http://web:3000; + set $web_upstream web:3000; + proxy_pass http://$web_upstream; include proxy.conf; } location /mcp { - proxy_pass http://api:5001; + set $api_upstream api:5001; + proxy_pass http://$api_upstream; include proxy.conf; } location /triggers { - proxy_pass http://api:5001; + set $api_upstream api:5001; + proxy_pass http://$api_upstream; include proxy.conf; } diff --git a/docker/ssrf_proxy/docker-agent-entrypoint.sh b/docker/ssrf_proxy/docker-agent-entrypoint.sh new file mode 100644 index 00000000000..29dee3c48e3 --- /dev/null +++ b/docker/ssrf_proxy/docker-agent-entrypoint.sh @@ -0,0 +1,25 @@ +#!/bin/bash + +tail -F /var/log/squid/access.log 2>/dev/null & +tail -F /var/log/squid/error.log 2>/dev/null & +tail -F /var/log/squid/store.log 2>/dev/null & +tail -F /var/log/squid/cache.log 2>/dev/null & + +expand_env() { + awk '{ + while(match($0, /\${[A-Za-z_][A-Za-z_0-9]*}/)) { + var = substr($0, RSTART+2, RLENGTH-3) + val = ENVIRON[var] + $0 = substr($0, 1, RSTART-1) val substr($0, RSTART+RLENGTH) + } + print + }' "$1" +} + +echo "[ENTRYPOINT] replacing environment variables in the templates" +expand_env /etc/squid/squid.conf.template > /etc/squid/squid.conf +expand_env /etc/squid/dify_common.conf.template > /etc/squid/dify_common.conf + +/usr/sbin/squid -Nz +echo "[ENTRYPOINT] starting squid" +/usr/sbin/squid -f /etc/squid/squid.conf -NYC 1 diff --git a/docker/ssrf_proxy/docker-entrypoint.sh b/docker/ssrf_proxy/docker-entrypoint.sh index a19f9818b24..36c3a60a9f8 100755 --- a/docker/ssrf_proxy/docker-entrypoint.sh +++ b/docker/ssrf_proxy/docker-entrypoint.sh @@ -74,16 +74,22 @@ if [ -n "${SSRF_SANDBOX_PROXY_PORT:-}" ]; then } >> "$SANDBOX_PROXY_CONF" fi -# Replace environment variables in the template and output to the squid.conf -echo "[ENTRYPOINT] replacing environment variables in the template" -awk '{ - while(match($0, /\${[A-Za-z_][A-Za-z_0-9]*}/)) { - var = substr($0, RSTART+2, RLENGTH-3) - val = ENVIRON[var] - $0 = substr($0, 1, RSTART-1) val substr($0, RSTART+RLENGTH) - } - print -}' /etc/squid/squid.conf.template > /etc/squid/squid.conf +# Replace environment variables in a template file. +expand_env() { + awk '{ + while(match($0, /\${[A-Za-z_][A-Za-z_0-9]*}/)) { + var = substr($0, RSTART+2, RLENGTH-3) + val = ENVIRON[var] + $0 = substr($0, 1, RSTART-1) val substr($0, RSTART+RLENGTH) + } + print + }' "$1" +} + +# Replace environment variables in the templates and output to squid.conf +echo "[ENTRYPOINT] replacing environment variables in the templates" +expand_env /etc/squid/squid.conf.template > /etc/squid/squid.conf +expand_env /etc/squid/dify_common.conf.template > /etc/squid/dify_common.conf /usr/sbin/squid -Nz echo "[ENTRYPOINT] starting squid" diff --git a/docker/ssrf_proxy/squid-agent.conf.template b/docker/ssrf_proxy/squid-agent.conf.template new file mode 100644 index 00000000000..3a00d73dc85 --- /dev/null +++ b/docker/ssrf_proxy/squid-agent.conf.template @@ -0,0 +1,21 @@ +# Dedicated Squid config for the dify-agent local_sandbox SSRF proxy. +# All traffic on this proxy is restricted to: +# - agent_backend /agent-stub/* endpoints +# - Dify API /files/* endpoints (signed upload/download URLs) +# External internet is allowed; all other private-network destinations are denied. + +include /etc/squid/dify_common.conf + +acl dst_agent_backend dstdomain agent_backend +acl dst_dify_api dstdomain api +acl path_files urlpath_regex -i ^/files/ +acl path_agent_stub urlpath_regex -i ^/agent-stub/ + +http_port ${HTTP_PORT} + +http_access deny !Safe_ports +http_access deny CONNECT !SSL_ports +http_access allow dst_agent_backend path_agent_stub +http_access allow dst_dify_api path_files +http_access deny to_private_networks +http_access allow all diff --git a/docker/ssrf_proxy/squid-common.conf.template b/docker/ssrf_proxy/squid-common.conf.template new file mode 100644 index 00000000000..d3620dd0b8b --- /dev/null +++ b/docker/ssrf_proxy/squid-common.conf.template @@ -0,0 +1,90 @@ +# Shared Squid configuration used by both ssrf_proxy and agent_ssrf_proxy. + +################################## ACL Definitions ################################ +acl client_localnet src 0.0.0.1-0.255.255.255 # RFC 1122 "this" network (LAN) +acl client_localnet src 10.0.0.0/8 # RFC 1918 local private network (LAN) +acl client_localnet src 100.64.0.0/10 # RFC 6598 shared address space (CGN) +acl client_localnet src 169.254.0.0/16 # RFC 3927 link-local (directly plugged) machines +acl client_localnet src 172.16.0.0/12 # RFC 1918 local private network (LAN) +acl client_localnet src 192.168.0.0/16 # RFC 1918 local private network (LAN) +acl client_localnet src fc00::/7 # RFC 4193 local private network range +acl client_localnet src fe80::/10 # RFC 4291 link-local (directly plugged) machines + +acl to_private_networks dst 0.0.0.0/8 +acl to_private_networks dst 10.0.0.0/8 +acl to_private_networks dst 100.64.0.0/10 +acl to_private_networks dst 127.0.0.0/8 +acl to_private_networks dst 169.254.0.0/16 +acl to_private_networks dst 172.16.0.0/12 +acl to_private_networks dst 192.168.0.0/16 +acl to_private_networks dst 224.0.0.0/4 +acl to_private_networks dst 240.0.0.0/4 +acl to_private_networks dst ::/128 +acl to_private_networks dst ::1/128 +acl to_private_networks dst ::ffff:0:0/96 # IPv4-mapped +acl to_private_networks dst ::/96 # deprecated IPv4-compatible +acl to_private_networks dst fc00::/7 +acl to_private_networks dst fe80::/10 + +acl SSL_ports port 443 +acl Safe_ports port 80 # http +acl Safe_ports port 21 # ftp +acl Safe_ports port 443 # https +acl Safe_ports port 70 # gopher +acl Safe_ports port 210 # wais +acl Safe_ports port 1025-65535 # unregistered ports +acl Safe_ports port 280 # http-mgmt +acl Safe_ports port 488 # gss-http +acl Safe_ports port 591 # filemaker +acl Safe_ports port 777 # multiling http +acl CONNECT method CONNECT + +################################## Common Parameters ################################ + +tcp_outgoing_address 0.0.0.0 + +################################## Proxy Server ################################ +coredump_dir ${COREDUMP_DIR} +refresh_pattern ^ftp: 1440 20% 10080 +refresh_pattern ^gopher: 1440 0% 1440 +refresh_pattern -i (/cgi-bin/|\?) 0 0% 0 +refresh_pattern \/(Packages|Sources)(|\.bz2|\.gz|\.xz)$ 0 0% 0 refresh-ims +refresh_pattern \/Release(|\.gpg)$ 0 0% 0 refresh-ims +refresh_pattern \/InRelease$ 0 0% 0 refresh-ims +refresh_pattern \/(Translation-.*)(|\.bz2|\.gz|\.xz)$ 0 0% 0 refresh-ims +refresh_pattern . 0 20% 4320 + +################################## Request Buffer ################################ +client_request_buffer_max_size 100 MB + +################################## Performance & Concurrency ############################### +max_filedescriptors 65536 +connect_timeout 30 seconds +request_timeout 2 minutes +read_timeout 2 minutes +client_lifetime 5 minutes +shutdown_lifetime 30 seconds + +server_persistent_connections on +client_persistent_connections on +persistent_request_timeout 30 seconds +pconn_timeout 1 minute + +client_db on +server_idle_pconn_timeout 2 minutes +client_idle_pconn_timeout 2 minutes + +quick_abort_min 16 KB +quick_abort_max 16 MB +quick_abort_pct 95 + +memory_cache_mode disk +cache_mem 256 MB +maximum_object_size_in_memory 512 KB + +dns_timeout 30 seconds +dns_retransmit_interval 5 seconds + +logformat dify_log %ts.%03tu %6tr %>a %Ss/%03>Hs %a %Ss/%03>Hs %/dev/null 2>&1 || true docker rm -f "$SANDBOX_CONTAINER_NAME" >/dev/null 2>&1 || true + docker rm -f "$AGENT_PROXY_CONTAINER_NAME" >/dev/null 2>&1 || true + docker rm -f "$API_CONTAINER_NAME" >/dev/null 2>&1 || true + docker rm -f "$AGENT_BACKEND_CONTAINER_NAME" >/dev/null 2>&1 || true docker network rm "$NETWORK_NAME" >/dev/null 2>&1 || true } @@ -33,6 +39,24 @@ http_code_for() { printf '%s\n' "$output" | awk '$1 ~ /^HTTP\// { code = $2 } END { print code }' } +http_code_for_post() { + local proxy_url="$1" + local target_url="$2" + local output + + output="$( + docker run \ + --rm \ + --network "$NETWORK_NAME" \ + --env "http_proxy=$proxy_url" \ + --env "https_proxy=$proxy_url" \ + "$CLIENT_IMAGE" \ + wget -S -O /dev/null -T 10 --post-data=file-bytes "$target_url" 2>&1 || true + )" + + printf '%s\n' "$output" | awk '$1 ~ /^HTTP\// { code = $2 } END { print code }' +} + direct_http_code_for() { local target_url="$1" local output @@ -74,6 +98,19 @@ assert_public_target_allowed() { fi } +assert_post_target_not_blocked() { + local proxy_url="$1" + local target_url="$2" + local status_code + + status_code="$(http_code_for_post "$proxy_url" "$target_url")" + if [[ -z "$status_code" || "$status_code" == "403" ]]; then + echo "Expected POST $target_url to pass the proxy ACL, got ${status_code:-no response}." + docker logs "$AGENT_PROXY_CONTAINER_NAME" >&2 || true + exit 1 + fi +} + assert_sandbox_bridge_allowed() { local target_url="$1" local status_code @@ -106,6 +143,7 @@ docker run \ --entrypoint sh \ --network "$NETWORK_NAME" \ --volume "$ROOT_DIR/docker/ssrf_proxy/squid.conf.template:/etc/squid/squid.conf.template:ro" \ + --volume "$ROOT_DIR/docker/ssrf_proxy/squid-common.conf.template:/etc/squid/dify_common.conf.template:ro" \ --volume "$ROOT_DIR/docker/ssrf_proxy/docker-entrypoint.sh:/docker-entrypoint-mount.sh:ro" \ --env HTTP_PORT=3128 \ --env COREDUMP_DIR=/var/spool/squid \ @@ -141,3 +179,81 @@ if [[ "$RUN_PUBLIC_CHECK" == "true" ]]; then fi assert_sandbox_bridge_allowed "http://$CONTAINER_NAME:8194/health" + +# --------------------------------------------------------------------------- +# agent_ssrf_proxy tests +# --------------------------------------------------------------------------- + +# Mock api server: serves /files/* (200) and everything else (404). +docker run \ + --detach \ + --name "$API_CONTAINER_NAME" \ + --network "$NETWORK_NAME" \ + --network-alias api \ + "$CLIENT_IMAGE" \ + sh -c "mkdir -p /www/files && echo file-ok > /www/files/test && echo denied > /www/index.html && httpd -f -p 5001 -h /www" \ + >/dev/null + +# Mock agent_backend server: serves /agent-stub/* (200) and everything else (404). +docker run \ + --detach \ + --name "$AGENT_BACKEND_CONTAINER_NAME" \ + --network "$NETWORK_NAME" \ + --network-alias agent_backend \ + "$CLIENT_IMAGE" \ + sh -c "mkdir -p /www/agent-stub && echo stub-ok > /www/agent-stub/config && echo denied > /www/index.html && httpd -f -p 5050 -h /www" \ + >/dev/null + +docker run \ + --detach \ + --name "$AGENT_PROXY_CONTAINER_NAME" \ + --entrypoint sh \ + --network "$NETWORK_NAME" \ + --volume "$ROOT_DIR/docker/ssrf_proxy/squid-agent.conf.template:/etc/squid/squid.conf.template:ro" \ + --volume "$ROOT_DIR/docker/ssrf_proxy/squid-common.conf.template:/etc/squid/dify_common.conf.template:ro" \ + --volume "$ROOT_DIR/docker/ssrf_proxy/docker-agent-entrypoint.sh:/docker-entrypoint-mount.sh:ro" \ + --env HTTP_PORT=3128 \ + --env COREDUMP_DIR=/var/spool/squid \ + "$IMAGE" \ + -c "cp /docker-entrypoint-mount.sh /docker-entrypoint.sh && sed -i 's/\r$//' /docker-entrypoint.sh && chmod +x /docker-entrypoint.sh && /docker-entrypoint.sh" \ + >/dev/null + +agent_proxy_url="http://$AGENT_PROXY_CONTAINER_NAME:3128" +for _ in {1..30}; do + agent_probe_status="$(http_code_for "$agent_proxy_url" "http://127.0.0.1:80/")" + if [[ -n "$agent_probe_status" ]]; then + break + fi + sleep 1 +done + +if [[ -z "${agent_probe_status:-}" ]]; then + echo "Agent SSRF proxy did not respond to probes." + docker logs "$AGENT_PROXY_CONTAINER_NAME" >&2 || true + exit 1 +fi + +# Private targets must be blocked. +assert_private_target_blocked "$agent_proxy_url" "http://127.0.0.1:80/" +assert_private_target_blocked "$agent_proxy_url" "http://169.254.169.254/latest/meta-data/" + +# agent_backend /agent-stub/* must be allowed. +assert_public_target_allowed "$agent_proxy_url" "http://agent_backend:5050/agent-stub/config" + +# agent_backend non-/agent-stub paths must be blocked (403 from Squid). +assert_private_target_blocked "$agent_proxy_url" "http://agent_backend:5050/index.html" + +# api /files/* must be allowed. +assert_public_target_allowed "$agent_proxy_url" "http://api:5001/files/test" +assert_public_target_allowed "$agent_proxy_url" "http://api:5001/files/test?timestamp=1&nonce=2&sign=3" +assert_post_target_not_blocked "$agent_proxy_url" "http://api:5001/files/upload/for-plugin?timestamp=1&nonce=2&sign=3" + +# api non-/files paths must be blocked. +assert_private_target_blocked "$agent_proxy_url" "http://api:5001/index.html" + +# External internet must be allowed. +if [[ "$RUN_PUBLIC_CHECK" == "true" ]]; then + assert_public_target_allowed "$agent_proxy_url" "http://example.com/" +fi + +echo "All SSRF proxy tests passed." diff --git a/docs/design/human-in-the-loop/hitl-form-file-upload-design.md b/docs/design/human-in-the-loop/hitl-form-file-upload-design.md deleted file mode 100644 index 86a81d70c95..00000000000 --- a/docs/design/human-in-the-loop/hitl-form-file-upload-design.md +++ /dev/null @@ -1,184 +0,0 @@ -# HITL Standalone Form File Upload Design - -## Context - -HITL standalone forms can be opened directly through a form link and do not require the -submitter to sign in through the Web App. After `file` and `file-list` inputs were introduced, -this standalone entry point also needed file upload support. - -This entry point has a different identity model from the existing upload paths: - -- Web App upload is backed by Web App authentication and an `EndUser` context. -- Service API upload is backed by API key authentication and the `user` parameter. -- HITL standalone form submission is link-based and anonymous from the product perspective. - The standalone submitter is not necessarily the workflow or chatflow initiator. - -The goal is therefore not to add another general-purpose upload channel. The goal is to -provide a constrained, short-lived upload capability that is scoped to one HITL form submission. - -## Goals - -- Support local file upload and remote URL upload from the HITL standalone page. -- Keep the standalone page independent from the Web App login flow. -- Avoid creating a technical HITL `EndUser`. -- Avoid changing the authentication model of existing Web App, Service API, or Console upload endpoints. -- Keep upload request parameters aligned with the equivalent Web App upload endpoints where possible. -- Invalidate upload capability once the form is submitted, expired, or timed out. -- Store uploaded files in a way that remains compatible with the existing workflow resume and file access model. - -## Decision - -HITL standalone upload uses a dedicated upload token that is bound to a form -recipient. The token authorizes file upload only while the related form is still valid. - -Files uploaded through the HITL standalone page are recorded under the workflow or -chatflow initiator, not under the anonymous standalone submitter and not under a technical HITL `EndUser`. - -This keeps workflow resume aligned with the existing execution model: one initiator owns the -workflow run context, and file restoration continues to resolve files through that initiator's -access scope. HITL-specific form/token/file relationships remain available for audit and tracing, -but they do not become the source of truth for file access control. - -## API Shape - -HITL standalone upload has three endpoint categories: - -| Purpose | HITL endpoint | Aligned Web App endpoint | -| --- | --- | --- | -| Issue upload token | `POST /api/form/human_input/{form_token}/upload-token` | No direct equivalent | -| Upload local file | `POST /api/form/human_input/files/upload` | `POST /api/files/upload` | -| Upload remote file | `POST /api/form/human_input/files/remote-upload` | `POST /api/remote-files/upload` | - -Local upload follows the Web App `POST /api/files/upload` parameter shape: - -- `multipart/form-data` -- Required `file` - -Remote upload follows the Web App `POST /api/remote-files/upload` parameter shape: - -- `application/json` -- Required `url` - -HITL upload endpoints do not accept the Service API `user` parameter. That parameter -belongs to the Service API `EndUser` mapping model and does not represent the anonymous -standalone form submitter. - -## Upload Token - -The upload token is issued through the form token: - -```http -POST /api/form/human_input/{form_token}/upload-token -``` - -Upload requests carry the token through the `Authorization` header: - -```http -Authorization: bearer hitl_upload_{random_value} -``` - -The `hitl_upload_` prefix only distinguishes this credential from other bearer token types. -Security comes from the high-entropy random value, server-side hash storage, and server-side state validation. - -The token is bound to at least: - -- The HITL form. -- The form recipient. -- The tenant. -- The app. - -The token must satisfy these rules: - -- It cannot outlive the form expiration. -- It cannot be used after the form is submitted, expired, or timed out. -- It is validated through the HITL upload path, not through the existing app token validation chain. - -## Why Authorization Header - -Putting `upload_token` in the request body would avoid additional CORS header configuration, but it has a bad failure mode for file upload. The server often needs to parse the multipart body before it can read a body token, so invalid requests can still consume upload parsing, temporary file, memory, or disk resources. - -Using `Authorization: bearer hitl_upload_{random_value}` keeps authentication before expensive business processing: - -- Invalid local upload requests can be rejected before reading the multipart body. -- Invalid remote upload requests can be rejected before any outbound network access. -- Bearer credential semantics are explicit and do not mix authentication with business fields. -- The token is not exposed through query strings, access logs, referrers, or browser history. - -The tradeoff is that cross-origin deployments must allow the `Authorization` header and accept browser preflight requests. This is a reasonable configuration cost for an earlier and clearer authentication boundary. - -## File Ownership - -The standalone submitter is not a reliable product identity. Assigning files to a technical `EndUser` would also conflict with workflow resume: existing file restoration expects files to be readable through the workflow or chatflow initiator's scope. - -The selected model is: - -- If the original run was started by an `Account`, standalone HITL uploads are stored under that `Account`. -- If the original run was started by an `EndUser`, standalone HITL uploads are stored under that `EndUser`. - -This means `UploadFile.created_by_role` and `UploadFile.created_by` continue to be the source of truth for file access control. HITL association records provide auditability but do not grant file access by themselves. - -## Persistence And Audit - -The HITL upload model has two responsibilities: - -- Upload tokens authorize a form recipient to upload files while the form remains valid. -- Upload-file association records trace which files were uploaded through which HITL upload token. - -These records are intentionally not tied to an `EndUser`. Their purpose is to preserve the HITL form/token/file relationship for audit and cleanup, not to define a separate file owner identity. - -## Local Upload Boundary - -Local upload should reuse the existing file upload semantics as much as possible: - -- Request parameters stay aligned with Web App local upload. -- Existing file size checks, extension restrictions, and document-type handling remain applicable. -- Response shape stays aligned with the existing file upload response. -- Token validation happens before reading the upload body. - -## Remote Upload Boundary - -Remote upload should reuse the existing remote upload semantics as much as possible: - -- Request parameters stay aligned with Web App remote upload. -- Token validation happens before outbound network access. -- Remote fetching continues to go through the existing SSRF-safe path. -- Remote filename, extension, MIME type, and file size inference stay aligned with existing behavior. -- Response shape stays aligned with the existing remote upload response. - -## Alternatives Considered - -### Unauthenticated Standalone Upload Endpoint - -This is the simplest implementation option and does not require Web App login, `EndUser`, or a new token model. It was not selected because it exposes a public file upload surface that can be abused in SaaS or internet-facing deployments. Adding authentication later would also change the endpoint contract after clients have integrated with it. - -### Reuse Web App Login And Upload - -This would maximize reuse of existing Web App upload behavior, but it would bind HITL standalone forms to Web App login, app code, Web App enablement, and enterprise SSO semantics. That coupling is undesirable because HITL forms can be reached from independent channels such as email. It also makes product behavior unclear when the Web App is disabled or the app code is reset. - -### Create `EndUser` From Form Token - -This would satisfy existing upload paths that require an `EndUser` context and would allow form state to limit upload capability. It was not selected because the created identity would be technical rather than a real submitter identity. It would also mix HITL standalone form behavior into the broader `EndUser` model already used by Web App, Service API, triggers, MCP, and other entry points. - -More importantly, files owned by this technical `EndUser` would not naturally be readable through the workflow initiator scope during workflow resume. - -### Technical `EndUser` With File Access Exception - -This would keep the technical `EndUser` as the file owner and add an access-control exception so the workflow initiator can read files uploaded through the same HITL form. It solves the immediate resume problem, but it pushes a HITL-specific rule into the general file access layer. Over time, that makes permission reasoning harder and increases the chance of accidental access expansion. - -Assigning files directly to the workflow or chatflow initiator avoids that bypass and keeps file access governed by the existing owner model. - -## Design Constraints - -- HITL standalone upload does not reuse the Web App login flow. -- HITL standalone upload does not create a technical HITL `EndUser`. -- Upload token validity is controlled by form state. -- File access control continues to use the `UploadFile` owner as the source of truth. -- HITL association records provide audit and traceability only. -- Workflow resume continues to restore submitted file values through the existing file restoration path. - -## Future Considerations - -- If endpoint parameters change, compare them with the corresponding Web App upload endpoint to avoid unnecessary drift. -- If file ownership changes, verify the workflow resume path and the file access model together. -- If HITL forms support multiple submissions or reopening, token invalidation semantics must be redefined. -- If remote upload policy expands, prefer extending the existing remote upload and SSRF-safe behavior instead of creating a HITL-only network path. diff --git a/docs/eu-ai-act-compliance.md b/docs/eu-ai-act-compliance.md deleted file mode 100644 index 5fa29eed3f3..00000000000 --- a/docs/eu-ai-act-compliance.md +++ /dev/null @@ -1,186 +0,0 @@ -# EU AI Act Compliance Guide for Dify Deployers - -Dify is an LLMOps platform for building RAG pipelines, agents, and AI workflows. If you deploy Dify in the EU — whether self-hosted or using a cloud provider — the EU AI Act applies to your deployment. This guide covers what the regulation requires and how Dify's architecture maps to those requirements. - -## Is your system in scope? - -The detailed obligations in Articles 12, 13, and 14 only apply to **high-risk AI systems** as defined in Annex III of the EU AI Act. A Dify application is high-risk if it is used for: - -- **Recruitment and HR** — screening candidates, evaluating employee performance, allocating tasks -- **Credit scoring and insurance** — assessing creditworthiness or setting premiums -- **Law enforcement** — profiling, criminal risk assessment, border control -- **Critical infrastructure** — managing energy, water, transport, or telecommunications systems -- **Education assessment** — grading students, determining admissions -- **Essential public services** — evaluating eligibility for benefits, housing, or emergency services - -Most Dify deployments (customer-facing chatbots, internal knowledge bases, content generation workflows) are **not** high-risk. If your Dify application does not fall into one of the categories above: - -- **Article 50** (end-user transparency) still applies if users interact with your application directly. See the [Article 50 section](#article-50-end-user-transparency) below. -- **GDPR** still applies if you process personal data. See the [GDPR section](#gdpr-considerations) below. -- The high-risk obligations (Articles 9-15) are less likely to apply, but risk classification is context-dependent. **Do not self-classify without legal review.** Focus on Article 50 (transparency) and GDPR (data protection) as your baseline obligations. - -If you are unsure whether your use case qualifies as high-risk, consult a qualified legal professional before proceeding. - -## Self-hosted vs cloud: different compliance profiles - -| Deployment | Your role | Dify's role | Who handles compliance? | -|-----------|----------|-------------|------------------------| -| **Self-hosted** | Provider and deployer | Framework provider — obligations under Article 25 apply only if Dify is placed on the market or put into service as part of a complete AI system bearing its name or trademark | You | -| **Dify Cloud** | Deployer | Provider and processor | Shared — Dify handles SOC 2 and GDPR for the platform; you handle AI Act obligations for your specific use case | - -Dify Cloud already has SOC 2 Type II and GDPR compliance for the platform itself. But the EU AI Act adds obligations specific to AI systems that SOC 2 does not cover: risk classification, technical documentation, transparency, and human oversight. - -## Supported providers and services - -Dify integrates with a broad range of AI providers and data stores. The following are the key ones relevant to compliance: - -- **AI providers:** HuggingFace (core), plus integrations with OpenAI, Anthropic, Google, and 100+ models via provider plugins -- **Model identifiers include:** gpt-4o, gpt-3.5-turbo, claude-3-opus, gemini-2.5-flash, whisper-1, and others -- **Vector database connections:** Extensive RAG infrastructure supporting numerous vector stores - -Dify's plugin architecture means actual provider usage depends on your configuration. Document which providers and models are active in your deployment. - -## Data flow diagram - -A typical Dify RAG deployment: - -```mermaid -graph LR - USER((User)) -->|query| DIFY[Dify Platform] - DIFY -->|prompts| LLM([LLM Provider]) - LLM -->|responses| DIFY - DIFY -->|documents| EMBED([Embedding Model]) - EMBED -->|vectors| DIFY - DIFY -->|store/retrieve| VS[(Vector Store)] - DIFY -->|knowledge| KB[(Knowledge Base)] - DIFY -->|response| USER - - classDef processor fill:#60a5fa,stroke:#1e40af,color:#000 - classDef controller fill:#4ade80,stroke:#166534,color:#000 - classDef app fill:#a78bfa,stroke:#5b21b6,color:#000 - classDef user fill:#f472b6,stroke:#be185d,color:#000 - - class USER user - class DIFY app - class LLM processor - class EMBED processor - class VS controller - class KB controller -``` - -**GDPR roles** (providers are typically processors for customer-submitted data, but the exact role depends on each provider's terms of service and processing purpose; deployers should review each provider's DPA): -- **Cloud LLM providers (OpenAI, Anthropic, Google)** typically act as processors — requires DPA. -- **Cloud embedding services** typically act as processors — requires DPA. -- **Self-hosted vector stores (Weaviate, Qdrant, pgvector):** Your organization remains the controller — no third-party transfer. -- **Cloud vector stores (Pinecone, Zilliz Cloud)** typically act as processors — requires DPA. -- **Knowledge base documents:** Your organization is the controller — stored in your infrastructure. - -## Article 11: Technical documentation - -High-risk systems need Annex IV documentation. For Dify deployments, key sections include: - -| Section | What Dify provides | What you must document | -|---------|-------------------|----------------------| -| General description | Platform capabilities, supported models | Your specific use case, intended users, deployment context | -| Development process | Dify's architecture, plugin system | Your RAG pipeline design, prompt engineering, knowledge base curation | -| Monitoring | Dify's built-in logging and analytics | Your monitoring plan, alert thresholds, incident response | -| Performance metrics | Dify's evaluation features | Your accuracy benchmarks, quality thresholds, bias testing | -| Risk management | — | Risk assessment for your specific use case | - -Some sections can be derived from Dify's architecture and your deployment configuration, as shown in the table above. The remaining sections require your input. - -## Article 12: Record-keeping - -Dify's built-in logging covers several Article 12 requirements: - -| Requirement | Dify Feature | Status | -|------------|-------------|--------| -| Conversation logs | Full conversation history with timestamps | **Covered** | -| Model tracking | Model name recorded per interaction | **Covered** | -| Token usage | Token counts per message | **Covered** | -| Cost tracking | Cost per conversation (if provider reports it) | **Partial** | -| Document retrieval | RAG source documents logged | **Covered** | -| User identification | User session tracking | **Covered** | -| Error logging | Failed generation logs | **Covered** | -| Data retention | Configurable | **Your responsibility** | - -**Retention periods:** The required retention period depends on your role under the Act. Article 18 requires **providers** of high-risk systems to retain logs and technical documentation for **10 years** after market placement. Article 26(6) requires **deployers** to retain logs for at least **6 months**. If you self-host Dify and have substantially modified the system, you may be classified as a provider rather than a deployer. Confirm the applicable retention period with legal counsel. - -## Article 13: Transparency to deployers - -Article 13 requires providers of high-risk AI systems to supply deployers with the information needed to understand and operate the system correctly. This is a **documentation obligation**, not a logging obligation. For Dify deployments, this means the upstream LLM and embedding providers must give you: - -- Instructions for use, including intended purpose and known limitations -- Accuracy metrics and performance benchmarks -- Known or foreseeable risks and residual risks after mitigation -- Technical specifications: input/output formats, training data characteristics, model architecture details - -As a deployer, collect model cards, system documentation, and accuracy reports from each AI provider your Dify application uses. Maintain these as part of your Annex IV technical documentation. - -Dify's platform features provide **supporting evidence** that can inform Article 13 documentation, but they do not satisfy Article 13 on their own: -- **Source attribution** — Dify's RAG citation feature shows which documents informed the response, supporting deployer-side auditing -- **Model identification** — Dify logs which LLM model generates responses, providing evidence for system documentation -- **Conversation logs** — execution history helps compile performance and behavior evidence - -You must independently produce system documentation covering how your specific Dify deployment uses AI, its intended purpose, performance characteristics, and residual risks. - -## Article 50: End-user transparency - -Article 50 requires deployers to inform end users that they are interacting with an AI system. This is a separate obligation from Article 13 and applies even to limited-risk systems. - -For Dify applications serving end users: - -1. **Disclose AI involvement** — tell users they are interacting with an AI system -2. **AI-generated content labeling** — identify AI-generated content as such (e.g., clear labeling in the UI) - -Dify's "citation" feature also supports end-user transparency by showing users which knowledge base documents informed the answer. - -> **Note:** Article 50 applies to chatbots and systems interacting directly with natural persons. It has a separate scope from the high-risk designation under Annex III — it applies even to limited-risk systems. - -## Article 14: Human oversight - -Article 14 requires that high-risk AI systems be designed so that natural persons can effectively oversee them. Dify provides **automated technical safeguards** that support human oversight, but they are not a substitute for it: - -| Dify Feature | What It Does | Oversight Role | -|-------------|-------------|----------------| -| Annotation/feedback system | Human review of AI outputs | **Direct oversight** — humans evaluate and correct AI responses | -| Content moderation | Built-in filtering before responses reach users | **Automated safeguard** — reduces harmful outputs but does not replace human judgment on edge cases | -| Rate limiting | Controls on API usage | **Automated safeguard** — bounds system behavior, supports overseer's ability to maintain control | -| Workflow control | Insert human review steps between AI generation and output | **Oversight enabler** — allows building approval gates into the pipeline | - -These automated controls are necessary building blocks, but Article 14 compliance requires **human oversight procedures** on top of them: -- **Escalation procedures** — define what happens when moderation triggers or edge cases arise (who is notified, what action is taken) -- **Human review pipeline** — for high-stakes decisions, route AI outputs to a qualified person before they take effect -- **Override mechanism** — a human must be able to halt AI responses or override the system's output -- **Competence requirements** — the human overseer must understand the system's capabilities, limitations, and the context of its outputs - -### Recommended pattern - -For high-risk use cases (HR, legal, medical), configure your Dify workflow to require human approval before the AI response is delivered to the end user or acted upon. - -## Knowledge base compliance - -Dify's knowledge base feature has specific compliance implications: - -1. **Data provenance:** Document where your knowledge base documents come from. Article 10 requires data governance for training data; knowledge bases are analogous. -2. **Update tracking:** When you add, remove, or update documents in the knowledge base, log the change. The AI system's behavior changes with its knowledge base. -3. **PII in documents:** If knowledge base documents contain personal data, GDPR applies to the entire RAG pipeline. Implement access controls and consider PII redaction before indexing. -4. **Copyright:** Ensure you have the right to use the documents in your knowledge base for AI-assisted generation. - -## GDPR considerations - -1. **Legal basis** (Article 6): Document why AI processing of user queries is necessary -2. **Data Processing Agreements** (Article 28): Required for each cloud LLM and embedding provider -3. **Data minimization:** Only include necessary context in prompts; avoid sending entire documents when a relevant excerpt suffices -4. **Right to erasure:** If a user requests deletion, ensure their conversations are removed from Dify's logs AND any vector store entries derived from their data -5. **Cross-border transfers:** Providers based outside the EEA — including US-based providers (OpenAI, Anthropic), and any other non-EEA providers you route to — require Standard Contractual Clauses (SCCs) or equivalent safeguards under Chapter V of the GDPR. Review each provider's transfer mechanism individually. - -## Resources - -- [EU AI Act full text](https://artificialintelligenceact.eu/) -- [Dify documentation](https://docs.dify.ai/) -- [Dify SOC 2 compliance](https://dify.ai/trust) - ---- - -*This is not legal advice. Consult a qualified professional for compliance decisions.* diff --git a/docs/hi-IN/README.md b/docs/hi-IN/README.md index e20cb9c9845..118f413b007 100644 --- a/docs/hi-IN/README.md +++ b/docs/hi-IN/README.md @@ -75,7 +75,7 @@ Dify एक मुक्त-स्रोत प्लेटफ़ॉर्म
-Dify सर्वर शुरू करने का सबसे आसान तरीका [Docker Compose](../..docker/docker-compose.yaml) के माध्यम से है। नीचे दिए गए कमांड्स से Dify चलाने से पहले, सुनिश्चित करें कि आपकी मशीन पर [Docker] (https://docs.docker.com/get-docker/) और [Docker Compose] (https://docs.docker.com/compose/install/) इंस्टॉल हैं।: +Dify सर्वर शुरू करने का सबसे आसान तरीका [Docker Compose](../../docker/docker-compose.yaml) के माध्यम से है। नीचे दिए गए कमांड्स से Dify चलाने से पहले, सुनिश्चित करें कि आपकी मशीन पर [Docker] (https://docs.docker.com/get-docker/) और [Docker Compose] (https://docs.docker.com/compose/install/) इंस्टॉल हैं।: ```bash cd dify diff --git a/docs/tr-TR/README.md b/docs/tr-TR/README.md index 6139fd16ca2..4a245cc7e27 100644 --- a/docs/tr-TR/README.md +++ b/docs/tr-TR/README.md @@ -132,7 +132,7 @@ Yapılandırmayı özelleştirmeniz gerekiyorsa, lütfen [.env.example](../../do Uygulamalar, kiracılar, mesajlar ve daha fazlasının granularitesinde metrikleri izlemek için Dify'nin PostgreSQL veritabanını veri kaynağı olarak kullanarak panoyu Grafana'ya aktarın. -- [@bowenliang123 tarafından Grafana Panosu](%E9%93%BE%E6%8E%A5) +- [@bowenliang123 tarafından Grafana Panosu](https://github.com/bowenliang123/dify-grafana-dashboard) ### Kubernetes ile Dağıtım diff --git a/docs/weaviate/WEAVIATE_MIGRATION_GUIDE/README.md b/docs/weaviate/WEAVIATE_MIGRATION_GUIDE/README.md deleted file mode 100644 index b2599e8c2e8..00000000000 --- a/docs/weaviate/WEAVIATE_MIGRATION_GUIDE/README.md +++ /dev/null @@ -1,187 +0,0 @@ -# Weaviate Migration Guide: v1.19 → v1.27 - -## Overview - -Dify has upgraded from Weaviate v1.19 to v1.27 with the Python client updated from v3.24 to v4.17. - -## What Changed - -### Breaking Changes - -1. **Weaviate Server**: `1.19.0` → `1.27.0` -1. **Python Client**: `weaviate-client~=3.24.0` → `weaviate-client==4.17.0` -1. **gRPC Required**: Weaviate v1.27 requires gRPC port `50051` (in addition to HTTP port `8080`) -1. **Docker Compose**: Added temporary entrypoint overrides for client installation - -### Key Improvements - -- Faster vector operations via gRPC -- Improved batch processing -- Better error handling - -## Migration Steps - -### For Docker Users - -#### Step 1: Backup Your Data - -```bash -cd docker -docker compose down -sudo cp -r ./volumes/weaviate ./volumes/weaviate_backup_$(date +%Y%m%d) -``` - -#### Step 2: Update Dify - -```bash -git pull origin main -docker compose pull -``` - -#### Step 3: Start Services - -```bash -docker compose up -d -sleep 30 -curl http://localhost:8080/v1/meta -``` - -#### Step 4: Verify Migration - -```bash -# Check both ports are accessible -curl http://localhost:8080/v1/meta -netstat -tulpn | grep 50051 - -# Test in Dify UI: -# 1. Go to Knowledge Base -# 2. Test search functionality -# 3. Upload a test document -``` - -### For Source Installation - -#### Step 1: Update Dependencies - -```bash -cd api -uv sync --dev -uv run python -c "import weaviate; print(weaviate.__version__)" -# Should show: 4.17.0 -``` - -#### Step 2: Update Weaviate Server - -```bash -cd docker -docker compose -f docker-compose.middleware.yaml --profile weaviate up -d weaviate -curl http://localhost:8080/v1/meta -netstat -tulpn | grep 50051 -``` - -## Troubleshooting - -### Error: "No module named 'weaviate.classes'" - -**Solution**: - -```bash -cd api -uv sync --reinstall-package weaviate-client -uv run python -c "import weaviate; print(weaviate.__version__)" -# Should show: 4.17.0 -``` - -### Error: "gRPC health check failed" - -**Solution**: - -```bash -# Check Weaviate ports -docker ps | grep weaviate -# Should show: 0.0.0.0:8080->8080/tcp, 0.0.0.0:50051->50051/tcp - -# If missing gRPC port, add to docker-compose: -# ports: -# - "8080:8080" -# - "50051:50051" -``` - -### Error: "Weaviate version 1.19.0 is not supported" - -**Solution**: - -```bash -# Update Weaviate image in docker-compose -# Change: semitechnologies/weaviate:1.19.0 -# To: semitechnologies/weaviate:1.27.0 -docker compose down -docker compose up -d -``` - -### Data Migration Failed - -**Solution**: - -```bash -cd docker -docker compose down -sudo rm -rf ./volumes/weaviate -sudo cp -r ./volumes/weaviate_backup_YYYYMMDD ./volumes/weaviate -docker compose up -d -``` - -## Rollback Instructions - -```bash -# 1. Stop services -docker compose down - -# 2. Restore data backup -sudo rm -rf ./volumes/weaviate -sudo cp -r ./volumes/weaviate_backup_YYYYMMDD ./volumes/weaviate - -# 3. Checkout previous version -git checkout - -# 4. Restart services -docker compose up -d -``` - -## Compatibility - -| Component | Old Version | New Version | Compatible | -|-----------|-------------|-------------|------------| -| Weaviate Server | 1.19.0 | 1.27.0 | ✅ Yes | -| weaviate-client | ~3.24.0 | ==4.17.0 | ✅ Yes | -| Existing Data | v1.19 format | v1.27 format | ✅ Yes | - -## Testing Checklist - -Before deploying to production: - -- [ ] Backup all Weaviate data -- [ ] Test in staging environment -- [ ] Verify existing collections are accessible -- [ ] Test vector search functionality -- [ ] Test document upload and retrieval -- [ ] Monitor gRPC connection stability -- [ ] Check performance metrics - -## Support - -If you encounter issues: - -1. Check GitHub Issues: https://github.com/langgenius/dify/issues -1. Create a bug report with: - - Error messages - - Docker logs: `docker compose logs weaviate` - - Dify version - - Migration steps attempted - -## Important Notes - -- **Data Safety**: Existing vector data remains fully compatible -- **No Re-indexing**: No need to rebuild vector indexes -- **Temporary Workaround**: The entrypoint overrides are temporary until next Dify release -- **Performance**: May see improved performance due to gRPC usage diff --git a/e2e/.env.example b/e2e/.env.example index a2b63e5c119..1c1e430848b 100644 --- a/e2e/.env.example +++ b/e2e/.env.example @@ -22,6 +22,8 @@ E2E_MODEL_PROVIDER_CREDENTIALS_JSON='{"openai_api_key":"replace-with-real-key"}' # from @external-model because other external-model scenarios may not use dify-agent. # When enabled, the E2E runner also starts the shellctl local sandbox required by # dify-agent's dify.shell/config runtime layer. +# Use E2E_AGENT_BACKEND_URL only for an already-running backend; do not set it +# together with E2E_START_AGENT_BACKEND. # E2E_START_AGENT_BACKEND=1 # E2E_AGENT_BACKEND_PORT=5050 # E2E_AGENT_BACKEND_URL=http://127.0.0.1:5050 diff --git a/e2e/AGENTS.md b/e2e/AGENTS.md index abbf2a27b4b..fd10b188b1a 100644 --- a/e2e/AGENTS.md +++ b/e2e/AGENTS.md @@ -1,381 +1,71 @@ # E2E -This package contains the repository-level end-to-end tests for Dify. +This package contains Dify's repository-level Cucumber scenarios with Playwright as the browser layer. This file owns current package architecture, runtime, session and tag semantics, seed, protocol, and cleanup contracts. The repo-local `e2e-cucumber-playwright` skill owns authoring and review methodology; feature-specific facts belong in the nearest feature `AGENTS.md`. -This file is the canonical package guide for `e2e/`. Keep detailed workflow, architecture, debugging, and reporting documentation here. Keep `README.md` as a minimal pointer to this file so the two documents do not drift. +## Commands -The suite uses Cucumber for scenario definitions and Playwright as the browser execution layer. +Run commands from the repository root. Install dependencies and browsers once with `pnpm install` and `pnpm -C e2e e2e:install`. Run only one local `pnpm -C e2e e2e*` process at a time because runners share ports, auth state, and log paths. -It tests: +- Existing initialized instance: `pnpm -C e2e e2e` +- Standalone automated WCAG Level A scan: `pnpm -C e2e e2e:accessibility:a` +- Standalone automated WCAG Level AA scan: `pnpm -C e2e e2e:accessibility:aa` +- One-page automated WCAG scan: `pnpm -C e2e exec tsx ./scripts/run-cucumber.ts --full -- --tags "@axe and @wcag-a and @wcag-page-studio"` (replace the level and page tag as needed) +- Reset, initialize, and run deterministic scenarios: `pnpm -C e2e e2e:full` +- Prepare and run scenarios backed by shared fixtures: `E2E_START_AGENT_BACKEND=1 pnpm -C e2e e2e:prepared` +- Tagged subset: `pnpm -C e2e e2e -- --tags @smoke` +- Headed debugging: `pnpm -C e2e e2e:headed -- --tags @smoke` +- Prepare and run external runtime scenarios: `E2E_START_AGENT_BACKEND=1 pnpm -C e2e e2e:external` +- Seed against existing middleware without running Cucumber: `pnpm -C e2e seed -- --profile ` +- Reset persisted E2E state: `pnpm -C e2e e2e:reset` +- Middleware lifecycle: `pnpm -C e2e e2e:middleware:up` and `pnpm -C e2e e2e:middleware:down` +- Scoped static checks: `vp check e2e` -- backend API started from source -- frontend served from the production artifact -- middleware services started from Docker +The runner reuses `web/.next/BUILD_ID` when present. Set `E2E_FORCE_WEB_BUILD=1` to force a frontend rebuild. Use `E2E_BROWSER=webkit` for focused cross-browser runs and `E2E_SLOW_MO=500` with a headed command for local action debugging. -## Prerequisites +## Runtime Ownership -- Node.js `^22.22.1` -- `pnpm` -- `uv` -- Docker +- `scripts/setup.ts` owns reset, middleware, backend, and frontend startup. +- `scripts/run-cucumber.ts` is the only E2E runtime orchestrator. It owns service lifetime, optional seed execution, Cucumber invocation, and teardown. +- `scripts/seed-runner.ts` owns fixture creation and verification against an already-running runtime; it never starts services. +- `support/web-server.ts` owns frontend reuse, readiness, and shutdown. +- `features/support/hooks.ts` owns shared auth bootstrap, scenario lifecycle, and diagnostics. +- `features/support/world.ts` owns `DifyWorld`, the per-scenario behavior `BrowserContext`, and its authenticated setup and cleanup client. Browser and API identities remain separate so unauthenticated and logout journeys cannot invalidate fixture ownership. +- Cross-actor scenarios keep each actor in a separate `BrowserContext` and typed `DifyWorld` state so diagnostics and cleanup cover every actor. +- `features/step-definitions/` contains capability-oriented glue; `common/` is reserved for genuinely cross-capability steps. +- Step definitions that access World state use `async function (this: DifyWorld, ...)`; arrow functions cannot receive Cucumber's bound World instance. -Run the following commands from the repository root. +An uninitialized instance is installed and authenticated lazily; an initialized instance signs in and reuses authenticated state. Full runs prove reset and bootstrap during setup rather than through a Gherkin scenario. Cucumber's exit status is the behavior gate, and the runner also requires at least one `testCaseStarted` message so an empty tag selection cannot pass. Do not replace this gate with scenario-count baselines or skipped-scenario allowlists. -Install Playwright browsers once: +## Tags And External Runtime -```bash -pnpm install -pnpm -C e2e e2e:install -``` +- Default scenarios use shared authenticated storage state. `@unauthenticated` creates a clean context; `@authenticated` is an intent and selection tag only. +- `@axe` identifies standalone automated WCAG scans and is excluded from the default functional suite and normal CI commands. `@wcag-a` and `@wcag-aa` qualify the independent level-specific scans, and commands selecting either level must also select `@axe`. Page selectors use `@wcag-page-` and are attached to the matching Examples blocks under `features/accessibility/`. The accessibility workflow is an opt-in manual audit rather than a regression gate. The PR author should run the AA/all path before merge when changing the audit workflow, page matrix, or readiness contracts. +- `@prepared` requires the prepared fixtures; the post-merge seed profile includes them. +- `@external-model` and `@external-tool` identify scenarios that call real external runtimes. Deterministic commands exclude these tags; external commands are opt-in. +- `@microphone` uses the checked-in fake audio fixture and an isolated Chromium context. +- `@browser-smoke` runs focused keyboard and navigation coverage in Chromium and WebKit CI lanes. +- Feature-owned services use their own tags. Agent v2 runtime scenarios use `@agent-backend-runtime` and require the explicit runtime-availability step. Set `E2E_START_AGENT_BACKEND=1` to start it locally, or provide `E2E_AGENT_BACKEND_URL` / `AGENT_BACKEND_BASE_URL`. -`pnpm install` is resolved through the repository workspace and uses the shared root lockfile plus `pnpm-workspace.yaml`. +Seed and Cucumber must share one runtime lifecycle. Combined commands own reset, middleware, services, seed, Cucumber, and teardown; CI must not reproduce that lifecycle in workflow YAML. `E2E_START_AGENT_BACKEND=1` starts a managed local backend before the API; it is mutually exclusive with an explicit Agent backend URL. -Run only one `pnpm -C e2e e2e*` process against a local workspace at a time. Separate runner processes share the frontend port, backend port, auth bootstrap state, and log paths; running them in parallel can create startup or authorization failures that are not scenario failures. +Do not overload runtime tags to imply unrelated services or silently skip behavior when a required fixture is missing. -Use root lint plus the package type check as the default local verification step after editing E2E TypeScript, Cucumber support code, or feature glue: +## Browser, API, And Contract Boundaries -```bash -vp lint --fix --quiet -pnpm -C e2e type-check -``` +The action under test belongs to the browser. APIs may prepare fixtures, poll persistence, and clean up; they do not replace the user's `When` action. Prefer a user-observable browser result unless persisted backend state is the contract under test. -Common commands: +For ordinary Console JSON and representable multipart operations, use the scenario- or process-owned generated oRPC client with request and response validation enabled. Call generated operations directly. Do not add handwritten endpoint URLs, duplicate DTOs or schemas, response casts, one-to-one forwarding wrappers, mutable cross-scenario clients, or TanStack Query caching. -```bash -# deterministic regression against an initialized instance -# expects backend API, frontend artifact, and middleware stack to already be running -pnpm -C e2e e2e +Keep helpers only when they own fixture construction, multi-operation orchestration, cleanup registries, invariants, eventual-consistency polling, narrowed test views, or a protocol adapter. SSE, binary downloads, redirect-only flows, external services, and infrastructure readiness may use centralized adapters under their real owner. -# reset, initialize, and run deterministic scenarios -# starts required middleware/dependencies for you -pnpm -C e2e e2e:full +Validation failures are contract failures. Trace them to the backend schema owner, update `api/controllers/API_SCHEMA_GUIDE.md` contracts when required, regenerate `@dify/contracts`, and keep the scenario aligned with the product's real state owner. Do not disable validation or add fallback schemas to make E2E pass. -# run a tagged subset -pnpm -C e2e e2e -- --tags @smoke +## Seeds, Cleanup, And Diagnostics -# prepare external runtime seed resources for opt-in external suites -pnpm -C e2e e2e:external:prepare +- Generate disposable resource names through `support/naming.ts` with an `E2E` prefix. +- Keep deterministic upload material in `fixtures/test-materials/` and resolve it through `support/test-materials.ts`. +- Seed scripts own shared long-lived fixtures; scenarios own disposable resources they create and must register cleanup. +- Use typed `DifyWorld` cleanup fields for known resource types and `registerCleanup(...)` for additional lifecycle owners. Registered callbacks run LIFO after typed cleanup queues. +- Remove child and referencing resources before owners. Attach cleanup failures to the report instead of swallowing them. -# run scenarios that call real external providers -pnpm -C e2e e2e:external - -# headed browser -pnpm -C e2e e2e:headed -- --tags @smoke - -# slow down browser actions for local debugging -E2E_SLOW_MO=500 pnpm -C e2e e2e:headed -- --tags @smoke - -# focused keyboard and cross-browser smoke coverage -E2E_BROWSER=webkit pnpm -C e2e e2e -- --tags @browser-smoke -``` - -Frontend artifact behavior: - -- if `web/.next/BUILD_ID` exists, E2E reuses the existing build by default -- if you set `E2E_FORCE_WEB_BUILD=1`, E2E rebuilds the frontend before starting it - -## Lifecycle - -```mermaid -flowchart TD - A["Start E2E run"] --> B["run-cucumber.ts orchestrates setup/API/frontend"] - B --> C["support/web-server.ts starts or reuses frontend directly"] - C --> D["Cucumber loads config, steps, and support modules"] - D --> E["The first Before hook lazily bootstraps shared auth state"] - E --> F{"Which command is running?"} - F -->|`pnpm -C e2e e2e`| G["Run deterministic scenarios; exclude @prepared and external runtime"] - F -->|`pnpm -C e2e e2e:full*`| H["Reset and run deterministic scenarios; exclude @prepared and external runtime"] - G --> I["Per-scenario BrowserContext from shared browser"] - H --> I - I --> J["Failure artifacts written to cucumber-report/artifacts"] -``` - -Ownership is split like this: - -- `scripts/setup.ts` is the single environment entrypoint for reset, middleware, backend, and frontend startup -- `run-cucumber.ts` orchestrates the E2E run and Cucumber invocation -- `support/web-server.ts` manages frontend reuse, startup, readiness, and shutdown -- `features/support/hooks.ts` manages auth bootstrap, scenario lifecycle, and diagnostics -- `features/support/world.ts` owns the per-scenario behavior BrowserContext and authenticated setup/cleanup client; their identities remain separate so unauthenticated and logout journeys cannot invalidate fixture ownership -- `features/step-definitions/` holds domain-oriented glue so the official VS Code Cucumber plugin works with default conventions when `e2e/` is opened as the workspace root - -Package layout: - -- `features/`: Gherkin scenarios grouped by capability -- `features/step-definitions/`: domain-oriented step definitions -- `features/support/hooks.ts`: suite lifecycle, auth-state bootstrap, diagnostics -- `features/support/world.ts`: shared scenario context -- `support/web-server.ts`: typed frontend startup/reuse logic -- `scripts/setup.ts`: reset and service lifecycle commands -- `scripts/run-cucumber.ts`: Cucumber orchestration entrypoint - -Behavior depends on instance state: - -- uninitialized instance: completes install and stores authenticated state -- initialized instance: signs in and reuses authenticated state - -The `pnpm -C e2e e2e:full*` flows prove reset and authentication bootstrap by failing setup when initialization cannot complete; they do not model bootstrap state as a Gherkin scenario. Deterministic runs exclude `@prepared`, `@external-model`, and `@external-tool`. Post-merge first seeds required fixtures, then runs prepared and external scenarios. - -Reset all persisted E2E state: - -```bash -pnpm -C e2e e2e:reset -``` - -This removes: - -- `docker/volumes/db/data` -- `docker/volumes/redis/data` -- `docker/volumes/weaviate` -- `docker/volumes/plugin_daemon` -- `e2e/.auth` -- `e2e/.logs` -- `e2e/.logs-non-external` -- `e2e/.logs-webkit` -- `e2e/cucumber-report` -- `e2e/cucumber-report-non-external` -- `e2e/cucumber-report-webkit` -- `e2e/seed-report` - -Start the full middleware stack: - -```bash -pnpm -C e2e e2e:middleware:up -``` - -Stop the full middleware stack: - -```bash -pnpm -C e2e e2e:middleware:down -``` - -The middleware stack includes: - -- PostgreSQL -- Redis -- Weaviate -- Sandbox -- SSRF proxy -- Plugin daemon - -Fresh install verification: - -```bash -pnpm -C e2e e2e:full -``` - -Run the Cucumber suite against an already running middleware stack: - -```bash -pnpm -C e2e e2e:middleware:up -pnpm -C e2e e2e -pnpm -C e2e e2e:middleware:down -``` - -Artifacts and diagnostics: - -- `cucumber-report/report.html`: HTML report -- `cucumber-report/report.ndjson`: Cucumber Messages report -- `cucumber-report/artifacts/`: failure screenshots and HTML captures -- `cucumber-report-non-external/`: Chromium core report preserved before later CI lanes -- `cucumber-report-webkit/`: focused WebKit keyboard/browser smoke report -- `.logs/cucumber-api.log`: backend startup log -- `.logs/cucumber-web.log`: frontend startup log -- `.logs-non-external/`: non-external logs preserved before an external CI run -- `.logs-webkit/`: focused WebKit lane logs -- `seed-report/`: JSON readiness reports emitted by external runtime seed packs - -Cucumber's exit status is the behavior gate. The runner also requires at least one -`testCaseStarted` message so an empty or broken tag selector cannot pass silently. Do not add -scenario-count baselines or skipped-scenario allowlists. - -Open the HTML report locally with: - -```bash -open cucumber-report/report.html -``` - -## Scenario admission and behavior ownership - -Add an E2E scenario only when it protects a critical user journey and a cross-boundary result that -cheaper owner-level tests do not already prove. A control changing its own label is not sufficient -E2E evidence when component or integration tests can own that contract. - -Start from product truth, including real defaults and actor roles. API fixtures may establish -preconditions, but they must not manufacture an opposite state merely to make the intended action -look meaningful. When a product default is part of the journey, make it explicit and observable. - -For cross-actor journeys, isolate each actor's browser state, keep their pages in typed `DifyWorld` -state, and include them in failure diagnostics and cleanup. Assert the downstream user-observable -effect, not only the initiating control's local state. - -When a run exposes behavior that conflicts with the intended product contract, identify the first -layer that misclassifies the business state. Fix that owner or report the mismatch explicitly; do -not make the E2E pass by encoding an accidental redirect, stale label, or misleading error state. - -## Writing new scenarios - -### Workflow - -1. Create a `.feature` file under `features//` -1. Add step definitions under `features/step-definitions//` -1. Reuse existing steps from `common/` and other definition files before writing new ones -1. Run with `pnpm -C e2e e2e -- --tags @your-tag` to verify -1. Run `vp lint --fix --quiet` from the repository root and `pnpm -C e2e type-check` before committing - -### Feature file conventions - -Tag every feature or scenario with a capability tag. Add auth tags only when they clarify intent or change the browser session behavior: - -```gherkin -@datasets @authenticated -Feature: Create dataset - Scenario: Create a new empty dataset - Given I am signed in as the default E2E admin - When I open the datasets page - ... -``` - -- Capability tags (`@apps`, `@auth`, `@datasets`, …) group related scenarios for selective runs -- Auth/session tags: - - default behavior — scenarios run with the shared authenticated storageState unless marked otherwise - - `@unauthenticated` — uses a clean BrowserContext with no cookies or storage - - `@authenticated` — optional intent tag for readability or selective runs; it does not currently change hook behavior on its own -- `@prepared` — deterministic user behavior that requires the strict post-merge seed profile -- `@external-model` — scenario execution can call a real model provider. Use this only for runtime requests, not for scenarios that only require an active model fixture. -- `@external-tool` — scenario execution can call a real third-party tool provider. Use this only for runtime tool execution, not for plugin installation, discovery, or local deterministic tools. -- `@microphone` — runs the scenario in an isolated Chromium instance backed by the checked-in fake audio fixture and grants microphone permission only to that scenario context. -- `@browser-smoke` — focused keyboard and navigation coverage that runs in Chromium with the core suite and again in WebKit on CI. - External runtime commands are opt-in. `pnpm -C e2e e2e:external:prepare` prepares the fixed Agent v2 external-runtime seed and `pnpm -C e2e e2e:external` runs every `@external-model` or `@external-tool` scenario. CI uses `e2e:post-merge:prepare` followed by `e2e:post-merge` to run `@prepared` and external scenarios against one strict seed. - -The Agent v2 external runtime seed also prepares the workspace default Speech-to-Text model. `E2E_SPEECH_TO_TEXT_MODEL_PROVIDER` and `E2E_SPEECH_TO_TEXT_MODEL_NAME` select an existing model or the model configured through `E2E_MODEL_PROVIDER_CREDENTIALS_JSON`; they default to `openai` and `gpt-4o-mini-transcribe`. - -Some external runtime scenarios need feature-owned services in addition to a real model or tool provider. Do not overload `@external-model` or `@external-tool` to mean those services are available. For Agent v2, scenarios that require the standalone `dify-agent` run server use the feature tag `@agent-backend-runtime` plus the explicit step `the Agent v2 runtime backend is available`. Run them with `E2E_START_AGENT_BACKEND=1` to let E2E start `dify-agent` and the shellctl local sandbox required by its `dify.config`/`dify.shell` runtime layers, or set `E2E_AGENT_BACKEND_URL`/`AGENT_BACKEND_BASE_URL` when an existing server should be reused. - -Keep scenarios short and declarative. Each step should describe **what** the user does, not **how** the UI works. - -### Step definition conventions - -```typescript -import type { DifyWorld } from '../../support/world' -import { Then, When } from '@cucumber/cucumber' -import { expect } from '@playwright/test' - -When('I open the datasets page', async function (this: DifyWorld) { - await this.getPage().goto('/datasets') -}) -``` - -Rules: - -- Always type `this` as `DifyWorld` for proper context access -- Use `async function` (not arrow functions — Cucumber binds `this`) -- One step = one user-visible action or one assertion -- Keep steps stateless across scenarios; use `DifyWorld` properties for in-scenario state - -### Locator priority - -Follow the Playwright recommended locator strategy, in order of preference: - -| Priority | Locator | Example | When to use | -| -------- | ------------------ | ----------------------------------------- | ----------------------------------------- | -| 1 | `getByRole` | `getByRole('button', { name: 'Create' })` | Default choice — accessible and resilient | -| 2 | `getByLabel` | `getByLabel('App name')` | Form inputs with visible labels | -| 3 | `getByPlaceholder` | `getByPlaceholder('Enter name')` | Inputs without visible labels | -| 4 | `getByText` | `getByText('Welcome')` | Static text content | -| 5 | `getByTestId` | `getByTestId('workflow-canvas')` | Only when no semantic locator works | - -Avoid raw CSS/XPath selectors. They break when the DOM structure changes. - -### Assertions - -Use `@playwright/test` `expect` — it auto-waits and retries until the condition is met or the timeout expires: - -```typescript -// URL assertion -await expect(page).toHaveURL(/\/datasets\/[a-f0-9-]+\/documents/) - -// Element visibility -await expect(page.getByRole('button', { name: 'Save' })).toBeVisible() - -// Element state -await expect(page.getByRole('button', { name: 'Submit' })).toBeEnabled() - -// Negation -await expect(page.getByText('Loading')).not.toBeVisible() -``` - -Do not use manual `waitForTimeout` or polling loops. If you need a longer wait for a specific assertion, pass `{ timeout: 30_000 }` to the assertion. - -### Cucumber expressions - -Use Cucumber expression parameter types to extract values from Gherkin steps: - -| Type | Pattern | Example step | -| ---------- | ------------- | ---------------------------------- | -| `{string}` | Quoted string | `I select the "Workflow" app type` | -| `{int}` | Integer | `I should see {int} items` | -| `{float}` | Decimal | `the progress is {float} percent` | -| `{word}` | Single word | `I click the {word} tab` | - -Prefer `{string}` for UI labels, names, and text content — it maps naturally to Gherkin's quoted values. - -### Scoping locators - -When the page has multiple similar elements, scope locators to a container: - -```typescript -When('I fill in the app name in the dialog', async function (this: DifyWorld) { - const dialog = this.getPage().getByRole('dialog') - await dialog.getByPlaceholder('Give your app a name').fill('My App') -}) -``` - -### Failure diagnostics - -The `After` hook automatically captures diagnostics for failed, ambiguous, pending, undefined, or unknown scenarios: - -- Full-page screenshot (PNG) -- Page HTML dump -- Console errors and page errors - -Artifacts are saved to `cucumber-report/artifacts/` and attached to the HTML report. No extra code needed in step definitions. - -### Seed and fixture contracts - -Use `support/naming.ts` for generated test resource names. New app, Agent, dataset, file, or credential seeds should start with `E2E` so local and shared environments can identify disposable resources. - -Use `fixtures/test-materials/` for checked-in files that scenarios upload, preview, index, or retrieve. Keep these fixtures small and deterministic, and use `support/test-materials.ts` to resolve their absolute paths. - -Seed scripts own long-lived models, plugins, datasets, and fixed apps. Selected scenarios may resolve and verify those fixtures through explicit `Given` steps, but a missing or drifted fixture must fail the scenario. Do not represent environment readiness as Gherkin scenarios and do not conditionally skip behavior. - -Keep package-level support limited to broadly reusable primitives such as API clients, naming, fixture path resolution, and cleanup helpers. Feature-specific seed and fixture contracts belong under the owning feature's support folder. - -Use generated API contracts for Console/Web/Service API request, response, and payload shapes. Import the concrete type directly from `@dify/contracts/.../types.gen` when it exists, and do not hand-write duplicate response shapes or wrap generated types in local aliases just to preserve an older helper name. Keep local E2E types only for scenario state, fixture registries, helper input options, and intentionally narrowed test view models that are not complete API responses. - -### Console API and protocol boundaries - -The action under test belongs to the browser. `When` steps must use Playwright to perform the user action; do not replace the action with an API request. `Given` setup, seed preparation, persistence polling, and `After` cleanup may use APIs when that makes the scenario faster and more deterministic. `Then` should prefer a user-observable browser result; an API read is appropriate only when persistence itself is the asserted contract and the endpoint owns that state. - -For ordinary Console JSON operations and multipart uploads represented by Console OpenAPI, use the generated oRPC router with generated request and response validation enabled. A scenario client belongs to its `DifyWorld` and uses a scenario-owned authenticated request context that is independent from the behavior browser; seed processes own a standalone client for their process lifetime. Do not create a mutable cross-scenario API client, add TanStack Query caching to Cucumber, hand-write Console endpoint URLs, cast response JSON to an API DTO, or duplicate a generated Zod schema. When a browser action's captured response must provide an ID for cleanup, parse it with the generated response schema. - -Do not add a helper that only renames or forwards one generated operation. Call the generated client directly from the owning step, hook, or fixture orchestration. Keep a helper only when it owns a real test concern such as constructing a valid domain fixture, coordinating multiple operations, maintaining an invariant or cleanup registry, polling eventual consistency, deriving a narrowed test view, or adapting a non-OpenAPI protocol. - -SSE/event streams, binary downloads, redirect-only flows, external services, and infrastructure health/readiness checks may use a dedicated protocol adapter. Keep each exception centralized under its real owner and continue to use generated payload types where the contract covers the request. Multipart is not an exception merely because it carries a file: fix the backend OpenAPI schema and regenerate when the operation can be represented. - -Request or response validation failures are contract failures. Do not suppress them with casts, permissive fallback schemas, disabled validation, swallowed cleanup errors, or a second handwritten request path. Trace the mismatch to the endpoint's backend schema owner, update it according to `api/controllers/API_SCHEMA_GUIDE.md`, regenerate `@dify/contracts`, and keep the E2E assertion aligned with the product's real state owner rather than an internal backing resource. - -Use typed cleanup fields on `DifyWorld` for resource types created by scenarios, and use `DifyWorld.registerCleanup(...)` when a scenario creates any resource type that is not covered by typed cleanup fields. Typed cleanup should remove child or referencing resources before their owners, such as Agent files before Agents and workflow apps before Agents they reference. Cleanup failures should be attached to the report instead of being swallowed silently. Cleanup callbacks run after typed cleanup queues, even when the scenario fails. - -Scenario-owned setup may create disposable apps, Agents, files, credentials, drafts, or access toggles when the scenario owns their lifecycle and cleanup. Do not use scenario setup to silently fix a shared fixture; a missing or drifted fixed resource is a seed failure. - -Feature-specific seed contracts, resource readiness rules, tags, and scenario ownership can be documented in one scoped `AGENTS.md` at the feature root when a module becomes large enough to need it. Do not add deeper `AGENTS.md` files unless the nested module becomes independently owned. - -## Reusing existing steps - -Before writing a new step definition, inspect the existing step definition files first. Reuse a matching step when the wording and behavior already fit, and only add a new step when the scenario needs a genuinely new user action or assertion. Steps in `common/` are designed for broad reuse across all features. - -Or browse the step definition files directly: - -- `features/step-definitions/common/` — auth guards and navigation assertions shared by all features -- `features/step-definitions//` — domain-specific steps scoped to a single feature area +Failures produce screenshots and HTML captures under `cucumber-report/artifacts/`; the HTML and Cucumber Messages reports live under `cucumber-report/`. Backend and frontend startup logs live under `.logs/`. Additional CI lanes preserve their own report and log directories. diff --git a/e2e/cucumber.config.ts b/e2e/cucumber.config.ts index edf96d3a0e2..81fbcacb7e2 100644 --- a/e2e/cucumber.config.ts +++ b/e2e/cucumber.config.ts @@ -2,7 +2,8 @@ import type { IConfiguration } from '@cucumber/cucumber' import './scripts/env-register' const hasCliTags = process.argv.some((arg) => arg === '--tags' || arg.startsWith('--tags=')) -const defaultNonExternalTags = 'not @prepared and not @external-model and not @external-tool' +const defaultNonExternalTags = + 'not @axe and not @prepared and not @external-model and not @external-tool' const defaultTags = process.env.E2E_CUCUMBER_TAGS || (hasCliTags ? undefined : defaultNonExternalTags) diff --git a/e2e/features/accessibility/main-pages.feature b/e2e/features/accessibility/main-pages.feature new file mode 100644 index 00000000000..4ae6a71f8e1 --- /dev/null +++ b/e2e/features/accessibility/main-pages.feature @@ -0,0 +1,59 @@ +@accessibility @axe @authenticated +Feature: Main page automated WCAG checks + Axe checks only automatically detectable issues and does not establish WCAG conformance. + + Scenario Outline: page has no automatically detectable WCAG Level violations + Given I am signed in as the default E2E admin + When I open the "" main page + Then I should be on the "" main page + And the current page should have no automatically detectable WCAG Level violations + + @wcag-a @wcag-page-agents + Examples: Agents at Level A + | page | level | + | Agents | A | + + @wcag-aa @wcag-page-agents + Examples: Agents at Level AA + | page | level | + | Agents | AA | + + @wcag-a @wcag-page-home + Examples: Home at Level A + | page | level | + | Home | A | + + @wcag-aa @wcag-page-home + Examples: Home at Level AA + | page | level | + | Home | AA | + + @wcag-a @wcag-page-integrations + Examples: Integrations at Level A + | page | level | + | Integrations | A | + + @wcag-aa @wcag-page-integrations + Examples: Integrations at Level AA + | page | level | + | Integrations | AA | + + @wcag-a @wcag-page-knowledge + Examples: Knowledge at Level A + | page | level | + | Knowledge | A | + + @wcag-aa @wcag-page-knowledge + Examples: Knowledge at Level AA + | page | level | + | Knowledge | AA | + + @wcag-a @wcag-page-studio + Examples: Studio at Level A + | page | level | + | Studio | A | + + @wcag-aa @wcag-page-studio + Examples: Studio at Level AA + | page | level | + | Studio | AA | diff --git a/e2e/features/accessibility/sign-in.feature b/e2e/features/accessibility/sign-in.feature new file mode 100644 index 00000000000..3080229ef01 --- /dev/null +++ b/e2e/features/accessibility/sign-in.feature @@ -0,0 +1,19 @@ +@accessibility @axe @unauthenticated +Feature: Sign-in page automated WCAG checks + Axe checks only automatically detectable issues and does not establish WCAG conformance. + + Scenario Outline: Sign-in page has no automatically detectable WCAG Level violations + Given I am not signed in + When I open the sign-in page + Then I should see the "Sign in" button + And the current page should have no automatically detectable WCAG Level violations + + @wcag-a @wcag-page-sign-in + Examples: Level A + | level | + | A | + + @wcag-aa @wcag-page-sign-in + Examples: Level AA + | level | + | AA | diff --git a/e2e/features/agent-v2/AGENTS.md b/e2e/features/agent-v2/AGENTS.md index ab8a20a53c2..1e27c084aa4 100644 --- a/e2e/features/agent-v2/AGENTS.md +++ b/e2e/features/agent-v2/AGENTS.md @@ -48,16 +48,15 @@ Use `the Agent v2 configuration should be saved automatically` for Configure aut ## Seed and fixture contract -Seed scripts create or update environment-owned models, plugins, datasets, Agents, and workflows. `fixtures.steps.ts` resolves and validates those resources before a dependent behavior runs. Missing, inactive, unindexed, or drifted fixtures must throw and fail the scenario; never return `skipped`. +Seed tasks create or update environment-owned models, plugins, datasets, Agents, and workflows. `fixtures.steps.ts` resolves and validates those resources before a dependent behavior runs. Missing, inactive, unindexed, or drifted fixtures must throw and fail the scenario; never return `skipped`. `@prepared` scenarios are excluded from deterministic PR core. Post-merge runs: ```bash -pnpm -C e2e e2e:post-merge:prepare -pnpm -C e2e e2e:post-merge +E2E_START_AGENT_BACKEND=1 pnpm -C e2e e2e:post-merge ``` -The strict seed must finish without blocked tasks. The concrete resource inventory and defaults belong to the seed profile and environment configuration rather than this guidance. +The command owns runtime setup, strict seed, Cucumber, and teardown. The strict seed must finish without blocked tasks. The concrete resource inventory and defaults belong to the seed profile and environment configuration rather than this guidance. Organize fixture helpers by the product resource or infrastructure capability they own, not by the feature file that happens to consume them. Keep runtime readiness adapters separate from Console resource fixtures, and keep all fixture state in the current `SeedContext` or scenario `DifyWorld` rather than module globals. diff --git a/e2e/features/agent-v2/access-point.feature b/e2e/features/agent-v2/access-point.feature index 53a775739b0..2870e5270f9 100644 --- a/e2e/features/agent-v2/access-point.feature +++ b/e2e/features/agent-v2/access-point.feature @@ -1,18 +1,18 @@ @agent-v2 @authenticated @access-point Feature: Agent v2 Access Point @core - Scenario: Access Point shows the available Agent v2 access surfaces + Scenario: Access Point keeps unpublished Agent v2 access unavailable Given I am signed in as the default E2E admin And an Agent v2 test agent has been created via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section - Then I should see the Agent v2 Access Point overview + Then the unpublished Agent v2 access surfaces should be unavailable @core @web-app-access Scenario: Web app access URL can be copied without changing orchestration Given I am signed in as the default E2E admin And a basic configured Agent v2 test agent has been created via API - And Agent v2 Web app access has been enabled via API + And the Agent v2 draft has been published via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section Then I should see the Agent v2 Web app access URL @@ -25,7 +25,6 @@ Feature: Agent v2 Access Point Given I am signed in as the default E2E admin And a basic configured Agent v2 test agent has been created via API And the Agent v2 draft has been published via API - And Agent v2 Web app access has been enabled via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section Then I should see the Agent v2 Web app access URL @@ -37,7 +36,7 @@ Feature: Agent v2 Access Point Scenario: Web app Embedded configuration opens from Access Point Given I am signed in as the default E2E admin And a basic configured Agent v2 test agent has been created via API - And Agent v2 Web app access has been enabled via API + And the Agent v2 draft has been published via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section And I open Agent v2 Embedded configuration @@ -48,7 +47,7 @@ Feature: Agent v2 Access Point Scenario: Web app customization opens from Access Point Given I am signed in as the default E2E admin And a basic configured Agent v2 test agent has been created via API - And Agent v2 Web app access has been enabled via API + And the Agent v2 draft has been published via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section And I open Agent v2 Web app customization @@ -59,7 +58,7 @@ Feature: Agent v2 Access Point Scenario: Web app settings open from Access Point without changing orchestration Given I am signed in as the default E2E admin And a basic configured Agent v2 test agent has been created via API - And Agent v2 Web app access has been enabled via API + And the Agent v2 draft has been published via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section And I open Agent v2 Web app settings @@ -71,15 +70,15 @@ Feature: Agent v2 Access Point Given I am signed in as the default E2E admin And a basic configured Agent v2 test agent has been created via API And the Agent v2 draft has been published via API - And Agent v2 Web app access has been enabled via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section And I disable Agent v2 Web app access Then Agent v2 Web app access should be out of service + When I republish the Agent v2 draft via API + And I refresh the current page + Then Agent v2 Web app access should be out of service When I enable Agent v2 Web app access Then Agent v2 Web app access should be in service - When I refresh the current page - Then Agent v2 Web app access should be in service @core @prepared @workflow-reference Scenario: Workflow access shows the referencing workflow @@ -96,7 +95,7 @@ Feature: Agent v2 Access Point Scenario: Backend service API endpoint can be copied Given I am signed in as the default E2E admin And an Agent v2 test agent has been created via API - And Agent v2 Backend service API access has been enabled via API + And the Agent v2 draft has been published via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section Then I should see the Agent v2 Backend service API endpoint @@ -107,7 +106,8 @@ Feature: Agent v2 Access Point Scenario: Backend service API keys are managed without exposing existing secrets Given I am signed in as the default E2E admin And an Agent v2 test agent has been created via API - And Agent v2 Backend service API access has been enabled with a key via API + And the Agent v2 draft has been published via API + And an Agent v2 Backend service API key has been created via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section And I open Agent v2 API key management @@ -123,7 +123,7 @@ Feature: Agent v2 Access Point Scenario: Backend service API Reference opens from Access Point Given I am signed in as the default E2E admin And an Agent v2 test agent has been created via API - And Agent v2 Backend service API access has been enabled via API + And the Agent v2 draft has been published via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section And I open the Agent v2 API Reference @@ -133,7 +133,7 @@ Feature: Agent v2 Access Point Scenario: Backend service API access can be disabled and restored from Access Point Given I am signed in as the default E2E admin And an Agent v2 test agent has been created via API - And Agent v2 Backend service API access has been enabled via API + And the Agent v2 draft has been published via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section And I disable Agent v2 Backend service API access @@ -149,10 +149,10 @@ Feature: Agent v2 Access Point And the Agent Builder stable chat model is available And the Agent v2 runtime backend is available And a runnable Agent v2 test agent has been created via API - And Agent v2 Backend service API access has been enabled with a key via API When I open the Agent v2 configure page And I publish the Agent v2 draft Then the Agent v2 draft should be published and up to date + Given an Agent v2 Backend service API key has been created via API When I send the Agent v2 Backend service API minimal request Then the Agent v2 Backend service API request should succeed with the normal E2E marker @@ -163,7 +163,7 @@ Feature: Agent v2 Access Point And the Agent v2 runtime backend is available And a runnable Agent v2 test agent has been created via API And the Agent v2 draft has been published via API - And Agent v2 Backend service API access has been enabled with a key via API + And an Agent v2 Backend service API key has been created via API When I open the Agent v2 configure page from the Agent Roster And I switch to the Agent v2 Access Point section And I disable Agent v2 Backend service API access diff --git a/e2e/features/agent-v2/agent-edit.feature b/e2e/features/agent-v2/agent-edit.feature index 7d3bd949da4..ca13b4fc845 100644 --- a/e2e/features/agent-v2/agent-edit.feature +++ b/e2e/features/agent-v2/agent-edit.feature @@ -17,7 +17,7 @@ Feature: Agent v2 Agent Edit page And the Agent Builder preseeded Agent "E2E New Agent Builder Full Config" includes the core fixture configuration And the preseeded Agent v2 "E2E New Agent Builder Full Config" has been published via API When I duplicate the preseeded Agent v2 "E2E New Agent Builder Full Config" from the Agent Roster - Then the duplicated Agent v2 should inherit the full-config fixture from "E2E New Agent Builder Full Config" + Then the duplicated Agent v2 should inherit the full-config fixture from "E2E New Agent Builder Full Config" without inheriting its publication state When I open the Agent v2 configure page And I fill the Agent v2 prompt editor with the updated E2E prompt Then the Agent v2 configuration should be saved automatically diff --git a/e2e/features/agent-v2/knowledge.feature b/e2e/features/agent-v2/knowledge.feature index be8796fc299..20184bf697c 100644 --- a/e2e/features/agent-v2/knowledge.feature +++ b/e2e/features/agent-v2/knowledge.feature @@ -31,13 +31,13 @@ Feature: Agent v2 Knowledge Retrieval And the Agent v2 runtime backend is available And the Agent Builder preseeded dataset "E2E Agent Knowledge Base" is indexed and ready And a runnable Agent v2 test agent using the agent-decision model has been created via API - And Agent v2 Backend service API access has been enabled with a key via API When I open the Agent v2 configure page And I add the Agent Builder knowledge base as an Agent decide Knowledge Retrieval Then the Agent v2 Agent decide Knowledge Retrieval should be saved in the Agent v2 draft And the Agent v2 configuration should be saved automatically When I publish the Agent v2 draft Then the Agent v2 draft should be published and up to date + Given an Agent v2 Backend service API key has been created via API When I send the Agent v2 Backend service API knowledge request Then the Agent v2 Backend service API response should include the knowledge E2E marker @@ -48,13 +48,13 @@ Feature: Agent v2 Knowledge Retrieval And the Agent v2 runtime backend is available And the Agent Builder preseeded dataset "E2E Agent Knowledge Base" is indexed and ready And a runnable Agent v2 test agent has been created via API - And Agent v2 Backend service API access has been enabled with a key via API When I open the Agent v2 configure page And I add the Agent Builder knowledge base as a Custom query Knowledge Retrieval Then the Agent v2 Custom query Knowledge Retrieval should be saved in the Agent v2 draft And the Agent v2 configuration should be saved automatically When I publish the Agent v2 draft Then the Agent v2 draft should be published and up to date + Given an Agent v2 Backend service API key has been created via API When I send the Agent v2 Backend service API knowledge request Then the Agent v2 Backend service API response should include the knowledge E2E marker diff --git a/e2e/features/agent-v2/publish.feature b/e2e/features/agent-v2/publish.feature index 88a84dc9381..30a8eabb799 100644 --- a/e2e/features/agent-v2/publish.feature +++ b/e2e/features/agent-v2/publish.feature @@ -18,6 +18,9 @@ Feature: Agent v2 publish When I open the Agent v2 configure page And I publish the Agent v2 draft Then the Agent v2 draft should be published and up to date + When I switch to the Agent v2 Access Point section + Then Agent v2 Web app access should be in service + And Agent v2 Backend service API access should be in service @core @prepared @stable-model Scenario: Publish action follows unpublished changes @@ -61,7 +64,6 @@ Feature: Agent v2 publish And the Agent Builder stable chat model is available And the Agent v2 runtime backend is available And a runnable Agent v2 test agent has been created via API - And Agent v2 Web app access has been enabled via API When I open the Agent v2 configure page And I publish the Agent v2 draft Then the Agent v2 draft should be published and up to date @@ -76,7 +78,6 @@ Feature: Agent v2 publish And the Agent Builder stable chat model is available And the Agent v2 runtime backend is available And a runnable Agent v2 test agent has been created via API - And Agent v2 Web app access has been enabled via API When I open the Agent v2 configure page And I publish the Agent v2 draft Then the Agent v2 draft should be published and up to date @@ -95,7 +96,6 @@ Feature: Agent v2 publish And the Agent Builder stable chat model is available And the Agent v2 runtime backend is available And a runnable Agent v2 test agent has been created via API - And Agent v2 Web app access has been enabled via API When I open the Agent v2 configure page And I publish the Agent v2 draft Then the Agent v2 draft should be published and up to date diff --git a/e2e/features/agent-v2/support/access-point.ts b/e2e/features/agent-v2/support/access-point.ts index 575cef3e662..014e4a3c53d 100644 --- a/e2e/features/agent-v2/support/access-point.ts +++ b/e2e/features/agent-v2/support/access-point.ts @@ -1,6 +1,5 @@ import type { AgentAppDetailWithSite } from '@dify/contracts/api/console/agent/types.gen' import type { ChatRequestPayloadWithUser } from '@dify/contracts/api/service/types.gen' -import type { ConsoleClient } from '../../../support/api/console-client' import { consumeServiceApiSse, SERVICE_API_STREAM_TIMEOUT_MS } from './service-api-sse' export type AgentServiceApiChatResult = { @@ -41,19 +40,6 @@ export function getAgentWebAppURL(agent: AgentAppDetailWithSite): string { return `${baseURL.replace(/\/$/, '')}/agent/${token}` } -export async function enableAgentWebApp(client: ConsoleClient, agentId: string): Promise { - const agent = await client.agent.byAgentId.get({ params: { agent_id: agentId } }) - const appId = agent.app_id ?? agent.backing_app_id - if (!appId) throw new Error(`Agent v2 ${agentId} does not expose a backing app ID.`) - - await client.apps.byAppId.siteEnable.post({ - body: { enable_site: true }, - params: { app_id: appId }, - }) - const updatedAgent = await client.agent.byAgentId.get({ params: { agent_id: agentId } }) - return getAgentWebAppURL(updatedAgent) -} - export async function sendAgentServiceApiChatMessage({ apiKey, query = 'Please reply with the test success marker.', diff --git a/e2e/features/agent-v2/support/workflow.ts b/e2e/features/agent-v2/support/workflow.ts index 19ca155f6d8..48a548787c0 100644 --- a/e2e/features/agent-v2/support/workflow.ts +++ b/e2e/features/agent-v2/support/workflow.ts @@ -60,7 +60,6 @@ export async function syncAgentV2WorkflowDraft( viewport: { x: 0, y: 0, zoom: 1 }, }, features: {}, - environment_variables: [], conversation_variables: [], } satisfies SyncDraftWorkflowPayload await client.apps.byAppId.workflows.draft.post({ body, params: { app_id: appId } }) diff --git a/e2e/features/agent-v2/tools.feature b/e2e/features/agent-v2/tools.feature index 262389b8f69..c84f4274148 100644 --- a/e2e/features/agent-v2/tools.feature +++ b/e2e/features/agent-v2/tools.feature @@ -19,10 +19,10 @@ Feature: Agent v2 tools And the Agent Builder stable chat model is available And the Agent Builder preseeded tool "JSON Process / JSON Replace" is available And a runnable Agent v2 test agent with the JSON Replace tool has been created via API - And Agent v2 Backend service API access has been enabled with a key via API When I open the Agent v2 configure page Then the Agent v2 JSON Replace tool should be saved in the Agent v2 draft When I publish the Agent v2 draft Then the Agent v2 draft should be published and up to date + Given an Agent v2 Backend service API key has been created via API When I send the Agent v2 Backend service API JSON Replace request Then the Agent v2 Backend service API response should include the JSON Replace E2E marker diff --git a/e2e/features/apps/web-app-service.feature b/e2e/features/apps/web-app-service.feature index f027ed3e9da..3dc55c23437 100644 --- a/e2e/features/apps/web-app-service.feature +++ b/e2e/features/apps/web-app-service.feature @@ -4,8 +4,7 @@ Feature: Manage Web App service Scenario: Disable and restore a published workflow Web App Given I am signed in as the default E2E admin And a new runnable workflow app has been published - When I navigate to the app overview page - And I open the app information panel + When I navigate to the app access point page Then the Web App should be in service When an anonymous visitor opens the Web App Then the published workflow Web App should be accessible diff --git a/e2e/features/step-definitions/accessibility/axe.steps.ts b/e2e/features/step-definitions/accessibility/axe.steps.ts new file mode 100644 index 00000000000..12bbe95d01a --- /dev/null +++ b/e2e/features/step-definitions/accessibility/axe.steps.ts @@ -0,0 +1,69 @@ +import type { DifyWorld } from '../../support/world' +import AxeBuilder from '@axe-core/playwright' +import { Then } from '@cucumber/cucumber' +import { expect } from '@playwright/test' + +type AxeResults = Awaited> +type WcagLevel = 'A' | 'AA' + +const wcagTagsByLevel = { + A: ['wcag2a', 'wcag21a'], + AA: ['wcag2a', 'wcag2aa', 'wcag21a', 'wcag21aa', 'wcag22aa'], +} satisfies Record + +const formatFindings = (findings: AxeResults['violations']) => + findings + .map( + (violation) => + `${violation.id} (${violation.impact ?? 'unknown impact'}): ${violation.help}\n${violation.helpUrl}\n${violation.nodes.map((node) => ` - ${node.target.join(' > ')}${node.failureSummary ? `\n ${node.failureSummary.replaceAll('\n', '\n ')}` : ''}`).join('\n')}`, + ) + .join('\n\n') + +const checkCurrentPage = async (world: DifyWorld, level: WcagLevel) => { + const results = await new AxeBuilder({ page: world.getPage() }) + .withTags(wcagTagsByLevel[level]) + .analyze() + const formattedViolations = formatFindings(results.violations) + const formattedIncomplete = formatFindings(results.incomplete) + + if (results.violations.length > 0) { + world.attach( + `WCAG Level ${level} violations for ${results.url}:\n\n${formattedViolations}`, + 'text/plain', + ) + } + + if (results.incomplete.length > 0) { + world.attach( + `WCAG Level ${level} items requiring manual review for ${results.url}:\n\n${formattedIncomplete}`, + 'text/plain', + ) + } + + if (results.violations.length > 0 || results.incomplete.length > 0) + world.attach( + JSON.stringify( + { + incomplete: results.incomplete, + level, + url: results.url, + violations: results.violations, + }, + null, + 2, + ), + 'application/json', + ) + + expect(results.violations, formattedViolations).toEqual([]) +} + +Then( + 'the current page should have no automatically detectable WCAG Level {word} violations', + async function (this: DifyWorld, level: string) { + if (!Object.hasOwn(wcagTagsByLevel, level)) + throw new Error(`Unsupported WCAG level "${level}". Expected A or AA.`) + + await checkCurrentPage(this, level as WcagLevel) + }, +) diff --git a/e2e/features/step-definitions/accessibility/navigation.steps.ts b/e2e/features/step-definitions/accessibility/navigation.steps.ts new file mode 100644 index 00000000000..af37b1ca6d1 --- /dev/null +++ b/e2e/features/step-definitions/accessibility/navigation.steps.ts @@ -0,0 +1,42 @@ +import type { Page } from '@playwright/test' +import type { DifyWorld } from '../../support/world' +import { Then, When } from '@cucumber/cucumber' +import { waitForAgentsConsole } from '../../../support/agents' +import { waitForAppsConsole } from '../../../support/apps' +import { waitForConsoleHome } from '../../../support/home' +import { waitForModelProviderIntegrations } from '../../../support/integrations' +import { waitForKnowledgeConsole } from '../../../support/knowledge' + +type AccessibilityPageConfig = { + path: string + waitUntilReady: (page: Page) => Promise +} + +const accessibilityPages = { + Agents: { path: '/agents', waitUntilReady: waitForAgentsConsole }, + Home: { path: '/', waitUntilReady: waitForConsoleHome }, + Integrations: { + path: '/integrations/model-provider', + waitUntilReady: waitForModelProviderIntegrations, + }, + Knowledge: { path: '/datasets', waitUntilReady: waitForKnowledgeConsole }, + Studio: { path: '/apps', waitUntilReady: waitForAppsConsole }, +} satisfies Record + +const getAccessibilityPage = (pageName: string): AccessibilityPageConfig => { + const config = accessibilityPages[pageName as keyof typeof accessibilityPages] + if (!config) + throw new Error( + `Unknown accessibility page "${pageName}". Expected one of: ${Object.keys(accessibilityPages).join(', ')}.`, + ) + + return config +} + +When('I open the {string} main page', async function (this: DifyWorld, pageName: string) { + await this.getPage().goto(getAccessibilityPage(pageName).path) +}) + +Then('I should be on the {string} main page', async function (this: DifyWorld, pageName: string) { + await getAccessibilityPage(pageName).waitUntilReady(this.getPage()) +}) diff --git a/e2e/features/step-definitions/agent-v2/access-point-service-api.steps.ts b/e2e/features/step-definitions/agent-v2/access-point-service-api.steps.ts index c6f62ccc6f8..330d7abca77 100644 --- a/e2e/features/step-definitions/agent-v2/access-point-service-api.steps.ts +++ b/e2e/features/step-definitions/agent-v2/access-point-service-api.steps.ts @@ -9,13 +9,12 @@ import { import { SERVICE_API_RUNTIME_STEP_TIMEOUT_MS } from '../../agent-v2/support/service-api-sse' import { getCurrentAgentId, getServiceApiCard } from './access-point-helpers' -async function enableAgentApiAccessWithKey(world: DifyWorld) { +const API_KEY_DIALOG_NAME = /^API Key$/i + +async function createAgentApiKey(world: DifyWorld) { const agentId = getCurrentAgentId(world) const client = world.getConsoleClient() - const apiAccess = await client.agent.byAgentId.apiEnable.post({ - body: { enable_api: true }, - params: { agent_id: agentId }, - }) + const apiAccess = await client.agent.byAgentId.apiAccess.get({ params: { agent_id: agentId } }) const apiKey = await client.agent.byAgentId.apiKeys.post({ params: { agent_id: agentId } }) world.agentBuilder.accessPoint.serviceApiBaseURL = apiAccess.service_api_base_url @@ -23,25 +22,25 @@ async function enableAgentApiAccessWithKey(world: DifyWorld) { } Given( - 'Agent v2 Backend service API access has been enabled with a key via API', + 'an Agent v2 Backend service API key has been created via API', async function (this: DifyWorld) { - await enableAgentApiAccessWithKey(this) + await createAgentApiKey(this) }, ) Then('I should see the Agent v2 Backend service API endpoint', async function (this: DifyWorld) { const serviceApiCard = getServiceApiCard(this) - - if (!this.agentBuilder.accessPoint.serviceApiBaseURL) - throw new Error('No Agent v2 service API endpoint found. Enable Backend service API first.') + const agentId = getCurrentAgentId(this) + const apiAccess = await this.getConsoleClient().agent.byAgentId.apiAccess.get({ + params: { agent_id: agentId }, + }) + this.agentBuilder.accessPoint.serviceApiBaseURL = apiAccess.service_api_base_url await expect(serviceApiCard.getByRole('heading', { name: 'Backend service API' })).toBeVisible({ timeout: 30_000, }) await expect(serviceApiCard.getByText('Service API Endpoint')).toBeVisible() - await expect( - serviceApiCard.getByText(this.agentBuilder.accessPoint.serviceApiBaseURL), - ).toBeVisible() + await expect(serviceApiCard.getByText(apiAccess.service_api_base_url)).toBeVisible() await expect(serviceApiCard.getByLabel('Copy service API endpoint')).toBeEnabled() }) @@ -64,7 +63,7 @@ When('I open Agent v2 API key management', async function (this: DifyWorld) { Then('Agent v2 API keys should not expose a secret by default', async function (this: DifyWorld) { const page = this.getPage() - const dialog = page.getByRole('dialog', { name: /API Secret key/i }) + const dialog = page.getByRole('dialog', { name: API_KEY_DIALOG_NAME }) const existingSecret = this.agentBuilder.accessPoint.generatedApiKey await expect(dialog).toBeVisible() @@ -79,14 +78,14 @@ Then('Agent v2 API keys should not expose a secret by default', async function ( }) When('I create a new Agent v2 API key', async function (this: DifyWorld) { - const dialog = this.getPage().getByRole('dialog', { name: /API Secret key/i }) + const dialog = this.getPage().getByRole('dialog', { name: API_KEY_DIALOG_NAME }) await dialog.getByRole('button', { name: 'Create new Secret key' }).click() }) Then('I should see the newly generated Agent v2 API key once', async function (this: DifyWorld) { const generatedKeyDialog = this.getPage() - .getByRole('dialog', { name: /API Secret key/i }) + .getByRole('dialog', { name: API_KEY_DIALOG_NAME }) .last() const generatedKey = generatedKeyDialog.getByText(/^app-/) @@ -104,7 +103,7 @@ Then('I should see the newly generated Agent v2 API key once', async function (t When('I copy the newly generated Agent v2 API key', async function (this: DifyWorld) { const generatedKeyDialog = this.getPage() - .getByRole('dialog', { name: /API Secret key/i }) + .getByRole('dialog', { name: API_KEY_DIALOG_NAME }) .last() await generatedKeyDialog.getByLabel('Copy').first().click() @@ -114,7 +113,7 @@ Then( 'the newly generated Agent v2 API key should show it was copied', async function (this: DifyWorld) { const generatedKeyDialog = this.getPage() - .getByRole('dialog', { name: /API Secret key/i }) + .getByRole('dialog', { name: API_KEY_DIALOG_NAME }) .last() await expect(generatedKeyDialog.getByLabel('Copied')).toBeVisible() @@ -123,7 +122,7 @@ Then( When('I close the newly generated Agent v2 API key', async function (this: DifyWorld) { const page = this.getPage() - const generatedKeyDialog = page.getByRole('dialog', { name: /API Secret key/i }).last() + const generatedKeyDialog = page.getByRole('dialog', { name: API_KEY_DIALOG_NAME }).last() await generatedKeyDialog.getByRole('button', { name: 'OK' }).click() await expect(page.getByText('Keep this key in a secure and accessible place.')).not.toBeVisible() @@ -135,7 +134,7 @@ Then( const fullSecret = this.agentBuilder.accessPoint.generatedApiKey if (!fullSecret) throw new Error('No generated Agent v2 API key found.') - const apiKeyDialog = this.getPage().getByRole('dialog', { name: /API Secret key/i }) + const apiKeyDialog = this.getPage().getByRole('dialog', { name: API_KEY_DIALOG_NAME }) await expect(apiKeyDialog).toBeVisible() await expect(apiKeyDialog.getByText(fullSecret, { exact: true })).not.toBeVisible() diff --git a/e2e/features/step-definitions/agent-v2/access-point-web-app.steps.ts b/e2e/features/step-definitions/agent-v2/access-point-web-app.steps.ts index 9182bc7ad5a..53eba93fdae 100644 --- a/e2e/features/step-definitions/agent-v2/access-point-web-app.steps.ts +++ b/e2e/features/step-definitions/agent-v2/access-point-web-app.steps.ts @@ -2,6 +2,7 @@ import type { Page } from '@playwright/test' import type { DifyWorld } from '../../support/world' import { Then, When } from '@cucumber/cucumber' import { expect } from '@playwright/test' +import { getAgentWebAppURL } from '../../agent-v2/support/access-point' import { agentBuilderExpectedTokens } from '../../agent-v2/support/agent-builder-resources' import { getCurrentAgentId, getDialog, getWebAppCard } from './access-point-helpers' @@ -21,7 +22,7 @@ Then('I should see the Agent v2 Web app access URL', async function (this: DifyW const webAppCard = getWebAppCard(this) await expect(webAppCard.getByRole('heading', { name: 'Web app' })).toBeVisible() - await expect(webAppCard.getByText('Access URL')).toBeVisible() + await expect(webAppCard.getByText('Web App URL')).toBeVisible() await expect(webAppCard.getByLabel('Copy access URL')).toBeEnabled() await expect(webAppCard.getByRole('link', { name: 'Launch' })).toBeVisible() }) @@ -48,8 +49,9 @@ When('I launch the Agent v2 Web app', async function (this: DifyWorld) { }) When('I open the Agent v2 Web app URL', async function (this: DifyWorld) { - const webAppURL = this.agentBuilder.accessPoint.webAppURL - if (!webAppURL) throw new Error('No Agent v2 Web app URL was recorded.') + const agentId = getCurrentAgentId(this) + const agent = await this.getConsoleClient().agent.byAgentId.get({ params: { agent_id: agentId } }) + const webAppURL = this.agentBuilder.accessPoint.webAppURL ?? getAgentWebAppURL(agent) if (!this.context) throw new Error('Playwright browser context has not been initialized.') const webAppPage = await this.context.newPage() diff --git a/e2e/features/step-definitions/agent-v2/access-point.steps.ts b/e2e/features/step-definitions/agent-v2/access-point.steps.ts index 79b0f06aa40..1490f252cfa 100644 --- a/e2e/features/step-definitions/agent-v2/access-point.steps.ts +++ b/e2e/features/step-definitions/agent-v2/access-point.steps.ts @@ -2,38 +2,22 @@ import type { DifyWorld } from '../../support/world' import type { AccessSurfaceName } from './access-point-helpers' import { Given, Then, When } from '@cucumber/cucumber' import { expect } from '@playwright/test' -import { enableAgentWebApp } from '../../agent-v2/support/access-point' import { publishAgentWithPublishableDraft } from '../../agent-v2/support/agent' import { - getAccessRegion, getAccessSurfaceCard, getCurrentAgentId, getPreseededResource, + getServiceApiCard, + getWebAppCard, } from './access-point-helpers' Given('the Agent v2 draft has been published via API', async function (this: DifyWorld) { await publishAgentWithPublishableDraft(this.getConsoleClient(), getCurrentAgentId(this)) }) -Given( - /^Agent v2 (Web app|Backend service API) access has been enabled via API$/, - async function (this: DifyWorld, surface: AccessSurfaceName) { - if (surface === 'Web app') { - this.agentBuilder.accessPoint.webAppURL = await enableAgentWebApp( - this.getConsoleClient(), - getCurrentAgentId(this), - ) - return - } - - const agentId = getCurrentAgentId(this) - const apiAccess = await this.getConsoleClient().agent.byAgentId.apiEnable.post({ - body: { enable_api: true }, - params: { agent_id: agentId }, - }) - this.agentBuilder.accessPoint.serviceApiBaseURL = apiAccess.service_api_base_url - }, -) +When('I republish the Agent v2 draft via API', async function (this: DifyWorld) { + await publishAgentWithPublishableDraft(this.getConsoleClient(), getCurrentAgentId(this)) +}) When( 'I open the preseeded Agent v2 Access Point page for {string} from the Agent Roster', @@ -61,34 +45,20 @@ When('I switch to the Agent v2 Access Point section', async function (this: Dify await expect(page.getByRole('region', { name: 'Access Point' })).toBeVisible() }) -Then('I should see the Agent v2 Access Point overview', async function (this: DifyWorld) { - const accessRegion = getAccessRegion(this) +Then( + 'the unpublished Agent v2 access surfaces should be unavailable', + async function (this: DifyWorld) { + const webAppCard = getWebAppCard(this) + const serviceApiCard = getServiceApiCard(this) - await expect(accessRegion).toBeVisible({ timeout: 30_000 }) - await expect(accessRegion.getByRole('heading', { name: 'Access Point' })).toBeVisible() - await expect(accessRegion.getByRole('heading', { name: 'Web app' })).toBeVisible() - await expect(accessRegion.getByText('Access URL')).toBeVisible() - await expect(accessRegion.getByLabel('Copy access URL')).toBeVisible() - await expect(accessRegion.getByLabel('Toggle Web app access')).toBeVisible() - await expect(accessRegion.getByRole('link', { name: 'Launch' })).toBeVisible() - await expect(accessRegion.getByRole('button', { name: 'Embedded' })).toBeVisible() - await expect(accessRegion.getByRole('button', { name: 'Custom Frontend' })).toBeVisible() - await expect(accessRegion.getByRole('button', { name: 'Branding' })).toBeVisible() - await expect(accessRegion.getByRole('heading', { name: 'Backend service API' })).toBeVisible() - await expect(accessRegion.getByText('Service API Endpoint')).toBeVisible() - await expect(accessRegion.getByLabel('Copy service API endpoint')).toBeVisible() - await expect(accessRegion.getByLabel('Toggle Backend service API access')).toBeVisible() - await expect(accessRegion.getByRole('button', { name: /^API Key\b/ })).toBeVisible() - await expect(accessRegion.getByRole('link', { name: 'API Reference' })).toBeVisible() - await expect(accessRegion.getByText(/^(?:In|Out of) service$/i)).toHaveCount(2) - await expect(accessRegion.getByRole('heading', { name: 'Workflow access' })).toBeVisible() - await expect(accessRegion.getByRole('columnheader', { name: 'Name' })).toBeVisible() - await expect(accessRegion.getByRole('columnheader', { name: 'Version' })).toBeVisible() - await expect(accessRegion.getByRole('columnheader', { name: 'Nodes' })).toBeVisible() - await expect(accessRegion.getByRole('columnheader', { name: 'Last updated' })).toBeVisible() - await expect(accessRegion.getByRole('columnheader', { name: 'Actions' })).toBeVisible() - await expect(accessRegion.getByText('No workflow references yet.')).toBeVisible() -}) + await expect(webAppCard.getByText('Out of service')).toBeVisible({ timeout: 30_000 }) + await expect(webAppCard.getByLabel('Toggle Web app access')).toBeDisabled() + await expect(webAppCard.getByRole('button', { name: 'Launch' })).toBeDisabled() + await expect(serviceApiCard.getByText('Out of service')).toBeVisible() + await expect(serviceApiCard.getByLabel('Toggle Backend service API access')).toBeDisabled() + await expect(serviceApiCard.getByRole('button', { name: /^API Key\b/ })).toBeDisabled() + }, +) When( /^I disable Agent v2 (Web app|Backend service API) access$/, diff --git a/e2e/features/step-definitions/agent-v2/agent-edit.steps.ts b/e2e/features/step-definitions/agent-v2/agent-edit.steps.ts index bce34d4ff54..698784b11f1 100644 --- a/e2e/features/step-definitions/agent-v2/agent-edit.steps.ts +++ b/e2e/features/step-definitions/agent-v2/agent-edit.steps.ts @@ -31,6 +31,7 @@ const getComposerInheritanceSnapshot = async (world: DifyWorld, agentId: string) const knowledgeSets = asArray(asRecord(soul.knowledge).sets) return { + activeConfigIsPublished: draft.active_config_is_published, fileNames: files .map((file) => asString(asRecord(file).name)) .filter(Boolean) @@ -178,7 +179,7 @@ Then('I should see the Agent v2 full-config fixture sections', async function (t }) Then( - 'the duplicated Agent v2 should inherit the full-config fixture from {string}', + 'the duplicated Agent v2 should inherit the full-config fixture from {string} without inheriting its publication state', async function (this: DifyWorld, agentName: string) { const sourceAgent = getPreseededAgent(this, agentName) const duplicatedAgentId = getCurrentAgentId(this) @@ -189,8 +190,7 @@ Then( ) const client = this.getConsoleClient() - const [sourceDetail, duplicatedDetail, sourceSnapshot, duplicatedSnapshot] = await Promise.all([ - client.agent.byAgentId.get({ params: { agent_id: sourceAgent.id } }), + const [duplicatedDetail, sourceSnapshot, duplicatedSnapshot] = await Promise.all([ client.agent.byAgentId.get({ params: { agent_id: duplicatedAgentId } }), getComposerInheritanceSnapshot(this, sourceAgent.id), getComposerInheritanceSnapshot(this, duplicatedAgentId), @@ -198,9 +198,8 @@ Then( expect(duplicatedDetail.id).toBe(duplicatedAgentId) expect(duplicatedDetail.name).toBe(this.lastCreatedAgentName) - expect(duplicatedDetail.active_config_is_published).toBe( - sourceDetail.active_config_is_published, - ) + expect(sourceSnapshot.activeConfigIsPublished).toBe(true) + expect(duplicatedSnapshot.activeConfigIsPublished).toBe(false) expect(duplicatedSnapshot.model).toEqual({ name: stableModel.name, provider: stableModel.provider, @@ -249,11 +248,12 @@ Then('I should see the Agent v2 tool state fixture tools', async function (this: const toolsSection = page.getByRole('region', { name: 'Tools' }) await expect(toolsSection).toBeVisible({ timeout: 30_000 }) - await expect( - toolsSection.getByRole('button', { exact: true, name: 'Not authorized' }), - ).toHaveCount(2) - const { action: jsonReplaceAction, tool: jsonTool } = await expectProviderToolActionVisible( + const { + action: jsonReplaceAction, + provider: jsonProvider, + tool: jsonTool, + } = await expectProviderToolActionVisible( toolsSection, agentBuilderPreseededResources.jsonReplaceTool, ) @@ -271,10 +271,16 @@ Then('I should see the Agent v2 tool state fixture tools', async function (this: }), ).toBeVisible() - await expectProviderToolActionVisible( + const { provider: tavilyProvider } = await expectProviderToolActionVisible( toolsSection, agentBuilderPreseededResources.tavilySearchTool, ) + await expect( + tavilyProvider.locator('..').getByRole('button', { exact: true, name: 'Not authorized' }), + ).toBeVisible() + await expect( + jsonProvider.locator('..').getByRole('button', { exact: true, name: 'Not authorized' }), + ).toHaveCount(0) }) Then('I should see the Agent v2 dual retrieval fixture settings', async function (this: DifyWorld) { diff --git a/e2e/features/step-definitions/agent-v2/configure-helpers.ts b/e2e/features/step-definitions/agent-v2/configure-helpers.ts index da76a65d44c..34faf827856 100644 --- a/e2e/features/step-definitions/agent-v2/configure-helpers.ts +++ b/e2e/features/step-definitions/agent-v2/configure-helpers.ts @@ -276,7 +276,7 @@ export const expectProviderToolActionVisible = async ( const action = toolsSection.getByText(tool.actionName, { exact: true }) await expect(action).toBeVisible() - return { action, tool } + return { action, provider, tool } } export const openAgentKnowledgeRetrievalDialog = async ( diff --git a/e2e/features/step-definitions/agent-v2/publish.steps.ts b/e2e/features/step-definitions/agent-v2/publish.steps.ts index dcfaaf69e5f..db47cf0225b 100644 --- a/e2e/features/step-definitions/agent-v2/publish.steps.ts +++ b/e2e/features/step-definitions/agent-v2/publish.steps.ts @@ -32,10 +32,10 @@ Then('the Agent v2 draft should remain unpublished', async function (this: DifyW .poll( async () => { const agentId = getCurrentAgentId(this) - const agent = await this.getConsoleClient().agent.byAgentId.get({ + const composer = await this.getConsoleClient().agent.byAgentId.composer.get({ params: { agent_id: agentId }, }) - return agent.active_config_is_published + return composer.active_config_is_published }, { timeout: 30_000 }, ) @@ -55,10 +55,10 @@ Then('the Agent v2 draft should be published and up to date', async function (th await expect(page.getByText('Up to date')).toBeVisible() await expect .poll(async () => { - const agent = await this.getConsoleClient().agent.byAgentId.get({ + const composer = await this.getConsoleClient().agent.byAgentId.composer.get({ params: { agent_id: agentId }, }) - return agent.active_config_is_published + return composer.active_config_is_published }) .toBe(true) }) diff --git a/e2e/features/step-definitions/agent-v2/speech-to-text.steps.ts b/e2e/features/step-definitions/agent-v2/speech-to-text.steps.ts index b0a0d36afbc..07187edf69d 100644 --- a/e2e/features/step-definitions/agent-v2/speech-to-text.steps.ts +++ b/e2e/features/step-definitions/agent-v2/speech-to-text.steps.ts @@ -50,10 +50,9 @@ When( const page = this.getPage() const agentId = getCurrentAgentId(this) - await expect(page.getByTestId('voice-input-timer')).toHaveText( - voiceInputTestMaterial.recordingDuration, - { timeout: 15_000 }, - ) + await expect(page.getByRole('timer')).toHaveText(voiceInputTestMaterial.recordingDuration, { + timeout: 15_000, + }) const responsePromise = page.waitForResponse( (response) => diff --git a/e2e/features/step-definitions/apps/app-detail-navigation.steps.ts b/e2e/features/step-definitions/apps/app-detail-navigation.steps.ts index 743aafbc9f1..6110d2f4305 100644 --- a/e2e/features/step-definitions/apps/app-detail-navigation.steps.ts +++ b/e2e/features/step-definitions/apps/app-detail-navigation.steps.ts @@ -1,7 +1,7 @@ import type { DifyWorld } from '../../support/world' import { When } from '@cucumber/cucumber' -When('I navigate to the app overview page', async function (this: DifyWorld) { +When('I navigate to the app access point page', async function (this: DifyWorld) { const appId = this.createdAppIds.at(-1) - await this.getPage().goto(`/app/${appId}/overview`) + await this.getPage().goto(`/app/${appId}/access-point`) }) diff --git a/e2e/features/step-definitions/apps/publish-app.steps.ts b/e2e/features/step-definitions/apps/publish-app.steps.ts index aca00b220e4..9b7c402a065 100644 --- a/e2e/features/step-definitions/apps/publish-app.steps.ts +++ b/e2e/features/step-definitions/apps/publish-app.steps.ts @@ -7,9 +7,8 @@ When('I open the publish panel', async function (this: DifyWorld) { }) When('I publish the app', async function (this: DifyWorld) { - await this.getPage() - .getByRole('button', { name: /Publish Update/ }) - .click() + const publishPanel = this.getPage().getByRole('dialog') + await publishPanel.getByRole('button', { name: 'Publish', exact: true }).click() }) Then('the app should be marked as published', async function (this: DifyWorld) { diff --git a/e2e/features/step-definitions/apps/web-app-service.steps.ts b/e2e/features/step-definitions/apps/web-app-service.steps.ts index d7f9d39d2e9..8bd2717a62e 100644 --- a/e2e/features/step-definitions/apps/web-app-service.steps.ts +++ b/e2e/features/step-definitions/apps/web-app-service.steps.ts @@ -23,15 +23,6 @@ Given('a new runnable workflow app has been published', async function (this: Di this.shareURL = getAppSiteURL(appDetail) }) -When('I open the app information panel', async function (this: DifyWorld) { - const appName = this.lastCreatedAppName - if (!appName) { - throw new Error('No app name available. Create an app before opening its information panel.') - } - - await this.getPage().getByRole('button', { name: appName }).click() -}) - const getWebAppSwitch = (world: DifyWorld) => { const webAppCard = world.getPage().getByRole('region', { name: 'Web App' }) return webAppCard.getByRole('switch', { name: 'Web App' }) @@ -73,7 +64,7 @@ When('I enable the Web App', async function (this: DifyWorld) { Then('the Web App should be in service', async function (this: DifyWorld) { const webAppCard = this.getPage().getByRole('region', { name: 'Web App' }) - await expect(webAppCard.getByText('In Service', { exact: true })).toBeVisible({ + await expect(webAppCard.getByText(/^In service$/i)).toBeVisible({ timeout: 10_000, }) }) diff --git a/e2e/package.json b/e2e/package.json index 562cfb6a46c..f936f76f383 100644 --- a/e2e/package.json +++ b/e2e/package.json @@ -4,22 +4,30 @@ "type": "module", "scripts": { "e2e": "tsx ./scripts/run-cucumber.ts", + "e2e:accessibility": "pnpm run e2e:accessibility:aa", + "e2e:accessibility:a": "tsx ./scripts/run-cucumber.ts --full -- --tags \"@axe and @wcag-a\"", + "e2e:accessibility:aa": "tsx ./scripts/run-cucumber.ts --full -- --tags \"@axe and @wcag-aa\"", "e2e:external": "tsx ./scripts/run-external-runtime.ts", - "e2e:external:prepare": "tsx ./scripts/prepare-external-runtime.ts", + "e2e:external:prepare": "tsx ./scripts/run-cucumber.ts --seed-only --profile external-runtime", "e2e:full": "tsx ./scripts/run-cucumber.ts --full", "e2e:full:headed": "tsx ./scripts/run-cucumber.ts --full --headed", "e2e:headed": "tsx ./scripts/run-cucumber.ts --headed", "e2e:install": "playwright install --with-deps chromium webkit", + "e2e:install:ci": "playwright install --with-deps --only-shell chromium webkit", + "e2e:install:ci:chromium": "playwright install --with-deps --only-shell chromium", "e2e:middleware:down": "tsx ./scripts/setup.ts middleware-down", "e2e:middleware:up": "tsx ./scripts/setup.ts middleware-up", "e2e:post-merge": "tsx ./scripts/run-post-merge.ts", - "e2e:post-merge:prepare": "tsx ./scripts/seed.ts --pack agent-v2 --profile post-merge", + "e2e:post-merge:prepare": "tsx ./scripts/run-cucumber.ts --seed-only --profile post-merge", + "e2e:prepared": "tsx ./scripts/run-prepared.ts", + "e2e:prepared:prepare": "tsx ./scripts/run-cucumber.ts --seed-only --profile prepared", "e2e:reset": "tsx ./scripts/setup.ts reset", - "seed": "tsx ./scripts/seed.ts", + "seed": "tsx ./scripts/run-cucumber.ts --seed-only", "test:unit": "vitest run", "type-check": "tsc" }, "devDependencies": { + "@axe-core/playwright": "catalog:", "@cucumber/cucumber": "catalog:", "@dify/contracts": "workspace:*", "@dify/tsconfig": "workspace:*", diff --git a/e2e/scripts/prepare-external-runtime.ts b/e2e/scripts/prepare-external-runtime.ts deleted file mode 100644 index a550fca5d5b..00000000000 --- a/e2e/scripts/prepare-external-runtime.ts +++ /dev/null @@ -1,17 +0,0 @@ -import { e2eDir, isMainModule, runCommandOrThrow } from './common' -import './env-register' - -const main = async () => { - await runCommandOrThrow({ - command: 'npx', - args: ['tsx', './scripts/seed.ts', '--pack', 'agent-v2', '--profile', 'external-runtime'], - cwd: e2eDir, - }) -} - -if (isMainModule(import.meta.url)) { - void main().catch((error) => { - console.error(error instanceof Error ? error.message : String(error)) - process.exit(1) - }) -} diff --git a/e2e/scripts/run-cucumber.ts b/e2e/scripts/run-cucumber.ts index 38cc22e6dfa..74f93b6d2b3 100644 --- a/e2e/scripts/run-cucumber.ts +++ b/e2e/scripts/run-cucumber.ts @@ -7,56 +7,17 @@ import { startLoggedProcess, stopManagedProcess, waitForUrl } from '../support/p import { startWebServer, stopWebServer } from '../support/web-server' import { apiURL, baseURL, reuseExistingWebServer } from '../test-env' import { e2eDir, isMainModule, runCommand } from './common' +import { parseRunOptions, shouldStartManagedAgentBackend } from './run-options' +import { runSeed } from './seed-runner' import { resetState, startMiddleware, stopMiddleware } from './setup' import './env-register' -type RunOptions = { - forwardArgs: string[] - full: boolean - headed: boolean -} - -const parseArgs = (argv: string[]): RunOptions => { - let full = false - let headed = false - const forwardArgs: string[] = [] - - for (const [index, arg] of argv.entries()) { - if (arg === '--') { - forwardArgs.push(...argv.slice(index + 1)) - return { forwardArgs, full, headed } - } - - if (arg === '--full') { - full = true - continue - } - - if (arg === '--headed') { - headed = true - continue - } - - forwardArgs.push(arg) - } - - return { forwardArgs, full, headed } -} - const hasCustomTags = (forwardArgs: string[]) => forwardArgs.some((arg) => arg === '--tags' || arg.startsWith('--tags=')) -const fullNonExternalTags = 'not @prepared and not @external-model and not @external-tool' - -const isTruthyEnv = (value: string | undefined) => value === '1' || value === 'true' - -const shouldStartAgentBackend = () => { - if (isTruthyEnv(process.env.E2E_START_AGENT_BACKEND)) return true - - if (process.env.E2E_AGENT_BACKEND_URL || process.env.AGENT_BACKEND_BASE_URL) return false - - return false -} +const fullNonExternalTags = + 'not @axe and not @prepared and not @external-model and not @external-tool' +const seedCeleryQueues = 'dataset,priority_dataset,workflow_based_app_execution' const readLogTail = async (logFilePath: string) => { const content = await readFile(logFilePath, 'utf8').catch(() => '') @@ -87,65 +48,41 @@ const waitForUnexpectedProcessExit = async ( throw new Error(`${label} exited before becoming ready. See ${logFilePath}.${logTailMessage}`) } +const waitForManagedProcess = async ({ + errorMessage, + managedProcess, + url, +}: { + errorMessage: string + managedProcess: ManagedProcess + url: string +}) => { + let waiting = true + try { + await Promise.race([ + waitForUrl(url, 180_000, 1_000), + waitForUnexpectedProcessExit(managedProcess, () => !waiting), + ]) + } catch (error) { + if (error instanceof Error && error.message.includes('exited before becoming ready')) + throw error + + throw new Error(`${errorMessage} See ${managedProcess.logFilePath}.`) + } finally { + waiting = false + } +} + const main = async () => { - const { forwardArgs, full, headed } = parseArgs(process.argv.slice(2)) - const startMiddlewareForRun = full - const resetStateForRun = full - const startAgentBackendForRun = shouldStartAgentBackend() - - if (resetStateForRun) await resetState() - - if (startMiddlewareForRun) await startMiddleware() - + const { forwardArgs, full, headed, seed, seedOnly } = parseRunOptions(process.argv.slice(2)) + const startAgentBackendForRun = shouldStartManagedAgentBackend() const cucumberReportDir = path.join(e2eDir, 'cucumber-report') const logDir = path.join(e2eDir, '.logs') - - await rm(cucumberReportDir, { force: true, recursive: true }) - await mkdir(logDir, { recursive: true }) - - const shellctlProcess = startAgentBackendForRun - ? await startLoggedProcess({ - command: 'npx', - args: ['tsx', './scripts/setup.ts', 'shellctl-sandbox'], - cwd: e2eDir, - label: 'shellctl sandbox', - logFilePath: path.join(logDir, 'cucumber-shellctl-sandbox.log'), - }) - : undefined - - const difyAgentProcess = startAgentBackendForRun - ? await startLoggedProcess({ - command: 'npx', - args: ['tsx', './scripts/setup.ts', 'agent-backend'], - cwd: e2eDir, - env: { - E2E_START_AGENT_BACKEND: '1', - }, - label: 'agent backend', - logFilePath: path.join(logDir, 'cucumber-agent-backend.log'), - }) - : undefined - - const apiProcess = await startLoggedProcess({ - command: 'npx', - args: ['tsx', './scripts/setup.ts', 'api'], - cwd: e2eDir, - env: startAgentBackendForRun - ? { - E2E_START_AGENT_BACKEND: '1', - } - : undefined, - label: 'api server', - logFilePath: path.join(logDir, 'cucumber-api.log'), - }) - - const celeryProcess = await startLoggedProcess({ - command: 'npx', - args: ['tsx', './scripts/setup.ts', 'celery'], - cwd: e2eDir, - label: 'celery worker', - logFilePath: path.join(logDir, 'cucumber-celery.log'), - }) + let apiProcess: ManagedProcess | undefined + let celeryProcess: ManagedProcess | undefined + let difyAgentProcess: ManagedProcess | undefined + let middlewareStarted = false + let shellctlProcess: ManagedProcess | undefined let cleanupPromise: Promise | undefined const cleanup = async () => { @@ -157,7 +94,7 @@ const main = async () => { { label: 'Stop API server', run: () => stopManagedProcess(apiProcess) }, { label: 'Stop agent backend', run: () => stopManagedProcess(difyAgentProcess) }, { label: 'Stop shellctl sandbox', run: () => stopManagedProcess(shellctlProcess) }, - ...(startMiddlewareForRun ? [{ label: 'Stop middleware', run: stopMiddleware }] : []), + ...(middlewareStarted ? [{ label: 'Stop middleware', run: stopMiddleware }] : []), ]) if (cleanupErrors.length > 0) @@ -182,60 +119,73 @@ const main = async () => { process.once('SIGTERM', onTerminate) try { - if (shellctlProcess) { - let waitingForShellctl = true - try { - const shellctlPort = process.env.E2E_SHELLCTL_PORT || '5004' - await Promise.race([ - waitForUrl(`http://127.0.0.1:${shellctlPort}/healthz`, 180_000, 1_000), - waitForUnexpectedProcessExit(shellctlProcess, () => !waitingForShellctl), - ]) - } catch (error) { - if (error instanceof Error && error.message.includes('exited before becoming ready')) - throw error + if (full) await resetState() - throw new Error( - `Shellctl sandbox did not become ready. See ${shellctlProcess.logFilePath}.`, - ) - } finally { - waitingForShellctl = false - } + if (full) { + middlewareStarted = true + await startMiddleware() } - if (difyAgentProcess) { - let waitingForAgentBackend = true - try { - const agentBackendPort = process.env.E2E_AGENT_BACKEND_PORT || '5050' - await Promise.race([ - waitForUrl(`http://127.0.0.1:${agentBackendPort}/openapi.json`, 180_000, 1_000), - waitForUnexpectedProcessExit(difyAgentProcess, () => !waitingForAgentBackend), - ]) - } catch (error) { - if (error instanceof Error && error.message.includes('exited before becoming ready')) - throw error + if (!seedOnly) await rm(cucumberReportDir, { force: true, recursive: true }) + await mkdir(logDir, { recursive: true }) - throw new Error(`Agent backend did not become ready. See ${difyAgentProcess.logFilePath}.`) - } finally { - waitingForAgentBackend = false - } + if (startAgentBackendForRun) { + shellctlProcess = await startLoggedProcess({ + command: 'npx', + args: ['tsx', './scripts/setup.ts', 'shellctl-sandbox'], + cwd: e2eDir, + label: 'shellctl sandbox', + logFilePath: path.join(logDir, 'cucumber-shellctl-sandbox.log'), + }) + const shellctlPort = process.env.E2E_SHELLCTL_PORT || '5004' + await waitForManagedProcess({ + errorMessage: 'Shellctl sandbox did not become ready.', + managedProcess: shellctlProcess, + url: `http://127.0.0.1:${shellctlPort}/healthz`, + }) + + difyAgentProcess = await startLoggedProcess({ + command: 'npx', + args: ['tsx', './scripts/setup.ts', 'agent-backend'], + cwd: e2eDir, + env: { E2E_START_AGENT_BACKEND: '1' }, + label: 'agent backend', + logFilePath: path.join(logDir, 'cucumber-agent-backend.log'), + }) + const agentBackendPort = process.env.E2E_AGENT_BACKEND_PORT || '5050' + await waitForManagedProcess({ + errorMessage: 'Agent backend did not become ready.', + managedProcess: difyAgentProcess, + url: `http://127.0.0.1:${agentBackendPort}/openapi.json`, + }) } - let waitingForApi = true - try { - await Promise.race([ - waitForUrl(`${apiURL}/health`, 180_000, 1_000), - waitForUnexpectedProcessExit(apiProcess, () => !waitingForApi), - ]) - } catch (error) { - if (error instanceof Error && error.message.includes('exited before becoming ready')) - throw error + apiProcess = await startLoggedProcess({ + command: 'npx', + args: ['tsx', './scripts/setup.ts', 'api'], + cwd: e2eDir, + env: startAgentBackendForRun ? { E2E_START_AGENT_BACKEND: '1' } : undefined, + label: 'api server', + logFilePath: path.join(logDir, 'cucumber-api.log'), + }) + await waitForManagedProcess({ + errorMessage: `API did not become ready at ${apiURL}/health.`, + managedProcess: apiProcess, + url: `${apiURL}/health`, + }) - throw new Error( - `API did not become ready at ${apiURL}/health. See ${apiProcess.logFilePath}.`, - ) - } finally { - waitingForApi = false - } + celeryProcess = await startLoggedProcess({ + command: 'npx', + args: [ + 'tsx', + './scripts/setup.ts', + 'celery', + ...(seed ? ['--queues', seedCeleryQueues] : []), + ], + cwd: e2eDir, + label: 'celery worker', + logFilePath: path.join(logDir, 'cucumber-celery.log'), + }) await startWebServer({ baseURL, @@ -247,33 +197,36 @@ const main = async () => { timeoutMs: 300_000, }) - const cucumberEnv: NodeJS.ProcessEnv = { - ...process.env, - CUCUMBER_HEADLESS: headed ? '0' : '1', + if (seed) await runSeed(seed) + + if (!seedOnly) { + const cucumberEnv: NodeJS.ProcessEnv = { + ...process.env, + CUCUMBER_HEADLESS: headed ? '0' : '1', + } + + if (full && !hasCustomTags(forwardArgs)) cucumberEnv.E2E_CUCUMBER_TAGS = fullNonExternalTags + + const result = await runCommand({ + command: 'npx', + args: [ + 'tsx', + './node_modules/@cucumber/cucumber/bin/cucumber.js', + '--config', + './cucumber.config.ts', + ...forwardArgs, + ], + cwd: e2eDir, + env: cucumberEnv, + }) + + if (result.exitCode === 0) { + const messages = await readFile(path.join(cucumberReportDir, 'report.ndjson'), 'utf8') + assertCucumberScenariosStarted(messages) + } + + process.exitCode = result.exitCode } - - if (startMiddlewareForRun && !hasCustomTags(forwardArgs)) - cucumberEnv.E2E_CUCUMBER_TAGS = fullNonExternalTags - - const result = await runCommand({ - command: 'npx', - args: [ - 'tsx', - './node_modules/@cucumber/cucumber/bin/cucumber.js', - '--config', - './cucumber.config.ts', - ...forwardArgs, - ], - cwd: e2eDir, - env: cucumberEnv, - }) - - if (result.exitCode === 0) { - const messages = await readFile(path.join(cucumberReportDir, 'report.ndjson'), 'utf8') - assertCucumberScenariosStarted(messages) - } - - process.exitCode = result.exitCode } finally { process.off('SIGINT', onTerminate) process.off('SIGTERM', onTerminate) diff --git a/e2e/scripts/run-external-runtime.ts b/e2e/scripts/run-external-runtime.ts index 13adf66f221..20c0435abe0 100644 --- a/e2e/scripts/run-external-runtime.ts +++ b/e2e/scripts/run-external-runtime.ts @@ -6,7 +6,16 @@ const defaultExternalRuntimeTags = '@external-model or @external-tool' const main = async () => { await runForegroundProcess({ command: 'npx', - args: ['tsx', './scripts/run-cucumber.ts', '--', '--tags', defaultExternalRuntimeTags], + args: [ + 'tsx', + './scripts/run-cucumber.ts', + '--full', + '--profile', + 'external-runtime', + '--', + '--tags', + defaultExternalRuntimeTags, + ], cwd: e2eDir, }) } diff --git a/e2e/scripts/run-options.ts b/e2e/scripts/run-options.ts new file mode 100644 index 00000000000..16b5c0a2deb --- /dev/null +++ b/e2e/scripts/run-options.ts @@ -0,0 +1,122 @@ +import type { SeedOptions } from './seed-runner' + +export type RunOptions = { + forwardArgs: string[] + full: boolean + headed: boolean + seed?: SeedOptions + seedOnly: boolean +} + +const readOptionValue = (argv: string[], index: number, option: string) => { + const value = argv[index + 1] + if (!value || value.startsWith('--')) throw new Error(`${option} requires a value.`) + + return value +} + +export const parseRunOptions = (argv: string[]): RunOptions => { + let allowBlocked = false + let dryRun = false + let full = false + let headed = false + let pack = 'agent-v2' + let profile: string | undefined + let seedOnly = false + const forwardArgs: string[] = [] + + for (let index = 0; index < argv.length; index += 1) { + const arg = argv[index]! + + if (arg === '--') { + forwardArgs.push(...argv.slice(index + 1)) + break + } + + if (arg === '--full') { + full = true + continue + } + + if (arg === '--headed') { + headed = true + continue + } + + if (arg === '--seed-only') { + seedOnly = true + continue + } + + if (arg === '--allow-blocked') { + allowBlocked = true + continue + } + + if (arg === '--dry-run') { + dryRun = true + continue + } + + if (arg === '--pack') { + pack = readOptionValue(argv, index, '--pack') + index += 1 + continue + } + + if (arg.startsWith('--pack=')) { + pack = arg.slice('--pack='.length) + if (!pack) throw new Error('--pack requires a value.') + continue + } + + if (arg === '--profile') { + profile = readOptionValue(argv, index, '--profile') + index += 1 + continue + } + + if (arg.startsWith('--profile=')) { + profile = arg.slice('--profile='.length) + if (!profile) throw new Error('--profile requires a value.') + continue + } + + forwardArgs.push(arg) + } + + const shouldSeed = seedOnly || profile !== undefined + if (!shouldSeed && (allowBlocked || dryRun || pack !== 'agent-v2')) + throw new Error('Seed options require --seed-only or --profile.') + if (dryRun && !seedOnly) throw new Error('--dry-run requires --seed-only.') + + return { + forwardArgs, + full, + headed, + seed: shouldSeed + ? { + allowBlocked, + dryRun, + pack, + profile: profile ?? 'post-merge', + } + : undefined, + seedOnly, + } +} + +const isTruthyEnv = (value: string | undefined) => value === '1' || value === 'true' + +export const shouldStartManagedAgentBackend = (env: NodeJS.ProcessEnv = process.env) => { + const shouldStart = isTruthyEnv(env.E2E_START_AGENT_BACKEND) + const externalUrl = env.E2E_AGENT_BACKEND_URL?.trim() || env.AGENT_BACKEND_BASE_URL?.trim() + + if (shouldStart && externalUrl) { + throw new Error( + 'E2E_START_AGENT_BACKEND cannot be enabled when E2E_AGENT_BACKEND_URL or AGENT_BACKEND_BASE_URL is set.', + ) + } + + return shouldStart +} diff --git a/e2e/scripts/run-post-merge.ts b/e2e/scripts/run-post-merge.ts index c53e4ee7be1..25e92bc1e98 100644 --- a/e2e/scripts/run-post-merge.ts +++ b/e2e/scripts/run-post-merge.ts @@ -6,7 +6,16 @@ const postMergeTags = '@prepared or @external-model or @external-tool' const main = async () => { await runForegroundProcess({ command: 'npx', - args: ['tsx', './scripts/run-cucumber.ts', '--', '--tags', postMergeTags], + args: [ + 'tsx', + './scripts/run-cucumber.ts', + '--full', + '--profile', + 'post-merge', + '--', + '--tags', + postMergeTags, + ], cwd: e2eDir, }) } diff --git a/e2e/scripts/run-prepared.ts b/e2e/scripts/run-prepared.ts new file mode 100644 index 00000000000..3cd641575e4 --- /dev/null +++ b/e2e/scripts/run-prepared.ts @@ -0,0 +1,23 @@ +import { e2eDir, isMainModule, runForegroundProcess } from './common' +import './env-register' + +const preparedTags = '@prepared' + +const main = async () => { + await runForegroundProcess({ + command: 'npx', + args: [ + 'tsx', + './scripts/run-cucumber.ts', + '--full', + '--profile', + 'prepared', + '--', + '--tags', + preparedTags, + ], + cwd: e2eDir, + }) +} + +if (isMainModule(import.meta.url)) void main() diff --git a/e2e/scripts/seed-runner.ts b/e2e/scripts/seed-runner.ts new file mode 100644 index 00000000000..b6cb013029e --- /dev/null +++ b/e2e/scripts/seed-runner.ts @@ -0,0 +1,54 @@ +import { chromium } from '@playwright/test' +import { createAgentV2SeedTasks } from '../features/agent-v2/support/seed' +import { ensureAuthenticatedState } from '../fixtures/auth' +import { createStandaloneConsoleSession } from '../support/api/console-session' +import { runSeedTasks, writeSeedReport } from '../support/seed' +import { baseURL } from '../test-env' + +export type SeedOptions = { + allowBlocked: boolean + dryRun: boolean + pack: string + profile: string +} + +const getTasks = (pack: string, profile: string) => { + if (pack === 'agent-v2') return createAgentV2SeedTasks(profile) + + throw new Error(`Unknown seed pack "${pack}".`) +} + +const ensureAuth = async () => { + const browser = await chromium.launch({ headless: true }) + try { + await ensureAuthenticatedState(browser, baseURL) + } finally { + await browser.close() + } +} + +export const runSeed = async ({ allowBlocked, dryRun, pack, profile }: SeedOptions) => { + console.warn(`[seed] bootstrapping auth state against ${baseURL}`) + await ensureAuth() + + const consoleSession = await createStandaloneConsoleSession() + try { + const results = await runSeedTasks(getTasks(pack, profile), { + consoleClient: consoleSession.client, + dryRun, + resources: new Map(), + }) + const reportName = `${pack}-${profile}` + const reportPath = await writeSeedReport(reportName, results) + const blockedCount = results.filter((result) => result.status === 'blocked').length + + console.warn(`[seed] report ${reportPath}`) + if (blockedCount > 0 && !allowBlocked) { + throw new Error( + `${blockedCount} seed task${blockedCount === 1 ? '' : 's'} blocked. Re-run with --allow-blocked only when partial readiness is intentional.`, + ) + } + } finally { + await consoleSession.dispose() + } +} diff --git a/e2e/scripts/seed.ts b/e2e/scripts/seed.ts deleted file mode 100644 index 9b1b13dc5c6..00000000000 --- a/e2e/scripts/seed.ts +++ /dev/null @@ -1,163 +0,0 @@ -import type { ManagedProcess } from '../support/process' -import { mkdir } from 'node:fs/promises' -import path from 'node:path' -import { chromium } from '@playwright/test' -import { createAgentV2SeedTasks } from '../features/agent-v2/support/seed' -import { ensureAuthenticatedState } from '../fixtures/auth' -import { createStandaloneConsoleSession } from '../support/api/console-session' -import { startLoggedProcess, stopManagedProcess, waitForUrl } from '../support/process' -import { runSeedTasks, writeSeedReport } from '../support/seed' -import { startWebServer, stopWebServer } from '../support/web-server' -import { apiURL, baseURL, reuseExistingWebServer } from '../test-env' -import { e2eDir, isMainModule } from './common' -import './env-register' - -type SeedOptions = { - allowBlocked: boolean - dryRun: boolean - pack: string - profile: string -} - -const parseArgs = (argv: string[]): SeedOptions => { - const options: SeedOptions = { - allowBlocked: false, - dryRun: false, - pack: 'agent-v2', - profile: 'post-merge', - } - - for (const [index, arg] of argv.entries()) { - if (arg === '--pack') { - options.pack = argv[index + 1] || options.pack - continue - } - if (arg.startsWith('--pack=')) { - options.pack = arg.slice('--pack='.length) - continue - } - if (arg === '--dry-run') options.dryRun = true - if (arg === '--allow-blocked') options.allowBlocked = true - if (arg === '--profile') { - options.profile = argv[index + 1] || options.profile - continue - } - if (arg.startsWith('--profile=')) { - options.profile = arg.slice('--profile='.length) - continue - } - } - - return options -} - -const getTasks = (pack: string, profile: string) => { - if (pack === 'agent-v2') return createAgentV2SeedTasks(profile) - - throw new Error(`Unknown seed pack "${pack}".`) -} - -const ensureAuth = async () => { - const browser = await chromium.launch({ headless: true }) - try { - await ensureAuthenticatedState(browser, baseURL) - } finally { - await browser.close() - } -} - -const startApiProcess = async (logDir: string) => { - try { - await waitForUrl(`${apiURL}/health`, 1_000, 250, 1_000) - return undefined - } catch { - // Start a local API process below. - } - - const apiProcess = await startLoggedProcess({ - command: 'npx', - args: ['tsx', './scripts/setup.ts', 'api'], - cwd: e2eDir, - label: 'api server', - logFilePath: path.join(logDir, 'seed-api.log'), - }) - - try { - await waitForUrl(`${apiURL}/health`, 180_000, 1_000) - return apiProcess - } catch (error) { - await stopManagedProcess(apiProcess) - throw error - } -} - -const startCeleryProcess = async (logDir: string) => - startLoggedProcess({ - command: 'npx', - args: [ - 'tsx', - './scripts/setup.ts', - 'celery', - '--queues', - 'dataset,priority_dataset,workflow_based_app_execution', - ], - cwd: e2eDir, - label: 'celery worker', - logFilePath: path.join(logDir, 'seed-celery.log'), - }) - -const main = async () => { - const options = parseArgs(process.argv.slice(2)) - const logDir = path.join(e2eDir, '.logs') - let apiProcess: ManagedProcess | undefined - let celeryProcess: ManagedProcess | undefined - let consoleSession: Awaited> | undefined - - await mkdir(logDir, { recursive: true }) - - try { - apiProcess = await startApiProcess(logDir) - celeryProcess = await startCeleryProcess(logDir) - await startWebServer({ - baseURL, - command: 'npx', - args: ['tsx', './scripts/setup.ts', 'web'], - cwd: e2eDir, - logFilePath: path.join(logDir, 'seed-web.log'), - reuseExistingServer: reuseExistingWebServer, - timeoutMs: 300_000, - }) - - console.warn(`[seed] bootstrapping auth state against ${baseURL}`) - await ensureAuth() - consoleSession = await createStandaloneConsoleSession() - - const results = await runSeedTasks(getTasks(options.pack, options.profile), { - consoleClient: consoleSession.client, - dryRun: options.dryRun, - resources: new Map(), - }) - const reportName = `${options.pack}-${options.profile}` - const reportPath = await writeSeedReport(reportName, results) - const blockedCount = results.filter((result) => result.status === 'blocked').length - - console.warn(`[seed] report ${reportPath}`) - if (blockedCount > 0 && !options.allowBlocked) { - throw new Error( - `${blockedCount} seed task${blockedCount === 1 ? '' : 's'} blocked. Re-run with --allow-blocked only when partial readiness is intentional.`, - ) - } - } finally { - await consoleSession?.dispose() - await stopWebServer() - await stopManagedProcess(celeryProcess) - await stopManagedProcess(apiProcess) - } -} - -if (isMainModule(import.meta.url)) { - void main().catch((error) => { - console.error(error instanceof Error ? error.message : String(error)) - process.exit(1) - }) -} diff --git a/e2e/scripts/setup.ts b/e2e/scripts/setup.ts index 3d8a8f5095a..5aa99884560 100644 --- a/e2e/scripts/setup.ts +++ b/e2e/scripts/setup.ts @@ -27,10 +27,12 @@ import { waitForCondition, webDir, } from './common' +import './env-register' const buildIdPath = path.join(webDir, '.next', 'BUILD_ID') const webBuildStampPath = path.join(webDir, '.next', 'e2e-web-build.sha256') -const apiHost = '127.0.0.1' +const apiLoopbackHost = '127.0.0.1' +const apiBindHost = '0.0.0.0' const apiPort = 5001 const agentBackendHost = '127.0.0.1' const agentBackendBindHost = '0.0.0.0' @@ -41,6 +43,7 @@ const shellctlContainerName = process.env.E2E_SHELLCTL_CONTAINER_NAME || 'dify-a const shellctlImage = process.env.E2E_SHELLCTL_IMAGE || 'dify-agent-local-sandbox:e2e' const shellctlUrl = `http://${shellctlHost}:${shellctlPort}` const agentStubApiBaseUrl = `http://host.docker.internal:${agentBackendPort}/agent-stub` +const sandboxFilesBaseUrl = `http://host.docker.internal:${apiPort}` const defaultPluginDaemonKey = 'lYkiYYT6owG+71oLerGzA7GXCgOT++6ovaezWAjpCjf+Sjc3ZtU+qUEi' const defaultInnerApiKeyForPlugin = 'QaHbTe77CtuXmsfyhR7+vRjI/+XbV1AaFy691iy+kGDv2Jvy0/eAh8Y1' const defaultAgentServerSecretKey = 'MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY' @@ -103,7 +106,7 @@ function getAgentBackendBaseUrl() { return undefined } -const getAgentBackendEnvironment = async () => { +export const getAgentBackendEnvironment = async () => { const apiEnv = await getApiEnvironment() const redisPassword = process.env.REDIS_PASSWORD || apiEnv.REDIS_PASSWORD || 'difyai123456' @@ -115,7 +118,9 @@ const getAgentBackendEnvironment = async () => { apiEnv.INNER_API_KEY_FOR_PLUGIN || defaultInnerApiKeyForPlugin, DIFY_AGENT_INNER_API_URL: - process.env.DIFY_AGENT_INNER_API_URL || `http://${apiHost}:${apiPort}`, + process.env.DIFY_AGENT_INNER_API_URL || `http://${apiLoopbackHost}:${apiPort}`, + DIFY_AGENT_SANDBOX_FILES_BASE_URL: + process.env.DIFY_AGENT_SANDBOX_FILES_BASE_URL || sandboxFilesBaseUrl, DIFY_AGENT_SERVER_SECRET_KEY: process.env.DIFY_AGENT_SERVER_SECRET_KEY || defaultAgentServerSecretKey, DIFY_AGENT_STUB_API_BASE_URL: process.env.DIFY_AGENT_STUB_API_BASE_URL || agentStubApiBaseUrl, @@ -303,12 +308,12 @@ export const startWeb = async () => { } export const startApi = async () => { - if (await isTcpPortReachable(apiHost, apiPort)) { + if (await isTcpPortReachable(apiLoopbackHost, apiPort)) { const listenerDescription = await getTcpPortListenerDescription(apiPort) const listenerMessage = listenerDescription ? `\n\nPort listener:\n${listenerDescription}` : '' throw new Error( - `Cannot start the E2E API server because ${apiHost}:${apiPort} is already in use.${listenerMessage}`, + `Cannot start the E2E API server because ${apiLoopbackHost}:${apiPort} is already in use.${listenerMessage}`, ) } @@ -331,7 +336,7 @@ export const startApi = async () => { 'flask', 'run', '--host', - apiHost, + apiBindHost, '--port', String(apiPort), ], diff --git a/e2e/support/agents.ts b/e2e/support/agents.ts new file mode 100644 index 00000000000..d5363f8ca21 --- /dev/null +++ b/e2e/support/agents.ts @@ -0,0 +1,20 @@ +import type { Page } from '@playwright/test' +import { expect } from '@playwright/test' + +const getExpectOptions = (timeout?: number) => (timeout === undefined ? undefined : { timeout }) + +export const waitForAgentsConsole = async (page: Page, timeout?: number) => { + const options = getExpectOptions(timeout) + + await expect(page).toHaveURL(/\/agents(?:\?.*)?$/, options) + await expect(page.getByRole('link', { name: 'Agents', exact: true })).toHaveAttribute( + 'aria-current', + 'page', + options, + ) + await expect(page.getByRole('heading', { name: 'Agents', exact: true })).toBeVisible(options) + + const roster = page.getByRole('region', { name: 'Agent list' }) + await expect(roster).toBeVisible(options) + await expect(roster).not.toHaveAttribute('aria-busy', 'true', options) +} diff --git a/e2e/support/api/workflows.ts b/e2e/support/api/workflows.ts index 00ee6544d4e..b47dbeeb624 100644 --- a/e2e/support/api/workflows.ts +++ b/e2e/support/api/workflows.ts @@ -19,7 +19,6 @@ export async function syncMinimalWorkflowDraft( viewport: { x: 0, y: 0, zoom: 1 }, }, features: {}, - environment_variables: [], conversation_variables: [], } satisfies SyncDraftWorkflowPayload await client.apps.byAppId.workflows.draft.post({ body, params: { app_id: appId } }) @@ -63,7 +62,6 @@ export async function syncRunnableWorkflowDraft( viewport: { x: 0, y: 0, zoom: 1 }, }, features: {}, - environment_variables: [], conversation_variables: [], } satisfies SyncDraftWorkflowPayload await client.apps.byAppId.workflows.draft.post({ body, params: { app_id: appId } }) diff --git a/e2e/support/apps.ts b/e2e/support/apps.ts index 95f0687f647..c26033ee7de 100644 --- a/e2e/support/apps.ts +++ b/e2e/support/apps.ts @@ -8,6 +8,7 @@ export const waitForAppsConsole = async (page: Page, timeout?: number) => { await expect(page).toHaveURL(/\/apps(?:\?.*)?$/, options) await expect(page.getByRole('heading', { name: 'Studio' })).toBeVisible(options) + await expect(page.getByRole('status', { name: 'Loading' })).toBeHidden(options) } export const openBlankAppCreation = async (page: Page) => { diff --git a/e2e/support/home.ts b/e2e/support/home.ts index f6729eb277d..1f37ed184c9 100644 --- a/e2e/support/home.ts +++ b/e2e/support/home.ts @@ -6,10 +6,11 @@ const getExpectOptions = (timeout?: number) => (timeout === undefined ? undefine export const waitForConsoleHome = async (page: Page, timeout?: number) => { const options = getExpectOptions(timeout) - await expect.poll(() => new URL(page.url()).pathname, options).toBe('/') + await expect(page).toHaveURL(/\/(?:\?.*)?$/, options) await expect(page.getByRole('link', { name: 'Home' })).toHaveAttribute( 'aria-current', 'page', options, ) + await expect(page.getByRole('heading', { name: 'Templates', exact: true })).toBeVisible(options) } diff --git a/e2e/support/integrations.ts b/e2e/support/integrations.ts new file mode 100644 index 00000000000..37275a724bf --- /dev/null +++ b/e2e/support/integrations.ts @@ -0,0 +1,16 @@ +import type { Page } from '@playwright/test' +import { expect } from '@playwright/test' + +const getExpectOptions = (timeout?: number) => (timeout === undefined ? undefined : { timeout }) + +export const waitForModelProviderIntegrations = async (page: Page, timeout?: number) => { + const options = getExpectOptions(timeout) + + await expect(page).toHaveURL(/\/integrations\/model-provider(?:\?.*)?$/, options) + await expect(page.getByRole('link', { name: 'Integrations', exact: true })).toHaveAttribute( + 'aria-current', + 'page', + options, + ) + await expect(page.getByRole('region', { name: 'Model Provider' })).toBeVisible(options) +} diff --git a/e2e/support/knowledge.ts b/e2e/support/knowledge.ts new file mode 100644 index 00000000000..0a63bd33970 --- /dev/null +++ b/e2e/support/knowledge.ts @@ -0,0 +1,17 @@ +import type { Page } from '@playwright/test' +import { expect } from '@playwright/test' + +const getExpectOptions = (timeout?: number) => (timeout === undefined ? undefined : { timeout }) + +export const waitForKnowledgeConsole = async (page: Page, timeout?: number) => { + const options = getExpectOptions(timeout) + + await expect(page).toHaveURL(/\/datasets(?:\?.*)?$/, options) + await expect(page.getByRole('link', { name: 'Knowledge', exact: true })).toHaveAttribute( + 'aria-current', + 'page', + options, + ) + await expect(page.getByRole('heading', { name: 'Knowledge', exact: true })).toBeVisible(options) + await expect(page.getByRole('status', { name: 'Loading' })).toBeHidden(options) +} diff --git a/e2e/support/marketplace-plugins.ts b/e2e/support/marketplace-plugins.ts index cb8b837ef11..902371411fb 100644 --- a/e2e/support/marketplace-plugins.ts +++ b/e2e/support/marketplace-plugins.ts @@ -100,7 +100,7 @@ const waitForPluginInstallTask = async ( const getMarketplaceDownloadUrl = (pluginUniqueIdentifier: string) => { const url = new URL( - '/api/v1/plugins/download', + '/api/v1/plugins/download-url', process.env.E2E_MARKETPLACE_API_URL || defaultMarketplaceApiUrl, ) url.searchParams.set('unique_identifier', pluginUniqueIdentifier) diff --git a/e2e/test-env.ts b/e2e/test-env.ts index 0af395909b9..ec2de36f5ec 100644 --- a/e2e/test-env.ts +++ b/e2e/test-env.ts @@ -1,3 +1,5 @@ +import './scripts/env-register' + export const defaultBaseURL = 'http://127.0.0.1:3000' export const defaultApiURL = 'http://127.0.0.1:5001' export const defaultLocale = 'en-US' diff --git a/e2e/tests/console-client.test.ts b/e2e/tests/console-client.test.ts index 40078b08554..ffd1c42fa7b 100644 --- a/e2e/tests/console-client.test.ts +++ b/e2e/tests/console-client.test.ts @@ -31,6 +31,17 @@ const createApiResponse = ({ status: () => status, statusText: () => statusText, text: async () => body, + timing: () => ({ + connectEnd: -1, + connectStart: -1, + domainLookupEnd: -1, + domainLookupStart: -1, + requestStart: -1, + responseEnd: -1, + responseStart: -1, + secureConnectionStart: -1, + startTime: -1, + }), url: () => url, [Symbol.asyncDispose]: async () => {}, } diff --git a/e2e/tests/managed-runtime-environment.test.ts b/e2e/tests/managed-runtime-environment.test.ts new file mode 100644 index 00000000000..88519ef8772 --- /dev/null +++ b/e2e/tests/managed-runtime-environment.test.ts @@ -0,0 +1,26 @@ +import { afterEach, describe, expect, it, vi } from 'vitest' +import { getAgentBackendEnvironment } from '../scripts/setup' + +describe('managed runtime environment', () => { + afterEach(() => { + vi.unstubAllEnvs() + }) + + it('uses separate host and Sandbox API addresses', async () => { + vi.stubEnv('DIFY_AGENT_INNER_API_URL', '') + vi.stubEnv('DIFY_AGENT_SANDBOX_FILES_BASE_URL', '') + + await expect(getAgentBackendEnvironment()).resolves.toMatchObject({ + DIFY_AGENT_INNER_API_URL: 'http://127.0.0.1:5001', + DIFY_AGENT_SANDBOX_FILES_BASE_URL: 'http://host.docker.internal:5001', + }) + }) + + it('preserves an explicit Sandbox-reachable API address', async () => { + vi.stubEnv('DIFY_AGENT_SANDBOX_FILES_BASE_URL', 'https://dify.example.test/base') + + await expect(getAgentBackendEnvironment()).resolves.toMatchObject({ + DIFY_AGENT_SANDBOX_FILES_BASE_URL: 'https://dify.example.test/base', + }) + }) +}) diff --git a/e2e/tests/marketplace-plugins.test.ts b/e2e/tests/marketplace-plugins.test.ts index dda4008adb7..4722212f657 100644 --- a/e2e/tests/marketplace-plugins.test.ts +++ b/e2e/tests/marketplace-plugins.test.ts @@ -58,13 +58,16 @@ describe('bootstrapMarketplacePlugins', () => { it('uses generated package upload when the API process cannot download from Marketplace', async () => { vi.stubEnv('E2E_TEST_MARKETPLACE_PLUGIN_IDS', '') - vi.spyOn(globalThis, 'fetch').mockResolvedValue(new Response('plugin-package')) + vi.stubEnv('E2E_MARKETPLACE_API_URL', 'https://marketplace.test') + const marketplaceFetch = vi + .spyOn(globalThis, 'fetch') + .mockResolvedValue(new Response('plugin-package')) const { consoleClient, installPackage, uploadPackage } = createMarketplaceConsoleClient( new ORPCError('INTERNAL_SERVER_ERROR', { data: { body: { message: - 'Reached maximum retries (3) for URL https://marketplace.test/plugins/download', + 'Reached maximum retries (3) for URL https://marketplace.test/api/v1/plugins/download-url', }, }, status: 500, @@ -73,6 +76,10 @@ describe('bootstrapMarketplacePlugins', () => { const result = await bootstrapTestPlugin(consoleClient) expect(result.status).toBe('verified') + expect(marketplaceFetch).toHaveBeenCalledOnce() + expect(marketplaceFetch).toHaveBeenCalledWith( + 'https://marketplace.test/api/v1/plugins/download-url?unique_identifier=langgenius%2Ftest%3A1.0.0%40marketplace', + ) expect(uploadPackage).toHaveBeenCalledWith({ body: { pkg: expect.any(File) }, }) diff --git a/e2e/tests/run-options.test.ts b/e2e/tests/run-options.test.ts new file mode 100644 index 00000000000..638d0091321 --- /dev/null +++ b/e2e/tests/run-options.test.ts @@ -0,0 +1,65 @@ +import { describe, expect, it } from 'vitest' +import { parseRunOptions, shouldStartManagedAgentBackend } from '../scripts/run-options' + +describe('E2E run options', () => { + it('forwards Cucumber arguments without requesting seed data', () => { + expect(parseRunOptions(['--tags', '@smoke'])).toEqual({ + forwardArgs: ['--tags', '@smoke'], + full: false, + headed: false, + seed: undefined, + seedOnly: false, + }) + }) + + it('uses the post-merge profile for the default seed command', () => { + expect(parseRunOptions(['--seed-only'])).toMatchObject({ + forwardArgs: [], + seed: { + allowBlocked: false, + dryRun: false, + pack: 'agent-v2', + profile: 'post-merge', + }, + seedOnly: true, + }) + }) + + it('seeds a named profile before forwarding Cucumber arguments', () => { + expect(parseRunOptions(['--profile', 'prepared', '--', '--tags', '@prepared'])).toMatchObject({ + forwardArgs: ['--tags', '@prepared'], + seed: { profile: 'prepared' }, + seedOnly: false, + }) + }) + + it('rejects seed-only options when no seed was requested', () => { + expect(() => parseRunOptions(['--allow-blocked'])).toThrow( + 'Seed options require --seed-only or --profile.', + ) + expect(() => parseRunOptions(['--dry-run', '--profile', 'prepared'])).toThrow( + '--dry-run requires --seed-only.', + ) + }) +}) + +describe('managed Agent backend selection', () => { + it.each(['1', 'true'])('starts a managed backend for %s', (value) => { + expect(shouldStartManagedAgentBackend({ E2E_START_AGENT_BACKEND: value })).toBe(true) + }) + + it('uses an explicitly configured backend without starting a managed one', () => { + expect( + shouldStartManagedAgentBackend({ E2E_AGENT_BACKEND_URL: 'http://agent.example.test' }), + ).toBe(false) + }) + + it('rejects two Agent backend owners', () => { + expect(() => + shouldStartManagedAgentBackend({ + AGENT_BACKEND_BASE_URL: 'http://agent.example.test', + E2E_START_AGENT_BACKEND: '1', + }), + ).toThrow('E2E_START_AGENT_BACKEND cannot be enabled') + }) +}) diff --git a/knip.config.ts b/knip.config.ts new file mode 100644 index 00000000000..63e60aeab63 --- /dev/null +++ b/knip.config.ts @@ -0,0 +1,58 @@ +import type { KnipConfig } from 'knip' + +/** + * @see https://knip.dev/reference/configuration + */ +const config: KnipConfig = { + compilers: { + mdx: true, + }, + workspaces: { + web: { + entry: [ + 'scripts/**/*.{js,ts,mjs}', + 'bin/**/*.{js,ts,mjs}', + 'tsslint.config.ts', + 'dev-proxy.config.ts', + 'plugins/eslint/index.js', + ], + project: [ + '**/*.{js,mjs,cjs,jsx,ts,tsx,mts,cts,css,mdx}!', + '!**/__mocks__/**!', + '!**/__tests__/**!', + '!**/*.stories.{js,jsx,ts,tsx,mdx}!', + '!.storybook/**!', + '!plugins/**!', + '!test/**!', + '!vitest.setup.ts!', + ], + ignore: ['public/**'], + ignoreFiles: [ + 'features/agent-v2/agent-detail/configure/components/orchestrate/memory.tsx', + 'features/agent-v2/agent-detail/configure/components/orchestrate/prompt-editor/option-menu.tsx', + 'i18n-config/locale-resources/*.ts', + ], + ignoreDependencies: ['@iconify-json/*', '@storybook/addon-onboarding'], + }, + }, + /// keep-sorted + rules: { + binaries: 'warn', + catalog: 'error', + dependencies: 'error', + devDependencies: 'warn', + duplicates: 'error', + enumMembers: 'error', + exports: 'error', + files: 'error', + namespaceMembers: 'error', + nsExports: 'error', + nsTypes: 'error', + optionalPeerDependencies: 'error', + types: 'error', + unlisted: 'error', + unresolved: 'error', + }, +} + +export default config diff --git a/lint.config.ts b/lint.config.ts index 6ee726761a9..b8d7d7f9a04 100644 --- a/lint.config.ts +++ b/lint.config.ts @@ -1,8 +1,114 @@ +/// + import type { OxlintConfig } from 'vite-plus/lint' import path from 'node:path' const rootDir = import.meta.dirname const difyUiPackageJson = path.resolve(rootDir, 'packages/dify-ui/package.json') +const enableTailwindCanonicalClasses = process.env.TAILWIND_CANONICAL_CLASSES === 'true' +const tailwindCanonicalClassesOverride = { + files: [ + 'web/**/*.{js,cjs,mjs,jsx,ts,cts,mts,tsx}', + 'packages/dify-ui/**/*.{js,cjs,mjs,jsx,ts,cts,mts,tsx}', + ], + rules: { + 'better-tailwindcss/enforce-canonical-classes': [ + 'warn', + { + collapse: false, + logical: false, + }, + ], + }, +} satisfies NonNullable[number] + +export const webJsxA11yRules = { + 'jsx-a11y/alt-text': 'error', + 'jsx-a11y/anchor-ambiguous-text': 'off', + 'jsx-a11y/anchor-has-content': 'error', + 'jsx-a11y/anchor-is-valid': 'error', + 'jsx-a11y/aria-activedescendant-has-tabindex': 'error', + 'jsx-a11y/aria-props': 'error', + 'jsx-a11y/aria-proptypes': 'error', + 'jsx-a11y/aria-role': 'error', + 'jsx-a11y/aria-unsupported-elements': 'error', + 'jsx-a11y/autocomplete-valid': 'error', + 'jsx-a11y/click-events-have-key-events': 'error', + 'jsx-a11y/heading-has-content': 'error', + 'jsx-a11y/html-has-lang': 'error', + 'jsx-a11y/iframe-has-title': 'error', + 'jsx-a11y/img-redundant-alt': 'error', + 'jsx-a11y/interactive-supports-focus': [ + 'error', + { + tabbable: ['button', 'checkbox', 'link', 'searchbox', 'spinbutton', 'switch', 'textbox'], + }, + ], + 'jsx-a11y/label-has-associated-control': 'error', + 'jsx-a11y/media-has-caption': 'error', + 'jsx-a11y/mouse-events-have-key-events': 'error', + 'jsx-a11y/no-access-key': 'error', + 'jsx-a11y/no-autofocus': 'error', + 'jsx-a11y/no-distracting-elements': 'error', + 'jsx-a11y/no-interactive-element-to-noninteractive-role': [ + 'error', + { + tr: ['none', 'presentation'], + canvas: ['img'], + }, + ], + 'jsx-a11y/no-noninteractive-element-interactions': [ + 'error', + { + handlers: [ + 'onClick', + 'onError', + 'onLoad', + 'onMouseDown', + 'onMouseUp', + 'onKeyPress', + 'onKeyDown', + 'onKeyUp', + ], + alert: ['onKeyUp', 'onKeyDown', 'onKeyPress'], + body: ['onError', 'onLoad'], + dialog: ['onKeyUp', 'onKeyDown', 'onKeyPress'], + iframe: ['onError', 'onLoad'], + img: ['onError', 'onLoad'], + }, + ], + 'jsx-a11y/no-noninteractive-element-to-interactive-role': [ + 'error', + { + ul: ['listbox', 'menu', 'menubar', 'radiogroup', 'tablist', 'tree', 'treegrid'], + ol: ['listbox', 'menu', 'menubar', 'radiogroup', 'tablist', 'tree', 'treegrid'], + li: ['menuitem', 'menuitemradio', 'menuitemcheckbox', 'option', 'row', 'tab', 'treeitem'], + table: ['grid'], + td: ['gridcell'], + fieldset: ['radiogroup', 'presentation'], + }, + ], + 'jsx-a11y/no-noninteractive-tabindex': [ + 'error', + { + tags: [], + roles: ['tabpanel'], + allowExpressionValues: true, + }, + ], + 'jsx-a11y/no-redundant-roles': 'error', + 'jsx-a11y/no-static-element-interactions': [ + 'error', + { + allowExpressionValues: true, + handlers: ['onClick', 'onMouseDown', 'onMouseUp', 'onKeyPress', 'onKeyDown', 'onKeyUp'], + }, + ], + 'jsx-a11y/role-has-required-aria-props': 'error', + 'jsx-a11y/role-supports-aria-props': 'error', + 'jsx-a11y/scope': 'error', + 'jsx-a11y/tabindex-no-positive': 'error', +} satisfies NonNullable /** * Oxlint equivalent of the ESLint configurations that were active before the migration. @@ -55,6 +161,7 @@ export const lintConfig = { jsPlugins: [ '@tanstack/eslint-plugin-query', 'eslint-plugin-antfu', + ...(enableTailwindCanonicalClasses ? ['eslint-plugin-better-tailwindcss'] : []), 'eslint-plugin-command', 'eslint-plugin-erasable-syntax-only', { @@ -86,6 +193,15 @@ export const lintConfig = { typeCheck: true, }, settings: { + ...(enableTailwindCanonicalClasses + ? { + 'better-tailwindcss': { + cwd: path.resolve(rootDir, 'web'), + entryPoint: 'app/styles/globals.css', + rootFontSize: 16, + }, + } + : {}), 'react-x': { additionalStateHooks: '/^use\\w*State(?:s)?|useAtom$/u', }, @@ -425,6 +541,7 @@ export const lintConfig = { 'no-undef': 'error', }, overrides: [ + ...(enableTailwindCanonicalClasses ? [tailwindCanonicalClassesOverride] : []), { files: ['**/*.{js,cjs,mjs,jsx,ts,cts,mts,tsx}'], rules: { @@ -667,109 +784,7 @@ export const lintConfig = { }, { files: ['web/**/*.tsx'], - rules: { - 'jsx-a11y/alt-text': 'error', - 'jsx-a11y/anchor-ambiguous-text': 'off', - 'jsx-a11y/anchor-has-content': 'error', - 'jsx-a11y/anchor-is-valid': 'error', - 'jsx-a11y/aria-activedescendant-has-tabindex': 'error', - 'jsx-a11y/aria-props': 'error', - 'jsx-a11y/aria-proptypes': 'error', - 'jsx-a11y/aria-role': 'error', - 'jsx-a11y/aria-unsupported-elements': 'error', - 'jsx-a11y/autocomplete-valid': 'error', - 'jsx-a11y/click-events-have-key-events': 'error', - 'jsx-a11y/heading-has-content': 'error', - 'jsx-a11y/html-has-lang': 'error', - 'jsx-a11y/iframe-has-title': 'error', - 'jsx-a11y/img-redundant-alt': 'error', - 'jsx-a11y/interactive-supports-focus': [ - 'error', - { - tabbable: [ - 'button', - 'checkbox', - 'link', - 'searchbox', - 'spinbutton', - 'switch', - 'textbox', - ], - }, - ], - 'jsx-a11y/label-has-associated-control': 'error', - 'jsx-a11y/media-has-caption': 'error', - 'jsx-a11y/mouse-events-have-key-events': 'error', - 'jsx-a11y/no-access-key': 'error', - 'jsx-a11y/no-autofocus': 'error', - 'jsx-a11y/no-distracting-elements': 'error', - 'jsx-a11y/no-interactive-element-to-noninteractive-role': [ - 'error', - { - tr: ['none', 'presentation'], - canvas: ['img'], - }, - ], - 'jsx-a11y/no-noninteractive-element-interactions': [ - 'error', - { - handlers: [ - 'onClick', - 'onError', - 'onLoad', - 'onMouseDown', - 'onMouseUp', - 'onKeyPress', - 'onKeyDown', - 'onKeyUp', - ], - alert: ['onKeyUp', 'onKeyDown', 'onKeyPress'], - body: ['onError', 'onLoad'], - dialog: ['onKeyUp', 'onKeyDown', 'onKeyPress'], - iframe: ['onError', 'onLoad'], - img: ['onError', 'onLoad'], - }, - ], - 'jsx-a11y/no-noninteractive-element-to-interactive-role': [ - 'error', - { - ul: ['listbox', 'menu', 'menubar', 'radiogroup', 'tablist', 'tree', 'treegrid'], - ol: ['listbox', 'menu', 'menubar', 'radiogroup', 'tablist', 'tree', 'treegrid'], - li: [ - 'menuitem', - 'menuitemradio', - 'menuitemcheckbox', - 'option', - 'row', - 'tab', - 'treeitem', - ], - table: ['grid'], - td: ['gridcell'], - fieldset: ['radiogroup', 'presentation'], - }, - ], - 'jsx-a11y/no-noninteractive-tabindex': [ - 'error', - { - tags: [], - roles: ['tabpanel'], - allowExpressionValues: true, - }, - ], - 'jsx-a11y/no-redundant-roles': 'error', - 'jsx-a11y/no-static-element-interactions': [ - 'error', - { - allowExpressionValues: true, - handlers: ['onClick', 'onMouseDown', 'onMouseUp', 'onKeyPress', 'onKeyDown', 'onKeyUp'], - }, - ], - 'jsx-a11y/role-has-required-aria-props': 'error', - 'jsx-a11y/role-supports-aria-props': 'error', - 'jsx-a11y/scope': 'error', - 'jsx-a11y/tabindex-no-positive': 'error', - }, + rules: webJsxA11yRules, }, { files: ['web/**/*.stories.{js,cjs,mjs,jsx,ts,tsx}', 'web/**/*.story.{js,cjs,mjs,jsx,ts,tsx}'], diff --git a/oxlint-suppressions.json b/oxlint-suppressions.json index 0569de71067..2bd94977d9e 100644 --- a/oxlint-suppressions.json +++ b/oxlint-suppressions.json @@ -4,26 +4,11 @@ "count": 2 } }, - "cli/test/e2e/helpers/retry.ts": { - "no-throw-literal": { - "count": 1 - } - }, "cli/test/e2e/suites/output/json-yaml-output.e2e.ts": { "prefer-const": { "count": 1 } }, - "e2e/support/web-server.ts": { - "no-throw-literal": { - "count": 1 - } - }, - "packages/dev-proxy/src/cli.spec.ts": { - "no-throw-literal": { - "count": 1 - } - }, "packages/migrate-no-unchecked-indexed-access/src/no-unchecked-indexed-access/migrate.ts": { "no-console": { "count": 11 @@ -59,44 +44,6 @@ "count": 3 } }, - "web/__mocks__/base-ui-dropdown-menu.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 3 - }, - "jsx_a11y/interactive-supports-focus": { - "count": 2 - }, - "jsx_a11y/role-has-required-aria-props": { - "count": 1 - } - }, - "web/__mocks__/base-ui-popover.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, - "web/__mocks__/base-ui-select.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/interactive-supports-focus": { - "count": 1 - }, - "jsx_a11y/role-has-required-aria-props": { - "count": 2 - } - }, - "web/__mocks__/base-ui-tooltip.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/__mocks__/zustand.ts": { "no-barrel-files/no-barrel-files": { "count": 1 @@ -117,14 +64,6 @@ "count": 3 } }, - "web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/time-range-picker/date-picker.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/(commonLayout)/app/(appDetailLayout)/[appId]/overview/tracing/field.tsx": { "no-restricted-imports": { "count": 1 @@ -149,21 +88,7 @@ "count": 1 } }, - "web/app/(shareLayout)/components/authenticated-layout.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/(shareLayout)/webapp-reset-password/check-code/page.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - }, "no-restricted-imports": { "count": 1 } @@ -184,12 +109,6 @@ } }, "web/app/(shareLayout)/webapp-signin/check-code/page.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - }, "no-restricted-imports": { "count": 1 } @@ -199,14 +118,6 @@ "count": 1 } }, - "web/app/(shareLayout)/webapp-signin/components/mail-and-password-auth.tsx": { - "jsx_a11y/tabindex-no-positive": { - "count": 3 - }, - "no-restricted-imports": { - "count": 1 - } - }, "web/app/(shareLayout)/webapp-signin/normalForm.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 2 @@ -215,14 +126,6 @@ "count": 2 } }, - "web/app/(shareLayout)/webapp-signin/page.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/account/(commonLayout)/account-page/email-change-modal.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 3 @@ -234,17 +137,6 @@ "count": 2 } }, - "web/app/account/(commonLayout)/account-page/index.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 2 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 2 - }, - "no-restricted-imports": { - "count": 1 - } - }, "web/app/account/(commonLayout)/delete-account/components/verify-email.tsx": { "eslint-react/set-state-in-effect": { "count": 1 @@ -253,27 +145,6 @@ "count": 1 } }, - "web/app/components/app-sidebar/app-info/app-info-modals.tsx": { - "jsx_a11y/label-has-associated-control": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, - "web/app/components/app-sidebar/app-info/app-operations.tsx": { - "eslint-react/set-state-in-effect": { - "count": 4 - } - }, - "web/app/components/app-sidebar/dataset-info/menu-item.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/app/annotation/add-annotation-modal/edit-item/index.tsx": { "erasable-syntax-only/enums": { "count": 1 @@ -378,14 +249,6 @@ "count": 1 } }, - "web/app/components/app/app-publisher/sections.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/app/configuration/config-prompt/__tests__/index.spec.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -473,7 +336,7 @@ "count": 2 }, "typescript/no-explicit-any": { - "count": 9 + "count": 7 } }, "web/app/components/app/configuration/config/agent/agent-tools/setting-built-in-tool.tsx": { @@ -490,20 +353,6 @@ "count": 4 } }, - "web/app/components/app/configuration/config/automatic/get-automatic-res.tsx": { - "eslint-react/set-state-in-effect": { - "count": 4 - }, - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - }, - "typescript/no-explicit-any": { - "count": 1 - } - }, "web/app/components/app/configuration/config/automatic/idea-output.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -532,14 +381,6 @@ "count": 1 } }, - "web/app/components/app/configuration/config/code-generator/get-code-generator-res.tsx": { - "eslint-react/set-state-in-effect": { - "count": 4 - }, - "typescript/no-explicit-any": { - "count": 2 - } - }, "web/app/components/app/configuration/dataset-config/context-var/var-picker.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -583,11 +424,6 @@ "count": 1 } }, - "web/app/components/app/configuration/debug/__tests__/chat-user-input.spec.tsx": { - "jsx_a11y/role-has-required-aria-props": { - "count": 1 - } - }, "web/app/components/app/configuration/debug/__tests__/index.spec.tsx": { "typescript/no-explicit-any": { "count": 1 @@ -622,9 +458,6 @@ } }, "web/app/components/app/configuration/debug/hooks.tsx": { - "no-restricted-globals": { - "count": 2 - }, "typescript/no-explicit-any": { "count": 3 } @@ -650,11 +483,6 @@ "count": 1 } }, - "web/app/components/app/create-app-dialog/app-list/index.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "web/app/components/app/create-app-dialog/app-list/sidebar.tsx": { "erasable-syntax-only/enums": { "count": 1 @@ -710,27 +538,11 @@ "count": 1 } }, - "web/app/components/app/overview/app-card-sections.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 2 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 2 - } - }, "web/app/components/app/overview/workflow-hidden-input-fields.tsx": { "no-restricted-imports": { "count": 1 } }, - "web/app/components/app/switch-app-modal/index.tsx": { - "eslint-react/set-state-in-effect": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "web/app/components/app/text-generate/item/index.tsx": { "typescript/no-explicit-any": { "count": 3 @@ -767,17 +579,6 @@ "count": 2 } }, - "web/app/components/apps/app-card.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-noninteractive-element-to-interactive-role": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/apps/import-from-marketplace-template-modal.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -921,14 +722,6 @@ "count": 3 } }, - "web/app/components/base/chat/chat-with-history/sidebar/__tests__/operation.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 2 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 2 - } - }, "web/app/components/base/chat/chat-with-history/sidebar/item.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 2 @@ -957,30 +750,6 @@ "count": 1 } }, - "web/app/components/base/chat/chat/answer/human-input-content/content-wrapper.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, - "web/app/components/base/chat/chat/answer/suggested-questions.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, - "web/app/components/base/chat/chat/answer/tool-detail.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/base/chat/chat/answer/workflow-process.tsx": { "eslint-react/set-state-in-effect": { "count": 1 @@ -994,12 +763,6 @@ "web/app/components/base/chat/chat/citation/index.tsx": { "eslint-react/set-state-in-effect": { "count": 1 - }, - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 } }, "web/app/components/base/chat/chat/hooks.ts": { @@ -1082,14 +845,6 @@ "count": 2 } }, - "web/app/components/base/date-and-time-picker/time-picker/__tests__/index.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/base/date-and-time-picker/time-picker/index.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -1217,14 +972,8 @@ } }, "web/app/components/base/file-uploader/audio-preview.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, "jsx_a11y/media-has-caption": { "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 } }, "web/app/components/base/file-uploader/dynamic-pdf-preview.tsx": { @@ -1276,14 +1025,6 @@ "count": 2 } }, - "web/app/components/base/file-uploader/pdf-preview.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/base/file-uploader/store.tsx": { "react/only-export-components": { "count": 4 @@ -1303,14 +1044,8 @@ } }, "web/app/components/base/file-uploader/video-preview.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, "jsx_a11y/media-has-caption": { "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 } }, "web/app/components/base/form/components/base/base-field.tsx": { @@ -1409,12 +1144,7 @@ }, "web/app/components/base/icons/src/public/common/index.ts": { "no-barrel-files/no-barrel-files": { - "count": 6 - } - }, - "web/app/components/base/icons/src/public/education/index.ts": { - "no-barrel-files/no-barrel-files": { - "count": 1 + "count": 5 } }, "web/app/components/base/icons/src/public/files/index.ts": { @@ -1559,7 +1289,7 @@ }, "web/app/components/base/icons/src/vender/solid/development/index.ts": { "no-barrel-files/no-barrel-files": { - "count": 4 + "count": 3 } }, "web/app/components/base/icons/src/vender/solid/editor/index.ts": { @@ -1569,7 +1299,7 @@ }, "web/app/components/base/icons/src/vender/solid/education/index.ts": { "no-barrel-files/no-barrel-files": { - "count": 2 + "count": 1 } }, "web/app/components/base/icons/src/vender/solid/files/index.ts": { @@ -1657,9 +1387,6 @@ } }, "web/app/components/base/image-uploader/image-preview.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, "jsx_a11y/no-static-element-interactions": { "count": 1 }, @@ -1691,11 +1418,6 @@ "count": 1 } }, - "web/app/components/base/logo/dify-logo.tsx": { - "react/only-export-components": { - "count": 2 - } - }, "web/app/components/base/markdown-blocks/__tests__/paragraph.spec.tsx": { "jsx_a11y/anchor-is-valid": { "count": 1 @@ -1706,11 +1428,6 @@ "count": 5 } }, - "web/app/components/base/markdown-blocks/button.tsx": { - "typescript/no-explicit-any": { - "count": 1 - } - }, "web/app/components/base/markdown-blocks/code-block.tsx": { "typescript/no-explicit-any": { "count": 9 @@ -2041,14 +1758,6 @@ "count": 1 } }, - "web/app/components/base/qrcode/index.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/base/search-input/index.stories.tsx": { "jsx_a11y/label-has-associated-control": { "count": 3 @@ -2062,14 +1771,6 @@ "count": 1 } }, - "web/app/components/base/tab-slider-new/index.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/base/tab-slider-plain/index.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -2078,17 +1779,6 @@ "count": 1 } }, - "web/app/components/base/tab-slider/index.tsx": { - "eslint-react/set-state-in-effect": { - "count": 2 - }, - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/base/tag-input/index.stories.tsx": { "jsx_a11y/label-has-associated-control": { "count": 12 @@ -2146,14 +1836,9 @@ "count": 1 } }, - "web/app/components/base/voice-input/recorder.ts": { - "no-throw-literal": { - "count": 1 - } - }, "web/app/components/billing/plan/assets/index.tsx": { "no-barrel-files/no-barrel-files": { - "count": 4 + "count": 3 } }, "web/app/components/billing/pricing/assets/index.tsx": { @@ -2169,14 +1854,6 @@ "count": 1 } }, - "web/app/components/billing/pricing/plan-switcher/tab.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/billing/pricing/plans/self-hosted-plan-item/button.tsx": { "eslint-react/static-components": { "count": 2 @@ -2187,21 +1864,11 @@ "count": 1 } }, - "web/app/components/billing/type.ts": { - "erasable-syntax-only/enums": { - "count": 4 - } - }, "web/app/components/datasets/chunk.tsx": { "jsx_a11y/label-has-associated-control": { "count": 2 } }, - "web/app/components/datasets/common/credential-icon.tsx": { - "jsx_a11y/alt-text": { - "count": 1 - } - }, "web/app/components/datasets/common/document-status-with-action/status-with-action.tsx": { "eslint-react/static-components": { "count": 2 @@ -2802,14 +2469,6 @@ "count": 1 } }, - "web/app/components/datasets/list/dataset-card/__tests__/index.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/datasets/list/dataset-card/components/operations-dropdown.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -2921,11 +2580,6 @@ "count": 1 } }, - "web/app/components/explore/learn-dify/item.tsx": { - "jsx_a11y/no-noninteractive-element-interactions": { - "count": 1 - } - }, "web/app/components/explore/try-app/app/text-generation.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -2984,12 +2638,6 @@ } }, "web/app/components/header/account-setting/members-page/transfer-ownership-modal/member-selector.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - }, "no-restricted-imports": { "count": 1 } @@ -3002,14 +2650,6 @@ "count": 4 } }, - "web/app/components/header/account-setting/model-provider-page/hooks.ts": { - "eslint-react/no-unnecessary-use-prefix": { - "count": 1 - }, - "typescript/no-explicit-any": { - "count": 2 - } - }, "web/app/components/header/account-setting/model-provider-page/model-auth/add-custom-model.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 2 @@ -3034,14 +2674,6 @@ "count": 1 } }, - "web/app/components/header/account-setting/model-provider-page/model-auth/config-model.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/header/account-setting/model-provider-page/model-auth/credential-selector.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -3085,11 +2717,6 @@ "count": 5 } }, - "web/app/components/header/account-setting/model-provider-page/model-parameter-modal/configuration-button.tsx": { - "typescript/no-explicit-any": { - "count": 2 - } - }, "web/app/components/header/account-setting/model-provider-page/model-parameter-modal/index.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -3230,9 +2857,6 @@ } }, "web/app/components/plugins/marketplace/hooks.ts": { - "@tanstack/query/exhaustive-deps": { - "count": 1 - }, "no-restricted-imports": { "count": 1 } @@ -3245,14 +2869,6 @@ "count": 1 } }, - "web/app/components/plugins/marketplace/plugin-type-switch.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/plugins/plugin-auth/authorized/index.tsx": { "no-restricted-imports": { "count": 1 @@ -3375,14 +2991,6 @@ "count": 2 } }, - "web/app/components/plugins/plugin-detail-panel/subscription-list/__tests__/selector-entry.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/plugins/plugin-detail-panel/subscription-list/create/hooks/use-common-modal-state.ts": { "erasable-syntax-only/enums": { "count": 1 @@ -3469,14 +3077,6 @@ "count": 1 } }, - "web/app/components/plugins/reference-setting-modal/auto-update-setting/__tests__/tool-picker.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/plugins/reference-setting-modal/auto-update-setting/tool-picker.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -3498,14 +3098,6 @@ "count": 25 } }, - "web/app/components/plugins/update-plugin/plugin-version-picker.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/rag-pipeline/components/__tests__/publish-as-knowledge-pipeline-modal.spec.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -3634,11 +3226,6 @@ "count": 1 } }, - "web/app/components/rag-pipeline/hooks/index.ts": { - "no-barrel-files/no-barrel-files": { - "count": 9 - } - }, "web/app/components/rag-pipeline/hooks/use-DSL.ts": { "typescript/no-explicit-any": { "count": 1 @@ -3816,14 +3403,6 @@ "count": 1 } }, - "web/app/components/tools/labels/__tests__/selector.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/tools/labels/filter.tsx": { "no-restricted-imports": { "count": 1 @@ -3839,14 +3418,6 @@ "count": 1 } }, - "web/app/components/tools/provider/tool-item.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/tools/setting/build-in/config-credentials.tsx": { "typescript/no-explicit-any": { "count": 3 @@ -3860,22 +3431,6 @@ "count": 4 } }, - "web/app/components/tools/workflow-tool/__tests__/configure-button.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, - "web/app/components/tools/workflow-tool/configure-button.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/tools/workflow-tool/index.tsx": { "no-restricted-imports": { "count": 1 @@ -3891,11 +3446,6 @@ "count": 3 } }, - "web/app/components/workflow-app/hooks/index.ts": { - "no-barrel-files/no-barrel-files": { - "count": 13 - } - }, "web/app/components/workflow-app/hooks/use-DSL.ts": { "typescript/no-explicit-any": { "count": 1 @@ -3988,14 +3538,6 @@ "count": 2 } }, - "web/app/components/workflow/header/scroll-to-selected-node-button.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/workflow/header/test-run-menu.tsx": { "erasable-syntax-only/enums": { "count": 1 @@ -4035,11 +3577,6 @@ "count": 1 } }, - "web/app/components/workflow/hooks/index.ts": { - "no-barrel-files/no-barrel-files": { - "count": 25 - } - }, "web/app/components/workflow/hooks/use-checklist.ts": { "typescript/no-empty-object-type": { "count": 1 @@ -4068,11 +3605,6 @@ "count": 1 } }, - "web/app/components/workflow/hooks/use-workflow-run-event/index.ts": { - "no-barrel-files/no-barrel-files": { - "count": 19 - } - }, "web/app/components/workflow/hooks/use-workflow-run-event/use-workflow-agent-log.ts": { "typescript/no-explicit-any": { "count": 1 @@ -4093,14 +3625,6 @@ "count": 1 } }, - "web/app/components/workflow/nodes/_base/components/__tests__/agent-strategy-selector.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/workflow/nodes/_base/components/__tests__/node-handle.spec.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -4985,9 +4509,6 @@ }, "web/app/components/workflow/nodes/llm/use-config.ts": { "eslint-react/set-state-in-effect": { - "count": 2 - }, - "typescript/no-explicit-any": { "count": 1 } }, @@ -5070,11 +4591,6 @@ "count": 3 } }, - "web/app/components/workflow/nodes/parameter-extractor/components/extract-parameter/__tests__/list.spec.tsx": { - "no-unused-vars": { - "count": 1 - } - }, "web/app/components/workflow/nodes/parameter-extractor/components/extract-parameter/item.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 2 @@ -5144,22 +4660,6 @@ "count": 5 } }, - "web/app/components/workflow/nodes/tool/components/__tests__/copy-id.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, - "web/app/components/workflow/nodes/tool/components/copy-id.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/workflow/nodes/tool/components/mixed-variable-text-input/index.tsx": { "typescript/no-explicit-any": { "count": 1 @@ -5318,11 +4818,6 @@ "count": 1 } }, - "web/app/components/workflow/note-node/note-editor/toolbar/__tests__/operator.spec.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - } - }, "web/app/components/workflow/note-node/note-editor/toolbar/color-picker.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 1 @@ -5436,14 +4931,6 @@ "count": 1 } }, - "web/app/components/workflow/panel/chat-variable-panel/index.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/workflow/panel/chat-variable-panel/type.ts": { "erasable-syntax-only/enums": { "count": 1 @@ -5480,14 +4967,6 @@ "count": 11 } }, - "web/app/components/workflow/panel/debug-and-preview/index.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/workflow/panel/debug-and-preview/user-input.tsx": { "jsx_a11y/no-autofocus": { "count": 1 @@ -5498,28 +4977,17 @@ "count": 2 } }, - "web/app/components/workflow/panel/env-panel/index.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/workflow/panel/env-panel/variable-modal.tsx": { "eslint-react/set-state-in-effect": { "count": 4 }, "jsx_a11y/click-events-have-key-events": { - "count": 4 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 4 - }, - "no-restricted-imports": { "count": 1 }, - "typescript/no-explicit-any": { + "jsx_a11y/no-static-element-interactions": { + "count": 1 + }, + "no-restricted-imports": { "count": 1 } }, @@ -5531,11 +4999,6 @@ "count": 1 } }, - "web/app/components/workflow/panel/human-input-form-list.tsx": { - "typescript/no-explicit-any": { - "count": 1 - } - }, "web/app/components/workflow/panel/inputs-panel.tsx": { "jsx_a11y/no-autofocus": { "count": 1 @@ -5544,30 +5007,6 @@ "count": 4 } }, - "web/app/components/workflow/panel/version-history-panel/filter/filter-item.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, - "web/app/components/workflow/panel/version-history-panel/index.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 2 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 2 - } - }, - "web/app/components/workflow/panel/version-history-panel/version-history-item.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/app/components/workflow/panel/workflow-preview.tsx": { "jsx_a11y/click-events-have-key-events": { "count": 4 @@ -5678,9 +5117,6 @@ } }, "web/app/components/workflow/run/output-panel.tsx": { - "jsx_a11y/no-noninteractive-tabindex": { - "count": 1 - }, "typescript/no-explicit-any": { "count": 3 } @@ -5762,7 +5198,7 @@ "count": 3 }, "typescript/no-explicit-any": { - "count": 8 + "count": 6 } }, "web/app/components/workflow/update-dsl-modal.tsx": { @@ -5810,11 +5246,6 @@ "count": 1 } }, - "web/app/components/workflow/variable-inspect/display-content.tsx": { - "typescript/no-explicit-any": { - "count": 1 - } - }, "web/app/components/workflow/variable-inspect/group.tsx": { "typescript/no-explicit-any": { "count": 2 @@ -5902,25 +5333,6 @@ "count": 1 } }, - "web/app/education-apply/role-selector.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, - "web/app/education-apply/search-input.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "web/app/forgot-password/ChangePasswordForm.tsx": { "no-restricted-imports": { "count": 1 @@ -5952,12 +5364,6 @@ } }, "web/app/reset-password/check-code/page.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - }, "no-restricted-imports": { "count": 1 } @@ -5977,17 +5383,6 @@ "count": 1 } }, - "web/app/signin/check-code/page.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "web/app/signin/layout.tsx": { "typescript/no-explicit-any": { "count": 1 @@ -5999,20 +5394,11 @@ } }, "web/app/signup/check-code/page.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - }, "no-restricted-imports": { "count": 1 } }, "web/app/signup/components/input-mail.tsx": { - "jsx_a11y/tabindex-no-positive": { - "count": 2 - }, "no-restricted-imports": { "count": 1 } @@ -6027,16 +5413,6 @@ "count": 1 } }, - "web/context/external-api-panel-context.tsx": { - "react/only-export-components": { - "count": 1 - } - }, - "web/context/external-knowledge-api-context.tsx": { - "react/only-export-components": { - "count": 1 - } - }, "web/context/hooks/use-trigger-events-limit-modal.ts": { "eslint-react/set-state-in-effect": { "count": 3 @@ -6050,14 +5426,6 @@ "count": 1 } }, - "web/features/tag-management/components/dataset-card-tags.tsx": { - "jsx_a11y/click-events-have-key-events": { - "count": 1 - }, - "jsx_a11y/no-static-element-interactions": { - "count": 1 - } - }, "web/features/tag-management/components/tag-item-editor.tsx": { "jsx_a11y/no-autofocus": { "count": 1 @@ -6106,11 +5474,6 @@ "count": 1 } }, - "web/models/app.ts": { - "erasable-syntax-only/enums": { - "count": 2 - } - }, "web/models/datasets.ts": { "erasable-syntax-only/enums": { "count": 7 @@ -6201,7 +5564,7 @@ "count": 1 }, "typescript/no-explicit-any": { - "count": 6 + "count": 4 } }, "web/service/base.ts": { @@ -6212,17 +5575,12 @@ "count": 3 } }, - "web/service/billing.ts": { - "no-restricted-imports": { - "count": 1 - } - }, "web/service/common.ts": { "no-restricted-imports": { "count": 1 }, "typescript/no-explicit-any": { - "count": 8 + "count": 6 } }, "web/service/datasets.ts": { @@ -6230,7 +5588,7 @@ "count": 1 }, "typescript/no-explicit-any": { - "count": 5 + "count": 3 } }, "web/service/debug.ts": { @@ -6241,11 +5599,6 @@ "count": 6 } }, - "web/service/fetch.spec.ts": { - "no-restricted-imports": { - "count": 1 - } - }, "web/service/fetch.ts": { "no-restricted-imports": { "count": 1 @@ -6335,21 +5688,11 @@ "count": 1 } }, - "web/service/use-apps.ts": { - "no-restricted-imports": { - "count": 1 - } - }, "web/service/use-datasource.ts": { "no-restricted-imports": { "count": 1 } }, - "web/service/use-education.ts": { - "no-restricted-imports": { - "count": 1 - } - }, "web/service/use-endpoints.ts": { "no-restricted-imports": { "count": 1 @@ -6374,9 +5717,6 @@ } }, "web/service/use-pipeline.ts": { - "@tanstack/query/exhaustive-deps": { - "count": 1 - }, "no-restricted-imports": { "count": 1 } @@ -6395,9 +5735,6 @@ } }, "web/service/use-workflow.ts": { - "@tanstack/query/exhaustive-deps": { - "count": 1 - }, "no-restricted-imports": { "count": 1 }, diff --git a/package.json b/package.json index 17d05fb488b..736d6afac10 100644 --- a/package.json +++ b/package.json @@ -6,10 +6,15 @@ "check": "vp check && pnpm lint:eslint", "check:fix": "pnpm lint:eslint:fix && vp check --fix", "dev": "concurrently -k -n vinext,proxy \"vp run dify-web#dev:vinext\" \"vp run dify-web#dev:proxy\"", + "knip": "knip --workspace web", + "knip:production": "knip --workspace web --production --include files", + "knip:production-unused-check": "node ./scripts/check-web-production-unused-after-knip-fix.mjs", "prepare": "vp config", "lint:oxlint": "vp lint", "lint:oxlint:fix": "vp lint --fix", "lint:oxlint:quiet": "vp lint --quiet", + "lint:tailwind": "TAILWIND_CANONICAL_CLASSES=true vp lint web packages/dify-ui", + "lint:tailwind:fix": "TAILWIND_CANONICAL_CLASSES=true vp lint --fix web packages/dify-ui", "lint:eslint": "eslint --concurrency=auto", "lint:eslint:fix": "eslint --fix --concurrency=auto", "lint:eslint:quiet": "eslint --quiet --concurrency=auto" @@ -21,12 +26,14 @@ "@iconify-json/heroicons": "catalog:", "@iconify-json/ri": "catalog:", "@tanstack/eslint-plugin-query": "catalog:", + "@types/node": "catalog:", "@typescript-eslint/parser": "catalog:", "@typescript/native": "catalog:", "concurrently": "catalog:", "eslint": "catalog:", "eslint-markdown": "catalog:", "eslint-plugin-antfu": "catalog:", + "eslint-plugin-better-tailwindcss": "catalog:", "eslint-plugin-command": "catalog:", "eslint-plugin-erasable-syntax-only": "catalog:", "eslint-plugin-hyoban": "catalog:", @@ -42,6 +49,7 @@ "eslint-plugin-toml": "catalog:", "eslint-plugin-unicorn": "catalog:", "eslint-plugin-yml": "catalog:", + "knip": "catalog:", "typescript": "catalog:", "vite": "catalog:", "vite-plus": "catalog:" @@ -56,5 +64,5 @@ "engines": { "node": "^22.22.1" }, - "packageManager": "pnpm@11.15.0" + "packageManager": "pnpm@11.20.0" } diff --git a/packages/contracts/console.ts b/packages/contracts/console.ts index af4c1f89f3e..ea712f2924a 100644 --- a/packages/contracts/console.ts +++ b/packages/contracts/console.ts @@ -1,7 +1,12 @@ import { consoleRouterContract as generatedConsoleRouterContract } from './generated/api/console/router.gen' +import { contract as enterpriseAppDeployContract } from './generated/enterprise-app-deploy/orpc.gen' import { contract as knowledgeFsContract } from './generated/knowledge-fs/orpc.gen' export const consoleRouterContract = { ...generatedConsoleRouterContract, + enterprise: { + ...generatedConsoleRouterContract.enterprise, + appDeploy: enterpriseAppDeployContract, + }, knowledgeFs: knowledgeFsContract, } diff --git a/packages/contracts/generated/api/console/agent/orpc.gen.ts b/packages/contracts/generated/api/console/agent/orpc.gen.ts index 12b58ef4188..7d10e128fa0 100644 --- a/packages/contracts/generated/api/console/agent/orpc.gen.ts +++ b/packages/contracts/generated/api/console/agent/orpc.gen.ts @@ -143,7 +143,6 @@ import { zPostAgentByAgentIdCopyBody, zPostAgentByAgentIdCopyPath, zPostAgentByAgentIdCopyResponse, - zPostAgentByAgentIdDebugConversationRefreshBody, zPostAgentByAgentIdDebugConversationRefreshPath, zPostAgentByAgentIdDebugConversationRefreshResponse, zPostAgentByAgentIdFeaturesBody, @@ -158,9 +157,9 @@ import { zPostAgentByAgentIdPublishBody, zPostAgentByAgentIdPublishPath, zPostAgentByAgentIdPublishResponse, - zPostAgentByAgentIdSandboxFilesUploadBody, - zPostAgentByAgentIdSandboxFilesUploadPath, - zPostAgentByAgentIdSandboxFilesUploadResponse, + zPostAgentByAgentIdSandboxFilesDownloadBody, + zPostAgentByAgentIdSandboxFilesDownloadPath, + zPostAgentByAgentIdSandboxFilesDownloadResponse, zPostAgentByAgentIdSkillsBySlugInferToolsPath, zPostAgentByAgentIdSkillsBySlugInferToolsResponse, zPostAgentByAgentIdSkillsUploadBody, @@ -853,12 +852,7 @@ export const post12 = oc path: '/agent/{agent_id}/debug-conversation/refresh', tags: ['console'], }) - .input( - z.object({ - body: zPostAgentByAgentIdDebugConversationRefreshBody.optional(), - params: zPostAgentByAgentIdDebugConversationRefreshPath, - }), - ) + .input(z.object({ params: zPostAgentByAgentIdDebugConversationRefreshPath })) .output(zPostAgentByAgentIdDebugConversationRefreshResponse) export const refresh = { @@ -1185,6 +1179,30 @@ export const referencingWorkflows = { get: get28, } +/** + * Create a ToolFile from one Agent App Binding file and return its download URL + */ +export const post17 = oc + .route({ + description: 'Create a ToolFile from one Agent App Binding file and return its download URL', + inputStructure: 'detailed', + method: 'POST', + operationId: 'postAgentByAgentIdSandboxFilesDownload', + path: '/agent/{agent_id}/sandbox/files/download', + tags: ['console'], + }) + .input( + z.object({ + body: zPostAgentByAgentIdSandboxFilesDownloadBody, + params: zPostAgentByAgentIdSandboxFilesDownloadPath, + }), + ) + .output(zPostAgentByAgentIdSandboxFilesDownloadResponse) + +export const download5 = { + post: post17, +} + /** * Read a text/binary preview file in an Agent App conversation sandbox */ @@ -1209,30 +1227,6 @@ export const read = { get: get29, } -/** - * Upload one Agent App sandbox file and return a signed download URL - */ -export const post17 = oc - .route({ - description: 'Upload one Agent App sandbox file and return a signed download URL', - inputStructure: 'detailed', - method: 'POST', - operationId: 'postAgentByAgentIdSandboxFilesUpload', - path: '/agent/{agent_id}/sandbox/files/upload', - tags: ['console'], - }) - .input( - z.object({ - body: zPostAgentByAgentIdSandboxFilesUploadBody, - params: zPostAgentByAgentIdSandboxFilesUploadPath, - }), - ) - .output(zPostAgentByAgentIdSandboxFilesUploadResponse) - -export const upload2 = { - post: post17, -} - /** * List a directory in an Agent App conversation sandbox */ @@ -1255,8 +1249,8 @@ export const get30 = oc export const files5 = { get: get30, + download: download5, read, - upload: upload2, } /** @@ -1300,7 +1294,7 @@ export const post18 = oc ) .output(zPostAgentByAgentIdSkillsUploadResponse) -export const upload3 = { +export const upload2 = { post: post18, } @@ -1344,7 +1338,7 @@ export const bySlug = { } export const skills3 = { - upload: upload3, + upload: upload2, bySlug, } diff --git a/packages/contracts/generated/api/console/agent/types.gen.ts b/packages/contracts/generated/api/console/agent/types.gen.ts index 98fda2ccab9..1a10436b734 100644 --- a/packages/contracts/generated/api/console/agent/types.gen.ts +++ b/packages/contracts/generated/api/console/agent/types.gen.ts @@ -23,7 +23,7 @@ export type AgentAppCreatePayload = { export type AgentAppDetailWithSite = { access_mode?: string | null - active_config_is_published?: boolean + access_ready?: boolean api_base_url?: string | null app_id?: string | null backing_app_id?: string | null @@ -79,6 +79,7 @@ export type AgentAppUpdatePayload = { } export type AgentApiAccessResponse = { + access_ready: boolean api_key_count: number api_rph: number api_rpm: number @@ -169,6 +170,7 @@ export type SuggestedQuestionsResponse = { } export type AgentAppComposerResponse = { + active_config_is_published: boolean active_config_snapshot?: AgentConfigSnapshotSummaryResponse | null agent: AgentComposerAgentResponse agent_soul: AgentSoulConfig @@ -282,10 +284,6 @@ export type AgentAppCopyPayload = { role?: string | null } -export type AgentDebugConversationRefreshPayload = { - draft_type?: AgentConfigDraftType -} - export type AgentDebugConversationRefreshResponse = { debug_conversation_has_messages?: boolean debug_conversation_id: string @@ -426,7 +424,6 @@ export type AgentReferencingWorkflowsResponse = { } export type SandboxInfoResponse = { - session_id: string workspace_cwd: string } @@ -436,6 +433,16 @@ export type SandboxListResponse = { truncated?: boolean } +export type AgentSandboxDownloadPayload = { + caller_id: string + caller_type: 'build_draft' | 'conversation' + path: string +} + +export type SandboxDownloadResponse = { + url: string +} + export type SandboxReadResponse = { binary: boolean path: string @@ -444,15 +451,6 @@ export type SandboxReadResponse = { truncated: boolean } -export type AgentSandboxUploadPayload = { - conversation_id: string - path: string -} - -export type SandboxUploadResponse = { - url: string -} - export type AgentSkillUploadResponse = { manifest: SkillManifest skill: AgentUploadedSkillResponse @@ -822,8 +820,6 @@ export type AgentConfigSkillMarkdownResponse = { truncated: boolean } -export type AgentConfigDraftType = 'debug_build' | 'draft' - export type AgentDriveItemResponse = { created_at?: number | null file_kind: string @@ -939,6 +935,8 @@ export type AgentLogMessageItemResponse = { created_at?: number | null currency: string error?: string | null + feedback_enabled?: boolean + feedbacks?: Array from_account_id?: string | null from_end_user_id?: string | null id: string @@ -1227,6 +1225,8 @@ export type AgentSoulToolsConfig = { dify_tools?: Array } +export type AgentConfigDraftType = 'debug_build' | 'draft' + export type DeclaredOutputConfig = { array_item?: DeclaredArrayItem | null check?: DeclaredOutputCheckConfig | null @@ -1363,6 +1363,12 @@ export type AgentSuggestedQuestionsAfterAnswerModelConfig = { [key: string]: unknown } +export type AgentLogFeedbackResponse = { + content?: string | null + from_source: 'admin' | 'user' + rating: 'dislike' | 'like' +} + export type SimpleAccount = { email: string id: string @@ -1580,6 +1586,7 @@ export type AgentSoulModelSettings = { stop?: Array | null temperature?: number | null top_p?: number | null + [key: string]: unknown } export type AgentSandboxProviderConfig = { @@ -1876,6 +1883,8 @@ export type AgentKnowledgeMetadataCondition = { | '≠' | '≤' | '≥' + id?: string | null + metadata_id?: string | null name: string value?: string | Array | number | null } @@ -1892,7 +1901,7 @@ export type AgentAppPaginationWritable = { export type AgentAppDetailWithSiteWritable = { access_mode?: string | null - active_config_is_published?: boolean + access_ready?: boolean api_base_url?: string | null app_id?: string | null backing_app_id?: string | null @@ -2757,7 +2766,7 @@ export type PostAgentByAgentIdCopyResponse = PostAgentByAgentIdCopyResponses[keyof PostAgentByAgentIdCopyResponses] export type PostAgentByAgentIdDebugConversationRefreshData = { - body?: AgentDebugConversationRefreshPayload + body?: never path: { agent_id: string } @@ -3078,7 +3087,8 @@ export type GetAgentByAgentIdSandboxData = { agent_id: string } query: { - conversation_id: string + caller_id: string + caller_type: 'build_draft' | 'conversation' } url: '/agent/{agent_id}/sandbox' } @@ -3096,7 +3106,8 @@ export type GetAgentByAgentIdSandboxFilesData = { agent_id: string } query: { - conversation_id: string + caller_id: string + caller_type: 'build_draft' | 'conversation' path?: string } url: '/agent/{agent_id}/sandbox/files' @@ -3109,13 +3120,30 @@ export type GetAgentByAgentIdSandboxFilesResponses = { export type GetAgentByAgentIdSandboxFilesResponse = GetAgentByAgentIdSandboxFilesResponses[keyof GetAgentByAgentIdSandboxFilesResponses] +export type PostAgentByAgentIdSandboxFilesDownloadData = { + body: AgentSandboxDownloadPayload + path: { + agent_id: string + } + query?: never + url: '/agent/{agent_id}/sandbox/files/download' +} + +export type PostAgentByAgentIdSandboxFilesDownloadResponses = { + 200: SandboxDownloadResponse +} + +export type PostAgentByAgentIdSandboxFilesDownloadResponse = + PostAgentByAgentIdSandboxFilesDownloadResponses[keyof PostAgentByAgentIdSandboxFilesDownloadResponses] + export type GetAgentByAgentIdSandboxFilesReadData = { body?: never path: { agent_id: string } query: { - conversation_id: string + caller_id: string + caller_type: 'build_draft' | 'conversation' path: string } url: '/agent/{agent_id}/sandbox/files/read' @@ -3128,22 +3156,6 @@ export type GetAgentByAgentIdSandboxFilesReadResponses = { export type GetAgentByAgentIdSandboxFilesReadResponse = GetAgentByAgentIdSandboxFilesReadResponses[keyof GetAgentByAgentIdSandboxFilesReadResponses] -export type PostAgentByAgentIdSandboxFilesUploadData = { - body: AgentSandboxUploadPayload - path: { - agent_id: string - } - query?: never - url: '/agent/{agent_id}/sandbox/files/upload' -} - -export type PostAgentByAgentIdSandboxFilesUploadResponses = { - 200: SandboxUploadResponse -} - -export type PostAgentByAgentIdSandboxFilesUploadResponse = - PostAgentByAgentIdSandboxFilesUploadResponses[keyof PostAgentByAgentIdSandboxFilesUploadResponses] - export type PostAgentByAgentIdSkillsUploadData = { body: { file: Blob | File diff --git a/packages/contracts/generated/api/console/agent/zod.gen.ts b/packages/contracts/generated/api/console/agent/zod.gen.ts index 4146d261ffc..d3fc03ef81e 100644 --- a/packages/contracts/generated/api/console/agent/zod.gen.ts +++ b/packages/contracts/generated/api/console/agent/zod.gen.ts @@ -6,6 +6,7 @@ import * as z from 'zod' * AgentApiAccessResponse */ export const zAgentApiAccessResponse = z.object({ + access_ready: z.boolean(), api_key_count: z.int(), api_rph: z.int(), api_rpm: z.int(), @@ -196,10 +197,25 @@ export const zAgentPublishPayload = z.object({ * SandboxInfoResponse */ export const zSandboxInfoResponse = z.object({ - session_id: z.string(), workspace_cwd: z.string(), }) +/** + * AgentSandboxDownloadPayload + */ +export const zAgentSandboxDownloadPayload = z.object({ + caller_id: z.string().min(1), + caller_type: z.enum(['build_draft', 'conversation']), + path: z.string().min(1), +}) + +/** + * SandboxDownloadResponse + */ +export const zSandboxDownloadResponse = z.object({ + url: z.string(), +}) + /** * SandboxReadResponse */ @@ -211,21 +227,6 @@ export const zSandboxReadResponse = z.object({ truncated: z.boolean(), }) -/** - * AgentSandboxUploadPayload - */ -export const zAgentSandboxUploadPayload = z.object({ - conversation_id: z.string().min(1), - path: z.string().min(1), -}) - -/** - * SandboxUploadResponse - */ -export const zSandboxUploadResponse = z.object({ - url: z.string(), -}) - /** * AgentConfigSnapshotRestoreResponse */ @@ -374,7 +375,7 @@ export const zWorkflowPartial = z.object({ */ export const zAgentAppDetailWithSite = z.object({ access_mode: z.string().nullish(), - active_config_is_published: z.boolean().optional().default(false), + access_ready: z.boolean().optional().default(false), api_base_url: z.string().nullish(), app_id: z.string().nullish(), backing_app_id: z.string().nullish(), @@ -651,45 +652,6 @@ export const zAgentConfigSkillInspectResponse = z.object({ warnings: z.array(z.string()).optional(), }) -/** - * AgentConfigDraftType - * - * Editable Agent Soul draft workspace type. - */ -export const zAgentConfigDraftType = z.enum(['debug_build', 'draft']) - -/** - * AgentDebugConversationRefreshPayload - */ -export const zAgentDebugConversationRefreshPayload = z.object({ - draft_type: zAgentConfigDraftType.optional().default('debug_build'), -}) - -/** - * AgentConfigDraftSummaryResponse - */ -export const zAgentConfigDraftSummaryResponse = z.object({ - account_id: z.string().nullish(), - agent_id: z.string(), - base_snapshot_id: z.string().nullish(), - created_at: z.int().nullish(), - created_by: z.string().nullish(), - draft_type: zAgentConfigDraftType, - id: z.string(), - updated_at: z.int().nullish(), - updated_by: z.string().nullish(), -}) - -/** - * AgentPublishResponse - */ -export const zAgentPublishResponse = z.object({ - active_config_snapshot: zAgentConfigSnapshotSummaryResponse.nullish(), - active_config_snapshot_id: z.string(), - draft: zAgentConfigDraftSummaryResponse.nullish(), - result: z.string(), -}) - /** * AgentDriveItemResponse */ @@ -872,40 +834,6 @@ export const zAgentLogListResponse = z.object({ total: z.int(), }) -/** - * AgentLogMessageItemResponse - */ -export const zAgentLogMessageItemResponse = z.object({ - answer: z.string(), - answer_tokens: z.int(), - conversation_id: z.string(), - created_at: z.int().nullish(), - currency: z.string(), - error: z.string().nullish(), - from_account_id: z.string().nullish(), - from_end_user_id: z.string().nullish(), - id: z.string(), - latency: z.number(), - message_id: z.string(), - message_tokens: z.int(), - query: z.string(), - status: z.string(), - total_price: z.string(), - total_tokens: z.int(), - updated_at: z.int().nullish(), -}) - -/** - * AgentLogMessageListResponse - */ -export const zAgentLogMessageListResponse = z.object({ - data: z.array(zAgentLogMessageItemResponse), - has_more: z.boolean(), - limit: z.int(), - page: z.int(), - total: z.int(), -}) - export const zJsonValue = z .union([ z.string(), @@ -1279,6 +1207,38 @@ export const zAgentSoulPromptConfig = z.object({ system_prompt: z.string().optional().default(''), }) +/** + * AgentConfigDraftType + * + * Editable Agent Soul draft workspace type. + */ +export const zAgentConfigDraftType = z.enum(['debug_build', 'draft']) + +/** + * AgentConfigDraftSummaryResponse + */ +export const zAgentConfigDraftSummaryResponse = z.object({ + account_id: z.string().nullish(), + agent_id: z.string(), + base_snapshot_id: z.string().nullish(), + created_at: z.int().nullish(), + created_by: z.string().nullish(), + draft_type: zAgentConfigDraftType, + id: z.string(), + updated_at: z.int().nullish(), + updated_by: z.string().nullish(), +}) + +/** + * AgentPublishResponse + */ +export const zAgentPublishResponse = z.object({ + active_config_snapshot: zAgentConfigSnapshotSummaryResponse.nullish(), + active_config_snapshot_id: z.string(), + draft: zAgentConfigDraftSummaryResponse.nullish(), + result: z.string(), +}) + /** * AgentHumanContactConfig */ @@ -1373,6 +1333,51 @@ export const zAgentSuggestedQuestionsAfterAnswerFeatureConfig = z.object({ prompt: z.string().nullish(), }) +/** + * AgentLogFeedbackResponse + */ +export const zAgentLogFeedbackResponse = z.object({ + content: z.string().nullish(), + from_source: z.enum(['admin', 'user']), + rating: z.enum(['dislike', 'like']), +}) + +/** + * AgentLogMessageItemResponse + */ +export const zAgentLogMessageItemResponse = z.object({ + answer: z.string(), + answer_tokens: z.int(), + conversation_id: z.string(), + created_at: z.int().nullish(), + currency: z.string(), + error: z.string().nullish(), + feedback_enabled: z.boolean().optional().default(false), + feedbacks: z.array(zAgentLogFeedbackResponse).optional(), + from_account_id: z.string().nullish(), + from_end_user_id: z.string().nullish(), + id: z.string(), + latency: z.number(), + message_id: z.string(), + message_tokens: z.int(), + query: z.string(), + status: z.string(), + total_price: z.string(), + total_tokens: z.int(), + updated_at: z.int().nullish(), +}) + +/** + * AgentLogMessageListResponse + */ +export const zAgentLogMessageListResponse = z.object({ + data: z.array(zAgentLogMessageItemResponse), + has_more: z.boolean(), + limit: z.int(), + page: z.int(), + total: z.int(), +}) + /** * SimpleAccount */ @@ -2032,6 +2037,13 @@ export const zAgentModelResponseFormatConfig = z.object({ /** * AgentSoulModelSettings + * + * Model parameters for the Agent Soul model. + * + * Model plugins can declare arbitrary parameters via ``parameter_rules`` + * (e.g. Qwen/Tongyi's ``enable_thinking``) beyond the common OpenAI-style + * fields typed below, so extra keys must round-trip through persistence + * rather than being dropped. */ export const zAgentSoulModelSettings = z.object({ frequency_penalty: z.number().nullish(), @@ -2352,6 +2364,14 @@ export const zAgentKnowledgeRetrievalConfig = z.object({ /** * AgentKnowledgeMetadataCondition + * + * One manual metadata filter clause. + * + * ``id`` and ``metadata_id`` are UI-only bookkeeping the composer sends on + * every save (a stable row key and a reference to the selected metadata + * field). They are persisted here for round-tripping the composer's draft + * state but are stripped before building the Agent runtime request, whose + * DTO only accepts ``name``/``comparison_operator``/``value``. */ export const zAgentKnowledgeMetadataCondition = z.object({ comparison_operator: z.enum([ @@ -2374,6 +2394,8 @@ export const zAgentKnowledgeMetadataCondition = z.object({ '≤', '≥', ]), + id: z.string().nullish(), + metadata_id: z.string().nullish(), name: z.string().min(1).max(255), value: z.union([z.string(), z.array(z.string()), z.number()]).nullish(), }) @@ -2494,6 +2516,7 @@ export const zComposerSavePayload = z.object({ * AgentAppComposerResponse */ export const zAgentAppComposerResponse = z.object({ + active_config_is_published: z.boolean(), active_config_snapshot: zAgentConfigSnapshotSummaryResponse.nullish(), agent: zAgentComposerAgentResponse, agent_soul: zAgentSoulConfig, @@ -2528,7 +2551,7 @@ export const zAgentConfigSnapshotDetailResponse = z.object({ * ValueSourceType * * ValueSourceType records whether the value comes from a static setting - * in form definiton, or a variable while the workflow is running. + * in form definition, or a variable while the workflow is running. */ export const zValueSourceType = z.enum(['constant', 'variable']) @@ -2730,7 +2753,7 @@ export const zAppDetailSiteResponseWritable = z.object({ */ export const zAgentAppDetailWithSiteWritable = z.object({ access_mode: z.string().nullish(), - active_config_is_published: z.boolean().optional().default(false), + access_ready: z.boolean().optional().default(false), api_base_url: z.string().nullish(), app_id: z.string().nullish(), backing_app_id: z.string().nullish(), @@ -3264,8 +3287,6 @@ export const zPostAgentByAgentIdCopyPath = z.object({ */ export const zPostAgentByAgentIdCopyResponse = zAgentAppDetailWithSite -export const zPostAgentByAgentIdDebugConversationRefreshBody = zAgentDebugConversationRefreshPayload - export const zPostAgentByAgentIdDebugConversationRefreshPath = z.object({ agent_id: z.uuid(), }) @@ -3472,7 +3493,8 @@ export const zGetAgentByAgentIdSandboxPath = z.object({ }) export const zGetAgentByAgentIdSandboxQuery = z.object({ - conversation_id: z.string().min(1), + caller_id: z.string().min(1), + caller_type: z.enum(['build_draft', 'conversation']), }) /** @@ -3485,7 +3507,8 @@ export const zGetAgentByAgentIdSandboxFilesPath = z.object({ }) export const zGetAgentByAgentIdSandboxFilesQuery = z.object({ - conversation_id: z.string().min(1), + caller_id: z.string().min(1), + caller_type: z.enum(['build_draft', 'conversation']), path: z.string().optional().default('.'), }) @@ -3494,12 +3517,24 @@ export const zGetAgentByAgentIdSandboxFilesQuery = z.object({ */ export const zGetAgentByAgentIdSandboxFilesResponse = zSandboxListResponse +export const zPostAgentByAgentIdSandboxFilesDownloadBody = zAgentSandboxDownloadPayload + +export const zPostAgentByAgentIdSandboxFilesDownloadPath = z.object({ + agent_id: z.uuid(), +}) + +/** + * Download URL returned + */ +export const zPostAgentByAgentIdSandboxFilesDownloadResponse = zSandboxDownloadResponse + export const zGetAgentByAgentIdSandboxFilesReadPath = z.object({ agent_id: z.uuid(), }) export const zGetAgentByAgentIdSandboxFilesReadQuery = z.object({ - conversation_id: z.string().min(1), + caller_id: z.string().min(1), + caller_type: z.enum(['build_draft', 'conversation']), path: z.string().min(1), }) @@ -3508,17 +3543,6 @@ export const zGetAgentByAgentIdSandboxFilesReadQuery = z.object({ */ export const zGetAgentByAgentIdSandboxFilesReadResponse = zSandboxReadResponse -export const zPostAgentByAgentIdSandboxFilesUploadBody = zAgentSandboxUploadPayload - -export const zPostAgentByAgentIdSandboxFilesUploadPath = z.object({ - agent_id: z.uuid(), -}) - -/** - * Uploaded - */ -export const zPostAgentByAgentIdSandboxFilesUploadResponse = zSandboxUploadResponse - export const zPostAgentByAgentIdSkillsUploadBody = z.object({ file: z.custom((value) => value instanceof Blob || value instanceof File), }) diff --git a/packages/contracts/generated/api/console/apps/orpc.gen.ts b/packages/contracts/generated/api/console/apps/orpc.gen.ts index d479394c028..ede274a1f60 100644 --- a/packages/contracts/generated/api/console/apps/orpc.gen.ts +++ b/packages/contracts/generated/api/console/apps/orpc.gen.ts @@ -272,11 +272,11 @@ import { zGetAppsByAppIdWorkflowsTriggersWebhookResponse, zGetAppsByResourceIdApiKeysPath, zGetAppsByResourceIdApiKeysResponse, - zGetAppsByServerIdServerRefreshPath, - zGetAppsByServerIdServerRefreshResponse, zGetAppsImportsByAppIdCheckDependenciesPath, zGetAppsImportsByAppIdCheckDependenciesResponse, zGetAppsQuery, + zGetAppsRecentQuery, + zGetAppsRecentResponse, zGetAppsResponse, zGetAppsStarredQuery, zGetAppsStarredResponse, @@ -373,6 +373,8 @@ import { zPostAppsByAppIdPublishToCreatorsPlatformResponse, zPostAppsByAppIdServerBody, zPostAppsByAppIdServerPath, + zPostAppsByAppIdServerRefreshPath, + zPostAppsByAppIdServerRefreshResponse, zPostAppsByAppIdServerResponse, zPostAppsByAppIdSiteAccessTokenResetPath, zPostAppsByAppIdSiteAccessTokenResetResponse, @@ -404,9 +406,9 @@ import { zPostAppsByAppIdWorkflowCommentsByCommentIdResolveResponse, zPostAppsByAppIdWorkflowCommentsPath, zPostAppsByAppIdWorkflowCommentsResponse, - zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadBody, - zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadPath, - zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadResponse, + zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadBody, + zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadPath, + zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadResponse, zPostAppsByAppIdWorkflowRunsTasksByTaskIdStopPath, zPostAppsByAppIdWorkflowRunsTasksByTaskIdStopResponse, zPostAppsByAppIdWorkflowsByWorkflowIdRestorePath, @@ -551,9 +553,31 @@ export const imports = { } /** - * Get applications starred by the current account + * Return the lightweight app cards needed by the Explore home page + * + * Get recently modified apps for the home Continue Work section */ export const get2 = oc + .route({ + description: 'Get recently modified apps for the home Continue Work section', + inputStructure: 'detailed', + method: 'GET', + operationId: 'getAppsRecent', + path: '/apps/recent', + summary: 'Return the lightweight app cards needed by the Explore home page', + tags: ['console'], + }) + .input(z.object({ query: zGetAppsRecentQuery.optional() })) + .output(zGetAppsRecentResponse) + +export const recent = { + get: get2, +} + +/** + * Get applications starred by the current account + */ +export const get3 = oc .route({ description: 'Get applications starred by the current account', inputStructure: 'detailed', @@ -566,7 +590,7 @@ export const get2 = oc .output(zGetAppsStarredResponse) export const starred = { - get: get2, + get: get3, } /** @@ -597,7 +621,7 @@ export const workflows = { * * Get advanced chat workflow runs count statistics */ -export const get3 = oc +export const get4 = oc .route({ description: 'Get advanced chat workflow runs count statistics', inputStructure: 'detailed', @@ -616,7 +640,7 @@ export const get3 = oc .output(zGetAppsByAppIdAdvancedChatWorkflowRunsCountResponse) export const count = { - get: get3, + get: get4, } /** @@ -624,7 +648,7 @@ export const count = { * * Get advanced chat workflow run list */ -export const get4 = oc +export const get5 = oc .route({ description: 'Get advanced chat workflow run list', inputStructure: 'detailed', @@ -643,7 +667,7 @@ export const get4 = oc .output(zGetAppsByAppIdAdvancedChatWorkflowRunsResponse) export const workflowRuns = { - get: get4, + get: get5, count, } @@ -839,7 +863,7 @@ export const advancedChat = { workflows: workflows2, } -export const get5 = oc +export const get6 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -856,10 +880,10 @@ export const get5 = oc .output(zGetAppsByAppIdAgentConfigFilesByNameDownloadResponse) export const download = { - get: get5, + get: get6, } -export const get6 = oc +export const get7 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -876,7 +900,7 @@ export const get6 = oc .output(zGetAppsByAppIdAgentConfigFilesByNamePreviewResponse) export const preview2 = { - get: get6, + get: get7, } export const delete_ = oc @@ -901,7 +925,7 @@ export const byName = { preview: preview2, } -export const get7 = oc +export const get8 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -936,12 +960,12 @@ export const post9 = oc .output(zPostAppsByAppIdAgentConfigFilesResponse) export const files = { - get: get7, + get: get8, post: post9, byName, } -export const get8 = oc +export const get9 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -958,7 +982,7 @@ export const get8 = oc .output(zGetAppsByAppIdAgentConfigManifestResponse) export const manifest = { - get: get8, + get: get9, } export const post10 = oc @@ -983,7 +1007,7 @@ export const upload = { post: post10, } -export const get9 = oc +export const get10 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1000,10 +1024,10 @@ export const get9 = oc .output(zGetAppsByAppIdAgentConfigSkillsByNameDownloadResponse) export const download2 = { - get: get9, + get: get10, } -export const get10 = oc +export const get11 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1015,10 +1039,10 @@ export const get10 = oc .output(zGetAppsByAppIdAgentConfigSkillsByNameFilesContentResponse) export const content = { - get: get10, + get: get11, } -export const get11 = oc +export const get12 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1035,10 +1059,10 @@ export const get11 = oc .output(zGetAppsByAppIdAgentConfigSkillsByNameFilesDownloadResponse) export const download3 = { - get: get11, + get: get12, } -export const get12 = oc +export const get13 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1055,7 +1079,7 @@ export const get12 = oc .output(zGetAppsByAppIdAgentConfigSkillsByNameFilesPreviewResponse) export const preview3 = { - get: get12, + get: get13, } export const files2 = { @@ -1064,7 +1088,7 @@ export const files2 = { preview: preview3, } -export const get13 = oc +export const get14 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1081,7 +1105,7 @@ export const get13 = oc .output(zGetAppsByAppIdAgentConfigSkillsByNameInspectResponse) export const inspect = { - get: get13, + get: get14, } export const delete2 = oc @@ -1107,7 +1131,7 @@ export const byName2 = { inspect, } -export const get14 = oc +export const get15 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1124,7 +1148,7 @@ export const get14 = oc .output(zGetAppsByAppIdAgentConfigSkillsResponse) export const skills = { - get: get14, + get: get15, upload, byName: byName2, } @@ -1138,7 +1162,7 @@ export const config = { /** * Time-limited external signed URL for one drive value (no streaming proxy) */ -export const get15 = oc +export const get16 = oc .route({ description: 'Time-limited external signed URL for one drive value (no streaming proxy)', inputStructure: 'detailed', @@ -1156,13 +1180,13 @@ export const get15 = oc .output(zGetAppsByAppIdAgentDriveFilesDownloadResponse) export const download4 = { - get: get15, + get: get16, } /** * Truncated text preview of one drive value (binary-safe; SKILL.md is the main case) */ -export const get16 = oc +export const get17 = oc .route({ description: 'Truncated text preview of one drive value (binary-safe; SKILL.md is the main case)', @@ -1181,13 +1205,13 @@ export const get16 = oc .output(zGetAppsByAppIdAgentDriveFilesPreviewResponse) export const preview4 = { - get: get16, + get: get17, } /** * List agent drive entries (read-only inspector; one endpoint for both tabs) */ -export const get17 = oc +export const get18 = oc .route({ description: 'List agent drive entries (read-only inspector; one endpoint for both tabs)', inputStructure: 'detailed', @@ -1205,7 +1229,7 @@ export const get17 = oc .output(zGetAppsByAppIdAgentDriveFilesResponse) export const files3 = { - get: get17, + get: get18, download: download4, preview: preview4, } @@ -1213,7 +1237,7 @@ export const files3 = { /** * Inspect one drive-backed skill for slash-menu hover/detail UI */ -export const get18 = oc +export const get19 = oc .route({ description: 'Inspect one drive-backed skill for slash-menu hover/detail UI', inputStructure: 'detailed', @@ -1231,7 +1255,7 @@ export const get18 = oc .output(zGetAppsByAppIdAgentDriveSkillsBySkillPathInspectResponse) export const inspect2 = { - get: get18, + get: get19, } export const bySkillPath = { @@ -1241,7 +1265,7 @@ export const bySkillPath = { /** * List drive-backed skills for the bound agent */ -export const get19 = oc +export const get20 = oc .route({ description: 'List drive-backed skills for the bound agent', inputStructure: 'detailed', @@ -1259,7 +1283,7 @@ export const get19 = oc .output(zGetAppsByAppIdAgentDriveSkillsResponse) export const skills2 = { - get: get19, + get: get20, bySkillPath, } @@ -1323,7 +1347,7 @@ export const files4 = { * * Get agent execution logs for an application */ -export const get20 = oc +export const get21 = oc .route({ description: 'Get agent execution logs for an application', inputStructure: 'detailed', @@ -1337,7 +1361,7 @@ export const get20 = oc .output(zGetAppsByAppIdAgentLogsResponse) export const logs = { - get: get20, + get: get21, } /** @@ -1439,7 +1463,7 @@ export const agent = { /** * Get status of annotation reply action job */ -export const get21 = oc +export const get22 = oc .route({ description: 'Get status of annotation reply action job', inputStructure: 'detailed', @@ -1452,7 +1476,7 @@ export const get21 = oc .output(zGetAppsByAppIdAnnotationReplyByActionStatusByJobIdResponse) export const byJobId = { - get: get21, + get: get22, } export const status = { @@ -1491,7 +1515,7 @@ export const annotationReply = { /** * Get annotation settings for an app */ -export const get22 = oc +export const get23 = oc .route({ description: 'Get annotation settings for an app', inputStructure: 'detailed', @@ -1504,7 +1528,7 @@ export const get22 = oc .output(zGetAppsByAppIdAnnotationSettingResponse) export const annotationSetting = { - get: get22, + get: get23, } /** @@ -1557,7 +1581,7 @@ export const batchImport = { /** * Get status of batch import job */ -export const get23 = oc +export const get24 = oc .route({ description: 'Get status of batch import job', inputStructure: 'detailed', @@ -1570,7 +1594,7 @@ export const get23 = oc .output(zGetAppsByAppIdAnnotationsBatchImportStatusByJobIdResponse) export const byJobId2 = { - get: get23, + get: get24, } export const batchImportStatus = { @@ -1580,7 +1604,7 @@ export const batchImportStatus = { /** * Get count of message annotations for the app */ -export const get24 = oc +export const get25 = oc .route({ description: 'Get count of message annotations for the app', inputStructure: 'detailed', @@ -1593,13 +1617,13 @@ export const get24 = oc .output(zGetAppsByAppIdAnnotationsCountResponse) export const count2 = { - get: get24, + get: get25, } /** * Export all annotations for an app with CSV injection protection */ -export const get25 = oc +export const get26 = oc .route({ description: 'Export all annotations for an app with CSV injection protection', inputStructure: 'detailed', @@ -1612,13 +1636,13 @@ export const get25 = oc .output(zGetAppsByAppIdAnnotationsExportResponse) export const export_ = { - get: get25, + get: get26, } /** * Get hit histories for an annotation */ -export const get26 = oc +export const get27 = oc .route({ description: 'Get hit histories for an annotation', inputStructure: 'detailed', @@ -1636,7 +1660,7 @@ export const get26 = oc .output(zGetAppsByAppIdAnnotationsByAnnotationIdHitHistoriesResponse) export const hitHistories = { - get: get26, + get: get27, } export const delete5 = oc @@ -1692,7 +1716,7 @@ export const delete6 = oc /** * Get annotations for an app with pagination */ -export const get27 = oc +export const get28 = oc .route({ description: 'Get annotations for an app with pagination', inputStructure: 'detailed', @@ -1729,7 +1753,7 @@ export const post18 = oc export const annotations = { delete: delete6, - get: get27, + get: get28, post: post18, batchImport, batchImportStatus, @@ -1797,7 +1821,7 @@ export const delete7 = oc /** * Get chat conversation details */ -export const get28 = oc +export const get29 = oc .route({ description: 'Get chat conversation details', inputStructure: 'detailed', @@ -1811,13 +1835,13 @@ export const get28 = oc export const byConversationId = { delete: delete7, - get: get28, + get: get29, } /** * Get chat conversations with pagination, filtering and summary */ -export const get29 = oc +export const get30 = oc .route({ description: 'Get chat conversations with pagination, filtering and summary', inputStructure: 'detailed', @@ -1835,14 +1859,14 @@ export const get29 = oc .output(zGetAppsByAppIdChatConversationsResponse) export const chatConversations = { - get: get29, + get: get30, byConversationId, } /** * Get suggested questions for a message */ -export const get30 = oc +export const get31 = oc .route({ description: 'Get suggested questions for a message', inputStructure: 'detailed', @@ -1855,7 +1879,7 @@ export const get30 = oc .output(zGetAppsByAppIdChatMessagesByMessageIdSuggestedQuestionsResponse) export const suggestedQuestions = { - get: get30, + get: get31, } export const byMessageId = { @@ -1888,7 +1912,7 @@ export const byTaskId = { /** * Get chat messages for a conversation with pagination */ -export const get31 = oc +export const get32 = oc .route({ description: 'Get chat messages for a conversation with pagination', inputStructure: 'detailed', @@ -1903,7 +1927,7 @@ export const get31 = oc .output(zGetAppsByAppIdChatMessagesResponse) export const chatMessages = { - get: get31, + get: get32, byMessageId, byTaskId, } @@ -1927,7 +1951,7 @@ export const delete8 = oc /** * Get completion conversation details with messages */ -export const get32 = oc +export const get33 = oc .route({ description: 'Get completion conversation details with messages', inputStructure: 'detailed', @@ -1941,13 +1965,13 @@ export const get32 = oc export const byConversationId2 = { delete: delete8, - get: get32, + get: get33, } /** * Get completion conversations with pagination and filtering */ -export const get33 = oc +export const get34 = oc .route({ description: 'Get completion conversations with pagination and filtering', inputStructure: 'detailed', @@ -1965,7 +1989,7 @@ export const get33 = oc .output(zGetAppsByAppIdCompletionConversationsResponse) export const completionConversations = { - get: get33, + get: get34, byConversationId: byConversationId2, } @@ -2020,7 +2044,7 @@ export const completionMessages = { /** * Get conversation variables for an application */ -export const get34 = oc +export const get35 = oc .route({ description: 'Get conversation variables for an application', inputStructure: 'detailed', @@ -2038,7 +2062,7 @@ export const get34 = oc .output(zGetAppsByAppIdConversationVariablesResponse) export const conversationVariables = { - get: get34, + get: get35, } /** @@ -2099,7 +2123,7 @@ export const copy = { * * Export application configuration as DSL */ -export const get35 = oc +export const get36 = oc .route({ description: 'Export application configuration as DSL', inputStructure: 'detailed', @@ -2115,13 +2139,13 @@ export const get35 = oc .output(zGetAppsByAppIdExportResponse) export const export2 = { - get: get35, + get: get36, } /** * Export user feedback data for Google Sheets */ -export const get36 = oc +export const get37 = oc .route({ description: 'Export user feedback data for Google Sheets', inputStructure: 'detailed', @@ -2139,7 +2163,7 @@ export const get36 = oc .output(zGetAppsByAppIdFeedbacksExportResponse) export const export3 = { - get: get36, + get: get37, } /** @@ -2184,7 +2208,7 @@ export const icon = { /** * Get message details by ID */ -export const get37 = oc +export const get38 = oc .route({ description: 'Get message details by ID', inputStructure: 'detailed', @@ -2197,7 +2221,7 @@ export const get37 = oc .output(zGetAppsByAppIdMessagesByMessageIdResponse) export const byMessageId2 = { - get: get37, + get: get38, } export const messages = { @@ -2266,10 +2290,29 @@ export const publishToCreatorsPlatform = { post: post30, } +/** + * Refresh MCP server configuration and regenerate server code + */ +export const post31 = oc + .route({ + description: 'Refresh MCP server configuration and regenerate server code', + inputStructure: 'detailed', + method: 'POST', + operationId: 'postAppsByAppIdServerRefresh', + path: '/apps/{app_id}/server/refresh', + tags: ['console'], + }) + .input(z.object({ params: zPostAppsByAppIdServerRefreshPath })) + .output(zPostAppsByAppIdServerRefreshResponse) + +export const refresh = { + post: post31, +} + /** * Get MCP server configuration for an application */ -export const get38 = oc +export const get39 = oc .route({ description: 'Get MCP server configuration for an application', inputStructure: 'detailed', @@ -2284,7 +2327,7 @@ export const get38 = oc /** * Create MCP server configuration for an application */ -export const post31 = oc +export const post32 = oc .route({ description: 'Create MCP server configuration for an application', inputStructure: 'detailed', @@ -2313,15 +2356,16 @@ export const put = oc .output(zPutAppsByAppIdServerResponse) export const server = { - get: get38, - post: post31, + get: get39, + post: post32, put, + refresh, } /** * Reset access token for application site */ -export const post32 = oc +export const post33 = oc .route({ description: 'Reset access token for application site', inputStructure: 'detailed', @@ -2334,13 +2378,13 @@ export const post32 = oc .output(zPostAppsByAppIdSiteAccessTokenResetResponse) export const accessTokenReset = { - post: post32, + post: post33, } /** * Update application site configuration */ -export const post33 = oc +export const post34 = oc .route({ description: 'Update application site configuration', inputStructure: 'detailed', @@ -2353,14 +2397,14 @@ export const post33 = oc .output(zPostAppsByAppIdSiteResponse) export const site = { - post: post33, + post: post34, accessTokenReset, } /** * Enable or disable app site */ -export const post34 = oc +export const post35 = oc .route({ description: 'Enable or disable app site', inputStructure: 'detailed', @@ -2373,7 +2417,7 @@ export const post34 = oc .output(zPostAppsByAppIdSiteEnableResponse) export const siteEnable = { - post: post34, + post: post35, } /** @@ -2394,7 +2438,7 @@ export const delete9 = oc /** * Star an application for the current account */ -export const post35 = oc +export const post36 = oc .route({ description: 'Star an application for the current account', inputStructure: 'detailed', @@ -2408,13 +2452,13 @@ export const post35 = oc export const star = { delete: delete9, - post: post35, + post: post36, } /** * Get average response time statistics for an application */ -export const get39 = oc +export const get40 = oc .route({ description: 'Get average response time statistics for an application', inputStructure: 'detailed', @@ -2432,13 +2476,13 @@ export const get39 = oc .output(zGetAppsByAppIdStatisticsAverageResponseTimeResponse) export const averageResponseTime = { - get: get39, + get: get40, } /** * Get average session interaction statistics for an application */ -export const get40 = oc +export const get41 = oc .route({ description: 'Get average session interaction statistics for an application', inputStructure: 'detailed', @@ -2456,13 +2500,13 @@ export const get40 = oc .output(zGetAppsByAppIdStatisticsAverageSessionInteractionsResponse) export const averageSessionInteractions = { - get: get40, + get: get41, } /** * Get daily conversation statistics for an application */ -export const get41 = oc +export const get42 = oc .route({ description: 'Get daily conversation statistics for an application', inputStructure: 'detailed', @@ -2480,13 +2524,13 @@ export const get41 = oc .output(zGetAppsByAppIdStatisticsDailyConversationsResponse) export const dailyConversations = { - get: get41, + get: get42, } /** * Get daily terminal/end-user statistics for an application */ -export const get42 = oc +export const get43 = oc .route({ description: 'Get daily terminal/end-user statistics for an application', inputStructure: 'detailed', @@ -2504,13 +2548,13 @@ export const get42 = oc .output(zGetAppsByAppIdStatisticsDailyEndUsersResponse) export const dailyEndUsers = { - get: get42, + get: get43, } /** * Get daily message statistics for an application */ -export const get43 = oc +export const get44 = oc .route({ description: 'Get daily message statistics for an application', inputStructure: 'detailed', @@ -2528,13 +2572,13 @@ export const get43 = oc .output(zGetAppsByAppIdStatisticsDailyMessagesResponse) export const dailyMessages = { - get: get43, + get: get44, } /** * Get daily token cost statistics for an application */ -export const get44 = oc +export const get45 = oc .route({ description: 'Get daily token cost statistics for an application', inputStructure: 'detailed', @@ -2552,13 +2596,13 @@ export const get44 = oc .output(zGetAppsByAppIdStatisticsTokenCostsResponse) export const tokenCosts = { - get: get44, + get: get45, } /** * Get tokens per second statistics for an application */ -export const get45 = oc +export const get46 = oc .route({ description: 'Get tokens per second statistics for an application', inputStructure: 'detailed', @@ -2576,13 +2620,13 @@ export const get45 = oc .output(zGetAppsByAppIdStatisticsTokensPerSecondResponse) export const tokensPerSecond = { - get: get45, + get: get46, } /** * Get user satisfaction rate statistics for an application */ -export const get46 = oc +export const get47 = oc .route({ description: 'Get user satisfaction rate statistics for an application', inputStructure: 'detailed', @@ -2600,7 +2644,7 @@ export const get46 = oc .output(zGetAppsByAppIdStatisticsUserSatisfactionRateResponse) export const userSatisfactionRate = { - get: get46, + get: get47, } export const statistics = { @@ -2617,7 +2661,7 @@ export const statistics = { /** * Get available TTS voices for a specific language */ -export const get47 = oc +export const get48 = oc .route({ description: 'Get available TTS voices for a specific language', inputStructure: 'detailed', @@ -2635,13 +2679,13 @@ export const get47 = oc .output(zGetAppsByAppIdTextToAudioVoicesResponse) export const voices = { - get: get47, + get: get48, } /** * Convert text to speech for chat messages */ -export const post36 = oc +export const post37 = oc .route({ description: 'Convert text to speech for chat messages', inputStructure: 'detailed', @@ -2656,7 +2700,7 @@ export const post36 = oc .output(zPostAppsByAppIdTextToAudioResponse) export const textToAudio = { - post: post36, + post: post37, voices, } @@ -2665,7 +2709,7 @@ export const textToAudio = { * * Get app tracing configuration */ -export const get48 = oc +export const get49 = oc .route({ description: 'Get app tracing configuration', inputStructure: 'detailed', @@ -2681,7 +2725,7 @@ export const get48 = oc /** * Update app tracing configuration */ -export const post37 = oc +export const post38 = oc .route({ description: 'Update app tracing configuration', inputStructure: 'detailed', @@ -2694,8 +2738,8 @@ export const post37 = oc .output(zPostAppsByAppIdTraceResponse) export const trace = { - get: get48, - post: post37, + get: get49, + post: post38, } /** @@ -2725,7 +2769,7 @@ export const delete10 = oc /** * Get tracing configuration for an application */ -export const get49 = oc +export const get50 = oc .route({ description: 'Get tracing configuration for an application', inputStructure: 'detailed', @@ -2764,7 +2808,7 @@ export const patch = oc * * Create a new tracing configuration for an application */ -export const post38 = oc +export const post39 = oc .route({ description: 'Create a new tracing configuration for an application', inputStructure: 'detailed', @@ -2782,15 +2826,15 @@ export const post38 = oc export const traceConfig = { delete: delete10, - get: get49, + get: get50, patch, - post: post38, + post: post39, } /** * Update app trigger (enable/disable) */ -export const post39 = oc +export const post40 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -2808,13 +2852,13 @@ export const post39 = oc .output(zPostAppsByAppIdTriggerEnableResponse) export const triggerEnable = { - post: post39, + post: post40, } /** * Get app triggers list */ -export const get50 = oc +export const get51 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2827,7 +2871,7 @@ export const get50 = oc .output(zGetAppsByAppIdTriggersResponse) export const triggers = { - get: get50, + get: get51, } /** @@ -2835,7 +2879,7 @@ export const triggers = { * * Get workflow application execution logs */ -export const get51 = oc +export const get52 = oc .route({ description: 'Get workflow application execution logs', inputStructure: 'detailed', @@ -2854,7 +2898,7 @@ export const get51 = oc .output(zGetAppsByAppIdWorkflowAppLogsResponse) export const workflowAppLogs = { - get: get51, + get: get52, } /** @@ -2862,7 +2906,7 @@ export const workflowAppLogs = { * * Get workflow archived execution logs */ -export const get52 = oc +export const get53 = oc .route({ description: 'Get workflow archived execution logs', inputStructure: 'detailed', @@ -2881,7 +2925,7 @@ export const get52 = oc .output(zGetAppsByAppIdWorkflowArchivedLogsResponse) export const workflowArchivedLogs = { - get: get52, + get: get53, } /** @@ -2889,7 +2933,7 @@ export const workflowArchivedLogs = { * * Get workflow runs count statistics */ -export const get53 = oc +export const get54 = oc .route({ description: 'Get workflow runs count statistics', inputStructure: 'detailed', @@ -2908,7 +2952,7 @@ export const get53 = oc .output(zGetAppsByAppIdWorkflowRunsCountResponse) export const count3 = { - get: get53, + get: get54, } /** @@ -2916,7 +2960,7 @@ export const count3 = { * * Stop running workflow task */ -export const post40 = oc +export const post41 = oc .route({ description: 'Stop running workflow task', inputStructure: 'detailed', @@ -2930,7 +2974,7 @@ export const post40 = oc .output(zPostAppsByAppIdWorkflowRunsTasksByTaskIdStopResponse) export const stop3 = { - post: post40, + post: post41, } export const byTaskId3 = { @@ -2944,7 +2988,7 @@ export const tasks = { /** * Generate a download URL for an archived workflow run. */ -export const get54 = oc +export const get55 = oc .route({ description: 'Generate a download URL for an archived workflow run.', inputStructure: 'detailed', @@ -2957,7 +3001,7 @@ export const get54 = oc .output(zGetAppsByAppIdWorkflowRunsByRunIdExportResponse) export const export4 = { - get: get54, + get: get55, } /** @@ -2965,7 +3009,7 @@ export const export4 = { * * Get workflow run node execution list */ -export const get55 = oc +export const get56 = oc .route({ description: 'Get workflow run node execution list', inputStructure: 'detailed', @@ -2979,7 +3023,7 @@ export const get55 = oc .output(zGetAppsByAppIdWorkflowRunsByRunIdNodeExecutionsResponse) export const nodeExecutions = { - get: get55, + get: get56, } /** @@ -2987,7 +3031,7 @@ export const nodeExecutions = { * * Get workflow run detail */ -export const get56 = oc +export const get57 = oc .route({ description: 'Get workflow run detail', inputStructure: 'detailed', @@ -3001,15 +3045,40 @@ export const get56 = oc .output(zGetAppsByAppIdWorkflowRunsByRunIdResponse) export const byRunId = { - get: get56, + get: get57, export: export4, nodeExecutions, } +/** + * Create a ToolFile from one workflow Agent Binding file and return its download URL + */ +export const post42 = oc + .route({ + description: + 'Create a ToolFile from one workflow Agent Binding file and return its download URL', + inputStructure: 'detailed', + method: 'POST', + operationId: 'postAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownload', + path: '/apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/download', + tags: ['console'], + }) + .input( + z.object({ + body: zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadBody, + params: zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadPath, + }), + ) + .output(zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadResponse) + +export const download5 = { + post: post42, +} + /** * Read a text/binary preview file in a workflow Agent node sandbox */ -export const get57 = oc +export const get58 = oc .route({ description: 'Read a text/binary preview file in a workflow Agent node sandbox', inputStructure: 'detailed', @@ -3027,37 +3096,13 @@ export const get57 = oc .output(zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesReadResponse) export const read = { - get: get57, -} - -/** - * Upload one workflow Agent sandbox file and return a signed download URL - */ -export const post41 = oc - .route({ - description: 'Upload one workflow Agent sandbox file and return a signed download URL', - inputStructure: 'detailed', - method: 'POST', - operationId: 'postAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUpload', - path: '/apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/upload', - tags: ['console'], - }) - .input( - z.object({ - body: zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadBody, - params: zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadPath, - }), - ) - .output(zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadResponse) - -export const upload3 = { - post: post41, + get: get58, } /** * List a directory in a workflow Agent node sandbox */ -export const get58 = oc +export const get59 = oc .route({ description: 'List a directory in a workflow Agent node sandbox', inputStructure: 'detailed', @@ -3069,16 +3114,15 @@ export const get58 = oc .input( z.object({ params: zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesPath, - query: - zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesQuery.optional(), + query: zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesQuery, }), ) .output(zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesResponse) export const files5 = { - get: get58, + get: get59, + download: download5, read, - upload: upload3, } export const sandbox = { @@ -3102,7 +3146,7 @@ export const byWorkflowRunId = { * * Get workflow run list */ -export const get59 = oc +export const get60 = oc .route({ description: 'Get workflow run list', inputStructure: 'detailed', @@ -3121,7 +3165,7 @@ export const get59 = oc .output(zGetAppsByAppIdWorkflowRunsResponse) export const workflowRuns2 = { - get: get59, + get: get60, count: count3, tasks, byRunId, @@ -3133,7 +3177,7 @@ export const workflowRuns2 = { * * Get all users in current tenant for mentions */ -export const get60 = oc +export const get61 = oc .route({ description: 'Get all users in current tenant for mentions', inputStructure: 'detailed', @@ -3147,7 +3191,7 @@ export const get60 = oc .output(zGetAppsByAppIdWorkflowCommentsMentionUsersResponse) export const mentionUsers = { - get: get60, + get: get61, } /** @@ -3202,7 +3246,7 @@ export const byReplyId = { * * Add a reply to a workflow comment */ -export const post42 = oc +export const post43 = oc .route({ description: 'Add a reply to a workflow comment', inputStructure: 'detailed', @@ -3222,7 +3266,7 @@ export const post42 = oc .output(zPostAppsByAppIdWorkflowCommentsByCommentIdRepliesResponse) export const replies = { - post: post42, + post: post43, byReplyId, } @@ -3231,7 +3275,7 @@ export const replies = { * * Resolve a workflow comment */ -export const post43 = oc +export const post44 = oc .route({ description: 'Resolve a workflow comment', inputStructure: 'detailed', @@ -3245,7 +3289,7 @@ export const post43 = oc .output(zPostAppsByAppIdWorkflowCommentsByCommentIdResolveResponse) export const resolve = { - post: post43, + post: post44, } /** @@ -3272,7 +3316,7 @@ export const delete12 = oc * * Get a specific workflow comment */ -export const get61 = oc +export const get62 = oc .route({ description: 'Get a specific workflow comment', inputStructure: 'detailed', @@ -3310,7 +3354,7 @@ export const put3 = oc export const byCommentId = { delete: delete12, - get: get61, + get: get62, put: put3, replies, resolve, @@ -3321,7 +3365,7 @@ export const byCommentId = { * * Get all comments for a workflow */ -export const get62 = oc +export const get63 = oc .route({ description: 'Get all comments for a workflow', inputStructure: 'detailed', @@ -3339,7 +3383,7 @@ export const get62 = oc * * Create a new workflow comment */ -export const post44 = oc +export const post45 = oc .route({ description: 'Create a new workflow comment', inputStructure: 'detailed', @@ -3359,8 +3403,8 @@ export const post44 = oc .output(zPostAppsByAppIdWorkflowCommentsResponse) export const comments = { - get: get62, - post: post44, + get: get63, + post: post45, mentionUsers, byCommentId, } @@ -3368,7 +3412,7 @@ export const comments = { /** * Get workflow average app interaction statistics */ -export const get63 = oc +export const get64 = oc .route({ description: 'Get workflow average app interaction statistics', inputStructure: 'detailed', @@ -3386,13 +3430,13 @@ export const get63 = oc .output(zGetAppsByAppIdWorkflowStatisticsAverageAppInteractionsResponse) export const averageAppInteractions = { - get: get63, + get: get64, } /** * Get workflow daily runs statistics */ -export const get64 = oc +export const get65 = oc .route({ description: 'Get workflow daily runs statistics', inputStructure: 'detailed', @@ -3410,13 +3454,13 @@ export const get64 = oc .output(zGetAppsByAppIdWorkflowStatisticsDailyConversationsResponse) export const dailyConversations2 = { - get: get64, + get: get65, } /** * Get workflow daily terminals statistics */ -export const get65 = oc +export const get66 = oc .route({ description: 'Get workflow daily terminals statistics', inputStructure: 'detailed', @@ -3434,13 +3478,13 @@ export const get65 = oc .output(zGetAppsByAppIdWorkflowStatisticsDailyTerminalsResponse) export const dailyTerminals = { - get: get65, + get: get66, } /** * Get workflow daily token cost statistics */ -export const get66 = oc +export const get67 = oc .route({ description: 'Get workflow daily token cost statistics', inputStructure: 'detailed', @@ -3458,7 +3502,7 @@ export const get66 = oc .output(zGetAppsByAppIdWorkflowStatisticsTokenCostsResponse) export const tokenCosts2 = { - get: get66, + get: get67, } export const statistics2 = { @@ -3478,7 +3522,7 @@ export const workflow = { * * Get default block configuration by type */ -export const get67 = oc +export const get68 = oc .route({ description: 'Get default block configuration by type', inputStructure: 'detailed', @@ -3497,7 +3541,7 @@ export const get67 = oc .output(zGetAppsByAppIdWorkflowsDefaultWorkflowBlockConfigsByBlockTypeResponse) export const byBlockType = { - get: get67, + get: get68, } /** @@ -3505,7 +3549,7 @@ export const byBlockType = { * * Get default block configurations for workflow */ -export const get68 = oc +export const get69 = oc .route({ description: 'Get default block configurations for workflow', inputStructure: 'detailed', @@ -3519,14 +3563,14 @@ export const get68 = oc .output(zGetAppsByAppIdWorkflowsDefaultWorkflowBlockConfigsResponse) export const defaultWorkflowBlockConfigs = { - get: get68, + get: get69, byBlockType, } /** * Get conversation variables for workflow */ -export const get69 = oc +export const get70 = oc .route({ description: 'Get conversation variables for workflow', inputStructure: 'detailed', @@ -3541,7 +3585,7 @@ export const get69 = oc /** * Update conversation variables for workflow draft */ -export const post45 = oc +export const post46 = oc .route({ description: 'Update conversation variables for workflow draft', inputStructure: 'detailed', @@ -3559,8 +3603,8 @@ export const post45 = oc .output(zPostAppsByAppIdWorkflowsDraftConversationVariablesResponse) export const conversationVariables2 = { - get: get69, - post: post45, + get: get70, + post: post46, } /** @@ -3568,7 +3612,7 @@ export const conversationVariables2 = { * * Get environment variables for workflow */ -export const get70 = oc +export const get71 = oc .route({ description: 'Get environment variables for workflow', inputStructure: 'detailed', @@ -3584,7 +3628,7 @@ export const get70 = oc /** * Update environment variables for workflow draft */ -export const post46 = oc +export const post47 = oc .route({ description: 'Update environment variables for workflow draft', inputStructure: 'detailed', @@ -3602,14 +3646,14 @@ export const post46 = oc .output(zPostAppsByAppIdWorkflowsDraftEnvironmentVariablesResponse) export const environmentVariables = { - get: get70, - post: post46, + get: get71, + post: post47, } /** * Update draft workflow features */ -export const post47 = oc +export const post48 = oc .route({ description: 'Update draft workflow features', inputStructure: 'detailed', @@ -3627,7 +3671,7 @@ export const post47 = oc .output(zPostAppsByAppIdWorkflowsDraftFeaturesResponse) export const features = { - post: post47, + post: post48, } /** @@ -3635,7 +3679,7 @@ export const features = { * * Test human input delivery for workflow */ -export const post48 = oc +export const post49 = oc .route({ description: 'Test human input delivery for workflow', inputStructure: 'detailed', @@ -3654,7 +3698,7 @@ export const post48 = oc .output(zPostAppsByAppIdWorkflowsDraftHumanInputNodesByNodeIdDeliveryTestResponse) export const deliveryTest = { - post: post48, + post: post49, } /** @@ -3662,7 +3706,7 @@ export const deliveryTest = { * * Get human input form preview for workflow */ -export const post49 = oc +export const post50 = oc .route({ description: 'Get human input form preview for workflow', inputStructure: 'detailed', @@ -3681,7 +3725,7 @@ export const post49 = oc .output(zPostAppsByAppIdWorkflowsDraftHumanInputNodesByNodeIdFormPreviewResponse) export const preview5 = { - post: post49, + post: post50, } /** @@ -3689,7 +3733,7 @@ export const preview5 = { * * Submit human input form preview for workflow */ -export const post50 = oc +export const post51 = oc .route({ description: 'Submit human input form preview for workflow', inputStructure: 'detailed', @@ -3708,7 +3752,7 @@ export const post50 = oc .output(zPostAppsByAppIdWorkflowsDraftHumanInputNodesByNodeIdFormRunResponse) export const run5 = { - post: post50, + post: post51, } export const form2 = { @@ -3734,7 +3778,7 @@ export const humanInput2 = { * * Run draft workflow iteration node */ -export const post51 = oc +export const post52 = oc .route({ description: 'Run draft workflow iteration node', inputStructure: 'detailed', @@ -3753,7 +3797,7 @@ export const post51 = oc .output(zPostAppsByAppIdWorkflowsDraftIterationNodesByNodeIdRunResponse) export const run6 = { - post: post51, + post: post52, } export const byNodeId6 = { @@ -3773,7 +3817,7 @@ export const iteration2 = { * * Run draft workflow loop node */ -export const post52 = oc +export const post53 = oc .route({ description: 'Run draft workflow loop node', inputStructure: 'detailed', @@ -3792,7 +3836,7 @@ export const post52 = oc .output(zPostAppsByAppIdWorkflowsDraftLoopNodesByNodeIdRunResponse) export const run7 = { - post: post52, + post: post53, } export const byNodeId7 = { @@ -3807,7 +3851,7 @@ export const loop2 = { nodes: nodes6, } -export const get71 = oc +export const get72 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3821,10 +3865,10 @@ export const get71 = oc .output(zGetAppsByAppIdWorkflowsDraftNodesByNodeIdAgentComposerCandidatesResponse) export const candidates = { - get: get71, + get: get72, } -export const post53 = oc +export const post54 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -3841,10 +3885,10 @@ export const post53 = oc .output(zPostAppsByAppIdWorkflowsDraftNodesByNodeIdAgentComposerCopyFromRosterResponse) export const copyFromRoster = { - post: post53, + post: post54, } -export const post54 = oc +export const post55 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -3861,10 +3905,10 @@ export const post54 = oc .output(zPostAppsByAppIdWorkflowsDraftNodesByNodeIdAgentComposerImpactResponse) export const impact = { - post: post54, + post: post55, } -export const post55 = oc +export const post56 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -3881,10 +3925,10 @@ export const post55 = oc .output(zPostAppsByAppIdWorkflowsDraftNodesByNodeIdAgentComposerSaveToRosterResponse) export const saveToRoster = { - post: post55, + post: post56, } -export const post56 = oc +export const post57 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -3901,10 +3945,10 @@ export const post56 = oc .output(zPostAppsByAppIdWorkflowsDraftNodesByNodeIdAgentComposerValidateResponse) export const validate = { - post: post56, + post: post57, } -export const get72 = oc +export const get73 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3937,7 +3981,7 @@ export const put4 = oc .output(zPutAppsByAppIdWorkflowsDraftNodesByNodeIdAgentComposerResponse) export const agentComposer = { - get: get72, + get: get73, put: put4, candidates, copyFromRoster, @@ -3949,7 +3993,7 @@ export const agentComposer = { /** * Get last run result for draft workflow node */ -export const get73 = oc +export const get74 = oc .route({ description: 'Get last run result for draft workflow node', inputStructure: 'detailed', @@ -3962,7 +4006,7 @@ export const get73 = oc .output(zGetAppsByAppIdWorkflowsDraftNodesByNodeIdLastRunResponse) export const lastRun = { - get: get73, + get: get74, } /** @@ -3970,7 +4014,7 @@ export const lastRun = { * * Run draft workflow node */ -export const post57 = oc +export const post58 = oc .route({ description: 'Run draft workflow node', inputStructure: 'detailed', @@ -3989,7 +4033,7 @@ export const post57 = oc .output(zPostAppsByAppIdWorkflowsDraftNodesByNodeIdRunResponse) export const run8 = { - post: post57, + post: post58, } /** @@ -3997,7 +4041,7 @@ export const run8 = { * * Poll for trigger events and execute single node when event arrives */ -export const post58 = oc +export const post59 = oc .route({ description: 'Poll for trigger events and execute single node when event arrives', inputStructure: 'detailed', @@ -4011,7 +4055,7 @@ export const post58 = oc .output(zPostAppsByAppIdWorkflowsDraftNodesByNodeIdTriggerRunResponse) export const run9 = { - post: post58, + post: post59, } export const trigger = { @@ -4037,7 +4081,7 @@ export const delete13 = oc /** * Get variables for a specific node */ -export const get74 = oc +export const get75 = oc .route({ description: 'Get variables for a specific node', inputStructure: 'detailed', @@ -4051,7 +4095,7 @@ export const get74 = oc export const variables = { delete: delete13, - get: get74, + get: get75, } export const byNodeId8 = { @@ -4071,7 +4115,7 @@ export const nodes7 = { * * Run draft workflow */ -export const post59 = oc +export const post60 = oc .route({ description: 'Run draft workflow', inputStructure: 'detailed', @@ -4090,13 +4134,13 @@ export const post59 = oc .output(zPostAppsByAppIdWorkflowsDraftRunResponse) export const run10 = { - post: post59, + post: post60, } /** * Server-Sent Events stream of inspector deltas for a draft workflow run. */ -export const get75 = oc +export const get76 = oc .route({ description: 'Server-Sent Events stream of inspector deltas for a draft workflow run.', inputStructure: 'detailed', @@ -4109,13 +4153,13 @@ export const get75 = oc .output(zGetAppsByAppIdWorkflowsDraftRunsByRunIdNodeOutputsEventsResponse) export const events = { - get: get75, + get: get76, } /** * Full value for one declared output, including signed download URL for files. */ -export const get76 = oc +export const get77 = oc .route({ description: 'Full value for one declared output, including signed download URL for files.', inputStructure: 'detailed', @@ -4132,7 +4176,7 @@ export const get76 = oc .output(zGetAppsByAppIdWorkflowsDraftRunsByRunIdNodeOutputsByNodeIdByOutputNamePreviewResponse) export const preview6 = { - get: get76, + get: get77, } export const byOutputName = { @@ -4142,7 +4186,7 @@ export const byOutputName = { /** * One node's declared outputs for a draft workflow run. */ -export const get77 = oc +export const get78 = oc .route({ description: "One node's declared outputs for a draft workflow run.", inputStructure: 'detailed', @@ -4155,14 +4199,14 @@ export const get77 = oc .output(zGetAppsByAppIdWorkflowsDraftRunsByRunIdNodeOutputsByNodeIdResponse) export const byNodeId9 = { - get: get77, + get: get78, byOutputName, } /** * Snapshot of every node's declared outputs for a draft workflow run. */ -export const get78 = oc +export const get79 = oc .route({ description: "Snapshot of every node's declared outputs for a draft workflow run.", inputStructure: 'detailed', @@ -4175,7 +4219,7 @@ export const get78 = oc .output(zGetAppsByAppIdWorkflowsDraftRunsByRunIdNodeOutputsResponse) export const nodeOutputs = { - get: get78, + get: get79, events, byNodeId: byNodeId9, } @@ -4191,7 +4235,7 @@ export const runs = { /** * Get system variables for workflow */ -export const get79 = oc +export const get80 = oc .route({ description: 'Get system variables for workflow', inputStructure: 'detailed', @@ -4204,7 +4248,7 @@ export const get79 = oc .output(zGetAppsByAppIdWorkflowsDraftSystemVariablesResponse) export const systemVariables = { - get: get79, + get: get80, } /** @@ -4212,7 +4256,7 @@ export const systemVariables = { * * Poll for trigger events and execute full workflow when event arrives */ -export const post60 = oc +export const post61 = oc .route({ description: 'Poll for trigger events and execute full workflow when event arrives', inputStructure: 'detailed', @@ -4231,7 +4275,7 @@ export const post60 = oc .output(zPostAppsByAppIdWorkflowsDraftTriggerRunResponse) export const run11 = { - post: post60, + post: post61, } /** @@ -4239,7 +4283,7 @@ export const run11 = { * * Full workflow debug when the start node is a trigger */ -export const post61 = oc +export const post62 = oc .route({ description: 'Full workflow debug when the start node is a trigger', inputStructure: 'detailed', @@ -4258,7 +4302,7 @@ export const post61 = oc .output(zPostAppsByAppIdWorkflowsDraftTriggerRunAllResponse) export const runAll = { - post: post61, + post: post62, } export const trigger2 = { @@ -4304,7 +4348,7 @@ export const delete14 = oc /** * Get a specific workflow variable */ -export const get80 = oc +export const get81 = oc .route({ description: 'Get a specific workflow variable', inputStructure: 'detailed', @@ -4338,7 +4382,7 @@ export const patch2 = oc export const byVariableId = { delete: delete14, - get: get80, + get: get81, patch: patch2, reset, } @@ -4364,7 +4408,7 @@ export const delete15 = oc * * Get draft workflow variables */ -export const get81 = oc +export const get82 = oc .route({ description: 'Get draft workflow variables', inputStructure: 'detailed', @@ -4384,7 +4428,7 @@ export const get81 = oc export const variables2 = { delete: delete15, - get: get81, + get: get82, byVariableId, } @@ -4393,7 +4437,7 @@ export const variables2 = { * * Get draft workflow for an application */ -export const get82 = oc +export const get83 = oc .route({ description: 'Get draft workflow for an application', inputStructure: 'detailed', @@ -4411,7 +4455,7 @@ export const get82 = oc * * Sync draft workflow configuration */ -export const post62 = oc +export const post63 = oc .route({ description: 'Sync draft workflow configuration', inputStructure: 'detailed', @@ -4430,8 +4474,8 @@ export const post62 = oc .output(zPostAppsByAppIdWorkflowsDraftResponse) export const draft2 = { - get: get82, - post: post62, + get: get83, + post: post63, conversationVariables: conversationVariables2, environmentVariables, features, @@ -4451,7 +4495,7 @@ export const draft2 = { * * Get published workflow for an application */ -export const get83 = oc +export const get84 = oc .route({ description: 'Get published workflow for an application', inputStructure: 'detailed', @@ -4467,7 +4511,7 @@ export const get83 = oc /** * Publish workflow */ -export const post63 = oc +export const post64 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -4485,14 +4529,14 @@ export const post63 = oc .output(zPostAppsByAppIdWorkflowsPublishResponse) export const publish = { - get: get83, - post: post63, + get: get84, + post: post64, } /** * Server-Sent Events stream of inspector deltas for a published workflow run. */ -export const get84 = oc +export const get85 = oc .route({ description: 'Server-Sent Events stream of inspector deltas for a published workflow run.', inputStructure: 'detailed', @@ -4505,13 +4549,13 @@ export const get84 = oc .output(zGetAppsByAppIdWorkflowsPublishedRunsByRunIdNodeOutputsEventsResponse) export const events2 = { - get: get84, + get: get85, } /** * Full value for one declared output of a published run. */ -export const get85 = oc +export const get86 = oc .route({ description: 'Full value for one declared output of a published run.', inputStructure: 'detailed', @@ -4532,7 +4576,7 @@ export const get85 = oc ) export const preview7 = { - get: get85, + get: get86, } export const byOutputName2 = { @@ -4542,7 +4586,7 @@ export const byOutputName2 = { /** * One node's declared outputs for a published workflow run. */ -export const get86 = oc +export const get87 = oc .route({ description: "One node's declared outputs for a published workflow run.", inputStructure: 'detailed', @@ -4555,14 +4599,14 @@ export const get86 = oc .output(zGetAppsByAppIdWorkflowsPublishedRunsByRunIdNodeOutputsByNodeIdResponse) export const byNodeId10 = { - get: get86, + get: get87, byOutputName: byOutputName2, } /** * Snapshot of every node's declared outputs for a published workflow run. */ -export const get87 = oc +export const get88 = oc .route({ description: "Snapshot of every node's declared outputs for a published workflow run.", inputStructure: 'detailed', @@ -4575,7 +4619,7 @@ export const get87 = oc .output(zGetAppsByAppIdWorkflowsPublishedRunsByRunIdNodeOutputsResponse) export const nodeOutputs2 = { - get: get87, + get: get88, events: events2, byNodeId: byNodeId10, } @@ -4595,7 +4639,7 @@ export const published = { /** * Get webhook trigger for a node */ -export const get88 = oc +export const get89 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4613,7 +4657,7 @@ export const get88 = oc .output(zGetAppsByAppIdWorkflowsTriggersWebhookResponse) export const webhook = { - get: get88, + get: get89, } export const triggers2 = { @@ -4623,7 +4667,7 @@ export const triggers2 = { /** * Restore a published workflow version into the draft workflow */ -export const post64 = oc +export const post65 = oc .route({ description: 'Restore a published workflow version into the draft workflow', inputStructure: 'detailed', @@ -4636,7 +4680,7 @@ export const post64 = oc .output(zPostAppsByAppIdWorkflowsByWorkflowIdRestoreResponse) export const restore = { - post: post64, + post: post65, } /** @@ -4689,7 +4733,7 @@ export const byWorkflowId = { * * Get all published workflows for an application */ -export const get89 = oc +export const get90 = oc .route({ description: 'Get all published workflows for an application', inputStructure: 'detailed', @@ -4708,7 +4752,7 @@ export const get89 = oc .output(zGetAppsByAppIdWorkflowsResponse) export const workflows3 = { - get: get89, + get: get90, defaultWorkflowBlockConfigs, draft: draft2, publish, @@ -4741,7 +4785,7 @@ export const delete17 = oc * * Get application details */ -export const get90 = oc +export const get91 = oc .route({ description: 'Get application details', inputStructure: 'detailed', @@ -4774,7 +4818,7 @@ export const put6 = oc export const byAppId2 = { delete: delete17, - get: get90, + get: get91, put: put6, advancedChat, agent, @@ -4843,7 +4887,7 @@ export const byApiKeyId = { * * Get all API keys for an app */ -export const get91 = oc +export const get92 = oc .route({ description: 'Get all API keys for an app', inputStructure: 'detailed', @@ -4861,7 +4905,7 @@ export const get91 = oc * * Create a new API key for an app */ -export const post65 = oc +export const post66 = oc .route({ description: 'Create a new API key for an app', inputStructure: 'detailed', @@ -4876,8 +4920,8 @@ export const post65 = oc .output(zPostAppsByResourceIdApiKeysResponse) export const apiKeys = { - get: get91, - post: post65, + get: get92, + post: post66, byApiKeyId, } @@ -4885,33 +4929,6 @@ export const byResourceId = { apiKeys, } -/** - * Refresh MCP server configuration and regenerate server code - */ -export const get92 = oc - .route({ - description: 'Refresh MCP server configuration and regenerate server code', - inputStructure: 'detailed', - method: 'GET', - operationId: 'getAppsByServerIdServerRefresh', - path: '/apps/{server_id}/server/refresh', - tags: ['console'], - }) - .input(z.object({ params: zGetAppsByServerIdServerRefreshPath })) - .output(zGetAppsByServerIdServerRefreshResponse) - -export const refresh = { - get: get92, -} - -export const server2 = { - refresh, -} - -export const byServerId = { - server: server2, -} - /** * Get app list * @@ -4935,7 +4952,7 @@ export const get93 = oc * * Create a new application */ -export const post66 = oc +export const post67 = oc .route({ description: 'Create a new application', inputStructure: 'detailed', @@ -4951,13 +4968,13 @@ export const post66 = oc export const apps = { get: get93, - post: post66, + post: post67, imports, + recent, starred, workflows, byAppId: byAppId2, byResourceId, - byServerId, } export const contract = { diff --git a/packages/contracts/generated/api/console/apps/types.gen.ts b/packages/contracts/generated/api/console/apps/types.gen.ts index abf5b16538a..03d2a09d35c 100644 --- a/packages/contracts/generated/api/console/apps/types.gen.ts +++ b/packages/contracts/generated/api/console/apps/types.gen.ts @@ -80,6 +80,10 @@ export type CheckDependenciesResult = { leaked_dependencies?: Array } +export type RecentAppListResponse = { + data: Array +} + export type WorkflowOnlineUsersPayload = { app_ids?: Array } @@ -861,6 +865,15 @@ export type SandboxListResponse = { truncated?: boolean } +export type WorkflowAgentSandboxDownloadPayload = { + node_execution_id: string + path: string +} + +export type SandboxDownloadResponse = { + url: string +} + export type SandboxReadResponse = { binary: boolean path: string @@ -869,15 +882,6 @@ export type SandboxReadResponse = { truncated: boolean } -export type WorkflowAgentSandboxUploadPayload = { - node_execution_id?: string | null - path: string -} - -export type SandboxUploadResponse = { - url: string -} - export type WorkflowCommentBasicList = { data: Array } @@ -1000,6 +1004,7 @@ export type WorkflowResponse = { updated_at: number updated_by?: SimpleAccountResponse | null version: string + version_number?: number | null } export type SyncDraftWorkflowPayload = { @@ -1007,9 +1012,7 @@ export type SyncDraftWorkflowPayload = { conversation_variables?: Array<{ [key: string]: unknown }> - environment_variables?: Array<{ - [key: string]: unknown - }> + environment_variable_patch?: SyncEnvironmentVariablePatchPayload | null features: { [key: string]: unknown } @@ -1038,7 +1041,9 @@ export type EnvironmentVariableListResponse = { } export type EnvironmentVariableUpdatePayload = { + deleted_environment_variable_ids?: Array environment_variables: Array + patch?: boolean } export type WorkflowFeaturesPayload = { @@ -1408,6 +1413,20 @@ export type PluginDependency = { value: Github | Marketplace | Package } +export type RecentAppResponse = { + author_name?: string | null + icon?: string | null + icon_background?: string | null + icon_type?: IconType | null + readonly icon_url: string | null + id: string + maintainer?: string | null + mode: 'advanced-chat' | 'agent-chat' | 'chat' | 'completion' | 'workflow' + name: string + permission_keys?: Array + updated_at: number +} + export type WorkflowOnlineUsersByApp = { app_id: string users: Array @@ -1959,6 +1978,13 @@ export type PipelineVariableResponse = { variable: string } +export type SyncEnvironmentVariablePatchPayload = { + deleted_environment_variable_ids?: Array + environment_variables?: Array<{ + [key: string]: unknown + }> +} + export type ConversationVariableItemPayload = { description?: string | null id?: string | null @@ -2821,6 +2847,7 @@ export type AgentSoulModelSettings = { stop?: Array | null temperature?: number | null top_p?: number | null + [key: string]: unknown } export type AgentSandboxProviderConfig = { @@ -3072,6 +3099,8 @@ export type AgentKnowledgeMetadataCondition = { | '≠' | '≤' | '≥' + id?: string | null + metadata_id?: string | null name: string value?: string | Array | number | null } @@ -3114,6 +3143,10 @@ export type AppDetailWithSiteWritable = { workflow?: WorkflowPartial | null } +export type RecentAppListResponseWritable = { + data: Array +} + export type GeneratedAppResponseWritable = JsonValue export type WorkflowCommentBasicListWritable = { @@ -3196,6 +3229,19 @@ export type AppDetailSiteResponseWritable = { use_icon_as_answer_icon?: boolean | null } +export type RecentAppResponseWritable = { + author_name?: string | null + icon?: string | null + icon_background?: string | null + icon_type?: IconType | null + id: string + maintainer?: string | null + mode: 'advanced-chat' | 'agent-chat' | 'chat' | 'completion' | 'workflow' + name: string + permission_keys?: Array + updated_at: number +} + export type WorkflowCommentBasicWritable = { content: string created_at?: number | null @@ -3356,6 +3402,21 @@ export type PostAppsImportsByImportIdConfirmResponses = { export type PostAppsImportsByImportIdConfirmResponse = PostAppsImportsByImportIdConfirmResponses[keyof PostAppsImportsByImportIdConfirmResponses] +export type GetAppsRecentData = { + body?: never + path?: never + query?: { + limit?: number + } + url: '/apps/recent' +} + +export type GetAppsRecentResponses = { + 200: RecentAppListResponse +} + +export type GetAppsRecentResponse = GetAppsRecentResponses[keyof GetAppsRecentResponses] + export type GetAppsStarredData = { body?: never path?: never @@ -4960,6 +5021,27 @@ export type PutAppsByAppIdServerResponses = { export type PutAppsByAppIdServerResponse = PutAppsByAppIdServerResponses[keyof PutAppsByAppIdServerResponses] +export type PostAppsByAppIdServerRefreshData = { + body?: never + path: { + app_id: string + } + query?: never + url: '/apps/{app_id}/server/refresh' +} + +export type PostAppsByAppIdServerRefreshErrors = { + 403: unknown + 404: unknown +} + +export type PostAppsByAppIdServerRefreshResponses = { + 200: AppMcpServerResponse +} + +export type PostAppsByAppIdServerRefreshResponse = + PostAppsByAppIdServerRefreshResponses[keyof PostAppsByAppIdServerRefreshResponses] + export type PostAppsByAppIdSiteData = { body: AppSiteUpdatePayload path: { @@ -5608,8 +5690,8 @@ export type GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFi node_id: string workflow_run_id: string } - query?: { - node_execution_id?: string + query: { + node_execution_id: string path?: string } url: '/apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files' @@ -5622,6 +5704,25 @@ export type GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFi export type GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesResponse = GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesResponses[keyof GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesResponses] +export type PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadData = { + body: WorkflowAgentSandboxDownloadPayload + path: { + app_id: string + node_id: string + workflow_run_id: string + } + query?: never + url: '/apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/download' +} + +export type PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadResponses = + { + 200: SandboxDownloadResponse + } + +export type PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadResponse = + PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadResponses[keyof PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadResponses] + export type GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesReadData = { body?: never path: { @@ -5630,7 +5731,7 @@ export type GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFi workflow_run_id: string } query: { - node_execution_id?: string + node_execution_id: string path: string } url: '/apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/read' @@ -5643,25 +5744,6 @@ export type GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFi export type GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesReadResponse = GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesReadResponses[keyof GetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesReadResponses] -export type PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadData = { - body: WorkflowAgentSandboxUploadPayload - path: { - app_id: string - node_id: string - workflow_run_id: string - } - query?: never - url: '/apps/{app_id}/workflow-runs/{workflow_run_id}/agent-nodes/{node_id}/sandbox/files/upload' -} - -export type PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadResponses = - { - 200: SandboxUploadResponse - } - -export type PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadResponse = - PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadResponses[keyof PostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadResponses] - export type GetAppsByAppIdWorkflowCommentsData = { body?: never path: { @@ -6951,24 +7033,3 @@ export type DeleteAppsByResourceIdApiKeysByApiKeyIdResponses = { export type DeleteAppsByResourceIdApiKeysByApiKeyIdResponse = DeleteAppsByResourceIdApiKeysByApiKeyIdResponses[keyof DeleteAppsByResourceIdApiKeysByApiKeyIdResponses] - -export type GetAppsByServerIdServerRefreshData = { - body?: never - path: { - server_id: string - } - query?: never - url: '/apps/{server_id}/server/refresh' -} - -export type GetAppsByServerIdServerRefreshErrors = { - 403: unknown - 404: unknown -} - -export type GetAppsByServerIdServerRefreshResponses = { - 200: AppMcpServerResponse -} - -export type GetAppsByServerIdServerRefreshResponse = - GetAppsByServerIdServerRefreshResponses[keyof GetAppsByServerIdServerRefreshResponses] diff --git a/packages/contracts/generated/api/console/apps/zod.gen.ts b/packages/contracts/generated/api/console/apps/zod.gen.ts index a41672798e6..12f765715d6 100644 --- a/packages/contracts/generated/api/console/apps/zod.gen.ts +++ b/packages/contracts/generated/api/console/apps/zod.gen.ts @@ -545,6 +545,21 @@ export const zWorkflowRunExportResponse = z.object({ status: z.string(), }) +/** + * WorkflowAgentSandboxDownloadPayload + */ +export const zWorkflowAgentSandboxDownloadPayload = z.object({ + node_execution_id: z.string().min(1), + path: z.string().min(1), +}) + +/** + * SandboxDownloadResponse + */ +export const zSandboxDownloadResponse = z.object({ + url: z.string(), +}) + /** * SandboxReadResponse */ @@ -556,21 +571,6 @@ export const zSandboxReadResponse = z.object({ truncated: z.boolean(), }) -/** - * WorkflowAgentSandboxUploadPayload - */ -export const zWorkflowAgentSandboxUploadPayload = z.object({ - node_execution_id: z.string().nullish(), - path: z.string().min(1), -}) - -/** - * SandboxUploadResponse - */ -export const zSandboxUploadResponse = z.object({ - url: z.string(), -}) - /** * WorkflowCommentCreatePayload */ @@ -651,18 +651,6 @@ export const zDefaultBlockConfigsResponse = z.array(z.record(z.string(), z.unkno */ export const zDefaultBlockConfigResponse = z.record(z.string(), z.unknown()) -/** - * SyncDraftWorkflowPayload - */ -export const zSyncDraftWorkflowPayload = z.object({ - _is_collaborative: z.boolean().optional().default(false), - conversation_variables: z.array(z.record(z.string(), z.unknown())).optional(), - environment_variables: z.array(z.record(z.string(), z.unknown())).optional(), - features: z.record(z.string(), z.unknown()), - graph: z.record(z.string(), z.unknown()), - hash: z.string().nullish(), -}) - /** * SyncDraftWorkflowResponse */ @@ -1076,6 +1064,30 @@ export const zAppImportResponse = z.object({ warnings: z.array(zDslImportWarning).optional(), }) +/** + * RecentAppResponse + */ +export const zRecentAppResponse = z.object({ + author_name: z.string().nullish(), + icon: z.string().nullish(), + icon_background: z.string().nullish(), + icon_type: zIconType.nullish(), + icon_url: z.string().nullable(), + id: z.string(), + maintainer: z.string().nullish(), + mode: z.enum(['advanced-chat', 'agent-chat', 'chat', 'completion', 'workflow']), + name: z.string(), + permission_keys: z.array(z.string()).optional(), + updated_at: z.int(), +}) + +/** + * RecentAppListResponse + */ +export const zRecentAppListResponse = z.object({ + data: z.array(zRecentAppResponse), +}) + export const zJsonValue = z .union([ z.string(), @@ -2017,6 +2029,7 @@ export const zWorkflowResponse = z.object({ updated_at: z.int(), updated_by: zSimpleAccountResponse.nullish(), version: z.string(), + version_number: z.int().nullish(), }) /** @@ -2029,6 +2042,26 @@ export const zWorkflowPaginationResponse = z.object({ page: z.int(), }) +/** + * SyncEnvironmentVariablePatchPayload + */ +export const zSyncEnvironmentVariablePatchPayload = z.object({ + deleted_environment_variable_ids: z.array(z.string()).optional(), + environment_variables: z.array(z.record(z.string(), z.unknown())).optional(), +}) + +/** + * SyncDraftWorkflowPayload + */ +export const zSyncDraftWorkflowPayload = z.object({ + _is_collaborative: z.boolean().optional().default(false), + conversation_variables: z.array(z.record(z.string(), z.unknown())).optional(), + environment_variable_patch: zSyncEnvironmentVariablePatchPayload.nullish(), + features: z.record(z.string(), z.unknown()), + graph: z.record(z.string(), z.unknown()), + hash: z.string().nullish(), +}) + /** * ConversationVariableItemPayload */ @@ -2085,7 +2118,9 @@ export const zEnvironmentVariableItemPayload = z.object({ * EnvironmentVariableUpdatePayload */ export const zEnvironmentVariableUpdatePayload = z.object({ + deleted_environment_variable_ids: z.array(z.string()).optional(), environment_variables: z.array(zEnvironmentVariableItemPayload), + patch: z.boolean().optional().default(false), }) /** @@ -3609,6 +3644,13 @@ export const zAgentModelResponseFormatConfig = z.object({ /** * AgentSoulModelSettings + * + * Model parameters for the Agent Soul model. + * + * Model plugins can declare arbitrary parameters via ``parameter_rules`` + * (e.g. Qwen/Tongyi's ``enable_thinking``) beyond the common OpenAI-style + * fields typed below, so extra keys must round-trip through persistence + * rather than being dropped. */ export const zAgentSoulModelSettings = z.object({ frequency_penalty: z.number().nullish(), @@ -3895,7 +3937,7 @@ export const zAgentKnowledgeRetrievalConfig = z.object({ * ValueSourceType * * ValueSourceType records whether the value comes from a static setting - * in form definiton, or a variable while the workflow is running. + * in form definition, or a variable while the workflow is running. */ export const zValueSourceType = z.enum(['constant', 'variable']) @@ -4014,6 +4056,14 @@ export const zMessageInfiniteScrollPaginationResponse = z.object({ /** * AgentKnowledgeMetadataCondition + * + * One manual metadata filter clause. + * + * ``id`` and ``metadata_id`` are UI-only bookkeeping the composer sends on + * every save (a stable row key and a reference to the selected metadata + * field). They are persisted here for round-tripping the composer's draft + * state but are stripped before building the Agent runtime request, whose + * DTO only accepts ``name``/``comparison_operator``/``value``. */ export const zAgentKnowledgeMetadataCondition = z.object({ comparison_operator: z.enum([ @@ -4036,6 +4086,8 @@ export const zAgentKnowledgeMetadataCondition = z.object({ '≤', '≥', ]), + id: z.string().nullish(), + metadata_id: z.string().nullish(), name: z.string().min(1).max(255), value: z.union([z.string(), z.array(z.string()), z.number()]).nullish(), }) @@ -4279,6 +4331,29 @@ export const zAppDetailWithSiteWritable = z.object({ workflow: zWorkflowPartial.nullish(), }) +/** + * RecentAppResponse + */ +export const zRecentAppResponseWritable = z.object({ + author_name: z.string().nullish(), + icon: z.string().nullish(), + icon_background: z.string().nullish(), + icon_type: zIconType.nullish(), + id: z.string(), + maintainer: z.string().nullish(), + mode: z.enum(['advanced-chat', 'agent-chat', 'chat', 'completion', 'workflow']), + name: z.string(), + permission_keys: z.array(z.string()).optional(), + updated_at: z.int(), +}) + +/** + * RecentAppListResponse + */ +export const zRecentAppListResponseWritable = z.object({ + data: z.array(zRecentAppResponseWritable), +}) + /** * AccountWithRoleResponse */ @@ -4442,6 +4517,15 @@ export const zPostAppsImportsByImportIdConfirmPath = z.object({ */ export const zPostAppsImportsByImportIdConfirmResponse = zImport +export const zGetAppsRecentQuery = z.object({ + limit: z.int().gte(1).lte(8).optional().default(8), +}) + +/** + * Success + */ +export const zGetAppsRecentResponse = zRecentAppListResponse + export const zGetAppsStarredQuery = z.object({ creator_ids: z.array(z.string()).optional(), is_created_by_me: z.boolean().optional(), @@ -5471,6 +5555,15 @@ export const zPutAppsByAppIdServerPath = z.object({ */ export const zPutAppsByAppIdServerResponse = zAppMcpServerResponse +export const zPostAppsByAppIdServerRefreshPath = z.object({ + app_id: z.uuid(), +}) + +/** + * MCP server refreshed successfully + */ +export const zPostAppsByAppIdServerRefreshResponse = zAppMcpServerResponse + export const zPostAppsByAppIdSiteBody = zAppSiteUpdatePayload export const zPostAppsByAppIdSitePath = z.object({ @@ -5875,7 +5968,7 @@ export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandbox export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesQuery = z.object({ - node_execution_id: z.string().optional(), + node_execution_id: z.string().min(1), path: z.string().optional().default('.'), }) @@ -5885,6 +5978,22 @@ export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandbox export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesResponse = zSandboxListResponse +export const zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadBody = + zWorkflowAgentSandboxDownloadPayload + +export const zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadPath = + z.object({ + app_id: z.uuid(), + node_id: z.string(), + workflow_run_id: z.uuid(), + }) + +/** + * Download URL returned + */ +export const zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesDownloadResponse = + zSandboxDownloadResponse + export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesReadPath = z.object({ app_id: z.uuid(), @@ -5894,7 +6003,7 @@ export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandbox export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesReadQuery = z.object({ - node_execution_id: z.string().optional(), + node_execution_id: z.string().min(1), path: z.string().min(1), }) @@ -5904,22 +6013,6 @@ export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandbox export const zGetAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesReadResponse = zSandboxReadResponse -export const zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadBody = - zWorkflowAgentSandboxUploadPayload - -export const zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadPath = - z.object({ - app_id: z.uuid(), - node_id: z.string(), - workflow_run_id: z.uuid(), - }) - -/** - * Uploaded - */ -export const zPostAppsByAppIdWorkflowRunsByWorkflowRunIdAgentNodesByNodeIdSandboxFilesUploadResponse = - zSandboxUploadResponse - export const zGetAppsByAppIdWorkflowCommentsPath = z.object({ app_id: z.uuid(), }) @@ -6716,12 +6809,3 @@ export const zDeleteAppsByResourceIdApiKeysByApiKeyIdPath = z.object({ * API key deleted successfully */ export const zDeleteAppsByResourceIdApiKeysByApiKeyIdResponse = z.void() - -export const zGetAppsByServerIdServerRefreshPath = z.object({ - server_id: z.uuid(), -}) - -/** - * MCP server refreshed successfully - */ -export const zGetAppsByServerIdServerRefreshResponse = zAppMcpServerResponse diff --git a/packages/contracts/generated/api/console/billing/types.gen.ts b/packages/contracts/generated/api/console/billing/types.gen.ts index dc7afb41d0f..7410402d897 100644 --- a/packages/contracts/generated/api/console/billing/types.gen.ts +++ b/packages/contracts/generated/api/console/billing/types.gen.ts @@ -16,6 +16,10 @@ export type BillingResponse = { [key: string]: unknown } +export type BillingSubscriptionResponse = { + url: string +} + export type GetBillingInvoicesData = { body?: never path?: never @@ -61,7 +65,7 @@ export type GetBillingSubscriptionData = { } export type GetBillingSubscriptionResponses = { - 200: BillingResponse + 200: BillingSubscriptionResponse } export type GetBillingSubscriptionResponse = diff --git a/packages/contracts/generated/api/console/billing/zod.gen.ts b/packages/contracts/generated/api/console/billing/zod.gen.ts index a25890905c0..87d75a92621 100644 --- a/packages/contracts/generated/api/console/billing/zod.gen.ts +++ b/packages/contracts/generated/api/console/billing/zod.gen.ts @@ -21,6 +21,13 @@ export const zPartnerTenantsPayload = z.object({ */ export const zBillingResponse = z.record(z.string(), z.unknown()) +/** + * BillingSubscriptionResponse + */ +export const zBillingSubscriptionResponse = z.object({ + url: z.string(), +}) + /** * Success */ @@ -45,4 +52,4 @@ export const zGetBillingSubscriptionQuery = z.object({ /** * Success */ -export const zGetBillingSubscriptionResponse = zBillingResponse +export const zGetBillingSubscriptionResponse = zBillingSubscriptionResponse diff --git a/packages/contracts/generated/api/console/datasets/types.gen.ts b/packages/contracts/generated/api/console/datasets/types.gen.ts index 0382f581700..41b885bf98f 100644 --- a/packages/contracts/generated/api/console/datasets/types.gen.ts +++ b/packages/contracts/generated/api/console/datasets/types.gen.ts @@ -1677,6 +1677,10 @@ export type PostDatasetsByDatasetIdDocumentsMetadataData = { url: '/datasets/{dataset_id}/documents/metadata' } +export type PostDatasetsByDatasetIdDocumentsMetadataErrors = { + 404: unknown +} + export type PostDatasetsByDatasetIdDocumentsMetadataResponses = { 204: void } diff --git a/packages/contracts/generated/api/console/datasets/zod.gen.ts b/packages/contracts/generated/api/console/datasets/zod.gen.ts index b809480eb51..8f617205b5b 100644 --- a/packages/contracts/generated/api/console/datasets/zod.gen.ts +++ b/packages/contracts/generated/api/console/datasets/zod.gen.ts @@ -44,6 +44,13 @@ export const zBatchImportPayload = z.object({ /** * ExternalDatasetCreatePayload + * + * Validated fields required to create an external dataset binding. + * + * The console controller owns HTTP concerns, but the service also needs this + * contract when creating the tenant-scoped dataset and external knowledge + * binding. Keep it outside controllers so service imports do not depend on + * Flask blueprint initialization. */ export const zExternalDatasetCreatePayload = z.object({ description: z.string().max(400).nullish(), @@ -861,7 +868,7 @@ export const zProcessRule = z.object({ * MetadataDetail */ export const zMetadataDetail = z.object({ - id: z.string(), + id: z.uuid(), name: z.string(), value: z.union([z.string(), z.int(), z.number()]).nullish(), }) @@ -870,7 +877,7 @@ export const zMetadataDetail = z.object({ * DocumentMetadataOperation */ export const zDocumentMetadataOperation = z.object({ - document_id: z.string(), + document_id: z.uuid(), metadata_list: z.array(zMetadataDetail), partial_update: z.boolean().optional().default(false), }) diff --git a/packages/contracts/generated/api/console/email-code-login/types.gen.ts b/packages/contracts/generated/api/console/email-code-login/types.gen.ts index a07f11dfd19..293fe010f7b 100644 --- a/packages/contracts/generated/api/console/email-code-login/types.gen.ts +++ b/packages/contracts/generated/api/console/email-code-login/types.gen.ts @@ -4,9 +4,10 @@ export type ClientOptions = { baseUrl: `${string}://${string}/console/api` | (string & {}) } -export type EmailPayload = { +export type EmailCodeSendPayload = { email: string language?: string | null + turnstile_token?: string | null } export type SimpleResultDataResponse = { @@ -27,7 +28,7 @@ export type SimpleResultResponse = { } export type PostEmailCodeLoginData = { - body: EmailPayload + body: EmailCodeSendPayload path?: never query?: never url: '/email-code-login' diff --git a/packages/contracts/generated/api/console/email-code-login/zod.gen.ts b/packages/contracts/generated/api/console/email-code-login/zod.gen.ts index af72ec33867..14c69a9076a 100644 --- a/packages/contracts/generated/api/console/email-code-login/zod.gen.ts +++ b/packages/contracts/generated/api/console/email-code-login/zod.gen.ts @@ -3,11 +3,12 @@ import * as z from 'zod' /** - * EmailPayload + * EmailCodeSendPayload */ -export const zEmailPayload = z.object({ +export const zEmailCodeSendPayload = z.object({ email: z.string(), language: z.string().nullish(), + turnstile_token: z.string().nullish(), }) /** @@ -36,7 +37,7 @@ export const zSimpleResultResponse = z.object({ result: z.string(), }) -export const zPostEmailCodeLoginBody = zEmailPayload +export const zPostEmailCodeLoginBody = zEmailCodeSendPayload /** * Success diff --git a/packages/contracts/generated/api/console/explore/types.gen.ts b/packages/contracts/generated/api/console/explore/types.gen.ts index 331815aaf56..55181dbc577 100644 --- a/packages/contracts/generated/api/console/explore/types.gen.ts +++ b/packages/contracts/generated/api/console/explore/types.gen.ts @@ -20,7 +20,7 @@ export type BannerListResponse = Array export type RecommendedAppResponse = { app?: RecommendedAppInfoResponse | null app_id: string - can_trial?: boolean | null + can_trial: boolean categories?: Array copyright?: string | null custom_disclaimer?: string | null @@ -31,7 +31,7 @@ export type RecommendedAppResponse = { } export type RecommendedAppDetailResponse = { - can_trial?: boolean | null + can_trial: boolean export_data: string icon?: string | null icon_background?: string | null @@ -41,12 +41,12 @@ export type RecommendedAppDetailResponse = { } export type BannerResponse = { - content: unknown - created_at?: string | null + content: BannerContentResponse + created_at: string id: string - link?: string | null + link: string sort: number - status: string + status: BannerStatus } export type RecommendedAppInfoResponse = { @@ -59,6 +59,15 @@ export type RecommendedAppInfoResponse = { name?: string | null } +export type BannerContentResponse = { + category: string + description: string + 'img-src': string + title: string +} + +export type BannerStatus = 'disabled' | 'enabled' + export type RecommendedAppListResponseWritable = { categories: Array recommended_apps: Array @@ -71,7 +80,7 @@ export type LearnDifyAppListResponseWritable = { export type RecommendedAppResponseWritable = { app?: RecommendedAppInfoResponseWritable | null app_id: string - can_trial?: boolean | null + can_trial: boolean categories?: Array copyright?: string | null custom_disclaimer?: string | null diff --git a/packages/contracts/generated/api/console/explore/zod.gen.ts b/packages/contracts/generated/api/console/explore/zod.gen.ts index b835f405436..efd87b24f91 100644 --- a/packages/contracts/generated/api/console/explore/zod.gen.ts +++ b/packages/contracts/generated/api/console/explore/zod.gen.ts @@ -6,7 +6,7 @@ import * as z from 'zod' * RecommendedAppDetailResponse */ export const zRecommendedAppDetailResponse = z.object({ - can_trial: z.boolean().nullish(), + can_trial: z.boolean(), export_data: z.string(), icon: z.string().nullish(), icon_background: z.string().nullish(), @@ -20,23 +20,6 @@ export const zRecommendedAppDetailResponse = z.object({ */ export const zRecommendedAppDetailNullableResponse = zRecommendedAppDetailResponse.nullable() -/** - * BannerResponse - */ -export const zBannerResponse = z.object({ - content: z.unknown(), - created_at: z.string().nullish(), - id: z.string(), - link: z.string().nullish(), - sort: z.int(), - status: z.string(), -}) - -/** - * BannerListResponse - */ -export const zBannerListResponse = z.array(zBannerResponse) - /** * RecommendedAppInfoResponse */ @@ -56,7 +39,7 @@ export const zRecommendedAppInfoResponse = z.object({ export const zRecommendedAppResponse = z.object({ app: zRecommendedAppInfoResponse.nullish(), app_id: z.string(), - can_trial: z.boolean().nullish(), + can_trial: z.boolean(), categories: z.array(z.string()).optional(), copyright: z.string().nullish(), custom_disclaimer: z.string().nullish(), @@ -81,6 +64,40 @@ export const zLearnDifyAppListResponse = z.object({ recommended_apps: z.array(zRecommendedAppResponse), }) +/** + * BannerContentResponse + */ +export const zBannerContentResponse = z.object({ + category: z.string(), + description: z.string(), + 'img-src': z.string().min(1), + title: z.string().min(1), +}) + +/** + * BannerStatus + * + * ExporleBanner status + */ +export const zBannerStatus = z.enum(['disabled', 'enabled']) + +/** + * BannerResponse + */ +export const zBannerResponse = z.object({ + content: zBannerContentResponse, + created_at: z.string(), + id: z.string(), + link: z.string(), + sort: z.int(), + status: zBannerStatus, +}) + +/** + * BannerListResponse + */ +export const zBannerListResponse = z.array(zBannerResponse) + /** * RecommendedAppInfoResponse */ @@ -99,7 +116,7 @@ export const zRecommendedAppInfoResponseWritable = z.object({ export const zRecommendedAppResponseWritable = z.object({ app: zRecommendedAppInfoResponseWritable.nullish(), app_id: z.string(), - can_trial: z.boolean().nullish(), + can_trial: z.boolean(), categories: z.array(z.string()).optional(), copyright: z.string().nullish(), custom_disclaimer: z.string().nullish(), diff --git a/packages/contracts/generated/api/console/features/types.gen.ts b/packages/contracts/generated/api/console/features/types.gen.ts index 52c6cf80402..bcc666d3f40 100644 --- a/packages/contracts/generated/api/console/features/types.gen.ts +++ b/packages/contracts/generated/api/console/features/types.gen.ts @@ -27,6 +27,12 @@ export type FeatureModel = { workspace_members: LicenseLimitationModel } +export type VectorSpaceLimitationModel = { + limit: number + size: number + usage_unknown?: boolean +} + export type LimitationModel = { limit: number size: number @@ -60,9 +66,11 @@ export type LicenseLimitationModel = { export type SubscriptionModel = { interval: string - plan: string + plan: CloudPlan } +export type CloudPlan = 'professional' | 'sandbox' | 'team' + export type GetFeaturesData = { body?: never path?: never @@ -84,7 +92,7 @@ export type GetFeaturesVectorSpaceData = { } export type GetFeaturesVectorSpaceResponses = { - 200: LimitationModel + 200: VectorSpaceLimitationModel } export type GetFeaturesVectorSpaceResponse = diff --git a/packages/contracts/generated/api/console/features/zod.gen.ts b/packages/contracts/generated/api/console/features/zod.gen.ts index a5a66a25782..0e248a8dee3 100644 --- a/packages/contracts/generated/api/console/features/zod.gen.ts +++ b/packages/contracts/generated/api/console/features/zod.gen.ts @@ -2,6 +2,15 @@ import * as z from 'zod' +/** + * VectorSpaceLimitationModel + */ +export const zVectorSpaceLimitationModel = z.object({ + limit: z.int(), + size: z.int(), + usage_unknown: z.boolean().optional().default(false), +}) + /** * LimitationModel */ @@ -47,12 +56,23 @@ export const zLicenseLimitationModel = z.object({ size: z.int().default(0), }) +/** + * CloudPlan + * + * Enum representing user plan types in the cloud platform. + * + * SANDBOX: Free/default plan with limited features + * PROFESSIONAL: Professional paid plan + * TEAM: Team collaboration paid plan + */ +export const zCloudPlan = z.enum(['professional', 'sandbox', 'team']) + /** * SubscriptionModel */ export const zSubscriptionModel = z.object({ interval: z.string().default(''), - plan: z.string().default('sandbox'), + plan: zCloudPlan.default('sandbox'), }) /** @@ -112,4 +132,4 @@ export const zGetFeaturesResponse = zFeatureModel /** * Success */ -export const zGetFeaturesVectorSpaceResponse = zLimitationModel +export const zGetFeaturesVectorSpaceResponse = zVectorSpaceLimitationModel diff --git a/packages/contracts/generated/api/console/files/types.gen.ts b/packages/contracts/generated/api/console/files/types.gen.ts index be26853f3bf..f68df4af336 100644 --- a/packages/contracts/generated/api/console/files/types.gen.ts +++ b/packages/contracts/generated/api/console/files/types.gen.ts @@ -16,7 +16,9 @@ export type UploadConfig = { file_upload_limit: number image_file_batch_limit: number image_file_size_limit: number + knowledge_file_size_limit: number single_chunk_attachment_limit: number + skill_file_size_limit: number video_file_size_limit: number workflow_file_upload_limit: number } diff --git a/packages/contracts/generated/api/console/files/zod.gen.ts b/packages/contracts/generated/api/console/files/zod.gen.ts index d3d35b401a3..ac4400b777e 100644 --- a/packages/contracts/generated/api/console/files/zod.gen.ts +++ b/packages/contracts/generated/api/console/files/zod.gen.ts @@ -20,7 +20,9 @@ export const zUploadConfig = z.object({ file_upload_limit: z.int(), image_file_batch_limit: z.int(), image_file_size_limit: z.int(), + knowledge_file_size_limit: z.int(), single_chunk_attachment_limit: z.int(), + skill_file_size_limit: z.int(), video_file_size_limit: z.int(), workflow_file_upload_limit: z.int(), }) diff --git a/packages/contracts/generated/api/console/info/orpc.gen.ts b/packages/contracts/generated/api/console/info/orpc.gen.ts deleted file mode 100644 index 2563a61ce4e..00000000000 --- a/packages/contracts/generated/api/console/info/orpc.gen.ts +++ /dev/null @@ -1,22 +0,0 @@ -// This file is auto-generated by @hey-api/openapi-ts - -import { oc } from '@orpc/contract' -import { zPostInfoResponse } from './zod.gen' - -export const post = oc - .route({ - inputStructure: 'detailed', - method: 'POST', - operationId: 'postInfo', - path: '/info', - tags: ['console'], - }) - .output(zPostInfoResponse) - -export const info = { - post, -} - -export const contract = { - info, -} diff --git a/packages/contracts/generated/api/console/info/types.gen.ts b/packages/contracts/generated/api/console/info/types.gen.ts deleted file mode 100644 index 1902b2ab6e7..00000000000 --- a/packages/contracts/generated/api/console/info/types.gen.ts +++ /dev/null @@ -1,39 +0,0 @@ -// This file is auto-generated by @hey-api/openapi-ts - -export type ClientOptions = { - baseUrl: `${string}://${string}/console/api` | (string & {}) -} - -export type TenantInfoResponse = { - created_at?: number | null - custom_config?: WorkspaceCustomConfigResponse | null - id: string - in_trial?: boolean | null - name?: string | null - next_credit_reset_date?: number | null - plan?: string | null - role?: string | null - status?: string | null - trial_credits?: number | null - trial_credits_exhausted_at?: number | null - trial_credits_used?: number | null - trial_end_reason?: string | null -} - -export type WorkspaceCustomConfigResponse = { - remove_webapp_brand?: boolean | null - replace_webapp_logo?: string | null -} - -export type PostInfoData = { - body?: never - path?: never - query?: never - url: '/info' -} - -export type PostInfoResponses = { - 200: TenantInfoResponse -} - -export type PostInfoResponse = PostInfoResponses[keyof PostInfoResponses] diff --git a/packages/contracts/generated/api/console/info/zod.gen.ts b/packages/contracts/generated/api/console/info/zod.gen.ts deleted file mode 100644 index e116e2dbe4a..00000000000 --- a/packages/contracts/generated/api/console/info/zod.gen.ts +++ /dev/null @@ -1,35 +0,0 @@ -// This file is auto-generated by @hey-api/openapi-ts - -import * as z from 'zod' - -/** - * WorkspaceCustomConfigResponse - */ -export const zWorkspaceCustomConfigResponse = z.object({ - remove_webapp_brand: z.boolean().nullish(), - replace_webapp_logo: z.string().nullish(), -}) - -/** - * TenantInfoResponse - */ -export const zTenantInfoResponse = z.object({ - created_at: z.int().nullish(), - custom_config: zWorkspaceCustomConfigResponse.nullish(), - id: z.string(), - in_trial: z.boolean().nullish(), - name: z.string().nullish(), - next_credit_reset_date: z.int().nullish(), - plan: z.string().nullish(), - role: z.string().nullish(), - status: z.string().nullish(), - trial_credits: z.int().nullish(), - trial_credits_exhausted_at: z.int().nullish(), - trial_credits_used: z.int().nullish(), - trial_end_reason: z.string().nullish(), -}) - -/** - * Success - */ -export const zPostInfoResponse = zTenantInfoResponse diff --git a/packages/contracts/generated/api/console/installed-apps/orpc.gen.ts b/packages/contracts/generated/api/console/installed-apps/orpc.gen.ts index 8179e0738ba..4a0941ab095 100644 --- a/packages/contracts/generated/api/console/installed-apps/orpc.gen.ts +++ b/packages/contracts/generated/api/console/installed-apps/orpc.gen.ts @@ -24,6 +24,8 @@ import { zGetInstalledAppsByInstalledAppIdMetaResponse, zGetInstalledAppsByInstalledAppIdParametersPath, zGetInstalledAppsByInstalledAppIdParametersResponse, + zGetInstalledAppsByInstalledAppIdPath, + zGetInstalledAppsByInstalledAppIdResponse, zGetInstalledAppsByInstalledAppIdSavedMessagesPath, zGetInstalledAppsByInstalledAppIdSavedMessagesQuery, zGetInstalledAppsByInstalledAppIdSavedMessagesResponse, @@ -520,6 +522,17 @@ export const delete3 = oc .input(z.object({ params: zDeleteInstalledAppsByInstalledAppIdPath })) .output(zDeleteInstalledAppsByInstalledAppIdResponse) +export const get8 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getInstalledAppsByInstalledAppId', + path: '/installed-apps/{installed_app_id}', + tags: ['console'], + }) + .input(z.object({ params: zGetInstalledAppsByInstalledAppIdPath })) + .output(zGetInstalledAppsByInstalledAppIdResponse) + export const patch3 = oc .route({ inputStructure: 'detailed', @@ -538,6 +551,7 @@ export const patch3 = oc export const byInstalledAppId = { delete: delete3, + get: get8, patch: patch3, audioToText, chatMessages, @@ -551,7 +565,7 @@ export const byInstalledAppId = { workflows, } -export const get8 = oc +export const get9 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -574,7 +588,7 @@ export const post12 = oc .output(zPostInstalledAppsResponse) export const installedApps = { - get: get8, + get: get9, post: post12, byInstalledAppId, } diff --git a/packages/contracts/generated/api/console/installed-apps/types.gen.ts b/packages/contracts/generated/api/console/installed-apps/types.gen.ts index 4abc67defa3..d2b2bc01f0a 100644 --- a/packages/contracts/generated/api/console/installed-apps/types.gen.ts +++ b/packages/contracts/generated/api/console/installed-apps/types.gen.ts @@ -5,7 +5,9 @@ export type ClientOptions = { } export type InstalledAppListResponse = { + has_more: boolean installed_apps: Array + next_cursor: string | null } export type InstalledAppCreatePayload = { @@ -16,6 +18,16 @@ export type SimpleMessageResponse = { message: string } +export type InstalledAppResponse = { + app: InstalledAppInfoResponse + app_owner_tenant_id: string + editable: boolean + id: string + is_pinned: boolean + last_used_at: number | null + uninstallable: boolean +} + export type InstalledAppUpdatePayload = { is_pinned?: boolean | null } @@ -169,14 +181,16 @@ export type WorkflowRunPayload = { } } -export type InstalledAppResponse = { - app: InstalledAppInfoResponse - app_owner_tenant_id: string - editable: boolean +export type InstalledAppInfoResponse = { + description: string + icon: string | null + icon_background: string | null + icon_type: IconType | null + readonly icon_url: string | null id: string - is_pinned: boolean - last_used_at?: number | null - uninstallable: boolean + mode: AppMode + name: string + use_icon_as_answer_icon: boolean } export type JsonValue = @@ -240,17 +254,17 @@ export type SavedMessageItem = { query: string } -export type InstalledAppInfoResponse = { - description?: string | null - icon?: string | null - icon_background?: string | null - icon_type?: string | null - readonly icon_url: string | null - id: string - mode?: string | null - name?: string | null - use_icon_as_answer_icon?: boolean | null -} +export type IconType = 'emoji' | 'image' | 'link' + +export type AppMode = + | 'advanced-chat' + | 'agent' + | 'agent-chat' + | 'channel' + | 'chat' + | 'completion' + | 'rag-pipeline' + | 'workflow' export type AgentThought = { answer?: string | null @@ -413,13 +427,9 @@ export type FileTransferMethod = 'datasource_file' | 'local_file' | 'remote_url' export type ValueSourceType = 'constant' | 'variable' export type InstalledAppListResponseWritable = { - installed_apps: Array -} - -export type ExploreMessageInfiniteScrollPaginationWritable = { - data: Array has_more: boolean - limit: number + installed_apps: Array + next_cursor: string | null } export type InstalledAppResponseWritable = { @@ -428,10 +438,27 @@ export type InstalledAppResponseWritable = { editable: boolean id: string is_pinned: boolean - last_used_at?: number | null + last_used_at: number | null uninstallable: boolean } +export type ExploreMessageInfiniteScrollPaginationWritable = { + data: Array + has_more: boolean + limit: number +} + +export type InstalledAppInfoResponseWritable = { + description: string + icon: string | null + icon_background: string | null + icon_type: IconType | null + id: string + mode: AppMode + name: string + use_icon_as_answer_icon: boolean +} + export type ExploreMessageListItemWritable = { agent_thoughts: Array answer: string @@ -457,22 +484,14 @@ export type ExploreMessageListItemWritable = { total_price?: string | null } -export type InstalledAppInfoResponseWritable = { - description?: string | null - icon?: string | null - icon_background?: string | null - icon_type?: string | null - id: string - mode?: string | null - name?: string | null - use_icon_as_answer_icon?: boolean | null -} - export type GetInstalledAppsData = { body?: never path?: never query?: { app_id?: string + cursor?: string + limit?: number + name?: string } url: '/installed-apps' } @@ -512,6 +531,22 @@ export type DeleteInstalledAppsByInstalledAppIdResponses = { export type DeleteInstalledAppsByInstalledAppIdResponse = DeleteInstalledAppsByInstalledAppIdResponses[keyof DeleteInstalledAppsByInstalledAppIdResponses] +export type GetInstalledAppsByInstalledAppIdData = { + body?: never + path: { + installed_app_id: string + } + query?: never + url: '/installed-apps/{installed_app_id}' +} + +export type GetInstalledAppsByInstalledAppIdResponses = { + 200: InstalledAppResponse +} + +export type GetInstalledAppsByInstalledAppIdResponse = + GetInstalledAppsByInstalledAppIdResponses[keyof GetInstalledAppsByInstalledAppIdResponses] + export type PatchInstalledAppsByInstalledAppIdData = { body: InstalledAppUpdatePayload path: { diff --git a/packages/contracts/generated/api/console/installed-apps/zod.gen.ts b/packages/contracts/generated/api/console/installed-apps/zod.gen.ts index f0acbe06c0e..753aa7d0160 100644 --- a/packages/contracts/generated/api/console/installed-apps/zod.gen.ts +++ b/packages/contracts/generated/api/console/installed-apps/zod.gen.ts @@ -229,19 +229,38 @@ export const zParameters = z.object({ user_input_form: z.array(zJsonObject), }) +/** + * IconType + */ +export const zIconType = z.enum(['emoji', 'image', 'link']) + +/** + * AppMode + */ +export const zAppMode = z.enum([ + 'advanced-chat', + 'agent', + 'agent-chat', + 'channel', + 'chat', + 'completion', + 'rag-pipeline', + 'workflow', +]) + /** * InstalledAppInfoResponse */ export const zInstalledAppInfoResponse = z.object({ - description: z.string().nullish(), - icon: z.string().nullish(), - icon_background: z.string().nullish(), - icon_type: z.string().nullish(), + description: z.string(), + icon: z.string().nullable(), + icon_background: z.string().nullable(), + icon_type: zIconType.nullable(), icon_url: z.string().nullable(), id: z.string(), - mode: z.string().nullish(), - name: z.string().nullish(), - use_icon_as_answer_icon: z.boolean().nullish(), + mode: zAppMode, + name: z.string(), + use_icon_as_answer_icon: z.boolean(), }) /** @@ -253,7 +272,7 @@ export const zInstalledAppResponse = z.object({ editable: z.boolean(), id: z.string(), is_pinned: z.boolean(), - last_used_at: z.int().nullish(), + last_used_at: z.int().nullable(), uninstallable: z.boolean(), }) @@ -261,7 +280,9 @@ export const zInstalledAppResponse = z.object({ * InstalledAppListResponse */ export const zInstalledAppListResponse = z.object({ + has_more: z.boolean(), installed_apps: z.array(zInstalledAppResponse), + next_cursor: z.string().nullable(), }) /** @@ -430,7 +451,7 @@ export const zFileListInputConfig = z.object({ * ValueSourceType * * ValueSourceType records whether the value comes from a static setting - * in form definiton, or a variable while the workflow is running. + * in form definition, or a variable while the workflow is running. */ export const zValueSourceType = z.enum(['constant', 'variable']) @@ -547,6 +568,42 @@ export const zExploreMessageInfiniteScrollPagination = z.object({ limit: z.int(), }) +/** + * InstalledAppInfoResponse + */ +export const zInstalledAppInfoResponseWritable = z.object({ + description: z.string(), + icon: z.string().nullable(), + icon_background: z.string().nullable(), + icon_type: zIconType.nullable(), + id: z.string(), + mode: zAppMode, + name: z.string(), + use_icon_as_answer_icon: z.boolean(), +}) + +/** + * InstalledAppResponse + */ +export const zInstalledAppResponseWritable = z.object({ + app: zInstalledAppInfoResponseWritable, + app_owner_tenant_id: z.string(), + editable: z.boolean(), + id: z.string(), + is_pinned: z.boolean(), + last_used_at: z.int().nullable(), + uninstallable: z.boolean(), +}) + +/** + * InstalledAppListResponse + */ +export const zInstalledAppListResponseWritable = z.object({ + has_more: z.boolean(), + installed_apps: z.array(zInstalledAppResponseWritable), + next_cursor: z.string().nullable(), +}) + /** * ExploreMessageListItem */ @@ -585,42 +642,11 @@ export const zExploreMessageInfiniteScrollPaginationWritable = z.object({ limit: z.int(), }) -/** - * InstalledAppInfoResponse - */ -export const zInstalledAppInfoResponseWritable = z.object({ - description: z.string().nullish(), - icon: z.string().nullish(), - icon_background: z.string().nullish(), - icon_type: z.string().nullish(), - id: z.string(), - mode: z.string().nullish(), - name: z.string().nullish(), - use_icon_as_answer_icon: z.boolean().nullish(), -}) - -/** - * InstalledAppResponse - */ -export const zInstalledAppResponseWritable = z.object({ - app: zInstalledAppInfoResponseWritable, - app_owner_tenant_id: z.string(), - editable: z.boolean(), - id: z.string(), - is_pinned: z.boolean(), - last_used_at: z.int().nullish(), - uninstallable: z.boolean(), -}) - -/** - * InstalledAppListResponse - */ -export const zInstalledAppListResponseWritable = z.object({ - installed_apps: z.array(zInstalledAppResponseWritable), -}) - export const zGetInstalledAppsQuery = z.object({ app_id: z.string().optional(), + cursor: z.string().optional(), + limit: z.int().gte(1).lte(100).optional().default(20), + name: z.string().max(100).optional(), }) /** @@ -644,6 +670,15 @@ export const zDeleteInstalledAppsByInstalledAppIdPath = z.object({ */ export const zDeleteInstalledAppsByInstalledAppIdResponse = z.void() +export const zGetInstalledAppsByInstalledAppIdPath = z.object({ + installed_app_id: z.uuid(), +}) + +/** + * Success + */ +export const zGetInstalledAppsByInstalledAppIdResponse = zInstalledAppResponse + export const zPatchInstalledAppsByInstalledAppIdBody = zInstalledAppUpdatePayload export const zPatchInstalledAppsByInstalledAppIdPath = z.object({ diff --git a/packages/contracts/generated/api/console/orpc.gen.ts b/packages/contracts/generated/api/console/orpc.gen.ts index eda2bbfcefd..db816dd03c4 100644 --- a/packages/contracts/generated/api/console/orpc.gen.ts +++ b/packages/contracts/generated/api/console/orpc.gen.ts @@ -34,7 +34,6 @@ export const contractLoaders = { forgotPassword: () => import('./forgot-password/orpc.gen').then(({ forgotPassword }) => ({ forgotPassword })), form: () => import('./form/orpc.gen').then(({ form }) => ({ form })), - info: () => import('./info/orpc.gen').then(({ info }) => ({ info })), init: () => import('./init/orpc.gen').then(({ init }) => ({ init })), installedApps: () => import('./installed-apps/orpc.gen').then(({ installedApps }) => ({ installedApps })), diff --git a/packages/contracts/generated/api/console/rag/types.gen.ts b/packages/contracts/generated/api/console/rag/types.gen.ts index 2777d150024..dedbd4ec55e 100644 --- a/packages/contracts/generated/api/console/rag/types.gen.ts +++ b/packages/contracts/generated/api/console/rag/types.gen.ts @@ -175,6 +175,7 @@ export type WorkflowResponse = { updated_at: number updated_by?: SimpleAccountResponse | null version: string + version_number?: number | null } export type DraftWorkflowSyncPayload = { @@ -584,6 +585,10 @@ export type PostRagPipelineCustomizedTemplatesByTemplateIdData = { url: '/rag/pipeline/customized/templates/{template_id}' } +export type PostRagPipelineCustomizedTemplatesByTemplateIdErrors = { + 404: unknown +} + export type PostRagPipelineCustomizedTemplatesByTemplateIdResponses = { 200: SimpleDataResponse } @@ -647,6 +652,10 @@ export type GetRagPipelineTemplatesByTemplateIdData = { url: '/rag/pipeline/templates/{template_id}' } +export type GetRagPipelineTemplatesByTemplateIdErrors = { + 404: unknown +} + export type GetRagPipelineTemplatesByTemplateIdResponses = { 200: PipelineTemplateDetailResponse } @@ -754,6 +763,10 @@ export type PostRagPipelinesTransformDatasetsByDatasetIdData = { url: '/rag/pipelines/transform/datasets/{dataset_id}' } +export type PostRagPipelinesTransformDatasetsByDatasetIdErrors = { + 404: unknown +} + export type PostRagPipelinesTransformDatasetsByDatasetIdResponses = { 200: RagPipelineOpaqueResponse } @@ -770,6 +783,10 @@ export type PostRagPipelinesByPipelineIdCustomizedPublishData = { url: '/rag/pipelines/{pipeline_id}/customized/publish' } +export type PostRagPipelinesByPipelineIdCustomizedPublishErrors = { + 404: unknown +} + export type PostRagPipelinesByPipelineIdCustomizedPublishResponses = { 204: void } diff --git a/packages/contracts/generated/api/console/rag/zod.gen.ts b/packages/contracts/generated/api/console/rag/zod.gen.ts index 3b9a4ad4a71..7ff07aabc10 100644 --- a/packages/contracts/generated/api/console/rag/zod.gen.ts +++ b/packages/contracts/generated/api/console/rag/zod.gen.ts @@ -480,6 +480,7 @@ export const zWorkflowResponse = z.object({ updated_at: z.int(), updated_by: zSimpleAccountResponse.nullish(), version: z.string(), + version_number: z.int().nullish(), }) /** diff --git a/packages/contracts/generated/api/console/router.gen.ts b/packages/contracts/generated/api/console/router.gen.ts index a28155f4eea..a036eec4e06 100644 --- a/packages/contracts/generated/api/console/router.gen.ts +++ b/packages/contracts/generated/api/console/router.gen.ts @@ -23,7 +23,6 @@ import { features } from './features/orpc.gen' import { files } from './files/orpc.gen' import { forgotPassword } from './forgot-password/orpc.gen' import { form } from './form/orpc.gen' -import { info } from './info/orpc.gen' import { init } from './init/orpc.gen' import { installedApps } from './installed-apps/orpc.gen' import { instructionGenerate } from './instruction-generate/orpc.gen' @@ -80,7 +79,6 @@ const communityContract = { files, forgotPassword, form, - info, init, installedApps, instructionGenerate, diff --git a/packages/contracts/generated/api/console/setup/orpc.gen.ts b/packages/contracts/generated/api/console/setup/orpc.gen.ts index 12e1631e7e8..e298a076637 100644 --- a/packages/contracts/generated/api/console/setup/orpc.gen.ts +++ b/packages/contracts/generated/api/console/setup/orpc.gen.ts @@ -31,7 +31,7 @@ export const get = oc * Initialize system setup with admin account. * * NOTE: This endpoint is unauthenticated by design for first-time bootstrap. - * Access is restricted by deployment mode (`SELF_HOSTED`), one-time setup guards, + * Access is restricted to self-hosted editions (`COMMUNITY` and `ENTERPRISE`), one-time setup guards, * and init-password validation rather than user session authentication. * */ @@ -43,7 +43,7 @@ export const post = oc path: '/setup', successStatus: 201, summary: - 'Initialize system setup with admin account.\n\n NOTE: This endpoint is unauthenticated by design for first-time bootstrap.\n Access is restricted by deployment mode (`SELF_HOSTED`), one-time setup guards,\n and init-password validation rather than user session authentication.\n ', + 'Initialize system setup with admin account.\n\n NOTE: This endpoint is unauthenticated by design for first-time bootstrap.\n Access is restricted to self-hosted editions (`COMMUNITY` and `ENTERPRISE`), one-time setup guards,\n and init-password validation rather than user session authentication.\n ', tags: ['console'], }) .input(z.object({ body: zPostSetupBody })) diff --git a/packages/contracts/generated/api/console/snippets/types.gen.ts b/packages/contracts/generated/api/console/snippets/types.gen.ts index d5b6a62cfc4..67e5d738e3e 100644 --- a/packages/contracts/generated/api/console/snippets/types.gen.ts +++ b/packages/contracts/generated/api/console/snippets/types.gen.ts @@ -71,6 +71,7 @@ export type SnippetWorkflowResponse = { updated_at: number updated_by?: SimpleAccountResponse | null version: string + version_number?: number | null } export type SnippetDraftSyncPayload = { @@ -924,6 +925,7 @@ export type AgentSoulModelSettings = { stop?: Array | null temperature?: number | null top_p?: number | null + [key: string]: unknown } export type AgentSandboxProviderConfig = { @@ -1130,6 +1132,8 @@ export type AgentKnowledgeMetadataCondition = { | '≠' | '≤' | '≥' + id?: string | null + metadata_id?: string | null name: string value?: string | Array | number | null } diff --git a/packages/contracts/generated/api/console/snippets/zod.gen.ts b/packages/contracts/generated/api/console/snippets/zod.gen.ts index 4474c465df9..958018613d8 100644 --- a/packages/contracts/generated/api/console/snippets/zod.gen.ts +++ b/packages/contracts/generated/api/console/snippets/zod.gen.ts @@ -317,6 +317,7 @@ export const zSnippetWorkflowResponse = z.object({ updated_at: z.int(), updated_by: zSimpleAccountResponse.nullish(), version: z.string(), + version_number: z.int().nullish(), }) /** @@ -1155,6 +1156,13 @@ export const zAgentModelResponseFormatConfig = z.object({ /** * AgentSoulModelSettings + * + * Model parameters for the Agent Soul model. + * + * Model plugins can declare arbitrary parameters via ``parameter_rules`` + * (e.g. Qwen/Tongyi's ``enable_thinking``) beyond the common OpenAI-style + * fields typed below, so extra keys must round-trip through persistence + * rather than being dropped. */ export const zAgentSoulModelSettings = z.object({ frequency_penalty: z.number().nullish(), @@ -1439,6 +1447,14 @@ export const zAgentKnowledgeRetrievalConfig = z.object({ /** * AgentKnowledgeMetadataCondition + * + * One manual metadata filter clause. + * + * ``id`` and ``metadata_id`` are UI-only bookkeeping the composer sends on + * every save (a stable row key and a reference to the selected metadata + * field). They are persisted here for round-tripping the composer's draft + * state but are stripped before building the Agent runtime request, whose + * DTO only accepts ``name``/``comparison_operator``/``value``. */ export const zAgentKnowledgeMetadataCondition = z.object({ comparison_operator: z.enum([ @@ -1461,6 +1477,8 @@ export const zAgentKnowledgeMetadataCondition = z.object({ '≤', '≥', ]), + id: z.string().nullish(), + metadata_id: z.string().nullish(), name: z.string().min(1).max(255), value: z.union([z.string(), z.array(z.string()), z.number()]).nullish(), }) diff --git a/packages/contracts/generated/api/console/system-features/types.gen.ts b/packages/contracts/generated/api/console/system-features/types.gen.ts index 076f98546cc..b3b4a11723f 100644 --- a/packages/contracts/generated/api/console/system-features/types.gen.ts +++ b/packages/contracts/generated/api/console/system-features/types.gen.ts @@ -18,7 +18,6 @@ export type SystemFeatureModel = { enable_marketplace: boolean enable_social_oauth_login: boolean enable_step_by_step_tour: boolean - enable_trial_app: boolean is_allow_register: boolean is_email_setup: boolean knowledge_fs_enabled: boolean @@ -26,12 +25,13 @@ export type SystemFeatureModel = { plugin_installation_permission: PluginInstallationPermissionModel rbac_enabled: boolean sso_enforced_for_signin: boolean - sso_enforced_for_signin_protocol: string + sso_enforced_for_signin_protocol: SsoProtocol | null webapp_auth: WebAppAuthModel } export type LicenseModel = { expired_at: string + license_expiry_notice_enabled: boolean seats: LicenseLimitationModel status: LicenseStatus workspaces: LicenseLimitationModel @@ -56,6 +56,8 @@ export type PluginInstallationPermissionModel = { restrict_to_marketplace_only: boolean } +export type SsoProtocol = 'oauth2' | 'oidc' | 'saml' + export type WebAppAuthModel = { allow_email_code_login: boolean allow_email_password_login: boolean @@ -80,7 +82,7 @@ export type PluginInstallationScope = | 'official_only' export type WebAppAuthSsoModel = { - protocol: string + protocol: SsoProtocol | null } export type GetSystemFeaturesData = { diff --git a/packages/contracts/generated/api/console/system-features/zod.gen.ts b/packages/contracts/generated/api/console/system-features/zod.gen.ts index 20cf33d3891..bec08d1c901 100644 --- a/packages/contracts/generated/api/console/system-features/zod.gen.ts +++ b/packages/contracts/generated/api/console/system-features/zod.gen.ts @@ -20,6 +20,11 @@ export const zBrandingModel = z.object({ */ export const zDeploymentEdition = z.enum(['CLOUD', 'COMMUNITY', 'ENTERPRISE']) +/** + * SSOProtocol + */ +export const zSsoProtocol = z.enum(['oauth2', 'oidc', 'saml']) + /** * LicenseLimitationModel * @@ -43,6 +48,7 @@ export const zLicenseStatus = z.enum(['active', 'expired', 'expiring', 'inactive */ export const zLicenseModel = z.object({ expired_at: z.string().default(''), + license_expiry_notice_enabled: z.boolean().default(false), seats: zLicenseLimitationModel.default({ enabled: false, limit: 0, @@ -85,7 +91,7 @@ export const zPluginInstallationPermissionModel = z.object({ * WebAppAuthSSOModel */ export const zWebAppAuthSsoModel = z.object({ - protocol: z.string().default(''), + protocol: zSsoProtocol.nullable(), }) /** @@ -97,7 +103,7 @@ export const zWebAppAuthModel = z.object({ allow_public_access: z.boolean().default(true), allow_sso: z.boolean().default(false), enabled: z.boolean().default(false), - sso_config: zWebAppAuthSsoModel.default({ protocol: '' }), + sso_config: zWebAppAuthSsoModel, }) /** @@ -125,7 +131,6 @@ export const zSystemFeatureModel = z.object({ enable_marketplace: z.boolean().default(false), enable_social_oauth_login: z.boolean().default(false), enable_step_by_step_tour: z.boolean().default(false), - enable_trial_app: z.boolean().default(false), is_allow_register: z.boolean().default(false), is_email_setup: z.boolean().default(false), knowledge_fs_enabled: z.boolean().default(false), @@ -136,15 +141,8 @@ export const zSystemFeatureModel = z.object({ }), rbac_enabled: z.boolean().default(false), sso_enforced_for_signin: z.boolean().default(false), - sso_enforced_for_signin_protocol: z.string().default(''), - webapp_auth: zWebAppAuthModel.default({ - allow_email_code_login: false, - allow_email_password_login: false, - allow_public_access: true, - allow_sso: false, - enabled: false, - sso_config: { protocol: '' }, - }), + sso_enforced_for_signin_protocol: zSsoProtocol.nullable(), + webapp_auth: zWebAppAuthModel, }) /** diff --git a/packages/contracts/generated/api/console/version/types.gen.ts b/packages/contracts/generated/api/console/version/types.gen.ts index b11ccd8a5fb..67bca5e9b8b 100644 --- a/packages/contracts/generated/api/console/version/types.gen.ts +++ b/packages/contracts/generated/api/console/version/types.gen.ts @@ -5,18 +5,10 @@ export type ClientOptions = { } export type VersionResponse = { - can_auto_update: boolean - features: VersionFeatures - release_date: string release_notes: string version: string } -export type VersionFeatures = { - can_replace_logo: boolean - model_load_balancing_enabled: boolean -} - export type GetVersionData = { body?: never path?: never diff --git a/packages/contracts/generated/api/console/version/zod.gen.ts b/packages/contracts/generated/api/console/version/zod.gen.ts index c73aa358f21..6bc8019709d 100644 --- a/packages/contracts/generated/api/console/version/zod.gen.ts +++ b/packages/contracts/generated/api/console/version/zod.gen.ts @@ -2,21 +2,10 @@ import * as z from 'zod' -/** - * VersionFeatures - */ -export const zVersionFeatures = z.object({ - can_replace_logo: z.boolean(), - model_load_balancing_enabled: z.boolean(), -}) - /** * VersionResponse */ export const zVersionResponse = z.object({ - can_auto_update: z.boolean(), - features: zVersionFeatures, - release_date: z.string(), release_notes: z.string(), version: z.string(), }) diff --git a/packages/contracts/generated/api/console/workspaces/orpc.gen.ts b/packages/contracts/generated/api/console/workspaces/orpc.gen.ts index 730d89367a7..292ad4b0ebe 100644 --- a/packages/contracts/generated/api/console/workspaces/orpc.gen.ts +++ b/packages/contracts/generated/api/console/workspaces/orpc.gen.ts @@ -69,8 +69,10 @@ import { zGetWorkspacesCurrentModelProvidersByProviderModelsParameterRulesResponse, zGetWorkspacesCurrentModelProvidersByProviderModelsPath, zGetWorkspacesCurrentModelProvidersByProviderModelsResponse, + zGetWorkspacesCurrentModelProvidersCreditsResponse, zGetWorkspacesCurrentModelProvidersQuery, zGetWorkspacesCurrentModelProvidersResponse, + zGetWorkspacesCurrentModelProvidersSummaryResponse, zGetWorkspacesCurrentModelsModelTypesByModelTypePath, zGetWorkspacesCurrentModelsModelTypesByModelTypeResponse, zGetWorkspacesCurrentPermissionResponse, @@ -86,6 +88,8 @@ import { zGetWorkspacesCurrentPluginFetchManifestResponse, zGetWorkspacesCurrentPluginIconQuery, zGetWorkspacesCurrentPluginIconResponse, + zGetWorkspacesCurrentPluginInstalledIdsQuery, + zGetWorkspacesCurrentPluginInstalledIdsResponse, zGetWorkspacesCurrentPluginListQuery, zGetWorkspacesCurrentPluginListResponse, zGetWorkspacesCurrentPluginMarketplacePkgQuery, @@ -147,6 +151,7 @@ import { zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPoliciesByPolicyIdRoleBindingsPath, zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPoliciesByPolicyIdRoleBindingsResponse, zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPolicyResponse, + zGetWorkspacesCurrentSummaryResponse, zGetWorkspacesCurrentToolLabelsResponse, zGetWorkspacesCurrentToolProviderApiGetQuery, zGetWorkspacesCurrentToolProviderApiGetResponse, @@ -201,6 +206,7 @@ import { zGetWorkspacesCurrentTriggerProviderByProviderSubscriptionsOauthAuthorizePath, zGetWorkspacesCurrentTriggerProviderByProviderSubscriptionsOauthAuthorizeResponse, zGetWorkspacesCurrentTriggersResponse, + zGetWorkspacesCustomConfigResponse, zGetWorkspacesResponse, zPatchWorkspacesCurrentCustomizedSnippetsBySnippetIdBody, zPatchWorkspacesCurrentCustomizedSnippetsBySnippetIdPath, @@ -315,7 +321,6 @@ import { zPostWorkspacesCurrentRbacRolesByRoleIdCopyPath, zPostWorkspacesCurrentRbacRolesByRoleIdCopyResponse, zPostWorkspacesCurrentRbacRolesResponse, - zPostWorkspacesCurrentResponse, zPostWorkspacesCurrentToolProviderApiAddBody, zPostWorkspacesCurrentToolProviderApiAddResponse, zPostWorkspacesCurrentToolProviderApiDeleteBody, @@ -1066,6 +1071,34 @@ export const members = { } export const get12 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getWorkspacesCurrentModelProvidersCredits', + path: '/workspaces/current/model-providers/credits', + tags: ['console'], + }) + .output(zGetWorkspacesCurrentModelProvidersCreditsResponse) + +export const credits = { + get: get12, +} + +export const get13 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getWorkspacesCurrentModelProvidersSummary', + path: '/workspaces/current/model-providers/summary', + tags: ['console'], + }) + .output(zGetWorkspacesCurrentModelProvidersSummaryResponse) + +export const summary = { + get: get13, +} + +export const get14 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1077,7 +1110,7 @@ export const get12 = oc .output(zGetWorkspacesCurrentModelProvidersByProviderCheckoutUrlResponse) export const checkoutUrl = { - get: get12, + get: get14, } export const post16 = oc @@ -1137,7 +1170,7 @@ export const delete5 = oc ) .output(zDeleteWorkspacesCurrentModelProvidersByProviderCredentialsResponse) -export const get13 = oc +export const get15 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1188,7 +1221,7 @@ export const put2 = oc export const credentials = { delete: delete5, - get: get13, + get: get15, post: post18, put: put2, switch: switch_, @@ -1252,7 +1285,7 @@ export const delete6 = oc ) .output(zDeleteWorkspacesCurrentModelProvidersByProviderModelsCredentialsResponse) -export const get14 = oc +export const get16 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1303,7 +1336,7 @@ export const put3 = oc export const credentials2 = { delete: delete6, - get: get14, + get: get16, post: post21, put: put3, switch: switch2, @@ -1407,7 +1440,7 @@ export const loadBalancingConfigs = { byConfigId, } -export const get15 = oc +export const get17 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1424,7 +1457,7 @@ export const get15 = oc .output(zGetWorkspacesCurrentModelProvidersByProviderModelsParameterRulesResponse) export const parameterRules = { - get: get15, + get: get17, } export const delete7 = oc @@ -1444,7 +1477,7 @@ export const delete7 = oc ) .output(zDeleteWorkspacesCurrentModelProvidersByProviderModelsResponse) -export const get16 = oc +export const get18 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1473,7 +1506,7 @@ export const post24 = oc export const models = { delete: delete7, - get: get16, + get: get18, post: post24, credentials: credentials2, disable: disable2, @@ -1509,7 +1542,7 @@ export const byProvider = { preferredProviderType, } -export const get17 = oc +export const get19 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1521,11 +1554,13 @@ export const get17 = oc .output(zGetWorkspacesCurrentModelProvidersResponse) export const modelProviders = { - get: get17, + get: get19, + credits, + summary, byProvider, } -export const get18 = oc +export const get20 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1537,7 +1572,7 @@ export const get18 = oc .output(zGetWorkspacesCurrentModelsModelTypesByModelTypeResponse) export const byModelType = { - get: get18, + get: get20, } export const modelTypes = { @@ -1553,7 +1588,7 @@ export const models2 = { * * Returns permission flags that control workspace features like member invitations and owner transfer. */ -export const get19 = oc +export const get21 = oc .route({ description: 'Returns permission flags that control workspace features like member invitations and owner transfer.', @@ -1567,10 +1602,10 @@ export const get19 = oc .output(zGetWorkspacesCurrentPermissionResponse) export const permission = { - get: get19, + get: get21, } -export const get20 = oc +export const get22 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1582,7 +1617,7 @@ export const get20 = oc .output(zGetWorkspacesCurrentPluginAssetResponse) export const asset = { - get: get20, + get: get22, } export const post26 = oc @@ -1615,7 +1650,7 @@ export const exclude = { post: post27, } -export const get21 = oc +export const get23 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1627,7 +1662,7 @@ export const get21 = oc .output(zGetWorkspacesCurrentPluginAutoUpgradeFetchResponse) export const fetch_ = { - get: get21, + get: get23, } export const autoUpgrade = { @@ -1636,7 +1671,7 @@ export const autoUpgrade = { fetch: fetch_, } -export const get22 = oc +export const get24 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1647,10 +1682,10 @@ export const get22 = oc .output(zGetWorkspacesCurrentPluginDebuggingKeyResponse) export const debuggingKey = { - get: get22, + get: get24, } -export const get23 = oc +export const get25 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1662,10 +1697,10 @@ export const get23 = oc .output(zGetWorkspacesCurrentPluginFetchManifestResponse) export const fetchManifest = { - get: get23, + get: get25, } -export const get24 = oc +export const get26 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1677,7 +1712,7 @@ export const get24 = oc .output(zGetWorkspacesCurrentPluginIconResponse) export const icon = { - get: get24, + get: get26, } export const post28 = oc @@ -1731,6 +1766,21 @@ export const install = { pkg, } +export const get27 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getWorkspacesCurrentPluginInstalledIds', + path: '/workspaces/current/plugin/installed-ids', + tags: ['console'], + }) + .input(z.object({ query: zGetWorkspacesCurrentPluginInstalledIdsQuery })) + .output(zGetWorkspacesCurrentPluginInstalledIdsResponse) + +export const installedIds = { + get: get27, +} + export const post31 = oc .route({ inputStructure: 'detailed', @@ -1765,7 +1815,7 @@ export const latestVersions = { post: post32, } -export const get25 = oc +export const get28 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1777,12 +1827,12 @@ export const get25 = oc .output(zGetWorkspacesCurrentPluginListResponse) export const list2 = { - get: get25, + get: get28, installations, latestVersions, } -export const get26 = oc +export const get29 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1794,14 +1844,14 @@ export const get26 = oc .output(zGetWorkspacesCurrentPluginMarketplacePkgResponse) export const pkg2 = { - get: get26, + get: get29, } export const marketplace2 = { pkg: pkg2, } -export const get27 = oc +export const get30 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1813,7 +1863,7 @@ export const get27 = oc .output(zGetWorkspacesCurrentPluginParametersDynamicOptionsResponse) export const dynamicOptions = { - get: get27, + get: get30, } /** @@ -1857,7 +1907,7 @@ export const change2 = { post: post34, } -export const get28 = oc +export const get31 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1868,7 +1918,7 @@ export const get28 = oc .output(zGetWorkspacesCurrentPluginPermissionFetchResponse) export const fetch2 = { - get: get28, + get: get31, } export const permission2 = { @@ -1876,7 +1926,7 @@ export const permission2 = { fetch: fetch2, } -export const get29 = oc +export const get32 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1888,7 +1938,7 @@ export const get29 = oc .output(zGetWorkspacesCurrentPluginReadmeResponse) export const readme = { - get: get29, + get: get32, } export const post35 = oc @@ -1936,7 +1986,7 @@ export const delete8 = { byIdentifier, } -export const get30 = oc +export const get33 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1948,11 +1998,11 @@ export const get30 = oc .output(zGetWorkspacesCurrentPluginTasksByTaskIdResponse) export const byTaskId = { - get: get30, + get: get33, delete: delete8, } -export const get31 = oc +export const get34 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -1964,7 +2014,7 @@ export const get31 = oc .output(zGetWorkspacesCurrentPluginTasksResponse) export const tasks = { - get: get31, + get: get34, deleteAll, byTaskId, } @@ -2069,7 +2119,7 @@ export const upload = { pkg: pkg3, } -export const get32 = oc +export const get35 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2086,7 +2136,7 @@ export const get32 = oc .output(zGetWorkspacesCurrentPluginByCategoryListResponse) export const list3 = { - get: get32, + get: get35, } export const byCategory = { @@ -2100,6 +2150,7 @@ export const plugin2 = { fetchManifest, icon, install, + installedIds, list: list2, marketplace: marketplace2, parameters, @@ -2139,7 +2190,7 @@ export const delete9 = oc .input(z.object({ params: zDeleteWorkspacesCurrentRbacAccessPoliciesByPolicyIdPath })) .output(zDeleteWorkspacesCurrentRbacAccessPoliciesByPolicyIdResponse) -export const get33 = oc +export const get36 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2163,12 +2214,12 @@ export const put4 = oc export const byPolicyId = { delete: delete9, - get: get33, + get: get36, put: put4, copy, } -export const get34 = oc +export const get37 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2190,7 +2241,7 @@ export const post45 = oc .output(zPostWorkspacesCurrentRbacAccessPoliciesResponse) export const accessPolicies = { - get: get34, + get: get37, post: post45, byPolicyId, } @@ -2250,7 +2301,7 @@ export const delete10 = oc ) .output(zDeleteWorkspacesCurrentRbacAppsByAppIdAccessPoliciesByPolicyIdMemberBindingsResponse) -export const get35 = oc +export const get38 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2267,10 +2318,10 @@ export const get35 = oc export const memberBindings = { delete: delete10, - get: get35, + get: get38, } -export const get36 = oc +export const get39 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2286,7 +2337,7 @@ export const get36 = oc .output(zGetWorkspacesCurrentRbacAppsByAppIdAccessPoliciesByPolicyIdRoleBindingsResponse) export const roleBindings = { - get: get36, + get: get39, } export const byPolicyId2 = { @@ -2298,7 +2349,7 @@ export const accessPolicies2 = { byPolicyId: byPolicyId2, } -export const get37 = oc +export const get40 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2315,10 +2366,10 @@ export const get37 = oc .output(zGetWorkspacesCurrentRbacAppsByAppIdAccessPolicyResponse) export const accessPolicy = { - get: get37, + get: get40, } -export const get38 = oc +export const get41 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2335,7 +2386,7 @@ export const get38 = oc .output(zGetWorkspacesCurrentRbacAppsByAppIdUserAccessPoliciesResponse) export const userAccessPolicies = { - get: get38, + get: get41, } export const put7 = oc @@ -2366,7 +2417,7 @@ export const users = { byTargetAccountId, } -export const get39 = oc +export const get42 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2394,7 +2445,7 @@ export const put8 = oc .output(zPutWorkspacesCurrentRbacAppsByAppIdWhitelistResponse) export const whitelist = { - get: get39, + get: get42, put: put8, } @@ -2430,7 +2481,7 @@ export const delete11 = oc zDeleteWorkspacesCurrentRbacDatasetsByDatasetIdAccessPoliciesByPolicyIdMemberBindingsResponse, ) -export const get40 = oc +export const get43 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2451,10 +2502,10 @@ export const get40 = oc export const memberBindings2 = { delete: delete11, - get: get40, + get: get43, } -export const get41 = oc +export const get44 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2470,7 +2521,7 @@ export const get41 = oc .output(zGetWorkspacesCurrentRbacDatasetsByDatasetIdAccessPoliciesByPolicyIdRoleBindingsResponse) export const roleBindings2 = { - get: get41, + get: get44, } export const byPolicyId3 = { @@ -2482,7 +2533,7 @@ export const accessPolicies4 = { byPolicyId: byPolicyId3, } -export const get42 = oc +export const get45 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2499,10 +2550,10 @@ export const get42 = oc .output(zGetWorkspacesCurrentRbacDatasetsByDatasetIdAccessPolicyResponse) export const accessPolicy2 = { - get: get42, + get: get45, } -export const get43 = oc +export const get46 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2519,7 +2570,7 @@ export const get43 = oc .output(zGetWorkspacesCurrentRbacDatasetsByDatasetIdUserAccessPoliciesResponse) export const userAccessPolicies2 = { - get: get43, + get: get46, } export const put9 = oc @@ -2550,7 +2601,7 @@ export const users2 = { byTargetAccountId: byTargetAccountId2, } -export const get44 = oc +export const get47 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2578,7 +2629,7 @@ export const put10 = oc .output(zPutWorkspacesCurrentRbacDatasetsByDatasetIdWhitelistResponse) export const whitelist2 = { - get: get44, + get: get47, put: put10, } @@ -2594,7 +2645,7 @@ export const datasets = { byDatasetId, } -export const get45 = oc +export const get48 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2622,7 +2673,7 @@ export const put11 = oc .output(zPutWorkspacesCurrentRbacMembersByMemberIdRbacRolesResponse) export const rbacRoles = { - get: get45, + get: get48, put: put11, } @@ -2634,7 +2685,7 @@ export const members2 = { byMemberId: byMemberId2, } -export const get46 = oc +export const get49 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2645,10 +2696,10 @@ export const get46 = oc .output(zGetWorkspacesCurrentRbacMyPermissionsResponse) export const myPermissions = { - get: get46, + get: get49, } -export const get47 = oc +export const get50 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2659,10 +2710,10 @@ export const get47 = oc .output(zGetWorkspacesCurrentRbacRolePermissionsCatalogAppResponse) export const app = { - get: get47, + get: get50, } -export const get48 = oc +export const get51 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2673,10 +2724,10 @@ export const get48 = oc .output(zGetWorkspacesCurrentRbacRolePermissionsCatalogDatasetResponse) export const dataset = { - get: get48, + get: get51, } -export const get49 = oc +export const get52 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2687,7 +2738,7 @@ export const get49 = oc .output(zGetWorkspacesCurrentRbacRolePermissionsCatalogResponse) export const catalog = { - get: get49, + get: get52, app, dataset, } @@ -2712,7 +2763,7 @@ export const copy2 = { post: post46, } -export const get50 = oc +export const get53 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2724,7 +2775,7 @@ export const get50 = oc .output(zGetWorkspacesCurrentRbacRolesByRoleIdMembersResponse) export const members3 = { - get: get50, + get: get53, } export const delete12 = oc @@ -2738,7 +2789,7 @@ export const delete12 = oc .input(z.object({ params: zDeleteWorkspacesCurrentRbacRolesByRoleIdPath })) .output(zDeleteWorkspacesCurrentRbacRolesByRoleIdResponse) -export const get51 = oc +export const get54 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2762,13 +2813,13 @@ export const put12 = oc export const byRoleId = { delete: delete12, - get: get51, + get: get54, put: put12, copy: copy2, members: members3, } -export const get52 = oc +export const get55 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2790,7 +2841,7 @@ export const post47 = oc .output(zPostWorkspacesCurrentRbacRolesResponse) export const roles = { - get: get52, + get: get55, post: post47, byRoleId, } @@ -2815,7 +2866,7 @@ export const bindings = { put: put13, } -export const get53 = oc +export const get56 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2831,10 +2882,10 @@ export const get53 = oc .output(zGetWorkspacesCurrentRbacWorkspaceAppsAccessPoliciesByPolicyIdMemberBindingsResponse) export const memberBindings3 = { - get: get53, + get: get56, } -export const get54 = oc +export const get57 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2850,7 +2901,7 @@ export const get54 = oc .output(zGetWorkspacesCurrentRbacWorkspaceAppsAccessPoliciesByPolicyIdRoleBindingsResponse) export const roleBindings3 = { - get: get54, + get: get57, } export const byPolicyId4 = { @@ -2863,7 +2914,7 @@ export const accessPolicies6 = { byPolicyId: byPolicyId4, } -export const get55 = oc +export const get58 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2874,7 +2925,7 @@ export const get55 = oc .output(zGetWorkspacesCurrentRbacWorkspaceAppsAccessPolicyResponse) export const accessPolicy3 = { - get: get55, + get: get58, } export const apps2 = { @@ -2902,7 +2953,7 @@ export const bindings2 = { put: put14, } -export const get56 = oc +export const get59 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2918,10 +2969,10 @@ export const get56 = oc .output(zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPoliciesByPolicyIdMemberBindingsResponse) export const memberBindings4 = { - get: get56, + get: get59, } -export const get57 = oc +export const get60 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2937,7 +2988,7 @@ export const get57 = oc .output(zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPoliciesByPolicyIdRoleBindingsResponse) export const roleBindings4 = { - get: get57, + get: get60, } export const byPolicyId5 = { @@ -2950,7 +3001,7 @@ export const accessPolicies7 = { byPolicyId: byPolicyId5, } -export const get58 = oc +export const get61 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2961,7 +3012,7 @@ export const get58 = oc .output(zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPolicyResponse) export const accessPolicy4 = { - get: get58, + get: get61, } export const datasets2 = { @@ -2986,7 +3037,21 @@ export const rbac = { workspace, } -export const get59 = oc +export const get62 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getWorkspacesCurrentSummary', + path: '/workspaces/current/summary', + tags: ['console'], + }) + .output(zGetWorkspacesCurrentSummaryResponse) + +export const summary2 = { + get: get62, +} + +export const get63 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -2997,7 +3062,7 @@ export const get59 = oc .output(zGetWorkspacesCurrentToolLabelsResponse) export const toolLabels = { - get: get59, + get: get63, } export const post48 = oc @@ -3030,7 +3095,7 @@ export const delete13 = { post: post49, } -export const get60 = oc +export const get64 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3041,11 +3106,11 @@ export const get60 = oc .input(z.object({ query: zGetWorkspacesCurrentToolProviderApiGetQuery })) .output(zGetWorkspacesCurrentToolProviderApiGetResponse) -export const get61 = { - get: get60, +export const get65 = { + get: get64, } -export const get62 = oc +export const get66 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3057,7 +3122,7 @@ export const get62 = oc .output(zGetWorkspacesCurrentToolProviderApiRemoteResponse) export const remote = { - get: get62, + get: get66, } export const post50 = oc @@ -3094,7 +3159,7 @@ export const test = { pre, } -export const get63 = oc +export const get67 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3106,7 +3171,7 @@ export const get63 = oc .output(zGetWorkspacesCurrentToolProviderApiToolsResponse) export const tools = { - get: get63, + get: get67, } export const post52 = oc @@ -3127,7 +3192,7 @@ export const update2 = { export const api = { add, delete: delete13, - get: get61, + get: get65, remote, schema, test, @@ -3155,7 +3220,7 @@ export const add2 = { post: post53, } -export const get64 = oc +export const get68 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3172,10 +3237,10 @@ export const get64 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderCredentialInfoResponse) export const info = { - get: get64, + get: get68, } -export const get65 = oc +export const get69 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3195,7 +3260,7 @@ export const get65 = oc ) export const byCredentialType = { - get: get65, + get: get69, } export const schema2 = { @@ -3207,7 +3272,7 @@ export const credential = { schema: schema2, } -export const get66 = oc +export const get70 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3224,7 +3289,7 @@ export const get66 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderCredentialsResponse) export const credentials3 = { - get: get66, + get: get70, } export const post54 = oc @@ -3267,7 +3332,7 @@ export const delete14 = { post: post55, } -export const get67 = oc +export const get71 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3279,10 +3344,10 @@ export const get67 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderIconResponse) export const icon2 = { - get: get67, + get: get71, } -export const get68 = oc +export const get72 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3294,10 +3359,10 @@ export const get68 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderInfoResponse) export const info2 = { - get: get68, + get: get72, } -export const get69 = oc +export const get73 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3311,7 +3376,7 @@ export const get69 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderOauthClientSchemaResponse) export const clientSchema = { - get: get69, + get: get73, } export const delete15 = oc @@ -3329,7 +3394,7 @@ export const delete15 = oc ) .output(zDeleteWorkspacesCurrentToolProviderBuiltinByProviderOauthCustomClientResponse) -export const get70 = oc +export const get74 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3360,7 +3425,7 @@ export const post56 = oc export const customClient = { delete: delete15, - get: get70, + get: get74, post: post56, } @@ -3369,7 +3434,7 @@ export const oauth = { customClient, } -export const get71 = oc +export const get75 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3381,7 +3446,7 @@ export const get71 = oc .output(zGetWorkspacesCurrentToolProviderBuiltinByProviderToolsResponse) export const tools2 = { - get: get71, + get: get75, } export const post57 = oc @@ -3436,7 +3501,7 @@ export const auth = { post: post58, } -export const get72 = oc +export const get76 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3448,14 +3513,14 @@ export const get72 = oc .output(zGetWorkspacesCurrentToolProviderMcpToolsByProviderIdResponse) export const byProviderId = { - get: get72, + get: get76, } export const tools3 = { byProviderId, } -export const get73 = oc +export const get77 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3467,7 +3532,7 @@ export const get73 = oc .output(zGetWorkspacesCurrentToolProviderMcpUpdateByProviderIdResponse) export const byProviderId2 = { - get: get73, + get: get77, } export const update4 = { @@ -3546,7 +3611,7 @@ export const delete17 = { post: post61, } -export const get74 = oc +export const get78 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3557,11 +3622,11 @@ export const get74 = oc .input(z.object({ query: zGetWorkspacesCurrentToolProviderWorkflowGetQuery.optional() })) .output(zGetWorkspacesCurrentToolProviderWorkflowGetResponse) -export const get75 = { - get: get74, +export const get79 = { + get: get78, } -export const get76 = oc +export const get80 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3573,7 +3638,7 @@ export const get76 = oc .output(zGetWorkspacesCurrentToolProviderWorkflowToolsResponse) export const tools4 = { - get: get76, + get: get80, } export const post62 = oc @@ -3594,7 +3659,7 @@ export const update5 = { export const workflow = { create: create2, delete: delete17, - get: get75, + get: get79, tools: tools4, update: update5, } @@ -3606,7 +3671,7 @@ export const toolProvider = { workflow, } -export const get77 = oc +export const get81 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3618,10 +3683,10 @@ export const get77 = oc .output(zGetWorkspacesCurrentToolProvidersResponse) export const toolProviders = { - get: get77, + get: get81, } -export const get78 = oc +export const get82 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3632,10 +3697,10 @@ export const get78 = oc .output(zGetWorkspacesCurrentToolsApiResponse) export const api2 = { - get: get78, + get: get82, } -export const get79 = oc +export const get83 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3646,10 +3711,10 @@ export const get79 = oc .output(zGetWorkspacesCurrentToolsBuiltinResponse) export const builtin2 = { - get: get79, + get: get83, } -export const get80 = oc +export const get84 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3660,10 +3725,10 @@ export const get80 = oc .output(zGetWorkspacesCurrentToolsMcpResponse) export const mcp2 = { - get: get80, + get: get84, } -export const get81 = oc +export const get85 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3674,7 +3739,7 @@ export const get81 = oc .output(zGetWorkspacesCurrentToolsWorkflowResponse) export const workflow2 = { - get: get81, + get: get85, } export const tools5 = { @@ -3684,7 +3749,7 @@ export const tools5 = { workflow: workflow2, } -export const get82 = oc +export const get86 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3696,13 +3761,13 @@ export const get82 = oc .output(zGetWorkspacesCurrentTriggerProviderByProviderIconResponse) export const icon3 = { - get: get82, + get: get86, } /** * Get info for a trigger provider */ -export const get83 = oc +export const get87 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3715,7 +3780,7 @@ export const get83 = oc .output(zGetWorkspacesCurrentTriggerProviderByProviderInfoResponse) export const info3 = { - get: get83, + get: get87, } /** @@ -3736,7 +3801,7 @@ export const delete18 = oc /** * Get OAuth client configuration for a provider */ -export const get84 = oc +export const get88 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3770,7 +3835,7 @@ export const post63 = oc export const client = { delete: delete18, - get: get84, + get: get88, post: post63, } @@ -3837,7 +3902,7 @@ export const create3 = { /** * Get the request logs for a subscription instance for a trigger provider */ -export const get85 = oc +export const get89 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3858,7 +3923,7 @@ export const get85 = oc ) export const bySubscriptionBuilderId2 = { - get: get85, + get: get89, } export const logs = { @@ -3932,7 +3997,7 @@ export const verifyAndUpdate = { /** * Get a subscription instance for a trigger provider */ -export const get86 = oc +export const get90 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3953,7 +4018,7 @@ export const get86 = oc ) export const bySubscriptionBuilderId5 = { - get: get86, + get: get90, } export const builder = { @@ -3968,7 +4033,7 @@ export const builder = { /** * List all trigger subscriptions for the current tenant's provider */ -export const get87 = oc +export const get91 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -3981,13 +4046,13 @@ export const get87 = oc .output(zGetWorkspacesCurrentTriggerProviderByProviderSubscriptionsListResponse) export const list4 = { - get: get87, + get: get91, } /** * Initiate OAuth authorization flow for a trigger provider */ -export const get88 = oc +export const get92 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4004,7 +4069,7 @@ export const get88 = oc .output(zGetWorkspacesCurrentTriggerProviderByProviderSubscriptionsOauthAuthorizeResponse) export const authorize = { - get: get88, + get: get92, } export const oauth3 = { @@ -4121,7 +4186,7 @@ export const triggerProvider = { /** * List all trigger providers for the current tenant */ -export const get89 = oc +export const get93 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4133,21 +4198,10 @@ export const get89 = oc .output(zGetWorkspacesCurrentTriggersResponse) export const triggers = { - get: get89, + get: get93, } -export const post71 = oc - .route({ - inputStructure: 'detailed', - method: 'POST', - operationId: 'postWorkspacesCurrent', - path: '/workspaces/current', - tags: ['console'], - }) - .output(zPostWorkspacesCurrentResponse) - export const current = { - post: post71, agentProvider, agentProviders, customizedSnippets, @@ -4160,6 +4214,7 @@ export const current = { permission, plugin: plugin2, rbac, + summary: summary2, toolLabels, toolProvider, toolProviders, @@ -4168,7 +4223,7 @@ export const current = { triggers, } -export const post72 = oc +export const post71 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -4181,14 +4236,24 @@ export const post72 = oc .output(zPostWorkspacesCustomConfigWebappLogoUploadResponse) export const upload2 = { - post: post72, + post: post71, } export const webappLogo = { upload: upload2, } -export const post73 = oc +export const get94 = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'getWorkspacesCustomConfig', + path: '/workspaces/custom-config', + tags: ['console'], + }) + .output(zGetWorkspacesCustomConfigResponse) + +export const post72 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -4200,11 +4265,12 @@ export const post73 = oc .output(zPostWorkspacesCustomConfigResponse) export const customConfig = { - post: post73, + get: get94, + post: post72, webappLogo, } -export const post74 = oc +export const post73 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -4216,10 +4282,10 @@ export const post74 = oc .output(zPostWorkspacesInfoResponse) export const info4 = { - post: post74, + post: post73, } -export const post75 = oc +export const post74 = oc .route({ inputStructure: 'detailed', method: 'POST', @@ -4231,10 +4297,10 @@ export const post75 = oc .output(zPostWorkspacesSwitchResponse) export const switch3 = { - post: post75, + post: post74, } -export const get90 = oc +export const get95 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4246,7 +4312,7 @@ export const get90 = oc .output(zGetWorkspacesByTenantIdModelProvidersByProviderByIconTypeByLangResponse) export const byLang = { - get: get90, + get: get95, } export const byIconType = { @@ -4265,7 +4331,7 @@ export const byTenantId = { modelProviders: modelProviders2, } -export const get91 = oc +export const get96 = oc .route({ inputStructure: 'detailed', method: 'GET', @@ -4276,7 +4342,7 @@ export const get91 = oc .output(zGetWorkspacesResponse) export const workspaces = { - get: get91, + get: get96, current, customConfig, info: info4, diff --git a/packages/contracts/generated/api/console/workspaces/types.gen.ts b/packages/contracts/generated/api/console/workspaces/types.gen.ts index 350b46dd1a2..a7c53ab7a4b 100644 --- a/packages/contracts/generated/api/console/workspaces/types.gen.ts +++ b/packages/contracts/generated/api/console/workspaces/types.gen.ts @@ -8,22 +8,6 @@ export type TenantListResponse = { workspaces: Array } -export type TenantInfoResponse = { - created_at?: number | null - custom_config?: WorkspaceCustomConfigResponse | null - id: string - in_trial?: boolean | null - name?: string | null - next_credit_reset_date?: number | null - plan?: string | null - role?: string | null - status?: string | null - trial_credits?: number | null - trial_credits_exhausted_at?: number | null - trial_credits_used?: number | null - trial_end_reason?: string | null -} - export type AgentProviderResponse = { [key: string]: unknown } @@ -227,6 +211,24 @@ export type ModelProviderListResponse = { data: Array } +export type ModelProviderCreditsResponse = { + exhausted_at: number | null + is_exhausted: boolean + is_unlimited: boolean + next_credit_reset_date: number | null + pool_type: 'paid' | 'trial' | null + quota_limit: number | null + quota_used: number | null + remaining_credits: number | null +} + +export type ModelProviderSummaryListResponse = { + data: Array + plugins: { + [key: string]: ModelProviderPluginSummaryResponse + } +} + export type ModelProviderPaymentCheckoutUrlResponse = { payment_link: string } @@ -417,6 +419,10 @@ export type ParserPluginIdentifiers = { plugin_unique_identifiers: Array } +export type PluginInstalledIdsResponse = { + plugin_ids: Array +} + export type PluginListResponse = { plugins: Array total: number @@ -475,6 +481,7 @@ export type PluginTaskResponse = { export type ParserUninstall = { plugin_installation_id: string + preserve_credentials?: boolean } export type ParserGithubUpgrade = { @@ -635,6 +642,14 @@ export type WorkspaceAccessMatrix = { pagination?: Pagination | null } +export type CurrentWorkspaceSummaryResponse = { + credits: number | null + id: string + name: string + plan: CloudPlan | null + role: TenantAccountRole +} + export type ToolLabelListResponse = Array export type ApiToolProviderAddPayload = { @@ -1000,6 +1015,11 @@ export type TriggerOAuthAuthorizeResponse = { export type TriggerProviderListResponse = Array +export type WorkspaceCustomConfigResponse = { + remove_webapp_brand?: boolean | null + replace_webapp_logo?: string | null +} + export type WorkspaceCustomConfigPayload = { remove_webapp_brand?: boolean | null replace_webapp_logo?: string | null @@ -1033,15 +1053,10 @@ export type TenantListItemResponse = { id: string last_opened_at?: number | null name?: string | null - plan?: string | null + plan?: CloudPlan | null status?: string | null } -export type WorkspaceCustomConfigResponse = { - remove_webapp_brand?: boolean | null - replace_webapp_logo?: string | null -} - export type SnippetListItemResponse = { author_name: string | null created_at: number @@ -1191,6 +1206,30 @@ export type ProviderResponse = { tenant_id: string } +export type ModelProviderSummaryResponse = { + configurate_methods: Array + custom_configuration: ModelProviderCustomConfigurationSummaryResponse + description?: I18nObject | null + icon_small?: I18nObject | null + icon_small_dark?: I18nObject | null + is_configured: boolean + label: I18nObject + plugin_id: string + preferred_provider_type: ProviderType + provider: string + supported_model_types: Array + system_configuration: ModelProviderSystemConfigurationSummaryResponse +} + +export type ModelProviderPluginSummaryResponse = { + installation_id: string + plugin_id: string + plugin_unique_identifier: string + runtime_type: string + source: PluginInstallationSource + version: string +} + export type ModelType = 'llm' | 'moderation' | 'rerank' | 'speech2text' | 'text-embedding' | 'tts' export type ModelWithProviderEntityResponse = { @@ -1502,6 +1541,10 @@ export type AccessPolicyRole = { role_tag?: string } +export type CloudPlan = 'professional' | 'sandbox' | 'team' + +export type TenantAccountRole = 'admin' | 'dataset_operator' | 'editor' | 'normal' | 'owner' + export type ToolLabel = { icon: string label: I18nObject @@ -1661,6 +1704,22 @@ export type TriggerProviderSubscriptionApiEntity = { workflows_in_use: number } +export type TenantInfoResponse = { + created_at?: number | null + custom_config?: WorkspaceCustomConfigResponse | null + id: string + in_trial?: boolean | null + name?: string | null + next_credit_reset_date?: number | null + plan?: CloudPlan | null + role?: string | null + status?: string | null + trial_credits?: number | null + trial_credits_exhausted_at?: number | null + trial_credits_used?: number | null + trial_end_reason?: string | null +} + export type PluginDependencyType = 'github' | 'marketplace' | 'package' export type Github = { @@ -1729,6 +1788,21 @@ export type SystemConfigurationResponse = { quota_configurations?: Array } +export type ModelProviderCustomConfigurationSummaryResponse = { + available_credentials: Array + current_credential_id?: string | null + current_credential_name?: string | null + current_credential_usable: boolean + has_custom_models: boolean + status: CustomConfigurationStatus +} + +export type ModelProviderSystemConfigurationSummaryResponse = { + enabled: boolean +} + +export type PluginInstallationSource = 'github' | 'marketplace' | 'package' | 'remote' + export type ModelFeature = | 'agent-thought' | 'audio' @@ -1899,8 +1973,6 @@ export type PluginInstallTaskPluginStatus = { export type PluginInstallTaskStatus = 'failed' | 'pending' | 'running' | 'success' -export type PluginInstallationSource = 'github' | 'marketplace' | 'package' | 'remote' - export type PluginDeclarationResponse = { agent_strategy?: { [key: string]: unknown @@ -2262,6 +2334,8 @@ export type ToolParameterType = | 'array' | 'boolean' | 'checkbox' + | 'date' + | 'date-range' | 'dynamic-select' | 'file' | 'files' @@ -2429,20 +2503,6 @@ export type GetWorkspacesResponses = { export type GetWorkspacesResponse = GetWorkspacesResponses[keyof GetWorkspacesResponses] -export type PostWorkspacesCurrentData = { - body?: never - path?: never - query?: never - url: '/workspaces/current' -} - -export type PostWorkspacesCurrentResponses = { - 200: TenantInfoResponse -} - -export type PostWorkspacesCurrentResponse = - PostWorkspacesCurrentResponses[keyof PostWorkspacesCurrentResponses] - export type GetWorkspacesCurrentAgentProviderByProviderNameData = { body?: never path: { @@ -3028,6 +3088,34 @@ export type GetWorkspacesCurrentModelProvidersResponses = { export type GetWorkspacesCurrentModelProvidersResponse = GetWorkspacesCurrentModelProvidersResponses[keyof GetWorkspacesCurrentModelProvidersResponses] +export type GetWorkspacesCurrentModelProvidersCreditsData = { + body?: never + path?: never + query?: never + url: '/workspaces/current/model-providers/credits' +} + +export type GetWorkspacesCurrentModelProvidersCreditsResponses = { + 200: ModelProviderCreditsResponse +} + +export type GetWorkspacesCurrentModelProvidersCreditsResponse = + GetWorkspacesCurrentModelProvidersCreditsResponses[keyof GetWorkspacesCurrentModelProvidersCreditsResponses] + +export type GetWorkspacesCurrentModelProvidersSummaryData = { + body?: never + path?: never + query?: never + url: '/workspaces/current/model-providers/summary' +} + +export type GetWorkspacesCurrentModelProvidersSummaryResponses = { + 200: ModelProviderSummaryListResponse +} + +export type GetWorkspacesCurrentModelProvidersSummaryResponse = + GetWorkspacesCurrentModelProvidersSummaryResponses[keyof GetWorkspacesCurrentModelProvidersSummaryResponses] + export type GetWorkspacesCurrentModelProvidersByProviderCheckoutUrlData = { body?: never path: { @@ -3574,6 +3662,22 @@ export type PostWorkspacesCurrentPluginInstallPkgResponses = { export type PostWorkspacesCurrentPluginInstallPkgResponse = PostWorkspacesCurrentPluginInstallPkgResponses[keyof PostWorkspacesCurrentPluginInstallPkgResponses] +export type GetWorkspacesCurrentPluginInstalledIdsData = { + body?: never + path?: never + query: { + category: 'agent-strategy' | 'datasource' | 'extension' | 'model' | 'tool' | 'trigger' + } + url: '/workspaces/current/plugin/installed-ids' +} + +export type GetWorkspacesCurrentPluginInstalledIdsResponses = { + 200: PluginInstalledIdsResponse +} + +export type GetWorkspacesCurrentPluginInstalledIdsResponse = + GetWorkspacesCurrentPluginInstalledIdsResponses[keyof GetWorkspacesCurrentPluginInstalledIdsResponses] + export type GetWorkspacesCurrentPluginListData = { body?: never path?: never @@ -3887,8 +3991,11 @@ export type GetWorkspacesCurrentPluginByCategoryListData = { category: string } query?: { + language?: 'en_US' | 'ja_JP' | 'pt_BR' | 'zh_Hans' page?: number page_size?: number + query?: string + tags?: Array } url: '/workspaces/current/plugin/{category}/list' } @@ -4625,6 +4732,24 @@ export type GetWorkspacesCurrentRbacWorkspaceDatasetsAccessPolicyResponses = { export type GetWorkspacesCurrentRbacWorkspaceDatasetsAccessPolicyResponse = GetWorkspacesCurrentRbacWorkspaceDatasetsAccessPolicyResponses[keyof GetWorkspacesCurrentRbacWorkspaceDatasetsAccessPolicyResponses] +export type GetWorkspacesCurrentSummaryData = { + body?: never + path?: never + query?: never + url: '/workspaces/current/summary' +} + +export type GetWorkspacesCurrentSummaryErrors = { + 409: unknown +} + +export type GetWorkspacesCurrentSummaryResponses = { + 200: CurrentWorkspaceSummaryResponse +} + +export type GetWorkspacesCurrentSummaryResponse = + GetWorkspacesCurrentSummaryResponses[keyof GetWorkspacesCurrentSummaryResponses] + export type GetWorkspacesCurrentToolLabelsData = { body?: never path?: never @@ -5524,6 +5649,20 @@ export type GetWorkspacesCurrentTriggersResponses = { export type GetWorkspacesCurrentTriggersResponse = GetWorkspacesCurrentTriggersResponses[keyof GetWorkspacesCurrentTriggersResponses] +export type GetWorkspacesCustomConfigData = { + body?: never + path?: never + query?: never + url: '/workspaces/custom-config' +} + +export type GetWorkspacesCustomConfigResponses = { + 200: WorkspaceCustomConfigResponse +} + +export type GetWorkspacesCustomConfigResponse = + GetWorkspacesCustomConfigResponses[keyof GetWorkspacesCustomConfigResponses] + export type PostWorkspacesCustomConfigData = { body: WorkspaceCustomConfigPayload path?: never diff --git a/packages/contracts/generated/api/console/workspaces/zod.gen.ts b/packages/contracts/generated/api/console/workspaces/zod.gen.ts index 77bfc7dfe19..8bcd255b3c9 100644 --- a/packages/contracts/generated/api/console/workspaces/zod.gen.ts +++ b/packages/contracts/generated/api/console/workspaces/zod.gen.ts @@ -158,6 +158,20 @@ export const zMemberRoleUpdatePayload = z.object({ role: z.string(), }) +/** + * ModelProviderCreditsResponse + */ +export const zModelProviderCreditsResponse = z.object({ + exhausted_at: z.int().nullable(), + is_exhausted: z.boolean(), + is_unlimited: z.boolean(), + next_credit_reset_date: z.int().nullable(), + pool_type: z.enum(['paid', 'trial']).nullable(), + quota_limit: z.int().nullable(), + quota_used: z.int().nullable(), + remaining_credits: z.int().nullable(), +}) + /** * ModelProviderPaymentCheckoutUrlResponse */ @@ -283,6 +297,13 @@ export const zParserPluginIdentifiers = z.object({ plugin_unique_identifiers: z.array(z.string()), }) +/** + * PluginInstalledIdsResponse + */ +export const zPluginInstalledIdsResponse = z.object({ + plugin_ids: z.array(z.string()), +}) + /** * ParserLatest */ @@ -314,6 +335,7 @@ export const zPluginReadmeResponse = z.object({ */ export const zParserUninstall = z.object({ plugin_installation_id: z.string(), + preserve_credentials: z.boolean().optional().default(false), }) /** @@ -564,6 +586,14 @@ export const zTriggerProviderErrorResponse = z.object({ error: z.string(), }) +/** + * WorkspaceCustomConfigResponse + */ +export const zWorkspaceCustomConfigResponse = z.object({ + remove_webapp_brand: z.boolean().nullish(), + replace_webapp_logo: z.string().nullish(), +}) + /** * WorkspaceCustomConfigPayload */ @@ -593,69 +623,6 @@ export const zSwitchWorkspacePayload = z.object({ tenant_id: z.string(), }) -/** - * TenantListItemResponse - */ -export const zTenantListItemResponse = z.object({ - created_at: z.int().nullish(), - current: z.boolean(), - id: z.string(), - last_opened_at: z.int().nullish(), - name: z.string().nullish(), - plan: z.string().nullish(), - status: z.string().nullish(), -}) - -/** - * TenantListResponse - */ -export const zTenantListResponse = z.object({ - workspaces: z.array(zTenantListItemResponse), -}) - -/** - * WorkspaceCustomConfigResponse - */ -export const zWorkspaceCustomConfigResponse = z.object({ - remove_webapp_brand: z.boolean().nullish(), - replace_webapp_logo: z.string().nullish(), -}) - -/** - * TenantInfoResponse - */ -export const zTenantInfoResponse = z.object({ - created_at: z.int().nullish(), - custom_config: zWorkspaceCustomConfigResponse.nullish(), - id: z.string(), - in_trial: z.boolean().nullish(), - name: z.string().nullish(), - next_credit_reset_date: z.int().nullish(), - plan: z.string().nullish(), - role: z.string().nullish(), - status: z.string().nullish(), - trial_credits: z.int().nullish(), - trial_credits_exhausted_at: z.int().nullish(), - trial_credits_used: z.int().nullish(), - trial_end_reason: z.string().nullish(), -}) - -/** - * WorkspaceTenantResultResponse - */ -export const zWorkspaceTenantResultResponse = z.object({ - result: z.string(), - tenant: zTenantInfoResponse, -}) - -/** - * SwitchWorkspaceResponse - */ -export const zSwitchWorkspaceResponse = z.object({ - new_tenant: zTenantInfoResponse, - result: z.string(), -}) - /** * IconInfo * @@ -1244,6 +1211,53 @@ export const zWorkspaceAccessMatrix = z.object({ pagination: zPagination.nullish(), }) +/** + * CloudPlan + * + * Enum representing user plan types in the cloud platform. + * + * SANDBOX: Free/default plan with limited features + * PROFESSIONAL: Professional paid plan + * TEAM: Team collaboration paid plan + */ +export const zCloudPlan = z.enum(['professional', 'sandbox', 'team']) + +/** + * TenantListItemResponse + */ +export const zTenantListItemResponse = z.object({ + created_at: z.int().nullish(), + current: z.boolean(), + id: z.string(), + last_opened_at: z.int().nullish(), + name: z.string().nullish(), + plan: zCloudPlan.nullish(), + status: z.string().nullish(), +}) + +/** + * TenantListResponse + */ +export const zTenantListResponse = z.object({ + workspaces: z.array(zTenantListItemResponse), +}) + +/** + * TenantAccountRole + */ +export const zTenantAccountRole = z.enum(['admin', 'dataset_operator', 'editor', 'normal', 'owner']) + +/** + * CurrentWorkspaceSummaryResponse + */ +export const zCurrentWorkspaceSummaryResponse = z.object({ + credits: z.int().nullable(), + id: z.string(), + name: z.string(), + plan: zCloudPlan.nullable(), + role: zTenantAccountRole, +}) + /** * ToolEmojiIcon */ @@ -1553,6 +1567,41 @@ export const zTriggerProviderSubscriptionListResponse = z.array( zTriggerProviderSubscriptionApiEntity, ) +/** + * TenantInfoResponse + */ +export const zTenantInfoResponse = z.object({ + created_at: z.int().nullish(), + custom_config: zWorkspaceCustomConfigResponse.nullish(), + id: z.string(), + in_trial: z.boolean().nullish(), + name: z.string().nullish(), + next_credit_reset_date: z.int().nullish(), + plan: zCloudPlan.nullish(), + role: z.string().nullish(), + status: z.string().nullish(), + trial_credits: z.int().nullish(), + trial_credits_exhausted_at: z.int().nullish(), + trial_credits_used: z.int().nullish(), + trial_end_reason: z.string().nullish(), +}) + +/** + * WorkspaceTenantResultResponse + */ +export const zWorkspaceTenantResultResponse = z.object({ + result: z.string(), + tenant: zTenantInfoResponse, +}) + +/** + * SwitchWorkspaceResponse + */ +export const zSwitchWorkspaceResponse = z.object({ + new_tenant: zTenantInfoResponse, + result: z.string(), +}) + /** * PluginDependencyType */ @@ -1612,6 +1661,30 @@ export const zConfigurateMethod = z.enum(['customizable-model', 'predefined-mode */ export const zProviderType = z.enum(['custom', 'system']) +/** + * ModelProviderSystemConfigurationSummaryResponse + */ +export const zModelProviderSystemConfigurationSummaryResponse = z.object({ + enabled: z.boolean(), +}) + +/** + * PluginInstallationSource + */ +export const zPluginInstallationSource = z.enum(['github', 'marketplace', 'package', 'remote']) + +/** + * ModelProviderPluginSummaryResponse + */ +export const zModelProviderPluginSummaryResponse = z.object({ + installation_id: z.string(), + plugin_id: z.string(), + plugin_unique_identifier: z.string(), + runtime_type: z.string(), + source: zPluginInstallationSource, + version: z.string(), +}) + /** * ModelFeature * @@ -1802,6 +1875,46 @@ export const zAvailableModelListResponse = z.object({ data: z.array(zProviderWithModelsResponse), }) +/** + * ModelProviderCustomConfigurationSummaryResponse + */ +export const zModelProviderCustomConfigurationSummaryResponse = z.object({ + available_credentials: z.array(zCredentialConfiguration), + current_credential_id: z.string().nullish(), + current_credential_name: z.string().nullish(), + current_credential_usable: z.boolean(), + has_custom_models: z.boolean(), + status: zCustomConfigurationStatus, +}) + +/** + * ModelProviderSummaryResponse + * + * Fields required to render the collapsed model-provider list. + */ +export const zModelProviderSummaryResponse = z.object({ + configurate_methods: z.array(zConfigurateMethod), + custom_configuration: zModelProviderCustomConfigurationSummaryResponse, + description: zI18nObject.nullish(), + icon_small: zI18nObject.nullish(), + icon_small_dark: zI18nObject.nullish(), + is_configured: z.boolean(), + label: zI18nObject, + plugin_id: z.string(), + preferred_provider_type: zProviderType, + provider: z.string(), + supported_model_types: z.array(zModelType), + system_configuration: zModelProviderSystemConfigurationSummaryResponse, +}) + +/** + * ModelProviderSummaryListResponse + */ +export const zModelProviderSummaryListResponse = z.object({ + data: z.array(zModelProviderSummaryResponse), + plugins: z.record(z.string(), zModelProviderPluginSummaryResponse), +}) + /** * TenantPluginAutoUpgradeStrategySetting */ @@ -1947,11 +2060,6 @@ export const zPluginTaskResponse = z.object({ task: zPluginInstallTask, }) -/** - * PluginInstallationSource - */ -export const zPluginInstallationSource = z.enum(['github', 'marketplace', 'package', 'remote']) - /** * PluginBundleDependencyType */ @@ -2381,6 +2489,8 @@ export const zToolParameterType = z.enum([ 'array', 'boolean', 'checkbox', + 'date', + 'date-range', 'dynamic-select', 'file', 'files', @@ -3402,11 +3512,6 @@ export const zAccountWithRoleListResponseWritable = z.object({ */ export const zGetWorkspacesResponse = zTenantListResponse -/** - * Success - */ -export const zPostWorkspacesCurrentResponse = zTenantInfoResponse - export const zGetWorkspacesCurrentAgentProviderByProviderNamePath = z.object({ provider_name: z.string(), }) @@ -3695,6 +3800,16 @@ export const zGetWorkspacesCurrentModelProvidersQuery = z.object({ */ export const zGetWorkspacesCurrentModelProvidersResponse = zModelProviderListResponse +/** + * Model provider credits retrieved successfully + */ +export const zGetWorkspacesCurrentModelProvidersCreditsResponse = zModelProviderCreditsResponse + +/** + * Model provider summaries retrieved successfully + */ +export const zGetWorkspacesCurrentModelProvidersSummaryResponse = zModelProviderSummaryListResponse + export const zGetWorkspacesCurrentModelProvidersByProviderCheckoutUrlPath = z.object({ provider: z.string(), }) @@ -4070,6 +4185,15 @@ export const zPostWorkspacesCurrentPluginInstallPkgBody = zParserPluginIdentifie */ export const zPostWorkspacesCurrentPluginInstallPkgResponse = zPluginInstallTaskStartResponse +export const zGetWorkspacesCurrentPluginInstalledIdsQuery = z.object({ + category: z.enum(['agent-strategy', 'datasource', 'extension', 'model', 'tool', 'trigger']), +}) + +/** + * Success + */ +export const zGetWorkspacesCurrentPluginInstalledIdsResponse = zPluginInstalledIdsResponse + export const zGetWorkspacesCurrentPluginListQuery = z.object({ page: z.int().gte(1).optional().default(1), page_size: z.int().gte(1).lte(256).optional().default(256), @@ -4240,8 +4364,11 @@ export const zGetWorkspacesCurrentPluginByCategoryListPath = z.object({ }) export const zGetWorkspacesCurrentPluginByCategoryListQuery = z.object({ + language: z.enum(['en_US', 'ja_JP', 'pt_BR', 'zh_Hans']).optional().default('en_US'), page: z.int().gte(1).optional().default(1), page_size: z.int().gte(1).lte(256).optional().default(256), + query: z.string().max(256).optional().default(''), + tags: z.array(z.string()).max(128).optional(), }) /** @@ -4694,6 +4821,11 @@ export const zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPoliciesByPolicyIdR */ export const zGetWorkspacesCurrentRbacWorkspaceDatasetsAccessPolicyResponse = zWorkspaceAccessMatrix +/** + * Success + */ +export const zGetWorkspacesCurrentSummaryResponse = zCurrentWorkspaceSummaryResponse + /** * Tool labels retrieved successfully */ @@ -5247,6 +5379,11 @@ export const zPostWorkspacesCurrentTriggerProviderBySubscriptionIdSubscriptionsU */ export const zGetWorkspacesCurrentTriggersResponse = zTriggerProviderListResponse +/** + * Success + */ +export const zGetWorkspacesCustomConfigResponse = zWorkspaceCustomConfigResponse + export const zPostWorkspacesCustomConfigBody = zWorkspaceCustomConfigPayload /** diff --git a/packages/contracts/generated/api/openapi/types.gen.ts b/packages/contracts/generated/api/openapi/types.gen.ts index 4c1f59b3ce7..d0a00673acf 100644 --- a/packages/contracts/generated/api/openapi/types.gen.ts +++ b/packages/contracts/generated/api/openapi/types.gen.ts @@ -125,6 +125,8 @@ export type CheckDependenciesResult = { leaked_dependencies?: Array } +export type DeploymentEdition = 'CLOUD' | 'COMMUNITY' | 'ENTERPRISE' + export type DeviceCodeRequest = { client_id: string device_label: string @@ -392,7 +394,7 @@ export type RevokeResponse = { } export type ServerVersionResponse = { - edition: 'CLOUD' | 'SELF_HOSTED' + edition: DeploymentEdition version: string } diff --git a/packages/contracts/generated/api/openapi/zod.gen.ts b/packages/contracts/generated/api/openapi/zod.gen.ts index 5c45490ba64..fb584d54e9d 100644 --- a/packages/contracts/generated/api/openapi/zod.gen.ts +++ b/packages/contracts/generated/api/openapi/zod.gen.ts @@ -51,7 +51,7 @@ export const zAppDescribeResponse = z.object({ */ export const zAppDslExportQuery = z.object({ include_secret: z.boolean().optional().default(false), - workflow_id: z.string().nullish(), + workflow_id: z.uuid().nullish(), }) /** @@ -141,6 +141,13 @@ export const zAppRunRequest = z.object({ workspace_id: z.string().nullish(), }) +/** + * DeploymentEdition + * + * Enum representing the deployment edition of the platform. + */ +export const zDeploymentEdition = z.enum(['CLOUD', 'COMMUNITY', 'ENTERPRISE']) + /** * DeviceCodeRequest */ @@ -499,7 +506,7 @@ export const zRevokeResponse = z.object({ * Meta endpoint payload for `GET /openapi/v1/_version` — no auth required. */ export const zServerVersionResponse = z.object({ - edition: z.enum(['CLOUD', 'SELF_HOSTED']), + edition: zDeploymentEdition, version: z.string(), }) @@ -577,7 +584,7 @@ export const zAppListQuery = z.object({ mode: zSupportedAppType.nullish(), name: z.string().max(200).nullish(), page: z.int().gte(1).optional().default(1), - workspace_id: z.string(), + workspace_id: z.uuid(), }) /** @@ -746,7 +753,7 @@ export const zGetAppsQuery = z.object({ mode: z.enum(['advanced-chat', 'agent-chat', 'chat', 'completion', 'workflow']).optional(), name: z.string().max(200).optional(), page: z.int().gte(1).optional().default(1), - workspace_id: z.string(), + workspace_id: z.uuid(), }) /** @@ -782,7 +789,7 @@ export const zGetAppsByAppIdDslPath = z.object({ export const zGetAppsByAppIdDslQuery = z.object({ include_secret: z.boolean().optional().default(false), - workflow_id: z.string().optional(), + workflow_id: z.uuid().optional(), }) /** diff --git a/packages/contracts/generated/api/service/types.gen.ts b/packages/contracts/generated/api/service/types.gen.ts index 671fa451c08..ca880a75653 100644 --- a/packages/contracts/generated/api/service/types.gen.ts +++ b/packages/contracts/generated/api/service/types.gen.ts @@ -722,6 +722,8 @@ export type DocumentStatusResponse = { completed_at: number | null completed_segments?: number | null error: string | null + error_code?: string | null + estimated_vector_space_mb?: number | null id: string indexing_status: string parsing_completed_at: number | null @@ -730,6 +732,7 @@ export type DocumentStatusResponse = { splitting_completed_at: number | null stopped_at: number | null total_segments?: number | null + vector_space_limit_mb?: number | null } export type DocumentTextCreatePayload = { @@ -2414,6 +2417,7 @@ export type PostDatasetsByDatasetIdDocumentCreateByFileErrors = { 400: unknown 401: unknown 403: unknown + 413: unknown } export type PostDatasetsByDatasetIdDocumentCreateByFileResponses = { @@ -2461,6 +2465,7 @@ export type PostDatasetsByDatasetIdDocumentCreateByFile2Errors = { 400: unknown 401: unknown 403: unknown + 413: unknown } export type PostDatasetsByDatasetIdDocumentCreateByFile2Responses = { @@ -2681,6 +2686,7 @@ export type PatchDatasetsByDatasetIdDocumentsByDocumentIdErrors = { 401: unknown 403: unknown 404: unknown + 413: unknown } export type PatchDatasetsByDatasetIdDocumentsByDocumentIdResponses = { @@ -2966,6 +2972,7 @@ export type PostDatasetsByDatasetIdDocumentsByDocumentIdUpdateByFileErrors = { 401: unknown 403: unknown 404: unknown + 413: unknown } export type PostDatasetsByDatasetIdDocumentsByDocumentIdUpdateByFileResponses = { @@ -3017,6 +3024,7 @@ export type PostDatasetsByDatasetIdDocumentsByDocumentIdUpdateByFile2Errors = { 401: unknown 403: unknown 404: unknown + 413: unknown } export type PostDatasetsByDatasetIdDocumentsByDocumentIdUpdateByFile2Responses = { diff --git a/packages/contracts/generated/api/service/zod.gen.ts b/packages/contracts/generated/api/service/zod.gen.ts index 78fcc65cdad..2857af1401f 100644 --- a/packages/contracts/generated/api/service/zod.gen.ts +++ b/packages/contracts/generated/api/service/zod.gen.ts @@ -872,6 +872,8 @@ export const zDocumentStatusResponse = z.object({ completed_at: z.int().nullable(), completed_segments: z.int().nullish(), error: z.string().nullable(), + error_code: z.string().nullish(), + estimated_vector_space_mb: z.int().nullish(), id: z.string(), indexing_status: z.string(), parsing_completed_at: z.int().nullable(), @@ -880,6 +882,7 @@ export const zDocumentStatusResponse = z.object({ splitting_completed_at: z.int().nullable(), stopped_at: z.int().nullable(), total_segments: z.int().nullish(), + vector_space_limit_mb: z.int().nullish(), }) /** @@ -1263,7 +1266,7 @@ export const zMetadataArgs = z.object({ * MetadataDetail */ export const zMetadataDetail = z.object({ - id: z.string(), + id: z.uuid(), name: z.string(), value: z.union([z.string(), z.int(), z.number()]).nullish(), }) @@ -1272,7 +1275,7 @@ export const zMetadataDetail = z.object({ * DocumentMetadataOperation */ export const zDocumentMetadataOperation = z.object({ - document_id: z.string(), + document_id: z.uuid(), metadata_list: z.array(zMetadataDetail), partial_update: z.boolean().optional().default(false), }) @@ -1894,7 +1897,7 @@ export const zUserActionConfig = z.object({ * ValueSourceType * * ValueSourceType records whether the value comes from a static setting - * in form definiton, or a variable while the workflow is running. + * in form definition, or a variable while the workflow is running. */ export const zValueSourceType = z.enum(['constant', 'variable']) diff --git a/packages/contracts/generated/api/web/types.gen.ts b/packages/contracts/generated/api/web/types.gen.ts index 14e9bbf2e53..f95ed066bb5 100644 --- a/packages/contracts/generated/api/web/types.gen.ts +++ b/packages/contracts/generated/api/web/types.gen.ts @@ -43,6 +43,16 @@ export type AppMetaResponse = { } } +export type AppMode = + | 'advanced-chat' + | 'agent' + | 'agent-chat' + | 'channel' + | 'chat' + | 'completion' + | 'rag-pipeline' + | 'workflow' + export type AppPermissionQuery = { appId: string } @@ -423,6 +433,8 @@ export type RetrieverResource = { word_count?: number | null } +export type SsoProtocol = 'oauth2' | 'oidc' | 'saml' + export type SavedMessageCreatePayload = { message_id: string } @@ -511,7 +523,6 @@ export type SystemFeatureModel = { enable_marketplace: boolean enable_social_oauth_login: boolean enable_step_by_step_tour: boolean - enable_trial_app: boolean is_allow_register: boolean is_email_setup: boolean knowledge_fs_enabled: boolean @@ -519,7 +530,7 @@ export type SystemFeatureModel = { plugin_installation_permission: PluginInstallationPermissionModel rbac_enabled: boolean sso_enforced_for_signin: boolean - sso_enforced_for_signin_protocol: string + sso_enforced_for_signin_protocol: SsoProtocol | null webapp_auth: WebAppAuthModel } @@ -562,7 +573,7 @@ export type WebAppAuthModel = { } export type WebAppAuthSsoModel = { - protocol: string + protocol: SsoProtocol | null } export type WebAppCustomConfigResponse = { @@ -576,6 +587,7 @@ export type WebAppSiteResponse = { custom_config?: WebAppCustomConfigResponse | null enable_site: boolean end_user_id?: string | null + mode: AppMode model_config?: WebModelConfigResponse | null plan: string site: WebSiteResponse diff --git a/packages/contracts/generated/api/web/zod.gen.ts b/packages/contracts/generated/api/web/zod.gen.ts index f41de0e0077..36ebdd700fc 100644 --- a/packages/contracts/generated/api/web/zod.gen.ts +++ b/packages/contracts/generated/api/web/zod.gen.ts @@ -39,6 +39,20 @@ export const zAppMetaResponse = z.object({ tool_icons: z.record(z.string(), z.unknown()).optional(), }) +/** + * AppMode + */ +export const zAppMode = z.enum([ + 'advanced-chat', + 'agent', + 'agent-chat', + 'channel', + 'chat', + 'completion', + 'rag-pipeline', + 'workflow', +]) + /** * AppPermissionQuery */ @@ -499,6 +513,11 @@ export const zRetrieverResource = z.object({ word_count: z.int().nullish(), }) +/** + * SSOProtocol + */ +export const zSsoProtocol = z.enum(['oauth2', 'oidc', 'saml']) + /** * SavedMessageCreatePayload */ @@ -641,7 +660,7 @@ export const zUserActionConfig = z.object({ * ValueSourceType * * ValueSourceType records whether the value comes from a static setting - * in form definiton, or a variable while the workflow is running. + * in form definition, or a variable while the workflow is running. */ export const zValueSourceType = z.enum(['constant', 'variable']) @@ -732,7 +751,7 @@ export const zVerificationTokenResponse = z.object({ * WebAppAuthSSOModel */ export const zWebAppAuthSsoModel = z.object({ - protocol: z.string().default(''), + protocol: zSsoProtocol.nullable(), }) /** @@ -744,7 +763,7 @@ export const zWebAppAuthModel = z.object({ allow_public_access: z.boolean().default(true), allow_sso: z.boolean().default(false), enabled: z.boolean().default(false), - sso_config: zWebAppAuthSsoModel.default({ protocol: '' }), + sso_config: zWebAppAuthSsoModel, }) /** @@ -772,7 +791,6 @@ export const zSystemFeatureModel = z.object({ enable_marketplace: z.boolean().default(false), enable_social_oauth_login: z.boolean().default(false), enable_step_by_step_tour: z.boolean().default(false), - enable_trial_app: z.boolean().default(false), is_allow_register: z.boolean().default(false), is_email_setup: z.boolean().default(false), knowledge_fs_enabled: z.boolean().default(false), @@ -783,15 +801,8 @@ export const zSystemFeatureModel = z.object({ }), rbac_enabled: z.boolean().default(false), sso_enforced_for_signin: z.boolean().default(false), - sso_enforced_for_signin_protocol: z.string().default(''), - webapp_auth: zWebAppAuthModel.default({ - allow_email_code_login: false, - allow_email_password_login: false, - allow_public_access: true, - allow_sso: false, - enabled: false, - sso_config: { protocol: '' }, - }), + sso_enforced_for_signin_protocol: zSsoProtocol.nullable(), + webapp_auth: zWebAppAuthModel, }) /** @@ -885,6 +896,7 @@ export const zWebAppSiteResponse = z.object({ custom_config: zWebAppCustomConfigResponse.nullish(), enable_site: z.boolean(), end_user_id: z.string().nullish(), + mode: zAppMode, model_config: zWebModelConfigResponse.nullish(), plan: z.string(), site: zWebSiteResponse, diff --git a/packages/contracts/generated/enterprise-app-deploy/orpc.gen.ts b/packages/contracts/generated/enterprise-app-deploy/orpc.gen.ts new file mode 100644 index 00000000000..6df55a62422 --- /dev/null +++ b/packages/contracts/generated/enterprise-app-deploy/orpc.gen.ts @@ -0,0 +1,307 @@ +// This file is auto-generated by @hey-api/openapi-ts + +import { oc } from '@orpc/contract' +import * as z from 'zod' +import { + zConsoleAccessServiceCreateEnvironmentApiKeyPath, + zConsoleAccessServiceCreateEnvironmentApiKeyResponse, + zConsoleAccessServiceDeleteEnvironmentApiKeyPath, + zConsoleAccessServiceDeleteEnvironmentApiKeyResponse, + zConsoleAccessServiceGetEnvironmentApiPath, + zConsoleAccessServiceGetEnvironmentApiResponse, + zConsoleAccessServiceGetEnvironmentMcpServerPath, + zConsoleAccessServiceGetEnvironmentMcpServerResponse, + zConsoleAccessServiceGetEnvironmentSitePath, + zConsoleAccessServiceGetEnvironmentSiteResponse, + zConsoleAccessServiceGetEnvironmentWebAppSubjectsPath, + zConsoleAccessServiceGetEnvironmentWebAppSubjectsResponse, + zConsoleAccessServiceListEnvironmentApiKeysPath, + zConsoleAccessServiceListEnvironmentApiKeysResponse, + zConsoleAccessServiceListEnvironmentTriggersPath, + zConsoleAccessServiceListEnvironmentTriggersResponse, + zConsoleAccessServiceResetEnvironmentSiteAccessTokenPath, + zConsoleAccessServiceResetEnvironmentSiteAccessTokenResponse, + zConsoleAccessServiceUpdateEnvironmentApiBody, + zConsoleAccessServiceUpdateEnvironmentApiPath, + zConsoleAccessServiceUpdateEnvironmentApiResponse, + zConsoleAccessServiceUpdateEnvironmentSiteBody, + zConsoleAccessServiceUpdateEnvironmentSitePath, + zConsoleAccessServiceUpdateEnvironmentSiteResponse, + zConsoleAccessServiceUpdateEnvironmentWebAppAccessModeBody, + zConsoleAccessServiceUpdateEnvironmentWebAppAccessModePath, + zConsoleAccessServiceUpdateEnvironmentWebAppAccessModeResponse, + zConsoleDeploymentServiceDeployWorkflowBody, + zConsoleDeploymentServiceDeployWorkflowPath, + zConsoleDeploymentServiceDeployWorkflowResponse, + zConsoleDeploymentServiceGetEnvironmentDeploymentPath, + zConsoleDeploymentServiceGetEnvironmentDeploymentResponse, + zConsoleDeploymentServiceGetWorkflowDeploymentOptionsPath, + zConsoleDeploymentServiceGetWorkflowDeploymentOptionsResponse, + zConsoleDeploymentServiceListAppEnvironmentsPath, + zConsoleDeploymentServiceListAppEnvironmentsResponse, + zConsoleDeploymentServiceListEnvironmentDeploymentsPath, + zConsoleDeploymentServiceListEnvironmentDeploymentsResponse, + zConsoleDeploymentServicePrecheckWorkflowDeploymentPath, + zConsoleDeploymentServicePrecheckWorkflowDeploymentResponse, + zConsoleDeploymentServiceUndeployWorkflowPath, + zConsoleDeploymentServiceUndeployWorkflowResponse, +} from './zod.gen' + +export const listAppEnvironments = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleDeploymentService_ListAppEnvironments', + path: '/enterprise/app-deploy/apps/{app_id}/environments', + tags: ['ConsoleDeploymentService'], + }) + .input(z.object({ params: zConsoleDeploymentServiceListAppEnvironmentsPath })) + .output(zConsoleDeploymentServiceListAppEnvironmentsResponse) + +export const listEnvironmentDeployments = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleDeploymentService_ListEnvironmentDeployments', + path: '/enterprise/app-deploy/apps/{app_id}/workflows/environment-deployments', + tags: ['ConsoleDeploymentService'], + }) + .input(z.object({ params: zConsoleDeploymentServiceListEnvironmentDeploymentsPath })) + .output(zConsoleDeploymentServiceListEnvironmentDeploymentsResponse) + +export const getEnvironmentDeployment = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleDeploymentService_GetEnvironmentDeployment', + path: '/enterprise/app-deploy/apps/{app_id}/workflows/environment-deployments/{environment_id}', + tags: ['ConsoleDeploymentService'], + }) + .input(z.object({ params: zConsoleDeploymentServiceGetEnvironmentDeploymentPath })) + .output(zConsoleDeploymentServiceGetEnvironmentDeploymentResponse) + +export const getWorkflowDeploymentOptions = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleDeploymentService_GetWorkflowDeploymentOptions', + path: '/enterprise/app-deploy/apps/{app_id}/workflows/{workflow_id}/environments/{environment_id}/deployment-options', + tags: ['ConsoleDeploymentService'], + }) + .input(z.object({ params: zConsoleDeploymentServiceGetWorkflowDeploymentOptionsPath })) + .output(zConsoleDeploymentServiceGetWorkflowDeploymentOptionsResponse) + +export const deployWorkflow = oc + .route({ + inputStructure: 'detailed', + method: 'POST', + operationId: 'ConsoleDeploymentService_DeployWorkflow', + path: '/enterprise/app-deploy/apps/{app_id}/workflows/{workflow_id}/environments/{environment_id}/deployment:deploy', + tags: ['ConsoleDeploymentService'], + }) + .input( + z.object({ + body: zConsoleDeploymentServiceDeployWorkflowBody, + params: zConsoleDeploymentServiceDeployWorkflowPath, + }), + ) + .output(zConsoleDeploymentServiceDeployWorkflowResponse) + +export const undeployWorkflow = oc + .route({ + inputStructure: 'detailed', + method: 'POST', + operationId: 'ConsoleDeploymentService_UndeployWorkflow', + path: '/enterprise/app-deploy/apps/{app_id}/workflows/{workflow_id}/environments/{environment_id}/deployment:undeploy', + tags: ['ConsoleDeploymentService'], + }) + .input(z.object({ params: zConsoleDeploymentServiceUndeployWorkflowPath })) + .output(zConsoleDeploymentServiceUndeployWorkflowResponse) + +export const precheckWorkflowDeployment = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleDeploymentService_PrecheckWorkflowDeployment', + path: '/enterprise/app-deploy/apps/{app_id}/workflows/{workflow_id}:precheck', + tags: ['ConsoleDeploymentService'], + }) + .input(z.object({ params: zConsoleDeploymentServicePrecheckWorkflowDeploymentPath })) + .output(zConsoleDeploymentServicePrecheckWorkflowDeploymentResponse) + +export const deploymentService = { + listAppEnvironments, + listEnvironmentDeployments, + getEnvironmentDeployment, + getWorkflowDeploymentOptions, + deployWorkflow, + undeployWorkflow, + precheckWorkflowDeployment, +} + +export const getEnvironmentApi = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleAccessService_GetEnvironmentAPI', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceGetEnvironmentApiPath })) + .output(zConsoleAccessServiceGetEnvironmentApiResponse) + +export const updateEnvironmentApi = oc + .route({ + inputStructure: 'detailed', + method: 'PATCH', + operationId: 'ConsoleAccessService_UpdateEnvironmentAPI', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api', + tags: ['ConsoleAccessService'], + }) + .input( + z.object({ + body: zConsoleAccessServiceUpdateEnvironmentApiBody, + params: zConsoleAccessServiceUpdateEnvironmentApiPath, + }), + ) + .output(zConsoleAccessServiceUpdateEnvironmentApiResponse) + +export const listEnvironmentApiKeys = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleAccessService_ListEnvironmentApiKeys', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api-keys', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceListEnvironmentApiKeysPath })) + .output(zConsoleAccessServiceListEnvironmentApiKeysResponse) + +export const createEnvironmentApiKey = oc + .route({ + inputStructure: 'detailed', + method: 'POST', + operationId: 'ConsoleAccessService_CreateEnvironmentApiKey', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api-keys', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceCreateEnvironmentApiKeyPath })) + .output(zConsoleAccessServiceCreateEnvironmentApiKeyResponse) + +export const deleteEnvironmentApiKey = oc + .route({ + inputStructure: 'detailed', + method: 'DELETE', + operationId: 'ConsoleAccessService_DeleteEnvironmentApiKey', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api-keys/{api_key_id}', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceDeleteEnvironmentApiKeyPath })) + .output(zConsoleAccessServiceDeleteEnvironmentApiKeyResponse) + +export const getEnvironmentMcpServer = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleAccessService_GetEnvironmentMCPServer', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/server', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceGetEnvironmentMcpServerPath })) + .output(zConsoleAccessServiceGetEnvironmentMcpServerResponse) + +export const getEnvironmentSite = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleAccessService_GetEnvironmentSite', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/site', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceGetEnvironmentSitePath })) + .output(zConsoleAccessServiceGetEnvironmentSiteResponse) + +export const updateEnvironmentSite = oc + .route({ + inputStructure: 'detailed', + method: 'PATCH', + operationId: 'ConsoleAccessService_UpdateEnvironmentSite', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/site', + tags: ['ConsoleAccessService'], + }) + .input( + z.object({ + body: zConsoleAccessServiceUpdateEnvironmentSiteBody, + params: zConsoleAccessServiceUpdateEnvironmentSitePath, + }), + ) + .output(zConsoleAccessServiceUpdateEnvironmentSiteResponse) + +export const resetEnvironmentSiteAccessToken = oc + .route({ + inputStructure: 'detailed', + method: 'POST', + operationId: 'ConsoleAccessService_ResetEnvironmentSiteAccessToken', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/site/access-token-reset', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceResetEnvironmentSiteAccessTokenPath })) + .output(zConsoleAccessServiceResetEnvironmentSiteAccessTokenResponse) + +export const listEnvironmentTriggers = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleAccessService_ListEnvironmentTriggers', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/triggers', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceListEnvironmentTriggersPath })) + .output(zConsoleAccessServiceListEnvironmentTriggersResponse) + +export const updateEnvironmentWebAppAccessMode = oc + .route({ + inputStructure: 'detailed', + method: 'POST', + operationId: 'ConsoleAccessService_UpdateEnvironmentWebAppAccessMode', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/webapp/access-mode', + tags: ['ConsoleAccessService'], + }) + .input( + z.object({ + body: zConsoleAccessServiceUpdateEnvironmentWebAppAccessModeBody, + params: zConsoleAccessServiceUpdateEnvironmentWebAppAccessModePath, + }), + ) + .output(zConsoleAccessServiceUpdateEnvironmentWebAppAccessModeResponse) + +export const getEnvironmentWebAppSubjects = oc + .route({ + inputStructure: 'detailed', + method: 'GET', + operationId: 'ConsoleAccessService_GetEnvironmentWebAppSubjects', + path: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/webapp/subjects', + tags: ['ConsoleAccessService'], + }) + .input(z.object({ params: zConsoleAccessServiceGetEnvironmentWebAppSubjectsPath })) + .output(zConsoleAccessServiceGetEnvironmentWebAppSubjectsResponse) + +export const accessService = { + getEnvironmentApi, + updateEnvironmentApi, + listEnvironmentApiKeys, + createEnvironmentApiKey, + deleteEnvironmentApiKey, + getEnvironmentMcpServer, + getEnvironmentSite, + updateEnvironmentSite, + resetEnvironmentSiteAccessToken, + listEnvironmentTriggers, + updateEnvironmentWebAppAccessMode, + getEnvironmentWebAppSubjects, +} + +export const contract = { + deploymentService, + accessService, +} diff --git a/packages/contracts/generated/enterprise-app-deploy/types.gen.ts b/packages/contracts/generated/enterprise-app-deploy/types.gen.ts new file mode 100644 index 00000000000..c64c433d4b6 --- /dev/null +++ b/packages/contracts/generated/enterprise-app-deploy/types.gen.ts @@ -0,0 +1,1326 @@ +// This file is auto-generated by @hey-api/openapi-ts + +export type ClientOptions = { + baseUrl: `${string}://${string}` | (string & {}) +} + +export const EnvironmentStatus = { + ENVIRONMENT_STATUS_UNSPECIFIED: 'ENVIRONMENT_STATUS_UNSPECIFIED', + ENVIRONMENT_STATUS_PENDING: 'ENVIRONMENT_STATUS_PENDING', + ENVIRONMENT_STATUS_READY: 'ENVIRONMENT_STATUS_READY', + ENVIRONMENT_STATUS_FAILED: 'ENVIRONMENT_STATUS_FAILED', + ENVIRONMENT_STATUS_DELETING: 'ENVIRONMENT_STATUS_DELETING', +} as const + +export type EnvironmentStatus = (typeof EnvironmentStatus)[keyof typeof EnvironmentStatus] + +export const ApplicationInteractionStatus = { + APPLICATION_INTERACTION_STATUS_UNSPECIFIED: 'APPLICATION_INTERACTION_STATUS_UNSPECIFIED', + APPLICATION_INTERACTION_STATUS_RUNNING: 'APPLICATION_INTERACTION_STATUS_RUNNING', + APPLICATION_INTERACTION_STATUS_SUCCEEDED: 'APPLICATION_INTERACTION_STATUS_SUCCEEDED', + APPLICATION_INTERACTION_STATUS_FAILED: 'APPLICATION_INTERACTION_STATUS_FAILED', + APPLICATION_INTERACTION_STATUS_PARTIAL_SUCCEEDED: + 'APPLICATION_INTERACTION_STATUS_PARTIAL_SUCCEEDED', +} as const + +export type ApplicationInteractionStatus = + (typeof ApplicationInteractionStatus)[keyof typeof ApplicationInteractionStatus] + +export const EnvironmentMode = { + ENVIRONMENT_MODE_UNSPECIFIED: 'ENVIRONMENT_MODE_UNSPECIFIED', + ENVIRONMENT_MODE_SHARED: 'ENVIRONMENT_MODE_SHARED', + ENVIRONMENT_MODE_ISOLATED: 'ENVIRONMENT_MODE_ISOLATED', +} as const + +export type EnvironmentMode = (typeof EnvironmentMode)[keyof typeof EnvironmentMode] + +export const PluginCategory = { + PLUGIN_CATEGORY_UNSPECIFIED: 'PLUGIN_CATEGORY_UNSPECIFIED', + PLUGIN_CATEGORY_MODEL: 'PLUGIN_CATEGORY_MODEL', + PLUGIN_CATEGORY_TOOL: 'PLUGIN_CATEGORY_TOOL', +} as const + +export type PluginCategory = (typeof PluginCategory)[keyof typeof PluginCategory] + +export const DeploymentOperationType = { + DEPLOYMENT_OPERATION_TYPE_UNSPECIFIED: 'DEPLOYMENT_OPERATION_TYPE_UNSPECIFIED', + DEPLOYMENT_OPERATION_TYPE_DEPLOY: 'DEPLOYMENT_OPERATION_TYPE_DEPLOY', + DEPLOYMENT_OPERATION_TYPE_UNDEPLOY: 'DEPLOYMENT_OPERATION_TYPE_UNDEPLOY', +} as const + +export type DeploymentOperationType = + (typeof DeploymentOperationType)[keyof typeof DeploymentOperationType] + +export const DeploymentOperationOutcome = { + DEPLOYMENT_OPERATION_OUTCOME_UNSPECIFIED: 'DEPLOYMENT_OPERATION_OUTCOME_UNSPECIFIED', + DEPLOYMENT_OPERATION_OUTCOME_IN_PROGRESS: 'DEPLOYMENT_OPERATION_OUTCOME_IN_PROGRESS', + DEPLOYMENT_OPERATION_OUTCOME_SUCCEEDED: 'DEPLOYMENT_OPERATION_OUTCOME_SUCCEEDED', + DEPLOYMENT_OPERATION_OUTCOME_FAILED: 'DEPLOYMENT_OPERATION_OUTCOME_FAILED', +} as const + +export type DeploymentOperationOutcome = + (typeof DeploymentOperationOutcome)[keyof typeof DeploymentOperationOutcome] + +export const DeploymentOperationFailureCode = { + DEPLOYMENT_OPERATION_FAILURE_CODE_UNSPECIFIED: 'DEPLOYMENT_OPERATION_FAILURE_CODE_UNSPECIFIED', + DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_IN_PROGRESS: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_IN_PROGRESS', + DEPLOYMENT_OPERATION_FAILURE_CODE_NOTHING_TO_UNDEPLOY: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_NOTHING_TO_UNDEPLOY', + DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_STATE_CHANGED: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_STATE_CHANGED', + DEPLOYMENT_OPERATION_FAILURE_CODE_APPLICATION_UNAVAILABLE: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_APPLICATION_UNAVAILABLE', + DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_UNAVAILABLE: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_UNAVAILABLE', + DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_REMOVED: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_REMOVED', + DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_INVALID: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_INVALID', + DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_ACCESS_DENIED: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_ACCESS_DENIED', + DEPLOYMENT_OPERATION_FAILURE_CODE_VERSION_UNAVAILABLE: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_VERSION_UNAVAILABLE', + DEPLOYMENT_OPERATION_FAILURE_CODE_VERSION_NOT_DEPLOYABLE: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_VERSION_NOT_DEPLOYABLE', + DEPLOYMENT_OPERATION_FAILURE_CODE_REQUIRED_CONFIGURATION_MISSING: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_REQUIRED_CONFIGURATION_MISSING', + DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_CONFIGURATION_INVALID: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_CONFIGURATION_INVALID', + DEPLOYMENT_OPERATION_FAILURE_CODE_CREDENTIAL_UNAVAILABLE: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_CREDENTIAL_UNAVAILABLE', + DEPLOYMENT_OPERATION_FAILURE_CODE_RUNTIME_UNAVAILABLE: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_RUNTIME_UNAVAILABLE', + DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_TIMEOUT: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_TIMEOUT', + DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_INTERRUPTED: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_INTERRUPTED', + DEPLOYMENT_OPERATION_FAILURE_CODE_INTERNAL_ERROR: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_INTERNAL_ERROR', + DEPLOYMENT_OPERATION_FAILURE_CODE_ENVIRONMENT_CPU_POOL_EXHAUSTED: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_ENVIRONMENT_CPU_POOL_EXHAUSTED', + DEPLOYMENT_OPERATION_FAILURE_CODE_APP_RUNNER_ENV_CPU_LIMIT_EXCEEDED: + 'DEPLOYMENT_OPERATION_FAILURE_CODE_APP_RUNNER_ENV_CPU_LIMIT_EXCEEDED', +} as const + +export type DeploymentOperationFailureCode = + (typeof DeploymentOperationFailureCode)[keyof typeof DeploymentOperationFailureCode] + +export const DeploymentOperationStatus = { + DEPLOYMENT_OPERATION_STATUS_UNSPECIFIED: 'DEPLOYMENT_OPERATION_STATUS_UNSPECIFIED', + DEPLOYMENT_OPERATION_STATUS_IN_PROGRESS: 'DEPLOYMENT_OPERATION_STATUS_IN_PROGRESS', + DEPLOYMENT_OPERATION_STATUS_SUCCEEDED: 'DEPLOYMENT_OPERATION_STATUS_SUCCEEDED', + DEPLOYMENT_OPERATION_STATUS_FAILED: 'DEPLOYMENT_OPERATION_STATUS_FAILED', +} as const + +export type DeploymentOperationStatus = + (typeof DeploymentOperationStatus)[keyof typeof DeploymentOperationStatus] + +export const EnvironmentBackend = { + ENVIRONMENT_BACKEND_UNSPECIFIED: 'ENVIRONMENT_BACKEND_UNSPECIFIED', + ENVIRONMENT_BACKEND_KUBERNETES: 'ENVIRONMENT_BACKEND_KUBERNETES', + ENVIRONMENT_BACKEND_EXTERNAL: 'ENVIRONMENT_BACKEND_EXTERNAL', +} as const + +export type EnvironmentBackend = (typeof EnvironmentBackend)[keyof typeof EnvironmentBackend] + +export const EnvironmentManagedBy = { + ENVIRONMENT_MANAGED_BY_UNSPECIFIED: 'ENVIRONMENT_MANAGED_BY_UNSPECIFIED', + ENVIRONMENT_MANAGED_BY_SYSTEM: 'ENVIRONMENT_MANAGED_BY_SYSTEM', + ENVIRONMENT_MANAGED_BY_USER: 'ENVIRONMENT_MANAGED_BY_USER', +} as const + +export type EnvironmentManagedBy = (typeof EnvironmentManagedBy)[keyof typeof EnvironmentManagedBy] + +export const EnvironmentDeployedAppStatus = { + ENVIRONMENT_DEPLOYED_APP_STATUS_UNSPECIFIED: 'ENVIRONMENT_DEPLOYED_APP_STATUS_UNSPECIFIED', + ENVIRONMENT_DEPLOYED_APP_STATUS_DEPLOYED: 'ENVIRONMENT_DEPLOYED_APP_STATUS_DEPLOYED', + ENVIRONMENT_DEPLOYED_APP_STATUS_DEPLOYING: 'ENVIRONMENT_DEPLOYED_APP_STATUS_DEPLOYING', + ENVIRONMENT_DEPLOYED_APP_STATUS_FAILED: 'ENVIRONMENT_DEPLOYED_APP_STATUS_FAILED', + ENVIRONMENT_DEPLOYED_APP_STATUS_UNDEPLOYED: 'ENVIRONMENT_DEPLOYED_APP_STATUS_UNDEPLOYED', +} as const + +export type EnvironmentDeployedAppStatus = + (typeof EnvironmentDeployedAppStatus)[keyof typeof EnvironmentDeployedAppStatus] + +export const DeploymentStatus = { + DEPLOYMENT_STATUS_UNSPECIFIED: 'DEPLOYMENT_STATUS_UNSPECIFIED', + DEPLOYMENT_STATUS_UNDEPLOYED: 'DEPLOYMENT_STATUS_UNDEPLOYED', + DEPLOYMENT_STATUS_DEPLOYING: 'DEPLOYMENT_STATUS_DEPLOYING', + DEPLOYMENT_STATUS_RUNNING: 'DEPLOYMENT_STATUS_RUNNING', + DEPLOYMENT_STATUS_UNDEPLOYING: 'DEPLOYMENT_STATUS_UNDEPLOYING', + DEPLOYMENT_STATUS_INVALID: 'DEPLOYMENT_STATUS_INVALID', + DEPLOYMENT_STATUS_FAILED: 'DEPLOYMENT_STATUS_FAILED', +} as const + +export type DeploymentStatus = (typeof DeploymentStatus)[keyof typeof DeploymentStatus] + +export const EnvVarValueSource = { + ENV_VAR_VALUE_SOURCE_UNSPECIFIED: 'ENV_VAR_VALUE_SOURCE_UNSPECIFIED', + ENV_VAR_VALUE_SOURCE_CONFIGURED: 'ENV_VAR_VALUE_SOURCE_CONFIGURED', + ENV_VAR_VALUE_SOURCE_LAST_DEPLOYED: 'ENV_VAR_VALUE_SOURCE_LAST_DEPLOYED', + ENV_VAR_VALUE_SOURCE_CUSTOM: 'ENV_VAR_VALUE_SOURCE_CUSTOM', +} as const + +export type EnvVarValueSource = (typeof EnvVarValueSource)[keyof typeof EnvVarValueSource] + +export const EnvVarValueType = { + ENV_VAR_VALUE_TYPE_UNSPECIFIED: 'ENV_VAR_VALUE_TYPE_UNSPECIFIED', + ENV_VAR_VALUE_TYPE_STRING: 'ENV_VAR_VALUE_TYPE_STRING', + ENV_VAR_VALUE_TYPE_NUMBER: 'ENV_VAR_VALUE_TYPE_NUMBER', + ENV_VAR_VALUE_TYPE_SECRET: 'ENV_VAR_VALUE_TYPE_SECRET', +} as const + +export type EnvVarValueType = (typeof EnvVarValueType)[keyof typeof EnvVarValueType] + +export const OperatorType = { + OPERATOR_TYPE_UNSPECIFIED: 'OPERATOR_TYPE_UNSPECIFIED', + OPERATOR_TYPE_END_USER: 'OPERATOR_TYPE_END_USER', + OPERATOR_TYPE_ACCOUNT: 'OPERATOR_TYPE_ACCOUNT', + OPERATOR_TYPE_SERVICE_ACCOUNT: 'OPERATOR_TYPE_SERVICE_ACCOUNT', + OPERATOR_TYPE_SYSTEM: 'OPERATOR_TYPE_SYSTEM', + OPERATOR_TYPE_DASHBOARD_ACCOUNT: 'OPERATOR_TYPE_DASHBOARD_ACCOUNT', +} as const + +export type OperatorType = (typeof OperatorType)[keyof typeof OperatorType] + +export const RouteTargetKind = { + ROUTE_TARGET_KIND_UNSPECIFIED: 'ROUTE_TARGET_KIND_UNSPECIFIED', + ROUTE_TARGET_KIND_K8S_SERVICE: 'ROUTE_TARGET_KIND_K8S_SERVICE', + ROUTE_TARGET_KIND_DIRECT_UPSTREAM: 'ROUTE_TARGET_KIND_DIRECT_UPSTREAM', +} as const + +export type RouteTargetKind = (typeof RouteTargetKind)[keyof typeof RouteTargetKind] + +export type AppEnvironment = { + id: string + display_name: string + description: string + status: EnvironmentStatus + in_use: boolean +} + +export type ApplicationInteraction = { + id: string + timestamp: string + workflowRunId: string + status: ApplicationInteractionStatus + durationSeconds: number + totalTokens: string + workspace: NamedRef + environment: NamedRef + app: NamedRef + operator?: Operator + invokeFrom: string + traceId: string + difyTraceId: string + deploymentVersionId: string + error?: string + body?: string + attributesJson?: string + resourceAttributesJson?: string +} + +export type BatchGetSourceVersionDeploymentsRequest = { + appId?: string + sourceVersionIds?: Array +} + +export type BatchGetSourceVersionDeploymentsResponse = { + items?: Array +} + +export type CheckAppDeletableRequest = { + appId?: string +} + +export type CheckAppDeletableResponse = { + deletable?: boolean + reason?: string +} + +export type CheckSourceVersionDeletableRequest = { + appId?: string + sourceVersionId?: string +} + +export type CheckSourceVersionDeletableResponse = { + deletable?: boolean + reason?: string +} + +export type CreateDockerComposeBundleRequest = { + environmentId?: string + dockerNetwork: string +} + +export type CreateDockerComposeBundleResponse = { + bundle: DockerComposeBundle +} + +export type CreateEnvironmentRequest = { + displayName: string + description?: string + mode: EnvironmentMode + cpuPool: number + namespace?: string + maxMemoryMib?: string +} + +export type CreateEnvironmentResponse = { + environment: Environment +} + +export type CredentialCandidate = { + credential_id: string + provider_id: string + category: PluginCategory + display_name: string + from_enterprise: boolean +} + +export type CredentialSelectionInput = { + provider_id: string + category: PluginCategory + credential_id: string +} + +export type CredentialSlot = { + provider_id: string + category: PluginCategory + candidates: Array + last_deployed_credential_id?: string + icon?: string + icon_dark?: string +} + +export type DashboardApp = { + id: string + workspaceId: string + displayName: string +} + +export type DeleteEnvironmentApiKeyResponse = { + [key: string]: unknown +} + +export type DeleteEnvironmentResponse = { + [key: string]: unknown +} + +export type DeployWorkflowResponse = { + operation: DeploymentOperationReceipt +} + +export type DeploymentEnvironment = { + id: string + display_name: string + status: EnvironmentStatus + description: string +} + +export type DeploymentOperation = { + id: string + type: DeploymentOperationType + outcome: DeploymentOperationOutcome + workspace: NamedRef + operator: Operator + app: NamedRef + environment: NamedRef + version: WorkflowVersion + failureCode?: DeploymentOperationFailureCode + failureMessage?: string + requestedAt: string + finalizedAt?: string + durationMilliseconds?: string +} + +export type DeploymentOperationReceipt = { + id: string + type: DeploymentOperationType + status: DeploymentOperationStatus +} + +export type DockerComposeBundle = { + filename: string + archive: string + composeYaml: string + configYaml: string +} + +export type Environment = { + id: string + displayName: string + description: string + mode: EnvironmentMode + backend: EnvironmentBackend + status: EnvironmentStatus + statusMessage: string + lastError?: Error + namespace?: string + managedBy?: EnvironmentManagedBy + cpuPool: number + createdAt: string + updatedAt: string + memory?: RunnerMemory + usage?: EnvironmentPoolUsage +} + +export type EnvironmentApi = { + enabled: boolean + base_url: string + api_key_count: number +} + +export type EnvironmentApiUpdate = { + enabled: boolean +} + +export type EnvironmentAccess = { + enable_site: boolean + enable_api: boolean +} + +export type EnvironmentApiKey = { + id: string + type: string + token: string + last_used_at?: number + created_at: number +} + +export type EnvironmentDeployedApp = { + deploymentId: string + workspace: NamedRef + app: NamedRef + status: EnvironmentDeployedAppStatus + currentVersion?: WorkflowVersion + deployedAt?: string + deployedBy?: Operator + latestAttempt?: EnvironmentDeployedAppAttempt + sizing?: RunnerSizing + occupiesPool?: boolean +} + +export type EnvironmentDeployedAppAttempt = { + operationId: string + type: DeploymentOperationType + outcome: DeploymentOperationOutcome + failureCode?: DeploymentOperationFailureCode + failureMessage?: string + requestedAt: string + finalizedAt?: string +} + +export type EnvironmentDeployedAppSummary = { + total: number + deployed: number + deploying: number + failed: number +} + +export type EnvironmentDeployment = { + environment: DeploymentEnvironment + deployment?: EnvironmentDeploymentState + access: EnvironmentAccess +} + +export type EnvironmentDeploymentOperation = { + id: string + type: DeploymentOperationType + status: DeploymentOperationStatus + target_version?: WorkflowVersion + operator: Operator + activity_at: number +} + +export type EnvironmentDeploymentState = { + status: DeploymentStatus + current_version?: WorkflowVersion + versions_behind?: number + deployed_at?: number + deployed_by?: Operator + latest_operation?: EnvironmentDeploymentOperation +} + +export type EnvironmentMcpServer = { + [key: string]: unknown +} + +export type EnvironmentPoolComposition = { + topApps?: Array + otherCpu?: number + otherAppCount?: number +} + +export type EnvironmentPoolShare = { + app: NamedRef + isolatedCpu: number +} + +export type EnvironmentPoolUsage = { + occupiedCpu: number + appCount: number +} + +export type EnvironmentSite = { + enabled: boolean + code: string + app_base_url: string + access_mode: string +} + +export type EnvironmentSiteUpdate = { + enabled: boolean +} + +export type EnvironmentTrigger = { + [key: string]: unknown +} + +export type EnvironmentVariableInput = { + key: string + value_source: EnvVarValueSource + value?: string +} + +export type EnvironmentVariableSlot = { + key: string + value_type: EnvVarValueType + description: string + has_configured_value: boolean + has_last_deployed_value: boolean + configured_value?: string + last_deployed_value?: string +} + +export type EnvironmentWebAppAccessModeUpdate = { + access_mode: string + subjects?: Array +} + +export type EnvironmentWebAppSubject = { + subject_id?: string + subject_type?: string + account_data?: EnvironmentWebAppSubjectAccountData + group_data?: EnvironmentWebAppSubjectGroupData +} + +export type EnvironmentWebAppSubjectAccountData = { + id?: string + name?: string + email?: string + avatar?: string +} + +export type EnvironmentWebAppSubjectGroupData = { + id?: string + name?: string + group_size?: number +} + +export type Error = { + reason?: + | 'APPDEPLOY_ERROR_REASON_UNSPECIFIED' + | 'APPDEPLOY_INTERNAL' + | 'APPDEPLOY_INVALID_ARGUMENT' + | 'APPDEPLOY_INVALID_TENANT_ID' + | 'APPDEPLOY_INVALID_APP_ID' + | 'APPDEPLOY_INVALID_SOURCE_VERSION_ID' + | 'APPDEPLOY_INVALID_ENVIRONMENT_ID' + | 'APPDEPLOY_INVALID_REVISION_ID' + | 'APPDEPLOY_INVALID_DEPLOYMENT_ID' + | 'APPDEPLOY_INVALID_CREDENTIAL_ID' + | 'APPDEPLOY_INVALID_SUBJECT_ID' + | 'APPDEPLOY_INVALID_API_KEY_ID' + | 'APPDEPLOY_INVALID_ACCOUNT' + | 'APPDEPLOY_ENVIRONMENT_MODE_REQUIRED' + | 'APPDEPLOY_ENVIRONMENT_BACKEND_REQUIRED' + | 'APPDEPLOY_APP_LOG_TIME_RANGE_REQUIRED' + | 'APPDEPLOY_APP_LOG_INVALID_TIME_RANGE' + | 'APPDEPLOY_APP_LOG_INVALID_CURSOR' + | 'APPDEPLOY_APP_LOG_CURSOR_FILTER_MISMATCH' + | 'APPDEPLOY_APP_LOG_ID_INVALID' + | 'APPDEPLOY_UNSUPPORTED_NODE_TYPE' + | 'APPDEPLOY_UNSUPPORTED_TOOL_PROVIDER_TYPE' + | 'APPDEPLOY_TOOL_PROVIDER_TYPE_INVALID' + | 'APPDEPLOY_TOOL_PROVIDER_NOT_INSTALLED' + | 'APPDEPLOY_MODEL_PROVIDER_NOT_INSTALLED' + | 'APPDEPLOY_MODEL_TYPE_UNSUPPORTED' + | 'APPDEPLOY_CREDENTIAL_MISMATCH' + | 'APPDEPLOY_APP_RUNNER_JOIN_TOKEN_REQUIRED' + | 'APPDEPLOY_INVALID_WORKFLOW_ID' + | 'APPDEPLOY_INVALID_DEPLOYMENT_VERSION_ID' + | 'APPDEPLOY_DEVELOPER_API_URL_NOT_CONFIGURED' + | 'APPDEPLOY_INVALID_DEPLOYMENT_OPERATION_ID' + | 'APPDEPLOY_APP_LOG_EXPORT_RANGE_TOO_WIDE' + | 'APPDEPLOY_APP_LOG_EXPORT_TOO_MANY_ROWS' + | 'APPDEPLOY_APP_LOG_EXPORT_TOO_LARGE' + | 'APPDEPLOY_UNAUTHORIZED' + | 'APPDEPLOY_FORBIDDEN' + | 'APPDEPLOY_APP_RUNNER_AUTH_REQUIRED' + | 'APPDEPLOY_APP_RUNNER_INVALID_JOIN_TOKEN' + | 'APPDEPLOY_APP_RUNNER_INVALID_CONTROL_TOKEN' + | 'APPDEPLOY_ENVIRONMENT_NOT_FOUND' + | 'APPDEPLOY_DEPLOYMENT_NOT_FOUND' + | 'APPDEPLOY_REVISION_NOT_FOUND' + | 'APPDEPLOY_ROLLBACK_TARGET_NOT_FOUND' + | 'APPDEPLOY_APP_NOT_FOUND' + | 'APPDEPLOY_CREDENTIAL_NOT_FOUND' + | 'APPDEPLOY_ENVIRONMENT_WEB_APP_SETTING_NOT_FOUND' + | 'APPDEPLOY_ACCESS_SUBJECT_NOT_FOUND' + | 'APPDEPLOY_API_KEY_NOT_FOUND' + | 'APPDEPLOY_SOURCE_VERSION_NOT_FOUND' + | 'APPDEPLOY_APP_LOG_NOT_FOUND' + | 'APPDEPLOY_WORKSPACE_NOT_FOUND' + | 'APPDEPLOY_APP_RUNNER_NOT_FOUND' + | 'APPDEPLOY_WORKFLOW_NOT_FOUND' + | 'APPDEPLOY_DEPLOYMENT_OPERATION_NOT_FOUND' + | 'APPDEPLOY_RUN_FILE_NOT_FOUND' + | 'APPDEPLOY_APPLICATION_UNAVAILABLE' + | 'APPDEPLOY_TARGET_ENVIRONMENT_REMOVED' + | 'APPDEPLOY_VERSION_UNAVAILABLE' + | 'APPDEPLOY_CONFLICT' + | 'APPDEPLOY_DEPLOYMENT_IN_PROGRESS' + | 'APPDEPLOY_ALREADY_UNDEPLOYED' + | 'APPDEPLOY_REVISION_ALREADY_TERMINAL' + | 'APPDEPLOY_NO_IN_FLIGHT_DEPLOYMENT' + | 'APPDEPLOY_ENVIRONMENT_CONFLICT' + | 'APPDEPLOY_ENVIRONMENT_BUSY' + | 'APPDEPLOY_ENV_VAR_KEY_CONFLICT' + | 'APPDEPLOY_APP_HAS_ACTIVE_DEPLOYMENTS' + | 'APPDEPLOY_SOURCE_VERSION_IN_USE' + | 'APPDEPLOY_NOTHING_TO_UNDEPLOY' + | 'APPDEPLOY_DEPLOYMENT_STATE_CHANGED' + | 'APPDEPLOY_ENVIRONMENT_HAS_ACTIVE_DEPLOYMENTS' + | 'APPDEPLOY_ENVIRONMENT_NOT_READY' + | 'APPDEPLOY_ENVIRONMENT_CAPACITY_EXCEEDED' + | 'APPDEPLOY_APP_RUNNER_ENV_CPU_LIMIT_EXCEEDED' + | 'APPDEPLOY_SOURCE_VERSION_NOT_DEPLOYABLE' + | 'APPDEPLOY_CREDENTIAL_GONE' + | 'APPDEPLOY_SLOT_UNSATISFIED' + | 'APPDEPLOY_WORKSPACE_INACTIVE' + | 'APPDEPLOY_ENVIRONMENT_BACKEND_UNAVAILABLE' + | 'APPDEPLOY_APP_RUNNER_REGISTRATION_NOT_ALLOWED' + | 'APPDEPLOY_WORKFLOW_NOT_DEPLOYABLE' + | 'APPDEPLOY_API_KEY_LIMIT_EXCEEDED' + | 'APPDEPLOY_ENVIRONMENT_NOT_FAILED' + | 'APPDEPLOY_TARGET_ENVIRONMENT_UNAVAILABLE' + | 'APPDEPLOY_TARGET_ENVIRONMENT_INVALID' + | 'APPDEPLOY_TARGET_ENVIRONMENT_ACCESS_DENIED' + | 'APPDEPLOY_VERSION_NOT_DEPLOYABLE' + | 'APPDEPLOY_REQUIRED_CONFIGURATION_MISSING' + | 'APPDEPLOY_DEPLOYMENT_CONFIGURATION_INVALID' + | 'APPDEPLOY_CREDENTIAL_UNAVAILABLE' + | 'APPDEPLOY_RUNTIME_UNAVAILABLE' + | 'APPDEPLOY_DEPLOYMENT_TIMEOUT' + | 'APPDEPLOY_DEPLOYMENT_INTERRUPTED' + | 'APPDEPLOY_ENVIRONMENT_CPU_POOL_EXHAUSTED' + | 'APPDEPLOY_RESOURCE_NOT_APPLICABLE_FOR_MODE' + | 'APPDEPLOY_ENVIRONMENT_CPU_POOL_BELOW_ALLOCATED' + | 'APPDEPLOY_APP_RUNNER_CONTROL_NOT_CONFIGURED' + | 'APPDEPLOY_RUNTIME_ASSIGNMENT_FAILED' + | 'APPDEPLOY_REVISION_TIMEOUT' + | 'APPDEPLOY_INTERNAL_ERROR' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_AUTH_REJECTED' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_NAMESPACE_MISSING' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_INSUFFICIENT_RBAC' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_DECRYPT_FAILED' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_INVALID_K8S_CONFIG' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_RESOURCE_COLLISION' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_APPLY_FAILED' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_SIGNING_KEY_MISSING' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_CONTROL_ENDPOINT_INVALID' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_K8S_UNREACHABLE' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_APP_RUNNER_NOT_READY' + | 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_SSAR_UNAVAILABLE' + message?: string + phase?: string + occurredAt?: string + detailCode?: string +} + +export type GetApplicationInteractionResponse = { + interaction: ApplicationInteraction +} + +export type GetDeploymentOperationResponse = { + operation: DeploymentOperation +} + +export type GetEnvironmentCapabilitiesResponse = { + backend: EnvironmentBackend + supportedModes: Array< + 'ENVIRONMENT_MODE_UNSPECIFIED' | 'ENVIRONMENT_MODE_SHARED' | 'ENVIRONMENT_MODE_ISOLATED' + > + namespaceRequired: boolean +} + +export type GetEnvironmentDeploymentResponse = { + environment_deployment: EnvironmentDeployment +} + +export type GetEnvironmentResponse = { + environment: Environment + pool?: EnvironmentPoolComposition +} + +export type GetEnvironmentWebAppSubjectsResponse = { + subjects: Array +} + +export type GetWebAppAccessModeResponse = { + accessMode?: string +} + +export type GetWebAppPermissionResponse = { + result?: boolean +} + +export type GetWorkflowDeploymentOptionsResponse = { + environment_variable_slots: Array + credential_slots: Array +} + +export type ListAppEnvironmentsResponse = { + data: Array +} + +export type ListApplicationInteractionsResponse = { + data: Array + pagination: Pagination +} + +export type ListAppsResponse = { + data: Array + pagination: Pagination +} + +export type ListDeploymentOperationsResponse = { + data: Array + pagination: Pagination +} + +export type ListEnvironmentApiKeysResponse = { + data: Array +} + +export type ListEnvironmentDeployedAppsResponse = { + data: Array + summary: EnvironmentDeployedAppSummary + pagination: Pagination +} + +export type ListEnvironmentDeploymentsResponse = { + environment_deployments: Array +} + +export type ListEnvironmentTriggersResponse = { + data: Array +} + +export type ListEnvironmentsResponse = { + data: Array + pagination: Pagination +} + +export type NamedRef = { + id: string + displayName: string +} + +export type Operator = { + type: OperatorType + id: string + display_name: string +} + +export type PrecheckWorkflowDeploymentResponse = { + unsupported_nodes: Array +} + +export type PrepareAppDeletionRequest = { + tenantId?: string + appId?: string +} + +export type ResolveApiTokenRouteRequest = { + token?: string +} + +export type ResolveApiTokenRouteResponse = { + environmentId?: string + namespace?: string + serviceName?: string + servicePort?: number + environmentStatus?: EnvironmentStatus + appId?: string + tenantId?: string + deploymentId?: string + servingRevisionId?: string + deploymentStatus?: DeploymentStatus + revoked?: boolean + unavailableReason?: string + targetKind?: RouteTargetKind + directUpstream?: string + deploymentGeneration?: string +} + +export type ResolveWebAppRouteRequest = { + appCode?: string + passport?: string +} + +export type ResolveWebAppRouteResponse = { + decision?: string + environmentId?: string + namespace?: string + serviceName?: string + servicePort?: number + environmentStatus?: EnvironmentStatus + appId?: string + tenantId?: string + deploymentId?: string + servingRevisionId?: string + deploymentStatus?: DeploymentStatus + unavailableReason?: string + targetKind?: RouteTargetKind + directUpstream?: string + deploymentGeneration?: string + endUserId?: string + authType?: string +} + +export type RetryEnvironmentBootstrapRequest = { + environmentId?: string +} + +export type RetryEnvironmentBootstrapResponse = { + environment: Environment +} + +export type RunnerMemory = { + maxMib?: string + readonly effectiveMaxMib?: string + readonly minMib?: string +} + +export type RunnerSizing = { + isolatedCpu: number + memory: RunnerMemory +} + +export type SimpleAccount = { + id: string + name?: string + email?: string +} + +export type SourceVersionDeployment = { + sourceVersionId?: string + environments?: Array +} + +export type TestConnectionRequest = { + environmentId?: string +} + +export type TestConnectionResponse = { + reachable?: boolean + message?: string +} + +export type UndeployEnvironmentAppResponse = { + operationId: string +} + +export type UndeployWorkflowResponse = { + operation: DeploymentOperationReceipt +} + +export type UnsupportedNode = { + id: string + type: string + title: string + provider?: UnsupportedNodeProvider +} + +export type UnsupportedNodeProvider = { + plugin_id: string + provider_id: string + provider_type: string + provider_name: string +} + +export type UpdateEnvironmentDeployedAppResourcesRequest = { + environmentId: string + deploymentId: string + isolatedCpu: number + maxMemoryMib?: string +} + +export type UpdateEnvironmentDeployedAppResourcesResponse = { + deploymentId: string + isolatedCpu: number + allocatedCpuCount: number + poolCpuCount: number + memory: RunnerMemory +} + +export type UpdateEnvironmentRequest = { + environmentId?: string + displayName?: string + description?: string + cpuPool?: number + maxMemoryMib?: string +} + +export type UpdateEnvironmentResponse = { + environment: Environment +} + +export type WorkflowDeploymentEnvironment = { + id?: string + name?: string +} + +export type WorkflowDeploymentInput = { + environment_variables?: Array + credentials?: Array +} + +export type WorkflowVersion = { + version?: string + marked_name?: string + id: string + marked_comment?: string + version_number?: number + created_at?: number + created_by?: SimpleAccount + dsl_hash?: string + deleted?: boolean +} + +export type Pagination = { + totalCount?: number + perPage?: number + currentPage?: number + totalPages?: number +} + +export type CreateEnvironmentResponseWritable = { + environment: EnvironmentWritable +} + +export type DeleteEnvironmentApiKeyResponseWritable = { + [key: string]: unknown +} + +export type DeleteEnvironmentResponseWritable = { + [key: string]: unknown +} + +export type EnvironmentWritable = { + id: string + displayName: string + description: string + mode: EnvironmentMode + backend: EnvironmentBackend + status: EnvironmentStatus + statusMessage: string + lastError?: Error + namespace?: string + managedBy?: EnvironmentManagedBy + cpuPool: number + createdAt: string + updatedAt: string + memory?: RunnerMemoryWritable + usage?: EnvironmentPoolUsage +} + +export type EnvironmentDeployedAppWritable = { + deploymentId: string + workspace: NamedRef + app: NamedRef + status: EnvironmentDeployedAppStatus + currentVersion?: WorkflowVersion + deployedAt?: string + deployedBy?: Operator + latestAttempt?: EnvironmentDeployedAppAttempt + sizing?: RunnerSizingWritable + occupiesPool?: boolean +} + +export type EnvironmentMcpServerWritable = { + [key: string]: unknown +} + +export type EnvironmentTriggerWritable = { + [key: string]: unknown +} + +export type GetEnvironmentResponseWritable = { + environment: EnvironmentWritable + pool?: EnvironmentPoolComposition +} + +export type ListEnvironmentDeployedAppsResponseWritable = { + data: Array + summary: EnvironmentDeployedAppSummary + pagination: Pagination +} + +export type ListEnvironmentsResponseWritable = { + data: Array + pagination: Pagination +} + +export type RetryEnvironmentBootstrapResponseWritable = { + environment: EnvironmentWritable +} + +export type RunnerMemoryWritable = { + maxMib?: string +} + +export type RunnerSizingWritable = { + isolatedCpu: number + memory: RunnerMemoryWritable +} + +export type UpdateEnvironmentDeployedAppResourcesResponseWritable = { + deploymentId: string + isolatedCpu: number + allocatedCpuCount: number + poolCpuCount: number + memory: RunnerMemoryWritable +} + +export type UpdateEnvironmentResponseWritable = { + environment: EnvironmentWritable +} + +export type ConsoleDeploymentServiceListAppEnvironmentsData = { + body?: never + path: { + app_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments' +} + +export type ConsoleDeploymentServiceListAppEnvironmentsResponses = { + 200: ListAppEnvironmentsResponse +} + +export type ConsoleDeploymentServiceListAppEnvironmentsResponse = + ConsoleDeploymentServiceListAppEnvironmentsResponses[keyof ConsoleDeploymentServiceListAppEnvironmentsResponses] + +export type ConsoleAccessServiceGetEnvironmentApiData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api' +} + +export type ConsoleAccessServiceGetEnvironmentApiResponses = { + 200: EnvironmentApi +} + +export type ConsoleAccessServiceGetEnvironmentApiResponse = + ConsoleAccessServiceGetEnvironmentApiResponses[keyof ConsoleAccessServiceGetEnvironmentApiResponses] + +export type ConsoleAccessServiceUpdateEnvironmentApiData = { + body: EnvironmentApiUpdate + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api' +} + +export type ConsoleAccessServiceUpdateEnvironmentApiResponses = { + 200: EnvironmentApi +} + +export type ConsoleAccessServiceUpdateEnvironmentApiResponse = + ConsoleAccessServiceUpdateEnvironmentApiResponses[keyof ConsoleAccessServiceUpdateEnvironmentApiResponses] + +export type ConsoleAccessServiceListEnvironmentApiKeysData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api-keys' +} + +export type ConsoleAccessServiceListEnvironmentApiKeysResponses = { + 200: ListEnvironmentApiKeysResponse +} + +export type ConsoleAccessServiceListEnvironmentApiKeysResponse = + ConsoleAccessServiceListEnvironmentApiKeysResponses[keyof ConsoleAccessServiceListEnvironmentApiKeysResponses] + +export type ConsoleAccessServiceCreateEnvironmentApiKeyData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api-keys' +} + +export type ConsoleAccessServiceCreateEnvironmentApiKeyResponses = { + 200: EnvironmentApiKey +} + +export type ConsoleAccessServiceCreateEnvironmentApiKeyResponse = + ConsoleAccessServiceCreateEnvironmentApiKeyResponses[keyof ConsoleAccessServiceCreateEnvironmentApiKeyResponses] + +export type ConsoleAccessServiceDeleteEnvironmentApiKeyData = { + body?: never + path: { + app_id: string + environment_id: string + api_key_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/api-keys/{api_key_id}' +} + +export type ConsoleAccessServiceDeleteEnvironmentApiKeyResponses = { + 200: DeleteEnvironmentApiKeyResponse +} + +export type ConsoleAccessServiceDeleteEnvironmentApiKeyResponse = + ConsoleAccessServiceDeleteEnvironmentApiKeyResponses[keyof ConsoleAccessServiceDeleteEnvironmentApiKeyResponses] + +export type ConsoleAccessServiceGetEnvironmentMcpServerData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/server' +} + +export type ConsoleAccessServiceGetEnvironmentMcpServerResponses = { + 200: EnvironmentMcpServer +} + +export type ConsoleAccessServiceGetEnvironmentMcpServerResponse = + ConsoleAccessServiceGetEnvironmentMcpServerResponses[keyof ConsoleAccessServiceGetEnvironmentMcpServerResponses] + +export type ConsoleAccessServiceGetEnvironmentSiteData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/site' +} + +export type ConsoleAccessServiceGetEnvironmentSiteResponses = { + 200: EnvironmentSite +} + +export type ConsoleAccessServiceGetEnvironmentSiteResponse = + ConsoleAccessServiceGetEnvironmentSiteResponses[keyof ConsoleAccessServiceGetEnvironmentSiteResponses] + +export type ConsoleAccessServiceUpdateEnvironmentSiteData = { + body: EnvironmentSiteUpdate + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/site' +} + +export type ConsoleAccessServiceUpdateEnvironmentSiteResponses = { + 200: EnvironmentSite +} + +export type ConsoleAccessServiceUpdateEnvironmentSiteResponse = + ConsoleAccessServiceUpdateEnvironmentSiteResponses[keyof ConsoleAccessServiceUpdateEnvironmentSiteResponses] + +export type ConsoleAccessServiceResetEnvironmentSiteAccessTokenData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/site/access-token-reset' +} + +export type ConsoleAccessServiceResetEnvironmentSiteAccessTokenResponses = { + 200: EnvironmentSite +} + +export type ConsoleAccessServiceResetEnvironmentSiteAccessTokenResponse = + ConsoleAccessServiceResetEnvironmentSiteAccessTokenResponses[keyof ConsoleAccessServiceResetEnvironmentSiteAccessTokenResponses] + +export type ConsoleAccessServiceListEnvironmentTriggersData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/triggers' +} + +export type ConsoleAccessServiceListEnvironmentTriggersResponses = { + 200: ListEnvironmentTriggersResponse +} + +export type ConsoleAccessServiceListEnvironmentTriggersResponse = + ConsoleAccessServiceListEnvironmentTriggersResponses[keyof ConsoleAccessServiceListEnvironmentTriggersResponses] + +export type ConsoleAccessServiceUpdateEnvironmentWebAppAccessModeData = { + body: EnvironmentWebAppAccessModeUpdate + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/webapp/access-mode' +} + +export type ConsoleAccessServiceUpdateEnvironmentWebAppAccessModeResponses = { + 200: EnvironmentSite +} + +export type ConsoleAccessServiceUpdateEnvironmentWebAppAccessModeResponse = + ConsoleAccessServiceUpdateEnvironmentWebAppAccessModeResponses[keyof ConsoleAccessServiceUpdateEnvironmentWebAppAccessModeResponses] + +export type ConsoleAccessServiceGetEnvironmentWebAppSubjectsData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/environments/{environment_id}/webapp/subjects' +} + +export type ConsoleAccessServiceGetEnvironmentWebAppSubjectsResponses = { + 200: GetEnvironmentWebAppSubjectsResponse +} + +export type ConsoleAccessServiceGetEnvironmentWebAppSubjectsResponse = + ConsoleAccessServiceGetEnvironmentWebAppSubjectsResponses[keyof ConsoleAccessServiceGetEnvironmentWebAppSubjectsResponses] + +export type ConsoleDeploymentServiceListEnvironmentDeploymentsData = { + body?: never + path: { + app_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/workflows/environment-deployments' +} + +export type ConsoleDeploymentServiceListEnvironmentDeploymentsResponses = { + 200: ListEnvironmentDeploymentsResponse +} + +export type ConsoleDeploymentServiceListEnvironmentDeploymentsResponse = + ConsoleDeploymentServiceListEnvironmentDeploymentsResponses[keyof ConsoleDeploymentServiceListEnvironmentDeploymentsResponses] + +export type ConsoleDeploymentServiceGetEnvironmentDeploymentData = { + body?: never + path: { + app_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/workflows/environment-deployments/{environment_id}' +} + +export type ConsoleDeploymentServiceGetEnvironmentDeploymentResponses = { + 200: GetEnvironmentDeploymentResponse +} + +export type ConsoleDeploymentServiceGetEnvironmentDeploymentResponse = + ConsoleDeploymentServiceGetEnvironmentDeploymentResponses[keyof ConsoleDeploymentServiceGetEnvironmentDeploymentResponses] + +export type ConsoleDeploymentServiceGetWorkflowDeploymentOptionsData = { + body?: never + path: { + app_id: string + workflow_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/workflows/{workflow_id}/environments/{environment_id}/deployment-options' +} + +export type ConsoleDeploymentServiceGetWorkflowDeploymentOptionsResponses = { + 200: GetWorkflowDeploymentOptionsResponse +} + +export type ConsoleDeploymentServiceGetWorkflowDeploymentOptionsResponse = + ConsoleDeploymentServiceGetWorkflowDeploymentOptionsResponses[keyof ConsoleDeploymentServiceGetWorkflowDeploymentOptionsResponses] + +export type ConsoleDeploymentServiceDeployWorkflowData = { + body: WorkflowDeploymentInput + path: { + app_id: string + workflow_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/workflows/{workflow_id}/environments/{environment_id}/deployment:deploy' +} + +export type ConsoleDeploymentServiceDeployWorkflowResponses = { + 200: DeployWorkflowResponse +} + +export type ConsoleDeploymentServiceDeployWorkflowResponse = + ConsoleDeploymentServiceDeployWorkflowResponses[keyof ConsoleDeploymentServiceDeployWorkflowResponses] + +export type ConsoleDeploymentServiceUndeployWorkflowData = { + body?: never + path: { + app_id: string + workflow_id: string + environment_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/workflows/{workflow_id}/environments/{environment_id}/deployment:undeploy' +} + +export type ConsoleDeploymentServiceUndeployWorkflowResponses = { + 200: UndeployWorkflowResponse +} + +export type ConsoleDeploymentServiceUndeployWorkflowResponse = + ConsoleDeploymentServiceUndeployWorkflowResponses[keyof ConsoleDeploymentServiceUndeployWorkflowResponses] + +export type ConsoleDeploymentServicePrecheckWorkflowDeploymentData = { + body?: never + path: { + app_id: string + workflow_id: string + } + query?: never + url: '/enterprise/app-deploy/apps/{app_id}/workflows/{workflow_id}:precheck' +} + +export type ConsoleDeploymentServicePrecheckWorkflowDeploymentResponses = { + 200: PrecheckWorkflowDeploymentResponse +} + +export type ConsoleDeploymentServicePrecheckWorkflowDeploymentResponse = + ConsoleDeploymentServicePrecheckWorkflowDeploymentResponses[keyof ConsoleDeploymentServicePrecheckWorkflowDeploymentResponses] diff --git a/packages/contracts/generated/enterprise-app-deploy/zod.gen.ts b/packages/contracts/generated/enterprise-app-deploy/zod.gen.ts new file mode 100644 index 00000000000..c823bce2015 --- /dev/null +++ b/packages/contracts/generated/enterprise-app-deploy/zod.gen.ts @@ -0,0 +1,1224 @@ +// This file is auto-generated by @hey-api/openapi-ts + +import * as z from 'zod' + +export const zEnvironmentStatus = z.enum([ + 'ENVIRONMENT_STATUS_UNSPECIFIED', + 'ENVIRONMENT_STATUS_PENDING', + 'ENVIRONMENT_STATUS_READY', + 'ENVIRONMENT_STATUS_FAILED', + 'ENVIRONMENT_STATUS_DELETING', +]) + +export const zApplicationInteractionStatus = z.enum([ + 'APPLICATION_INTERACTION_STATUS_UNSPECIFIED', + 'APPLICATION_INTERACTION_STATUS_RUNNING', + 'APPLICATION_INTERACTION_STATUS_SUCCEEDED', + 'APPLICATION_INTERACTION_STATUS_FAILED', + 'APPLICATION_INTERACTION_STATUS_PARTIAL_SUCCEEDED', +]) + +export const zEnvironmentMode = z.enum([ + 'ENVIRONMENT_MODE_UNSPECIFIED', + 'ENVIRONMENT_MODE_SHARED', + 'ENVIRONMENT_MODE_ISOLATED', +]) + +export const zPluginCategory = z.enum([ + 'PLUGIN_CATEGORY_UNSPECIFIED', + 'PLUGIN_CATEGORY_MODEL', + 'PLUGIN_CATEGORY_TOOL', +]) + +export const zDeploymentOperationType = z.enum([ + 'DEPLOYMENT_OPERATION_TYPE_UNSPECIFIED', + 'DEPLOYMENT_OPERATION_TYPE_DEPLOY', + 'DEPLOYMENT_OPERATION_TYPE_UNDEPLOY', +]) + +export const zDeploymentOperationOutcome = z.enum([ + 'DEPLOYMENT_OPERATION_OUTCOME_UNSPECIFIED', + 'DEPLOYMENT_OPERATION_OUTCOME_IN_PROGRESS', + 'DEPLOYMENT_OPERATION_OUTCOME_SUCCEEDED', + 'DEPLOYMENT_OPERATION_OUTCOME_FAILED', +]) + +export const zDeploymentOperationFailureCode = z.enum([ + 'DEPLOYMENT_OPERATION_FAILURE_CODE_UNSPECIFIED', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_IN_PROGRESS', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_NOTHING_TO_UNDEPLOY', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_STATE_CHANGED', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_APPLICATION_UNAVAILABLE', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_UNAVAILABLE', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_REMOVED', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_INVALID', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_TARGET_ENVIRONMENT_ACCESS_DENIED', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_VERSION_UNAVAILABLE', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_VERSION_NOT_DEPLOYABLE', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_REQUIRED_CONFIGURATION_MISSING', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_CONFIGURATION_INVALID', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_CREDENTIAL_UNAVAILABLE', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_RUNTIME_UNAVAILABLE', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_TIMEOUT', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_DEPLOYMENT_INTERRUPTED', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_INTERNAL_ERROR', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_ENVIRONMENT_CPU_POOL_EXHAUSTED', + 'DEPLOYMENT_OPERATION_FAILURE_CODE_APP_RUNNER_ENV_CPU_LIMIT_EXCEEDED', +]) + +export const zDeploymentOperationStatus = z.enum([ + 'DEPLOYMENT_OPERATION_STATUS_UNSPECIFIED', + 'DEPLOYMENT_OPERATION_STATUS_IN_PROGRESS', + 'DEPLOYMENT_OPERATION_STATUS_SUCCEEDED', + 'DEPLOYMENT_OPERATION_STATUS_FAILED', +]) + +export const zEnvironmentBackend = z.enum([ + 'ENVIRONMENT_BACKEND_UNSPECIFIED', + 'ENVIRONMENT_BACKEND_KUBERNETES', + 'ENVIRONMENT_BACKEND_EXTERNAL', +]) + +export const zEnvironmentManagedBy = z.enum([ + 'ENVIRONMENT_MANAGED_BY_UNSPECIFIED', + 'ENVIRONMENT_MANAGED_BY_SYSTEM', + 'ENVIRONMENT_MANAGED_BY_USER', +]) + +export const zEnvironmentDeployedAppStatus = z.enum([ + 'ENVIRONMENT_DEPLOYED_APP_STATUS_UNSPECIFIED', + 'ENVIRONMENT_DEPLOYED_APP_STATUS_DEPLOYED', + 'ENVIRONMENT_DEPLOYED_APP_STATUS_DEPLOYING', + 'ENVIRONMENT_DEPLOYED_APP_STATUS_FAILED', + 'ENVIRONMENT_DEPLOYED_APP_STATUS_UNDEPLOYED', +]) + +export const zDeploymentStatus = z.enum([ + 'DEPLOYMENT_STATUS_UNSPECIFIED', + 'DEPLOYMENT_STATUS_UNDEPLOYED', + 'DEPLOYMENT_STATUS_DEPLOYING', + 'DEPLOYMENT_STATUS_RUNNING', + 'DEPLOYMENT_STATUS_UNDEPLOYING', + 'DEPLOYMENT_STATUS_INVALID', + 'DEPLOYMENT_STATUS_FAILED', +]) + +export const zEnvVarValueSource = z.enum([ + 'ENV_VAR_VALUE_SOURCE_UNSPECIFIED', + 'ENV_VAR_VALUE_SOURCE_CONFIGURED', + 'ENV_VAR_VALUE_SOURCE_LAST_DEPLOYED', + 'ENV_VAR_VALUE_SOURCE_CUSTOM', +]) + +export const zEnvVarValueType = z.enum([ + 'ENV_VAR_VALUE_TYPE_UNSPECIFIED', + 'ENV_VAR_VALUE_TYPE_STRING', + 'ENV_VAR_VALUE_TYPE_NUMBER', + 'ENV_VAR_VALUE_TYPE_SECRET', +]) + +export const zOperatorType = z.enum([ + 'OPERATOR_TYPE_UNSPECIFIED', + 'OPERATOR_TYPE_END_USER', + 'OPERATOR_TYPE_ACCOUNT', + 'OPERATOR_TYPE_SERVICE_ACCOUNT', + 'OPERATOR_TYPE_SYSTEM', + 'OPERATOR_TYPE_DASHBOARD_ACCOUNT', +]) + +export const zRouteTargetKind = z.enum([ + 'ROUTE_TARGET_KIND_UNSPECIFIED', + 'ROUTE_TARGET_KIND_K8S_SERVICE', + 'ROUTE_TARGET_KIND_DIRECT_UPSTREAM', +]) + +export const zAppEnvironment = z.object({ + id: z.string(), + display_name: z.string(), + description: z.string(), + status: zEnvironmentStatus, + in_use: z.boolean(), +}) + +export const zBatchGetSourceVersionDeploymentsRequest = z.object({ + appId: z.string().optional(), + sourceVersionIds: z.array(z.string()).optional(), +}) + +export const zCheckAppDeletableRequest = z.object({ + appId: z.string().optional(), +}) + +export const zCheckAppDeletableResponse = z.object({ + deletable: z.boolean().optional(), + reason: z.string().optional(), +}) + +export const zCheckSourceVersionDeletableRequest = z.object({ + appId: z.string().optional(), + sourceVersionId: z.string().optional(), +}) + +export const zCheckSourceVersionDeletableResponse = z.object({ + deletable: z.boolean().optional(), + reason: z.string().optional(), +}) + +export const zCreateDockerComposeBundleRequest = z.object({ + environmentId: z.string().optional(), + dockerNetwork: z.string(), +}) + +export const zCreateEnvironmentRequest = z.object({ + displayName: z.string(), + description: z.string().optional(), + mode: zEnvironmentMode, + cpuPool: z.number(), + namespace: z.string().optional(), + maxMemoryMib: z.string().optional(), +}) + +export const zCredentialCandidate = z.object({ + credential_id: z.string(), + provider_id: z.string(), + category: zPluginCategory, + display_name: z.string(), + from_enterprise: z.boolean(), +}) + +export const zCredentialSelectionInput = z.object({ + provider_id: z.string(), + category: zPluginCategory, + credential_id: z.string(), +}) + +export const zCredentialSlot = z.object({ + provider_id: z.string(), + category: zPluginCategory, + candidates: z.array(zCredentialCandidate), + last_deployed_credential_id: z.string().optional(), + icon: z.string().optional(), + icon_dark: z.string().optional(), +}) + +export const zDashboardApp = z.object({ + id: z.string(), + workspaceId: z.string(), + displayName: z.string(), +}) + +export const zDeleteEnvironmentApiKeyResponse = z.record(z.string(), z.unknown()) + +export const zDeleteEnvironmentResponse = z.record(z.string(), z.unknown()) + +export const zDeploymentEnvironment = z.object({ + id: z.string(), + display_name: z.string(), + status: zEnvironmentStatus, + description: z.string(), +}) + +export const zDeploymentOperationReceipt = z.object({ + id: z.string(), + type: zDeploymentOperationType, + status: zDeploymentOperationStatus, +}) + +export const zDeployWorkflowResponse = z.object({ + operation: zDeploymentOperationReceipt, +}) + +export const zDockerComposeBundle = z.object({ + filename: z.string(), + archive: z.string(), + composeYaml: z.string(), + configYaml: z.string(), +}) + +export const zCreateDockerComposeBundleResponse = z.object({ + bundle: zDockerComposeBundle, +}) + +export const zEnvironmentApi = z.object({ + enabled: z.boolean(), + base_url: z.string(), + api_key_count: z + .int() + .min(0, { error: 'Invalid value: Expected uint32 to be >= 0' }) + .max(4294967295, { error: 'Invalid value: Expected uint32 to be <= 4294967295' }), +}) + +export const zEnvironmentApiUpdate = z.object({ + enabled: z.boolean(), +}) + +export const zEnvironmentAccess = z.object({ + enable_site: z.boolean(), + enable_api: z.boolean(), +}) + +export const zEnvironmentApiKey = z.object({ + id: z.string(), + type: z.string(), + token: z.string(), + last_used_at: z.number().optional(), + created_at: z.number(), +}) + +export const zEnvironmentDeployedAppAttempt = z.object({ + operationId: z.string(), + type: zDeploymentOperationType, + outcome: zDeploymentOperationOutcome, + failureCode: zDeploymentOperationFailureCode.optional(), + failureMessage: z.string().optional(), + requestedAt: z.iso.datetime(), + finalizedAt: z.iso.datetime().optional(), +}) + +export const zEnvironmentDeployedAppSummary = z.object({ + total: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }), + deployed: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }), + deploying: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }), + failed: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }), +}) + +export const zEnvironmentMcpServer = z.record(z.string(), z.unknown()) + +export const zEnvironmentPoolUsage = z.object({ + occupiedCpu: z.number(), + appCount: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }), +}) + +export const zEnvironmentSite = z.object({ + enabled: z.boolean(), + code: z.string(), + app_base_url: z.string(), + access_mode: z.string(), +}) + +export const zEnvironmentSiteUpdate = z.object({ + enabled: z.boolean(), +}) + +export const zEnvironmentTrigger = z.record(z.string(), z.unknown()) + +export const zEnvironmentVariableInput = z.object({ + key: z.string(), + value_source: zEnvVarValueSource, + value: z.string().optional(), +}) + +export const zEnvironmentVariableSlot = z.object({ + key: z.string(), + value_type: zEnvVarValueType, + description: z.string(), + has_configured_value: z.boolean(), + has_last_deployed_value: z.boolean(), + configured_value: z.string().optional(), + last_deployed_value: z.string().optional(), +}) + +export const zEnvironmentWebAppSubjectAccountData = z.object({ + id: z.string().optional(), + name: z.string().optional(), + email: z.string().optional(), + avatar: z.string().optional(), +}) + +export const zEnvironmentWebAppSubjectGroupData = z.object({ + id: z.string().optional(), + name: z.string().optional(), + group_size: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), +}) + +export const zEnvironmentWebAppSubject = z.object({ + subject_id: z.string().optional(), + subject_type: z.string().optional(), + account_data: zEnvironmentWebAppSubjectAccountData.optional(), + group_data: zEnvironmentWebAppSubjectGroupData.optional(), +}) + +export const zEnvironmentWebAppAccessModeUpdate = z.object({ + access_mode: z.string(), + subjects: z.array(zEnvironmentWebAppSubject).optional(), +}) + +export const zError = z.object({ + reason: z + .enum([ + 'APPDEPLOY_ERROR_REASON_UNSPECIFIED', + 'APPDEPLOY_INTERNAL', + 'APPDEPLOY_INVALID_ARGUMENT', + 'APPDEPLOY_INVALID_TENANT_ID', + 'APPDEPLOY_INVALID_APP_ID', + 'APPDEPLOY_INVALID_SOURCE_VERSION_ID', + 'APPDEPLOY_INVALID_ENVIRONMENT_ID', + 'APPDEPLOY_INVALID_REVISION_ID', + 'APPDEPLOY_INVALID_DEPLOYMENT_ID', + 'APPDEPLOY_INVALID_CREDENTIAL_ID', + 'APPDEPLOY_INVALID_SUBJECT_ID', + 'APPDEPLOY_INVALID_API_KEY_ID', + 'APPDEPLOY_INVALID_ACCOUNT', + 'APPDEPLOY_ENVIRONMENT_MODE_REQUIRED', + 'APPDEPLOY_ENVIRONMENT_BACKEND_REQUIRED', + 'APPDEPLOY_APP_LOG_TIME_RANGE_REQUIRED', + 'APPDEPLOY_APP_LOG_INVALID_TIME_RANGE', + 'APPDEPLOY_APP_LOG_INVALID_CURSOR', + 'APPDEPLOY_APP_LOG_CURSOR_FILTER_MISMATCH', + 'APPDEPLOY_APP_LOG_ID_INVALID', + 'APPDEPLOY_UNSUPPORTED_NODE_TYPE', + 'APPDEPLOY_UNSUPPORTED_TOOL_PROVIDER_TYPE', + 'APPDEPLOY_TOOL_PROVIDER_TYPE_INVALID', + 'APPDEPLOY_TOOL_PROVIDER_NOT_INSTALLED', + 'APPDEPLOY_MODEL_PROVIDER_NOT_INSTALLED', + 'APPDEPLOY_MODEL_TYPE_UNSUPPORTED', + 'APPDEPLOY_CREDENTIAL_MISMATCH', + 'APPDEPLOY_APP_RUNNER_JOIN_TOKEN_REQUIRED', + 'APPDEPLOY_INVALID_WORKFLOW_ID', + 'APPDEPLOY_INVALID_DEPLOYMENT_VERSION_ID', + 'APPDEPLOY_DEVELOPER_API_URL_NOT_CONFIGURED', + 'APPDEPLOY_INVALID_DEPLOYMENT_OPERATION_ID', + 'APPDEPLOY_APP_LOG_EXPORT_RANGE_TOO_WIDE', + 'APPDEPLOY_APP_LOG_EXPORT_TOO_MANY_ROWS', + 'APPDEPLOY_APP_LOG_EXPORT_TOO_LARGE', + 'APPDEPLOY_UNAUTHORIZED', + 'APPDEPLOY_FORBIDDEN', + 'APPDEPLOY_APP_RUNNER_AUTH_REQUIRED', + 'APPDEPLOY_APP_RUNNER_INVALID_JOIN_TOKEN', + 'APPDEPLOY_APP_RUNNER_INVALID_CONTROL_TOKEN', + 'APPDEPLOY_ENVIRONMENT_NOT_FOUND', + 'APPDEPLOY_DEPLOYMENT_NOT_FOUND', + 'APPDEPLOY_REVISION_NOT_FOUND', + 'APPDEPLOY_ROLLBACK_TARGET_NOT_FOUND', + 'APPDEPLOY_APP_NOT_FOUND', + 'APPDEPLOY_CREDENTIAL_NOT_FOUND', + 'APPDEPLOY_ENVIRONMENT_WEB_APP_SETTING_NOT_FOUND', + 'APPDEPLOY_ACCESS_SUBJECT_NOT_FOUND', + 'APPDEPLOY_API_KEY_NOT_FOUND', + 'APPDEPLOY_SOURCE_VERSION_NOT_FOUND', + 'APPDEPLOY_APP_LOG_NOT_FOUND', + 'APPDEPLOY_WORKSPACE_NOT_FOUND', + 'APPDEPLOY_APP_RUNNER_NOT_FOUND', + 'APPDEPLOY_WORKFLOW_NOT_FOUND', + 'APPDEPLOY_DEPLOYMENT_OPERATION_NOT_FOUND', + 'APPDEPLOY_RUN_FILE_NOT_FOUND', + 'APPDEPLOY_APPLICATION_UNAVAILABLE', + 'APPDEPLOY_TARGET_ENVIRONMENT_REMOVED', + 'APPDEPLOY_VERSION_UNAVAILABLE', + 'APPDEPLOY_CONFLICT', + 'APPDEPLOY_DEPLOYMENT_IN_PROGRESS', + 'APPDEPLOY_ALREADY_UNDEPLOYED', + 'APPDEPLOY_REVISION_ALREADY_TERMINAL', + 'APPDEPLOY_NO_IN_FLIGHT_DEPLOYMENT', + 'APPDEPLOY_ENVIRONMENT_CONFLICT', + 'APPDEPLOY_ENVIRONMENT_BUSY', + 'APPDEPLOY_ENV_VAR_KEY_CONFLICT', + 'APPDEPLOY_APP_HAS_ACTIVE_DEPLOYMENTS', + 'APPDEPLOY_SOURCE_VERSION_IN_USE', + 'APPDEPLOY_NOTHING_TO_UNDEPLOY', + 'APPDEPLOY_DEPLOYMENT_STATE_CHANGED', + 'APPDEPLOY_ENVIRONMENT_HAS_ACTIVE_DEPLOYMENTS', + 'APPDEPLOY_ENVIRONMENT_NOT_READY', + 'APPDEPLOY_ENVIRONMENT_CAPACITY_EXCEEDED', + 'APPDEPLOY_APP_RUNNER_ENV_CPU_LIMIT_EXCEEDED', + 'APPDEPLOY_SOURCE_VERSION_NOT_DEPLOYABLE', + 'APPDEPLOY_CREDENTIAL_GONE', + 'APPDEPLOY_SLOT_UNSATISFIED', + 'APPDEPLOY_WORKSPACE_INACTIVE', + 'APPDEPLOY_ENVIRONMENT_BACKEND_UNAVAILABLE', + 'APPDEPLOY_APP_RUNNER_REGISTRATION_NOT_ALLOWED', + 'APPDEPLOY_WORKFLOW_NOT_DEPLOYABLE', + 'APPDEPLOY_API_KEY_LIMIT_EXCEEDED', + 'APPDEPLOY_ENVIRONMENT_NOT_FAILED', + 'APPDEPLOY_TARGET_ENVIRONMENT_UNAVAILABLE', + 'APPDEPLOY_TARGET_ENVIRONMENT_INVALID', + 'APPDEPLOY_TARGET_ENVIRONMENT_ACCESS_DENIED', + 'APPDEPLOY_VERSION_NOT_DEPLOYABLE', + 'APPDEPLOY_REQUIRED_CONFIGURATION_MISSING', + 'APPDEPLOY_DEPLOYMENT_CONFIGURATION_INVALID', + 'APPDEPLOY_CREDENTIAL_UNAVAILABLE', + 'APPDEPLOY_RUNTIME_UNAVAILABLE', + 'APPDEPLOY_DEPLOYMENT_TIMEOUT', + 'APPDEPLOY_DEPLOYMENT_INTERRUPTED', + 'APPDEPLOY_ENVIRONMENT_CPU_POOL_EXHAUSTED', + 'APPDEPLOY_RESOURCE_NOT_APPLICABLE_FOR_MODE', + 'APPDEPLOY_ENVIRONMENT_CPU_POOL_BELOW_ALLOCATED', + 'APPDEPLOY_APP_RUNNER_CONTROL_NOT_CONFIGURED', + 'APPDEPLOY_RUNTIME_ASSIGNMENT_FAILED', + 'APPDEPLOY_REVISION_TIMEOUT', + 'APPDEPLOY_INTERNAL_ERROR', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_AUTH_REJECTED', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_NAMESPACE_MISSING', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_INSUFFICIENT_RBAC', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_DECRYPT_FAILED', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_INVALID_K8S_CONFIG', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_RESOURCE_COLLISION', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_APPLY_FAILED', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_SIGNING_KEY_MISSING', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_CONTROL_ENDPOINT_INVALID', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_K8S_UNREACHABLE', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_APP_RUNNER_NOT_READY', + 'APPDEPLOY_ENVIRONMENT_BOOTSTRAP_SSAR_UNAVAILABLE', + ]) + .optional(), + message: z.string().optional(), + phase: z.string().optional(), + occurredAt: z.iso.datetime().optional(), + detailCode: z.string().optional(), +}) + +export const zGetEnvironmentCapabilitiesResponse = z.object({ + backend: zEnvironmentBackend, + supportedModes: z.array( + z.enum([ + 'ENVIRONMENT_MODE_UNSPECIFIED', + 'ENVIRONMENT_MODE_SHARED', + 'ENVIRONMENT_MODE_ISOLATED', + ]), + ), + namespaceRequired: z.boolean(), +}) + +export const zGetEnvironmentWebAppSubjectsResponse = z.object({ + subjects: z.array(zEnvironmentWebAppSubject), +}) + +export const zGetWebAppAccessModeResponse = z.object({ + accessMode: z.string().optional(), +}) + +export const zGetWebAppPermissionResponse = z.object({ + result: z.boolean().optional(), +}) + +export const zGetWorkflowDeploymentOptionsResponse = z.object({ + environment_variable_slots: z.array(zEnvironmentVariableSlot), + credential_slots: z.array(zCredentialSlot), +}) + +export const zListAppEnvironmentsResponse = z.object({ + data: z.array(zAppEnvironment), +}) + +export const zListEnvironmentApiKeysResponse = z.object({ + data: z.array(zEnvironmentApiKey), +}) + +export const zListEnvironmentTriggersResponse = z.object({ + data: z.array(zEnvironmentTrigger), +}) + +export const zNamedRef = z.object({ + id: z.string(), + displayName: z.string(), +}) + +export const zEnvironmentPoolShare = z.object({ + app: zNamedRef, + isolatedCpu: z.number(), +}) + +/** + * EnvironmentPoolComposition names the largest few apps drawing on a pool and + * sums the rest, because a bar with forty segments says less than one with five + * and a remainder. + */ +export const zEnvironmentPoolComposition = z.object({ + topApps: z.array(zEnvironmentPoolShare).optional(), + otherCpu: z.number().optional(), + otherAppCount: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), +}) + +export const zOperator = z.object({ + type: zOperatorType, + id: z.string(), + display_name: z.string(), +}) + +export const zApplicationInteraction = z.object({ + id: z.string(), + timestamp: z.iso.datetime(), + workflowRunId: z.string(), + status: zApplicationInteractionStatus, + durationSeconds: z.number(), + totalTokens: z.string(), + workspace: zNamedRef, + environment: zNamedRef, + app: zNamedRef, + operator: zOperator.optional(), + invokeFrom: z.string(), + traceId: z.string(), + difyTraceId: z.string(), + deploymentVersionId: z.string(), + error: z.string().optional(), + body: z.string().optional(), + attributesJson: z.string().optional(), + resourceAttributesJson: z.string().optional(), +}) + +export const zGetApplicationInteractionResponse = z.object({ + interaction: zApplicationInteraction, +}) + +export const zPrepareAppDeletionRequest = z.object({ + tenantId: z.string().optional(), + appId: z.string().optional(), +}) + +export const zResolveApiTokenRouteRequest = z.object({ + token: z.string().optional(), +}) + +export const zResolveApiTokenRouteResponse = z.object({ + environmentId: z.string().optional(), + namespace: z.string().optional(), + serviceName: z.string().optional(), + servicePort: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), + environmentStatus: zEnvironmentStatus.optional(), + appId: z.string().optional(), + tenantId: z.string().optional(), + deploymentId: z.string().optional(), + servingRevisionId: z.string().optional(), + deploymentStatus: zDeploymentStatus.optional(), + revoked: z.boolean().optional(), + unavailableReason: z.string().optional(), + targetKind: zRouteTargetKind.optional(), + directUpstream: z.string().optional(), + deploymentGeneration: z.string().optional(), +}) + +export const zResolveWebAppRouteRequest = z.object({ + appCode: z.string().optional(), + passport: z.string().optional(), +}) + +export const zResolveWebAppRouteResponse = z.object({ + decision: z.string().optional(), + environmentId: z.string().optional(), + namespace: z.string().optional(), + serviceName: z.string().optional(), + servicePort: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), + environmentStatus: zEnvironmentStatus.optional(), + appId: z.string().optional(), + tenantId: z.string().optional(), + deploymentId: z.string().optional(), + servingRevisionId: z.string().optional(), + deploymentStatus: zDeploymentStatus.optional(), + unavailableReason: z.string().optional(), + targetKind: zRouteTargetKind.optional(), + directUpstream: z.string().optional(), + deploymentGeneration: z.string().optional(), + endUserId: z.string().optional(), + authType: z.string().optional(), +}) + +export const zRetryEnvironmentBootstrapRequest = z.object({ + environmentId: z.string().optional(), +}) + +/** + * RunnerMemory sizes one app-runner. max_mib is absent when the operator chose + * nothing, in which case the runner gets effective_max_mib. + */ +export const zRunnerMemory = z.object({ + maxMib: z.string().optional(), + effectiveMaxMib: z.string().readonly().optional(), + minMib: z.string().readonly().optional(), +}) + +export const zEnvironment = z.object({ + id: z.string(), + displayName: z.string(), + description: z.string(), + mode: zEnvironmentMode, + backend: zEnvironmentBackend, + status: zEnvironmentStatus, + statusMessage: z.string(), + lastError: zError.optional(), + namespace: z.string().optional(), + managedBy: zEnvironmentManagedBy.optional(), + cpuPool: z.number(), + createdAt: z.iso.datetime(), + updatedAt: z.iso.datetime(), + memory: zRunnerMemory.optional(), + usage: zEnvironmentPoolUsage.optional(), +}) + +export const zCreateEnvironmentResponse = z.object({ + environment: zEnvironment, +}) + +export const zGetEnvironmentResponse = z.object({ + environment: zEnvironment, + pool: zEnvironmentPoolComposition.optional(), +}) + +export const zRetryEnvironmentBootstrapResponse = z.object({ + environment: zEnvironment, +}) + +/** + * RunnerSizing describes a runner an app owns outright, so it is present only + * in an isolated environment; elsewhere the runner belongs to the environment. + */ +export const zRunnerSizing = z.object({ + isolatedCpu: z.number(), + memory: zRunnerMemory, +}) + +/** + * SimpleAccount mirrors the account shape Dify's console API returns, so a + * client can render an account the same way whichever API it came from. + */ +export const zSimpleAccount = z.object({ + id: z.string(), + name: z.string().optional(), + email: z.string().optional(), +}) + +export const zTestConnectionRequest = z.object({ + environmentId: z.string().optional(), +}) + +export const zTestConnectionResponse = z.object({ + reachable: z.boolean().optional(), + message: z.string().optional(), +}) + +export const zUndeployEnvironmentAppResponse = z.object({ + operationId: z.string(), +}) + +export const zUndeployWorkflowResponse = z.object({ + operation: zDeploymentOperationReceipt, +}) + +export const zUnsupportedNodeProvider = z.object({ + plugin_id: z.string(), + provider_id: z.string(), + provider_type: z.string(), + provider_name: z.string(), +}) + +export const zUnsupportedNode = z.object({ + id: z.string(), + type: z.string(), + title: z.string(), + provider: zUnsupportedNodeProvider.optional(), +}) + +export const zPrecheckWorkflowDeploymentResponse = z.object({ + unsupported_nodes: z.array(zUnsupportedNode), +}) + +export const zUpdateEnvironmentDeployedAppResourcesRequest = z.object({ + environmentId: z.string(), + deploymentId: z.string(), + isolatedCpu: z.number(), + maxMemoryMib: z.string().optional(), +}) + +export const zUpdateEnvironmentDeployedAppResourcesResponse = z.object({ + deploymentId: z.string(), + isolatedCpu: z.number(), + allocatedCpuCount: z.number(), + poolCpuCount: z.number(), + memory: zRunnerMemory, +}) + +export const zUpdateEnvironmentRequest = z.object({ + environmentId: z.string().optional(), + displayName: z.string().optional(), + description: z.string().optional(), + cpuPool: z.number().optional(), + maxMemoryMib: z.string().optional(), +}) + +export const zUpdateEnvironmentResponse = z.object({ + environment: zEnvironment, +}) + +export const zWorkflowDeploymentEnvironment = z.object({ + id: z.string().optional(), + name: z.string().optional(), +}) + +export const zSourceVersionDeployment = z.object({ + sourceVersionId: z.string().optional(), + environments: z.array(zWorkflowDeploymentEnvironment).optional(), +}) + +export const zBatchGetSourceVersionDeploymentsResponse = z.object({ + items: z.array(zSourceVersionDeployment).optional(), +}) + +export const zWorkflowDeploymentInput = z.object({ + environment_variables: z.array(zEnvironmentVariableInput).optional(), + credentials: z.array(zCredentialSelectionInput).optional(), +}) + +export const zWorkflowVersion = z.object({ + version: z.string().optional(), + marked_name: z.string().optional(), + id: z.string(), + marked_comment: z.string().optional(), + version_number: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), + created_at: z.number().optional(), + created_by: zSimpleAccount.optional(), + dsl_hash: z.string().optional(), + deleted: z.boolean().optional(), +}) + +export const zDeploymentOperation = z.object({ + id: z.string(), + type: zDeploymentOperationType, + outcome: zDeploymentOperationOutcome, + workspace: zNamedRef, + operator: zOperator, + app: zNamedRef, + environment: zNamedRef, + version: zWorkflowVersion, + failureCode: zDeploymentOperationFailureCode.optional(), + failureMessage: z.string().optional(), + requestedAt: z.iso.datetime(), + finalizedAt: z.iso.datetime().optional(), + durationMilliseconds: z.string().optional(), +}) + +export const zEnvironmentDeployedApp = z.object({ + deploymentId: z.string(), + workspace: zNamedRef, + app: zNamedRef, + status: zEnvironmentDeployedAppStatus, + currentVersion: zWorkflowVersion.optional(), + deployedAt: z.iso.datetime().optional(), + deployedBy: zOperator.optional(), + latestAttempt: zEnvironmentDeployedAppAttempt.optional(), + sizing: zRunnerSizing.optional(), + occupiesPool: z.boolean().optional(), +}) + +export const zEnvironmentDeploymentOperation = z.object({ + id: z.string(), + type: zDeploymentOperationType, + status: zDeploymentOperationStatus, + target_version: zWorkflowVersion.optional(), + operator: zOperator, + activity_at: z.number(), +}) + +export const zEnvironmentDeploymentState = z.object({ + status: zDeploymentStatus, + current_version: zWorkflowVersion.optional(), + versions_behind: z + .int() + .min(0, { error: 'Invalid value: Expected uint32 to be >= 0' }) + .max(4294967295, { error: 'Invalid value: Expected uint32 to be <= 4294967295' }) + .optional(), + deployed_at: z.number().optional(), + deployed_by: zOperator.optional(), + latest_operation: zEnvironmentDeploymentOperation.optional(), +}) + +export const zEnvironmentDeployment = z.object({ + environment: zDeploymentEnvironment, + deployment: zEnvironmentDeploymentState.optional(), + access: zEnvironmentAccess, +}) + +export const zGetDeploymentOperationResponse = z.object({ + operation: zDeploymentOperation, +}) + +export const zGetEnvironmentDeploymentResponse = z.object({ + environment_deployment: zEnvironmentDeployment, +}) + +export const zListEnvironmentDeploymentsResponse = z.object({ + environment_deployments: z.array(zEnvironmentDeployment), +}) + +/** + * Pagination : Just for pagination by page + */ +export const zPagination = z.object({ + totalCount: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), + perPage: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), + currentPage: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), + totalPages: z + .int() + .min(-2147483648, { error: 'Invalid value: Expected int32 to be >= -2147483648' }) + .max(2147483647, { error: 'Invalid value: Expected int32 to be <= 2147483647' }) + .optional(), +}) + +export const zListApplicationInteractionsResponse = z.object({ + data: z.array(zApplicationInteraction), + pagination: zPagination, +}) + +export const zListAppsResponse = z.object({ + data: z.array(zDashboardApp), + pagination: zPagination, +}) + +export const zListDeploymentOperationsResponse = z.object({ + data: z.array(zDeploymentOperation), + pagination: zPagination, +}) + +export const zListEnvironmentDeployedAppsResponse = z.object({ + data: z.array(zEnvironmentDeployedApp), + summary: zEnvironmentDeployedAppSummary, + pagination: zPagination, +}) + +export const zListEnvironmentsResponse = z.object({ + data: z.array(zEnvironment), + pagination: zPagination, +}) + +export const zDeleteEnvironmentApiKeyResponseWritable = z.record(z.string(), z.unknown()) + +export const zDeleteEnvironmentResponseWritable = z.record(z.string(), z.unknown()) + +export const zEnvironmentMcpServerWritable = z.record(z.string(), z.unknown()) + +export const zEnvironmentTriggerWritable = z.record(z.string(), z.unknown()) + +/** + * RunnerMemory sizes one app-runner. max_mib is absent when the operator chose + * nothing, in which case the runner gets effective_max_mib. + */ +export const zRunnerMemoryWritable = z.object({ + maxMib: z.string().optional(), +}) + +export const zEnvironmentWritable = z.object({ + id: z.string(), + displayName: z.string(), + description: z.string(), + mode: zEnvironmentMode, + backend: zEnvironmentBackend, + status: zEnvironmentStatus, + statusMessage: z.string(), + lastError: zError.optional(), + namespace: z.string().optional(), + managedBy: zEnvironmentManagedBy.optional(), + cpuPool: z.number(), + createdAt: z.iso.datetime(), + updatedAt: z.iso.datetime(), + memory: zRunnerMemoryWritable.optional(), + usage: zEnvironmentPoolUsage.optional(), +}) + +export const zCreateEnvironmentResponseWritable = z.object({ + environment: zEnvironmentWritable, +}) + +export const zGetEnvironmentResponseWritable = z.object({ + environment: zEnvironmentWritable, + pool: zEnvironmentPoolComposition.optional(), +}) + +export const zListEnvironmentsResponseWritable = z.object({ + data: z.array(zEnvironmentWritable), + pagination: zPagination, +}) + +export const zRetryEnvironmentBootstrapResponseWritable = z.object({ + environment: zEnvironmentWritable, +}) + +/** + * RunnerSizing describes a runner an app owns outright, so it is present only + * in an isolated environment; elsewhere the runner belongs to the environment. + */ +export const zRunnerSizingWritable = z.object({ + isolatedCpu: z.number(), + memory: zRunnerMemoryWritable, +}) + +export const zEnvironmentDeployedAppWritable = z.object({ + deploymentId: z.string(), + workspace: zNamedRef, + app: zNamedRef, + status: zEnvironmentDeployedAppStatus, + currentVersion: zWorkflowVersion.optional(), + deployedAt: z.iso.datetime().optional(), + deployedBy: zOperator.optional(), + latestAttempt: zEnvironmentDeployedAppAttempt.optional(), + sizing: zRunnerSizingWritable.optional(), + occupiesPool: z.boolean().optional(), +}) + +export const zListEnvironmentDeployedAppsResponseWritable = z.object({ + data: z.array(zEnvironmentDeployedAppWritable), + summary: zEnvironmentDeployedAppSummary, + pagination: zPagination, +}) + +export const zUpdateEnvironmentDeployedAppResourcesResponseWritable = z.object({ + deploymentId: z.string(), + isolatedCpu: z.number(), + allocatedCpuCount: z.number(), + poolCpuCount: z.number(), + memory: zRunnerMemoryWritable, +}) + +export const zUpdateEnvironmentResponseWritable = z.object({ + environment: zEnvironmentWritable, +}) + +export const zConsoleDeploymentServiceListAppEnvironmentsPath = z.object({ + app_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleDeploymentServiceListAppEnvironmentsResponse = zListAppEnvironmentsResponse + +export const zConsoleAccessServiceGetEnvironmentApiPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceGetEnvironmentApiResponse = zEnvironmentApi + +export const zConsoleAccessServiceUpdateEnvironmentApiBody = zEnvironmentApiUpdate + +export const zConsoleAccessServiceUpdateEnvironmentApiPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceUpdateEnvironmentApiResponse = zEnvironmentApi + +export const zConsoleAccessServiceListEnvironmentApiKeysPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceListEnvironmentApiKeysResponse = zListEnvironmentApiKeysResponse + +export const zConsoleAccessServiceCreateEnvironmentApiKeyPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceCreateEnvironmentApiKeyResponse = zEnvironmentApiKey + +export const zConsoleAccessServiceDeleteEnvironmentApiKeyPath = z.object({ + app_id: z.string(), + environment_id: z.string(), + api_key_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceDeleteEnvironmentApiKeyResponse = zDeleteEnvironmentApiKeyResponse + +export const zConsoleAccessServiceGetEnvironmentMcpServerPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceGetEnvironmentMcpServerResponse = zEnvironmentMcpServer + +export const zConsoleAccessServiceGetEnvironmentSitePath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceGetEnvironmentSiteResponse = zEnvironmentSite + +export const zConsoleAccessServiceUpdateEnvironmentSiteBody = zEnvironmentSiteUpdate + +export const zConsoleAccessServiceUpdateEnvironmentSitePath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceUpdateEnvironmentSiteResponse = zEnvironmentSite + +export const zConsoleAccessServiceResetEnvironmentSiteAccessTokenPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceResetEnvironmentSiteAccessTokenResponse = zEnvironmentSite + +export const zConsoleAccessServiceListEnvironmentTriggersPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceListEnvironmentTriggersResponse = zListEnvironmentTriggersResponse + +export const zConsoleAccessServiceUpdateEnvironmentWebAppAccessModeBody = + zEnvironmentWebAppAccessModeUpdate + +export const zConsoleAccessServiceUpdateEnvironmentWebAppAccessModePath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceUpdateEnvironmentWebAppAccessModeResponse = zEnvironmentSite + +export const zConsoleAccessServiceGetEnvironmentWebAppSubjectsPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleAccessServiceGetEnvironmentWebAppSubjectsResponse = + zGetEnvironmentWebAppSubjectsResponse + +export const zConsoleDeploymentServiceListEnvironmentDeploymentsPath = z.object({ + app_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleDeploymentServiceListEnvironmentDeploymentsResponse = + zListEnvironmentDeploymentsResponse + +export const zConsoleDeploymentServiceGetEnvironmentDeploymentPath = z.object({ + app_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleDeploymentServiceGetEnvironmentDeploymentResponse = + zGetEnvironmentDeploymentResponse + +export const zConsoleDeploymentServiceGetWorkflowDeploymentOptionsPath = z.object({ + app_id: z.string(), + workflow_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleDeploymentServiceGetWorkflowDeploymentOptionsResponse = + zGetWorkflowDeploymentOptionsResponse + +export const zConsoleDeploymentServiceDeployWorkflowBody = zWorkflowDeploymentInput + +export const zConsoleDeploymentServiceDeployWorkflowPath = z.object({ + app_id: z.string(), + workflow_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleDeploymentServiceDeployWorkflowResponse = zDeployWorkflowResponse + +export const zConsoleDeploymentServiceUndeployWorkflowPath = z.object({ + app_id: z.string(), + workflow_id: z.string(), + environment_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleDeploymentServiceUndeployWorkflowResponse = zUndeployWorkflowResponse + +export const zConsoleDeploymentServicePrecheckWorkflowDeploymentPath = z.object({ + app_id: z.string(), + workflow_id: z.string(), +}) + +/** + * OK + */ +export const zConsoleDeploymentServicePrecheckWorkflowDeploymentResponse = + zPrecheckWorkflowDeploymentResponse diff --git a/packages/contracts/openapi-ts.enterprise-app-deploy.config.ts b/packages/contracts/openapi-ts.enterprise-app-deploy.config.ts new file mode 100644 index 00000000000..f3e76397e08 --- /dev/null +++ b/packages/contracts/openapi-ts.enterprise-app-deploy.config.ts @@ -0,0 +1,338 @@ +import fs from 'node:fs' +import path from 'node:path' +import { fileURLToPath } from 'node:url' +import { defineConfig } from '@hey-api/openapi-ts' +import { loadOpenApiYaml } from './openapi-yaml' + +type JsonObject = Record + +type OpenApiDocument = JsonObject & { + components?: OpenApiComponents + paths?: Record +} + +type OpenApiComponents = JsonObject & { + schemas?: Record +} + +type OpenApiMediaType = JsonObject & { + schema?: unknown +} + +type OpenApiOperation = JsonObject & { + operationId?: string + responses?: Record +} + +type OpenApiPathItem = Record + +type OpenApiResponse = JsonObject & { + content?: Record +} + +type OpenApiSchema = JsonObject & { + enum?: unknown[] + format?: string + properties?: Record + type?: string | string[] +} + +type ContractOperation = { + id: string + operationId?: string + tags?: readonly string[] +} + +const currentDir = path.dirname(fileURLToPath(import.meta.url)) +const enterpriseServerDir = process.env.DIFY_ENTERPRISE_SERVER + ? path.resolve(process.env.DIFY_ENTERPRISE_SERVER) + : path.resolve(currentDir, '../../../dify-enterprise/server') +const enterpriseOpenApiPath = path.join(enterpriseServerDir, 'pkg/apis/appdeploy/openapi.yaml') +const operationMethods = new Set(['delete', 'get', 'patch', 'post', 'put']) + +const isConsoleApiPath = (routePath: string) => routePath.startsWith('/console/api/') + +const isObject = (value: unknown): value is JsonObject => { + return !!value && typeof value === 'object' && !Array.isArray(value) +} + +const isOpenApiSchema = (value: unknown): value is OpenApiSchema => { + return isObject(value) +} + +const asOpenApiOperation = (value: unknown): OpenApiOperation | undefined => { + return isObject(value) ? (value as OpenApiOperation) : undefined +} + +const asOpenApiResponse = (value: unknown): OpenApiResponse | undefined => { + return isObject(value) ? (value as OpenApiResponse) : undefined +} + +const asOpenApiMediaType = (value: unknown): OpenApiMediaType | undefined => { + return isObject(value) ? (value as OpenApiMediaType) : undefined +} + +const stripConsoleApiPrefix = (routePath: string) => { + if (isConsoleApiPath(routePath)) return routePath.replace('/console/api', '') + + return routePath +} + +const stripSchemaNamePrefix = (schemaName: string) => { + return schemaName + .replace(/^dify\.enterprise\.api\.enterprise\./, '') + .replace(/^dify\.enterprise\.api\.appdeploy\.v1\./, '') + .replace(/^dify\.enterprise\.api\.appdeploy\./, '') + .replace(/^pagination\./, '') +} + +const contractTagSegment = (tag?: string) => { + if (tag === 'ConsoleDeploymentService') return 'DeploymentService' + if (tag === 'ConsoleAccessService') return 'AccessService' + + return tag || 'default' +} + +const contractNameSegments = (operation: ContractOperation) => { + const operationId = operation.operationId || operation.id + const tag = operation.tags?.[0] + const tagPrefixPattern = tag ? new RegExp(`^${tag}[._/-]`) : undefined + const name = tagPrefixPattern ? operationId.replace(tagPrefixPattern, '') : operationId + const segments = name.split(/[._/-]+/).filter(Boolean) + + return segments.length > 0 ? segments : [operationId] +} + +const contractPathSegments = (operation: ContractOperation) => { + return [contractTagSegment(operation.tags?.[0]), ...contractNameSegments(operation)] +} + +const hasSchemaLessResponseContent = (operation: OpenApiOperation) => { + if (!isObject(operation.responses)) return false + + return Object.values(operation.responses).some((response) => { + const openApiResponse = asOpenApiResponse(response) + if (!openApiResponse || !isObject(openApiResponse.content)) return false + + return Object.values(openApiResponse.content).some((mediaType) => { + const openApiMediaType = asOpenApiMediaType(mediaType) + return !!openApiMediaType && !('schema' in openApiMediaType) + }) + }) +} + +// protoc-gen-openapi emits google.api.HttpBody responses as `*/*: {}`. Skip these +// raw download operations until the source OpenAPI exposes an explicit schema. +const stripSchemaLessResponseOperations = (pathItem: OpenApiPathItem) => { + return Object.fromEntries( + Object.entries(pathItem).filter(([method, operation]) => { + if (!operationMethods.has(method.toLowerCase())) return true + + const openApiOperation = asOpenApiOperation(operation) + return !openApiOperation || !hasSchemaLessResponseContent(openApiOperation) + }), + ) +} + +const toWords = (value: string) => { + return value + .replace(/[{}]/g, '') + .replace(/([a-z0-9])([A-Z])/g, '$1 $2') + .split(/[^a-z0-9]+/i) + .filter(Boolean) +} + +const toPascalCase = (words: string[]) => { + return words.map((word) => `${word.charAt(0).toUpperCase()}${word.slice(1)}`).join('') +} + +const commonWordPrefix = (values: string[]) => { + const wordLists = values.map((value) => value.split('_')) + const firstWords = wordLists[0] ?? [] + const prefix: string[] = [] + + for (const [index, word] of firstWords.entries()) { + if (!wordLists.every((words) => words[index] === word)) break + + prefix.push(word) + } + + return prefix +} + +const enumSchemaNameFromValues = (values: unknown[]) => { + if (values.length === 0 || !values.every((value) => typeof value === 'string')) return undefined + + const prefix = commonWordPrefix(values) + if (prefix.length < 2) return undefined + + return toPascalCase(prefix.map((word) => word.toLowerCase())) +} + +const findSchemaEntry = ( + schemas: Record, + schemaName: string, +): [string, OpenApiSchema] | undefined => { + return Object.entries(schemas).find(([name]) => stripSchemaNamePrefix(name) === schemaName) +} + +const enumValuesKey = (values: unknown[]) => JSON.stringify(values) + +const reusableEnumSchema = (propertySchema: OpenApiSchema): OpenApiSchema => ({ + ...(propertySchema.format ? { format: propertySchema.format } : {}), + enum: propertySchema.enum, + type: propertySchema.type ?? 'string', +}) + +const enumSchemaKey = ( + schemas: Record, + preferredName: string, + valuesKey: string, + valuesToSchemaKey: Map, + schemaName: string, + propertyName: string, +) => { + const existingKey = valuesToSchemaKey.get(valuesKey) + if (existingKey) return existingKey + + const existingEnumEntry = findSchemaEntry(schemas, preferredName) + if (!existingEnumEntry) return preferredName + + const existingEnumValues = existingEnumEntry[1].enum + if (Array.isArray(existingEnumValues) && enumValuesKey(existingEnumValues) === valuesKey) + return existingEnumEntry[0] + + return `${stripSchemaNamePrefix(schemaName)}${toPascalCase(toWords(propertyName))}` +} + +const promoteInlineEnumSchema = ( + schemas: Record, + schemaName: string, + properties: Record, + propertyName: string, + propertySchema: OpenApiSchema, + valuesToSchemaKey: Map, +) => { + if (!Array.isArray(propertySchema.enum)) return + + const preferredName = enumSchemaNameFromValues(propertySchema.enum) + if (!preferredName) return + + const valuesKey = enumValuesKey(propertySchema.enum) + const key = enumSchemaKey( + schemas, + preferredName, + valuesKey, + valuesToSchemaKey, + schemaName, + propertyName, + ) + + if (!schemas[key]) schemas[key] = reusableEnumSchema(propertySchema) + + valuesToSchemaKey.set(valuesKey, key) + properties[propertyName] = { + $ref: `#/components/schemas/${key}`, + } +} + +// gnostic's protoc-gen-openapi inlines proto enum schemas into every field. +// Promote prefixable inline enums to reusable schemas so Hey API can emit +// runtime enum objects from the generated contract. +const promoteReusableEnumSchemasForHeyApi = (document: OpenApiDocument) => { + const schemas = document.components?.schemas + if (!schemas) return + + const valuesToSchemaKey = new Map() + + Object.entries(schemas).forEach(([schemaName, schema]) => { + const properties = schema.properties + if (!properties) return + + Object.entries(properties).forEach(([propertyName, propertySchema]) => { + if (!isOpenApiSchema(propertySchema)) return + + promoteInlineEnumSchema( + schemas, + schemaName, + properties, + propertyName, + propertySchema, + valuesToSchemaKey, + ) + }) + }) +} + +const normalizeEnterpriseOpenApi = () => { + const openApi = loadOpenApiYaml(fs.readFileSync(enterpriseOpenApiPath, 'utf8')) + + if (!openApi || typeof openApi !== 'object' || Array.isArray(openApi)) + throw new Error(`Invalid enterprise OpenAPI document: ${enterpriseOpenApiPath}`) + + const document = openApi as OpenApiDocument + const paths = document.paths ?? {} + + document.paths = Object.fromEntries( + Object.entries(paths) + .filter(([routePath]) => isConsoleApiPath(routePath)) + .map(([routePath, pathItem]) => { + if (!isObject(pathItem)) return [stripConsoleApiPrefix(routePath), pathItem] + + return [stripConsoleApiPrefix(routePath), stripSchemaLessResponseOperations(pathItem)] + }) + .filter(([, pathItem]) => !isObject(pathItem) || Object.keys(pathItem).length > 0), + ) + + promoteReusableEnumSchemasForHeyApi(document) + + return document +} + +export default defineConfig({ + input: normalizeEnterpriseOpenApi(), + output: { + entryFile: false, + path: 'generated/enterprise-app-deploy', + fileName: { + suffix: '.gen', + }, + postProcess: [ + { + command: 'vp', + args: ['fmt', '{{path}}'], + }, + ], + }, + parser: { + transforms: { + schemaName: stripSchemaNamePrefix, + }, + }, + plugins: [ + { + name: '@hey-api/typescript', + comments: false, + enums: { + mode: 'javascript', + }, + }, + 'zod', + { + name: 'orpc', + contracts: { + strategy: 'single', + contractName: { + name: '{{name}}', + casing: 'camelCase', + }, + nesting: contractPathSegments, + segmentName: { + name: '{{name}}', + casing: 'camelCase', + }, + }, + validator: 'zod', + }, + ], +}) diff --git a/packages/contracts/package.json b/packages/contracts/package.json index ac0c28e7911..b7af1f88ad3 100644 --- a/packages/contracts/package.json +++ b/packages/contracts/package.json @@ -20,6 +20,10 @@ "types": "./generated/enterprise/*.ts", "import": "./generated/enterprise/*.ts" }, + "./enterprise-app-deploy/*": { + "types": "./generated/enterprise-app-deploy/*.ts", + "import": "./generated/enterprise-app-deploy/*.ts" + }, "./knowledge-fs/*": { "types": "./generated/knowledge-fs/*.ts", "import": "./generated/knowledge-fs/*.ts" @@ -28,6 +32,7 @@ "scripts": { "gen-api-contract": "uv run --project ../../api ../../api/dev/generate_swagger_specs.py --output-dir openapi && uv run --project ../../api ../../api/dev/generate_fastopenapi_specs.py --output-dir openapi && node -e \"fs.rmSync('generated/api', { recursive: true, force: true })\" && openapi-ts -f openapi-ts.api.config.ts && vp fmt generated/api", "gen-enterprise-contract": "openapi-ts -f openapi-ts.enterprise.config.ts", + "gen-enterprise-app-deploy-contract": "openapi-ts -f openapi-ts.enterprise-app-deploy.config.ts", "gen-knowledge-fs-contract": "node scripts/generate-knowledge-fs-contract.mjs", "test": "vp test", "type-check": "tsc" diff --git a/packages/contracts/sandbox-contract.smoke.test.ts b/packages/contracts/sandbox-contract.smoke.test.ts index c1e20b62905..b7f272b1ded 100644 --- a/packages/contracts/sandbox-contract.smoke.test.ts +++ b/packages/contracts/sandbox-contract.smoke.test.ts @@ -9,6 +9,7 @@ describe('generated sandbox contracts', () => { ])('exposes the %s file operations', (_, sandbox) => { expect(sandbox.files.get).toBeDefined() expect(sandbox.files.read.get).toBeDefined() - expect(sandbox.files.upload.post).toBeDefined() + expect(sandbox.files.download.post).toBeDefined() + expect('upload' in sandbox.files).toBe(false) }) }) diff --git a/packages/dify-ui/AGENTS.md b/packages/dify-ui/AGENTS.md index d37f44e934d..255d56f6bed 100644 --- a/packages/dify-ui/AGENTS.md +++ b/packages/dify-ui/AGENTS.md @@ -1,91 +1,31 @@ # @langgenius/dify-ui -Shared design tokens, the `cn()` utility, CSS-first Tailwind styles, and headless primitive components consumed by `web/`. +This package owns shared design tokens, CSS-first Tailwind styles, the `cn()` utility, and headless primitives consumed by `web/`. Read only the matching topic in [`README.md`] for public imports, forms, typed values, pickers, overlays, tokens, or tests. ## Component Authoring Rules -- Use `@base-ui/react` primitives + `cva` + `cn`. -- Inside dify-ui, cross-component imports use relative paths (`../button`). External consumers use subpath exports (`@langgenius/dify-ui/button`). -- No imports from `web/`. No dependencies on next / i18next / ky / jotai / zustand. -- One component per folder: `src//index.tsx`, optional `index.stories.tsx` and `__tests__/index.spec.tsx`. Add a matching `./` subpath to `package.json#exports`. -- Name the canonical public boundary and its associated public types after the primitive without a `Root` suffix (`Select` / `SelectProps`). Keep `Root` only when the subpath also exports a higher-level convenience component that must be distinguished from the low-level anatomy root (`CheckboxRoot` / `Checkbox`). Preserve the upstream anatomy in implementation type sources such as `BaseSelect.Root.Props`. -- Props pattern: `Omit & VariantProps & { /* custom */ }`. -- Use plain `Omit<...>` only for non-union Base UI props. When a prop changes the valid shape of related props (for example `value` / `defaultValue`, `multiple` / `value`, or `clearable` / `onChange`), model that relationship with an explicit discriminated union or a distributive helper instead of flattening the props. -- Preserve Base UI generic value contracts in wrappers. If the upstream primitive is generic, expose the same generic parameters and pass them through to the Base UI part, such as `Select.Root`, `RadioGroup`, or `Radio.Root`. -- Do not hard-code selection wrappers to `string` unless the upstream primitive is string-only. Select and radio wrappers that carry selected values should preserve the primitive's generic value contract. -- When a component accepts a prop typed from a shared internal module, `export type` it from that component so consumers import it from the component subpath. -- Prefer Base UI data attributes and CSS variables for visual states; do not mirror state in React solely to add classes. -- When a Base UI API or selector contract is unclear, read the docs linked from `README.md` and the local `@base-ui/react` `.d.ts` files before coding. +- Build primitives from `@base-ui/react`, `cva`, and `cn`. +- Use relative cross-component imports inside the package and subpath exports such as `@langgenius/dify-ui/button` from consumers. Add a matching `package.json#exports` entry for each public primitive. +- Keep one primitive per `src//` folder with optional colocated stories and tests. +- Do not import from `web/` or depend on Next.js, i18n, application state, or data-fetching libraries. +- Preserve upstream Base UI anatomy and generic value contracts. Use discriminated unions when one prop changes the valid shape of related props; do not flatten those relationships or hard-code selectable values to `string`. +- Export shared public types from the owning component subpath. +- Prefer Base UI data attributes and CSS variables for visual states; do not mirror primitive state in React solely to add classes. +- When a Base UI API or selector contract is unclear, read the current official documentation and local `@base-ui/react` type declarations before coding. -## Overlay Primitive Selection: Tooltip vs PreviewCard vs Popover +Use the README sections as the detailed owners: -Pick by the **trigger's purpose** and **a11y reach**, not visual richness. +- [Imports and public boundaries] +- [Typed value contracts] +- [Search and picker selection] +- [Tailwind and Figma radius mapping] +- [Overlay and portal contracts] +- [Development and test boundaries] -| Primitive | Opens on | Trigger's purpose | Content | Reachable on touch / SR? | -| ------------- | --------------------- | --------------------- | ------------------------- | ------------------------ | -| `Tooltip` | hover / focus | has its own action | short plain-text label | ❌ (label only) | -| `PreviewCard` | hover / focus | navigate through link | link destination preview | ❌ (visual enhancement) | -| `Popover` | click / tap (+ hover) | **open the popup** | anything, incl. long text | ✅ | - -Base UI decision rule ([docs]): - -> _"If the trigger's purpose is to open the popup itself, it's a popover. -> If the trigger's purpose is unrelated to opening the popup, it's a tooltip."_ - -Apply this first, then narrow: - -- `Tooltip` — ephemeral visual label. Trigger must already carry its own `aria-label` / visible text; tooltip mirrors it for sighted mouse/keyboard users. No interactive UI, no multi-line prose. Not dwell-able. -- `PreviewCard` — a visual enhancement for a link that previews its destination. Prefer the canonical anchor trigger and keep the popup non-interactive. Do not place unique or essential information or actions in the popup unless they are also available at the linked destination; touch and screen reader users cannot access the preview. If opening the popup is itself the trigger's purpose, or its content must be accessible across input modes, use `Popover` instead. -- `Popover` — any popup with its own interactions, or any "infotip" (`?` / `(i)` glyph whose sole purpose is to reveal help text). Pass `openOnHover` on `PopoverTrigger` for the infotip case — unlike `Tooltip` / `PreviewCard`, this stays accessible to touch and SR users because the popover still opens on tap and focus. - -Product-level polymorphic trigger compositions are local feature decisions. Document them in the owning feature and do not broaden or weaken the shared primitive contract to match one business workflow. - -## Border Radius: Figma Token → Tailwind Class Mapping - -The Figma design system uses `--radius/*` tokens whose scale is **offset by one step** from Tailwind CSS v4 defaults. When translating Figma specs to code, always use this mapping — never use `radius-*` as a CSS class, and never extend `borderRadius` in the preset. - -| Figma Token | Value | Tailwind Class | -| --------------- | ----- | ---------------- | -| `--radius/2xs` | 2px | `rounded-xs` | -| `--radius/xs` | 4px | `rounded-sm` | -| `--radius/sm` | 6px | `rounded-md` | -| `--radius/md` | 8px | `rounded-lg` | -| `--radius/lg` | 10px | `rounded-[10px]` | -| `--radius/xl` | 12px | `rounded-xl` | -| `--radius/2xl` | 16px | `rounded-2xl` | -| `--radius/3xl` | 20px | `rounded-[20px]` | -| `--radius/6xl` | 28px | `rounded-[28px]` | -| `--radius/full` | 999px | `rounded-full` | - -### Rules - -- **Do not** add custom `borderRadius` theme values. We use Tailwind v4 defaults and arbitrary values (`rounded-[Npx]`) for sizes without a standard equivalent. -- **Do not** use `radius-*` as CSS class names. The old `@utility radius-*` definitions have been removed. -- When the Figma MCP returns `rounded-[var(--radius/sm, 6px)]`, convert it to the standard Tailwind class from the table above (e.g. `rounded-md`). -- For values without a standard Tailwind equivalent (10px, 20px, 28px), use arbitrary values like `rounded-[10px]`. - -## Search / Picker Primitive Selection: Autocomplete vs Combobox vs Select - -Pick by whether the user is entering free-form text, choosing a remembered value, or selecting from a closed list. - -Base UI decision rules: - -- [Autocomplete docs]: use `Combobox` instead of `Autocomplete` if the selection should be remembered and the input value cannot be custom. -- [Combobox docs]: do not use `Combobox` for simple search widgets that require unrestricted text entry; use `Autocomplete` instead. - -Apply this split in Dify UI: - -- `Autocomplete` — free-form text input with optional suggestions or completions. The input value may be custom and does not necessarily become a selected option. Use for search boxes, command-style suggestions, tag suggestions, and async text completion. -- `Combobox` — searchable picker whose value is one or more selected items from a collection. The chosen value is remembered by the root, and free-form text is not the final value. Use for model pickers, user pickers, dataset/document pickers, and multi-select chips. -- `Select` — closed-list picker without text entry. Use when the option set is small or already scannable and filtering is unnecessary. - -Composition rules: - -- Keep Base UI primitive semantics visible in the public API. Export compound parts such as `ComboboxInputGroup`, `ComboboxInput`, `ComboboxContent`, `ComboboxList`, `ComboboxItem`, and `ComboboxItemIndicator` instead of wrapping them into one business component. -- For `Combobox` multiple selection, follow the official chips pattern: `ComboboxInputGroup` contains `ComboboxChips`, `ComboboxValue` renders `ComboboxChip` items, and `ComboboxInput` remains inside the chips row. Chips should wrap and let the input group grow vertically instead of forcing horizontal overflow. -- Content primitives must own their Base UI `Portal` and use `z-50` on `Positioner`, matching the overlay contract in `README.md`. Toast owns `z-60`. -- Use `w-(--anchor-width)` with viewport-aware max-width for `Autocomplete` and `Combobox` popups. Do not add `min-w-(--anchor-width)` when it would defeat available-width clamping. - -[Autocomplete docs]: https://base-ui.com/react/components/autocomplete.md#usage-guidelines -[Combobox docs]: https://base-ui.com/react/components/combobox.md#usage-guidelines -[docs]: https://base-ui.com/react/components/tooltip#infotips +[Development and test boundaries]: README.md#development +[Imports and public boundaries]: README.md#imports +[Overlay and portal contracts]: README.md#overlay--portal-contract +[Search and picker selection]: README.md#search-and-picker-selection +[Tailwind and Figma radius mapping]: README.md#tailwind-css-v4-integration +[Typed value contracts]: README.md#typed-value-contracts +[`README.md`]: README.md diff --git a/packages/dify-ui/README.md b/packages/dify-ui/README.md index 6345c4e010e..0a06fa1f2d9 100644 --- a/packages/dify-ui/README.md +++ b/packages/dify-ui/README.md @@ -43,6 +43,22 @@ Importing from `@langgenius/dify-ui` (no subpath) is intentionally not supported The canonical boundary exported from a primitive subpath uses the primitive name without a `Root` suffix, and its public types follow the same name (`Select` / `SelectProps`, `Drawer` / `DrawerProps`). Keep `Root` only when the subpath exposes both a low-level anatomy root and a higher-level convenience component, such as `CheckboxRoot` / `Checkbox` or `PaginationRoot` / `Pagination`. Implementation code should continue to reference the upstream Base UI anatomy explicitly through names such as `BaseSelect.Root.Props`. +### Public type contracts + +Every runtime component exported from a primitive subpath must have an accurate, importable props type with the matching name (`DialogContent` / `DialogContentProps`). Define Dify-authored composite props at the Dify UI boundary; use direct aliases for unchanged Base UI parts instead of copying their shapes. + +Treat each `src//index.tsx` as an explicit public API boundary. Keep component, factory, hook, and type declarations module-local, then publish the complete surface through separate `export { ... }` and `export type { ... }` manifests at the bottom of the file. Do not mix scattered inline exports with the manifest or use a wildcard export; omission from the manifest keeps implementation helpers private. + +Preserve generic relationships end to end. Generic public components and their props must carry the same caller-owned type parameters, including picker `Value` / `Multiple`, form values, radio values, slider values, and overlay payloads and handles. Never erase those relationships with `any` or a hard-coded `string`; use `unknown` only as the safe default for an independently consumed anatomy part whose value cannot be inferred from its parent through JSX. + +Do not add a root-only generic when independently rendered JSX anatomy can produce values outside that root's inferred type. Preserve the upstream contract until the complete component family can enforce one value type; otherwise the generic gives callbacks a narrower type than the runtime can guarantee. `Tabs` intentionally follows Base UI's value contract for this reason. + +Keep the public type surface smaller than the upstream Base UI namespace. A type is not public merely because Base UI provides a name for it, because Dify gives an internal contract a descriptive alias, or because an earlier implementation happened to export it. In addition to matching component props, export a type only when it pairs with a public factory or has a concrete workspace consumer that cannot express the contract clearly through the matching props type. Remove legacy aliases that have no matching runtime API or real consumer; package-local tests and stories can derive narrow values from canonical props instead of preserving an otherwise unused export. Dify-authored options, states, and controlled/uncontrolled branches stay private when the matching props type already expresses them. + +State, event details and reasons, actions, and controlled/uncontrolled composition helpers are private by default. Public props already provide contextual typing for inline render and event callbacks. Export one of these narrower contracts only when Dify defines the state or event itself, a public factory requires the named type, or an external consumer needs to name it independently. Apply the same rule to Dify-authored components such as FileTree, Pagination, ProgressCircle, StatusDot, and Toast; being implemented locally does not justify a broader API. + +Keep implementation-only render helpers, context values, styling helpers, and upstream passthrough aliases private. When a wrapper consumes `className` through `cn()`, omit the upstream state-callback form and expose `className?: string`; public types must describe behavior the wrapper actually implements. + ## Primitives | Category | Subpath | Notes | @@ -65,6 +81,11 @@ Utilities: ## Button loading and disabled contract +`Button` owns the spacing between direct children. Regular (`medium`) and Large sizes use +4px and 6px gaps. Small uses 3px for Primary and 4px for the other variants. Do not add icon +margins or a standard `gap-*` at call sites; use a Button `className` override only for a +documented layout exception. + `Button` keeps normal `disabled` controls native-disabled by default so unavailable actions are removed from the keyboard focus order. When `loading` is true, `Button` defaults `focusableWhenDisabled` to true. Loading represents an action that has already been triggered and is temporarily pending, so the button remains focusable while Base UI still suppresses click, pointer, keyboard activation, and submit-button activation. Pass `focusableWhenDisabled={false}` only when a loading button should use native disabled behavior. @@ -107,7 +128,7 @@ Use `Fieldset` and `FieldsetLegend` when one field is represented by a group of Selection primitives should preserve the caller's domain value type instead of widening values to `string`. Use `Select`, `RadioGroup`, `Radio`, and `RadioItem` when the selected value is an enum, union, boolean, number, object, or nullable placeholder value. -Root-level generics type `value`, `defaultValue`, `onValueChange`, and collection props such as `items`, but JSX children do not automatically inherit the parent generic. For non-string radio groups, type the child item too: +Root-level generics type `value`, `defaultValue`, and value-dependent callbacks such as `onValueChange`. Upstream collection props such as `items` can still be broader than the root value type, and JSX children do not automatically inherit the parent generic. Type independently consumed anatomy parts when their value cannot be inferred locally; for non-string radio groups, type the child item too: ```tsx value={promptMode} onValueChange={setPromptMode}> @@ -121,7 +142,22 @@ Root-level generics type `value`, `defaultValue`, `onValueChange`, and collectio Use `Radio` for the default Dify control. Use `RadioItem` when custom UI should be the radio item; place `RadioControl` inside it for the standard visual dot. `RadioControl` is a Dify visual part, not a Base UI anatomy export. -For select labels and display values, prefer the Base UI `items` collection pattern so the root, value display, and item list share the same typed values. Avoid helpers that stringify values only to recover labels later; convert values to strings only at real boundaries such as form submission, URL/search params, or legacy APIs that require strings. +For `Select` and `Combobox`, a literal multiple-value contract must match the runtime mode: ` multiple>`. Their value display can still receive `null` before a value is selected, including in multiple mode. Because JSX does not pass the root's generic to its children, repeat the domain type on independently consumed render anatomy instead of annotating callback parameters: + +```tsx + multiple value={subjects} onValueChange={setSubjects}> + > + {(selectedSubjects) => selectedSubjects?.map((subject) => subject.name).join(', ') ?? 'Anyone'} + + > + {(subject) => {subject.name}} + + +``` + +`AutocompleteList` follows the same rule. A `ComboboxGroup` or `AutocompleteGroup` can infer its local value type from `items`; a nested `Collection` is a separate JSX boundary and should use `` when its render callback needs the domain type. Dynamic `multiple={condition}` remains supported and gives callbacks the corresponding single-or-multiple union. + +For select labels and display values, prefer the Base UI `items` collection pattern so the root, value display, and item list share one runtime source of truth for values and labels. Avoid helpers that stringify values only to recover labels later; convert values to strings only at real boundaries such as form submission, URL/search params, or legacy APIs that require strings. `CheckboxGroup` follows the Base UI contract and uses `string[]`. Do not add a generic checkbox-group wrapper unless the underlying primitive contract changes; if different business IDs need stronger separation, model that at the feature/domain type boundary. @@ -129,6 +165,16 @@ For complex business forms, keep state ownership outside these primitives. TanSt Migration rule for `web/`: if a UI has a save/submit action, do not leave it as unrelated `Input` and `Button` pieces. Give it a real submit boundary with `Form` or a native `
`, attach visible field names through the appropriate label primitive (`FieldLabel`, `SelectLabel`, `SliderLabel`, or `FieldsetLegend`), expose helper/error text through `FieldDescription` / `FieldError`, and keep non-submit buttons as `type="button"`. +## Search and picker selection + +Choose the primitive by its value contract: + +- `Autocomplete` accepts free-form text with optional suggestions or completions. +- `Combobox` selects and remembers one or more values from a searchable collection. +- `Select` chooses from a closed, scannable list without text entry. + +Keep Base UI anatomy visible in public APIs instead of wrapping a picker into one business component. Multiple-selection comboboxes follow the official chips composition: chips and input share the input group, chips wrap, and the group grows vertically. Autocomplete and Combobox popups own their portals, use the package overlay layer, and size from `--anchor-width` with viewport-aware maximum width; do not force a minimum width that defeats viewport clamping. + ## Tailwind CSS v4 integration This package uses Tailwind CSS v4's CSS-first configuration model. Consumers should import Tailwind from their own root stylesheet, then import this package's CSS entry: @@ -144,6 +190,23 @@ If a consumer uses Dify UI source files through the workspace, add an explicit s @source '../packages/dify-ui/src'; ``` +Figma radius tokens are offset by one step from Tailwind CSS v4 defaults. Use this mapping rather than adding custom theme values or `radius-*` utilities: + +| Figma token | Tailwind class | +| --------------- | ---------------- | +| `--radius/2xs` | `rounded-xs` | +| `--radius/xs` | `rounded-sm` | +| `--radius/sm` | `rounded-md` | +| `--radius/md` | `rounded-lg` | +| `--radius/lg` | `rounded-[10px]` | +| `--radius/xl` | `rounded-xl` | +| `--radius/2xl` | `rounded-2xl` | +| `--radius/3xl` | `rounded-[20px]` | +| `--radius/6xl` | `rounded-[28px]` | +| `--radius/full` | `rounded-full` | + +Convert Figma output such as `rounded-[var(--radius/sm, 6px)]` to the mapped Tailwind class. Use an arbitrary value only when no standard class matches. + ## Overlay & portal contract Overlay primitives render their floating surfaces inside a [Base UI Portal] attached to `document.body`. This is the Base UI default — see the upstream [Portals][Base UI Portal] docs for the underlying behavior. Convenience content components such as `DialogContent`, `PopoverContent`, and `SelectContent` own their portal internally; primitives with explicit portal anatomy such as `Drawer` expose the matching `DrawerPortal` part so consumers can compose the full Base UI structure. @@ -189,7 +252,7 @@ See `[web/docs/overlay.md](../../web/docs/overlay.md)` for the web app overlay b ## Development -- `vp run @langgenius/dify-ui#lint` (from the repository root) — strict Oxlint checks for component source, stories, tests, and package configuration. +- `vp check packages/dify-ui` (from the repository root) — formatting, lint, and TypeScript diagnostics for the package. - `pnpm -C packages/dify-ui test` — Vitest unit tests for primitives. - `pnpm -C packages/dify-ui storybook` — Storybook on the default port. Each primitive has `index.stories.tsx`. - `pnpm -C packages/dify-ui test:storybook` — Storybook component tests in Vitest browser mode. Stories without `play` are render and a11y smoke tests; stories with `play` should cover public UI contracts such as opening overlays, keyboard navigation, disabled/loading guards, form submission, and controlled state updates. @@ -228,7 +291,6 @@ Set the Base UI test flag in a Vitest setup file to skip those waits: See `[AGENTS.md](./AGENTS.md)` for: - Component authoring rules (one-component-per-folder, `cva` + `cn`, relative imports inside the package, subpath imports from consumers). -- Figma `--radius/`_ token → Tailwind `rounded-_` class mapping. ## Not part of this package diff --git a/packages/dify-ui/src/alert-dialog/index.tsx b/packages/dify-ui/src/alert-dialog/index.tsx index d0a3cb1e729..451c360c7d3 100644 --- a/packages/dify-ui/src/alert-dialog/index.tsx +++ b/packages/dify-ui/src/alert-dialog/index.tsx @@ -7,10 +7,15 @@ import { Button } from '../button' import { cn } from '../cn' import { modalBackdropClassName, modalPopupAnimationClassName } from '../overlay-shared' -export const AlertDialog = BaseAlertDialog.Root -export const AlertDialogTrigger = BaseAlertDialog.Trigger -export const AlertDialogTitle = BaseAlertDialog.Title -export const AlertDialogDescription = BaseAlertDialog.Description +const AlertDialog = BaseAlertDialog.Root +const AlertDialogTrigger = BaseAlertDialog.Trigger +const AlertDialogTitle = BaseAlertDialog.Title +const AlertDialogDescription = BaseAlertDialog.Description + +type AlertDialogProps = BaseAlertDialog.Root.Props +type AlertDialogTriggerProps = BaseAlertDialog.Trigger.Props +type AlertDialogTitleProps = BaseAlertDialog.Title.Props +type AlertDialogDescriptionProps = BaseAlertDialog.Description.Props type AlertDialogContentProps = { children: React.ReactNode @@ -19,7 +24,7 @@ type AlertDialogContentProps = { backdropProps?: Omit } -export function AlertDialogContent({ +function AlertDialogContent({ children, className, backdropClassName, @@ -46,7 +51,7 @@ export function AlertDialogContent({ type AlertDialogActionsProps = React.ComponentProps<'div'> -export function AlertDialogActions({ className, ...props }: AlertDialogActionsProps) { +function AlertDialogActions({ className, ...props }: AlertDialogActionsProps) { return (
& { closeProps?: Omit } -export function AlertDialogCancelButton({ +function AlertDialogCancelButton({ children, closeProps, ...buttonProps @@ -74,10 +79,32 @@ export function AlertDialogCancelButton({ type AlertDialogConfirmButtonProps = ButtonProps -export function AlertDialogConfirmButton({ +function AlertDialogConfirmButton({ variant = 'primary', tone = 'destructive', ...props }: AlertDialogConfirmButtonProps) { return + ) }, } diff --git a/packages/dify-ui/src/button/index.tsx b/packages/dify-ui/src/button/index.tsx index cd8d0a8b7bb..79892eb982e 100644 --- a/packages/dify-ui/src/button/index.tsx +++ b/packages/dify-ui/src/button/index.tsx @@ -6,48 +6,46 @@ import { Button as BaseButton } from '@base-ui/react/button' import { cva } from 'class-variance-authority' import { cn } from '../cn' -export const buttonVariants = cva( - 'inline-flex cursor-pointer items-center justify-center whitespace-nowrap outline-hidden focus-visible:ring-2 focus-visible:ring-state-accent-solid data-[disabled]:cursor-not-allowed', +const buttonVariants = cva( + 'inline-flex cursor-pointer items-center justify-center overflow-hidden whitespace-nowrap focus-visible:ring-2 focus-visible:ring-state-accent-solid focus-visible:outline-hidden data-disabled:cursor-not-allowed', { variants: { variant: { primary: [ - 'border-components-button-primary-border bg-components-button-primary-bg text-components-button-primary-text shadow-sm', - 'hover:border-components-button-primary-border-hover hover:bg-components-button-primary-bg-hover', - 'data-[disabled]:border-components-button-primary-border-disabled data-[disabled]:bg-components-button-primary-bg-disabled data-[disabled]:text-components-button-primary-text-disabled data-[disabled]:shadow-none', + 'bg-components-button-primary-bg text-components-button-primary-text shadow-primary-button inset-ring-[0.5px] inset-ring-components-button-primary-border', + 'hover:bg-components-button-primary-bg-hover hover:shadow-xs hover:shadow-shadow-shadow-3 hover:inset-ring-components-button-primary-border-hover', + 'data-disabled:bg-components-button-primary-bg-disabled data-disabled:text-components-button-primary-text-disabled data-disabled:shadow-none data-disabled:inset-ring-components-button-primary-border-disabled', ], secondary: [ - 'border-[0.5px] shadow-xs backdrop-blur-[5px]', - 'border-components-button-secondary-border bg-components-button-secondary-bg text-components-button-secondary-text', - 'hover:border-components-button-secondary-border-hover hover:bg-components-button-secondary-bg-hover', - 'data-[disabled]:border-components-button-secondary-border-disabled data-[disabled]:bg-components-button-secondary-bg-disabled data-[disabled]:text-components-button-secondary-text-disabled data-[disabled]:backdrop-blur-xs', + 'bg-components-button-secondary-bg text-components-button-secondary-text shadow-xs inset-ring-[0.5px] shadow-shadow-shadow-3 inset-ring-components-button-secondary-border backdrop-blur-[5px]', + 'hover:bg-components-button-secondary-bg-hover hover:inset-ring-components-button-secondary-border-hover', + 'data-disabled:bg-components-button-secondary-bg-disabled data-disabled:text-components-button-secondary-text-disabled data-disabled:shadow-none data-disabled:inset-ring-components-button-secondary-border-disabled data-disabled:backdrop-blur-xs', ], 'secondary-accent': [ - 'border-[0.5px] shadow-xs', - 'border-components-button-secondary-border bg-components-button-secondary-bg text-components-button-secondary-accent-text', - 'hover:border-components-button-secondary-border-hover hover:bg-components-button-secondary-bg-hover', - 'data-[disabled]:border-components-button-secondary-border-disabled data-[disabled]:bg-components-button-secondary-bg-disabled data-[disabled]:text-components-button-secondary-accent-text-disabled', + 'bg-components-button-secondary-bg text-components-button-secondary-accent-text shadow-xs inset-ring-[0.5px] shadow-shadow-shadow-3 inset-ring-components-button-secondary-border backdrop-blur-[5px]', + 'hover:bg-components-button-secondary-bg-hover hover:inset-ring-components-button-secondary-border-hover', + 'data-disabled:bg-components-button-secondary-bg-disabled data-disabled:text-components-button-secondary-accent-text-disabled data-disabled:shadow-none data-disabled:inset-ring-components-button-secondary-border-disabled data-disabled:backdrop-blur-xs', ], tertiary: [ 'bg-components-button-tertiary-bg text-components-button-tertiary-text', 'hover:bg-components-button-tertiary-bg-hover', - 'data-[disabled]:bg-components-button-tertiary-bg-disabled data-[disabled]:text-components-button-tertiary-text-disabled', + 'data-disabled:bg-components-button-tertiary-bg-disabled data-disabled:text-components-button-tertiary-text-disabled', ], ghost: [ 'text-components-button-ghost-text', 'hover:bg-components-button-ghost-bg-hover', - 'data-[disabled]:text-components-button-ghost-text-disabled', + 'data-disabled:text-components-button-ghost-text-disabled', ], 'ghost-accent': [ 'text-components-button-secondary-accent-text', 'hover:bg-state-accent-hover', - 'data-[disabled]:text-components-button-secondary-accent-text-disabled', + 'data-disabled:text-components-button-secondary-accent-text-disabled', ], }, size: { - small: 'h-6 rounded-md px-2 text-xs font-medium', - medium: 'h-8 rounded-lg px-3.5 text-[13px] leading-4 font-medium', - large: 'h-9 rounded-[10px] px-4 text-sm font-semibold', + small: 'h-6 gap-1 rounded-md px-[9px] text-xs font-medium', + medium: 'h-8 gap-1 rounded-lg px-3.5 text-[13px] leading-4 font-medium', + large: 'h-9 gap-1.5 rounded-[10px] px-4 text-sm font-semibold', }, tone: { default: '', @@ -55,22 +53,27 @@ export const buttonVariants = cva( }, }, compoundVariants: [ + { + variant: 'primary', + size: 'small', + class: 'gap-[3px] px-2', + }, { variant: 'primary', tone: 'destructive', class: [ - 'border-components-button-destructive-primary-border bg-components-button-destructive-primary-bg text-components-button-destructive-primary-text', - 'hover:border-components-button-destructive-primary-border-hover hover:bg-components-button-destructive-primary-bg-hover', - 'data-[disabled]:border-components-button-destructive-primary-border-disabled data-[disabled]:bg-components-button-destructive-primary-bg-disabled data-[disabled]:text-components-button-destructive-primary-text-disabled data-[disabled]:shadow-none', + 'bg-components-button-destructive-primary-bg text-components-button-destructive-primary-text inset-ring-components-button-destructive-primary-border', + 'hover:bg-components-button-destructive-primary-bg-hover hover:inset-ring-components-button-destructive-primary-border-hover', + 'data-disabled:bg-components-button-destructive-primary-bg-disabled data-disabled:text-components-button-destructive-primary-text-disabled data-disabled:shadow-none data-disabled:inset-ring-components-button-destructive-primary-bg-disabled', ], }, { variant: 'secondary', tone: 'destructive', class: [ - 'border-components-button-destructive-secondary-border bg-components-button-destructive-secondary-bg text-components-button-destructive-secondary-text', - 'hover:border-components-button-destructive-secondary-border-hover hover:bg-components-button-destructive-secondary-bg-hover', - 'data-[disabled]:border-components-button-destructive-secondary-border-disabled data-[disabled]:bg-components-button-destructive-secondary-bg-disabled data-[disabled]:text-components-button-destructive-secondary-text-disabled', + 'bg-components-button-destructive-secondary-bg text-components-button-destructive-secondary-text inset-ring-components-button-destructive-secondary-border', + 'hover:bg-components-button-destructive-secondary-bg-hover hover:inset-ring-components-button-destructive-secondary-border-hover', + 'data-disabled:text-components-button-destructive-secondary-text-disabled', ], }, { @@ -79,7 +82,7 @@ export const buttonVariants = cva( class: [ 'bg-components-button-destructive-tertiary-bg text-components-button-destructive-tertiary-text', 'hover:bg-components-button-destructive-tertiary-bg-hover', - 'data-[disabled]:bg-components-button-destructive-tertiary-bg-disabled data-[disabled]:text-components-button-destructive-tertiary-text-disabled', + 'data-disabled:bg-components-button-destructive-tertiary-bg-disabled data-disabled:text-components-button-destructive-tertiary-text-disabled', ], }, { @@ -88,7 +91,7 @@ export const buttonVariants = cva( class: [ 'text-components-button-destructive-ghost-text', 'hover:bg-components-button-destructive-ghost-bg-hover', - 'data-[disabled]:text-components-button-destructive-ghost-text-disabled', + 'data-disabled:text-components-button-destructive-ghost-text-disabled', ], }, ], @@ -100,13 +103,13 @@ export const buttonVariants = cva( }, ) -export type ButtonProps = Omit & +type ButtonProps = Omit & VariantProps & { loading?: boolean className?: string } -export function Button({ +function Button({ className, variant, size, @@ -129,10 +132,14 @@ export function Button({ {children} {loading && (